"""Mac 侧跨副本一致性 + 具名股覆盖（只读，无 API）。

输出:
  1) 具名股在各池副本的行数 / 首末 trade_date（应 @default-researchagent 要求）
  2) 多副本股 (code,date) 并集 / pairwise 行对
  3) 不一致行对数: adj_factor / pct_chg / raw_close（可扩展到 10 列）
用法: ./zt_venv/bin/python code/chk_cross_pool_mac.py
"""
import glob, os, itertools
from collections import defaultdict
import pandas as pd

ROOT = "backtest_zt_full"
POOLS = ["hs300", "zz500", "zz1000", "zz2000"]
COLS = ['raw_open', 'raw_high', 'raw_low', 'raw_close', 'adj_factor',
        'pre_close', 'change', 'pct_chg', 'vol', 'amount']

print("=== 1) 具名股覆盖 ===")
for code in ['002812.SZ', '600705.SH']:
    for p in POOLS:
        f = f"{ROOT}/daily_{p}/{code}.csv"
        if not os.path.exists(f):
            continue
        d = pd.read_csv(f, usecols=['trade_date'])
        print(f"  {code} @{p}: {len(d)} 行, {int(d['trade_date'].min())}~{int(d['trade_date'].max())}")

print("\n=== 2/3) 跨副本一致性 ===")
by_code = defaultdict(dict)
for p in POOLS:
    for f in glob.glob(f"{ROOT}/daily_{p}/*.csv"):
        by_code[os.path.basename(f)[:-4]][p] = f

multi = {c: v for c, v in by_code.items() if len(v) >= 2}
print(f"多副本股: {len(multi)} 只 | 副本数分布: "
      f"{ {k: sum(1 for v in multi.values() if len(v)==k) for k in sorted({len(v) for v in multi.values()})} }")

frames = {}
union_rows = pairwise_rows = star_rows = 0
mismatch = {c: 0 for c in COLS}
codes_bad = {c: set() for c in COLS}
pair_dist = defaultdict(int)
for code, pools in multi.items():
    dfs = {}
    for p, f in pools.items():
        d = pd.read_csv(f)
        d['trade_date'] = d['trade_date'].astype(int)
        dfs[p] = d.set_index('trade_date')
    keys = set.intersection(*[set(d.index) for d in dfs.values()])
    uni = set.union(*[set(d.index) for d in dfs.values()])
    union_rows += len(uni)
    pairs = list(itertools.combinations(sorted(dfs), 2))
    pairwise_rows += sum(len(keys) for _ in pairs)
    base = pairs[0][0]
    star_rows += sum(len(keys) for _ in pairs)
    for p1, p2 in pairs:
        a, b = dfs[p1].loc[sorted(keys)], dfs[p2].loc[sorted(keys)]
        for col in COLS:
            if col not in a.columns or col not in b.columns:
                continue
            diff = ~((a[col] - b[col]).abs() <= 1e-9)
            n = int(diff.sum())
            if n:
                mismatch[col] += n
                codes_bad[col].add(code)
                pair_dist[f"{p1}-{p2}"] += n

print(f"(code,date) 并集: {union_rows} | pairwise 行对: {pairwise_rows} | star 行对: {star_rows}")
for col in COLS:
    print(f"  {col:<12} 不一致行对 {mismatch[col]:>8} | 涉及股数 {len(codes_bad[col]):>5}")
print("按池配对分布(adj_factor):", dict(sorted(pair_dist.items(), key=lambda kv: -kv[1])[:6]))
