"""烛龙 v2.0：从“赌反转”转到“赌延续”
右侧交易：趋势突破回踩确认 + 动量延续"""
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.resample("1h").agg({"open":"first","high":"max","low":"min","close":"last","volume":"sum"}).dropna()

O,H,L,C,V = d["open"].values,d["high"].values,d["low"].values,d["close"].values,d["volume"].values
O1,H1,L1,C1 = np.roll(O,1),np.roll(H,1),np.roll(L,1),np.roll(C,1)
O2,H2,L2,C2 = np.roll(O,2),np.roll(H,2),np.roll(L,2),np.roll(C,2)
n = len(d)

# ── 日线趋势 ──
dd = df_raw.resample("1D").agg({"close":"last"}).dropna()
dd["ma20"] = dd["close"].rolling(20).mean()
dd["trend_up"] = dd["close"] > dd["ma20"]
d_ts = d.index
trend_1h = np.array([dd["trend_up"].reindex([ts], method="ffill").values[0] if ts >= dd.index[0] else True for ts in d_ts])

# ── 关键位 ──
prev_l20 = pd.Series(L).shift(1).rolling(20).min().values
prev_h20 = pd.Series(H).shift(1).rolling(20).max().values

# ── 动量指标 ──
ret_1 = C/C1 - 1  # 1小时收益
ret_3 = C/np.roll(C,3) - 1  # 3小时收益
ret_6 = C/np.roll(C,6) - 1  # 6小时收益

# 波动率
atr_14 = pd.Series(np.maximum(H-L, np.maximum(abs(H-C1), abs(L-C1)))).rolling(14).mean().values

split = int(n * 0.67)
SL, TP, MB = 0.015, 0.045, 48
FEE = 0.001

# ══════════════════════════════════════════
# 策略A：突破回踩确认（右侧交易核心）
# 前提：价格突破前20根高点→回踩不破前高→继续涨→做多
# ══════════════════════════════════════════
def strategy_breakout_retest():
    """突破回踩确认：价格突破前高→回踩到前高附近→不破→继续涨"""
    signals = []
    
    for i in range(30, n-5):
        if not trend_1h[i]: continue  # 只在上升趋势中做多
        
        # 突破阶段：价格冲破前20根高点
        breakout = H[i-1] > prev_h20[i-2]  # 前一根K线突破了前高
        if not breakout: continue
        
        # 回踩阶段：价格回到了前高附近（但在前高之上）
        retest_low = L[i] >= prev_h20[i-2] * 0.998  # 回踩没跌破前高
        retest_close = C[i] > prev_h20[i-2]  # 收盘在前高之上
        
        if not (retest_low and retest_close): continue
        
        # 确认阶段：下一根K线继续涨
        confirm_up = C[i+1] > H[i] or C[i+1] > C[i] * 1.002
        if not confirm_up: continue
        
        signals.append({
            "i": i+1, "dirc": 1, "entry": C[i+1],
            "type": "突破回踩", "price_at_signal": C[i],
            "breakout_level": prev_h20[i-2]
        })
    
    return signals

# ══════════════════════════════════════════
# 策略B：动量延续（强势K线后继续追）
# 前提：连续3根阳线 + 每根涨幅递增 → 动量在加速→做多
# ══════════════════════════════════════════
def strategy_momentum():
    """动量延续：连续强势K线后，动量仍在加速"""
    signals = []
    
    for i in range(5, n-3):
        if not trend_1h[i]: continue
        
        # 连续3根阳线且每根涨幅递增（动量加速）
        bull1 = C[i-2] > O[i-2] and (C[i-2]/C[i-3]-1) > 0.001
        bull2 = C[i-1] > O[i-1] and (C[i-1]/C[i-2]-1) > (C[i-2]/C[i-3]-1)
        bull3 = C[i] > O[i] and (C[i]/C[i-1]-1) > (C[i-1]/C[i-2]-1)
        
        if not (bull1 and bull2 and bull3): continue
        
        # 放量确认
        vol_ok = V[i] > pd.Series(V).rolling(20).mean().values[i] * 1.2
        
        signals.append({
            "i": i, "dirc": 1, "entry": C[i],
            "type": "动量加速", "vol_ok": vol_ok
        })
    
    return signals

