#!/usr/bin/env python3
"""回填前后模型口径自检: 权重逐键对照 + 分池 rho / 分档统计（沙盒, 只读）

用法: python chk_pool_model_report.py --new /tmp/baseline_backfilled.csv --old <prod baseline> --model <model.json> --out /tmp/model_compare.md
"""
import argparse, json, os, sys
import numpy as np
import pandas as pd

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from build_strength_model import FCOLS, weights_from, tables_from, PN


def tier_stats(strength, ret):
    s = pd.Series(strength); r = pd.Series(ret)
    idx = s.dropna().index.intersection(r.dropna().index)
    s, r = s[idx], r[idx]
    if len(s) < 10:
        return None
    lo, hi = s.quantile(1 / 3), s.quantile(2 / 3)
    out = {}
    for nm, mask in [('强(前1/3)', s >= hi), ('中', (s > lo) & (s < hi)), ('弱(后1/3)', s <= lo)]:
        rr = r[mask]
        out[nm] = (int(mask.sum()), float(rr.mean()) if len(rr) else np.nan,
                   float((rr > 0).mean() * 100) if len(rr) else np.nan)
    return out, len(s)


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('--new', default='/tmp/baseline_backfilled.csv')
    ap.add_argument('--old', default=None, help='回填前 baseline（对照）')
    ap.add_argument('--model', default='/Users/xpresso/zt_app/backtest_zt_full/strength_model.json')
    ap.add_argument('--out', default='/tmp/model_compare.md')
    a = ap.parse_args()
    model = json.load(open(a.model))
    w_rec = {k: float(v) for k, v in model['weights'].items()}

    L = ['# 池级特征回填 · 模型口径自检（Mac 侧沙盒）\n']
    srcs = [('新(tb 口径回填)', a.new)] + ([('旧(现状)', a.old)] if a.old else [])
    for tag, path in srcs:
        df = pd.read_csv(path, dtype={'buy_date': str, 'sell_date': str})
        L.append(f'\n## {tag} — `{path}`（{len(df)} 行）\n')
        w_re = weights_from(df, None)
        L.append('### ① 权重逐键对照（记录版权重 vs 从本表重算）\n')
        L.append('| 特征 | 记录版权重 | 本表重算 | 差 | 参与样本(非空&有ret) |')
        L.append('|---|---|---|---|---|')
        wdf = df[df['ret'].notna()]
        for fc in FCOLS:
            n = int(wdf[fc].notna().sum()) if fc in wdf.columns else 0
            d = (w_re[fc] - w_rec[fc]) if (fc in w_re and fc in w_rec and not np.isnan(w_re[fc])) else np.nan
            L.append(f'| {fc} | {w_rec.get(fc, float("nan")):.4f} | {w_re.get(fc, float("nan")):.4f} | {d:+.4f} | {n} |')
        scored, pools_m = tables_from(df, w_rec)
        L.append('\n### ② 分池 rho（strength vs ret_net, 全样本）与分档统计\n')
        L.append('| pool | 行数 | 有 ret 行 | rho | 强(前1/3) 均收/胜率 | 中 | 弱(后1/3) |')
        L.append('|---|---|---|---|---|---|---|')
        from scipy import stats
        for pool in ['hs300', 'zz500', 'zz1000', 'zz2000']:
            sub = scored[scored['pool'] == pool]
            ret = sub['ret']
            raw = sub['raw_score']
            both = ret.notna() & raw.notna()
            rho = stats.spearmanr(raw[both], ret[both]).statistic if both.sum() > 5 else np.nan
            raw2 = (raw - raw.min()) / (raw.max() - raw.min()) * 100 if raw.notna().any() and raw.max() > raw.min() else raw
            t = tier_stats(raw2, ret)
            if t is None:
                L.append(f'| {PN[pool]} | {len(sub)} | {int(ret.notna().sum())} | {rho:.4f} | - | - | - |')
                continue
            stats_, n = t
            cells = [f'{stats_[k][1]:.2f}%/{stats_[k][2]:.1f}%' for k in ['强(前1/3)', '中', '弱(后1/3)']]
            L.append(f'| {PN[pool]} | {len(sub)} | {int(ret.notna().sum())} | {rho:.4f} | {cells[0]} | {cells[1]} | {cells[2]} |')
    open(a.out, 'w').write('\n'.join(L) + '\n')
    print('\n'.join(L))
    print(f'\n报告: {a.out}')


if __name__ == '__main__':
    main()
