"""Frozen research candidates inspired by open-source trend systems.

These classes are intentionally kept outside the production strategy directory.
Every class inherits the deployed RiskCap implementation and changes exactly one
entry factor so paired backtests remain attributable.
"""

import talib.abstract as ta

from freqtrade.strategy import merge_informative_pair

from v5_d48_riskcap import VolatilityBreakoutRiskCap


class VolatilityBreakoutRiskCapMatchedBaseline(VolatilityBreakoutRiskCap):
    """No-op matched baseline loaded through the same research strategy path."""

    def version(self) -> str:
        return "research-20260810-d48-riskcap-matched-baseline"


class VolatilityBreakoutRiskCapD36(VolatilityBreakoutRiskCap):
    """Single-factor candidate: shorten Donchian entry lookback from 48 to 36."""

    donchian_period = 36

    def version(self) -> str:
        return "research-20260810-d36-riskcap-v1"


class VolatilityBreakoutRiskCapD42(VolatilityBreakoutRiskCap):
    """Pre-registered fallback: shorten Donchian entry lookback from 48 to 42."""

    donchian_period = 42

    def version(self) -> str:
        return "research-20260810-d42-riskcap-v1"


class _PairDailyBase(VolatilityBreakoutRiskCap):
    """Add causal, completed-daily-bar columns for the signal pair itself."""

    def informative_pairs(self):
        pairs = set(super().informative_pairs())
        pairs.update((pair, "1d") for pair in self.dp.current_whitelist())
        return sorted(pairs)

    def _merge_pair_daily(self, dataframe, metadata):
        pair_daily = self.dp.get_pair_dataframe(metadata["pair"], "1d").copy()
        pair_daily["pair_daily_adx"] = ta.ADX(pair_daily, timeperiod=14)
        pair_daily["pair_daily_ema50"] = ta.EMA(pair_daily, timeperiod=50)
        pair_daily["pair_daily_ema200"] = ta.EMA(pair_daily, timeperiod=200)
        pair_daily["pair_daily_ema50_slope5"] = (
            pair_daily["pair_daily_ema50"]
            - pair_daily["pair_daily_ema50"].shift(5)
        )
        return merge_informative_pair(
            dataframe,
            pair_daily[
                [
                    "date",
                    "close",
                    "pair_daily_adx",
                    "pair_daily_ema50",
                    "pair_daily_ema200",
                    "pair_daily_ema50_slope5",
                ]
            ],
            self.timeframe,
            "1d",
            ffill=True,
        )

    def populate_indicators(self, dataframe, metadata):
        dataframe = super().populate_indicators(dataframe, metadata)
        dataframe = self._merge_pair_daily(dataframe, metadata)
        dataframe.drop(columns=["date_1d"], inplace=True, errors="ignore")
        return dataframe


class VolatilityBreakoutRiskCapPairDailyAdx20(_PairDailyBase):
    """Single-factor candidate: require the pair's completed daily ADX > 20."""

    def version(self) -> str:
        return "research-20260810-pair-daily-adx20-v1"

    def populate_entry_trend(self, dataframe, metadata):
        dataframe = super().populate_entry_trend(dataframe, metadata)
        daily_regime = dataframe["pair_daily_adx_1d"] > 20.0
        dataframe.loc[~daily_regime, ["enter_long", "enter_short"]] = 0
        return dataframe


class VolatilityBreakoutRiskCapPairDailyNoOp(_PairDailyBase):
    """Implementation-isolation control: merge pair daily data but change no signal."""

    def version(self) -> str:
        return "research-20260810-pair-daily-noop-v1"


class VolatilityBreakoutRiskCapPairDailyEma50Direction(_PairDailyBase):
    """Gate mature pairs by completed daily close versus EMA50.

    Before EMA50 is mature the baseline signal passes through unchanged, making
    listing age explicit instead of silently excluding recently listed pairs.
    """

    def version(self) -> str:
        return "research-20260810-pair-daily-ema50-direction-v1"

    def populate_entry_trend(self, dataframe, metadata):
        dataframe = super().populate_entry_trend(dataframe, metadata)
        ready = dataframe["pair_daily_ema50_1d"].notna()
        long_mismatch = ready & (
            dataframe["close_1d"] <= dataframe["pair_daily_ema50_1d"]
        )
        short_mismatch = ready & (
            dataframe["close_1d"] >= dataframe["pair_daily_ema50_1d"]
        )
        dataframe.loc[long_mismatch, "enter_long"] = 0
        dataframe.loc[short_mismatch, "enter_short"] = 0
        return dataframe


class VolatilityBreakoutRiskCapAdxScaledRisk(VolatilityBreakoutRiskCap):
    """Scale stake continuously with entry ADX, capped around baseline risk.

    The frozen multiplier is 0.75 at ADX 25, 1.00 at ADX 35 and 1.25 at
    ADX 45+, following the forecast-scaling/capping pattern used by systematic
    trend portfolios. Entry and exit signals are unchanged.
    """

    def version(self) -> str:
        return "research-20260810-adx-scaled-risk-v1"

    def custom_stake_amount(
        self,
        pair,
        current_time,
        current_rate,
        proposed_stake,
        min_stake,
        max_stake,
        leverage,
        entry_tag,
        side,
        **kwargs,
    ):
        stake = super().custom_stake_amount(
            pair=pair,
            current_time=current_time,
            current_rate=current_rate,
            proposed_stake=proposed_stake,
            min_stake=min_stake,
            max_stake=max_stake,
            leverage=leverage,
            entry_tag=entry_tag,
            side=side,
            **kwargs,
        )
        if not stake:
            return stake
        candle = self._last_candle(pair)
        adx = None if candle is None else candle.get("adx")
        if adx is None or adx != adx:
            return 0
        multiplier = min(1.25, max(0.75, 0.75 + (float(adx) - 25.0) / 40.0))
        scaled = min(float(stake) * multiplier, max_stake)
        if min_stake is not None and scaled < min_stake:
            return 0
        return scaled


