#!/usr/bin/env python3
"""fix_adj_factor.py — 用 tushare adj_factor 修准库内 `adj_factor` 列(按日全市场拉取)。

背景(2026-09-13 三方复核): `adj_factor` 是唯一进判定链(`hfq = raw × factor`)的脏列 ——
同一 `(code,date)` 在不同池副本里因子不同(例 002812.SZ hs300 3.492085 vs zz500 3.495246,
1313/1366 行差>1e-4), 逐池比 tushare: 20210802 hs300 80.8% / zz500 33.1% / zz1000 35.5% /
zz2000 41.9%, 2024 之后只剩 hs300 副本脏(16.6%~57.9%)。修法:
  **按日 `pro.adj_factor(trade_date=d)` 全市场拉取(≈1381 次覆盖 5.7 年, 不动 raw、不改列名口径);
  不要逐只无边界拉取 —— `pro.*` 无边界返回最近 6000 行, 会截掉最老那段(recon 实测)。**

两步(可分开跑):
  --build   建/补 `backtest_zt_full/adj_factor_all.csv` 累积因子表(与现有表 concat + 去重留最后, 不覆盖式重写)
  --apply   按表修各池 `adj_factor` 列(只动这一列; 变更前整池备份到 _adjfactor_repair_backup/, 明细落 manifest)
  --dry     只统计不写盘
  --days 1381 限制天数(试跑)   --pools hs300,zz500

验收: `code/chk_pool_factor_vs_tushare.py`(逐池 vs tushare) + `code/chk_cross_pool_mac.py`(跨副本)。
"""
import os, sys, glob, time, shutil, datetime
import numpy as np
import pandas as pd
import tushare as ts

BASE = os.environ.get('ZT_BASE', '/Users/xpresso/zt_app/backtest_zt_full')
FAC_TABLE = f'{BASE}/adj_factor_all.csv'
BAK = os.path.join(os.path.dirname(BASE), '_adjfactor_repair_backup')
CODE_DIR = os.path.dirname(os.path.abspath(__file__))
POOLS = ['hs300', 'zz500', 'zz1000', 'zz2000']
DRY = '--dry' in sys.argv
BUILD = '--build' in sys.argv
APPLY = '--apply' in sys.argv
LIMITDAYS = int(sys.argv[sys.argv.index('--days') + 1]) if '--days' in sys.argv else None
if '--pools' in sys.argv:
    POOLS = sys.argv[sys.argv.index('--pools') + 1].split(',')
STAMP = datetime.datetime.now().strftime('%Y%m%d_%H%M')
MANIFEST = f'{CODE_DIR}/_adjfactor_manifest_{STAMP}.csv'
TOL = 1e-6
if '--tol' in sys.argv:      # 0 → 逐位对齐因子表(把 ≤1e-6 的副本间残留也抹平, 跨副本对账才能严格归零)
    TOL = float(sys.argv[sys.argv.index('--tol') + 1])
pro = ts.pro_api()

files = [(p, f) for p in POOLS for f in sorted(glob.glob(f'{BASE}/daily_{p}/*.csv'))]
print(f'池 {POOLS} | 文件 {len(files)} | BUILD={BUILD} APPLY={APPLY} DRY={DRY}')

