#!/usr/bin/env python3
"""唯一强度模型构建器（ADJUDICATION 2026-09-14 §4.6：v3/v4 合并为单一事实源）

历史遗留两个构建器互相打架（`build_strength_model_v3.rebuild()` 会整文件重写 baseline、
`build_strength_v4.py` 也会整文件重写 baseline）→ 0913/0914 两次事故同源。本文件取代表二者：

- `feats()`                链特征（原 v4）：11 链特征，与 scan 的 find_tb_today 同口径
- `generate_backtest_rows()`回测交易 → 特征表（原 v4 主流程；**不写任何文件**）
- `weights_from()`         权重 = 分特征 spearman(feature, ret)（权重主场在研究侧，本地一般不重算）
- `tables_from()`          分池 percentiles + raw_scores 分布
- `build_model()`          **只写 strength_model.json**（绝不写 baseline）
- `merge_into_baseline()` 回测区写入 baseline 时**按 (code,buy_date) 合并**，近端 scan 行（buy_date ≥ 模拟盘起始日）
                           一律保留；写前备份 + 写后行数/键覆盖校验，缺一即拒绝写盘
- 守卫: 模型权重键集合必须 == FCOLS，否则拒绝写模型

用法:
  python build_strength_model.py --model-from $BASE/strength_baseline.csv   # 只重算表(权重沿用现有模型)
  python build_strength_model.py --model-from <baseline> --model-out /tmp/x.json --weights-from <model.json>
  python build_strength_model.py --backtest-rows-out /tmp/rows.csv          # 只生成回测行(不写库)
"""
import argparse, json, os, shutil, 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
import ideal_engine as ie
from plot_ideal_top10 import chain_for_buy


def _pools():
    """延迟取 COMBO/SELL_COMBO（避免 scan_daily_tb ↔ 本模块 的循环导入）"""
    import scan_daily_tb as sdt
    return sdt.COMBO, sdt.SELL_COMBO

BASE = g.BASE
MODEL = f'{BASE}/strength_model.json'
BASELINE = f'{BASE}/strength_baseline.csv'
# 模型世代标记(2026-09-15, 与研究侧 c9d0db6 约定): 取值集合固定 = 'v1-record'(记录版世代) | 'data-v2'(0919 生效世代)。
# 必须由**构建器**写入, 否则每次 rebuild 覆盖 model 时会丢标记 -> 对账报字符串差。落库切换日改这一处常量。
MODEL_VERSION = os.environ.get('ZT_MODEL_VERSION', 'v1-record')
SIM_START = '20260901'          # 近端 scan 追加行起点（模拟盘起始日）: 这些行不许被回测区覆盖
FCOLS = ['f1_t0_pct', 'f2_t0_vol_ratio', 'f3_ta_shrink', 'f4_ta_drawdown', 'f5_ta_gap',
         'f6_tb_strength', 'f7_tb_gap', 'f8_chain_len', 'f9_tb_vol', 'f10_tb_ma5_slope', 'f11_tb_ma_align',
         'f12_crowd_pool', 'f13_pool_zt', 'f14_pool_avgret']
PN = {'hs300': '沪深300', 'zz500': '中证500', 'zz1000': '中证1000', 'zz2000': '中证2000'}


