"""烛龙 v2.2：双策略融合
策略A（原版）：吞没@支撑+确认 → 贡献笔数
策略B（扫荡前置）：扫荡+形态+确认 → 贡献质量
两个独立运行，互不冲突"""
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

# ── 形态 ──
bull_engulf = (C1<O1) & (C>O) & (O<=C1) & (C>=O1)
piercing = (C1<O1) & (C>O) & (O<L1) & (C>(O1+C1)/2) & (C<O1)
bear_engulf = (C1>O1) & (C<O) & (O>=C1) & (C<=O1)
dark_cloud = (C1>O1) & (C<O) & (O>H1) & (C<(O1+C1)/2) & (C>O1)

nbull = np.roll(C,-1) > np.roll(O,-1)
nbear = np.roll(C,-1) < np.roll(O,-1)

# ── 扫荡 ──
sweep_down = (L < prev_l20) & (C > prev_l20)
sweep_up = (H > prev_h20) & (C < prev_h20)

# ── 成交量 ──
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

# ══════════════════════════════════════════
# 策略A：原版烛龙（吞没+支撑+确认）
# ══════════════════════════════════════════
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_a = bull_engulf & near_s & nbull & trend_1h
short_a = bear_engulf & near_r & nbear & (~trend_1h)

# ══════════════════════════════════════════
# 策略B：扫荡前置（扫荡+N根内出现反转+确认）
# ══════════════════════════════════════════
SEARCH_WINDOW = 5

signals_b = []  # 记录每笔，便于去重
for i in range(10, n-5):
    # 检查前面N根有没有扫荡
    sweep_found = False; sweep_dir = 0
    for lb in range(1, SEARCH_WINDOW+1):
        if i-lb < 0: continue
        if sweep_down[i-lb]: sweep_found=True; sweep_dir=1; break
        if sweep_up[i-lb]: sweep_found=True; sweep_dir=-1; break
    if not sweep_found: continue
    
    if sweep_dir == 1:  # 做多
        if not trend_1h[i]: continue
        has_good = bull_engulf[i] or piercing[i]
        if not has_good or not nbull[i]: continue
        signals_b.append(i)
    
    else:  # 做空
        if trend_1h[i]: continue
        has_good = bear_engulf[i] or dark_cloud[i]
        if not has_good or not nbear[i]: continue
        signals_b.append(i)

# 去重
signals_b_unique = sorted(set(signals_b))

# 检查AB重叠
overlap = 0
for i in range(n):
    if (long_a[i] or short_a[i]) and (i in signals_b_unique):
        overlap += 1
print(f"策略A信号: {(long_a|short_a).sum()}笔")
print(f"策略B信号: {len(signals_b_unique)}笔")
print(f"重叠信号: {overlap}笔")

# ══════════════════════════════════════════
# 回测：单独跑 + 融合跑
# ══════════════════════════════════════════
def backtest_signal_list(sig_indices, dirc_array):
    """给一个信号index列表和一个方向数组跑回测"""
    results={"in":[],"out":[]}
    for idx in sig_indices:
        i = idx
        dirc = dirc_array[i] if i < len(dirc_array) else 1
        if i+MB>=n: continue
        period="in" if i<split else "out"
        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 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
            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

def backtest_mask(mask, dirc):
    results={"in":[],"out":[]}
    for period,st,en in [("in",0,split),("out",split,n)]:
        for i in range(st,min(en,n)):
            if not mask[i]: continue
            if i+MB>=n: continue
            d=dirc[i]; 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)*d
                if 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
                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)*d
                results[period].append(fr*pos_size-0.001)
    return results

def ps(results,label):
    print(f"\n▶ {label}")
    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.2：双策略融合（原版+扫荡前置）")
print("="*70)
print(f"SL={SL*100:.1f}% TP={TP*100:.1f}% MB={MB}h")

# 策略A单独
dirc_a = np.where(long_a, 1, np.where(short_a, -1, 0))
rA = backtest_mask(long_a | short_a, dirc_a)
ps(rA, "策略A（原版烛龙）")

# 策略B单独
dirc_b = np.zeros(n, dtype=int)
for idx in signals_b_unique:
    # 确定方向
    for lb in range(1, SEARCH_WINDOW+1):
        if idx-lb < 0: continue
        if sweep_down[idx-lb]: dirc_b[idx]=1; break
        if sweep_up[idx-lb]: dirc_b[idx]=-1; break

mask_b = np.zeros(n, dtype=bool)
for idx in signals_b_unique:
    mask_b[idx] = True

rB = backtest_signal_list(signals_b_unique, dirc_b)
ps(rB, "策略B（扫荡前置）")

# 融合：取并集
combined_sigs = sorted(set(np.where(long_a|short_a)[0].tolist() + signals_b_unique))
dirc_c = np.zeros(n, dtype=int)
for idx in combined_sigs:
    if idx < n:
        if long_a[idx]: dirc_c[idx]=1
        elif short_a[idx]: dirc_c[idx]=-1
        elif dirc_b[idx]!=0: dirc_c[idx]=dirc_b[idx]

rC = backtest_signal_list(combined_sigs, dirc_c)
ps(rC, "融合版（A+B去重）")

# 统计每个策略在融合中的占比
in_a = sum(1 for i in combined_sigs if i < n and (long_a[i] or short_a[i]))
in_b = sum(1 for i in combined_sigs if i < n and mask_b[i])
only_a = sum(1 for i in combined_sigs if i < n and (long_a[i] or short_a[i]) and not mask_b[i])
only_b = sum(1 for i in combined_sigs if i < n and mask_b[i] and not (long_a[i] or short_a[i]))
both = sum(1 for i in combined_sigs if i < n and (long_a[i] or short_a[i]) and mask_b[i])
print(f"\n  融合构成: A独有={only_a}   B独有={only_b}   重叠={both}")

# 按信号来源拆分统计融合版的收益
# 找到样本外中哪些是A、哪些是B
out_sigs = [i for i in combined_sigs if i >= split and i+MB < n]
out_a = [i for i in out_sigs if i < n and (long_a[i] or short_a[i])]
out_b = [i for i in out_sigs if i < n and mask_b[i]]
out_only_a = [i for i in out_sigs if i < n and (long_a[i] or short_a[i]) and not mask_b[i]]
out_only_b = [i for i in out_sigs if i < n and mask_b[i] and not (long_a[i] or short_a[i])]
out_both = [i for i in out_sigs if i < n and (long_a[i] or short_a[i]) and mask_b[i]]
print(f"\n  样本外: A={len(out_a)} B={len(out_b)} A独有={len(out_only_a)} B独有={len(out_only_b)} 重叠={len(out_both)}")

# 对比融合vs单独
print(f"\n{'='*70}")
print("📊 样本外对比（统一加仓模式）")
for name,res in [("策略A原版",rA),("策略B扫荡",rB),("融合版",rC)]:
    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:16s}: {len(tr):>3d}笔 wr={wr:>5.1%} cum={cum:>7.3f} 周盈≈{pw:>5.1%} MDD={mdd:>6.1%} Sharpe={sp:>5.2f}")