#!/usr/bin/env python3
"""池级缓存 × 2（RULING c33b091 §3/§4；Mac 侧产物）——沙盒生成器

口径（**经对端实测复核后更正**，勿再改）:
  成员集 = **该池日线目录当日全部有行情的文件**，不是 `{pool}_members.pkl` 的成分名单。
  证据（对端独立复核）: zz1000 20260903 f14 用全目录 1938 只 → -0.112262（与 baseline 存值 Δ=3.8e-07）；
  用当时名单 998 只 → +0.0216（Δ=1.3e-01）。本端 `chk_pool_stat_variants.py` 的 all_buy 变体亦逐日重合。

输出（--out 目录）:
  ① pool_signal_count.csv  pool,trade_date,n_sig,c_mean250
     n_sig       = C(d): 该池 d 日"链触发条目数"（同 code 多链计多次；= scan 当日 all_sig 同语义）
     c_mean250   = 严格早于 d 的过去 250 个交易日的 C 均值（不足 250 日则用可得交易日）
     ⇒ f12(tb) = C(tb) / max(1, mean(C(tb-250..tb-1)))（分子分母同源、不读 baseline、幂等）
  ② pool_daily_stats.csv   pool,trade_date,n_zt,avg_ret,n_codes
     n_zt    = 当日 pct_chg >= 9.8 的家数（f13 分子 = n_zt/20）
     avg_ret = 当日有行情文件 pct_chg 均值（百分数单位，f14）
     n_codes = 当日有行情的文件数（缺失判据用: 当日 0 只 → 该行池级特征 NaN）

用法: python build_pool_caches.py --out /tmp/pool_cache_v1     # 沙盒
      python build_pool_caches.py --out $BASE --append           # 落库（非交易日, 待 Owner 签批后）
"""
import argparse, os, sys, time
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
from plot_ideal_top10 import chain_for_buy

ZT_PCT = 9.8      # 裁定 §3: f13 分子阈值
CMEAN_WIN = 250   # 裁定 §4(c): f12 分母窗口（交易日）


def pool_daily_stats(pool, dailies):
    """② 池×日统计（该池目录当日全部有行情文件）"""
    parts = [d[['trade_date', 'pct_chg']].assign(code=c) for c, d in dailies.items()]
    df = pd.concat(parts, ignore_index=True)
    df['pct_chg'] = pd.to_numeric(df['pct_chg'], errors='coerce')
    df = df.dropna(subset=['pct_chg'])
    gp = df.groupby('trade_date')['pct_chg']
    out = pd.DataFrame({'n_zt': gp.apply(lambda s: int((s >= ZT_PCT).sum())),
                        'avg_ret': gp.mean(), 'n_codes': gp.size()}).reset_index()
    out.insert(0, 'pool', pool)
    return out[['pool', 'trade_date', 'n_zt', 'avg_ret', 'n_codes']]


def pool_signal_count(pool, dailies, bp, cal):
    """① C(d) = 链触发条目数（按 tb 归属；同 code 多链计多次）+ 过去 250 交易日滚动均值"""
    sig = g.build_sig(dailies, bp)
    cnt = {}
    for code, bs in sig.items():
        d = dailies[code].sort_values('trade_date').reset_index(drop=True)
        for (bd, t0d) in bs:
            t0i, tai, tbi = chain_for_buy(d, code, bp, str(bd))
            if t0i is None:
                continue
            tb = str(d['trade_date'].iloc[tbi])
            cnt[tb] = cnt.get(tb, 0) + 1
    cal = sorted(cal)
    rows = []
    for i, d0 in enumerate(cal):
        hist = cal[max(0, i - CMEAN_WIN):i]
        c = cnt.get(d0, 0)
        cm = float(np.mean([cnt.get(x, 0) for x in hist])) if hist else np.nan
        rows.append((pool, d0, c, cm))
    return pd.DataFrame(rows, columns=['pool', 'trade_date', 'n_sig', 'c_mean250'])


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('--out', default='/tmp/pool_cache_v1')
    ap.add_argument('--pools', default='hs300,zz500,zz1000,zz2000')
    ap.add_argument('--append', action='store_true')
    a = ap.parse_args()
    os.makedirs(a.out, exist_ok=True)
    pools = [p for p in a.pools.split(',') if p]
    t0 = time.time()
    sig_parts, stat_parts = [], []
    for pool in pools:
        _, dailies = g.load_pool(pool)
        st = pool_daily_stats(pool, dailies)
        sc = pool_signal_count(pool, dailies, COMBO[pool], st['trade_date'])
        stat_parts.append(st); sig_parts.append(sc)
        print(f'{pool}: 目录文件 {len(dailies)} | 统计 {len(st)} 日({st["trade_date"].min()}~{st["trade_date"].max()}, '
              f'当日文件数 {int(st["n_codes"].min())}~{int(st["n_codes"].max())}) | C 合计 {int(sc["n_sig"].sum())} 条, '
              f'非零日 {int((sc["n_sig"] > 0).sum())}', flush=True)
    stats = pd.concat(stat_parts, ignore_index=True)
    sigs = pd.concat(sig_parts, ignore_index=True)

    def write(df, name, cols):
        fp = f'{a.out}/{name}'
        if a.append and os.path.exists(fp):
            old = pd.read_csv(fp, dtype={'trade_date': str})
            have = set(zip(old['pool'], old['trade_date']))
            add = df[~df.apply(lambda r: (r['pool'], r['trade_date']) in have, axis=1)]
            out = pd.concat([old, add], ignore_index=True) if len(add) else old
            out.sort_values(['pool', 'trade_date'])[cols].to_csv(fp, index=False)
            print(f'{name}: 追加 {len(add)} 行 -> 共 {len(out)} 行')
        else:
            df.sort_values(['pool', 'trade_date'])[cols].to_csv(fp, index=False)
            print(f'{name}: {len(df)} 行 -> {fp}')

    write(sigs, 'pool_signal_count.csv', ['pool', 'trade_date', 'n_sig', 'c_mean250'])
    write(stats, 'pool_daily_stats.csv', ['pool', 'trade_date', 'n_zt', 'avg_ret', 'n_codes'])
    print(f'完成, 用时 {time.time() - t0:.0f}s')


if __name__ == '__main__':
    main()