"""
蜡烛图 + 确认逻辑回测
测试不同确认方式对胜率的影响
"""
import pandas as pd, numpy as np

df_raw = pd.read_parquet("data/btc_multidim.parquet")
d = df_raw.resample("1h").agg({"open":"first","close":"last","high":"max","low":"min","volume":"sum"}).dropna()
s = pd.read_parquet("data/candle_signals.parquet")

FEE, SL, TP = 0.001, 0.01, 0.03

def backtest(signal_mask, direction, label):
    """三重屏障回测"""
    mask = signal_mask.values if hasattr(signal_mask, 'values') else signal_mask
    wins = losses = timeouts = 0
    rets = []
    for i in np.where(mask)[0]:
        if i + 48 >= len(d): continue
        entry = d["close"].iloc[i]
        for j in range(1, 49):
            if i+j >= len(d): break
            ret = (d["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 = (d["close"].iloc[min(i+48, len(d)-1)] / entry - 1) * direction
            timeouts += 1; rets.append(r - FEE)
    n = wins + losses + timeouts
    if n == 0: return None
    return {"n": n, "wr": wins/(wins+losses) if (wins+losses)>0 else 0,
            "avg_bps": np.mean(rets)*10000, "cum": np.prod([1+r for r in rets]),
            "wins": wins, "losses": losses}

# ── 确认规则 ──
# 1. 吞没后收阳确认 (做多)
bull_engulf = s["bullish_engulfing"].values.astype(bool) & s["near_support"].values.astype(bool)
next_candle_bull = (d["close"] > d["open"]).shift(-1).values.astype(bool)  # 下一根阳线
next_candle_bear = (d["close"] < d["open"]).shift(-1).values.astype(bool)

# 2. 吞没 + 放量确认
vol_spike = (d["volume"] > d["volume"].rolling(20).mean() * 1.3).values.astype(bool)

# 3. 吞没 + 扫荡双重信号
sweep_both = (s["sweep_down"].values.astype(bool) & s["bullish_engulfing"].values.astype(bool))

# 4. 不同确认强度
long_configs = [
    ("吞没@支撑(无确认)", bull_engulf, 1),
    ("+ 下一根阳线", bull_engulf & next_candle_bull, 1),
    ("+ 放量(1.3x)", bull_engulf & vol_spike, 1),
    ("+ 阳线+放量", bull_engulf & next_candle_bull & vol_spike, 1),
    ("吞没+扫荡双重", sweep_both & s["near_support"].values.astype(bool), 1),
]

bear_engulf = s["bearish_engulfing"].values.astype(bool) & s["near_resistance"].values.astype(bool)
short_configs = [
    ("吞没@阻力(无确认)", bear_engulf, -1),
    ("+ 下一根阴线", bear_engulf & next_candle_bear, -1),
    ("+ 放量(1.3x)", bear_engulf & vol_spike, -1),
    ("+ 阴线+放量", bear_engulf & next_candle_bear & vol_spike, -1),
]

# ── 测试 ──
print(f"{'做多确认':25s} {'笔数':>5s} {'胜率':>7s} {'均益bps':>9s} {'累计':>8s} {'W/L':>8s}")
print("-"*68)
for label, mask, dirc in long_configs:
    r = backtest(mask, dirc, label)
    if r:
        print(f"{label:25s} {r['n']:5d} {r['wr']:6.1%} {r['avg_bps']:+8.0f} {r['cum']:8.4f} {r['wins']:4d}/{r['losses']:4d}")

print(f"\n{'做空确认':25s} {'笔数':>5s} {'胜率':>7s} {'均益bps':>9s} {'累计':>8s} {'W/L':>8s}")
print("-"*68)
for label, mask, dirc in short_configs:
    r = backtest(mask, dirc, label)
    if r:
        print(f"{label:25s} {r['n']:5d} {r['wr']:6.1%} {r['avg_bps']:+8.0f} {r['cum']:8.4f} {r['wins']:4d}/{r['losses']:4d}")

# ── 找最优 ──
# 扫荡本身在阻力位就很好，加确认会怎样
sweep_up_res = s["sweep_up"].values.astype(bool) & s["near_resistance"].values.astype(bool)
sweep_down_sup = s["sweep_down"].values.astype(bool) & s["near_support"].values.astype(bool)

print(f"\n{'扫荡+确认':25s} {'笔数':>5s} {'胜率':>7s} {'均益bps':>9s} {'累计':>8s} {'W/L':>8s}")
print("-"*68)
sweep_configs = [
    ("扫荡空@阻力(无)", sweep_up_res, -1),
    ("+ 下一根阴线", sweep_up_res & next_candle_bear, -1),
    ("+ 放量", sweep_up_res & vol_spike, -1),
    ("+ 阴线+放量", sweep_up_res & next_candle_bear & vol_spike, -1),
    ("扫荡多@支撑(无)", sweep_down_sup, 1),
    ("+ 下一根阳线", sweep_down_sup & next_candle_bull, 1),
    ("+ 阳线+放量", sweep_down_sup & next_candle_bull & vol_spike, 1),
]
for label, mask, dirc in sweep_configs:
    r = backtest(mask, dirc, label)
    if r:
        print(f"{label:25s} {r['n']:5d} {r['wr']:6.1%} {r['avg_bps']:+8.0f} {r['cum']:8.4f} {r['wins']:4d}/{r['losses']:4d}")

# ── 组合策略: 做多(吞没+阳线@支撑) + 做空(扫荡+阴线@阻力) ──
print(f"\n{'='*68}")
print(f"组合策略回测 (做多:吞没+阳线@支撑 | 做空:扫荡+阴线@阻力)")
print(f"{'='*68}")

long_final = bull_engulf & next_candle_bull
short_final = sweep_up_res & next_candle_bear

all_trades = []
for mask, dirc in [(long_final, 1), (short_final, -1)]:
    m = mask if hasattr(mask, 'values') else mask
    for i in np.where(m)[0]:
        if i + 48 >= len(d): continue
        entry = d["close"].iloc[i]
        for j in range(1, 49):
            if i+j >= len(d): break
            ret = (d["close"].iloc[i+j] / entry - 1) * dirc
            if ret <= -SL:
                all_trades.append({"dir": dirc, "ret": -SL-FEE, "win": 0, "ts": d.index[i]})
                break
            elif ret >= TP:
                all_trades.append({"dir": dirc, "ret": TP-FEE, "win": 1, "ts": d.index[i]})
                break
        else:
            r = (d["close"].iloc[min(i+48, len(d)-1)] / entry - 1) * dirc
            all_trades.append({"dir": dirc, "ret": r-FEE, "win": 1 if r>0 else 0, "ts": d.index[i]})

tdf = pd.DataFrame(all_trades).sort_values("ts")
n = len(tdf)
wr = tdf["win"].mean()
cum = (1 + tdf["ret"]).prod()
avg = tdf["ret"].mean() * 10000

print(f"总交易: {n} 笔 | 胜率: {wr:.1%} | 均益: {avg:+.0f} bps | 累计: {cum:.4f} ({(cum-1)*100:+.1f}%)")

# 月度统计
tdf["month"] = pd.to_datetime(tdf["ts"]).dt.to_period("M")
monthly = tdf.groupby("month").agg(
    n=("ret","count"), wr=("win","mean"), 
    avg=("ret", lambda x: x.mean()*10000),
    cum=("ret", lambda x: (1+x).prod()-1)
)
print(f"\n月度:\n{monthly.round(3).to_string()}")

# 周统计
tdf["week"] = pd.to_datetime(tdf["ts"]).dt.to_period("W")
weekly = tdf.groupby("week").agg(
    n=("ret","count"), pnl=("ret","sum")
)
profitable_weeks = (weekly["pnl"] > 0).sum()
total_weeks = len(weekly)
print(f"\n周统计: {total_weeks}周, {profitable_weeks}周盈利 ({profitable_weeks/total_weeks:.1%})")
print(f"周均交易: {weekly['n'].mean():.0f}笔, 周均盈亏: {weekly['pnl'].mean()*100:.2f}%")
