"""生产口径重测除权守卫：raw_close/raw_prev/pct_chg 全部取**新拉的 pro.daily**
（与 update_daily.py 的 row 同源），不用本地文件里可能是脏列的 pct_chg。
输出：B-only 判据的命中/漏判/误报（真值=因子真变动），以及误报集中在哪些股票。
用法: ./zt_venv/bin/python code/_chk_guard_prod_eval.py
"""
import os, glob, time, random, pickle
import numpy as np
import pandas as pd

os.environ.setdefault('HTTP_PROXY', 'http://127.0.0.1:7897')
os.environ.setdefault('HTTPS_PROXY', 'http://127.0.0.1:7897')
import tushare as ts
ts.set_token('edf6739fe1a4de0d747600cc753a8b4bf335cf27ef0f5aea2d2aa64c')
pro = ts.pro_api()
ROOT = "backtest_zt_full"

files = []
for p in ["hs300", "zz500", "zz1000", "zz2000"]:
    files += glob.glob(f"{ROOT}/daily_{p}/*.csv")
random.seed(7)
sample = random.sample(files, 80)
tfac = pickle.load(open('/tmp/_true_fac_bounded.pkl', 'rb'))

def pull_daily(code, s, e, retries=3):
    for _ in range(retries):
        try:
            q = pro.daily(ts_code=code, start_date=s, end_date=e)
            if q is not None and len(q):
                return q
        except Exception:
            pass
        time.sleep(2)
    return None

tp = fn = fp = tn = 0
fp_by = {}
fn_by = []
loc_vs_ts_pct = []
for f in sample:
    c = os.path.basename(f)[:-4]
    if c not in tfac:
        continue
    lf = pd.read_csv(f)
    s, e = str(int(lf['trade_date'].min())), str(int(lf['trade_date'].max()))
    q = pull_daily(c, s, e)
    if q is None:
        print("no daily", c); continue
    q = q.sort_values('trade_date')
    q['trade_date'] = q['trade_date'].astype(int)
    t = tfac[c].sort_values('trade_date')
    t['prev'] = t['adj_factor'].shift(1)
    t['chg'] = t['adj_factor'].diff().abs() > 1e-9
    ex = {int(r.trade_date): float(r.adj_factor) / float(r.prev) - 1 for r in t[t['chg']].itertuples()}
    q['raw_prev'] = q['close'].shift(1)
    # 本地 pct_chg 与 tushare pct_chg 的偏差（诊断脏列）
    m = lf[['trade_date', 'pct_chg']].merge(q[['trade_date', 'pct_chg']], on='trade_date',
                                            suffixes=('_loc', '_ts'))
    loc_vs_ts_pct.append((m['pct_chg_loc'] - m['pct_chg_ts']).abs().values)
    for i in range(1, len(q)):
        r = q.iloc[i]
        rp = float(r['raw_prev'])
        if not rp > 0:
            continue
        B = (float(r['close']) / rp - 1) * 100
        got = abs(B - float(r['pct_chg'])) >= 0.05
        isex = int(r['trade_date']) in ex
        if isex and got:
            tp += 1
        elif isex and not got:
            fn += 1
            fn_by.append(abs(ex[int(r['trade_date'])]))
        elif not isex and got:
            fp += 1
            fp_by[c] = fp_by.get(c, 0) + 1
        else:
            tn += 1
    time.sleep(0.3)

d = np.concatenate(loc_vs_ts_pct)
print(f"本地 pct_chg vs tushare pct_chg: >1e-4 的行 {(d>1e-4).mean():.2%} | >0.05pp {(d>0.05).mean():.2%} | max {d.max():.3f}")
print(f"\n[生产口径] 真除权 {tp+fn} -> 命中 {tp} ({tp/max(tp+fn,1):.1%}) | 漏 {fn}"
      f" (跳幅中位 {np.median(fn_by) if fn_by else 0:.4%}, max {max(fn_by) if fn_by else 0:.3%})")
print(f"[生产口径] 非除权 {fp+tn} -> 误报 {fp} ({fp/max(fp+tn,1):.3%}) | 集中 {len(fp_by)} 只")
for c, n in sorted(fp_by.items(), key=lambda kv: -kv[1])[:10]:
    print(f"    {c}: {n}")
