"""隐秘枢轴（Rumers魔力线）独立回测
15分钟K线 + 四线关键区 + 三K确认
"""
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]

# 15分钟天然就是我们要的时间周期
d = df_raw.copy()

# BTC 24小时交易，用UTC日期作为自然日边界
d["date"] = d.index.date

O, H, L, C, V = d["open"].values, d["high"].values, d["low"].values, d["close"].values, d["volume"].values
idx = d.index  # DatetimeIndex

n = len(d)

# ── 逐日计算关键位 ──
# 前一日高低点
daily = d.resample("1D").agg({"high":"max","low":"min","close":"last"}).dropna()
daily["prev_high"] = daily["high"].shift(1)
daily["prev_low"] = daily["low"].shift(1)
# 摆动高低点（前10日）
daily["swing_high"] = daily["high"].rolling(10).max().shift(1)
daily["swing_low"] = daily["low"].rolling(10).min().shift(1)

# 把日级数据映射回15分钟
d["day_prev_high"] = d["date"].map(lambda x: daily.loc[pd.Timestamp(x).date(), "prev_high"] if pd.Timestamp(x).date() in daily.index else np.nan)
d["day_prev_low"] = d["date"].map(lambda x: daily.loc[pd.Timestamp(x).date(), "prev_low"] if pd.Timestamp(x).date() in daily.index else np.nan)
d["day_swing_high"] = d["date"].map(lambda x: daily.loc[pd.Timestamp(x).date(), "swing_high"] if pd.Timestamp(x).date() in daily.index else np.nan)
d["day_swing_low"] = d["date"].map(lambda x: daily.loc[pd.Timestamp(x).date(), "swing_low"] if pd.Timestamp(x).date() in daily.index else np.nan)

# ── 策略参数 ──
SL_PCT = 0.008  # 止损0.8%（关键位外侧一点）
TP_PCT = 0.02   # 止盈2%（对面关键位的距离）
MB = 48         # 最大持仓48根15分K = 12小时

# ── 检查每日第1-3根K线位置 ──
# 每天从UTC 00:00开始的前3根K线（即00:00, 00:15, 00:30）
# 检查价格是否接近关键位

# 接近关键位的阈值
SUPPORT_ZONE_PCT = 0.003  # 接近0.3%以内算"触达"
RESISTANCE_ZONE_PCT = 0.003

# 信号记录
signals = []

for i in range(n):
    # 跳过无关键位的日子
    if pd.isna(d["day_prev_high"].iloc[i]) or pd.isna(d["day_swing_high"].iloc[i]):
        continue
    
    # 当前时间
    t = idx[i]
    # 这天的午夜
    day_start = pd.Timestamp(t.date())
    
    # 计算这是今天的第几根K线（从午夜开始算索引）
    k_idx = int((t - day_start).total_seconds() / 900)  # 900秒 = 15分钟
    
    # 只关心每天前3根K线
    if k_idx > 2:
        continue
    
    # 如果这是第3根K线（k_idx=2），检查信号
    if k_idx == 2:
        # 回顾之前的K线
        bar1_idx = i - 2
        bar2_idx = i - 1
        bar3_idx = i
        
        bar1 = {"o": O[bar1_idx], "h": H[bar1_idx], "l": L[bar1_idx], "c": C[bar1_idx]}
        bar2 = {"o": O[bar2_idx], "h": H[bar2_idx], "l": L[bar2_idx], "c": C[bar2_idx]}
        bar3 = {"o": O[bar3_idx], "h": H[bar3_idx], "l": L[bar3_idx], "c": C[bar3_idx]}
        
        prev_high = d["day_prev_high"].iloc[i]
        prev_low = d["day_prev_low"].iloc[i]
        swing_high = d["day_swing_high"].iloc[i]
        swing_low = d["day_swing_low"].iloc[i]
        
        # ── 做多检查：价格在支撑区（前日低点或摆动低点附近）──
        price_near_prevlow = abs(bar1.l - prev_low) / prev_low < 0.005 or abs(bar1.c - prev_low) / prev_low < 0.005
        price_near_swinglow = abs(bar1.l - swing_low) / swing_low < 0.005 or abs(bar1.c - swing_low) / swing_low < 0.005
        
        if price_near_prevlow or price_near_swinglow:
            # 三K确认：第3根K线突破第2根K线的高点 → 确认做多
            if bar3.c > bar2.h:
                # 用最接近的关键位作为入场逻辑支撑
                sup = prev_low if price_near_prevlow else swing_low
                res = prev_high  # 目标：对面阻力位
                
                signals.append({
                    "i": i, "dir": 1, "sup": sup, "res": res,
                    "entry": bar3.c, "support_type": "prev_low" if price_near_prevlow else "swing_low",
                    "SL_pct": SL_PCT, "TP_pct": (res/bar3.c - 1) if res > bar3.c else TP_PCT
                })
        
        # ── 做空检查：价格在阻力区（前日高点或摆动高点附近）──
        price_near_prevhigh = abs(bar1.h - prev_high) / prev_high < 0.005 or abs(bar1.c - prev_high) / prev_high < 0.005
        price_near_swinghigh = abs(bar1.h - swing_high) / swing_high < 0.005 or abs(bar1.c - swing_high) / swing_high < 0.005
        
        if price_near_prevhigh or price_near_swinghigh:
            if bar3.c < bar2.l:
                sup = prev_low
                res = prev_high if price_near_prevhigh else swing_high
                
                signals.append({
                    "i": i, "dir": -1, "sup": sup, "res": res,
                    "entry": bar3.c, "resistance_type": "prev_high" if price_near_prevhigh else "swing_high",
                    "SL_pct": SL_PCT, "TP_pct": (1 - bar3.c/sup) if sup < bar3.c else TP_PCT
                })

