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);H2=np.roll(H,2);L2=np.roll(L,2);C2=np.roll(C,2)
body=abs(C-O);body1=abs(C1-O1);body2=abs(C2-O2);ra=H-L
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

b1=(C1<O1)&(C>O)&(O<=C1)&(C>=O1)
b2=(C2<O2)&(body1<ra*0.3)&(C>O)&(C>(O2+C2)/2)
rev_long=(b1|b2)&near_s&nbull
b3=(C1>O1)&(C<O)&(O>=C1)&(C<=O1)
b4=(C2>O2)&(body1<ra*0.3)&(C<O)&(C<(O2+C2)/2)
rev_short=(b3|b4)&near_r&nbear

tr=np.maximum(H-L,np.maximum(abs(H-np.roll(C,1)),abs(L-np.roll(C,1))))
atr=pd.Series(tr).rolling(14).mean().values
up=H-np.roll(H,1);dn=np.roll(L,1)-L
pdm=np.where((up>dn)&(up>0),up,0);mdm=np.where((dn>up)&(dn>0),dn,0)
pdi=100*pd.Series(pdm).rolling(14).mean().values/atr
mdi=100*pd.Series(mdm).rolling(14).mean().values/atr
adx=100*pd.Series(abs(pdi-mdi)/(pdi+mdi+1e-9)).rolling(14).mean().values

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

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

print("ADX过滤对样本外的影响:")
for adx_limit in [0,25,30,35]:
    lf=rev_long&td;sf=rev_short&(~td)
    if adx_limit>0:
        safe_adx=np.nan_to_num(adx,nan=0)
        lf=lf&(safe_adx<adx_limit)
        sf=sf&(safe_adx<adx_limit)
    t=[]
    for mask,dirc in [(lf,1),(sf,-1)]:
        for i in range(split,min(n,len(mask))):
            if not mask[i] or i+MB>=n: continue
            entry=C[i];w=l=False
            for j in range(1,min(MB,n-i-1)):
                ret=(C[i+j]/entry-1)*dirc
                if ret<=-SL:t.append(-SL-FEE);l=True;break
                elif ret>=TP:t.append(TP-FEE);w=True;break
            if not w and not l:t.append((C[min(i+MB,n-1)]/entry-1)*dirc-FEE)
    if len(t)>5:
        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
        tag="无过滤" if adx_limit==0 else f"ADX<{adx_limit}"
        print(f"  {tag:10s}: {len(t):4d}笔 wr={wr:.1%} cum={cum:.3f} 周盈={pw:.1%}")
