"""隐秘枢轴（Rumers魔力线）独立回测 v2
修复日期映射问题
"""
import pandas as pd, numpy as np

df_raw = pd.read_parquet("data/btc_multidim.parquet")
df_raw.columns = [c.lower() for c in df_raw.columns]
d = df_raw.copy()

O, H, L, C, V = d["open"].values, d["high"].values, d["low"].values, d["close"].values, d["volume"].values
n = len(d)

# ── 日级关键位 ──
daily = d.resample("1D").agg({"high":"max","low":"min"}).dropna()
daily["prev_high"] = daily["high"].shift(1)
daily["prev_low"] = daily["low"].shift(1)
daily["swing_high"] = daily["high"].rolling(10).max().shift(1)
daily["swing_low"] = daily["low"].rolling(10).min().shift(1)

# 每日往前填充到15分钟级别
d["prev_high"] = daily["prev_high"].reindex(d.index, method="ffill").values
d["prev_low"] = daily["prev_low"].reindex(d.index, method="ffill").values
d["swing_high"] = daily["swing_high"].reindex(d.index, method="ffill").values
d["swing_low"] = daily["swing_low"].reindex(d.index, method="ffill").values

# ── 参数 ──
FEE = 0.001
MB = 96  # 最大96根15分K = 24小时

# 把每天分为若干"交易时段"（适配24小时交易的BTC）
# 策略原版看开盘前3根K线，我们改成每6小时检查一次"前3根模式"
HOURS = [0, 6, 12, 18]  # UTC时间，每6小时检查

signals = []

for hour in HOURS:
    # 找到每天该小时开始的索引
    day_groups = d.groupby(d.index.date)
    
    for day, group in day_groups:
        day_start = group.index[0]
        # 目标时段开始
        target = pd.Timestamp(day_start.date()) + pd.Timedelta(hours=hour)
        
        # 确保数据足够
        if target not in d.index: continue
        
        target_idx = d.index.get_loc(target)
        if target_idx + 3 >= n: continue
        
        # 前3根K线
        bar1_idx = target_idx
        bar2_idx = target_idx + 1
        bar3_idx = target_idx + 2
        
        bar1_c = C[bar1_idx]; bar1_o = O[bar1_idx]; bar1_h = H[bar1_idx]; bar1_l = L[bar1_idx]
        bar2_c = C[bar2_idx]; bar2_o = O[bar2_idx]; bar2_h = H[bar2_idx]; bar2_l = L[bar2_idx]
        bar3_c = C[bar3_idx]; bar3_o = O[bar3_idx]; bar3_h = H[bar3_idx]; bar3_l = L[bar3_idx]
        
        prev_high = d["prev_high"].iloc[bar1_idx]
        prev_low = d["prev_low"].iloc[bar1_idx]
        swing_high = d["swing_high"].iloc[bar1_idx]
        swing_low = d["swing_low"].iloc[bar1_idx]
        
        if pd.isna(prev_high) or pd.isna(swing_high):
            continue
        
        # ── 做多检查 ──
        price_near_sup = (abs(bar1_l - prev_low) / prev_low < 0.005 or
                          abs(bar1_l - swing_low) / swing_low < 0.005 or
                          abs(bar1_c - prev_low) / prev_low < 0.005)
        
        if price_near_sup and bar3_c > bar2_h:
            sup = prev_low if abs(bar1_l - prev_low)/prev_low < 0.005 else swing_low
            res = max(prev_high, swing_high)
            tp_pct = min(res / bar3_c - 1, 0.03) if res > bar3_c else 0.02
            
            signals.append({
                "i": target_idx, "dir": 1, "entry": bar3_c,
                "tp_pct": tp_pct, "sl_pct": 0.008,
                "sup": sup, "res": res
            })
            continue  # 同组不做两个方向
        
        # ── 做空检查 ──
        price_near_res = (abs(bar1_h - prev_high) / prev_high < 0.005 or
                          abs(bar1_h - swing_high) / swing_high < 0.005 or
                          abs(bar1_c - prev_high) / prev_high < 0.005)
        
        if price_near_res and bar3_c < bar2_l:
            res_line = prev_high if abs(bar1_h - prev_high)/prev_high < 0.005 else swing_high
            sup = min(prev_low, swing_low)
            tp_pct = min(1 - bar3_c / sup, 0.03) if sup < bar3_c else 0.02
            
            signals.append({
                "i": target_idx, "dir": -1, "entry": bar3_c,
                "tp_pct": tp_pct, "sl_pct": 0.008,
                "sup": sup, "res": res_line
            })

print(f"信号总数: {len(signals)}")
if signals:
    longs = sum(1 for s in signals if s["dir"] == 1)
    shorts = len(signals) - longs
    print(f"做多: {longs}  做空: {shorts}")

# ── 回测 ──
results = {"in": [], "out": []}
split_pos = int(n * 0.67)

for sig in signals:
    i = sig["i"]
    if i + MB >= n: continue
    period = "in" if i < split_pos else "out"
    entry = sig["entry"]
    dirc = sig["dir"]
    sl_pct = sig["sl_pct"]
    tp_pct = sig["tp_pct"]
    closed = False
    
    for j in range(1, MB + 1):
        if i + j >= n: break
        ret = (C[i+j] / entry - 1) * dirc
        
        if ret >= tp_pct:
            results[period].append(tp_pct - FEE)
            closed = True; break
        if ret <= -sl_pct:
            results[period].append(-sl_pct - FEE)
            closed = True; break
    if not closed:
        final_ret = (C[min(i+MB, n-1)] / entry - 1) * dirc
        results[period].append(final_ret - FEE)

print("\n" + "=" * 60)
print("隐秘枢轴（Rumers魔力线）独立回测")
print("=" * 60)
print(f"周期: 15分钟  检查频率: 每6小时  SL=0.8%")
print(f"数据: BTC-USDT {len(d)}根K线 ({d.index[0].date()} ~ {d.index[-1].date()})")

for nm, key in [("样本内", "in"), ("样本外", "out")]:
    tr = results[key]
    if len(tr) < 3:
        print(f"\n{nm}: 仅{len(tr)}笔, 不统计"); continue
    wr = sum(1 for r in tr if r > 0) / len(tr)
    cum = np.prod([1 + r for r in tr])
    avg = np.mean(tr) * 100
    running = 1.0; peak = 1.0; mdd = 0
    for r in tr:
        running *= (1 + r); peak = max(peak, running)
        mdd = min(mdd, (running - peak) / peak)
    sharpe = np.mean(tr) / (np.std(tr) + 1e-9) * np.sqrt(365*96/MB) if np.std(tr) > 0 else 0
    print(f"\n{nm}: {len(tr)}笔 wr={wr:.1%} cum={cum:.3f} avg={avg:+.2f}% "
          f"MDD={mdd:.1%} Sharpe={sharpe:.2f}")
    print(f"  盈利: {sum(1 for r in tr if r>0)}笔  亏损: {sum(1 for r in tr if r<0)}笔")

print(f"\n{'='*60}")
print("总结：")