"""
多时间框架特征探索
枚举 OHLCV 能算的所有特征，看哪个时间框架上有预测力
"""

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

def calc_features(o, h, l, c, v, timeframe_name):
    """计算一个时间框架的所有特征和未来收益"""
    N = len(c)
    if N < 200:
        return None
    
    # ── 特征工程 ──
    body_pct = (c / o - 1) * 100
    range_pct = (h / l - 1) * 100
    upper_wick = (h - np.maximum(c, o)) / (h - l + 1e-9) * 100
    lower_wick = (np.minimum(c, o) - l) / (h - l + 1e-9) * 100
    ret_1 = (c / np.roll(c, 1) - 1) * 100
    ret_3 = (c / np.roll(c, 3) - 1) * 100
    ret_5 = (c / np.roll(c, 5) - 1) * 100
    ret_10 = (c / np.roll(c, 10) - 1) * 100
    ret_20 = (c / np.roll(c, 20) - 1) * 100
    
    # 均线
    ma5 = pd.Series(c).rolling(5).mean().values
    ma10 = pd.Series(c).rolling(10).mean().values
    ma20 = pd.Series(c).rolling(20).mean().values
    ma50 = pd.Series(c).rolling(50).mean().values
    ma100 = pd.Series(c).rolling(100).mean().values
    ma200 = pd.Series(c).rolling(200).mean().values
    
    dist_ma10 = (c / ma10 - 1) * 100
    dist_ma20 = (c / ma20 - 1) * 100
    dist_ma50 = (c / ma50 - 1) * 100
    dist_ma100 = (c / ma100 - 1) * 100
    
    # 波动率
    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
    
    # RSI
    delta = pd.Series(c).diff()
    gain = delta.clip(lower=0).rolling(14).mean()
    loss = (-delta.clip(upper=0)).rolling(14).mean()
    rsi14 = (100 - 100 / (1 + gain / (loss + 1e-9))).values
    
    # 成交量
    vol_ma5 = pd.Series(v).rolling(5).mean().values
    vol_ma10 = pd.Series(v).rolling(10).mean().values
    vol_ma20 = pd.Series(v).rolling(20).mean().values
    vol_ratio_5 = v / np.maximum(vol_ma5, 0.01)
    vol_ratio_10 = v / np.maximum(vol_ma10, 0.01)
    vol_ratio_20 = v / np.maximum(vol_ma20, 0.01)
    
    # 高低点突破（前N根）
    high_10 = pd.Series(h).rolling(10).max().values
    low_10 = pd.Series(l).rolling(10).min().values
    high_20 = pd.Series(h).rolling(20).max().values
    low_20 = pd.Series(l).rolling(20).min().values
    high_50 = pd.Series(h).rolling(50).max().values
    low_50 = pd.Series(l).rolling(50).min().values
    
    break_high_10 = np.roll(high_10, 1) > 0  # 避免除零
    break_low_10 = np.roll(low_10, 1) > 0
    for i in range(1, N):
        break_high_10[i] = h[i] > high_10[i-1] * 1.0001 and h[i] > c[i-1]
        break_low_10[i] = l[i] < low_10[i-1] * 0.9999 and l[i] < c[i-1]
    
    # 假突破
    fake_break_high = np.zeros(N, dtype=bool)
    fake_break_low = np.zeros(N, dtype=bool)
    for i in range(1, N):
        if break_high_10[i-1] and c[i] < high_10[i-1]:
            fake_break_high[i] = True
        if break_low_10[i-1] and c[i] > low_10[i-1]:
            fake_break_low[i] = True
    
    # 连续涨跌
    down_3 = (ret_3 < -1) & (ret_1 < 0) & (ret_3 > -20)
    down_5 = (ret_5 < -2) & (ret_3 < -1) & (ret_5 > -30)
    up_3 = (ret_3 > 1) & (ret_1 > 0) & (ret_3 < 20)
    
    # 远低于/高于均线
    far_below_ma50 = dist_ma50 < -5
    far_above_ma50 = dist_ma50 > 5
    near_ma20 = np.abs(dist_ma20) < 0.5
    
    # RSI极端
    rsi_oversold = rsi14 < 30
    rsi_overbought = rsi14 > 70
    rsi_mid = (rsi14 > 40) & (rsi14 < 60)
    
    # 成交量极端
    vol_surge = vol_ratio_20 > 2.0
    vol_quiet = vol_ratio_20 < 0.5
    
    # 未来收益（适应不同时间框架的步长）
    fwd_1 = (np.roll(c, -1) / c - 1) * 100
    fwd_3 = (np.roll(c, -3) / c - 1) * 100
    fwd_5 = (np.roll(c, -5) / c - 1) * 100
    fwd_10 = (np.roll(c, -10) / c - 1) * 100
    fwd_20 = (np.roll(c, -20) / c - 1) * 100
    
    # ── 特征字典 ──
    features = {}
    
    # 单根K线形态
    features['小实体(|body|<0.1%)'] = np.abs(body_pct) < 0.1
    features['大阳线(body>1%)'] = body_pct > 1
    features['大阴线(body<-1%)'] = body_pct < -1
    features['下影线长'] = lower_wick > 70
    features['上影线长'] = upper_wick > 70
    
    # 成交量
    features['放量>2x'] = vol_surge
    features['缩量<0.5x'] = vol_quiet
    features['放量上涨'] = vol_surge & (body_pct > 0)
    features['放量下跌'] = vol_surge & (body_pct < 0)
    features['缩量企稳'] = vol_quiet & (np.abs(body_pct) < 0.2)
    
    # 突破
    features['突破前10高'] = break_high_10
    features['跌破前10低'] = break_low_10
    
    # 假突破
    features['假突破上方'] = fake_break_high
    features['假跌破下方(Spring)'] = fake_break_low
    features['放量假突破'] = fake_break_high & vol_surge
    features['放量Spring'] = fake_break_low & vol_surge
    
    # RSI
    features['RSI<30超卖'] = rsi_oversold
    features['RSI>70超买'] = rsi_overbought
    features['RSI 40-60'] = rsi_mid
    features['RSI超卖+放量'] = rsi_oversold & vol_surge
    
    # 均线
    features['远低于MA50(>5%)'] = far_below_ma50
    features['远高于MA50(>5%)'] = far_above_ma50
    features['在MA20附近'] = near_ma20
    
    # 动量
    features['连跌3根'] = down_3
    features['连跌5根'] = down_5
    features['连涨3根'] = up_3
    
    # 组合
    features['Spring+RSI超卖'] = fake_break_low & rsi_oversold
    features['远低MA50+RSI超卖'] = far_below_ma50 & rsi_oversold
    features['远高MA50+RSI超买'] = far_above_ma50 & rsi_overbought
    features['放量Spring+超卖'] = fake_break_low & vol_surge & rsi_oversold
    
    return {
        'features': features,
        'fwd_5': fwd_5,      # 未来5根K线收益
        'fwd_10': fwd_10,    # 未来10根
        'body_pct': body_pct,
        'rsi14': rsi14,
        'dist_ma20': dist_ma20,
        'dist_ma50': dist_ma50,
        'ret_5': ret_5,
        'ret_10': ret_10,
        'vol_ratio_20': vol_ratio_20,
        'range_pct': range_pct,
        'N': N,
    }


