"""生产口径下 B-only vs A∧B 判据对比（输入全部新拉 pro.daily，真值=因子真变动）。
用法: ./zt_venv/bin/python code/_chk_guard_ab_vs_b.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'))
CACHE = '/tmp/_ts_daily_sample.pkl'
qcache = pickle.load(open(CACHE, 'rb')) if os.path.exists(CACHE) else {}

for f in sample:
    c = os.path.basename(f)[:-4]
    if c in qcache:
        continue
    lf = pd.read_csv(f, usecols=['trade_date'])
    s, e = str(int(lf['trade_date'].min())), str(int(lf['trade_date'].max()))
    for _ in range(3):
        try:
            q = pro.daily(ts_code=c, start_date=s, end_date=e)
            if q is not None and len(q):
                break
        except Exception:
            pass
        time.sleep(2)
    if q is None or not len(q):
        continue
    q = q.sort_values('trade_date')
    q['trade_date'] = q['trade_date'].astype(int)
    qcache[c] = q
    time.sleep(0.3)
pickle.dump(qcache, open(CACHE, 'wb'))

res = {}
for name in ['B-only', 'A∧B']:
    res[name] = {'tp': 0, 'fn': 0, 'fp': 0, 'tn': 0, 'js': [], 'fps': {}}
for f in sample:
    c = os.path.basename(f)[:-4]
    if c not in tfac or c not in qcache:
        continue
    q = qcache[c]
    t = tfac[c].sort_values('trade_date')
    t['chg'] = t['adj_factor'].diff().abs() > 1e-9
    t['prev'] = t['adj_factor'].shift(1)
    ex = {int(r.trade_date): abs(float(r.adj_factor) / float(r.prev) - 1) for r in t[t['chg']].itertuples()}
    q = q.copy()
    q['raw_prev'] = q['close'].shift(1)
    for i in range(1, len(q)):
        r = q.iloc[i]
        rp = float(r['raw_prev']); pc = float(r['pre_close']); pct = float(r['pct_chg'])
        if not rp > 0 or not pc > 0:
            continue
        B = (float(r['close']) / rp - 1) * 100
        A = (float(r['close']) / pc - 1) * 100
        isex = int(r['trade_date']) in ex
        for name, got in [('B-only', abs(B - pct) >= 0.05),
                          ('A∧B', abs(A - pct) < 0.05 and abs(B - pct) >= 0.05)]:
            d = res[name]
            if isex and got:
                d['tp'] += 1
            elif isex and not got:
                d['fn'] += 1; d['js'].append(ex[int(r['trade_date'])])
            elif not isex and got:
                d['fp'] += 1; d['fps'][c] = d['fps'].get(c, 0) + 1
            else:
                d['tn'] += 1

for name, d in res.items():
    tpn, fn, fp, tn = d['tp'], d['fn'], d['fp'], d['tn']
    print(f"[{name}] 命中 {tpn}/{tpn+fn} = {tpn/max(tpn+fn,1):.1%} | 漏跳幅中位 "
          f"{np.median(d['js']) if d['js'] else 0:.4%} max {max(d['js']) if d['js'] else 0:.3%} | "
          f"误报 {fp}/{fp+tn} = {fp/max(fp+tn,1):.3%}")
    if d['fps']:
        print("     误报集中:", sorted(d['fps'].items(), key=lambda kv: -kv[1])[:5])
    print(f"     漏判中跳幅≥0.5%: {sum(1 for j in d['js'] if j>=5e-3)} 条")