# ---------------- 特征（原 build_strength_v4.feats） ----------------
def feats(d, t0i, tai, tbi, t0d=None):
    """11 链特征, 与 scan_daily_tb.find_tb_today 一致（判定=后复权 hfq_*）"""
    close = d['hfq_close'].astype(float).values
    vol = d['vol'].astype(float).values
    pct = d['pct_chg'].astype(float).values
    ma5 = pd.Series(close).rolling(5).mean().values
    ma10 = pd.Series(close).rolling(10).mean().values
    ma20 = pd.Series(close).rolling(20).mean().values
    v0 = vol[t0i]
    m10_t0 = vol[max(0, t0i - 10):t0i].mean() if t0i >= 10 else vol[:t0i].mean()
    m10_tb = vol[max(0, tbi - 10):tbi].mean() if tbi >= 10 else vol[:tbi].mean()
    return {
        'f1_t0_pct': float(pct[t0i]),
        'f2_t0_vol_ratio': float(v0 / m10_t0) if m10_t0 > 0 else np.nan,
        'f3_ta_shrink': float(vol[tai] / v0) if v0 > 0 else np.nan,
        'f4_ta_drawdown': float(close[tai] / close[t0i] - 1),
        'f5_ta_gap': float(tai - t0i),
        'f6_tb_strength': float(close[tbi] / close[t0i] - 1),
        'f7_tb_gap': float(tbi - t0i),
        'f8_chain_len': float(tbi + 1 - t0i),
        'f9_tb_vol': float(vol[tbi] / m10_tb) if m10_tb > 0 else np.nan,
        'f10_tb_ma5_slope': float((ma5[tbi] - ma5[tbi - 3]) / ma5[tbi - 3]) if tbi >= 3 and ma5[tbi - 3] > 0 else np.nan,
        'f11_tb_ma_align': float(1 if (ma5[tbi] > ma10[tbi] > ma20[tbi]) else (-1 if (ma5[tbi] < ma10[tbi] < ma20[tbi]) else 0)),
    }


# ---------------- 回测行（原 v4 主流程；不写文件） ----------------
def generate_backtest_rows(cache=None, verbose=True):
    from scipy import stats  # noqa: F401  (保持依赖显式)
    COMBO, SELL_COMBO = _pools()
    rows = []
    for pool in COMBO:
        bp = COMBO[pool]; sc = SELL_COMBO[pool]
        members, dailies = g.load_pool(pool)
        sig = g.build_sig(dailies, bp)
        trades = ie.ideal_backtest(pool, members, dailies, sig, sc['profit'], sc['dd'],
                                   mode=sc['mode'], stop=sc['stop'], hold=sc['hold'], start='20240201')
        for t in trades:
            d = dailies[t['code']].sort_values('trade_date').reset_index(drop=True)
            t0i, tai, tbi = chain_for_buy(d, t['code'], bp, str(t['buy_date']))
            if t0i is None:
                continue
            f = feats(d, t0i, tai, tbi)
            rows.append({'code': t['code'], 'pool': pool, 'buy_date': str(t['buy_date']),
                         'sell_date': str(t['sell_date']), 'ret': t['ret_net'], **f})
        if verbose:
            print(f'{PN[pool]}: {len(trades)}笔', flush=True)
    return pd.DataFrame(rows)


# ---------------- 模型（权重 + 分池表） ----------------
def weights_from(df, weight_source='model'):
    """权重主场在研究侧。默认沿用 weight_source 模型(或 dict)的权重，避免本地重算漂移；
    weight_source=None 时才从 df 重算(需含 ret)。"""
    if isinstance(weight_source, dict) and 'weights' in weight_source:
        weight_source = weight_source['weights']
    if isinstance(weight_source, dict):
        return {k: float(v) for k, v in weight_source.items()}
    from scipy import stats
    wdf = df[df['ret'].notna()]
    w = {}
    for fc in FCOLS:
        sub = wdf[wdf[fc].notna()][fc] if fc in wdf.columns else pd.Series(dtype=float)
        if len(sub) < 4:
            w[fc] = float('nan'); continue
        rho, _ = stats.spearmanr(sub, wdf.loc[sub.index, 'ret'])
        w[fc] = round(float(rho), 4)
    return w


def tables_from(df, weights):
    COMBO, _ = _pools()
    pools_m = {}
    for pool in COMBO:
        sub = df[df['pool'] == pool]
        per = {}
        for fc in FCOLS:
            qs2 = sub[sub[fc].notna()][fc] if fc in sub.columns else pd.Series(dtype=float)
            per[fc] = [float(np.nanpercentile(qs2, q)) for q in range(101)] if len(qs2) > 3 else None
        pools_m[pool] = {'n': int(len(sub)), 'percentiles': per, 'raw_scores': None}
    def raw_score(row):
        s = 0.0
        for fc in FCOLS:
            v = row.get(fc, np.nan)
            qs = pools_m[row['pool']]['percentiles'].get(fc)
            if qs is None or v is None or (isinstance(v, float) and np.isnan(v)):
                continue
            s += weights[fc] * (np.searchsorted(qs, v) / 100.0)
        return s
    df = df.copy()
    df['raw_score'] = df.apply(raw_score, axis=1)
    for pool in COMBO:
        pools_m[pool]['raw_scores'] = sorted(df[df['pool'] == pool]['raw_score'].tolist())
    return df, pools_m


