"""
🎯 High Win-Rate Ensemble v2
============================
目标：胜率 > 60%
方法：
  1. 规则状态识别 (比HMM更稳)
  2. 历史回测反推最优权重矩阵
  3. 只在"状态×策略"高适配时开仓
  4. 增加过滤：低波动/低成交量时休息
"""

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

print("="*80)
print("🎯 High Win-Rate Ensemble v2")
print("="*80)

# =====================================================================
# 数据加载
# =====================================================================
df_daily = pd.read_parquet("/root/quant_pipeline/data/btc_daily.parquet")
df_15m = pd.read_parquet("/root/quant_pipeline/data/btc_15m_binance.parquet")
df_15m['time'] = pd.to_datetime(df_15m['ts'])
df_15m = df_15m.set_index('time').sort_index()

for df in [df_daily, df_15m]:
    if 'o' in df.columns:
        df.rename(columns={'o':'open','h':'high','l':'low','c':'close','v':'volume'}, inplace=True)

print(f"日线: {len(df_daily)}根")
print(f"15分: {len(df_15m)}根")

# =====================================================================
# 工具函数
# =====================================================================
def roll_mean(arr, n): return pd.Series(arr).rolling(n).mean().values
def ewm(arr, span): return pd.Series(arr).ewm(span=span, adjust=False).mean().values
def rsi(arr, period=14):
    delta = pd.Series(arr).diff()
    gain = delta.clip(lower=0).rolling(period).mean()
    loss = (-delta.clip(upper=0)).rolling(period).mean()
    return (100 - 100/(1+gain/(loss+1e-9))).values

# =====================================================================
# 1. 规则状态识别器 (比HMM更稳)
# =====================================================================
class RuleRegimeDetector:
    """基于明确规则的状态识别"""
    
    def __init__(self):
        pass
    
    def detect_regime(self, df_daily):
        c = df_daily['close'].values
        v = df_daily['volume'].values
        
        ma20 = pd.Series(c).rolling(20).mean().values
        ma50 = pd.Series(c).rolling(50).mean().values
        ma200 = pd.Series(c).rolling(200).mean().values
        
        rsi14 = rsi(c, 14)
        
        # 成交量趋势
        vol_ma20 = pd.Series(v).rolling(20).mean().values
        vol_trend = np.where(v > vol_ma20 * 1.1, 1, np.where(v < vol_ma20 * 0.9, -1, 0))
        
        regimes = []
        for i in range(200, len(c)):
            # BULL: 价格>MA20>MA50, RSI>55, 成交量放大
            if (c[i] > ma20[i] > ma50[i] and 
                rsi14[i] > 55 and 
                vol_trend[i] == 1):
                regimes.append("BULL")
            
            # BEAR: 价格<MA20<MA50, RSI<45, 成交量放大
            elif (c[i] < ma20[i] < ma50[i] and 
                  rsi14[i] < 45 and 
                  vol_trend[i] == 1):
                regimes.append("BEAR")
            
            # SIDEWAYS_HIGH_VOL: 价格在MA附近，但成交量大（即将突破）
            elif (abs(c[i] - ma20[i]) / ma20[i] < 0.02 and 
                  vol_trend[i] == 1):
                regimes.append("BREAKOUT_IMMINENT")
            
            # SIDEWAYS_LOW_VOL: 价格在MA附近，成交量小（真震荡）
            elif (abs(c[i] - ma20[i]) / ma20[i] < 0.03):
                regimes.append("SIDEWAYS")
            
            # TRANSITION: 其他过渡状态
            else:
                regimes.append("TRANSITION")
        
        # 填充前200天
        regimes = ["UNKNOWN"] * 200 + regimes
        
        return regimes


# =====================================================================
# 2. 四个基础策略 (精简版)
# =====================================================================

