"""
数据探索：15分钟BTC的预测性特征
不做任何预设，枚举所有特征，让数据说话
"""

import pandas as pd
import numpy as np
import warnings
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线  {df.index[0]} ~ {df.index[-1]}")
print(f"📊 约{len(df)//96}个交易日\n")

# ── 基础指标（所有特征基于当前K线已知信息）──
# 当前K线特征
body = c - o                      # 实体
body_pct = (c / o - 1) * 100      # 实体百分比
upper_wick = h - np.maximum(c, o) # 上影线
lower_wick = np.minimum(c, o) - l # 下影线
range_pct = (h / l - 1) * 100     # 振幅%
vol_ratio = np.ones(N)            # 初始化

# 前一根K线的特征
prev_close = np.roll(c, 1)
prev_open = np.roll(o, 1)
prev_high = np.roll(h, 1)
prev_low = np.roll(l, 1)
prev_body_pct = np.roll(body_pct, 1)
prev_vol = np.roll(v, 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 roll_max(arr, n): return pd.Series(arr).rolling(n).max().values
def roll_min(arr, n): return pd.Series(arr).rolling(n).min().values

# 均值
ma5 = roll_mean(c, 5)
ma20 = roll_mean(c, 20)
ma50 = roll_mean(c, 50)
ma100 = roll_mean(c, 100)
vol_ma5 = roll_mean(v, 5)
vol_ma20 = roll_mean(v, 20)
vol_ma50 = roll_mean(v, 50)

# 相对于均线的位置
dist_ma5 = (c / ma5 - 1) * 100
dist_ma20 = (c / ma20 - 1) * 100
dist_ma50 = (c / ma50 - 1) * 100

# 波动率
atr14 = roll_mean(np.maximum(h - l, np.maximum(np.abs(h - np.roll(c, 1)), np.abs(l - np.roll(c, 1)))), 14)
atr5 = roll_mean(np.maximum(h - l, np.maximum(np.abs(h - np.roll(c, 1)), np.abs(l - np.roll(c, 1)))), 5)

# 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_ratio_5 = v / np.maximum(vol_ma5, 1)
vol_ratio_20 = v / np.maximum(vol_ma20, 1)

# 价格动量
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

# 高低点突破
high_10 = roll_max(h, 10)
low_10 = roll_min(l, 10)
high_20 = roll_max(h, 20)
low_20 = roll_min(l, 20)

# 突破检测
break_high_10 = h > np.roll(high_10, 1)  # 突破前10根高点
break_low_10 = l < np.roll(low_10, 1)    # 跌破前10根低点

# 假突破检测（先突破后收回）
# 前一根突破了，这一根收回来了
fake_break_high = np.roll(break_high_10, 1) & (c < np.roll(high_10, 1))  # 前一根突破前高，这一根收回去
fake_break_low = np.roll(break_low_10, 1) & (c > np.roll(low_10, 1))     # 前一根跌破前低，这一根收回去

# 成交量放量突破
vol_surge = v > vol_ma20 * 2.0

# 时间特征
hour = df.index.hour.values
minute = df.index.minute.values

# 会话过滤
is_asia = ((0 <= hour) & (hour < 8)).astype(int)
is_europe = ((8 <= hour) & (hour < 16)).astype(int)
is_us = ((16 <= hour) | (hour < 0)).astype(int)  # UTC 16-24 = US

# ── 未来收益（预测目标）──
fwd_1 = (np.roll(c, -1) / c - 1) * 100   # 下一根K线收益
fwd_5 = (np.roll(c, -5) / c - 1) * 100   # 未来5根（75分钟）
fwd_10 = (np.roll(c, -10) / c - 1) * 100 # 未来10根（2.5小时）
fwd_20 = (np.roll(c, -20) / c - 1) * 100 # 未来20根（5小时）

fwd_high_5 = np.zeros(N)  # 未来5根最高价相对当前
fwd_low_5 = np.zeros(N)
for i in range(5, N - 5):
    fwd_high_5[i] = (np.max(h[i+1:i+6]) / c[i] - 1) * 100
    fwd_low_5[i] = (np.min(l[i+1:i+6]) / c[i] - 1) * 100

print("="*80)
print("📊 特征探索：什么样的条件下做多/做空期望值为正？")
print("="*80)

# ═══ 方法：对每个特征分段，计算未来收益 ═══

def analyze_feature(name, feature_values, bins=10, min_samples=30):
    """把一个特征分成bins段，算每段的未来收益"""
    valid = ~(np.isnan(feature_values) | np.isinf(feature_values))
    fv = feature_values[valid]
    
    if len(fv) < min_samples:
        return
    
    # 分bin
    percentiles = np.percentile(fv, np.linspace(0, 100, bins+1))
    
    print(f"\n  📌 {name} (共{len(fv):,}个样本)")
    
    for b in range(bins):
        lo, hi = percentiles[b], percentiles[b+1]
        if b == bins - 1:
            mask = (feature_values >= lo) & (feature_values <= hi) & valid
        else:
            mask = (feature_values >= lo) & (feature_values < hi) & valid
        
        count = mask.sum()
        if count < min_samples:
            continue
        
        # 未来收益
        mean_ret_5 = fwd_5[mask].mean()
        win_rate_5 = (fwd_5[mask] > 0).mean() * 100
        mean_ret_10 = fwd_10[mask].mean()
        
        # 是否显著
        t_stat = mean_ret_5 / (fwd_5[mask].std() / np.sqrt(count)) if fwd_5[mask].std() > 0 else 0
        
        range_str = f"[{lo:.2f}, {hi:.2f})" if b < bins - 1 else f"[{lo:.2f}, {hi:.2f}]"
        print(f"    {range_str:>16} | {count:>5}次 | 5K胜率{win_rate_5:>5.1f}% | 5K均{mean_ret_5:>+6.2f}% | 10K均{mean_ret_10:>+6.2f}% | t={t_stat:+.1f}")


def analyze_binary(name, condition, min_samples=20):
    """二元特征分析"""
    valid = ~np.isnan(condition)
    true_count = condition.sum()
    false_count = (~condition & valid).sum()
    
    if true_count < min_samples:
        print(f"\n  ⏭️ {name}: 仅{true_count}次正样本，跳过")
        return
    
    true_ret_5 = fwd_5[condition].mean()
    true_wr_5 = (fwd_5[condition] > 0).mean() * 100
    false_ret_5 = fwd_5[~condition & valid].mean()
    false_wr_5 = (fwd_5[~condition & valid] > 0).mean() * 100
    
    diff_ret = true_ret_5 - false_ret_5
    diff_wr = true_wr_5 - false_wr_5
    
    flag = "✅" if diff_ret > 0.02 else ("⚠️" if diff_ret > 0 else "❌")
    
    print(f"\n  {flag} {name}")
    print(f"     条件为真: {true_count:>5}次 | 5K均{true_ret_5:>+6.2f}% | 胜率{true_wr_5:>5.1f}%")
    print(f"     条件为假: {false_count:>5}次 | 5K均{false_ret_5:>+6.2f}% | 胜率{false_wr_5:>5.1f}%")
    print(f"     delta: 收益{diff_ret:+.2f}% | 胜率{diff_wr:+.1f}%")


# ── 1. K线形态特征 ──
print("\n" + "═"*80)
print("🔴 1. 单根K线特征")
print("═"*80)

# 实体大小
analyze_feature("实体百分比(body_pct)", body_pct, bins=10)

# 上下影线比例
wick_ratio = np.where(upper_wick + lower_wick > 0, 
                      (upper_wick - lower_wick) / (upper_wick + lower_wick + 1), 0)
analyze_feature("影线偏度(+上影长,-下影长)", wick_ratio, bins=10)

# 振幅
analyze_feature("振幅百分比(range_pct)", range_pct, bins=10)

# ── 2. 成交量特征 ──
print("\n" + "═"*80)
print("🔴 2. 成交量特征")
print("═"*80)

analyze_feature("成交量/20日均量(vol_ratio)", vol_ratio_20, bins=10)
analyze_feature("成交量/5日均量(vol_ratio_5)", vol_ratio_5, bins=10)

# 放量 + 方向
vol_up = (v > vol_ma20 * 1.5) & (body_pct > 0)
vol_down = (v > vol_ma20 * 1.5) & (body_pct < 0)
analyze_binary("放量上涨(vol>1.5x + 阳线)", vol_up)
analyze_binary("放量下跌(vol>1.5x + 阴线)", vol_down)

# 缩量
vol_quiet = v < vol_ma20 * 0.5
analyze_binary("极度缩量(vol<0.5x均量)", vol_quiet)

# ── 3. 支撑/阻力突破 ──
print("\n" + "═"*80)
print("🔴 3. 支撑/阻力突破")
print("═"*80)

analyze_binary("突破前10根高点(break_high_10)", break_high_10)
analyze_binary("跌破前10根低点(break_low_10)", break_low_10)

# ── 4. 假突破（Wyckoff Spring/Upthrust）──
print("\n" + "═"*80)
print("🔴 4. 假突破（流动性抓取）")
print("═"*80)

analyze_binary("假突破上方(fake_break_high)", fake_break_high)
analyze_binary("假跌破下方(fake_break_low - Spring)", fake_break_low)

# 区分放量的假突破
fake_high_vol = fake_break_high & (v > vol_ma20)
fake_low_vol = fake_break_low & (v > vol_ma20)
analyze_binary("放量假突破上方", fake_high_vol)
analyze_binary("放量假跌破下方", fake_low_vol)

# ── 5. RSI ──
print("\n" + "═"*80)
print("🔴 5. RSI")
print("═"*80)

analyze_feature("RSI14", rsi14, bins=10)

# RSI极端
rsi_oversold = rsi14 < 30
rsi_overbought = rsi14 > 70
analyze_binary("RSI<30（超卖）", rsi_oversold)
analyze_binary("RSI>70（超买）", rsi_overbought)

# RSI+成交量
rsi_oversold_vol = rsi_oversold & (v > vol_ma20)
analyze_binary("RSI超卖+放量", rsi_oversold_vol)

# ── 6. 相对于均线位置 ──
print("\n" + "═"*80)
print("🔴 6. 相对于均线位置")
print("═"*80)

analyze_feature("距离MA20%", dist_ma20, bins=10)
analyze_feature("距离MA50%", dist_ma50, bins=10)

near_ma20 = np.abs(dist_ma20) < 0.5
far_below = dist_ma20 < -3
far_above = dist_ma20 > 3
analyze_binary("在MA20附近(<0.5%)", near_ma20)
analyze_binary("远低于MA20(>3%)", far_below)
analyze_binary("远高于MA20(>3%)", far_above)

# ── 7. 价格动量 ──
print("\n" + "═"*80)
print("🔴 7. 价格动量")
print("═"*80)

analyze_feature("前3根K线收益%", ret_3, bins=10)
analyze_feature("前5根K线收益%", ret_5, bins=10)
analyze_feature("前10根K线收益%", ret_10, bins=10)

# 连续下跌/上涨
down_3 = (ret_3 < -1) & (ret_1 < 0)
down_5 = (ret_5 < -2) & (ret_3 < -1)
up_3 = (ret_3 > 1) & (ret_1 > 0)
analyze_binary("连续3根下跌(累计>1%)", down_3)
analyze_binary("连续5根下跌(累计>2%)", down_5)
analyze_binary("连续3根上涨(累计>1%)", up_3)

# ── 8. 组合特征 ──
print("\n" + "═"*80)
print("🔴 8. 组合特征（多条件同时满足）")
print("═"*80)

# Spring + RSI超卖
spring_oversold = fake_break_low & rsi_oversold
analyze_binary("Spring+RSI超卖", spring_oversold)

# 放量下跌后缩量企稳（卖力耗尽）
for i in range(5, N):
    climax_mask = np.zeros(N, dtype=bool)
    for look in range(5):
        idx = i - look
        if idx < 1: continue
        if c[idx] < c[idx-1] and v[idx] > vol_ma20[idx] * 2.0:
            # 之后3根缩量
            if i > idx and np.mean(v[idx+1:i+1]) < vol_ma20[i] * 0.8:
                climax_mask[i] = True

analyze_binary("卖力耗尽(大跌放量后缩量)", climax_mask)

# 回踩均线
pullback_to_ma20 = (ret_3 > 1) & (np.abs(dist_ma20) < 0.3) & (body_pct > 0)
analyze_binary("上涨后回踩MA20企稳", pullback_to_ma20)

# ── 9. 交易时段 ──
print("\n" + "═"*80)
print("🔴 9. 交易时段")
print("═"*80)

analyze_binary("亚洲盘(UTC 0-8)", is_asia.astype(bool))
analyze_binary("欧洲盘(UTC 8-16)", is_europe.astype(bool))
analyze_binary("美洲盘(UTC 16-24)", is_us.astype(bool))

# ── 10. 最佳信号汇总 ──
print("\n" + "═"*80)
print("🏆 所有信号汇总（按5K未来收益排序）")
print("═"*80)

# 收集所有二元特征对比
all_features = [
    ("放量假跌破(Spring+Vol)", fake_low_vol),
    ("放量假突破", fake_high_vol),
    ("Spring", fake_break_low),
    ("假突破上方", fake_break_high),
    ("RSI超卖+放量", rsi_oversold_vol),
    ("RSI<30超卖", rsi_oversold),
    ("RSI>70超买", rsi_overbought),
    ("放量上涨", vol_up),
    ("放量下跌", vol_down),
    ("极度缩量", vol_quiet),
    ("突破前高", break_high_10),
    ("跌破前低", break_low_10),
    ("远低于MA20", far_below),
    ("远高于MA20", far_above),
    ("在MA20附近", near_ma20),
    ("连跌3根", down_3),
    ("连跌5根", down_5),
    ("连涨3根", up_3),
    ("Spring+RSI超卖", spring_oversold),
    ("卖力耗尽", climax_mask),
    ("回踩MA20", pullback_to_ma20),
]

results_list = []
for feat_name, feat_mask in all_features:
    if feat_mask.sum() < 20:
        continue
    tr = fwd_5[feat_mask].mean()
    wr = (fwd_5[feat_mask] > 0).mean() * 100
    results_list.append((tr, wr, feat_mask.sum(), feat_name))

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

print(f"\n{'排名':>4} {'信号名称':<24} {'出现次数':>8} {'5K均收益':>10} {'5K胜率':>8}")
print("-" * 60)
for i, (tr, wr, cnt, name) in enumerate(results_list[:15]):
    print(f"{i+1:>4} {name:<24} {cnt:>8} {tr:>+9.3f}% {wr:>7.1f}%")

# 再看做空（负收益越大越好）
print(f"\n{'排名':>4} {'做空信号':<24} {'出现次数':>8} {'5K均收益':>10} {'5K胜率':>8}")
print("-" * 60)
results_short = [(r, w, c, n) for r, w, c, n in results_list if r < 0]
results_short.sort(key=lambda x: x[0])  # 最小的（最负的）排前面
for i, (tr, wr, cnt, name) in enumerate(results_short[:10]):
    print(f"{i+1:>4} {name:<24} {cnt:>8} {tr:>+9.3f}% {wr:>7.1f}%")

# ── 找出信号发生后的最佳持仓时间 ──
print(f"\n" + "═"*80)
print("⏱ 最佳持仓时间分析：Spring信号后的收益随时间变化")
print("═"*80)

# 取最好的信号，看不同持仓时间
best_mask = fake_break_low  # Spring
if best_mask.sum() > 20:
    print(f"\n  Spring信号共{best_mask.sum()}次，出场时间分析:")
    print(f"  {'出场(K线数)':<15} {'平均收益':>10} {'胜率':>8} {'最大收益':>10} {'最小收益':>10}")
    print("  " + "-" * 55)
    for fwd_bars in [1, 3, 5, 10, 20, 40, 80]:
        fwd_ret = (np.roll(c, -fwd_bars) / c - 1)[best_mask] * 100
        wr = (fwd_ret > 0).mean() * 100
        print(f"  {fwd_bars:>4}根({fwd_bars*0.25:>4.1f}h): {fwd_ret.mean():>+9.2f}% {wr:>7.1f}% {fwd_ret.max():>+9.2f}% {fwd_ret.min():>+9.2f}%")
