"""
1H 简化模型 + 回测
只用 ret_20（IC=0.16）+ fr_level 过滤
含手续费 0.05% + 资金费率
"""
import pandas as pd, numpy as np

df = pd.read_parquet("data/btc_multidim.parquet")

# Resample → 1H
d = df.resample("1h").agg({
    "open": "first", "high": "max", "low": "min", "close": "last",
    "volume": "sum", "taker_pct": "mean", "funding_rate": "last"
}).dropna()

# 1个信号 + 1个过滤器
d["ret_20"] = d["close"].pct_change(20)
d["fr_level"] = d["funding_rate"]

# 信号: ret_20 > 0 → 做多; ret_20 < 0 → 做空
# 过滤: |fr_level| 极端时不做（资金费率极端说明市场失衡）
d["signal"] = 0
d.loc[d["ret_20"] > 0, "signal"] = 1    # 做多
d.loc[d["ret_20"] < 0, "signal"] = -1   # 做空

# 极端资金费率反向过滤
d.loc[(d["fr_level"] > 0.001) & (d["signal"] == 1), "signal"] = 0   # 多头过热不跟
d.loc[(d["fr_level"] < -0.001) & (d["signal"] == -1), "signal"] = 0 # 空头恐慌不跟

# 未来收益 (持有4小时)
d["fwd_close"] = d["close"].shift(-4)
d["fwd_ret"] = d["fwd_close"] / d["close"] - 1

d = d.dropna()

# ── 回测 ──
FEE = 0.0005  # 0.05% 手续费
FR_COST = 0.0001 / 8  # 每8小时约0.01%, 摊到每小时

trades = []
pos = 0  # 0=none, 1=long, -1=short
entry_price = 0

for i in range(len(d) - 4):
    row = d.iloc[i]
    
    # 开仓
    if pos == 0 and row["signal"] != 0:
        pos = row["signal"]
        entry_price = row["close"]
        trades.append({"ts": d.index[i], "type": "open", "side": "long" if pos==1 else "short",
                       "price": entry_price, "signal_ret20": row["ret_20"],
                       "fr": row["fr_level"]})
    
    # 平仓 (4小时后)
    elif pos != 0 and i % 4 == 3:
        exit_price = d.iloc[i]["close"]
        raw_ret = (exit_price / entry_price - 1) * pos
        fee_cost = 2 * FEE  # 开+平
        fr_cost = FR_COST * 4  # 持有4小时
        net_ret = raw_ret - fee_cost - fr_cost
        
        trades.append({"ts": d.index[i], "type": "close", "side": "long" if pos==1 else "short",
                       "price": exit_price, "raw_ret": raw_ret, "net_ret": net_ret})
        pos = 0

# ── 统计 ──
closes = [t for t in trades if t["type"] == "close"]
if not closes:
    print("无交易信号")
    exit()

rets = [t["net_ret"] for t in closes]
raw_rets = [t["raw_ret"] for t in closes]
wins = sum(1 for r in rets if r > 0)
total = len(rets)

cum = 1.0
for r in rets:
    cum *= (1 + r)

print(f"=== 1H 简化回测 ===\n")
print(f"信号: ret_20方向 + fr极端过滤")
print(f"持仓: 4小时固定")
print(f"交易: {total} 笔")
print(f"胜率: {wins}/{total} = {wins/total:.1%}")
print(f"平均毛收益: {np.mean(raw_rets)*10000:+.1f} bps")
print(f"平均净收益: {np.mean(rets)*10000:+.1f} bps")
print(f"最大单笔: {max(rets)*10000:+.1f} bps / 最小: {min(rets)*10000:+.1f} bps")
print(f"累计: {cum:.4f} ({(cum-1)*100:+.1f}%)")

# 分段统计
half = len(rets) // 2
first_half = rets[:half]
second_half = rets[half:]
print(f"\n前半段: avg={np.mean(first_half)*10000:+.1f} bps, wins={(sum(1 for r in first_half if r>0))/len(first_half):.1%}")
print(f"后半段: avg={np.mean(second_half)*10000:+.1f} bps, wins={(sum(1 for r in second_half if r>0))/len(second_half):.1%}")

# 时间分布
df_trades = pd.DataFrame(closes)
df_trades["month"] = pd.to_datetime(df_trades["ts"]).dt.to_period("M")
monthly = df_trades.groupby("month")["net_ret"].agg(["mean","count","sum"])
print(f"\n月度:\n{monthly.to_string()}")