if BUILD:
    days = set()
    for _, f in files:
        days |= set(pd.read_csv(f, dtype={'trade_date': str}, usecols=['trade_date'])['trade_date'])
    days = sorted(days)
    if LIMITDAYS:
        days = days[-LIMITDAYS:]
    print(f'交易日 {len(days)} ({days[0]}~{days[-1]}) → 按日拉取 adj_factor')
    rows = []
    t0 = time.time()
    for i, d in enumerate(days, 1):
        for attempt in range(4):
            try:
                r = pro.adj_factor(trade_date=d)
                break
            except Exception as ex:
                if attempt == 3:
                    print(f'  ✗ {d}: {ex}'); r = None
                time.sleep(1.5 * (attempt + 1))
        if r is None or not len(r):
            print(f'  ⚠ {d}: 返回 0 行(跳过)'); continue
        r['trade_date'] = r['trade_date'].astype(str)      # tushare 返回 str, 统一保 str
        rows.append(r[['ts_code', 'trade_date', 'adj_factor']])
        if i % 100 == 0:
            print(f'  {i}/{len(days)} | {time.time()-t0:.0f}s | 累计 {sum(len(x) for x in rows)} 行')
        time.sleep(0.12)
    new = pd.concat(rows, ignore_index=True) if rows else pd.DataFrame()
    if len(new):
        if os.path.exists(FAC_TABLE):
            old = pd.read_csv(FAC_TABLE, dtype={'trade_date': str})
            new = pd.concat([old, new], ignore_index=True)
        new = new.drop_duplicates(subset=['ts_code', 'trade_date'], keep='last')
        print(f'因子表: {FAC_TABLE} → {len(new)} 行 / {new["trade_date"].nunique()} 交易日 '
              f'({new["trade_date"].min()}~{new["trade_date"].max()})')
        if DRY:
            print('[DRY] 未写因子表')
        else:
            new.to_csv(FAC_TABLE, index=False)
            print('因子表已写盘')
    else:
        print('未取到任何因子, 退出'); sys.exit(1)

if APPLY:
    tbl = pd.read_csv(FAC_TABLE, dtype={'trade_date': str})
    tbl = tbl.drop_duplicates(subset=['ts_code', 'trade_date'], keep='last')
    key = {(r.ts_code, r.trade_date): float(r.adj_factor) for r in tbl.itertuples()}
    print(f'因子表载入 {len(key)} 键')
    manifest = []
    stat = dict(files=0, changed_files=0, rows=0, patch=0, missing=0, err=0)
    skip = []
    for n, (pool, fpath) in enumerate(files, 1):
        code = os.path.basename(fpath)[:-4]
        try:
            d = pd.read_csv(fpath, dtype={'trade_date': str})
            ref = np.array([key.get((code, t), np.nan) for t in d['trade_date']])
            cur = d['adj_factor'].values.astype(float)
            miss = int(np.isnan(ref).sum())
            mask = np.logical_and(np.isfinite(ref), np.abs(cur - ref) > TOL)
            k = int(mask.sum())
            stat['files'] += 1; stat['rows'] += len(d); stat['missing'] += miss
            if miss:
                skip.append((code, f'因子表缺 {miss} 日'))
            if not k:
                continue
            stat['patch'] += k; stat['changed_files'] += 1
            for i in np.where(mask)[0]:
                manifest.append((code, d['trade_date'].values[i], float(cur[i]), float(ref[i])))
            if DRY:
                continue
            bdir = f'{BAK}/{pool}'
            os.makedirs(bdir, exist_ok=True)
            bdst = f'{bdir}/{code}.csv'
            if not os.path.exists(bdst):
                shutil.copy2(fpath, bdst)
            d.loc[mask, 'adj_factor'] = ref[mask]
            d.to_csv(fpath, index=False)
        except Exception as ex:
            stat['err'] += 1
            skip.append((code, f'异常 {type(ex).__name__}: {ex}'))
        if n % 500 == 0:
            print(f'  {n}/{len(files)} | 改文件 {stat["changed_files"]} | 改行 {stat["patch"]}')
    print(f'\n=== 汇总 === 处理 {stat["files"]} 文件 / {stat["rows"]} 行 | 需改文件 {stat["changed_files"]} '
          f'| 需改行 {stat["patch"]} | 因子表缺失行 {stat["missing"]} | 异常 {stat["err"]}')
    if skip:
        print(f'提示(前 10): {skip[:10]}')
    if manifest and not DRY:
        pd.DataFrame(manifest, columns=['code', 'trade_date', 'old', 'new']).to_csv(MANIFEST, index=False)
        print(f'变更明细: {MANIFEST} ({len(manifest)} 行)')
    print('DRY, 未写盘' if DRY else f'完成, 备份 {BAK}')