"""自由能预测误差 v2: 阶梯期望 + 波动率自适应 + 动态减仓"""
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)

# ── ATR ──
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
atr_pct = atr / C * 100  # ATR占价格的百分比

# ── 形态信号 ──
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
])

lsig = long_sig & trend_1h
ssig = short_sig & (~trend_1h)

split = int(n * 0.67)
SL, TP, MB = 0.015, 0.045, 48

# ── 配置4种离场方法 ──
configs = [
    ("fixed", None),                              # 原版
    ("阶梯期望(无ATR)", 0),                        # 固定阈值期望
    ("阶梯期望(+ATR)", 1),                         # ATR自适应
    ("阶梯期望+半仓ATR", 2),                        # ATR + 减半仓
]

def backtest_v2(lsig, ssig, trend, mode, sub_mode):
    results = {"in": [], "out": [], "early_exit": 0}
    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
                
                for j in range(1, MB + 1):
                    if i + j >= n: break
                    ret = (C[i+j] / entry - 1) * dirc
                    
                    # TP/SL 优先
                    if ret >= TP:
                        results[period].append(TP - 0.001)
                        closed = True; break
                    if ret <= -SL:
                        results[period].append(-SL - 0.001)
                        closed = True; break
                    
                    # ── 自由能离场 ──
                    if mode == "fixed":
                        continue  # 不做任何自由能处理
                    
                    # 阶梯期望：根据持仓时长判断"是否达到预期"
                    expected = {1: 0.001, 2: 0.003, 3: 0.005, 4: 0.008}
                    
                    if sub_mode == 0:
                        # 固定期望：第1根涨0.1%, 第2根0.3%, 第3根0.5%, 第4根0.8%
                        if j in expected and ret < expected[j]:
                            # 连续低于期望→提前走
                            results[period].append(ret - 0.001)
                            results["early_exit"] += 1
                            closed = True; break
                    
                    elif sub_mode == 1:
                        # ATR自适应期望
                        atr_val = atr_pct[i+j] if i+j < n else atr_pct[min(i+j, n-1)]
                        atr_factor = max(0.3, min(0.8, atr_val / 0.3 * 0.5))  # 波动率越大要求越高
                        adaptive_expected = {1: 0.001 * atr_factor, 2: 0.003 * atr_factor,
                                             3: 0.005 * atr_factor, 4: 0.008 * atr_factor}
                        if j in adaptive_expected and ret < adaptive_expected[j]:
                            results[period].append(ret - 0.001)
                            results["early_exit"] += 1
                            closed = True; break
                    
                    elif sub_mode == 2:
                        # ATR期望 + 半仓处理
                        atr_val = atr_pct[i+j] if i+j < n else atr_pct[min(i+j, n-1)]
                        atr_factor = max(0.3, min(0.8, atr_val / 0.3 * 0.5))
                        adaptive_expected = {1: 0.001 * atr_factor, 2: 0.003 * atr_factor,
                                             3: 0.005 * atr_factor, 4: 0.008 * atr_factor}
                        if j in adaptive_expected and ret < adaptive_expected[j]:
                            # 半仓处理（这里模拟：减半结果）
                            half_ret = ret * 0.5  # 半仓止损,半仓继续
                            results[period].append(half_ret - 0.001)
                            results["early_exit"] += 1
                            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
        running = 1.0; peak = 1.0; max_dd = 0
        for r in tr:
            running *= (1 + r); peak = max(peak, running)
            dd = (running - peak) / peak; max_dd = min(max_dd, dd)
        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_dd:.1%} sharpe={sharpe:.2f} "
              f"早退={results['early_exit']}")
    return results

print("=" * 80)
print("自由能 v2: 阶梯期望 + 波动率自适应 + 动态减仓")
print("=" * 80)
print(f"SL={SL} TP={TP} 最大持仓{MB}h")

for name, sub in configs:
    print(f"\n{'─'*70}")
    print(f"▶ {name}")
    backtest_v2(lsig, ssig, trend_1h, name if name == "fixed" else "fe", sub)

# ── 再跑一个更细致的：分阶段期望（前8根逐根检查）──
print(f"\n{'─'*70}")
print("▶ 补充: 更精细的3阶段期望")
print("  阶段1(1-2根): 期望+0.15%  阶段2(3-5根): 期望+0.4%")
print("  阶段3(6-10根): 期望+0.8%  否则提前离")

results3 = {"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
            for j in range(1, MB + 1):
                if i + j >= n: break
                ret = (C[i+j] / entry - 1) * dirc
                if ret >= TP:
                    results3[period].append(TP - 0.001); closed = True; break
                if ret <= -SL:
                    results3[period].append(-SL - 0.001); closed = True; break
                # 3阶段检查
                if 1 <= j <= 2 and ret < 0.0015:  # 前2根涨不到0.15%
                    results3[period].append(ret - 0.001); closed = True; break
                elif 3 <= j <= 5 and ret < 0.004:  # 3-5根涨不到0.4%
                    results3[period].append(ret - 0.001); closed = True; break
                elif 6 <= j <= 10 and ret < 0.008:  # 6-10根涨不到0.8%
                    results3[period].append(ret - 0.001); closed = True; break
            if not closed:
                final_ret = (C[min(i+MB, n-1)] / entry - 1) * dirc
                results3[period].append(final_ret - 0.001)

for nm, tr in results3.items():
    if len(tr) < 5: 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
    running = 1.0; peak = 1.0; max_dd = 0
    for r in tr:
        running *= (1 + r); peak = max(peak, running)
        dd = (running - peak) / peak; max_dd = min(max_dd, dd)
    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_dd:.1%} sharpe={sharpe:.2f}")