"""
🧬 Ensemble + Regime Switching Framework
==========================================
核心架构：
  基础策略池 (4个不同类型) → 状态识别器 → 动态权重分配 → 组合执行
  
策略池：
  1. TrendFollowing (日线动量，C1)     - 擅长牛市趋势
  2. MeanReversion (15分均值回归)       - 擅长震荡市
  3. Breakout (15分突破)               - 擅长启动行情
  4. MomentumReversal (15分大跌反转)    - 擅长超跌反弹

状态识别：
  HMM 3状态 (牛市/熊市/震荡) + 规则兜底

权重分配：
  状态 × 策略适应性矩阵 → 软最大化归一化
"""

import pandas as pd
import numpy as np
import warnings
from itertools import product
from hmmlearn import hmm
from sklearn.preprocessing import StandardScaler
warnings.filterwarnings('ignore')

print("="*80)
print("🧬 Ensemble + Regime Switching Framework v1.0")
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)}根  {df_daily.index[0].date()} ~ {df_daily.index[-1].date()}")
print(f"15分: {len(df_15m)}根  {df_15m.index[0]} ~ {df_15m.index[-1]}")

# =====================================================================
# 基础工具函数
# =====================================================================
def roll_mean(arr, n): return pd.Series(arr).rolling(n).mean().values
def roll_std(arr, n): return pd.Series(arr).rolling(n).std().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

def atr(h, l, c, period=14):
    tr = np.maximum(h - l, np.maximum(np.abs(h - np.roll(c, 1)), np.abs(l - np.roll(c, 1))))
    return pd.Series(tr).rolling(period).mean().values

# =====================================================================
# 1. 基础策略基类
# =====================================================================
class BaseStrategy:
    """所有策略的基类"""
    def __init__(self, name, timeframe='15m'):
        self.name = name
        self.timeframe = timeframe
        self.description = ""
    
    def generate_signals(self, df):
        """返回: 1=做多, -1=做空, 0=空仓"""
        raise NotImplementedError
    
    def get_position_size(self, df, signal, capital, risk_pct=0.02):
        """计算仓位大小"""
        return signal * capital * risk_pct

# =====================================================================
# 2. 四个基础策略实现
# =====================================================================

class TrendFollowingDaily(BaseStrategy):
    """策略1: 日线趋势动量 (C1方案)"""
    def __init__(self):
        super().__init__("TrendFollowingDaily", "1d")
        self.description = "RSI>75 + 远高于MA50 + 放量，日线级别趋势跟踪"
    
    def generate_signals(self, df_daily, df_15m=None):
        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)
        vol_ma20 = pd.Series(v).rolling(20).mean().values
        vol_ratio = v / np.maximum(vol_ma20, 1)
        
        signals = np.zeros(N)
        for i in range(100, N):
            # C1入场条件：RSI>75 + 价格>MA50*1.05 + 放量
            if (rsi14[i] > 75 and 
                c[i] > ma50[i] * 1.05 and 
                v[i] > v[i-1] * 1.3):
                signals[i] = 1
            # 出场：跌破MA50
            elif c[i] < ma50[i]:
                signals[i] = 0
            else:
                signals[i] = signals[i-1] if i > 0 else 0
        
        return signals


class MeanReversion15m(BaseStrategy):
    """策略2: 15分钟均值回归 (布林带 + RSI)"""
    def __init__(self):
        super().__init__("MeanReversion15m", "15m")
        self.description = "BBands + RSI超卖/超买，适合震荡市"
    
    def generate_signals(self, df):
        c = df['close'].values
        h = df['high'].values
        l = df['low'].values
        v = df['volume'].values
        N = len(c)
        
        # 布林带
        ma20 = pd.Series(c).rolling(20).mean().values
        std20 = pd.Series(c).rolling(20).std().values
        upper = ma20 + 2 * std20
        lower = ma20 - 2 * std20
        
        rsi14 = rsi(c, 14)
        vol_ma20 = pd.Series(df['volume'].values).rolling(20).mean().values
        vol_ratio = v / np.maximum(pd.Series(df['volume'].values).rolling(20).mean().values, 1)
        
        signals = np.zeros(N)
        position = 0
        
        for i in range(50, len(c)):
            if position == 0:
                # 做多：触及下轨 + RSI<30 + 放量
                if c[i] <= lower[i] and rsi14[i] < 35 and df['volume'].values[i] > df['volume'].values[i-1] * 1.2:
                    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] or c[i] <= lower[i] * 0.995:
                    position = 0
                    signals[i] = 0
                else:
                    signals[i] = 1
            
            elif position == -1:
                if c[i] <= ma20[i] or c[i] >= upper[i] * 1.005:
                    position = 0
                    signals[i] = 0
                else:
                    signals[i] = -1
        
        return signals