print(f"信号总数: {len(signals)}")

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

for sig in signals:
    i = sig["i"]
    if i + MB >= n: continue
    
    entry = sig["entry"]
    dirc = sig["dir"]
    
    period = "in" if i < split else "out"
    
    # 动态TP/SL
    sl_pct = SL_PCT
    tp_pct = sig.get("TP_pct", TP_PCT)
    tp_pct = min(tp_pct, 0.03)  # 不超过3%
    
    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"SL={SL_PCT*100:.1f}%  TP固定2%  最大持仓{MB*15//60}h (15分钟K线)")

for nm, key in [("样本内", "in"), ("样本外", "out")]:
    tr = results[key]
    if len(tr) < 5:
        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])
    ch = [tr[i:i+5] for i in range(0, len(tr), 5)]
    pw = sum(1 for c in ch if sum(c) > 0) / len(ch) if ch else 0
    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*24*4/MB) if np.std(tr) > 0 else 0
    print(f"\n{nm}: {len(tr)}笔 wr={wr:.1%} cum={cum:.3f} avg={avg:+.2f}% "
          f"周盈≈{pw:.1%} 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)}")

# 按方向分析
for nm, key in [("样本内", "in"), ("样本外", "out")]:
    long_tr = []
    short_tr = []
    for sig in signals:
        i = sig["i"]
        if i + MB >= n: continue
        period = "in" if i < split else "out"
        if period != key: continue
        
        entry = sig["entry"]
        dirc = sig["dir"]
        sl_pct = SL_PCT
        tp_pct = min(sig.get("TP_pct", TP_PCT), 0.03)
        
        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:
                (long_tr if dirc==1 else short_tr).append(tp_pct - fee)
                closed = True; break
            if ret <= -sl_pct:
                (long_tr if dirc==1 else short_tr).append(-sl_pct - fee)
                closed = True; break
        if not closed:
            final_ret = (C[min(i+MB, n-1)] / entry - 1) * dirc
            (long_tr if dirc==1 else short_tr).append(final_ret - fee)
    
    print(f"\n{nm} 细分:")
    for label, tr in [("做多", long_tr), ("做空", short_tr)]:
        if len(tr) >= 3:
            wr = sum(1 for r in tr if r>0)/len(tr)
            avg = np.mean(tr)*100
            cum = np.prod([1+r for r in tr])
            print(f"  {label}: {len(tr)}笔 wr={wr:.1%} avg={avg:+.2f}% cum={cum:.3f}")
        else:
            print(f"  {label}: {len(tr)}笔 数据不足")