"""
🌀 状态机策略 — 彻底跳出"信号确认"框架
核心思想：
  市场每时每刻都处于某种"状态"
  状态本身有惯性（趋势延续）
  状态转换时有最大概率优势
  
不做"预测"，只做"跟随+提前半步"
"""

import pandas as pd
import numpy as np
import warnings
from itertools import product
warnings.filterwarnings('ignore')

df = pd.read_parquet("/root/quant_pipeline/data/btc_15m_binance.parquet")
df['time'] = pd.to_datetime(df['ts'])
df = df.set_index('time').sort_index()

o = df['open'].values.astype(float)
h = df['high'].values.astype(float)
l = df['low'].values.astype(float)
c = df['close'].values.astype(float)
v = df['volume'].values.astype(float)
N = len(df)

print(f"📊 数据: {N:,}根15分K线 ≈ {N//96}天")
print("="*100)

# ═══ 定义市场状态 ═══
def classify_state(idx, lookback=8):
    """
    把当前市场归类到某个状态
    返回: (状态名, 状态强度)
    
    状态定义：
    - STRONG_UP: 连续上涨，高点越来越高
    - STRONG_DOWN: 连续下跌，低点越来越低
    - RANGING: 横盘震荡
    - PUMP: 突然放量拉升
    - DUMP: 突然放量砸盘
    - SQUEEZE: 波动率压缩到极致（即将突破）
    - REVERSAL_UP: 下跌后出现反转迹象
    - REVERSAL_DOWN: 上涨后出现反转迹象
    """
    if idx < lookback + 3:
        return None, 0
    
    # 基础统计
    recent_closes = c[idx-lookback:idx+1]
    recent_highs = h[idx-lookback:idx+1]
    recent_lows = l[idx-lookback:idx+1]
    recent_vols = v[idx-lookback:idx+1]
    
    # 涨跌幅
    total_ret = (recent_closes[-1] / recent_closes[0] - 1) * 100
    
    # 波动率
    range_pct = (np.max(recent_highs) / np.min(recent_lows) - 1) * 100
    
    # 成交量变化
    vol_ratio = np.mean(recent_vols[-3:]) / (np.mean(recent_vols[:-3]) + 1e-9)
    
    # 创新高/新低
    new_high = recent_closes[-1] > np.max(recent_closes[:-1])
    new_low = recent_closes[-1] < np.min(recent_closes[:-1])
    
    # 连续涨跌
    up_streak = 0
    down_streak = 0
    for i in range(len(recent_closes)-1, 0, -1):
        if recent_closes[i] > recent_closes[i-1]:
            up_streak += 1
            down_streak = 0
        elif recent_closes[i] < recent_closes[i-1]:
            down_streak += 1
            up_streak = 0
        else:
            break
    
    # ── 状态判断 ──
    state = "UNKNOWN"
    strength = 0
    
    # STRONG_UP
    if up_streak >= 4 and total_ret > 1.5 and new_high:
        state = "STRONG_UP"
        strength = min(up_streak / 6, 1.0)
    
    # STRONG_DOWN
    elif down_streak >= 4 and total_ret < -1.5 and new_low:
        state = "STRONG_DOWN"
        strength = min(down_streak / 6, 1.0)
    
    # PUMP
    elif vol_ratio > 2.5 and total_ret > 1.0 and up_streak >= 2:
        state = "PUMP"
        strength = min(vol_ratio / 4, 1.0)
    
    # DUMP
    elif vol_ratio > 2.5 and total_ret < -1.0 and down_streak >= 2:
        state = "DUMP"
        strength = min(vol_ratio / 4, 1.0)
    
    # SQUEEZE (波动率压缩)
    elif range_pct < 0.5 and np.std(recent_closes) / np.mean(recent_closes) < 0.003:
        state = "SQUEEZE"
        strength = 1.0 - range_pct * 2
    
    # RANGING
    elif abs(total_ret) < 0.5 and range_pct < 1.0:
        state = "RANGING"
        strength = 1.0 - abs(total_ret)
    
    # REVERSAL_UP (急跌后阳线反包)
    elif down_streak >= 3 and recent_closes[-1] > recent_closes[-2] and \
         (recent_closes[-1] - recent_closes[-2]) > (recent_closes[-2] - recent_closes[-3]):
        state = "REVERSAL_UP"
        strength = min(up_streak / 4, 1.0)
    
    # REVERSAL_DOWN
    elif up_streak >= 3 and recent_closes[-1] < recent_closes[-2] and \
         (recent_closes[-2] - recent_closes[-1]) > (recent_closes[-3] - recent_closes[-2]):
        state = "REVERSAL_DOWN"
        strength = min(down_streak / 4, 1.0)
    
    # 默认
    else:
        if total_ret > 0.3:
            state = "WEAK_UP"
            strength = abs(total_ret) * 2
        elif total_ret < -0.3:
            state = "WEAK_DOWN"
            strength = abs(total_ret) * 2
        else:
            state = "NEUTRAL"
            strength = 0.3
    
    return state, strength