class Breakout15m(BaseStrategy):
    """策略3: 15分钟突破 (唐奇安通道 + 成交量)"""
    def __init__(self):
        super().__init__("Breakout15m", "15m")
        self.description = "唐奇安通道突破 + 成交量确认，捕捉启动行情"
    
    def generate_signals(self, df):
        c = df['close'].values
        h = df['high'].values
        l = df['low'].values
        v = df['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(df['volume'].values).rolling(20).mean().values
        vol_ratio = v / np.maximum(vol_ma20, 1)
        
        signals = np.zeros(N)
        position = 0
        
        for i in range(50, len(c)):
            if position == 0:
                # 向上突破
                if c[i] > upper[i-1] and v[i] > v[i-1] * 1.5:
                    position = 1
                    signals[i] = 1
                # 向下突破
                elif c[i] < lower[i-1] and v[i] > v[i-1] * 1.5:
                    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


class MomentumReversal15m(BaseStrategy):
    """策略4: 15分钟动量反转 (大跌>2% + 4层确认)"""
    def __init__(self):
        super().__init__("MomentumReversal15m", "15m")
        self.description = "大跌>2%后4层确认反转，之前验证61%胜率"
    
    def generate_signals(self, df):
        c = df['close'].values
        o = df['open'].values
        h = df['high'].values
        l = df['low'].values
        v = df['volume'].values
        N = len(c)
        
        vol_ma20 = pd.Series(v).rolling(20).mean().values
        ema9 = pd.Series(c).ewm(span=9, adjust=False).mean().values
        ema21 = pd.Series(c).ewm(span=21, adjust=False).mean().values
        ema50 = pd.Series(c).ewm(span=50, adjust=False).mean().values
        
        signals = np.zeros(N)
        
        for i in range(60, len(c) - 5):
            # 1. 1H趋势向上
            if ema21[i] < ema50[i]:
                continue
            
            # 2. 检查过去12根是否有大跌>2%
            recent_high = np.max(h[i-12:i+1])
            drop_pct = (c[i] / recent_high - 1) * 100
            if drop_pct > -2.0:
                continue
            
            # 3. 等待反转确认 (最多等6根)
            confirmed = False
            confirm_idx = -1
            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] * 0.9: 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.4: confirms += 1
                if l[ci] > l[ci-1] * 0.999: confirms += 1
                if c[ci] > ema21[ci]: confirms += 1
                
                if confirms >= 4:
                    confirmed = True
                    break
            
            if confirmed:
                signals[i] = 1
        
        return signals


# =====================================================================
# 3. 状态识别器 - HMM + 规则兜底
# =====================================================================
class RegimeDetector:
    """市场状态识别器：HMM + 规则兜底"""
    
    def __init__(self, n_states=3):
        self.n_states = n_states
        self.hmm = hmm.GaussianHMM(n_components=n_states, covariance_type="full", n_iter=100, random_state=42)
        self.scaler = StandardScaler()
        self.fitted = False
        self.state_map = {0: "BEAR", 1: "SIDEWAYS", 2: "BULL"}  # 默认映射
    
    def prepare_features(self, df_daily):
        """准备HMM特征"""
        c = df_daily['close'].values
        v = df_daily['volume'].values
        
        # 特征：收益率、波动率、成交量变化
        ret = pd.Series(c).pct_change().values
        vol_20 = pd.Series(c).pct_change().rolling(20).std().values
        vol_ratio = v / np.maximum(pd.Series(df_daily['volume'].values).rolling(20).mean().values, 1)
        
        # 趋势：简化计算
        ma20 = pd.Series(c).rolling(20).mean().values
        trend = np.where(c > ma20, 1, -1)
        
        features = np.column_stack([ret, vol_20, vol_ratio, trend])
        features = np.nan_to_num(features, nan=0, posinf=0, neginf=0)
        
        return features
    
    def fit(self, df_daily):
        """训练HMM"""
        features = self.prepare_features(df_daily)
        valid = ~np.any(np.isnan(features), axis=1)
        features_clean = features[valid]
        
        if len(features_clean) < 100:
            return False
        
        scaled = self.scaler.fit_transform(features_clean)
        self.hmm.fit(scaled)
        self.fitted = True
        
        # 根据均值收益率重新映射状态
        means = self.hmm.means_[:, 0]  # 第一个特征是收益率
        order = np.argsort(means)
        self.state_map = {order[0]: "BEAR", order[1]: "SIDEWAYS", order[2]: "BULL"}
        
        return True
    
    def predict_regime(self, df_daily):
        """预测当前状态"""
        if not self.fitted:
            # 规则兜底
            return self._rule_based_regime(df_daily)
        
        features = self.prepare_features(df_daily)
        scaled = self.scaler.transform(features)
        hidden_states = self.hmm.predict(scaled)
        
        regimes = [self.state_map.get(s, "UNKNOWN") for s in hidden_states]
        return regimes
    
    def _rule_based_regime(self, df_daily):
        """规则兜底：MA200 + RSI + 成交量"""
        c = df_daily['close'].values
        v = df_daily['volume'].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_ratio = v / np.maximum(pd.Series(v).rolling(20).mean().values, 1)
        
        regimes = []
        for i in range(len(c)):
            if i < 200:
                regimes.append("UNKNOWN")
                continue
            
            # 牛市：价格>MA50>MA200，RSI>50
            if c[i] > ma50[i] > ma200[i] and rsi14[i] > 50:
                regimes.append("BULL")
            # 熊市：价格<MA50<MA200，RSI<50
            elif c[i] < ma50[i] < ma200[i] and rsi14[i] < 50:
                regimes.append("BEAR")
            # 震荡：其他
            else:
                regimes.append("SIDEWAYS")
        
        return regimes
    
    def get_current_regime(self, df_daily):
        """获取最新的状态"""
        regimes = self.predict_regime(df_daily)
        return regimes[-1] if regimes else "UNKNOWN"


