"""adj_factor 对账第二层：全历史 level 曲线 + hfq 收益污染量化（只读，仅打 tushare）。

用法:
  ./zt_venv/bin/python code/_chk_factor_curve.py curve            # 每月抽一个交易日, 全市场 level 比对
  ./zt_venv/bin/python code/_chk_factor_curve.py impact 20220401 20250401
                                                                  # 指定锚点起 20 交易日的 hfq 收益差

坑: pro.adj_factor 连打会返回**空表**(限频, 不是无数据) —— 空表会被误读成"零偏差",
    必须重试+间隔, 并区分「空返回」与「0 只超阈」。
"""
import os, sys, glob, time
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 = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "backtest_zt_full")
POOLS = ["hs300", "zz500", "zz1000", "zz2000"]


def load_factors():
    loc = {}
    for p in POOLS:
        for f in glob.glob(f"{ROOT}/daily_{p}/*.csv"):
            code = os.path.basename(f)[:-4]
            d = pd.read_csv(f, usecols=['trade_date', 'adj_factor'])
            loc[code] = dict(zip(d['trade_date'].astype(int), d['adj_factor']))
    return loc


def true_factor(dt, retries=4, delay=3):
    for _ in range(retries):
        try:
            t = pro.adj_factor(trade_date=dt)
            if t is not None and len(t) > 0:
                return {k: float(v) for k, v in zip(t['ts_code'], t['adj_factor'])}
        except Exception:
            pass
        time.sleep(delay)
    return None


def curve():
    ref = pd.read_csv(f"{ROOT}/daily_hs300/000001.SZ.csv", usecols=['trade_date'])
    byym = {}
    for d in ref['trade_date'].astype(int).astype(str):
        byym.setdefault(d[:6], d)
    loc = load_factors()
    for dt in sorted(byym.values()):
        tf = true_factor(dt)
        if tf is None:
            print(f"{dt} EMPTY(限频/无数据)"); continue
        e = np.array([abs(float(v) / tf[c] - 1) for c, d in loc.items()
                      if (v := d.get(int(dt))) is not None and c in tf and tf[c] > 0])
        print(f"{dt} n={len(e)} >1e-4={np.mean(e>1e-4):.2%} >1e-3={np.mean(e>1e-3):.2%} max={e.max():.2e}")


def impact(anchor, n_len=20):
    n_len = int(n_len)
    ref = pd.read_csv(f"{ROOT}/daily_hs300/000001.SZ.csv", usecols=['trade_date'])
    days = ref['trade_date'].astype(int).astype(str).tolist()
    i = days.index(anchor)
    days = days[max(0, i - 1):i + n_len]
    tfs = {}
    for dt in days:
        tf = true_factor(dt)
        if tf: tfs[int(dt)] = tf
        else: print(f"{dt} EMPTY")
    ti = sorted(tfs)
    stats = []
    for p in POOLS:
        for f in glob.glob(f"{ROOT}/daily_{p}/*.csv"):
            c = os.path.basename(f)[:-4]
            d = pd.read_csv(f, usecols=['trade_date', 'raw_close', 'adj_factor'])
            d = d[d.trade_date.isin(ti)].sort_values('trade_date')
            if len(d) < 3:
                continue
            try:
                tf = np.array([tfs[int(x)][c] for x in d['trade_date'].values])
            except KeyError:
                continue
            rc = d['raw_close'].values
            hl = rc * d['adj_factor'].values
            ht = rc * tf
            stats.append(np.abs((hl[1:] / hl[:-1] - 1) - (ht[1:] / ht[:-1] - 1)))
    a = np.concatenate(stats)
    print(f"[{anchor}..{ti[-1]}] n={len(a)} Δ>1e-4={np.mean(a>1e-4):.2%} "
          f"Δ>1e-3={np.mean(a>1e-3):.3%} Δ>5e-3={np.mean(a>5e-3):.4%} max={a.max():.3e}")


if __name__ == '__main__':
    mode = sys.argv[1] if len(sys.argv) > 1 else 'curve'
    if mode == 'curve':
        curve()
    else:
        impact(*sys.argv[2:])