# ═══ 状态转移矩阵 ═══
def predict_next_move(idx, current_state, strength):
    """
    根据当前状态和强度，预测下一根K线的方向
    返回: (方向, 置信度)
    方向: 1=做多, -1=做空, 0=观望
    """
    if current_state is None:
        return 0, 0
    
    # 基于状态惯性的简单规则
    if current_state == "STRONG_UP":
        # 强上涨 → 继续多，但强度太高就小心回调
        if strength > 0.8:
            return 0, 0.5  # 太高了，观望
        else:
            return 1, 0.6 + strength * 0.3
    
    elif current_state == "STRONG_DOWN":
        if strength > 0.8:
            return 0, 0.5
        else:
            return -1, 0.6 + strength * 0.3
    
    elif current_state == "PUMP":
        # 放量拉升 → 惯性继续冲一下
        return 1, 0.55
    
    elif current_state == "DUMP":
        return -1, 0.55
    
    elif current_state == "SQUEEZE":
        # 波动率压缩 → 即将突破，但不知道方向，观望
        return 0, 0
    
    elif current_state == "RANGING":
        # 震荡 → 高抛低吸
        # 看当前位置：接近区间上轨做空，下轨做多
        recent = c[idx-8:idx+1]
        high = np.max(h[idx-8:idx+1])
        low = np.min(l[idx-8:idx+1])
        mid = (high + low) / 2
        
        if c[idx] > mid * 1.002:
            return -1, 0.5  # 靠近上轨，做空
        elif c[idx] < mid * 0.998:
            return 1, 0.5   # 靠近下轨，做多
        else:
            return 0, 0
    
    elif current_state == "REVERSAL_UP":
        # 反转向上 → 做多
        return 1, 0.65
    
    elif current_state == "REVERSAL_DOWN":
        return -1, 0.65
    
    elif current_state == "WEAK_UP":
        # 弱上涨 → 顺势小多
        return 1, 0.45
    
    elif current_state == "WEAK_DOWN":
        return -1, 0.45
    
    else:  # NEUTRAL
        return 0, 0


# ═══ 回测 ═══
def backtest_state_machine(name, 
                           min_confidence=0.55,
                           sl_atr=1.0, tp_atr=2.0,
                           max_hold=8,
                           position_sizing="fixed"):
    """
    基于状态机的策略
    """
    trades = []
    
    prev_state = None
    prev_strength = 0
    
    for i in range(50, N - 5):
        # 获取当前状态
        state, strength = classify_state(i, lookback=8)
        
        if state is None:
            continue
        
        # 预测
        direction, confidence = predict_next_move(i, state, strength)
        
        if direction == 0 or confidence < min_confidence:
            prev_state = state
            prev_strength = strength
            continue
        
        # 状态转换检测（只做转换瞬间）
        if prev_state == state:
            # 状态没变，可能是惯性已经跑了一半了，跳过
            prev_state = state
            prev_strength = strength
            continue
        
        # 入场
        entry = c[i]
        cur_atr = atr14[i] if 'atr14' in globals() else np.std(c[max(0,i-14):i+1])
        
        if direction == 1:  # 做多
            sl = entry - cur_atr * sl_atr
            tp = entry + cur_atr * tp_atr
        else:  # 做空
            sl = entry + cur_atr * sl_atr
            tp = entry - cur_atr * tp_atr
        
        hit_sl = hit_tp = False
        exit_price = entry
        exit_idx = i
        
        for j in range(i + 1, min(i + max_hold + 1, N)):
            dh, dl = h[j], l[j]
            
            if direction == 1:
                if dh >= tp:
                    exit_price = tp; hit_tp = True; exit_idx = j; break
                if dl <= sl:
                    exit_price = sl; hit_sl = True; exit_idx = j; break
            else:
                if dl <= tp:
                    exit_price = tp; hit_tp = True; exit_idx = j; break
                if dh >= sl:
                    exit_price = sl; hit_sl = True; exit_idx = j; break
            
            if (j - i) >= max_hold:
                exit_price = c[j]; exit_idx = j; break
        
        net = round((exit_price / entry - 1) * 100 * direction - 0.06, 2)
        
        trades.append({
            "time": df.index[i],
            "state": state,
            "prev_state": prev_state,
            "transition": f"{prev_state}→{state}",
            "direction": "L" if direction == 1 else "S",
            "confidence": confidence,
            "entry": entry,
            "exit": exit_price,
            "bars": exit_idx - i,
            "hit_sl": hit_sl, "hit_tp": hit_tp,
            "net": net,
        })
        
        prev_state = state
        prev_strength = strength
    
    if not trades or len(trades) < 10:
        return {"name": name, "trades": 0}
    
    df_t = pd.DataFrame(trades)
    wins = sum(1 for t in trades if t['net'] > 0)
    total = len(trades)
    avg_ret = np.mean([t['net'] for t in trades])
    
    pnl = np.array([t['net'] for t in trades]) / 100
    cum = (1 + pnl).prod()
    
    eq = [1.0]
    for r in pnl:
        eq.append(eq[-1] * (1 + r))
    peak = np.maximum.accumulate(eq)
    dd = (np.array(eq) - peak) / peak * 100
    max_dd = np.min(dd)
    
    sharpe = (np.mean(pnl) / np.std(pnl) * np.sqrt(96*365)) if np.std(pnl) > 0 else 0
    
    # 按状态统计
    state_stats = df_t.groupby('state').agg(
        笔数=('net', 'count'),
        胜率=('net', lambda x: (x>0).mean()*100),
        平均=('net', 'mean'),
        合计=('net', 'sum')
    ).round(2)
    
    # 按转换类型统计
    trans_stats = df_t.groupby('transition').agg(
        笔数=('net', 'count'),
        胜率=('net', lambda x: (x>0).mean()*100),
        平均=('net', 'mean')
    ).round(2)
    
    # 日频统计
    df_t['day'] = df_t['time'].dt.strftime('%Y-%m-%d')
    daily = df_t.groupby('day').agg(
        笔数=('net', 'count'),
        日收益=('net', 'sum'),
        胜率=('net', lambda x: (x>0).mean()*100)
    )
    
    profit_days = (daily['日收益'] > 0).sum()
    total_days = len(daily)
    avg_daily = daily['日收益'].mean()
    
    return {
        "name": name, "trades": total, "wins": wins,
        "win_rate": wins/total*100, "avg_ret": avg_ret,
        "cum": cum, "max_dd": max_dd,
        "sharpe": round(sharpe, 2),
        "profit_days": profit_days, "total_days": total_days,
        "daily_profit_rate": profit_days/total_days*100,
        "avg_daily": avg_daily,
        "state_stats": state_stats,
        "trans_stats": trans_stats,
        "df": df_t,
    }


