"""Independent research-only 1h multi-speed EWMAC trend strategy.

This is not a D48 filter and does not inherit D48 signals, gates or exits.  It
uses completed 1h candles, equal-weights capped 8/32, 16/64 and 32/128 EWMAC
forecasts, enters on +/-5 crossings and exits on a zero-axis crossing.
"""

from __future__ import annotations

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

import talib.abstract as ta
from pandas import DataFrame

from freqtrade.strategy import IStrategy

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 ewmac_1h_math import hysteresis_signals, multi_speed_ewmac  # noqa: E402
from risk_cap_math import risk_capped_stake  # noqa: E402

logger = logging.getLogger(__name__)


class MultiSpeedEwmac1hV1(IStrategy):
    """Standalone 1h EWMAC sleeve with risk-capped 2x futures sizing."""

    INTERFACE_VERSION = 3
    timeframe = "1h"
    process_only_new_candles = True
    can_short = True
    startup_candle_count = 256

    minimal_roi = {"0": 100}
    stoploss = -0.06
    trailing_stop = False
    use_exit_signal = True
    exit_profit_only = False
    ignore_roi_if_entry_signal = False

    ewmac_vol_span_hours = 72
    ewmac_forecast_cap = 20.0
    ewmac_entry_threshold = 5.0
    ewmac_exit_threshold = 0.0

    risk_per_trade = 0.005
    atr_period = 24
    atr_stop_multiple = 3.0
    max_stake_wallet_fraction = 0.15

    def version(self) -> str:
        return "research-1h-ewmac-8-32-16-64-32-128-v1"

    @property
    def protections(self):
        # 1h: two stoplosses in 24h pause all entries for 6h.
        return [
            {
                "method": "StoplossGuard",
                "lookback_period_candles": 24,
                "trade_limit": 2,
                "stop_duration_candles": 6,
                "required_profit": 0.0,
                "only_per_pair": False,
                "only_per_side": False,
            }
        ]

    def populate_indicators(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        del metadata
        forecast = multi_speed_ewmac(
            dataframe["close"],
            vol_span_hours=self.ewmac_vol_span_hours,
            forecast_cap=self.ewmac_forecast_cap,
        )
        for column in forecast.columns:
            dataframe[column] = forecast[column]
        dataframe["atr"] = ta.ATR(dataframe, timeperiod=self.atr_period)
        return dataframe

    def populate_entry_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        del metadata
        signals = hysteresis_signals(
            dataframe["ewmac_forecast"],
            entry_threshold=self.ewmac_entry_threshold,
            exit_threshold=self.ewmac_exit_threshold,
        )
        valid_market = (dataframe["volume"] > 0) & dataframe["atr"].notna()
        dataframe.loc[signals["enter_long"] & valid_market, "enter_long"] = 1
        dataframe.loc[signals["enter_short"] & valid_market, "enter_short"] = 1
        return dataframe

    def populate_exit_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        del metadata
        signals = hysteresis_signals(
            dataframe["ewmac_forecast"],
            entry_threshold=self.ewmac_entry_threshold,
            exit_threshold=self.ewmac_exit_threshold,
        )
        valid_market = dataframe["volume"] > 0
        dataframe.loc[signals["exit_long"] & valid_market, "exit_long"] = 1
        dataframe.loc[signals["exit_short"] & valid_market, "exit_short"] = 1
        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()

    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:
        del current_time, proposed_stake, entry_tag, side, kwargs
        try:
            candle = self._last_candle(pair)
            atr = None if candle is None else float(candle.get("atr"))
            if atr is None or atr <= 0 or atr != atr or current_rate <= 0:
                return 0
            stake = risk_capped_stake(
                wallet=float(self.wallets.get_total_stake_amount()),
                risk_fraction=self.risk_per_trade,
                atr_loss_fraction=(atr * self.atr_stop_multiple / current_rate)
                * leverage,
                hard_stop_fraction=abs(self.stoploss),
                wallet_cap_fraction=self.max_stake_wallet_fraction,
                max_stake=max_stake,
            )
        except Exception as exc:
            logger.warning("计算 %s EWMAC 风险仓位失败(%s)，fail-closed", pair, exc)
            return 0
        if min_stake is not None and stake < min_stake:
            return 0
        return stake

    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:
        del (
            pair,
            current_time,
            current_rate,
            proposed_leverage,
            entry_tag,
            side,
            kwargs,
        )
        return min(2.0, max_leverage)