# ══════════════════════════════════════════
# 策略C：趋势回调结束（我们的原版+确认）
# 对比基准
# ══════════════════════════════════════════

# ══════════════════════════════════════════
# 测试函数
# ══════════════════════════════════════════
def backtest_signals(signals):
    results = {"in": [], "out": []}
    stats = {"total": 0, "breakout": 0, "momentum": 0}
    
    for sig in signals:
        i = sig["i"]
        if i + MB >= n: continue
        period = "in" if i < split else "out"
        entry = sig["entry"]
        dirc = sig["dirc"]
        
        # 统计类型
        if sig["type"] == "突破回踩": stats["breakout"] += 1
        elif sig["type"] == "动量加速": stats["momentum"] += 1
        stats["total"] += 1
        
        closed = False
        for j in range(1, MB+1):
            if i+j >= n: break
            ret = (C[i+j]/entry-1)*dirc
            if ret >= TP:
                results[period].append(TP-FEE); closed=True; break
            if ret <= -SL:
                results[period].append(-SL-FEE); closed=True; break
        if not closed:
            fr = (C[min(i+MB,n-1)]/entry-1)*dirc
            results[period].append(fr-FEE)
    
    return results, stats

def ps(results, stats, label):
    print(f"\n▶ {label}")
    if stats.get("total"):
        print(f"  [统计] 共{stats['total']}笔 突破回踩={stats.get('breakout',0)} 动量={stats.get('momentum',0)}")
    for nm,key in [("样本内","in"),("样本外","out")]:
        tr=results[key]
        if len(tr)<5: print(f"  {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)
        sp=np.mean(tr)/(np.std(tr)+1e-9)*np.sqrt(365*24/MB) if np.std(tr)>0 else 0
        print(f"  {nm}: {len(tr)}笔 wr={wr:.1%} cum={cum:.3f} avg={avg:+.2f}% 周盈≈{pw:.1%} MDD={mdd:.1%} Sharpe={sp:.2f}")

print("="*70)
print("烛龙 v2.0：右侧交易——赌延续，不赌反转")
print("="*70)
print(f"SL={SL*100:.1f}% TP={TP*100:.1f}% MB={MB}h")

sigs_a = strategy_breakout_retest()
r_a, s_a = backtest_signals(sigs_a)
ps(r_a, s_a, "策略A：突破回踩确认")
if sigs_a:
    print(f"  突破回踩信号示例:")
    for s in sigs_a[:3]:
        t = d.index[s['i']]
        print(f"    {t}: {s['type']} @ {s['entry']:.1f}")

sigs_b = strategy_momentum()
r_b, s_b = backtest_signals(sigs_b)
ps(r_b, s_b, "策略B：动量延续")
if sigs_b:
    print(f"  动量信号示例:")
    for s in sigs_b[:3]:
        t = d.index[s['i']]
        print(f"    {t}: {s['type']} @ {s['entry']:.1f}")

# 单独看突破回踩的成功率
print(f"\n{'='*70}")
print("📊 样本外对比：")
for name, res, st in [("突破回踩",r_a,s_a), ("动量延续",r_b,s_b)]:
    tr=res["out"]
    if len(tr)<5: 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
    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)
    sp=np.mean(tr)/(np.std(tr)+1e-9)*np.sqrt(365*24/MB) if np.std(tr)>0 else 0
    print(f"  {name:12s}: {len(tr):>3d}笔 wr={wr:>5.1%} cum={cum:>7.3f} 周盈≈{pw:>5.1%} MDD={mdd:>6.1%} Sharpe={sp:>5.2f}")

# 再跑一个原版 v1.7 做对比
print(f"\n{'='*70}")
print("📊 与烛龙原版对比：")