"""主动推理版：每3小时评估不确定性，动态调整仓位
完整版：融合回测 + 推理逻辑 + 仓位管理
"""
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)

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

# ── 主动推理配置 ──
configs = [
    ("fixed", None),               # 原版固定止损
    ("active_inference", 1),       # 主动推理版
]

def backtest_with_inference(lsig, ssig, trend, mode):
    results = {"in": [], "out": [], "total_trades": 0, "inference_count": 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
                pos_size = 1.0  # 初始仓位 1x
                
                # 模拟持仓过程
                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 == "active_inference" and j % 3 == 0:  # 每3小时评估一次
                        expected = 0.001  # 预期每3小时涨0.1%
                        error = abs(ret - expected)
                        
                        # 状态判断
                        if error < 0.0015:  # 高确定性 → 加仓
                            pos_size = min(2.0, pos_size * 1.2)
                            results["inference_count"] += 1
                        elif 0.0015 <= error < 0.004:  # 中等不确定性 → 保持
                            pass
                        else:  # 低确定性 → 减仓
                            pos_size = max(0.2, pos_size * 0.8)
                            results["inference_count"] += 1
                        
                        # 检查是否减仓到0 → 平仓
                        if pos_size <= 0.2:
                            results[period].append(ret * pos_size - 0.001)
                            closed = True; break
                    
                if not closed:
                    final_ret = (C[min(i+MB, n-1)] / entry - 1) * dirc
                    results[period].append(final_ret * pos_size - 0.001)
                    results["total_trades"] += 1
    
    # 打印统计
    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['inference_count']}")
    
    return results

print("=" * 80)
print("主动推理版：每3小时评估不确定性，动态调整仓位")
print("=" * 80)
print(f"SL={SL} TP={TP} 最大持仓{MB}h")

for name, mode in configs:
    print(f"\n{'─'*70}")
    print(f"▶ {name}")
    backtest_with_inference(lsig, ssig, trend_1h, mode)