"""VolatilityBreakout Controller：1:1 移植生产策略 v5_d48（5.3-d48-shortgate）。

入场：Donchian48 突破 + ATR 扩张 + ADX>25 + EMA200 顺势；
     开多需 BTC 日线 > 200DMA，开空需 BTC 日线 ≤ 200DMA（对称闸门）。
出场：交给 PositionExecutor 三重屏障（硬止损 -5% / trailing 3% 激活 5%）。
定仓：ATR 风险定仓（每笔风险 0.75% 钱包、单笔保证金 ≤18% 钱包）。
敞口：同方向 ≤2 仓；BTC/ETH/SOL/BNB 高相关组同方向 ≤2 仓。
熔断：2h 内 2 次止损 → 全局暂停开仓 1h（StoplossGuard 等价物）。
"""
from __future__ import annotations

import logging
import time
from dataclasses import dataclass
from typing import Optional

import pandas as pd

from .. import indicators as ind
from ..executors.base import ExecutorBase
from ..executors.position import PositionExecutor, TripleBarrier
from ..models import CloseType, Side
from .base import ControllerBase

logger = logging.getLogger(__name__)

MAJOR_GROUP = frozenset(
    {"BTC/USDT:USDT", "ETH/USDT:USDT", "SOL/USDT:USDT", "BNB/USDT:USDT"}
)


@dataclass
class VolBreakoutParams:
    donchian_period: int = 48
    atr_period: int = 14
    vol_baseline_period: int = 50
    adx_threshold: float = 25.0
    risk_per_trade: float = 0.0075
    max_stake_wallet_fraction: float = 0.18
    atr_stop_multiple: float = 3.0
    leverage: float = 2.0
    max_same_side: int = 2
    max_same_side_major: int = 2
    stoploss_pct: float = 0.05
    trailing_activation: float = 0.05
    trailing_distance: float = 0.03
    # StoplossGuard：lookback 2h 内 2 次止损 → 暂停 1h
    guard_lookback_s: float = 2 * 3600
    guard_trade_limit: int = 2
    guard_pause_s: float = 3600


class VolatilityBreakoutController(ControllerBase):
    warmup_candles = 210

    def __init__(self, pairs: list[str], params: Optional[VolBreakoutParams] = None,
                 btc_pair: str = "BTC/USDT:USDT"):
        super().__init__(pairs, btc_pair)
        self.p = params or VolBreakoutParams()
        self._pause_until: float = 0.0
        # 由引擎注入：提供钱包与成交日志（风险定仓与熔断用）
        self.ctx = None

    def prepare_dataframe(self, pair: str, df: pd.DataFrame,
                          btc_1d: Optional[pd.DataFrame] = None) -> pd.DataFrame:
        df = df.copy()
        df["atr"] = ind.atr(df, self.p.atr_period)
        df["atr_baseline"] = df["atr"].rolling(self.p.vol_baseline_period).mean()
        df["vol_expansion"] = df["atr"] > df["atr_baseline"]
        df["donchian_high"] = ind.donchian_high(df, self.p.donchian_period)
        df["donchian_low"] = ind.donchian_low(df, self.p.donchian_period)
        df["adx"] = ind.adx(df, 14)
        df["ema_trend"] = ind.ema(df["close"], 200)
        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]:
        if ts < self._pause_until:
            return []
        if any(ex.pair == pair and ex.is_active for ex in active):
            return []   # 同 pair 不叠加仓位
        row = self.last_row(pair)
        if row is None or pd.isna(row.get("donchian_high")) or pd.isna(row.get("atr")):
            return []
        if row["volume"] <= 0:
            return []

        adx_ok = row["adx"] > self.p.adx_threshold
        trend_up = row["close"] > row["ema_trend"]
        side: Optional[Side] = None
        if (row["close"] > row["donchian_high"] and row["vol_expansion"]
                and adx_ok and trend_up and row["btc_bull"]):
            side = Side.BUY
        elif (row["close"] < row["donchian_low"] and row["vol_expansion"]
                and adx_ok and not trend_up and not row["btc_bull"]):
            side = Side.SELL
        if side is None:
            return []

        if not self._passes_exposure(side, active):
            return []

        stake = self._stake_amount(float(row["atr"]), float(row["close"]))
        if stake <= 0:
            return []

        barrier = TripleBarrier(
            stop_loss_pct=self.p.stoploss_pct,
            trailing_activation_pct=self.p.trailing_activation,
            trailing_distance_pct=self.p.trailing_distance,
        )
        ex = PositionExecutor(pair=pair, side=side, barrier=barrier,
                              leverage=self.p.leverage, tag="vol_breakout")
        ex.planned_stake = stake          # 引擎入场时使用的保证金
        ex.entry_at_next_open = True      # 信号在收盘产生，次根开盘成交
        logger.info("[%s] %s突破信号 close=%.2f stake=%.2f", pair, side.value,
                    row["close"], stake)
        return [ex]

    # ---- 风控（对齐 v5_d48 回调语义，fail-closed） ----------------------
    def _passes_exposure(self, side: Side, active: list[ExecutorBase]) -> bool:
        same_side = [ex for ex in active
                     if ex.is_active and ex.position.amount > 0
                     and ((ex.side is Side.SELL) == (side is Side.SELL))]
        if len(same_side) >= self.p.max_same_side:
            logger.info("entry_blocked reason=same_side_limit side=%s", side.value)
            return False
        group = [ex for ex in same_side if ex.pair in MAJOR_GROUP]
        if len(group) >= self.p.max_same_side_major:
            logger.info("entry_blocked reason=major_group_limit side=%s", side.value)
            return False
        return True

    def _stake_amount(self, atr: float, price: float) -> float:
        if self.ctx is None or atr <= 0 or price <= 0:
            return 0.0
        wallet = self.ctx.free_wallet
        price_risk = (atr * self.p.atr_stop_multiple) / price
        leveraged_risk = price_risk * self.p.leverage
        if leveraged_risk <= 0:
            return 0.0
        risk_stake = (wallet * self.p.risk_per_trade) / leveraged_risk
        wallet_cap = wallet * self.p.max_stake_wallet_fraction
        return min(risk_stake, wallet_cap)

    # ---- StoplossGuard：由引擎在平仓成交后调用 --------------------------
    def notify_trade_closed(self, record, ts: float) -> None:
        if record.close_type is not CloseType.STOP_LOSS or self.ctx is None:
            return
        recent = [t for t in self.ctx.trade_log
                  if t.close_type is CloseType.STOP_LOSS
                  and ts - t.exit_ts <= self.p.guard_lookback_s]
        if len(recent) >= self.p.guard_trade_limit:
            self._pause_until = ts + self.p.guard_pause_s
            logger.warning("StoplossGuard 触发：暂停开仓至 %s",
                           time.strftime("%H:%M:%S", time.gmtime(self._pause_until)))