def strat_trend_following(df_daily):
    """趋势跟踪：只做多，只在明确牛市"""
    c = df_daily['close'].values
    v = df_daily['volume'].values
    N = len(c)
    
    ma50 = pd.Series(c).rolling(50).mean().values
    ma200 = pd.Series(c).rolling(200).mean().values
    rsi14 = rsi(c, 14)
    
    signals = np.zeros(N)
    position = 0
    
    for i in range(200, N):
        # 入场：牛市确认 + RSI适中
        if (c[i] > ma50[i] > ma200[i] and 
            50 < rsi14[i] < 75 and 
            position == 0):
            position = 1
            signals[i] = 1
        
        # 出场：跌破MA50或RSI超买
        elif (c[i] < ma50[i] or rsi14[i] > 80) and position == 1:
            position = 0
            signals[i] = 0
        
        else:
            signals[i] = position
    
    return signals


def strat_mean_reversion(df_15m):
    """均值回归：只做震荡市，高抛低吸"""
    c = df_15m['close'].values
    h = df_15m['high'].values
    l = df_15m['low'].values
    v = df_15m['volume'].values
    N = len(c)
    
    ma20 = pd.Series(c).rolling(20).mean().values
    std20 = pd.Series(c).rolling(20).std().values
    upper = ma20 + 1.8 * std20
    lower = ma20 - 1.8 * std20
    rsi14 = rsi(c, 14)
    
    signals = np.zeros(N)
    position = 0
    
    for i in range(50, N):
        if position == 0:
            # 做多：触及下轨 + RSI<35
            if c[i] <= lower[i] and rsi14[i] < 35:
                position = 1
                signals[i] = 1
            # 做空：触及上轨 + RSI>65
            elif c[i] >= upper[i] and rsi14[i] > 65:
                position = -1
                signals[i] = -1
        
        elif position == 1:
            if c[i] >= ma20[i]:  # 回归中轨止盈
                position = 0
                signals[i] = 0
            elif c[i] <= lower[i] * 0.99:  # 止损
                position = 0
                signals[i] = 0
            else:
                signals[i] = 1
        
        elif position == -1:
            if c[i] <= ma20[i]:
                position = 0
                signals[i] = 0
            elif c[i] >= upper[i] * 1.01:
                position = 0
                signals[i] = 0
            else:
                signals[i] = -1
    
    return signals


def strat_breakout(df_15m):
    """突破策略：只做成交量确认的突破"""
    c = df_15m['close'].values
    h = df_15m['high'].values
    l = df_15m['low'].values
    v = df_15m['volume'].values
    N = len(c)
    
    period = 20
    upper = pd.Series(h).rolling(period).max().values
    lower = pd.Series(l).rolling(period).min().values
    mid = (upper + lower) / 2
    vol_ma20 = pd.Series(v).rolling(20).mean().values
    
    signals = np.zeros(N)
    position = 0
    
    for i in range(50, N):
        if position == 0:
            # 向上突破：价格>上轨 + 成交量>1.8倍
            if c[i] > upper[i-1] and v[i] > vol_ma20[i] * 1.8:
                position = 1
                signals[i] = 1
            # 向下突破
            elif c[i] < lower[i-1] and v[i] > vol_ma20[i] * 1.8:
                position = -1
                signals[i] = -1
        
        elif position == 1:
            if c[i] < mid[i]:
                position = 0
                signals[i] = 0
            else:
                signals[i] = 1
        
        elif position == -1:
            if c[i] > mid[i]:
                position = 0
                signals[i] = 0
            else:
                signals[i] = -1
    
    return signals


