"""烛龙 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)
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
near_s = abs(L - prev_l20) / (prev_l20 + 1e-9) < 0.008
near_r = abs(H - prev_h20) / (prev_h20 + 1e-9) < 0.008

bull_engulf = (C1<O1) & (C>O) & (O<=C1) & (C>=O1)
nbull = np.roll(C,-1) > np.roll(O,-1)
long_raw = bull_engulf & near_s & nbull

swup = (H>pd.Series(H).shift(1).rolling(20).max().values) & (C<pd.Series(H).shift(1).rolling(20).max().values)
nbear = np.roll(C,-1) < np.roll(O,-1)
short_raw = swup & near_r & nbear

lsig = long_raw & trend_1h
ssig = short_raw & (~trend_1h)

vol_mean20 = pd.Series(V).rolling(20).mean().values

# ── 动量指标 ──
ret_2 = C/np.roll(C,2)-1  # 2小时收益
ret_3 = C/np.roll(C,3)-1

# 动量加速：连续3根K线涨幅递增
mom_accel = np.full(n, False)
for i in range(4, n):
    r1 = (C[i-1]/C[i-2]-1)
    r2 = (C[i]/C[i-1]-1)
    if C[i-1] > O[i-1] and C[i] > O[i]:
        if r2 > r1 and r2 > 0.001:
            mom_accel[i] = True

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

def resonance(i, dirc, use_momentum=False):
    s = 0
    if V[i] > vol_mean20[i] * 1.5: s += 15
    elif V[i] > vol_mean20[i] * 1.2: s += 8
    
    if use_momentum:
        if dirc == 1 and mom_accel[i]:
            s += 25  # 动量加速→高确信度
        elif dirc == -1:
            # 空头动量：连续下跌
            chk = i
            if (C[chk-1] < O[chk-1] and C[chk] < O[chk-1]):
                ret3 = abs(C[chk]/C[chk-3]-1)
                if ret3 > 0.01:
                    s += 25
    
    return s

def backtest(lsig, ssig, mode, use_momentum=False):
    results = {"in": [], "out": []}
    stats = {"add": 0}
    
    for period, st, en in [("in",0,split), ("out",split,n)]:
        for mask, dirc in [(lsig,1), (ssig,-1)]:
            for i in range(st, min(en,n)):
                if not mask[i]: continue
                if i+MB >= n: continue
                entry=C[i]; closed=False; pos_size=1.0
                add_count=0; last_eval=0
                res=resonance(i,dirc,use_momentum)
                
                for j in range(1,MB+1):
                    if i+j>=n: break
                    ret=(C[i+j]/entry-1)*dirc
                    
                    if mode != "fixed":
                        if j%3==0 and j!=last_eval:
                            last_eval=j
                            if mode=="addonly":
                                if add_count<3 and pos_size<2.0:
                                    pos_size=min(2.0,pos_size+0.3)
                                    add_count+=1;stats["add"]+=1
                            else:
                                cr=ret
                                if res>=30: bm=4;sz=0.5
                                elif res>=15: bm=3;sz=0.35
                                else: bm=2;sz=0.25
                                if cr>0.005: mxa=min(4,bm+1)
                                elif cr>0: mxa=bm
                                elif cr>-0.005: mxa=max(0,bm-1)
                                else: mxa=0;sz=0
                                if add_count<mxa and pos_size<2.0:
                                    pos_size=min(2.0,pos_size+sz)
                                    add_count+=1;stats["add"]+=1
                    
                    if ret>=TP:
                        results[period].append(TP*pos_size-0.001);closed=True;break
                    if ret<=-SL:
                        results[period].append(-SL*pos_size-0.001);closed=True;break
                
                if not closed:
                    fr=(C[min(i+MB,n-1)]/entry-1)*dirc
                    results[period].append(fr*pos_size-0.001)
    
    return results, stats

def ps(results,stats,label):
    print(f"\n▶ {label}")
    if stats.get("add"): print(f"  [统计] 加仓={stats['add']}")
    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)

r0,_=backtest(lsig,ssig,"fixed")
ps(r0,{},"固定止损（原版）")

r1,s1=backtest(lsig,ssig,"addonly")
ps(r1,s1,"v1.2 只加不减")

r7,s7=backtest(lsig,ssig,"full")
ps(r7,s7,"v1.7（无动量）")

r8,s8=backtest(lsig,ssig,"full",use_momentum=True)
ps(r8,s8,"v2.0 动量共震版")

print(f"\n{'='*70}")
print("📊 样本外对比：")
for name,res,st in [("固定止损",r0,{}),("v1.2",r1,s1),("v1.7",r7,s7),("v2.0动量",r8,s8)]:
    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}")