import pandas as pd, numpy as np
df_raw=pd.read_parquet("data/btc_multidim.parquet")
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;n=len(d)

# 蜡烛图
O1=np.roll(O,1);H1=np.roll(H,1);L1=np.roll(L,1);C1=np.roll(C,1)
O2=np.roll(O,2);C2=np.roll(C,2);body=abs(C-O);body1=abs(C1-O1)
nbull=np.roll(C,-1)>np.roll(O,-1);nbear=np.roll(C,-1)<np.roll(O,-1)
near_s=abs(L-pd.Series(L).shift(1).rolling(20).min().values)/(pd.Series(L).shift(1).rolling(20).min().values+1e-9)<0.008
near_r=abs(H-pd.Series(H).shift(1).rolling(20).max().values)/(pd.Series(H).shift(1).rolling(20).max().values+1e-9)<0.008
rev_l=((C1<O1)&(C>O)&(O<=C1)&(C>=O1)|(C2<O2)&(body1<(H-L)*0.3)&(C>O)&(C>(O2+C2)/2))&near_s&nbull
rev_s=((C1>O1)&(C<O)&(O>=C1)&(C<=O1)|(C2>O2)&(body1<(H-L)*0.3)&(C<O)&(C<(O2+C2)/2))&near_r&nbear

# ADX
tr=np.maximum(H-L,np.maximum(abs(H-np.roll(C,1)),abs(L-np.roll(C,1))))
a=pd.Series(tr).rolling(14).mean().values
pdi=100*pd.Series(np.where((H-np.roll(H,1)>np.roll(L,1)-L)&(H-np.roll(H,1)>0),H-np.roll(H,1),0)).rolling(14).mean().values/a
mdi=100*pd.Series(np.where((np.roll(L,1)-L>H-np.roll(H,1))&(np.roll(L,1)-L>0),np.roll(L,1)-L,0)).rolling(14).mean().values/a
adx=100*pd.Series(abs(pdi-mdi)/(pdi+mdi+1e-9)).rolling(14).mean().values
no_trend=np.nan_to_num(adx,nan=0)<30

# 价值区间 calc
di=d.index;dates=pd.Series(di.date,index=di)
vah_arr=np.zeros(n);val_arr=np.zeros(n)

for dt in dates.unique():
    mask=dates==dt
    px=(H[mask]+L[mask]+C[mask])/3;vl=V[mask]
    if len(px)<2 or vl.sum()==0: continue
    mn,mx=px.min(),px.max()
    if mx-mn<1: continue
    bins=np.linspace(mn,mx,50);vp=np.zeros(49)
    for i in range(len(px)):
        p=px[i];v=vl[i]
        for j in range(49):
            if bins[j]<=p<bins[j+1]:vp[j]+=v;break
    total=vp.sum()
    if total==0: continue
    best_s,best_e,best_sum=0,0,0
    for s in range(49):
        cs=0
        for e in range(s,49): cs+=vp[e]; 
        if cs>best_sum and cs<=total*0.7: best_sum=cs;best_s,best_e=s,e
    vah_arr[mask]=bins[best_e+1] if best_e+1<50 else bins[-1]
    val_arr[mask]=bins[best_s]

beyond_va=(C>vah_arr)|(C<val_arr)

# trend
dd=df_raw.resample("1D").agg({"close":"last"}).dropna();dd["ma20"]=dd["close"].rolling(20).mean()
du=dd["close"]>dd["ma20"];td=np.array([du.reindex([ts],method="ffill").values[0] if ts>=dd.index[0] else True for ts in di])

# backtest
SL,TP,MB,FEE=0.015,0.045,72,0.001;split=int(n*0.67)

configs=[
    ("L0 纯形态",rev_l&td,rev_s&(~td)),
    ("L1 +ADX<30",rev_l&td&no_trend,rev_s&(~td)&no_trend),
    ("L2 +偏离VA",rev_l&td&beyond_va,rev_s&(~td)&beyond_va),
    ("L3 全部(ADX+VA)",rev_l&td&no_trend&beyond_va,rev_s&(~td)&no_trend&beyond_va),
]

print(f"{'':18s} {'笔数':>5s} {'胜率':>7s} {'累计':>8s} {'周盈':>7s} {'月均':>5s}")
print("-"*50)
for label,ls,ss in configs:
    t=[]
    for mask,dirc in [(ls,1),(ss,-1)]:
        for i in range(split,min(n,len(mask))):
            if not mask[i] or i+MB>=n:continue
            e=C[i];w=l=False
            for j in range(1,min(MB,n-i-1)):
                r=(C[i+j]/e-1)*dirc
                if r<=-SL:t.append(-SL-FEE);l=True;break
                elif r>=TP:t.append(TP-FEE);w=True;break
            if not w and not l:t.append((C[min(i+MB,n-1)]/e-1)*dirc-FEE)
    if len(t)>3:
        wr=sum(1 for r in t if r>0)/len(t);cum=np.prod([1+r for r in t])
        ch=[t[i:i+5] for i in range(0,len(t),5)]
        pw=sum(1 for c in ch if sum(c)>0)/len(ch) if ch else 0
        print(f"{label:18s} {len(t):5d} {wr:6.1%} {cum:8.3f} {pw:6.1%} {len(t)/4:5.0f}")