class _AdxScaledRiskBudgetProfile(VolatilityBreakoutRiskCapAdxScaledRisk):
    """Research-only global risk-budget profile around the frozen ADX formula.

    The signal, leverage, exits, protections and ADX multiplier are unchanged.
    Both the per-trade risk budget and wallet stake ceiling move by the same
    factor so the latter does not silently clip the intended experiment.
    """


class VolatilityBreakoutRiskCapAdxScaledRisk2x(_AdxScaledRiskBudgetProfile):
    risk_per_trade = 0.015
    max_stake_wallet_fraction = 0.36

    def version(self) -> str:
        return "research-20260810-adx-risk-budget-2x"


class VolatilityBreakoutRiskCapAdxScaledRisk3x(_AdxScaledRiskBudgetProfile):
    risk_per_trade = 0.0225
    max_stake_wallet_fraction = 0.54

    def version(self) -> str:
        return "research-20260810-adx-risk-budget-3x"


class VolatilityBreakoutRiskCapAdxScaledRisk4x(_AdxScaledRiskBudgetProfile):
    risk_per_trade = 0.03
    max_stake_wallet_fraction = 0.72

    def version(self) -> str:
        return "research-20260810-adx-risk-budget-4x"


class _AdxScaledRiskSensitivity(VolatilityBreakoutRiskCap):
    """Sensitivity-only family; not eligible to replace the frozen V1 formula."""

    min_multiplier = 0.75
    max_multiplier = 1.25
    midpoint = 35.0
    denominator = 40.0

    def custom_stake_amount(
        self,
        pair,
        current_time,
        current_rate,
        proposed_stake,
        min_stake,
        max_stake,
        leverage,
        entry_tag,
        side,
        **kwargs,
    ):
        stake = VolatilityBreakoutRiskCap.custom_stake_amount(
            self,
            pair=pair,
            current_time=current_time,
            current_rate=current_rate,
            proposed_stake=proposed_stake,
            min_stake=min_stake,
            max_stake=max_stake,
            leverage=leverage,
            entry_tag=entry_tag,
            side=side,
            **kwargs,
        )
        if not stake:
            return stake
        candle = self._last_candle(pair)
        adx = None if candle is None else candle.get("adx")
        if adx is None or adx != adx:
            return 0
        multiplier = min(
            self.max_multiplier,
            max(
                self.min_multiplier,
                1.0 + (float(adx) - self.midpoint) / self.denominator,
            ),
        )
        scaled = min(float(stake) * multiplier, max_stake)
        if min_stake is not None and scaled < min_stake:
            return 0
        return scaled


class VolatilityBreakoutRiskCapAdxScaledRiskNarrow(_AdxScaledRiskSensitivity):
    min_multiplier = 0.80
    max_multiplier = 1.20
    denominator = 50.0

    def version(self) -> str:
        return "research-20260810-adx-scaled-risk-narrow-sensitivity"


class VolatilityBreakoutRiskCapAdxScaledRiskWide(_AdxScaledRiskSensitivity):
    min_multiplier = 0.65
    max_multiplier = 1.35
    denominator = 28.5714285714

    def version(self) -> str:
        return "research-20260810-adx-scaled-risk-wide-sensitivity"


class VolatilityBreakoutRiskCapAdxScaledRiskMid40(_AdxScaledRiskSensitivity):
    midpoint = 40.0
    denominator = 60.0

    def version(self) -> str:
        return "research-20260810-adx-scaled-risk-mid40-sensitivity"


class VolatilityBreakoutRiskCapPairDailyTrend(_PairDailyBase):
    """Single-factor candidate: pair daily EMA structure and slope align with side."""

    def version(self) -> str:
        return "research-20260810-pair-daily-trend-v1"

    def populate_entry_trend(self, dataframe, metadata):
        dataframe = super().populate_entry_trend(dataframe, metadata)
        long_daily = (
            (dataframe["close_1d"] > dataframe["pair_daily_ema50_1d"])
            & (dataframe["pair_daily_ema50_1d"] > dataframe["pair_daily_ema200_1d"])
            & (dataframe["pair_daily_ema50_slope5_1d"] > 0)
        )
        short_daily = (
            (dataframe["close_1d"] < dataframe["pair_daily_ema50_1d"])
            & (dataframe["pair_daily_ema50_1d"] < dataframe["pair_daily_ema200_1d"])
            & (dataframe["pair_daily_ema50_slope5_1d"] < 0)
        )
        dataframe.loc[~long_daily, "enter_long"] = 0
        dataframe.loc[~short_daily, "enter_short"] = 0
        return dataframe