def analyze_timeframe(name, o, h, l, c, v, fwd_bars=5):
    """分析一个时间框架"""
    data = calc_features(o, h, l, c, v, name)
    if data is None:
        return None
    
    features = data['features']
    N = data['N']
    fwd = data['fwd_5'] if fwd_bars == 5 else data['fwd_10']
    
    print(f"\n{'='*70}")
    print(f"📊 时间框架: {name}  ({N:,}根K线)")
    print(f"{'='*70}")
    
    results = []
    
    for feat_name, feat_mask in features.items():
        true_count = feat_mask.sum()
        if true_count < 10:
            continue
        
        true_ret = fwd[feat_mask].mean()
        true_wr = (fwd[feat_mask] > 0).mean() * 100
        
        # 相对于全样本的收益和胜率（作为baseline）
        all_ret = np.nanmean(fwd)
        all_wr = np.nanmean((fwd > 0).astype(float)) * 100
        
        diff_ret = true_ret - all_ret
        diff_wr = true_wr - all_wr
        
        # 统计显著性（简化的z-score）
        if true_count > 1 and np.nanstd(fwd[feat_mask]) > 0:
            z = true_ret / (np.nanstd(fwd[feat_mask]) / np.sqrt(true_count))
        else:
            z = 0
        
        results.append((diff_ret, diff_wr, true_count, true_ret, true_wr, z, feat_name))
    
    # 排序
    results.sort(key=lambda x: x[0], reverse=True)
    
    # 打印
    print(f"{'特征':<28} {'次数':>6} {'收益':>8} {'胜率':>6} {'vs基准收益':>10} {'vs基准胜率':>10} {'z':>6}")
    print("-" * 78)
    
    for diff_ret, diff_wr, cnt, tr, twr, z, fn in results[:15]:
        flag = "🟢" if diff_ret > 0.03 and z > 1 else ("🔴" if diff_ret < -0.03 else "⚪")
        print(f"{flag} {fn:<26} {cnt:>6} {tr:>+7.2f}% {twr:>5.1f}% {diff_ret:>+8.3f}% {diff_wr:>+8.1f}% {z:>+5.1f}")
    
    # 做空信号（负收益显著的）
    print(f"\n  {'做空信号':<28} {'次数':>6} {'收益':>8} {'胜率':>6} {'vs基准':>10} {'z':>10}")
    print("  " + "-" * 70)
    shorts = [r for r in results if r[0] < -0.02]
    shorts.sort(key=lambda x: x[0])
    for diff_ret, diff_wr, cnt, tr, twr, z, fn in shorts[:10]:
        print(f"  {fn:<26} {cnt:>6} {tr:>+7.2f}% {twr:>5.1f}% {diff_ret:>+8.3f}% {z:>+6.1f}")
    
    return results


