#!/usr/bin/env python3
"""update_daily.py — 每日增量拉四池行情, 写为【前复权】口径。
关键: tushare pro.daily 返回未复权 close 但 pct_chg 已复权正确(除权日平滑)。
为保证与历史前复权重建序列连续, 当日 close = 昨日文件尾前复权close × (1+当日pct_chg/100);
open/high/low 用 (当日open/high/low / raw当日close) × 前复权close 同比例缩放。
"""
import tushare as ts
import pandas as pd
import os, sys, datetime

BASE = '/Users/xpresso/zt_app/backtest_zt_full'
NEW_DATE = '20260831'  # 兜底
ts.set_token('edf6739fe1a4de0d747600cc753a8b4bf335cf27ef0f5aea2d2aa64c')
pro = ts.pro_api()

try:
    cal = pro.trade_cal(exchange='SSE', start_date='20260101',
                        end_date=datetime.date.today().strftime('%Y%m%d'), is_open='1')
    if cal is not None and len(cal) > 0:
        NEW_DATE = str(cal['cal_date'].max())
except Exception:
    pass
print(f'目标交易日: {NEW_DATE}')

df = pro.daily(trade_date=NEW_DATE, fields='ts_code,trade_date,open,high,low,close,pre_close,change,pct_chg,vol,amount')
if df is None or len(df) == 0:
    print(f'{NEW_DATE} 行情尚未发布, 无法更新')
    sys.exit(1)
print(f'全市场 {NEW_DATE}: {len(df)}只')

df = df.set_index('ts_code')

pools = ['hs300', 'zz500', 'zz1000', 'zz2000']
all_codes = set()
for pool in pools:
    members = pd.read_pickle(f'{BASE}/{pool}_members.pkl')
    for mset in members.values():
        all_codes |= set(mset)

# === 最新复权因子(全市场一次拉取, 用于后复权 adj_close) ===
latest_fac = {}
try:
    _af = pro.adj_factor(trade_date=NEW_DATE)
    latest_fac = {r['ts_code']: float(r['adj_factor']) for _, r in _af.iterrows()}
    print(f'最新复权因子: {len(latest_fac)}只')
except Exception as e:
    print(f'adj_factor 拉取失败(adj_close将空缺): {e}')

for pool in pools:
    ddir = f'{BASE}/daily_{pool}'
    members = pd.read_pickle(f'{BASE}/{pool}_members.pkl')
    codes = set()
    for mset in members.values():
        codes |= set(mset)
    updated, skipped = 0, 0
    for code in sorted(codes):
        fpath = f'{ddir}/{code}.csv'
        if not os.path.exists(fpath):
            skipped += 1
            continue
        if code not in df.index:
            continue  # 停牌/无行情
        row = df.loc[code]
        old = pd.read_csv(fpath, dtype={'trade_date': str})
        if len(old) > 0 and old['trade_date'].iloc[-1] >= NEW_DATE:
            skipped += 1
            continue
        # === 前复权增量: 用昨日文件尾前复权close ×(1+pct/100) 推当日 ===
        last_qfq_close = float(old['close'].iloc[-1]) if len(old) else None
        pct = float(row['pct_chg'])
        raw_close = float(row['close'])
        raw_open, raw_high, raw_low = float(row['open']), float(row['high']), float(row['low'])
        if last_qfq_close is None or pd.isna(pct):
            # 无昨日或缺失pct: 退化为直接用未复权close(极少, 通常次日补齐)
            qfq_close = raw_close
        else:
            qfq_close = last_qfq_close * (1 + pct / 100)
        # open/high/low 同交易日按与close的比例缩放(复权窗口内比例恒定)
        rx = qfq_close / raw_close if raw_close else 1.0
        new_row = {
            'ts_code': row['ts_code'], 'trade_date': NEW_DATE,
            'open': round(raw_open * rx, 4), 'high': round(raw_high * rx, 4),
            'low': round(raw_low * rx, 4), 'close': round(qfq_close, 4),
            'pre_close': row['pre_close'], 'change': row['change'],
            'pct_chg': pct, 'vol': row['vol'], 'amount': row['amount'],
            'adj_close': round(qfq_close * latest_fac.get(row['ts_code'], 1.0), 4),
            'raw_close': raw_close, 'raw_open': raw_open, 'raw_high': raw_high, 'raw_low': raw_low,
        }
        new_df = pd.DataFrame([new_row])
        merged = pd.concat([old, new_df], ignore_index=True)
        merged.to_csv(fpath, index=False)
        updated += 1
    print(f'{pool}: 更新{updated}只 (跳过{skipped})')
print('完成(前复权口径)')