#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""data-v2 落库/联合验收核验装置 (0919 用)

用法(正常验收):
    python code/v2_acceptance_check.py --local backtest_zt_full --recv <研究侧包目录>
用法(自检, 含负向注入):
    python code/v2_acceptance_check.py --selftest

核验五项(判据源自研究侧 CLARIFY_research_20260916_scope-three-items.md §①):
  1. 模型字节同    : 两端 strength_model.json sha256 相同, 且各自 == 本侧 manifest.model_sha256
  2. 键集合双向0差 : baseline 的 (code,pool,buy_date) 集合两端一致 (预期 5786)
  3. strength/raw  : 逐键 |Δ| ≤ 1e-4
  4. 池级特征      : 逐键 f13_pool_zt / f14_pool_avgret |Δ| ≤ 1e-2 (f12 生效后纳入, 预期归零)
  5. 池级缓存      : pool_daily_stats.csv 逐 (pool,trade_date) ≤1e-2; pool_signal_count.csv n_sig 双向0差
  (附带) 候选固化  : scan_today_candidates.json 有则比 (排除 generated_at)
"""
import argparse, hashlib, json, os, sys
import pandas as pd

TOL_STRENGTH = 1e-4
TOL_POOL = 1e-2
KEY = ['code', 'pool', 'buy_date']
MODEL = 'strength_model.json'
MANIFEST = 'model_manifest.json'
BASELINE = 'strength_baseline.csv'
CAND = 'scan_today_candidates.json'
PSTAT = 'pool_daily_stats.csv'
PSIG = 'pool_signal_count.csv'

RESULTS = []


def _sha256(p):
    with open(p, 'rb') as f:
        return hashlib.sha256(f.read()).hexdigest()


def _find(d, name):
    """在目录(含一层子目录)里找文件"""
    p = os.path.join(d, name)
    if os.path.isfile(p):
        return p
    for root, _dirs, files in os.walk(d):
        if name in files:
            return os.path.join(root, name)
    return None


def _norm_key(df):
    df = df.copy()
    df['buy_date'] = df['buy_date'].apply(
        lambda v: str(int(v)) if isinstance(v, float) and not pd.isna(v) else str(v).split('.')[0])
    return df


def _rec(label, ok, detail):
    RESULTS.append((label, ok, detail))
    print(f"{'✅' if ok else '❌'} {label}\n   {detail}")


def check_model(local, recv):
    lm, rm = _find(local, MODEL), _find(recv, MODEL)
    if not (lm and rm):
        _rec('1 模型字节同', False, f'缺文件 local={lm} recv={rm}')
        return
    ls, rs = _sha256(lm), _sha256(rm)
    lmf, rmf = _find(local, MANIFEST), _find(recv, MANIFEST)
    det = [f'local sha256 = {ls}', f'recv  sha256 = {rs}', f'字节同 = {ls == rs}']
    ok = ls == rs
    if lmf and rmf:
        j1, j2 = json.load(open(lmf)), json.load(open(rmf))
        for f in ('model_version', 'effective_from'):
            same = j1.get(f) == j2.get(f)
            det.append(f'{f}: local={j1.get(f)} recv={j2.get(f)} {"一致" if same else "不一致"}')
            ok = ok and same
        for tag, j, s in (('local', j1, ls), ('recv', j2, rs)):
            m = str(j.get('model_sha256', '')).lower() == s
            det.append(f'{tag} manifest.model_sha256 vs 自身 model 文件: {"一致" if m else "不一致"}')
            ok = ok and m
    else:
        det.append(f'⚠️ manifest 缺失(local={bool(lmf)} recv={bool(rmf)}), 跳过交叉校验')
        ok = False
    _rec('1 模型字节同 + manifest 交叉校验', ok, ' | '.join(det))


def check_baseline(local, recv, expect_keys=None):
    lb, rb = _find(local, BASELINE), _find(recv, BASELINE)
    if not (lb and rb):
        _rec('2 键集合双向0差', False, f'缺 baseline local={lb} recv={rb}')
        return None
    a = _norm_key(pd.read_csv(lb, low_memory=False))
    b = _norm_key(pd.read_csv(rb, low_memory=False))
    ka = set(map(tuple, a[KEY].astype(str).values))
    kb = set(map(tuple, b[KEY].astype(str).values))
    onlyA, onlyB = ka - kb, kb - ka
    det = f'local {len(ka)} 键 / recv {len(kb)} 键 / 本地独有 {len(onlyA)} / 对方独有 {len(onlyB)}'
    if expect_keys is not None:
        det += f' / 预期 {expect_keys}'
    ok = not onlyA and not onlyB and (expect_keys is None or len(ka) == expect_keys)
    if onlyA | onlyB:
        det += '\n   样例(各≤5): ' + str(sorted(onlyA)[:5]) + ' | ' + str(sorted(onlyB)[:5])
    _rec('2 baseline 键集合双向 0 差', ok, det)
    return a, b


def check_values(a, b, cols, tol, label):
    if a is None or b is None:
        _rec(label, False, 'baseline 缺失, 无法比较')
        return
    cols = [c for c in cols if c in a.columns and c in b.columns]
    if not cols:
        _rec(label, False, f'两端均无列 {cols}')
        return
    a2 = _norm_key(a)[KEY + cols]
    b2 = _norm_key(b)[KEY + cols]
    m = a2.merge(b2, on=KEY, suffixes=('_L', '_R'), how='inner')
    worst, ok = {}, True
    for c in cols:
        L = pd.to_numeric(m[f'{c}_L'], errors='coerce')
        R = pd.to_numeric(m[f'{c}_R'], errors='coerce')
        d = (L - R).abs()
        one_nan = int((L.isna() ^ R.isna()).sum())      # 一侧 NaN 另一侧有值 = 必须报错
        both_nan = int((L.isna() & R.isna()).sum())
        over = int((d > tol).sum())                      # NaN 不会计入(比较为 False)
        ncov = int(L.notna().sum())
        worst[c] = (float(d.max()) if d.notna().any() else float('nan'), over, one_nan, both_nan, ncov)
        ok = ok and over == 0 and one_nan == 0
    det = f'比 {len(m)} 键; 容差 {tol:g}; ' + '; '.join(
        f'{c}: max|Δ|={w[0]:.3e}, 超差 {w[1]}, 单侧NaN {w[2]}, 双侧NaN {w[3]}, 本侧有效 {w[4]}'
        for c, w in worst.items())
    if any(w[2] for w in worst.values()):
        det += ' ⚠️ 存在「一侧有值/一侧 NaN」⇒ 覆盖率不一致, 必须查明'
    _rec(label, ok, det)


def check_pool_caches(local, recv):
    lp, rp = _find(local, PSTAT), _find(recv, PSTAT)
    if lp and rp:
        a, b = pd.read_csv(lp), pd.read_csv(rp)
        m = a.merge(b, on=['pool', 'trade_date'], suffixes=('_L', '_R'), how='inner')
        det, ok = [], True
        for c in ('n_zt', 'avg_ret'):
            if f'{c}_L' in m and f'{c}_R' in m:
                d = (m[f'{c}_L'].astype(float) - m[f'{c}_R'].astype(float)).abs()
                n = int((d > TOL_POOL).sum())
                det.append(f'{c}: max|Δ|={d.max():.3e}, 超差 {n}')
                ok = ok and n == 0
        det.append(f'比 {len(m)} 行 (local {len(a)} / recv {len(b)})')
        _rec('5a 池级缓存 pool_daily_stats 逐(pool,date)', ok, '; '.join(det))
    else:
        _rec('5a 池级缓存 pool_daily_stats', False, f'缺文件 local={bool(lp)} recv={bool(rp)}')
    ls, rs = _find(local, PSIG), _find(recv, PSIG)
    if ls and rs:
        a, b = pd.read_csv(ls), pd.read_csv(rs)
        ka = set(map(tuple, a[['pool', 'trade_date']].astype(str).values))
        kb = set(map(tuple, b[['pool', 'trade_date']].astype(str).values))
        m = a.merge(b, on=['pool', 'trade_date'], suffixes=('_L', '_R'), how='inner')
        mism = int((m['n_sig_L'] != m['n_sig_R']).sum()) if 'n_sig_L' in m else -1
        ok = not (ka - kb) and not (kb - ka) and mism == 0
        _rec('5b 池级缓存 pool_signal_count n_sig 双向0差', ok,
             f'local {len(ka)} / recv {len(kb)} / 键差 {len(ka-kb)}+{len(kb-ka)} / 值不一致 {mism}')
    else:
        _rec('5b 池级缓存 pool_signal_count', False, f'缺文件 local={bool(ls)} recv={bool(rs)}')


def check_candidates(local, recv):
    lc, rc = _find(local, CAND), _find(recv, CAND)
    if not (lc and rc):
        _rec('6 候选固化(附带项)', False, f'缺文件 local={bool(lc)} recv={bool(rc)}')
        return
    j1, j2 = json.load(open(lc)), json.load(open(rc))
    c1, c2 = j1.get('candidates', []), j2.get('candidates', [])
    k1 = sorted(tuple(str(c.get(x)) for x in ('code', 'pool', 'tb')) for c in c1)
    k2 = sorted(tuple(str(c.get(x)) for x in ('code', 'pool', 'tb')) for c in c2)
    ok = k1 == k2
    det = f'local {len(c1)} 条 / recv {len(c2)} 条 / 集合相同 = {ok}'
    if not ok:
        det += f'\n   本地独有 {sorted(set(k1)-set(k2))[:5]} / 对方独有 {sorted(set(k2)-set(k1))[:5]}'
    _rec('6 候选固化(附带项)', ok, det)


def run(local, recv, expect_keys):
    print(f'\n===== data-v2 验收核验 =====\nlocal = {local}\nrecv  = {recv}\n')
    check_model(local, recv)
    a, b = check_baseline(local, recv, expect_keys)
    check_values(a, b, ['strength', 'raw_score'], TOL_STRENGTH, '3 strength/raw_score 逐键 ≤1e-4')
    check_values(a, b, ['f13_pool_zt', 'f14_pool_avgret'], TOL_POOL,
                 '4 池级特征 f13/f14 逐键 ≤1e-2 (f12 生效后纳入)')
    check_pool_caches(local, recv)
    check_candidates(local, recv)
    bad = [r for r in RESULTS if not r[1]]
    print(f'\n===== 结论: {"全部通过 ✅" if not bad else str(len(bad)) + " 项未通过 ❌"} =====')
    for lbl, _ok, det in bad:
        print(f'  ❌ {lbl}')
    return 0 if not bad else 1


def _synth_pool_caches(src_dir, dst_dir, base):
    """从 baseline 合成 pool_daily_stats.csv / pool_signal_count.csv 供自检使用"""
    df = _norm_key(pd.read_csv(base, low_memory=False))
    df['d'] = df['buy_date'].astype(str)
    agg = df.groupby(['pool', 'd']).agg(n_zt=('code', 'count'), avg_ret=('ret', 'mean'),
                                        n_members=('code', 'nunique')).reset_index()
    agg = agg.rename(columns={'d': 'trade_date'})[['pool', 'trade_date', 'n_zt', 'avg_ret', 'n_members']]
    agg.to_csv(os.path.join(dst_dir, PSTAT), index=False)
    sg = df.groupby(['pool', 'd']).size().reset_index(name='n_sig').rename(columns={'d': 'trade_date'})
    sg[['pool', 'trade_date', 'n_sig']].to_csv(os.path.join(dst_dir, PSIG), index=False)


def selftest(local):
    """自检: 两端同源应全绿, 再逐项注入扰动应各自报错"""
    global RESULTS
    import shutil
    import tempfile
    sbx = tempfile.mkdtemp(prefix='v2acc_')
    loc = os.path.join(sbx, 'loc')
    recv = os.path.join(sbx, 'recv')
    # 两端都做真副本(可写), 保留 daily_*/pkl 以外的文件
    for d in (loc, recv):
        os.makedirs(d)
    for f in (MODEL, MANIFEST, BASELINE, CAND):
        p = _find(local, f)
        if p:
            shutil.copy2(os.path.realpath(p), os.path.join(loc, f))
            shutil.copy2(os.path.realpath(p), os.path.join(recv, f))
    base = os.path.join(loc, BASELINE)
    _synth_pool_caches(loc, loc, base)
    for f in (PSTAT, PSIG):
        shutil.copy2(os.path.join(loc, f), os.path.join(recv, f))

    print('\n########## 自检 A: 两端同源(含合成池级缓存) 应全绿 ##########')
    rc0 = run(loc, recv, None)

    def inject(tag, fn, expect_label, expect_only=None):
        print(f'\n########## 注入「{tag}」-> 期望 {expect_label} 报错 ##########')
        global RESULTS
        keep = RESULTS
        RESULTS = []
        fn()
        bad = [r[0] for r in RESULTS if not r[1]]
        hit = any(expect_label in b for b in bad)
        extra = [b for b in bad if expect_label not in b]
        print(f'>>> {"捕获 ✅" if hit else "漏检 ❌"} | 报错项={bad}'
              + (f' | ⚠️ 附带报错(应为空)={extra}' if extra else ''))
        RESULTS = keep

    def _perturb(fname, mutate):
        p = _find(recv, fname)
        mutate(p, p + '.tmp')
        shutil.move(p + '.tmp', p)
        run(loc, recv, None)

    def _restore(fname):
        shutil.copy2(os.path.join(loc, fname), _find(recv, fname))

    # ① 模型字节不同
    inject('模型字节不同(权重 +1e-9)',
           lambda: _perturb(MODEL, lambda p, t: json.dump(
               {**json.load(open(p)), 'weights': {**json.load(open(p))['weights'], 'f1_t0_pct': json.load(open(p))['weights']['f1_t0_pct'] + 1e-9}},
               open(t, 'w'))),
           '1 模型字节同')
    _restore(MODEL)

    # ② 键集合少一个
    inject('baseline 少一个键',
           lambda: _perturb(BASELINE, lambda p, t: pd.read_csv(p, low_memory=False).iloc[:-1].to_csv(t, index=False)),
           '2 baseline 键集合')
    _restore(BASELINE)

    # ③ strength 越 1e-4 (应只报 3)
    def m3(p, t):
        df = pd.read_csv(p, low_memory=False)
        df.loc[df.index[0], 'strength'] += 1e-2
        df.to_csv(t, index=False)
    inject('strength +1e-2', lambda: _perturb(BASELINE, m3), '3 strength/raw_score')
    _restore(BASELINE)

    # ④ 池级 f14 越 1e-2 (应只报 4; 验证池级与逐键两层判据互不掩盖)
    #    注意: f12/f13/f14 在 v1 baseline 上仅部分行有值, 必须挑有效行注入, 否则注入到 NaN 上等于没注入
    def _first_valid(col):
        s = pd.to_numeric(pd.read_csv(os.path.join(loc, BASELINE), low_memory=False)[col], errors='coerce')
        v = s[s.notna()]
        return int(v.index[0]) if len(v) else 0

    def m4(p, t):
        df = pd.read_csv(p, low_memory=False)
        i = _first_valid('f14_pool_avgret')
        df.loc[i, 'f14_pool_avgret'] += 2e-2
        df.to_csv(t, index=False)
    inject(f'f14_pool_avgret(有效行 idx={_first_valid("f14_pool_avgret")}) +2e-2 (越池级容差)',
           lambda: _perturb(BASELINE, m4), '4 池级特征')
    _restore(BASELINE)

    # ④b 一侧有值 / 一侧 NaN (覆盖率不一致, 必须报错 —— 旧版比较器会静默放过)
    def m4b(p, t):
        df = pd.read_csv(p, low_memory=False)
        df.loc[_first_valid('f14_pool_avgret'), 'f14_pool_avgret'] = float('nan')
        df.to_csv(t, index=False)
    inject('f14_pool_avgret 一侧置 NaN(覆盖率不一致)', lambda: _perturb(BASELINE, m4b), '4 池级特征')
    _restore(BASELINE)

    # ⑤ 池级 f14 只动 5e-3 (容差内, 应全绿: 证明不误报)
    def m5(p, t):
        df = pd.read_csv(p, low_memory=False)
        df.loc[_first_valid('f14_pool_avgret'), 'f14_pool_avgret'] += 5e-3
        df.to_csv(t, index=False)
    print('\n########## 注入「f14_pool_avgret +5e-3(容差内)」-> 期望 全绿无报错 ##########')
    keep = RESULTS
    RESULTS = []
    _perturb(BASELINE, m5)
    bad = [r[0] for r in RESULTS if not r[1]]
    print(f'>>> {"正确未报 ✅" if not bad else "误报 ❌"} | 报错项={bad}')
    RESULTS = keep
    _restore(BASELINE)

    # ⑥ 池级缓存值不一致 (应只报 5a)
    def m6(p, t):
        df = pd.read_csv(p)
        df.loc[df.index[0], 'n_zt'] += 1
        df.to_csv(t, index=False)
    inject('pool_daily_stats n_zt +1', lambda: _perturb(PSTAT, m6), '5a 池级缓存')
    _restore(PSTAT)

    # ⑦ 信号计数不一致 (应只报 5b)
    def m7(p, t):
        df = pd.read_csv(p)
        df.loc[df.index[0], 'n_sig'] += 1
        df.to_csv(t, index=False)
    inject('pool_signal_count n_sig +1', lambda: _perturb(PSIG, m7), '5b 池级缓存')
    _restore(PSIG)

    print(f'\n自检完成。同源基线结论 rc={rc0} (0=全绿)')
    print(f'沙盒目录: {sbx}')
    return 0 if rc0 == 0 else 1


if __name__ == '__main__':
    ap = argparse.ArgumentParser()
    ap.add_argument('--local', default='backtest_zt_full')
    ap.add_argument('--recv', default=None)
    ap.add_argument('--expect-keys', type=int, default=None)
    ap.add_argument('--selftest', action='store_true')
    a = ap.parse_args()
    if a.selftest:
        sys.exit(selftest(a.local))
    if not a.recv:
        print('需要 --recv <研究侧包目录>'); sys.exit(2)
    sys.exit(run(a.local, a.recv, a.expect_keys))