def continuous_feature(name, values, fwd, bins=10, feat_label="值"):
    """分析连续特征（如RSI、dist_ma20等）"""
    valid = ~(np.isnan(values) | np.isinf(values))
    fv = values[valid]
    fwd_v = fwd[valid]
    
    if len(fv) < 50:
        return
    
    print(f"\n  📈 {name}:")
    percentiles = np.percentile(fv[~np.isnan(fv)], np.linspace(0, 100, bins+1))
    
    for b in range(bins):
        lo, hi = percentiles[b], percentiles[b+1]
        mask = (fv >= lo) & (fv <= hi) & ~np.isnan(fv)
        cnt = mask.sum()
        if cnt < 10:
            continue
        
        mean_ret = np.nanmean(fwd_v[mask])
        wr = np.nanmean((fwd_v[mask] > 0).astype(float)) * 100
        
        lo_s = f"{lo:.2f}" if abs(lo) < 100 else f"{lo:.1f}"
        hi_s = f"{hi:.2f}" if abs(hi) < 100 else f"{hi:.1f}"
        
        flag = "✅" if mean_ret > 0.02 else ("❌" if mean_ret < -0.02 else "·")
        print(f"    [{lo_s:>6}, {hi_s:>6})  {cnt:>5}次  {flag} 均{mean_ret:>+6.3f}% 胜率{wr:>5.1f}%")


