"""波动率突破策略。

当前部署版继承了原 V5 逻辑：
- Donchian48 突破 + ATR 扩张，过滤盘整假突破。
- ADX>25 + EMA200 顺势过滤，只做趋势方向。
- 开多需 BTC 日线收盘 > 200DMA，并叠加 regime.json 实时开多总闸。
- 开空需 BTC 日线收盘 ≤ 200DMA（对称闸门）：2026 熊市样本回测与无闸门一字不差，
  作用是市况翻多后不再追空回调；2020~2025 四个牛市窗口回测合计亏损约减六成。
- 不设 ROI 止盈上限，移动止损 3%（激活线 5%）让利润奔跑。
- ATR 风险定仓，减少 2x 满仓导致的账户回撤。
- 同方向最多 2 个未平仓，BTC/ETH/SOL/BNB 高相关组同方向最多 2 个未平仓。
- 2 小时内两次止损后全局暂停开仓 1 小时，避免异常行情连续追单。

注意：config 里的 timeframe 会覆盖策略类属性，两处需保持一致（都为 5m）。
"""

import json
import logging
from datetime import datetime
from pathlib import Path

import talib.abstract as ta
from pandas import DataFrame

from freqtrade.persistence import Trade
from freqtrade.strategy import IStrategy, merge_informative_pair

from entry_block_observability import (
    candidate_signal_snapshot,
    occupied_trade_snapshots,
    replacement_shadow_decision,
)

logger = logging.getLogger(__name__)

# v5_d48.py lives at user_data/strategies/volatility_breakout/.
REGIME_FILE = Path(__file__).resolve().parents[2] / "regime.json"
REGIME_MAX_AGE_H = 3
REGIME_MIN_SOURCES = 6

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


