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)
bull_engulf = (C1<O1)&(C>O)&(O<=C1)&(C>=O1)
near_sup = abs(L-pd.Series(L).shift(1).rolling(20).min().values)/(pd.Series(L).shift(1).rolling(20).min().values+1e-9)<0.005
next_bull = (np.roll(C,-1)>np.roll(O,-1))
sweep_up = (H>pd.Series(H).shift(1).rolling(20).max().values)&(C<pd.Series(H).shift(1).rolling(20).max().values)
near_res = abs(H-pd.Series(H).shift(1).rolling(20).max().values)/(pd.Series(H).shift(1).rolling(20).max().values+1e-9)<0.005
next_bear = (np.roll(C,-1)<np.roll(O,-1))

long_sig = bull_engulf & near_sup & next_bull
short_sig = sweep_up & near_res & next_bear

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
FEE=0.001

def run_fixed(sl,tp,mb):
    trades=[]
    for mask,dirc in [(long_sig,1),(short_sig,-1)]:
        for i in np.where(mask)[0]:
            if 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: trades.append(-sl-FEE); l=True; break
                elif ret>=tp: trades.append(tp-FEE); w=True; break
            if not w and not l: trades.append((C[min(i+mb,n-1)]/entry-1)*dirc-FEE)
    return trades

def run_atr(sl_mul,tp_mul,mb):
    trades=[]
    for mask,dirc in [(long_sig,1),(short_sig,-1)]:
        for i in np.where(mask)[0]:
            if i+mb>=n or np.isnan(atr[i]) or atr[i]==0: continue
            sl = (sl_mul*atr[i])/C[i]; tp = (tp_mul*atr[i])/C[i]
            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: trades.append(-sl-FEE); l=True; break
                elif ret>=tp: trades.append(tp-FEE); w=True; break
            if not w and not l: trades.append((C[min(i+mb,n-1)]/entry-1)*dirc-FEE)
    return trades

print(f"{'参数':22s} {'n':>5s} {'胜率':>7s} {'均益':>8s} {'累计':>8s} {'周盈率':>7s}")
print("-"*62)

# 周统计函数
def weekly(trades_list):
    # 简化：按每7笔分组模拟周
    chunks = [trades_list[i:i+7] for i in range(0,len(trades_list),7)]
    pos_weeks = sum(1 for c in chunks if sum(c)>0)
    return pos_weeks/len(chunks) if chunks else 0

all_configs = [
    ("固定 SL1% TP3%", lambda: run_fixed(0.01,0.03,48)),
    ("固定 SL1.5% TP4.5%", lambda: run_fixed(0.015,0.045,72)),
    ("ATR 1x/3x", lambda: run_atr(1.0,3.0,48)),
    ("ATR 1.5x/4.5x", lambda: run_atr(1.5,4.5,72)),
    ("ATR 1x/4x", lambda: run_atr(1.0,4.0,72)),
    ("ATR 0.8x/3.2x", lambda: run_atr(0.8,3.2,48)),
    ("ATR 0.5x/2x", lambda: run_atr(0.5,2.0,48)),
]

for label, fn in all_configs:
    t = fn()
    if len(t)<10: continue
    wr = sum(1 for r in t if r>0)/len(t)
    cum = np.prod([1+r for r in t])
    wk = weekly(t)
    print(f"{label:22s} {len(t):5d} {wr:6.1%} {np.mean(t)*10000:+7.0f}bps {cum:8.4f} {wk:6.1%}")