def get_ohlcv(tf_df):
    """统一获取OHLCV列，兼容大小写"""
    col_map = {}
    for col in ['open', 'high', 'low', 'close', 'volume']:
        if col in tf_df.columns:
            col_map[col] = col
        elif col[0] in tf_df.columns:
            col_map[col] = col[0]
        else:
            raise KeyError(f"找不到列 {col} 或 {col[0]}")
    o = tf_df[col_map['open']].values.astype(float)
    h = tf_df[col_map['high']].values.astype(float)
    l = tf_df[col_map['low']].values.astype(float)
    c = tf_df[col_map['close']].values.astype(float)
    v = tf_df[col_map['volume']].values.astype(float)
    return o, h, l, c, v

# ═══════════════════════════════════════════
# 加载数据
# ═══════════════════════════════════════════

# 15m数据（原始）
d15 = pd.read_parquet("/root/quant_pipeline/data/btc_15m_binance.parquet")
d15['time'] = pd.to_datetime(d15['ts']).values
print(f"15m原始: {len(d15):,}行  {d15.iloc[0]['time']} ~ {d15.iloc[-1]['time']}")

# 重采样到1H和4H
df_15m = d15.copy()
df_15m['ts'] = pd.to_datetime(df_15m['ts'])
df_15m = df_15m.set_index('ts').sort_index()

ohlc_dict = {'open': 'first', 'high': 'max', 'low': 'min', 'close': 'last', 'volume': 'sum'}

df_1h = df_15m.resample('1h').agg(ohlc_dict).dropna()
df_4h = df_15m.resample('4h').agg(ohlc_dict).dropna()
df_6h = df_15m.resample('6h').agg(ohlc_dict).dropna()
df_12h = df_15m.resample('12h').agg(ohlc_dict).dropna()

print(f"1H:  {len(df_1h):,}行  {df_1h.index[0]} ~ {df_1h.index[-1]}")
print(f"4H:  {len(df_4h):,}行  {df_4h.index[0]} ~ {df_4h.index[-1]}")
print(f"6H:  {len(df_6h):,}行  {df_6h.index[0]} ~ {df_6h.index[-1]}")
print(f"12H: {len(df_12h):,}行  {df_12h.index[0]} ~ {df_12h.index[-1]}")

# 日线和周线（2017-2024）
dd = pd.read_parquet("/root/quant_pipeline/data/btc_daily.parquet")
dw = pd.read_parquet("/root/quant_pipeline/data/btc_weekly.parquet")
print(f"日线: {len(dd):,}行  {dd.index[0].date()} ~ {dd.index[-1].date()}")
print(f"周线: {len(dw):,}行  {dw.index[0].date()} ~ {dw.index[-1].date()}")


# ═══════════════════════════════════════════
# 主分析
# ═══════════════════════════════════════════

timeframes = [
    ("1H", df_1h),
    ("4H", df_4h),
    ("6H", df_6h),
    ("12H", df_12h),
    ("日线", dd),
    ("周线", dw),
]

all_tf_data = {}

# 先填充all_tf_data
for tf_name, tf_df in timeframes:
    o, h, l, c, v = get_ohlcv(tf_df)
    data = calc_features(o, h, l, c, v, tf_name)
    if data:
        all_tf_data[tf_name] = data


# ── 每个时间框架的离散特征分析 ──
print("\n" + "═" * 110)
print("PART 1: 离散特征分析 — 每个时间框架上最有预测力的条件")
print("═" * 110)

for tf_name, tf_df in timeframes:
    o, h, l, c, v = get_ohlcv(tf_df)
    analyze_timeframe(tf_name, o, h, l, c, v, fwd_bars=5)


# ── 连续特征对比：RSI × 时间框架 ──
print("\n" + "═" * 110)
print("PART 2: RSI 预测力 按时间框架对比")
print("═" * 110)

for tf_name, data in all_tf_data.items():
    print(f"\n📊 {tf_name}  RSI14 vs 未来5K收益:")
    continuous_feature("RSI14", data['rsi14'], data['fwd_5'], bins=10)


# ── 连续特征对比：距离MA20 ──
print("\n" + "═" * 110)
print("PART 3: 距MA20位置 预测力 按时间框架对比")
print("═" * 110)