def build_model(baseline_path=BASELINE, model_path=MODEL, weights_from_path=None, verbose=True):
    df = pd.read_csv(baseline_path, dtype={'buy_date': str, 'sell_date': str})
    if weights_from_path:
        src = json.load(open(weights_from_path))
    elif os.path.exists(model_path):
        src = json.load(open(model_path))
        if verbose:
            print(f'权重沿用现有模型: {model_path}')
    else:
        src = None
    weights = weights_from(df, src if src is not None else None)
    miss = [fc for fc in FCOLS if fc not in weights or weights[fc] is None or (isinstance(weights[fc], float) and np.isnan(weights[fc]))]
    extra = [fc for fc in weights if fc not in FCOLS]
    if miss or extra:
        raise SystemExit(f'❌ 权重键与 FCOLS 不一致, 拒绝写模型: 缺 {miss}, 多 {extra}')
    scored, pools_m = tables_from(df, weights)
    model = {'weights': weights, 'pools': pools_m, 'n_trades': int(len(df)),
             'model_version': MODEL_VERSION,
             'method': '唯一构建器 build_strength_model.py（权重主场=研究侧；分池 percentiles/raw_scores 本地重算）'}
    json.dump(model, open(model_path, 'w'), ensure_ascii=False, default=str)
    if verbose:
        print(f'模型已写: {model_path} | 权重键 {len(weights)} | 行数 {len(df)}')
        for pool in pools_m:
            print(f'  {PN.get(pool, pool)}: n={pools_m[pool]["n"]} raw_scores={len(pools_m[pool]["raw_scores"])}')
    return model, scored


# ---------------- baseline 合并写（代替整文件覆盖） ----------------
def merge_into_baseline(rows, baseline_path=BASELINE, backup=True, allow_drop=False):
    old = pd.read_csv(baseline_path, dtype={'buy_date': str, 'sell_date': str})
    near = old[old['buy_date'] >= SIM_START]
    far = old[old['buy_date'] < SIM_START]
    merged = pd.concat([rows, near], ignore_index=True)
    merged = merged.sort_values(['pool', 'buy_date']).drop_duplicates(['code', 'buy_date'], keep='first')
    lost = set(zip(near['code'], near['buy_date'])) - set(zip(merged['code'], merged['buy_date']))
    if lost and not allow_drop:
        raise SystemExit(f'❌ 合并会丢失 {len(lost)} 个近端 scan 行, 拒绝写盘（示例 {list(lost)[:3]}）')
    if backup:
        b = f'{baseline_path}.bak_{time.strftime("%Y%m%d_%H%M%S")}'
        shutil.copy(baseline_path, b)
        print(f'备份: {b}')
    merged.to_csv(baseline_path, index=False)
    print(f'baseline 合并写: 回测区 {len(rows)} + 近端 {len(near)} -> {len(merged)} 行'
          f'（原 {len(old)} 行, 回测区旧行 {len(far)} 行被替换）')
    return merged


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('--model-from', default=None, help='用于算表的 baseline(默认取已有 baseline)')
    ap.add_argument('--model-out', default=MODEL)
    ap.add_argument('--weights-from', default=None)
    ap.add_argument('--backtest-rows-out', default=None, help='只生成回测行到该 csv(不写库)')
    ap.add_argument('--write-baseline', action='store_true', help='把回测行合并写回 baseline(需显式开启)')
    a = ap.parse_args()
    if a.backtest_rows_out:
        rows = generate_backtest_rows()
        rows.to_csv(a.backtest_rows_out, index=False)
        print(f'回测行已写: {a.backtest_rows_out} ({len(rows)} 行)')
        return
    df = a.model_from or BASELINE
    build_model(df, a.model_out, a.weights_from)
    if a.write_baseline:
        rows = generate_backtest_rows()
        merge_into_baseline(rows)


if __name__ == '__main__':
    main()