# 计算ATR
atr14 = pd.Series(np.maximum(h-l, np.maximum(np.abs(h-np.roll(c,1)), np.abs(l-np.roll(c,1))))).rolling(14).mean().values

# ═══ 方案 ═══
scenarios = [
    ("A 标准", 0.55, 1.0, 2.0, 8),
    ("B 高门槛", 0.65, 1.0, 2.0, 8),
    ("C 低门槛", 0.45, 1.0, 2.0, 8),
    ("D 宽止损", 0.55, 1.5, 2.5, 10),
    ("E 窄止损", 0.55, 0.8, 1.5, 6),
    ("F 超短", 0.55, 0.8, 1.2, 4),
    ("G 长持", 0.55, 1.2, 3.0, 16),
    ("H 只做反转", 0.60, 1.0, 2.0, 8),  # 需要改代码过滤
]

results = []
for name, min_conf, sl_m, tp_m, hold in scenarios:
    r = backtest_state_machine(name, min_conf, sl_m, tp_m, hold)
    results.append(r)
    if r["trades"] > 0:
        flag = "✅" if r["sharpe"] > 1.0 else ("⚠️" if r["sharpe"] > 0 else "❌")
        print(f"  {flag} {name}: {r['trades']}笔 {r['win_rate']:.0f}%胜率 平均{r['avg_ret']:+.2f}% Sharpe={r['sharpe']:.1f} 复利{r['cum']:.2f}x 盈利天数{r['daily_profit_rate']:.0f}%")
    else:
        print(f"  ⏭️ {name}: 0笔")

# ── 主表 ──
print(f"\n{'='*120}")
print(f"{'📊 状态机策略对比':^120}")
print(f"{'='*120}")
print(f"{'方案':<18} {'笔数':>6} {'胜率':>7} {'平均%':>8} {'复利':>10} {'回撤':>8} {'最好':>7} {'最差':>7} {'Sharpe':>7} {'盈利天%':>9}")
print(f"{'-'*120}")

for r in results:
    if r.get("trades", 0) > 0:
        wc = "🟢" if r["sharpe"] > 2.0 else ("🟡" if r["sharpe"] > 0.5 else "🔴")
        print(f"{wc} {r['name']:<16} {r['trades']:>6d} {r['win_rate']:>5.0f}% {r['avg_ret']:>+7.2f}% {r['cum']:>8.2f}x {r['max_dd']:>+7.1f}% {r.get('best',0):>+6.2f}% {r.get('worst',0):>+6.2f}% {r['sharpe']:>6.1f} {r['daily_profit_rate']:>6.0f}%")

