#!/usr/bin/env python3
"""模拟盘监控图: 净值 vs 四指数(买入持有, 统一起点) + 持仓仓位
- 输入: sim_trading_nav.csv (date,cash,mv,n_pos,total) + sim_trading_state.json
- 指数: tushare index_daily (000300.SH/000905.SH/000852.SH/932000.CSI)
- 输出: backtest_zt_full/backtest_zt_sim_live_compare.png
"""
import pandas as pd
import numpy as np
import json, sys, re
sys.path.insert(0, '/Users/xpresso/zt_app/code')
import grid_or as g
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
from matplotlib import font_manager

# 中文字体
for f in ['/System/Library/Fonts/STHeiti Medium.ttc',
          '/System/Library/Fonts/STHeiti Medium.ttc',
          '/System/Library/Fonts/STHeiti Medium.ttc']:
    try:
        font_manager.fontManager.addfont(f)
        plt.rcParams['font.family'] = font_manager.FontProperties(fname=f).get_name()
        break
    except Exception:
        continue
plt.rcParams['axes.unicode_minus'] = False

BASE = g.BASE  # backtest_zt_full(图输出目录)
QUANT_DIR = '/Users/xpresso/zt_app'  # 状态文件在quant根(与sim_live_daily.py一致)
STATE = f'{QUANT_DIR}/sim_trading_state.json'
NAV_CSV = f'{QUANT_DIR}/sim_trading_nav.csv'
OUT = f'{BASE}/backtest_zt_sim_live_compare.png'
IDX = [('hs300', '000300.SH', '#8c8c8c'), ('zz500', '000905.SH', '#4a90d9'),
       ('zz1000', '000852.SH', '#f5a623'), ('zz2000', '932000.CSI', '#7ed321')]

# ===== 模拟盘净值 =====
nav = pd.read_csv(NAV_CSV, dtype={'date': str})
nav.columns = ['date', 'cash', 'mv', 'n_pos', 'total']
for col in ['cash', 'mv', 'n_pos', 'total']:
    nav[col] = pd.to_numeric(nav[col], errors='coerce')
st = json.load(open(STATE))
init = st.get('init', 100000.0)
dates0 = nav['date'].astype(str).tolist()
nav_mult0 = nav['total'].values / init          # 相对初始倍数
pos_pct0 = nav['mv'].values / nav['total'].values  # 仓位%
# 基准日=模拟盘启动前一交易日(资金到位日): 插入虚拟点(1.0, 仓位0)
start_day = str(int(dates0[0]) - 1)  # 近似前一日(交易日历由指数数据对齐修正)


# ===== 四指数(基准日=启动前一交易日, 0901口径) =====
import tushare as ts
src = open('/Users/xpresso/zt_app/code/update_daily.py').read()
tok = re.search(r'ts\.set_token\([\"\x27]([^\"\x27]+)', src)
pro = ts.pro_api(tok.group(1))
end = dates0[-1].replace('-', '')
# 防非法日期: int减法跨月会出20260891这种非法串, 统一用datetime回退15天
import datetime as _dt
_d0 = _dt.datetime.strptime(dates0[0].replace('-',''), '%Y%m%d')
early = (_d0 - _dt.timedelta(days=15)).strftime('%Y%m%d')
ref_ix = pro.index_daily(ts_code='000300.SH', start_date=early, end_date=dates0[0].replace('-', ''))
ref_ix = ref_ix.sort_values('trade_date')
base_day = str(ref_ix[ref_ix['trade_date'] < dates0[0]]['trade_date'].iloc[-1])  # 如20260901
dates_full = [base_day] + dates0  # 五条线共同x轴: 基准日+模拟盘各日
fig, axes = plt.subplots(2, 1, figsize=(13, 9), gridspec_kw={'height_ratios': [3, 1]})
ax1, ax2 = axes

# 主面板: 净值对比(基准日=1.0: 指数0901收盘买入持有; 模拟盘0901资金到位10万, 0902建仓)
for nm, code, col in IDX:
    try:
        ix = pro.index_daily(ts_code=code, start_date=early, end_date=end)
        ix = ix.sort_values('trade_date').reset_index(drop=True)
        ix = ix[ix['trade_date'].isin(dates_full)].reset_index(drop=True)
        if len(ix) == 0:
            continue
        base = float(ix.iloc[0]['close'])  # 基准日收盘=1.0
        m = ix['close'].values / base
        ax1.plot(range(len(m)), m, color=col, lw=1.3, alpha=0.9, label=f'{nm}买入持有 {m[-1]*100-100:+.1f}%')
        print(f"  {nm}({code}): {len(m)}天(含基准日{base_day}) 期末{m[-1]*100-100:+.2f}%", flush=True)
    except Exception as e:
        print(f"  {nm} 指数拉取失败: {e}", flush=True)

nav_mult = np.concatenate([[1.0], nav_mult0])   # 基准日资金到位=1.0
pos_pct = np.concatenate([[0.0], pos_pct0])     # 基准日空仓
n_pos_cur = len(st.get('positions', {})) if st.get('positions') else int(nav['n_pos'].iloc[-1])
pos_pct_cur = pos_pct0[-1]
print(f"模拟盘: {len(nav_mult0)}天 ({dates0[0]}~{dates0[-1]}) 最新净值{nav['total'].iloc[-1]:.0f} ({nav_mult0[-1]*100-100:+.2f}%)", flush=True)
ax1.plot(range(len(nav_mult)), nav_mult, color='#d0021b', lw=2.2, label=f'模拟盘(10万) {nav_mult0[-1]*100-100:+.1f}%')
ax1.axhline(1.0, color='#999', lw=0.8, ls='--')
ax1.set_title(f'模拟盘净值 vs 四指数买入持有(基准{base_day}=1.0)', fontsize=13)
ax1.set_ylabel('相对基准倍数')
ax1.legend(loc='upper left', fontsize=10)
ax1.grid(alpha=0.25)
xt = range(len(dates_full))
ax1.set_xticks(list(xt))
ax1.set_xticklabels(dates_full, fontsize=9)

# 下面板: 仓位%
ax2.fill_between(range(len(pos_pct)), pos_pct * 100, alpha=0.35, color='#d0021b')
ax2.plot(range(len(pos_pct)), pos_pct * 100, color='#d0021b', lw=1.8)
ax2.set_ylabel('仓位%', fontsize=10)
ax2.set_ylim(0, 105)
ax2.set_title(f'持仓仓位(持仓市值/总资产) | 当前{int(n_pos_cur)}只 仓位{pos_pct_cur*100:.0f}%', fontsize=10)
ax2.grid(alpha=0.25)
ax2.set_xticks(list(xt))
ax2.set_xticklabels(dates_full, fontsize=9)

fig.suptitle(f'模拟盘实盘跟踪 (资金{base_day}到位·{dates0[0]}建仓, 初始10万, 最多{st.get("params", {}).get("max_pos", 15)}只, 强度≥{st.get("params", {}).get("min_strength", 85)})', fontsize=13)
plt.tight_layout(rect=[0, 0, 1, 0.96])
plt.savefig(OUT, dpi=110, bbox_inches='tight')
print(f"图已保存: {OUT}", flush=True)