def strat_momentum_reversal(df_15m):
    """动量反转：大跌后4层确认"""
    c = df_15m['close'].values
    o = df_15m['open'].values
    h = df_15m['high'].values
    l = df_15m['low'].values
    v = df_15m['volume'].values
    N = len(c)
    
    ema9 = ewm(c, 9)
    ema21 = ewm(c, 21)
    ema50 = ewm(c, 50)
    
    signals = np.zeros(N)
    
    for i in range(60, len(c) - 10):
        # 1H趋势向上
        if ema21[i] < ema50[i]:
            continue
        
        # 大跌>2%
        recent_high = np.max(h[i-12:i+1])
        drop_pct = (c[i] / recent_high - 1) * 100
        if drop_pct > -2.0:
            continue
        
        # 等待确认
        confirmed = False
        for offset in range(1, 7):
            ci = i + offset
            if ci >= len(c) - 5:
                break
            
            confirms = 0
            if c[ci] > o[ci]: confirms += 1
            if c[ci] > ema9[ci]: confirms += 1
            if c[ci] > h[ci-1]: confirms += 1
            if v[ci] >= v[ci-1]: confirms += 1
            if c[ci] > o[ci]:
                body = c[ci] - o[ci]
                prev_body = abs(c[ci-1] - o[ci-1])
                if body >= prev_body * 0.5: confirms += 1
            if l[ci] > l[ci-1] * 1.001: confirms += 1
            
            if confirms >= 4:
                confirmed = True
                break
        
        if confirmed:
            signals[i] = 1
    
    return signals


# =====================================================================
# 3. 生成所有策略信号
# =====================================================================
print("\n生成策略信号...")

sig_tf = strat_trend_following(df_daily)
sig_mr = strat_mean_reversion(df_15m)
sig_bo = strat_breakout(df_15m)
sig_mrev = strat_momentum_reversal(df_15m)

print(f"  TrendFollowing: {np.sum(sig_tf != 0)}个信号")
print(f"  MeanReversion: {np.sum(sig_mr != 0)}个信号")
print(f"  Breakout: {np.sum(sig_bo != 0)}个信号")
print(f"  MomentumReversal: {np.sum(sig_mrev != 0)}个信号")

# =====================================================================
# 4. 状态识别
# =====================================================================
print("识别市场状态...")
regime_detector = RuleRegimeDetector()
regimes = regime_detector.detect_regime(df_daily)
regime_counts = pd.Series(regimes).value_counts()
print(f"状态分布: {regime_counts.to_dict()}")

# =====================================================================
# 5. 高胜率组合逻辑
# =====================================================================
"""
核心思想：
  - 只在"状态×策略"高适配时开仓
  - 每个状态只选1-2个最佳策略
  - 低置信度时休息
"""

# 状态→策略映射 (基于历史经验)
STATE_STRATEGY_MAP = {
    "BULL": ["TrendFollowing"],           # 牛市只做趋势
    "BEAR": ["MomentumReversal"],         # 熊市只做反转
    "SIDEWAYS": ["MeanReversion"],        # 震荡只做均值回归
    "BREAKOUT_IMMINENT": ["Breakout"],    # 即将突破只做突破
    "TRANSITION": [],                      # 过渡期休息
    "UNKNOWN": [],
}

# 将日线状态映射到15分钟
# 简单方法：用最近的日线状态
regime_map_15m = {}
daily_dates = df_daily.index
for i, date in enumerate(daily_dates):
    if i < len(regimes):
        # 找到这一天对应的15分钟K线范围
        mask = (df_15m.index.date == date.date())
        for idx in df_15m[mask].index:
            regime_map_15m[idx] = regimes[i]

# 填充剩余部分
last_regime = regimes[-1]
for idx in df_15m.index:
    if idx not in regime_map_15m:
        regime_map_15m[idx] = last_regime

regimes_15m = [regime_map_15m.get(idx, "UNKNOWN") for idx in df_15m.index]

# =====================================================================
# 6. 组合回测
# =====================================================================
print("\n运行高胜率组合回测...")

N_15m = len(df_15m)
c_15m = df_15m['close'].values

portfolio_signals = np.zeros(N_15m)
win_count = 0
loss_count = 0
trade_results = []

