"""Research-only D48 RiskCap with causal Turtle-style winner pyramiding.

The production signal, first-entry sizing, exits, protections, exposure gates and
2x leverage are inherited unchanged.  The only experiment factor is up to two
position increases after favourable moves of 0.5 and 1.0 initial-entry ATR.
"""

from __future__ import annotations

import logging
import sys
from datetime import datetime
from pathlib import Path

RESEARCH_STRATEGY_DIR = Path(__file__).resolve().parent
VOLATILITY_STRATEGY_DIR = RESEARCH_STRATEGY_DIR.parent / "volatility_breakout"
for strategy_dir in (RESEARCH_STRATEGY_DIR, VOLATILITY_STRATEGY_DIR):
    if str(strategy_dir) not in sys.path:
        sys.path.insert(0, str(strategy_dir))

from d48_turtle_pyramid_math import (  # noqa: E402
    planned_risk_fraction,
    pyramid_addition,
)
from v5_d48_riskcap import VolatilityBreakoutRiskCap  # noqa: E402

logger = logging.getLogger(__name__)


class VolatilityBreakoutRiskCapTurtlePyramidV1(VolatilityBreakoutRiskCap):
    """D48 RiskCap plus two 0.5-ATR-spaced increases in winning trades."""

    position_adjustment_enable = True
    max_entry_position_adjustment = 2

    pyramid_spacing_atr = 0.5
    pyramid_add_stake_fraction = 0.25
    pyramid_max_additions = 2
    pyramid_max_planned_risk_fraction = 0.01125

    _initial_atr_key = "d48_turtle_initial_atr_v1"
    _initial_stake_key = "d48_turtle_initial_stake_v1"

    def version(self) -> str:
        return "research-d48-riskcap-turtle-pyramid-v1"

    def _freeze_initial_state(self, pair: str, trade) -> bool:
        """Persist entry-time ATR and stake; fail closed if either is unavailable."""
        try:
            candle = self._last_candle(pair)
            atr = None if candle is None else float(candle.get("atr"))
            initial_stake = float(trade.stake_amount)
            if (
                atr is None
                or atr <= 0
                or atr != atr
                or initial_stake <= 0
                or initial_stake != initial_stake
            ):
                return False
            trade.set_custom_data(key=self._initial_atr_key, value=atr)
            trade.set_custom_data(key=self._initial_stake_key, value=initial_stake)
            return True
        except Exception as exc:
            logger.warning(
                "冻结 %s Turtle 初始 ATR/stake 失败(%s)，fail-closed 不加仓",
                pair,
                exc,
            )
            return False

    def order_filled(
        self,
        pair: str,
        trade,
        order,
        current_time: datetime,
        **kwargs,
    ) -> None:
        """Freeze causal state on the initial successful entry fill."""
        del current_time, kwargs
        try:
            if trade.get_custom_data(key=self._initial_atr_key) is not None:
                return
            entry_side = getattr(
                trade, "entry_side", "sell" if trade.is_short else "buy"
            )
            if getattr(order, "ft_order_side", None) != entry_side:
                return
            # Never mistake a later increase for the initial stake after a restart.
            if int(trade.nr_of_successful_entries) != 1:
                return
            self._freeze_initial_state(pair, trade)
        except Exception as exc:
            logger.warning(
                "处理 %s Turtle 初始成交失败(%s)，fail-closed 不加仓", pair, exc
            )

    def adjust_trade_position(
        self,
        trade,
        current_time: datetime,
        current_rate: float,
        current_profit: float,
        min_stake: float | None,
        max_stake: float,
        current_entry_rate: float,
        current_exit_rate: float,
        current_entry_profit: float,
        current_exit_profit: float,
        **kwargs,
    ):
        """Add 25% of initial stake at +0.5/+1.0 frozen ATR, at most twice."""
        del (
            current_time,
            current_entry_rate,
            current_exit_rate,
            current_entry_profit,
            current_exit_profit,
            kwargs,
        )
        if current_profit <= 0 or getattr(trade, "has_open_orders", False):
            return None

        successful_entries = int(trade.nr_of_successful_entries)
        if successful_entries > self.pyramid_max_additions:
            return None

        try:
            initial_atr = trade.get_custom_data(key=self._initial_atr_key)
            initial_stake = trade.get_custom_data(key=self._initial_stake_key)
            if initial_atr is None or initial_stake is None:
                # Backtest/live engines normally call order_filled.  This fallback is
                # only valid before any add and freezes the earliest available state.
                if successful_entries != 1 or not self._freeze_initial_state(
                    trade.pair, trade
                ):
                    return None
                initial_atr = trade.get_custom_data(key=self._initial_atr_key)
                initial_stake = trade.get_custom_data(key=self._initial_stake_key)

            decision = pyramid_addition(
                initial_stake=float(initial_stake),
                initial_rate=float(trade.open_rate),
                initial_atr=float(initial_atr),
                current_rate=float(current_rate),
                successful_entries=successful_entries,
                is_short=bool(trade.is_short),
                add_stake_fraction=self.pyramid_add_stake_fraction,
                spacing_atr=self.pyramid_spacing_atr,
                max_additions=self.pyramid_max_additions,
                min_stake=min_stake,
                max_stake=max_stake,
            )
        except Exception as exc:
            logger.warning(
                "计算 %s Turtle 加仓失败(%s)，fail-closed 不加仓", trade.pair, exc
            )
            return None

        if decision is None:
            return None

        risk_bound = planned_risk_fraction(
            initial_risk_fraction=self.risk_per_trade,
            add_stake_fraction=self.pyramid_add_stake_fraction,
            max_additions=self.pyramid_max_additions,
        )
        if risk_bound > self.pyramid_max_planned_risk_fraction + 1e-12:
            logger.error(
                "%s Turtle 计划风险 %.6f 超过上限 %.6f，fail-closed 不加仓",
                trade.pair,
                risk_bound,
                self.pyramid_max_planned_risk_fraction,
            )
            return None

        return decision.stake, f"turtle_add_{decision.addition_number}"
