#!/usr/bin/env python3
"""池级统计口径差异分解（Mac 侧沙盒）: 把「存疑的大偏差」拆成 ①日错位 ②代码集 ③修准 三部分

对 170 行涉及的 (pool, buy_date, tb) 组合，逐一计算:
  A = 全目录代码 / 9.8 / 买入日   ← 复刻旧 scan 实现(target=买入日)
  B = 全目录代码 / 9.8 / tb
  C = 时点成分股 / 9.8 / tb       ← 新口径(裁定 §1/§2)
  D = 时点成分股 / 9.8 / 买入日
并与 baseline 里存的值对照（旧值=当时实时值，含 0913 修正前数据）。
用法: python chk_pool_stat_variants.py --cache-dir /tmp/pool_cache_sbx --out /tmp/pool_stat_variants.csv
"""
import argparse, os, sys
import numpy as np
import pandas as pd

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import grid_or as g
from scan_daily_tb import COMBO


def snapshots(pool):
    m = pd.read_pickle(f'{g.BASE}/{pool}_members.pkl')
    return [(k, set(m[k])) for k in sorted(k for k in m if len(m[k]) > 0)]


def members_asof(snaps, date):
    """≤ date 的最近月末快照（仅用于 mem_* 对照变体，非本口径）"""
    pick = snaps[0][1]
    for k, s in snaps:
        if k <= date:
            pick = s
        else:
            break
    return pick

POOLS = ['hs300', 'zz500', 'zz1000', 'zz2000']
ZT = 9.8


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('--baseline', default=f'{g.BASE}/strength_baseline.csv')
    ap.add_argument('--diffs', default='/tmp/pool_backfill_diffs.csv')
    ap.add_argument('--out', default='/tmp/pool_stat_variants.csv')
    ap.add_argument('--report', default='/tmp/pool_stat_variants.md')
    a = ap.parse_args()
    d = pd.read_csv(a.diffs, dtype={'buy_date': str, 'tb': str})
    base = pd.read_csv(a.baseline, dtype={'buy_date': str, 'sell_date': str})
    rows = []
    for pool in POOLS:
        sub = d[d['pool'] == pool][['buy_date', 'tb']].drop_duplicates()
        if sub.empty:
            continue
        sn = snapshots(pool)
        _, dailies = g.load_pool(pool)
        # 全目录 + 成分股 两套日度序列
        parts = [dl[['trade_date', 'pct_chg']].assign(code=c) for c, dl in dailies.items()]
        alld = pd.concat(parts, ignore_index=True)
        alld['pct_chg'] = pd.to_numeric(alld['pct_chg'], errors='coerce')
        alld = alld.dropna(subset=['pct_chg'])
        for r in sub.itertuples():
            for tag, date in [('buy', r.buy_date), ('tb', r.tb)]:
                if not isinstance(date, str) or date == 'nan':
                    continue
                mset = members_asof(sn, date)
                for setname, s in [('all', alld[alld['trade_date'] == date]),
                                   ('mem', alld[(alld['trade_date'] == date) & (alld['code'].isin(mset))])]:
                    rows.append({'pool': pool, 'buy_date': r.buy_date, 'tb': r.tb, 'variant': f'{setname}_{tag}',
                                 'n_codes': int(len(s)), 'n_zt': int((s['pct_chg'] >= ZT).sum()),
                                 'avg_ret': float(s['pct_chg'].mean()) if len(s) else np.nan,
                                 'f13': (s['pct_chg'] >= ZT).sum() / 20.0 if len(s) else np.nan})
        print(pool, 'done', flush=True)
    v = pd.DataFrame(rows)
    v.to_csv(a.out, index=False)
    # 与 baseline 存值对照（每个 pool×buy_date 取存值）
    oldmap = d.groupby(['pool', 'buy_date'])[['f13_pool_zt_old', 'f14_pool_avgret_old']].first()
    L = ['# 池级统计口径差异分解（Mac 侧沙盒）\n']
    L.append('变体: `all_buy`=全目录/买入日(=旧 scan 实现复刻), `all_tb`, `mem_buy`, `mem_tb`(=新口径)\n')
    L.append('## 与 baseline 存值(旧实时值)对照: f13 = n_zt/20\n')
    L.append('| pool | buy_date | 存值 f13 | all_buy | all_tb | mem_buy | mem_tb | 存值 f14 | all_buy f14 | mem_tb f14 |')
    L.append('|---|---|---|---|---|---|---|---|---|---|')
    for pool in POOLS:
        sub = d[d['pool'] == pool][['buy_date']].drop_duplicates().sort_values('buy_date')
        for r in sub.itertuples():
            pv = v[(v['pool'] == pool) & (v['buy_date'] == r.buy_date)]
            get = lambda var, col: (pv[pv['variant'] == var][col].iloc[0] if len(pv[pv['variant'] == var]) else np.nan)
            ov = oldmap.loc[(pool, r.buy_date)] if (pool, r.buy_date) in oldmap.index else None
            o13 = float(ov.iloc[0]) if ov is not None else np.nan
            o14 = float(ov.iloc[1]) if ov is not None else np.nan
            L.append(f'| {pool} | {r.buy_date} | {o13:.2f} | {get("all_buy","f13"):.2f} | {get("all_tb","f13"):.2f} | '
                     f'{get("mem_buy","f13"):.2f} | {get("mem_tb","f13"):.2f} | {o14:.3f} | {get("all_buy","avg_ret"):.3f} | {get("mem_tb","avg_ret"):.3f} |')
    open(a.report, 'w').write('\n'.join(L) + '\n')
    print('\n'.join(L))
    print(f'\n明细: {a.out}\n报告: {a.report}')


if __name__ == '__main__':
    main()