"""趋势网格 Controller：EMA48 中心 + ATR 格距 + 趋势/市况过滤。

⚠️ 研究否决警告见 grid.py 头注与 README。过滤器参数严格对齐
trend-grid-strategy-iterations.md：
- ADX 12~35（排除无趋势与过强趋势）
- EMA200 方向/斜率过滤
- BTC 日线 200DMA：多头网格需收在上方，空头网格需在下方
- 12h 滚动窗口（5m×144 根）判断区间是否适合网格
- 同 pair 同时最多 1 个活动网格
"""
from __future__ import annotations

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

import pandas as pd

from .. import indicators as ind
from ..executors.base import ExecutorBase
from ..executors.grid import GridConfig, GridExecutor
from ..models import Side
from .base import ControllerBase

logger = logging.getLogger(__name__)


@dataclass
class TrendGridParams:
    ema_center: int = 48
    atr_period: int = 14
    ema_trend: int = 200
    adx_min: float = 12.0
    adx_max: float = 35.0
    range_window: int = 144          # 12h×5m 滚动区间窗口
    range_min_pct: float = 0.02      # 区间振幅下限（太平不做）
    range_max_pct: float = 0.15      # 区间振幅上限（太乱不做）
    step_atr_mult: float = 1.0       # 格距 = step_atr_mult × ATR
    n_levels: int = 4
    weights: tuple = (3, 2, 1, 1)
    bound_mult: float = 3.25
    margin_per_grid: float = 90.0    # 单网格保证金预算（quote）
    leverage: float = 2.0
    allowed_sides: tuple = ("long", "short")
    dual: bool = False               # 多空双开（中性网格）：跳过方向过滤，两侧同时挂
    breach_cooldown_s: float = 12 * 3600   # 越界后同 pair+方向冷却期


class TrendGridController(ControllerBase):
    warmup_candles = 210

    def __init__(self, pairs: list[str], params: Optional[TrendGridParams] = None,
                 btc_pair: str = "BTC/USDT:USDT"):
        super().__init__(pairs, btc_pair)
        self.p = params or TrendGridParams()
        self._breach_until: dict[tuple[str, str], float] = {}   # (pair, side) → 冷却截止 ts
        self.ctx = None

    def prepare_dataframe(self, pair: str, df: pd.DataFrame,
                          btc_1d: Optional[pd.DataFrame] = None) -> pd.DataFrame:
        df = df.copy()
        df["ema_center"] = ind.ema(df["close"], self.p.ema_center)
        df["atr"] = ind.atr(df, self.p.atr_period)
        df["ema_trend"] = ind.ema(df["close"], self.p.ema_trend)
        df["ema_slope"] = df["ema_trend"].diff(12)     # 1h 斜率（5m×12）
        df["adx"] = ind.adx(df, 14)
        df["range_pct"] = ind.rolling_range_pct(df, self.p.range_window)
        # BTC 日线 200DMA 闸门（ffill 到 5m，仅使用已收盘日线，无未来数据）
        if btc_1d is not None and not btc_1d.empty:
            b = btc_1d[["ts", "close"]].copy()
            b["btc_bull"] = b["close"] > ind.sma(b["close"], 200)
            df = pd.merge_asof(df.sort_values("ts"), b[["ts", "btc_bull"]].sort_values("ts"),
                               on="ts", direction="backward")
            df["btc_bull"] = df["btc_bull"].ffill().fillna(False).astype(bool)
        else:
            df["btc_bull"] = False
        self.dataframes[pair] = df
        return df

    def on_candle(self, pair: str, ts: float,
                  active: list[ExecutorBase]) -> list[ExecutorBase]:
        row = self.last_row(pair)
        if row is None:
            return []
        if pd.isna(row.get("atr")) or row["atr"] <= 0 or pd.isna(row.get("ema_slope")):
            return []

        # 市况过滤：ADX 区间 + 滚动振幅窗口（双开模式同样生效）
        if not (self.p.adx_min <= row["adx"] <= self.p.adx_max):
            return []
        rp = row.get("range_pct")
        if pd.isna(rp) or not (self.p.range_min_pct <= rp <= self.p.range_max_pct):
            return []

        if self.p.dual:
            # 中性网格：不做方向选择，两侧各最多一个活动网格。
            # 方向过滤（EMA200斜率/BTC200DMA）跳过——代价是单边趋势时
            # 总有一侧逆势累积库存，由越界退出兜底。
            sides = []
            if "long" in self.p.allowed_sides:
                sides.append(Side.BUY)
            if "short" in self.p.allowed_sides:
                sides.append(Side.SELL)
            return [self._spawn(pair, row, s) for s in sides
                    if not self._in_cooldown(pair, s, ts)
                    and not any(ex.pair == pair and ex.is_active and ex.side is s
                                for ex in active)]

        # 单向模式：同 pair 最多一个活动网格
        if any(ex.pair == pair and ex.is_active for ex in active):
            return []
        # 方向选择：价格在 EMA200 上方且斜率向上 → 做多网格；下方且向下 → 做空
        trend_up = row["close"] > row["ema_trend"] and row["ema_slope"] > 0
        trend_dn = row["close"] < row["ema_trend"] and row["ema_slope"] < 0
        side: Optional[Side] = None
        if trend_up and row["btc_bull"] and "long" in self.p.allowed_sides:
            side = Side.BUY
        elif trend_dn and not row["btc_bull"] and "short" in self.p.allowed_sides:
            side = Side.SELL
        if side is None or self._in_cooldown(pair, side, ts):
            return []
        return [self._spawn(pair, row, side)]

    # ---- 越界冷却：由引擎在平仓成交后回调 -----------------------------
    def _in_cooldown(self, pair: str, side: Side, ts: float) -> bool:
        key = (pair, "long" if side is Side.BUY else "short")
        return ts < self._breach_until.get(key, 0.0)

    def notify_trade_closed(self, record, ts: float) -> None:
        from ..models import CloseType
        if record.close_type is CloseType.BOUNDARY_BREACH:
            key = (record.pair, "long" if record.side is Side.BUY else "short")
            self._breach_until[key] = ts + self.p.breach_cooldown_s

    def _spawn(self, pair: str, row: pd.Series, side: Side) -> GridExecutor:
        cfg = GridConfig(
            center=float(row["ema_center"]), step=float(row["atr"]) * self.p.step_atr_mult,
            n_levels=self.p.n_levels, weights=self.p.weights,
            bound_mult=self.p.bound_mult, total_margin=self.p.margin_per_grid,
            leverage=self.p.leverage,
        )
        logger.info("[%s] 新建%s网格 center=%.2f step=%.2f lev=%.1fx",
                    pair, side.value, cfg.center, cfg.step, cfg.leverage)
        return GridExecutor(pair=pair, side=side, config=cfg, tag="trend_grid")