class VolatilityBreakout(IStrategy):
    INTERFACE_VERSION = 3

    def version(self) -> str:
        return "5.3-d48-shortgate"

    timeframe = "5m"
    can_short = True

    # 不设 ROI 止盈上限，出场交给移动止损。
    minimal_roi = {"0": 100}
    stoploss = -0.05

    trailing_stop = True
    trailing_stop_positive = 0.03
    trailing_stop_positive_offset = 0.05
    trailing_only_offset_is_reached = True

    startup_candle_count = 210

    donchian_period = 48
    atr_period = 14
    vol_baseline_period = 50
    adx_threshold = 25
    risk_per_trade = 0.0075
    max_stake_wallet_fraction = 0.18
    atr_stop_multiple = 3.0
    max_same_side_open_trades = 2
    max_same_side_major_group_trades = 2
    entry_block_log_cooldown_s = 300
    # 仅用于 shadow 观测，不改变 max_same_side_open_trades 或真实交易动作。
    replacement_shadow_min_breakout_atr = 0.10
    replacement_shadow_min_adx = 30.0
    replacement_shadow_max_weakest_r = 0.0

    @property
    def protections(self):
        # 5m: 24 根=2h；两次止损后 12 根=1h 内不再开新仓。
        # 全币种、全方向共享熔断，定位为异常行情保护而非信号过滤器。
        return [
            {
                "method": "StoplossGuard",
                "lookback_period_candles": 24,
                "trade_limit": 2,
                "stop_duration_candles": 12,
                "required_profit": 0.0,
                "only_per_pair": False,
                "only_per_side": False,
            }
        ]

    def informative_pairs(self):
        return [(BTC_PAIR, "1d")]

    def populate_indicators(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        dataframe["atr"] = ta.ATR(dataframe, timeperiod=self.atr_period)
        dataframe["atr_baseline"] = dataframe["atr"].rolling(self.vol_baseline_period).mean()
        dataframe["vol_expansion"] = dataframe["atr"] > dataframe["atr_baseline"]

        # shift(1) 排除当前 K 线，避免用未来数据。
        dataframe["donchian_high"] = (
            dataframe["high"].rolling(self.donchian_period).max().shift(1)
        )
        dataframe["donchian_low"] = (
            dataframe["low"].rolling(self.donchian_period).min().shift(1)
        )
        dataframe["adx"] = ta.ADX(dataframe, timeperiod=14)
        dataframe["ema_trend"] = ta.EMA(dataframe, timeperiod=200)

        # BTC 日线 200DMA 开多闸门（merge_informative_pair 已处理未来数据对齐）。
        btc_1d = self.dp.get_pair_dataframe(BTC_PAIR, "1d")
        btc_1d["btc_bull"] = btc_1d["close"] > ta.SMA(btc_1d, timeperiod=200)
        dataframe = merge_informative_pair(
            dataframe, btc_1d[["date", "btc_bull"]], self.timeframe, "1d", ffill=True
        )
        dataframe["btc_bull_1d"] = dataframe["btc_bull_1d"].fillna(False).astype(bool)
        # merge 产生的 date_1d 只是元数据且不参与信号；删除它，避免递归审计把
        # 不同截断窗口的末端 informative 日期误报为指标未来偏差。
        dataframe.drop(columns=["date_1d"], inplace=True, errors="ignore")
        return dataframe

    def populate_entry_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        adx_ok = dataframe["adx"] > self.adx_threshold
        trend_up = dataframe["close"] > dataframe["ema_trend"]

        dataframe.loc[
            (dataframe["close"] > dataframe["donchian_high"])
            & dataframe["vol_expansion"]
            & adx_ok
            & trend_up
            & dataframe["btc_bull_1d"]
            & (dataframe["volume"] > 0),
            "enter_long",
        ] = 1

        dataframe.loc[
            (dataframe["close"] < dataframe["donchian_low"])
            & dataframe["vol_expansion"]
            & adx_ok
            & ~trend_up
            & ~dataframe["btc_bull_1d"]
            & (dataframe["volume"] > 0),
            "enter_short",
        ] = 1
        return dataframe

    def populate_exit_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        # 出场交给止损 / 移动止损。
        return dataframe

    def _last_candle(self, pair: str):
        dataframe, _ = self.dp.get_analyzed_dataframe(pair=pair, timeframe=self.timeframe)
        if dataframe.empty:
            return None
        return dataframe.iloc[-1].squeeze()

    @staticmethod
    def _trade_side(trade) -> str:
        return "short" if trade.is_short else "long"

    def _log_entry_blocked(
        self,
        *,
        reason: str,
        pair: str,
        side: str,
        current_time: datetime,
        open_trades: list,
    ) -> None:
        # confirm_trade_entry may be retried every bot loop while the same candle's
        # signal remains active.  Keep the operational evidence without flooding logs.
        key = (reason, pair, side)
        last_logged = getattr(self, "_entry_block_log_times", {}).get(key)
        now_ts = current_time.timestamp()
        if last_logged is not None and now_ts - last_logged < self.entry_block_log_cooldown_s:
            return
        if not hasattr(self, "_entry_block_log_times"):
            self._entry_block_log_times = {}
        self._entry_block_log_times[key] = now_ts
        occupied = ",".join(
            sorted(f"{trade.pair}#{trade.id}" for trade in open_trades)
        ) or "-"
        candidate = {}
        occupied_state = []
        shadow = {
            "eligible": False,
            "reason": "snapshot_error",
            "weakest": None,
        }
        try:
            candle = self._last_candle(pair)
            candidate = candidate_signal_snapshot(candle if candle is not None else {}, side)

            def mark_lookup(occupied_pair: str) -> float | None:
                occupied_candle = self._last_candle(occupied_pair)
                return None if occupied_candle is None else occupied_candle.get("close")

            occupied_state = occupied_trade_snapshots(
                open_trades,
                current_time=current_time,
                mark_lookup=mark_lookup,
                initial_risk=abs(self.stoploss),
            )
            shadow = replacement_shadow_decision(
                candidate,
                occupied_state,
                min_breakout_atr=self.replacement_shadow_min_breakout_atr,
                min_adx=self.replacement_shadow_min_adx,
                max_weakest_r=self.replacement_shadow_max_weakest_r,
            )
        except Exception as exc:
            # 观测增强不得改变原有 fail-closed 风控判定。
            shadow["error"] = type(exc).__name__
        logger.info(
            "entry_blocked reason=%s pair=%s side=%s occupied=%s "
            "candidate=%s occupied_state=%s replacement_shadow=%s",
            reason,
            pair,
            side,
            occupied,
            json.dumps(candidate, separators=(",", ":"), sort_keys=True),
            json.dumps(occupied_state, separators=(",", ":"), sort_keys=True),
            json.dumps(shadow, separators=(",", ":"), sort_keys=True),
        )

    def _log_entry_approved(
        self,
        *,
        pair: str,
        side: str,
        current_time: datetime,
        gates: dict,
    ) -> None:
        """Freeze decision-time inputs; this log is observational and fail-open."""
        try:
            if self.config["runmode"].value not in ("live", "dry_run"):
                return
            key = (pair, side)
            now_ts = current_time.timestamp()
            last_logged = getattr(self, "_entry_approved_log_times", {}).get(key)
            if last_logged is not None and now_ts - last_logged < self.entry_block_log_cooldown_s:
                return
            if not hasattr(self, "_entry_approved_log_times"):
                self._entry_approved_log_times = {}
            self._entry_approved_log_times[key] = now_ts

            candle = self._last_candle(pair)
            candidate = candidate_signal_snapshot(candle if candle is not None else {}, side)
            open_trades = [
                trade
                for trade in Trade.get_open_trades()
                if self._trade_side(trade) == side
            ]

            def mark_lookup(occupied_pair: str) -> float | None:
                occupied_candle = self._last_candle(occupied_pair)
                return None if occupied_candle is None else occupied_candle.get("close")

            occupied_state = occupied_trade_snapshots(
                open_trades,
                current_time=current_time,
                mark_lookup=mark_lookup,
                initial_risk=abs(self.stoploss),
            )
            logger.info(
                "entry_approved pair=%s side=%s candidate=%s occupied_state=%s gates=%s",
                pair,
                side,
                json.dumps(candidate, separators=(",", ":"), sort_keys=True),
                json.dumps(occupied_state, separators=(",", ":"), sort_keys=True),
                json.dumps(gates, separators=(",", ":"), sort_keys=True),
            )
        except Exception as exc:
            # 审计快照缺失不能反向改变已通过的真实交易决策。
            logger.warning("记录 %s %s entry_approved 快照失败(%s)", pair, side, exc)

    def _passes_exposure_limits(
        self, pair: str, side: str, current_time: datetime
    ) -> bool:
        try:
            same_side_trades = [
                trade
                for trade in Trade.get_open_trades()
                if self._trade_side(trade) == side
            ]
        except Exception as e:
            logger.error("读取未平仓失败(%s)，fail-closed 拦下 %s %s", e, pair, side)
            return False

        if len(same_side_trades) >= self.max_same_side_open_trades:
            self._log_entry_blocked(
                reason="same_side_limit",
                pair=pair,
                side=side,
                current_time=current_time,
                open_trades=same_side_trades,
            )
            return False

        if pair in MAJOR_CORRELATED_GROUP:
            group_trades = [
                trade for trade in same_side_trades if trade.pair in MAJOR_CORRELATED_GROUP
            ]
            if len(group_trades) >= self.max_same_side_major_group_trades:
                self._log_entry_blocked(
                    reason="major_group_limit",
                    pair=pair,
                    side=side,
                    current_time=current_time,
                    open_trades=group_trades,
                )
                return False

        return True

    def custom_stake_amount(
        self,
        pair: str,
        current_time: datetime,
        current_rate: float,
        proposed_stake: float,
        min_stake: float | None,
        max_stake: float,
        leverage: float,
        entry_tag: str | None,
        side: str,
        **kwargs,
    ) -> float | None:
        try:
            candle = self._last_candle(pair)
        except Exception as e:
            logger.error("读取 %s 分析 K 线失败(%s)，fail-closed 跳过开仓", pair, e)
            return 0
        if candle is None:
            logger.warning("%s 无可用分析 K 线，fail-closed 跳过开仓", pair)
            return 0

        atr = candle.get("atr")
        if atr is None or atr <= 0 or atr != atr or current_rate <= 0:
            logger.warning("%s ATR/价格无效，fail-closed 跳过开仓", pair)
            return 0

        price_risk = (atr * self.atr_stop_multiple) / current_rate
        leveraged_risk = price_risk * leverage
        if leveraged_risk <= 0:
            logger.warning("%s 杠杆风险无效，fail-closed 跳过开仓", pair)
            return 0

        try:
            wallet = self.wallets.get_total_stake_amount()
        except Exception as e:
            logger.error("读取钱包失败(%s)，fail-closed 跳过 %s 开仓", e, pair)
            return 0

        risk_budget = wallet * self.risk_per_trade
        risk_stake = risk_budget / leveraged_risk
        wallet_cap = wallet * self.max_stake_wallet_fraction
        stake = min(risk_stake, wallet_cap, max_stake)

        if min_stake is not None and stake < min_stake:
            return 0
        return stake

    def confirm_trade_entry(
        self,
        pair: str,
        order_type: str,
        amount: float,
        rate: float,
        time_in_force: str,
        current_time: datetime,
        entry_tag: str | None,
        side: str,
        **kwargs,
    ) -> bool:
        if not self._passes_exposure_limits(pair, side, current_time):
            return False

        # regime.json 实时开多总闸：只闸开多；文件缺失/陈旧/降级/异常一律 fail-closed。
        # 与 btc_bull_1d 回测闸门互补：前者实时多因子，后者可回测。
        if side != "long":
            self._log_entry_approved(
                pair=pair,
                side=side,
                current_time=current_time,
                gates={"exposure": "passed", "btc_200dma": "passed", "regime": "not_applicable"},
            )
            return True
        # regime.json 是“当前时刻”的外部快照，不能用于历史回测。官方建议只在
        # live/dry_run 访问远端/实时数据；历史模式由 btc_bull_1d 代理闸门负责。
        if self.config["runmode"].value not in ("live", "dry_run"):
            return True
        regime_snapshot = {}
        try:
            data = json.loads(REGIME_FILE.read_text())
            regime_snapshot = {
                "regime": data.get("regime"),
                "score": data.get("score"),
                "allow_long": data.get("allow_long"),
                "sources_ok": data.get("sources_ok"),
                "sources_failed": data.get("sources_failed"),
                "ts": data.get("ts"),
            }
            age_h = (
                current_time - datetime.fromisoformat(data["ts"])
            ).total_seconds() / 3600
            if age_h < 0 or age_h > REGIME_MAX_AGE_H:
                logger.warning(
                    "regime.json 时间异常/已过期 %.1fh，fail-closed 拦下 %s 开多",
                    age_h, pair,
                )
                return False
            if data.get("sources_ok", 0) < REGIME_MIN_SOURCES or data.get(
                "sources_failed", 0
            ):
                logger.warning(
                    "regime.json 数据源不完整 %s/%d，fail-closed 拦下 %s 开多",
                    data.get("sources_ok", 0), REGIME_MIN_SOURCES, pair,
                )
                return False
            if not data.get("allow_long", True):
                logger.info(
                    "regime=%s score=%+d → 拦下 %s 开多",
                    data.get("regime"), data.get("score", 0), pair,
                )
                return False
        except Exception as e:
            logger.error("读取 regime.json 失败(%s)，fail-closed 拦下 %s 开多", e, pair)
            return False
        self._log_entry_approved(
            pair=pair,
            side=side,
            current_time=current_time,
            gates={
                "exposure": "passed",
                "btc_200dma": "passed",
                "regime": "passed",
                "regime_snapshot": regime_snapshot,
            },
        )
        return True

    def leverage(
        self,
        pair: str,
        current_time: datetime,
        current_rate: float,
        proposed_leverage: float,
        max_leverage: float,
        entry_tag: str | None,
        side: str,
        **kwargs,
    ) -> float:
        return 2.0
