"""
缠论级别过滤: "看大做小"
在1H信号上加4H/日线方向过滤
"""
import pandas as pd, numpy as np

df = pd.read_parquet("data/btc_multidim.parquet")

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

# 4H数据
d4h = df.resample("4h").agg({"open":"first","close":"last","high":"max","low":"min"}).dropna()

# 日线数据
dd = df.resample("1D").agg({"open":"first","close":"last","high":"max","low":"min"}).dropna()

# ── 大级别方向定义 ──
# 4H趋势: close相对50周期MA
d4h["ma50"] = d4h["close"].rolling(50).mean()
d4h["trend_up"] = d4h["close"] > d4h["ma50"]
d4h["trend_dn"] = d4h["close"] < d4h["ma50"]

# 日线趋势: 同理
dd["ma20"] = dd["close"].rolling(20).mean()
dd["trend_up"] = dd["close"] > dd["ma20"]
dd["trend_dn"] = dd["close"] < dd["ma20"]

# 映射到1H (前向填充)
d1h["d_trend_up"] = dd["trend_up"].reindex(d1h.index, method="ffill").fillna(False)
d1h["d_trend_dn"] = dd["trend_dn"].reindex(d1h.index, method="ffill").fillna(False)
d1h["h4_trend_up"] = d4h["trend_up"].reindex(d1h.index, method="ffill").fillna(False)
d1h["h4_trend_dn"] = d4h["trend_dn"].reindex(d1h.index, method="ffill").fillna(False)

# ── 1H蜡烛图信号 (复用之前的逻辑) ──
O, H, L, C = d1h["open"], d1h["high"], d1h["low"], d1h["close"]
body = abs(C - O)
O1, H1, L1, C1 = O.shift(1), H.shift(1), L.shift(1), C.shift(1)

LB = 20
prev_high = H.shift(1).rolling(LB).max()
prev_low = L.shift(1).rolling(LB).min()
near_support = (L - prev_low).abs() / (prev_low + 1e-9) < 0.005
near_resistance = (H - prev_high).abs() / (prev_high + 1e-9) < 0.005

bullish_engulfing = (C1 < O1) & (C > O) & (O <= C1) & (C >= O1)
sweep_up = (H > prev_high) & (C < prev_high)
next_bull = (C.shift(-1) > O.shift(-1))
next_bear = (C.shift(-1) < O.shift(-1))

# 基础信号
long_base = bullish_engulfing & near_support & next_bull
short_base = sweep_up & near_resistance & next_bear

# ── 回测 ──
FEE, SL, TP, MAX_BARS = 0.001, 0.01, 0.03, 48

def run(sig_mask, direction, label):
    mask = sig_mask.values
    rets, wins, losses, timeouts = [], 0, 0, 0
    for i in np.where(mask)[0]:
        if i + MAX_BARS >= len(d1h): continue
        entry = d1h["close"].iloc[i]
        for j in range(1, MAX_BARS+1):
            if i+j >= len(d1h): break
            ret = (d1h["close"].iloc[i+j] / entry - 1) * direction
            if ret <= -SL:
                losses += 1; rets.append(-SL - FEE); break
            elif ret >= TP:
                wins += 1; rets.append(TP - FEE); break
        else:
            r = (d1h["close"].iloc[min(i+MAX_BARS, len(d1h)-1)] / entry - 1) * direction
            timeouts += 1; rets.append(r - FEE)
    n = wins + losses + timeouts
    if n == 0: return None
    wr = wins/(wins+losses) if (wins+losses)>0 else 0
    cum = np.prod([1+r for r in rets])
    return {"label": label, "n": n, "wr": wr, "cum": cum, "avg": np.mean(rets)*10000}

# ── 测试: 无过滤 vs 4H过滤 vs 日线过滤 vs 双重过滤 ──
tests = [
    ("无过滤(基准)", long_base, short_base, None, None),
    ("+4H方向过滤", long_base, short_base, d1h["h4_trend_up"], d1h["h4_trend_dn"]),
    ("+日线方向过滤", long_base, short_base, d1h["d_trend_up"], d1h["d_trend_dn"]),
]

print("="*60)
print("缠论 级别过滤: '看大做小'")
print("="*60)
print(f"{'过滤方式':20s} {'交易':>5s} {'胜率':>7s} {'累计':>8s} {'年均bps':>9s}")
print("-"*56)

for label, long_s, short_s, lt_filter, st_filter in tests:
    l_mask = long_s if lt_filter is None else long_s & lt_filter
    s_mask = short_s if st_filter is None else short_s & st_filter
    
    rl = run(l_mask, 1, "")
    rs = run(s_mask, -1, "")
    if rl is None or rs is None: continue
    
    all_n = rl["n"] + rs["n"]
    all_cum = rl["cum"] * rs["cum"]
    all_wr = (rl["wr"]*rl["n"] + rs["wr"]*rs["n"]) / all_n
    
    print(f"{label:20s} {all_n:5d} {all_wr:6.1%} {all_cum:8.4f} {rl['avg']+rs['avg']:+8.0f}")

# 看大级别到底过滤了什么
print(f"\n--- 信号过滤详情 ---")
for name, base, trend_f, side in [
    ("做多(吞没+阳线)", long_base, d1h["d_trend_up"], "多"),
    ("做空(扫荡+阴线)", short_base, d1h["d_trend_dn"], "空"),
]:
    raw_n = base.sum()
    filtered_n = (base & trend_f).sum()
    dropped = raw_n - filtered_n
    print(f"  {name}: {raw_n}→{filtered_n} (过滤了{dropped}笔逆势信号)")
