"""
缠论快速版: 跳过包含处理，直接在原始K线上找分型+笔
实际交易中缠论使用者经常跳过包含处理也能得到有效信号
"""
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"}).dropna()
H, L, C, O = d["high"].values, d["low"].values, d["close"].values, d["open"].values
n = len(d)

# 4H方向
d4h = df_raw.resample("4h").agg({"close":"last"}).dropna()
d4h["ma50"] = d4h["close"].rolling(50).mean()
d4h_trend = d4h["close"] > d4h["ma50"]

# ── 1. 分型 (直接在原始K线上) ──
fx_list = []
for i in range(1, n-1):
    # 顶分型
    if H[i] > H[i-1] and H[i] > H[i+1] and L[i] > L[i-1] and L[i] > L[i+1]:
        fx_list.append({"idx": i, "type": "顶", "price": H[i]})
    # 底分型
    elif L[i] < L[i-1] and L[i] < L[i+1] and H[i] < H[i-1] and H[i] < H[i+1]:
        fx_list.append({"idx": i, "type": "底", "price": L[i]})

print(f"分型: {len(fx_list)}")

# ── 2. 笔 ──
bi_list = []
i = 0
while i < len(fx_list) - 1:
    f1 = fx_list[i]
    j = i + 1
    while j < len(fx_list):
        f2 = fx_list[j]
        if f2["type"] != f1["type"] and f2["idx"] - f1["idx"] >= 2:
            # 上升笔: 底→顶
            if f1["type"] == "底" and f2["price"] > f1["price"]:
                bi_list.append({"start": f1["idx"], "end": f2["idx"], 
                               "dir": "up", "start_px": f1["price"], "end_px": f2["price"]})
                i = j; break
            # 下降笔: 顶→底
            elif f1["type"] == "顶" and f2["price"] < f1["price"]:
                bi_list.append({"start": f1["idx"], "end": f2["idx"],
                               "dir": "down", "start_px": f1["price"], "end_px": f2["price"]})
                i = j; break
        j += 1
    if j == len(fx_list): break

print(f"笔: {len(bi_list)} (上升{sum(1 for b in bi_list if b['dir']=='up')}, 下降{sum(1 for b in bi_list if b['dir']=='down')})")

# ── 3. 交易回测 ──
FEE = 0.001

def run(sl, tp, max_bars, use_htf, label):
    trades = []
    for b in bi_list:
        # 底分型+上升笔起点 → 潜在做多
        if b["dir"] == "up":
            start_idx = b["start"]
            # 等价格回到分型低点附近
            for j in range(start_idx + 1, min(start_idx + 15, n)):
                if abs(C[j] - b["start_px"]) / b["start_px"] > 0.01:
                    continue
                # 4H过滤
                if use_htf:
                    ts = d.index[j]
                    htf_val = d4h_trend.reindex([ts], method="ffill").values[0]
                    if not htf_val: continue
                
                direction = 1
                entry = C[j]
                win, loss = False, False
                for k in range(1, min(max_bars, n - j - 1)):
                    ret = (C[j+k] / entry - 1) * direction
                    if ret <= -sl:
                        trades.append(-sl - FEE); loss = True; break
                    elif ret >= tp:
                        trades.append(tp - FEE); win = True; break
                if not win and not loss:
                    trades.append((C[min(j+max_bars, n-1)]/entry-1)*direction - FEE)
        
        # 顶分型 → 做空
        if b["dir"] == "down":
            start_idx = b["start"]
            for j in range(start_idx + 1, min(start_idx + 15, n)):
                if abs(C[j] - b["start_px"]) / b["start_px"] > 0.01:
                    continue
                if use_htf:
                    ts = d.index[j]
                    htf_val = d4h_trend.reindex([ts], method="ffill").values[0]
                    if htf_val: continue  # 4H上升不做空
                
                direction = -1
                entry = C[j]
                win, loss = False, False
                for k in range(1, min(max_bars, n - j - 1)):
                    ret = (C[j+k] / entry - 1) * direction
                    if ret <= -sl:
                        trades.append(-sl - FEE); loss = True; break
                    elif ret >= tp:
                        trades.append(tp - FEE); win = True; break
                if not win and not loss:
                    trades.append((C[min(j+max_bars, n-1)]/entry-1)*direction - FEE)
    
    if len(trades) < 10: return None
    wr = sum(1 for r in trades if r > 0) / len(trades)
    cum = np.prod([1+r for r in trades])
    return {"label": label, "n": len(trades), "wr": wr, "cum": cum, "avg": np.mean(trades)*10000}

# 测试
configs = [
    (0.01, 0.03, 48, False, "SL1% TP3%"),
    (0.01, 0.03, 48, True, "SL1% TP3% +4H"),
    (0.015, 0.045, 72, False, "SL1.5% TP4.5%"),
    (0.015, 0.045, 72, True, "SL1.5% TP4.5% +4H"),
    (0.02, 0.06, 96, False, "SL2% TP6%"),
    (0.02, 0.06, 96, True, "SL2% TP6% +4H"),
    (0.02, 0.04, 48, True, "SL2% TP4% +4H"),
    (0.015, 0.03, 48, True, "SL1.5% TP3% +4H"),
]

print(f"\n{'='*52}")
print(f"{'参数':22s} {'n':>5s} {'胜率':>7s} {'均益':>8s} {'累计':>8s}")
print(f"{'-'*52}")
best = None
for sl, tp, mb, htf, label in configs:
    r = run(sl, tp, mb, htf, label)
    if r:
        print(f"{label:22s} {r['n']:5d} {r['wr']:6.1%} {r['avg']:+7.0f}bps {r['cum']:8.4f}")
        if best is None or r["cum"] > best[1]:
            best = (label, r["cum"], r)

print(f"\n🏆 最佳: {best[0]} 累计={best[1]:.4f}")
