"""自由能原理离场 vs 固定止损 回测对比
核心逻辑：开仓后如果预测误差（价格没朝预期方向走）持续不降，提前离场
"""
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()

dd = df_raw.resample("1D").agg({"close":"last"}).dropna()
dd["ma20"] = dd["close"].rolling(20).mean()
dd["trend_up"] = dd["close"] > dd["ma20"]

O,H,L,C,V = d["open"].values,d["high"].values,d["low"].values,d["close"].values,d["volume"].values
n = len(d)

# ── 形态信号（烛龙 v2 标准版）──
O1=np.roll(O,1);H1=np.roll(H,1);L1=np.roll(L,1);C1=np.roll(C,1)

# 做多信号: 看涨吞没@支撑 + 下一根阳线确认
bull = (C1<O1) & (C>O) & (O<=C1) & (C>=O1)
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.005
nbull = np.roll(C,-1) > np.roll(O,-1)
long_sig = bull & near_s & nbull

# 做空信号: 向上扫荡@阻力 + 下一根阴线确认
swup = (H>pd.Series(H).shift(1).rolling(20).max().values) & (C<pd.Series(H).shift(1).rolling(20).max().values)
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.005
nbear = np.roll(C,-1) < np.roll(O,-1)
short_sig = swup & near_r & nbear

# 日线趋势
d_ts = d.index
trend_1h = np.array([
    dd["trend_up"].reindex([ts], method="ffill").values[0]
    if ts >= dd.index[0] else True
    for ts in d_ts
])

# 交易公共参数
split = int(n * 0.67)
SL, TP, MB = 0.015, 0.045, 48  # 至少48小时最大持仓

# ── 新增参数：各离场方法的配置 ──
# fixed: 原版固定SL/TP
# fe_n: 连续N根K线预测误差(未盈利)则提前离场
# fe_ret: N根K线后收益为负则提前离场
configs = [
    # (名, 离场函数参数)
    ("fixed", None),
    ("fe_3bar", 3),       # 连续3根K线未触达TP方向 → 提前离
    ("fe_5bar", 5),       # 连续5根
    ("fe_8bar", 8),       # 连续8根
]

def backtest(lsig, ssig, trend, exit_mode, exit_param):
    """通用回测函数，exit_mode控制离场方式"""
    results = {"in": [], "out": []}
    for period, st, en in [("in", 0, split), ("out", split, n)]:
        for mask, dirc in [(lsig, 1), (ssig, -1)]:
            for i in range(st, min(en, n)):
                if not mask[i]: continue
                if i + MB >= n: continue
                
                entry = C[i]
                closed = False
                dir_sign = 1 if dirc == 1 else -1
                
                # 记录预测误差（每根K线是否朝预期方向走）
                prediction_error = []
                
                for j in range(1, MB + 1):
                    if i + j >= n: break
                    ret = (C[i+j] / entry - 1) * dirc
                    
                    # 记录当前K线的预测误差
                    # 正误差 = 朝预期方向(好)  负误差 = 反方向(不好)
                    prediction_error.append(ret)
                    
                    # 检查是否触达TP
                    if ret >= TP:
                        results[period].append(TP - 0.001)
                        closed = True
                        break
                    
                    # 检查是否触达SL
                    if ret <= -SL:
                        results[period].append(-SL - 0.001)
                        closed = True
                        break
                    
                    # ── 自由能离场检查 ──
                    if exit_mode == "fe_3bar" and j >= 3:
                        # 最近3根K线都没盈利 → 预测误差持续不降
                        recent_2 = prediction_error[-3:]
                        if all(r <= 0 for r in recent_2):
                            results[period].append(ret - 0.001)
                            closed = True
                            break
                    
                    elif exit_mode == "fe_5bar" and j >= 5:
                        recent_2 = prediction_error[-5:]
                        if all(r <= 0 for r in recent_2):
                            results[period].append(ret - 0.001)
                            closed = True
                            break
                    
                    elif exit_mode == "fe_8bar" and j >= 8:
                        recent_2 = prediction_error[-8:]
                        if all(r <= 0 for r in recent_2):
                            results[period].append(ret - 0.001)
                            closed = True
                            break
                
                # 没触发任何条件 → 到期平仓
                if not closed:
                    final_ret = (C[min(i+MB, n-1)] / entry - 1) * dirc
                    results[period].append(final_ret - 0.001)
    
    # 打印结果
    for nm, tr in [("样本内", results["in"]), ("样本外", results["out"])]:
        if len(tr) < 5:
            print(f"  {nm}: 仅{len(tr)}笔, 不统计")
            continue
        wr = sum(1 for r in tr if r > 0) / len(tr)
        cum = np.prod([1 + r for r in tr])
        
        # 周度胜率
        ch = [tr[i:i+5] for i in range(0, len(tr), 5)]
        pw = sum(1 for c in ch if sum(c) > 0) / len(ch) if ch else 0
        
        avg_ret = np.mean(tr) * 100
        max_drawdown = 0
        running = 1.0
        peak = 1.0
        for r in tr:
            running *= (1 + r)
            peak = max(peak, running)
            dd = (running - peak) / peak
            max_drawdown = min(max_drawdown, dd)
        
        # Sharpe (简化)
        sharpe = np.mean(tr) / (np.std(tr) + 1e-9) * np.sqrt(365*24/MB) if np.std(tr) > 0 else 0
        
        print(f"  {nm}: {len(tr)}笔 wr={wr:.1%} cum={cum:.3f} avg={avg_ret:+.2f}% "
              f"周盈≈{pw:.1%} MDD={max_drawdown:.1%} sharpe={sharpe:.2f}")


# 过滤日线趋势
lsig = long_sig & trend_1h
ssig = short_sig & (~trend_1h)

print("=" * 70)
print("自由能预测误差离场 vs 固定止损 回测对比")
print("=" * 70)
print(f"\n信号: 看涨吞没@支撑+阳线(多) / 向上扫荡@阻力+阴线(空)")
print(f"SL=1.5% TP=4.5% 最大持仓{MB}h")
print(f"数据: BTC 1H, {n}根K线")
print(f"样本内: {split}根 ({d.index[0].strftime('%Y-%m-%d')} ~ {d.index[split].strftime('%Y-%m-%d')})")
print(f"样本外: {n-split}根 ({d.index[split].strftime('%Y-%m-%d')} ~ {d.index[-1].strftime('%Y-%m-%d')})")

for name, param in configs:
    print(f"\n{'─'*60}")
    if param is None:
        print(f"▶ {name}: 固定 SL/TP")
    else:
        print(f"▶ {name}: 连续{param}根K线未盈利 → 提前离场")
    backtest(lsig, ssig, trend_1h, name, param)

print("\n" + "=" * 70)
print("对比总结")
print("=" * 70)
print("\n固定止损 vs 自由能（连续N根未盈利离场）")
print("预期: 自由能版笔数更多(会提前割肉), 但最大回撤更低, 总体收益可能更好或持平")
print("关键是找到N的平衡点: N太小→频繁被震下车, N太大→跟固定止损没区别")