"""网格执行器：实现 trend-grid-strategy-iterations.md 的设计原则。

⚠️ 注意：该策略族在项目回测中已被整体否决（最好变体验证段 +1.48%/PF 1.20，
费率 0.035%→0.05% 即转负）。本实现仅用于复现研究与 paper 验证，不建议实盘。

设计原则（与文档对齐）：
1. 空仓时由 Controller 以 EMA48 为中心、ATR 为格距创建；首单成交后【冻结】
   中心、格距与上下沿——持仓期间不把网格追着亏损方向平移。
2. 越过冻结边界（向上或向下）直接整体退出；超过 time_limit 未完成的也退出。
3. 固定总风险预算内的库存递减权重（默认 3:2:1:1），最多 n_levels 层，
   绝不无限马丁。
4. 每格成交后挂 +tp_mult*step 的反弹止盈（LIFO 语义：每层独立记账、独立释放）。
5. 若创建后 recenter_candles 内零成交，自我撤销，让 Controller 用新中心重建。
"""
from __future__ import annotations

import logging
from dataclasses import dataclass, field
from enum import Enum
from typing import Optional

import pandas as pd

from ..models import CloseType, Side
from .base import ExecutionContext, ExecutorBase

logger = logging.getLogger(__name__)


class LevelState(str, Enum):
    PENDING = "PENDING"    # 限价单已挂未成交
    OPEN = "OPEN"          # 已成交持有中，止盈单已挂
    CLOSED = "CLOSED"      # 该层已完成往返


@dataclass
class GridLevel:
    idx: int
    price: float           # 入场限价
    amount: float          # base 数量
    state: LevelState = LevelState.PENDING
    tp_price: float = 0.0
    entry_ts: float = 0.0


@dataclass
class GridConfig:
    center: float
    step: float
    n_levels: int = 4
    weights: tuple = (3, 2, 1, 1)     # 库存递减（研究最佳候选）
    bound_mult: float = 3.25          # 冻结边界 = center ± bound_mult*step
    tp_mult: float = 1.0              # 单层反弹止盈 = +tp_mult*step
    time_limit_s: float = 48 * 3600   # 48h 强制退出
    recenter_candles: int = 144       # 5m×144=12h 零成交则撤销重建
    total_margin: float = 100.0       # 本网格总保证金预算（quote）
    leverage: float = 2.0