for tf_name, data in all_tf_data.items():
    print(f"\n📊 {tf_name}  距离MA20% vs 未来5K收益:")
    continuous_feature("距MA20%", data['dist_ma20'], data['fwd_5'], bins=10)


# ── 连续特征对比：成交量 ──
print("\n" + "═" * 110)
print("PART 4: 成交量比 预测力 按时间框架对比")
print("═" * 110)

for tf_name, data in all_tf_data.items():
    print(f"\n📊 {tf_name}  Vol/MA20 vs 未来5K收益:")
    continuous_feature("Vol/MA20", data['vol_ratio_20'], data['fwd_5'], bins=10)


# ── 连续特征对比：前5根收益（均值回归测试）──
print("\n" + "═" * 110)
print("PART 5: 动量均值回归测试（前5K收益 vs 未来5K收益）")
print("═" * 110)

for tf_name, data in all_tf_data.items():
    print(f"\n📊 {tf_name}  前5K收益% vs 未来5K收益%:")
    continuous_feature("前5K收益%", data['ret_5'], data['fwd_5'], bins=10)


# ── 所有时间框架汇总对比 ──
print("\n" + "═" * 110)
print("🏆 所有TF × 所有特征 汇总对比（按收益提升排序）")
print("═" * 110)

all_results = []
for tf_name, tf_df in timeframes:
    o, h, l, c, v = get_ohlcv(tf_df)
    
    data = calc_features(o, h, l, c, v, tf_name)
    if data is None:
        continue
    
    fwd = data['fwd_5']
    all_ret = np.nanmean(fwd)
    
    for feat_name, feat_mask in data['features'].items():
        cnt = feat_mask.sum()
        if cnt < 10:
            continue
        
        tr = np.nanmean(fwd[feat_mask])
        twr = np.nanmean((fwd[feat_mask] > 0).astype(float)) * 100
        diff = tr - all_ret
        
        if cnt > 1 and np.nanstd(fwd[feat_mask]) > 0:
            z = tr / (np.nanstd(fwd[feat_mask]) / np.sqrt(cnt))
        else:
            z = 0
        
        all_results.append((diff, abs(z), cnt, tr, twr, z, tf_name, feat_name))

all_results.sort(key=lambda x: x[0], reverse=True)

print(f"\n{'TF':<6} {'特征':<28} {'次数':>6} {'均收益':>8} {'胜率':>6} {'提升':>8} {'|z|':>6}")
print("-" * 75)

# 做多信号 TOP20
print(f"\n🔴 做多信号 TOP 20:")
for diff, absz, cnt, tr, twr, z, tf, fn in all_results[:20]:
    if diff > 0:
        flag = "🟢" if z > 2 else "🟡"
        print(f"{flag} {tf:<5} {fn:<26} {cnt:>6} {tr:>+7.2f}% {twr:>5.1f}% {diff:>+7.3f}% {z:>+5.1f}")

# 做空信号 TOP20
print(f"\n🔵 做空信号 TOP 20:")
for diff, absz, cnt, tr, twr, z, tf, fn in reversed(all_results[-20:]):
    if diff < 0:
        flag = "🔵" if absz > 2 else "🔷"
        print(f"{flag} {tf:<5} {fn:<26} {cnt:>6} {tr:>+7.2f}% {twr:>5.1f}% {diff:>+7.3f}% {z:>+5.1f}")

# ── 核心结论 ──
print("\n" + "═" * 110)
print("💡 跨时间框架核心结论")
print("═" * 110)

# 计算每个时间框架的平均预测力
tf_power = {}
for diff, absz, cnt, tr, twr, z, tf, fn in all_results:
    if tf not in tf_power:
        tf_power[tf] = []
    tf_power[tf].append(abs(diff))

print(f"\n各时间框架平均特征提升（越大越好）:")
for tf, vals in sorted(tf_power.items(), key=lambda x: np.mean(x[1]), reverse=True):
    print(f"  {tf:>4}: 平均提升 {np.mean(vals)*100:.2f}bp")