for i in range(100, N_15m - 10):
    current_regime = regimes_15m[i] if i < len(regimes_15m) else "UNKNOWN"
    
    # 获取当前状态允许的策略
    allowed_strats = STATE_STRATEGY_MAP.get(current_regime, [])
    
    if not allowed_strats:
        portfolio_signals[i] = 0
        continue
    
    # 聚合允许策略的信号
    signal_sum = 0
    count = 0
    
    if "TrendFollowing" in allowed_strats:
        # 日线信号映射到15分钟 (简化：用前一天的信号)
        daily_idx = min(i // 96, len(sig_tf) - 1)
        signal_sum += sig_tf[daily_idx]
        count += 1
    
    if "MeanReversion" in allowed_strats:
        signal_sum += sig_mr[i]
        count += 1
    
    if "Breakout" in allowed_strats:
        signal_sum += sig_bo[i]
        count += 1
    
    if "MomentumReversal" in allowed_strats:
        signal_sum += sig_mrev[i]
        count += 1
    
    # 只有当所有允许的策略方向一致时才开仓
    if count > 0:
        avg_signal = signal_sum / count
        if abs(avg_signal) > 0.5:  # 高置信度
            portfolio_signals[i] = np.sign(avg_signal)
        else:
            portfolio_signals[i] = 0
    else:
        portfolio_signals[i] = 0

# 计算收益
equity = 1000
equity_curve = [1000]
daily_pnl = []

for i in range(100, N_15m - 1):
    if portfolio_signals[i] != 0:
        # 持仓一天
        ret = (c_15m[i+1] / c_15m[i] - 1) * portfolio_signals[i]
        ret -= 0.0003 * abs(portfolio_signals[i])  # 手续费
        
        equity *= (1 + ret)
        equity_curve.append(equity)
        daily_pnl.append(ret)
        
        if ret > 0:
            win_count += 1
        else:
            loss_count += 1

# 统计
total_trades = win_count + loss_count
win_rate = win_count / total_trades * 100 if total_trades > 0 else 0
total_ret = (equity_curve[-1] / 1000 - 1) * 100

eq_arr = np.array(equity_curve)
peak = np.maximum.accumulate(eq_arr)
dd = (eq_arr - peak) / peak * 100
max_dd = np.min(dd)

pnl_series = np.array(daily_pnl)
sharpe = np.mean(pnl_series) / np.std(pnl_series) * np.sqrt(252) if np.std(pnl_series) > 0 else 0

# =====================================================================
# 7. 输出结果
# =====================================================================
print(f"\n{'='*80}")
print(f"📊 高胜率组合回测结果")
print(f"{'='*80}")
print(f"总交易次数: {total_trades}")
print(f"盈利次数: {win_count}")
print(f"亏损次数: {loss_count}")
print(f"{'='*40}")
print(f"✅ 胜率: {win_rate:.1f}%")
print(f"总收益: {total_ret:+.1f}%")
print(f"Sharpe: {sharpe:.2f}")
print(f"最大回撤: {max_dd:.1f}%")
print(f"最终权益: ${equity_curve[-1]:.0f}")

# 按状态统计
print(f"\n{'='*80}")
print(f"📊 按状态统计胜率")
print(f"{'='*80}")

state_stats = {}
for i in range(100, N_15m - 1):
    if portfolio_signals[i] == 0:
        continue
    
    regime = regimes_15m[i] if i < len(regimes_15m) else "UNKNOWN"
    ret = (c_15m[i+1] / c_15m[i] - 1) * portfolio_signals[i]
    
    if regime not in state_stats:
        state_stats[regime] = {"wins": 0, "losses": 0, "total_ret": 0}
    
    state_stats[regime]["total_ret"] += ret
    if ret > 0:
        state_stats[regime]["wins"] += 1
    else:
        state_stats[regime]["losses"] += 1

for regime, stats in sorted(state_stats.items()):
    total = stats["wins"] + stats["losses"]
    wr = stats["wins"] / total * 100 if total > 0 else 0
    print(f"  {regime:<20}: {wr:>5.1f}%胜率 ({stats['wins']}/{total}) 累计收益{stats['total_ret']*100:+.1f}%")

print(f"\n{'='*80}")
if win_rate >= 60:
    print("🎉 达成目标！胜率 >= 60%")
elif win_rate >= 50:
    print("⚠️ 胜率接近50%，需要继续优化")
else:
    print("❌ 胜率不足，需要调整策略映射")
print(f"{'='*80}")
