"""烛龙 v2.1：扫荡前置策略
1. 先扫荡（流动性被干掉）
2. 再出现反转形态（机构进场）
3. 确认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]
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

# ── 扫荡定义 ──
sweep_down = (L < prev_l20) & (C > prev_l20)  # 跌破前低→收回
sweep_up = (H > prev_h20) & (C < prev_h20)  # 突破前高→收回

# ── 反转形态定义 ──
bull_engulf = (C1<O1) & (C>O) & (O<=C1) & (C>=O1)
piercing = (C1<O1) & (C>O) & (O<L1) & (C>(O1+C1)/2) & (C<O1)
hammer = (C>O) & ((O-L)>=abs(C-O)*2) & ((H-C)<=abs(C-O)*0.3) & (abs(C-O)/(H-L+1e-9)<0.4)
morning_star = (C2<O2) & (abs(C2-O2)/(H2-L2+1e-9)<0.3) & (C>O) & (C>(O2+C2)/2)

bear_engulf = (C1>O1) & (C<O) & (O>=C1) & (C<=O1)
dark_cloud = (C1>O1) & (C<O) & (O>H1) & (C<(O1+C1)/2) & (C>O1)
shooting = (C<O) & ((H-C)>=abs(C-O)*2) & ((O-L)<=abs(C-O)*0.3) & (abs(C-O)/(H-L+1e-9)<0.4)
evening_star = (C2>O2) & (abs(C2-O2)/(H2-L2+1e-9)<0.3) & (C<O) & (C<(O2+C2)/2)

# ── 确认K线 ──
nbull = np.roll(C,-1) > np.roll(O,-1)
nbear = np.roll(C,-1) < np.roll(O,-1)

# ── 成交量 ──
vol_ma20 = pd.Series(V).rolling(20).mean().values

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

# ── 扫荡前置策略 ──
# 搜索窗口：扫荡后N根K线内出现反转形态
SEARCH_WINDOW = 5

signals = []
for i in range(10, n-5):
    # 检查前面N根有没有扫荡
    sweep_occurred = False
    sweep_dir = 0
    sweep_idx = 0
    
    for lookback in range(1, SEARCH_WINDOW+1):
        if i-lookback < 0: continue
        if sweep_down[i-lookback]:
            sweep_occurred = True
            sweep_dir = 1  # 扫荡向下→做多
            sweep_idx = i-lookback
            break
        if sweep_up[i-lookback]:
            sweep_occurred = True
            sweep_dir = -1  # 扫荡向上→做空
            sweep_idx = i-lookback
            break
    
    if not sweep_occurred: continue
    
    # 检查当前K线是否有同向反转形态
    has_pattern = False
    if sweep_dir == 1:  # 做多
        # 需要日线趋势向上
        if not trend_1h[i]: continue
        has_pattern = bull_engulf[i] or piercing[i] or hammer[i] or morning_star[i]
        if not has_pattern: continue
        # 确认K线
        if not nbull[i]: continue
        
        # 加分：放量、实体占比
        score_vol = 15 if V[i] > vol_ma20[i]*1.5 else (8 if V[i] > vol_ma20[i]*1.2 else 0)
        body_r = abs(C[i]-O[i]) / (H[i]-L[i]+1e-9)
        score_body = 10 if body_r > 0.7 else (5 if body_r > 0.5 else 0)
        
        signals.append({
            "i": i, "dir": 1, "entry": C[i],
            "type": "sweep2long", "score": score_vol+score_body,
            "sweep_dist": i - sweep_idx  # 扫荡到形态的距离（根数）
        })
    
    else:  # 做空
        if trend_1h[i]: continue  # 需要日线趋势向下
        has_pattern = bear_engulf[i] or dark_cloud[i] or shooting[i] or evening_star[i]
        if not has_pattern: continue
        if not nbear[i]: continue
        
        score_vol = 15 if V[i] > vol_ma20[i]*1.5 else (8 if V[i] > vol_ma20[i]*1.2 else 0)
        body_r = abs(C[i]-O[i]) / (H[i]-L[i]+1e-9)
        score_body = 10 if body_r > 0.7 else (5 if body_r > 0.5 else 0)
        
        signals.append({
            "i": i, "dir": -1, "entry": C[i],
            "type": "sweep2short", "score": score_vol+score_body,
            "sweep_dist": i - sweep_idx
        })

print(f"扫荡前置策略信号: {len(signals)}笔")
longs = sum(1 for s in signals if s["dir"]==1)
shorts = len(signals)-longs
print(f"做多: {longs}  做空: {shorts}")

# ── 原版烛龙信号（对比用）──
near_s = abs(L - prev_l20) / (prev_l20 + 1e-9) < 0.008
near_r = abs(H - prev_h20) / (prev_h20 + 1e-9) < 0.008
long_orig = bull_engulf & near_s & nbull & trend_1h
short_orig = bear_engulf & near_r & nbear & (~trend_1h)
print(f"原版烛龙信号: {long_orig.sum() + short_orig.sum()}笔 (做多{long_orig.sum()} 做空{short_orig.sum()})")

# ── 回测函数 ──
def backtest(sigs, mode="fixed"):
    results = {"in":[],"out":[]}
    stats={"add":0}
    
    for s in sigs:
        i=s["i"]; entry=s["entry"]; dirc=s["dir"]
        if i+MB>=n: continue
        period="in" if i<split else "out"
        score=s.get("score",0)
        pos_size=1.0; add_count=0; last_eval=0
        closed=False
        
        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 score>=20: bm=4;sz=0.5
                        elif score>=10: 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 backtest_orig(long_mask, short_mask, mode="fixed"):
    results={"in":[],"out":[]}; stats={"add":0}
    for period,st,en in [("in",0,split),("out",split,n)]:
        for mask,dirc in [(long_mask,1),(short_mask,-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
                for j in range(1,MB+1):
                    if i+j>=n: break
                    ret=(C[i+j]/entry-1)*dirc
                    if mode!="fixed" and j%3==0 and j!=last_eval:
                        last_eval=j
                        if add_count<3 and pos_size<2.0:
                            pos_size=min(2.0,pos_size+0.3); 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.1：扫荡前置策略 vs 原版")
print("="*70)

# 原版
r0,s0=backtest_orig(long_orig,short_orig,"fixed")
ps(r0,s0,"原版固定止损")

r1,s1=backtest_orig(long_orig,short_orig,"addonly")
ps(r1,s1,"原版+只加不减")

# 扫荡前置
r2,s2=backtest(signals,"fixed")
ps(r2,s2,"扫荡前置固定止损")

r3,s3=backtest(signals,"addonly")
ps(r3,s3,"扫荡前置+只加不减")

print(f"\n{'='*70}")
print("📊 样本外对比：")
for name,res,st in [("原版固定",r0,{}),("原版加仓",r1,s1),("扫荡固定",r2,{}),("扫荡加仓",r3,s3)]:
    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}")