# =====================================================================
# 4. 权重分配器
# =====================================================================
class WeightAllocator:
    """根据状态动态分配策略权重"""
    
    # 状态 × 策略 适应性矩阵
    # 行：策略， 列：状态
    ADAPTABILITY = {
        # BEAR, SIDEWAYS, BULL
        "TrendFollowingDaily":    [0.1, 0.3, 0.9],   # 趋势策略：牛市强
        "MeanReversion15m":       [0.3, 0.9, 0.2],   # 均值回归：震荡市强
        "Breakout15m":            [0.2, 0.4, 0.8],   # 突破：牛市强
        "MomentumReversal15m":    [0.6, 0.5, 0.3],   # 反转：熊市强
    }
    
    STATE_ORDER = ["BEAR", "SIDEWAYS", "BULL"]
    
    def __init__(self, temp=1.0):
        self.temp = temp  # 温度参数，控制权重集中度
    
    def get_weights(self, regime):
        """获取当前状态下的策略权重"""
        if regime not in self.STATE_ORDER:
            # 未知状态：均匀分布
            return {s: 0.25 for s in self.ADAPTABILITY.keys()}
        
        state_idx = self.STATE_ORDER.index(regime)
        
        # 获取每个策略在这个状态下的适应性
        adapt = {}
        for strat, values in self.ADAPTABILITY.items():
            adapt[strat] = values[state_idx]
        
        # Softmax归一化
        weights = self._softmax(adapt)
        return weights
    
    def _softmax(self, d):
        vals = np.array(list(d.values()))
        keys = list(d.keys())
        exp_vals = np.exp(vals / self.temp)
        probs = exp_vals / exp_vals.sum()
        return dict(zip(keys, probs))


# =====================================================================
# 5. 组合执行器
# =====================================================================
class EnsemblePortfolio:
    """组合执行器：管理多策略组合"""
    
    def __init__(self, strategies, regime_detector, allocator, capital=1000):
        self.strategies = {s.name: s for s in strategies}
        self.regime_detector = regime_detector
        self.allocator = allocator
        self.capital = capital
        self.positions = {s.name: 0 for s in strategies}
        self.history = []
    
    def run_backtest(self, df_daily, df_15m):
        """完整回测"""
        # 1. 生成每个策略的信号
        print("生成策略信号...")
        signals = {}
        for name, strat in self.strategies.items():
            if strat.timeframe == '1d':
                signals[name] = strat.generate_signals(df_daily)
            else:
                signals[name] = strat.generate_signals(df_15m)
            print(f"  {name}: {np.sum(signals[name] != 0)}个信号")
        
        # 2. 识别每日状态
        print("识别市场状态...")
        regimes = self.regime_detector.predict_regime(df_daily)
        
        # 3. 每日重新平衡
        print("运行组合回测...")
        results = self._run_portfolio(df_daily, df_15m, signals, regimes)
        
        return results
    
    def _run_portfolio(self, df_daily, df_15m, signals, regimes):
        """核心回测循环"""
        # 简化版：按日线步进，每日根据权重分配
        N = len(df_daily)
        c_daily = df_daily['close'].values
        dates = df_daily.index
        
        portfolio_value = 1000
        equity_curve = [1000]
        daily_returns = []
        positions = {name: 0 for name in self.strategies.keys()}
        cash = 1000
        
        for i in range(200, N - 1):
            # 当前状态
            regime = regimes[i] if i < len(regimes) else "SIDEWAYS"
            
            # 获取权重
            weights = self.allocator.get_weights(regime)
            
            # 计算每个策略的目标仓位
            for name, weight in weights.items():
                if i < len(signals[name]):
                    signal = signals[name][i]
                    # 简化：信号 * 权重 * 杠杆
                    target_pos = signal * weight * 1.0  # 1x杠杆
                    positions[name] = target_pos
            
            # 计算当日收益
            if i < len(c_daily) - 1:
                ret = (c_daily[i+1] / c_daily[i] - 1)
                daily_pnl = sum(positions[name] * ret for name in positions)
                daily_pnl -= 0.0003 * sum(abs(positions[name]) for name in positions)  # 手续费
                
                portfolio_value *= (1 + daily_pnl)
                equity_curve.append(portfolio_value)
                daily_returns.append(daily_pnl * 100)
        
        # 统计
        eq = np.array(equity_curve)
        ret_series = np.array(daily_returns) / 100
        
        peak = np.maximum.accumulate(eq)
        dd = (eq - peak) / peak * 100
        max_dd = np.min(dd)
        
        total_ret = (eq[-1] / eq[0] - 1) * 100
        sharpe = np.mean(ret_series) / np.std(ret_series) * np.sqrt(252) if np.std(ret_series) > 0 else 0
        win_rate = np.mean(np.array(daily_returns) > 0) * 100
        
        return {
            "equity_curve": eq,
            "total_ret": total_ret,
            "sharpe": round(sharpe, 2),
            "max_dd": round(max_dd, 1),
            "win_rate": round(win_rate, 1),
            "daily_returns": daily_returns,
        }