# ── 最佳方案深度 ──
print(f"\n{'='*120}")
print(f"🏆 最佳方案分析")
print(f"{'='*120}")

best = max(results, key=lambda x: x.get("sharpe", 0) if x.get("trades", 0) > 20 else 0)

if best and best["trades"] > 0:
    print(f"\n🏆 {best['name']}")
    print(f"   总交易: {best['trades']}笔")
    print(f"   胜率: {best['wins']}/{best['trades']} = {best['win_rate']:.1f}%")
    print(f"   平均: {best['avg_ret']:+.2f}%")
    print(f"   复利: {best['cum']:.2f}x  回撤: {best['max_dd']:.1f}%")
    print(f"   Sharpe: {best['sharpe']}")
    print(f"   盈利天数: {best['profit_days']}/{best['total_days']} = {best['daily_profit_rate']:.0f}%")
    print(f"   平均日收益: {best['avg_daily']:+.2f}%")
    
    # 按状态统计
    print(f"\n   📊 各状态表现:")
    print(f"   {'状态':<18} {'笔数':>5} {'胜率':>7} {'平均%':>8} {'合计%':>8}")
    for state, row in best['state_stats'].iterrows():
        tag = "🟢" if row['平均'] > 0.2 else ("🟡" if row['平均'] > 0 else "🔴")
        print(f"   {tag} {state:<16} {row['笔数']:>5.0f} {row['胜率']:>5.0f}% {row['平均']:>+7.2f}% {row['合计']:>+7.0f}%")
    
    # 按转换统计（Top 10）
    print(f"\n   📊 最佳状态转换 (Top 10):")
    top_trans = best['trans_stats'].sort_values('平均', ascending=False).head(10)
    print(f"   {'转换':<25} {'笔数':>5} {'胜率':>7} {'平均%':>8}")
    for trans, row in top_trans.iterrows():
        tag = "🟢" if row['平均'] > 0.3 else ("🟡" if row['平均'] > 0 else "🔴")
        print(f"   {tag} {trans:<23} {row['笔数']:>5.0f} {row['胜率']:>5.0f}% {row['平均']:>+7.2f}%")
    
    # 日报表
    df_best = best['df']
    df_best['day'] = df_best['time'].dt.strftime('%Y-%m-%d')
    daily = df_best.groupby('day').agg(
        笔数=('net', 'count'),
        日收益=('net', 'sum'),
        胜率=('net', lambda x: (x>0).mean()*100)
    )
    
    print(f"\n   📅 前30天:")
    print(f"   {'日期':<12} {'笔数':>4} {'日收益':>8} {'胜率':>6}")
    print(f"   {'-'*33}")
    for d, row in daily.head(30).iterrows():
        tag = "🟢" if row['日收益'] > 0 else "🔴"
        print(f"   {d:<12} {row['笔数']:>4.0f} {tag} {row['日收益']:>+6.2f}% {row['胜率']:>5.0f}%")
    
    # 逐笔
    print(f"\n   📋 前20笔:")
    print(f"   {'时间':<16} {'状态转换':<20} {'方向':>4} {'入场':>8} {'出场':>8} {'收益':>7}")
    print(f"   {'-'*70}")
    for _, t in df_best.head(20).iterrows():
        tag = "🟢" if t['net'] > 0 else "🔴"
        print(f"   {t['time'].strftime('%m-%d %H:%M'):<16} {t['transition']:<20} {t['direction']:>4} {t['entry']:>8.0f} {t['exit']:>8.0f} {tag} {t['net']:>+5.2f}%")

# ── 结论 ──
print(f"\n{'='*120}")
print(f"💡 状态机方法的洞察")
print(f"{'='*120}")
print(f"""
状态机 vs 传统信号确认：

传统方法：等条件 → 赌方向
状态机：识别状态 → 跟随惯性 → 抓转换瞬间

关键差异：
  ① 不预测顶底，只识别"现在是什么状态"
  ② 不做所有行情，只做"状态转换"那一刻
  ③ 胜率来自"状态的惯性"，不是来自"指标的准确"

从结果看：
  盈利天数：{best['daily_profit_rate']:.0f}% （目标是>60%才算接近日复利）
  平均日收益：{best['avg_daily']:+.2f}%
  
{"+" * 60}
下一步优化方向：
  ① 增加更多状态细分（比如区分"真突破"和"假突破"）
  ② 用机器学习训练状态分类器（而不是手写规则）
  ③ 加入成交量分布、订单簿快照等更高维特征
{"+" * 60}
""")
