"""事件驱动回测引擎。

时序语义（避免未来数据）：
- 指标在完整历史上预先计算（全部为因果指标：ewm/rolling/shift/merge_asof backward）。
- Controller 在第 t 根【收盘】产生信号。
- PositionExecutor 于第 t+1 根【开盘价】成交（entry_at_next_open）。
- GridExecutor 的限价单自第 t+1 根起生效，触及 high/low 即成交
  （保守假设：不含 maker 排队优势，成交价为限价）。
- 止损/移动止盈按 K 线极值触发，跳空按更差的开盘价成交。

所有成交双边计 fee_rate（默认 0.00035，压力测试用 0.0005）。
"""
from __future__ import annotations

import logging
from typing import Optional

import pandas as pd

from bot.executors.base import ExecutionContext, ExecutorBase
from bot.models import CloseType, TradeRecord
from bot.strategies.base import ControllerBase

logger = logging.getLogger(__name__)


class BacktestContext(ExecutionContext):
    def __init__(self, fee_rate: float, wallet: float):
        super().__init__(fee_rate, wallet)
        self.trade_log: list[TradeRecord] = []

    def apply_close_fill(self, ex, price, amount, ts, close_type, tag=""):
        n_before = len(ex.trades)
        super().apply_close_fill(ex, price, amount, ts, close_type, tag)
        if len(ex.trades) > n_before:
            self.trade_log.extend(ex.trades[n_before:])


class BacktestEngine:
    def __init__(self, controller: ControllerBase, data: dict[str, pd.DataFrame],
                 btc_1d: Optional[pd.DataFrame] = None,
                 fee_rate: float = 0.00035, wallet: float = 1000.0,
                 timeframe_s: int = 300):
        """
        data: {pair: DataFrame[ts,open,high,low,close,volume]}，ts 为秒级时间戳。
        """
        self.controller = controller
        self.fee_rate = fee_rate
        self.timeframe_s = timeframe_s
        self.ctx = BacktestContext(fee_rate, wallet)
        controller.ctx = self.ctx  # 供策略做风险定仓 / StoplossGuard

        # 预计算指标
        self.data: dict[str, pd.DataFrame] = {}
        for pair, df in data.items():
            self.data[pair] = controller.prepare_dataframe(
                pair, df.sort_values("ts").reset_index(drop=True), btc_1d)

        self.active: list[ExecutorBase] = []
        self.done: list[ExecutorBase] = []
        self._pending_entries: list[ExecutorBase] = []
        self.equity_curve: list[tuple[float, float]] = []

    # ------------------------------------------------------------------
    def run(self) -> "BacktestContext":
        timeline = sorted(set().union(*[set(df["ts"]) for df in self.data.values()]))
        # 每 pair 的 ts → 行索引
        idx_maps = {p: {ts: i for i, ts in enumerate(df["ts"])}
                    for p, df in self.data.items()}
        step_i = 0

        for ts in timeline:
            step_i += 1
            for pair, df in self.data.items():
                i = idx_maps[pair].get(ts)
                if i is None:
                    continue
                candle = df.iloc[i]

                # 1) 待入场 executor：以本根开盘价成交
                for ex in list(self._pending_entries):
                    if ex.pair != pair:
                        continue
                    stake = getattr(ex, "planned_stake", 0.0)
                    open_px = float(candle["open"])
                    amount = stake * ex.leverage / open_px if open_px > 0 else 0.0
                    if amount > 0 and stake <= self.ctx.free_wallet:
                        self.ctx.apply_open_fill(ex, open_px, amount, ts)
                        ex.created_ts = ts
                        self.active.append(ex)
                    self._pending_entries.remove(ex)

                # 2) 活动 executor 状态机推进（网格撮合 / 屏障检查）
                for ex in list(self.active):
                    if ex.pair != pair or not ex.is_active:
                        continue
                    n_trades_before = len(ex.trades)
                    ex.on_candle(candle, self.ctx)
                    self._check_liquidation(ex, candle, ts)
                    # StoplossGuard 等熔断通知
                    for t in ex.trades[n_trades_before:]:
                        notify = getattr(self.controller, "notify_trade_closed", None)
                        if notify:
                            notify(t, ts)
                    if not ex.is_active:
                        self.done.append(ex)
                        self.active.remove(ex)

                # 3) Controller 信号（收盘后产生，次根入场）
                if i >= self.controller.warmup_candles:
                    # 截止当前的因果视图（iloc 切片为 view，无拷贝）
                    self.controller.dataframes[pair] = df.iloc[: i + 1]
                    new_exs = self.controller.on_candle(pair, float(ts), self.active)
                    for ex in new_exs:
                        if getattr(ex, "entry_at_next_open", False):
                            self._pending_entries.append(ex)
                        else:
                            ex.created_ts = float(ts)
                            self.active.append(ex)

            # 4) 权益曲线（每 12 步≈1小时记一点，控制内存）
            if step_i % 12 == 0:
                self.equity_curve.append((ts, self._equity(ts, idx_maps)))

        # 收尾：强制平掉所有活动仓位（按最后收盘价）
        for ex in list(self.active):
            df = self.data[ex.pair]
            last = df.iloc[-1]
            n_before = len(ex.trades)
            self.ctx.apply_close_fill(ex, float(last["close"]), ex.position.amount,
                                      float(last["ts"]), close_type=CloseType.FORCED,
                                      tag="end_of_data")
            self.done.append(ex)
        self.active.clear()
        return self.ctx

    # ------------------------------------------------------------------
    def _check_liquidation(self, ex: ExecutorBase, candle, ts: float) -> None:
        """逐仓强平近似：浮动亏损 ≥ 90% 占用保证金 → 按强平价强制平仓。

        高杠杆（如 5x）下这是真实风险，必须在回测中建模，否则会高估收益。
        """
        if not ex.is_active or ex.position.amount <= 0 or ex.margin_reserved <= 0:
            return
        long = ex.side.name == "BUY"
        # 浮动亏损 = -0.9×保证金 时的价格
        loss_px = 0.9 * ex.margin_reserved / ex.position.amount
        liq_px = (ex.position.entry_price - loss_px if long
                  else ex.position.entry_price + loss_px)
        worst = float(candle["low"]) if long else float(candle["high"])
        breached = worst <= liq_px if long else worst >= liq_px
        if breached:
            # 跳空则更差：开盘价已越过强平价按开盘价计
            open_px = float(candle["open"])
            fill = min(open_px, liq_px) if long else max(open_px, liq_px)
            self.ctx.apply_close_fill(ex, fill, ex.position.amount, float(ts),
                                      CloseType.FORCED, tag="liquidation")
            logger.warning("[%s] 强平 side=%s entry=%.4f liq=%.4f fill=%.4f lev=%.1fx",
                           ex.pair, ex.side.value, ex.position.entry_price,
                           liq_px, fill, ex.leverage)

    def _equity(self, ts: float, idx_maps: dict) -> float:
        eq = self.ctx.free_wallet
        for ex in self.active:
            df = self.data[ex.pair]
            i = idx_maps[ex.pair].get(ts)
            mark = float(df.iloc[i]["close"]) if i is not None else (
                ex.position.entry_price or 0.0)
            eq += ex.margin_reserved + ex.unrealized_pnl(mark)
        return eq