class GridExecutor(ExecutorBase):
    """均值回归网格：long 在中心下方分层接跌；short 在中心上方分层接涨。"""

    def __init__(self, pair: str, side: Side, config: GridConfig, tag: str = ""):
        super().__init__(pair=pair, side=side, leverage=config.leverage, tag=tag)
        self.cfg = config
        self.levels: list[GridLevel] = self._build_levels()
        self.ceiling, self.floor = self._frozen_bounds()
        self.first_fill_ts: Optional[float] = None
        self.candles_waited: int = 0
        self._frozen = False   # 首单后冻结（边界在创建时已按文档设为固定值）

    # ---- 构建 ----------------------------------------------------------
    def _build_levels(self) -> list[GridLevel]:
        c, s = self.cfg.center, self.cfg.step
        w = self.cfg.weights[: self.cfg.n_levels]
        wsum = sum(w)
        levels = []
        for i in range(self.cfg.n_levels):
            # long：挂在中心下方；short：挂在中心上方
            price = c - (i + 1) * s if self.side is Side.BUY else c + (i + 1) * s
            margin_i = self.cfg.total_margin * w[i] / wsum
            notional_i = margin_i * self.cfg.leverage
            levels.append(GridLevel(idx=i, price=price, amount=notional_i / price))
        return levels

    def _frozen_bounds(self) -> tuple[float, float]:
        c, s, m = self.cfg.center, self.cfg.step, self.cfg.bound_mult
        return c + m * s, c - m * s

    # ---- 状态机 --------------------------------------------------------
    def on_candle(self, candle: pd.Series, ctx: ExecutionContext) -> None:
        ts = float(candle["ts"])
        if self.created_ts == 0.0:
            self.created_ts = ts

        # 1) 检查层级限价单成交（保守：价格触及即成交，不含 maker 排队优势）
        for lv in self.levels:
            if lv.state is LevelState.PENDING and self._limit_touched(candle, lv.price, entry=True):
                self._fill_entry(lv, lv.price, ts, ctx)
            elif lv.state is LevelState.OPEN and self._limit_touched(candle, lv.tp_price, entry=False):
                self._fill_tp(lv, lv.tp_price, ts, ctx)

        if not self.is_active:
            return

        n_filled = sum(1 for lv in self.levels if lv.state is not LevelState.PENDING)
        if n_filled == 0:
            # 2) 零成交超时 → 自我撤销（空仓移动窗口语义）
            self.candles_waited += 1
            if self.candles_waited >= self.cfg.recenter_candles:
                logger.info("[%s] 网格零成交超时撤销 center=%.2f", self.pair, self.cfg.center)
                self.finish(ts, CloseType.SIGNAL)
            return

        close = float(candle["close"])
        # 3) 越界退出：突破冻结边界，整体市价平掉
        if close > self.ceiling or close < self.floor:
            logger.info("[%s] 越界退出 close=%.2f bounds=[%.2f, %.2f]",
                        self.pair, close, self.floor, self.ceiling)
            self._close_all(close, ts, CloseType.BOUNDARY_BREACH, ctx)
            return
        # 4) 时间止损：首单后 48h 未完成 → 退出
        if self.first_fill_ts and ts - self.first_fill_ts >= self.cfg.time_limit_s:
            logger.info("[%s] 48h 时间止损", self.pair)
            self._close_all(close, ts, CloseType.TIME_LIMIT, ctx)

    # ---- 成交处理 ------------------------------------------------------
    def _limit_touched(self, candle: pd.Series, price: float, entry: bool) -> bool:
        if self.side is Side.BUY:
            # 买入（开仓接跌 / 止盈卖在上方？）——long 网格：entry=买，tp=卖
            return float(candle["low"]) <= price if entry else float(candle["high"]) >= price
        return float(candle["high"]) >= price if entry else float(candle["low"]) <= price

    def _fill_entry(self, lv: GridLevel, price: float, ts: float, ctx: ExecutionContext) -> None:
        if not ctx.apply_open_fill(self, price, lv.amount, ts):
            return   # 保证金不足被交易所拒单，保持 PENDING
        lv.state = LevelState.OPEN
        lv.entry_ts = ts
        lv.tp_price = (price + self.cfg.tp_mult * self.cfg.step if self.side is Side.BUY
                       else price - self.cfg.tp_mult * self.cfg.step)
        if self.first_fill_ts is None:
            self.first_fill_ts = ts
            self._frozen = True   # 首单后冻结（bounds 已固定，此处仅记录语义）
        logger.debug("[%s] 层%d 成交 @%.2f amount=%.6f", self.pair, lv.idx, price, lv.amount)

    def _fill_tp(self, lv: GridLevel, price: float, ts: float, ctx: ExecutionContext) -> None:
        lv.state = LevelState.CLOSED
        ctx.apply_close_fill(self, price, lv.amount, ts, CloseType.TAKE_PROFIT,
                             tag=f"grid_l{lv.idx}")
        logger.debug("[%s] 层%d 止盈 @%.2f", self.pair, lv.idx, price)

    def _close_all(self, price: float, ts: float, close_type: CloseType,
                   ctx: ExecutionContext) -> None:
        for lv in self.levels:
            if lv.state is LevelState.OPEN:
                lv.state = LevelState.CLOSED
                ctx.apply_close_fill(self, price, lv.amount, ts, close_type,
                                     tag=f"grid_l{lv.idx}")
        # 未成交层直接作废；position 归零时 apply_close_fill 已调用 finish
        self.finish(ts, close_type)

    # ---- 状态 ----------------------------------------------------------
    def status(self) -> dict:
        return {
            "pair": self.pair, "side": self.side.value, "center": self.cfg.center,
            "step": self.cfg.step, "bounds": [self.floor, self.ceiling],
            "levels": [{"i": lv.idx, "px": lv.price, "st": lv.state.value} for lv in self.levels],
            "net_pnl": round(self.net_pnl, 4), "active": self.is_active,
        }
