"""烛龙 v1.9：多时间框架共震（1H信号 + 4H趋势 + 15min确认）"""
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]

# ── 1H基础数据 ──
d1 = df_raw.resample("1h").agg({"open":"first","high":"max","low":"min","close":"last","volume":"sum"}).dropna()

# ── 4H数据 ──
d4 = df_raw.resample("4h").agg({"open":"first","high":"max","low":"min","close":"last","volume":"sum"}).dropna()
d4["ma20"] = d4["close"].rolling(20).mean()
d4["trend_up"] = d4["close"] > d4["ma20"]

# ── 15min原始数据 ──
d15 = df_raw.copy()
# 15min关键位
prev_l20_15 = d15["low"].rolling(20).min().shift(1)
prev_h20_15 = d15["high"].rolling(20).max().shift(1)

O1,H1,L1,C1 = d15["open"].shift(1), d15["high"].shift(1), d15["low"].shift(1), d15["close"].shift(1)
O2,H2,L2,C2 = d15["open"].shift(2), d15["high"].shift(2), d15["low"].shift(2), d15["close"].shift(2)

# 15min反转形态
body15 = (d15["close"] - d15["open"]).abs()
rk15 = d15["high"] - d15["low"]
body_r15 = body15 / (rk15 + 1e-9)

bull15 = (C1 < O1) & (d15["close"] > d15["open"]) & (d15["open"] <= C1) & (d15["close"] >= O1) & \
         ((d15["low"] - prev_l20_15).abs() / (prev_l20_15 + 1e-9) < 0.008)
bear15 = (C1 > O1) & (d15["close"] < d15["open"]) & (d15["open"] >= C1) & (d15["close"] <= O1) & \
         ((d15["high"] - prev_h20_15).abs() / (prev_h20_15 + 1e-9) < 0.008)

# 15min确认（下一根阳线/阴线）
nbull15 = d15["close"].shift(-1) > d15["open"].shift(-1)
nbear15 = d15["close"].shift(-1) < d15["open"].shift(-1)
sig15_long = bull15 & nbull15
sig15_short = bear15 & nbear15

# ── 回到1H级别 ──
O,H,L,C,V = d1["open"].values,d1["high"].values,d1["low"].values,d1["close"].values,d1["volume"].values
O1,H1,L1,C1 = np.roll(O,1),np.roll(H,1),np.roll(L,1),np.roll(C,1)
n = len(d1)

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

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 = d1.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])

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

vol_mean20 = pd.Series(V).rolling(20).mean().values
split = int(n * 0.67)
SL, TP, MB = 0.015, 0.045, 48

# ── 4H趋势映射回1H ──
def get_4h_trend(ts):
    """获取某个时间点的4H趋势"""
    aligned = d4.index[d4.index <= ts]
    if len(aligned) == 0: return True
    nearest = aligned[-1]
    return d4.loc[nearest, "trend_up"]

# ── 15min信号映射回1H ──
def has_15min_signal(ts, dirc, lookback=2):
    """过去2根15min内是否有同向信号"""
    window = sig15_long if dirc == 1 else sig15_short
    mask = window.loc[ts - pd.Timedelta(minutes=30):ts]
    return mask.any()

# ── 共振打分 ──
def resonance_v7(i, dirc):
    """v1.7 原始打分"""
    s = 0
    if V[i] > vol_mean20[i] * 1.5: s += 15
    elif V[i] > vol_mean20[i] * 1.2: s += 8
    return s

def resonance_v9(i, dirc):
    """v1.9 多时间框架共震"""
    s = resonance_v7(i, dirc)
    ts = d1.index[i]
    
    # 4H趋势共震
    trend_4h = get_4h_trend(ts)
    if dirc == 1 and trend_4h:
        s += 20  # 4H趋势也看多→强烈确认
    elif dirc == -1 and not trend_4h:
        s += 20  # 4H趋势看空→强烈确认
    elif dirc == 1 and not trend_4h:
        s -= 10  # 4H趋势看空→冲突
    
    # 15min小级别共震
    if has_15min_signal(ts, dirc):
        s += 10  # 15分钟也出现了同向信号→小级别确认
    
    return s

def backtest(lsig, ssig, mode, resonance_fn=None):
    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_fn(i,dirc) if resonance_fn else 0
                
                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("烛龙 v1.9：多时间框架共震（1H+4H+15min）")
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",resonance_v7)
ps(r7,s7,"v1.7 共振+Value Area")

r9,s9=backtest(lsig,ssig,"full",resonance_v9)
ps(r9,s9,"v1.9 多时间框架共震")

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