# =====================================================================
# 6. 主流程
# =====================================================================
print("\n" + "="*80)
print("🚀 初始化 Ensemble + Regime 框架")
print("="*80)

# 初始化策略
strategies = [
    TrendFollowingDaily(),
    MeanReversion15m(),
    Breakout15m(),
    MomentumReversal15m(),
]

# 初始化状态识别器
regime_detector = RegimeDetector(n_states=3)

# 先在日线数据上拟合HMM
print("训练HMM状态识别器...")
regime_detector.fit(df_daily)

# 预测所有状态
regimes = regime_detector.predict_regime(df_daily)
print(f"状态分布: {pd.Series(regimes).value_counts().to_dict()}")

# 初始化权重分配器
allocator = WeightAllocator(temp=1.0)

# 初始化组合
portfolio = EnsemblePortfolio(strategies, regime_detector, allocator)

# 运行回测
print("\n" + "="*80)
print("🚀 运行 Ensemble 回测")
print("="*80)

results = portfolio.run_backtest(df_daily, df_15m)

# 输出结果
print(f"\n{'='*80}")
print(f"📊 Ensemble 回测结果")
print(f"{'='*80}")
print(f"总收益: {results['total_ret']:+.1f}%")
print(f"Sharpe: {results['sharpe']:.2f}")
print(f"最大回撤: {results['max_dd']:.1f}%")
print(f"胜率: {results['win_rate']:.1f}%")
print(f"最终权益: {results['equity_curve'][-1]:.0f}")

# 单独测试每个策略
print(f"\n{'='*80}")
print(f"📊 单策略对比 (独立回测)")
print(f"{'='*80}")

for strat in strategies:
    if strat.timeframe == '1d':
        sig = strat.generate_signals(df_daily)
        # 简单统计
        active = np.sum(sig != 0)
        print(f"  {strat.name}: {active}个信号 ({active/len(df_daily)*100:.1f}%时间在仓)")
    else:
        sig = strat.generate_signals(df_15m)
        active = np.sum(sig != 0)
        print(f"  {strat.name}: {active}个信号 ({active/len(df_15m)*100:.1f}%时间在仓)")

# 状态分析
print(f"\n{'='*80}")
print(f"📊 状态分析")
print(f"{'='*80}")
regime_series = pd.Series(regimes)
print(f"状态分布: {regime_series.value_counts().to_dict()}")

# 最后可用的状态
current = regime_detector.get_current_regime(df_daily)
weights = allocator.get_weights(current)
print(f"\n当前状态: {current}")
print(f"建议权重:")
for s, w in sorted(weights.items(), key=lambda x: x[1], reverse=True):
    print(f"  {s}: {w*100:.0f}%")

print(f"\n{'='*80}")
print("✅ 框架搭建完成！可以在此基础上继续优化：")
print("  1. 调整 ADAPTABILITY 矩阵")
print("  2. 优化 HMM 特征工程")
print("  3. 加入更多策略 (套利、资金费率等)")
print("  3. 增加风控模块 (VaR, 相关性控制)")
print("="*80)