# source: https://raw.githubusercontent.com/ma-pony/crypto_trade_service/b9792ef228b3b7ab0305fca88e1d859e501d093b/backend/app/strategies/production/multi_factor.py
"""
多因子策略

综合多个技术指标的多因子交易策略。
通过因子打分系统决定入场方向和强度。

核心逻辑：
- 趋势因子：EMA 方向
- 动量因子：RSI 位置
- 波动因子：ATR 水平
- 成交量因子：放量程度
"""

from __future__ import annotations

from typing import TYPE_CHECKING, Any

from freqtrade.strategy import IStrategy, DecimalParameter, IntParameter
from app.strategies.production import indicators as ta

if TYPE_CHECKING:
    from pandas import DataFrame


class Github_ma_pony_crypto_trade_service__multi_factor__20260131_075609(IStrategy):
    """多因子策略

    综合趋势、动量、波动、成交量四个因子。
    """

    # 策略元数据
    strategy_id = "multi_factor_v1"
    strategy_version = "1.0.0"

    # 时间周期
    timeframe = "4h"

    # 风控配置
    minimal_roi = {
        "0": 0.15,
        "120": 0.08,
        "240": 0.04,
    }
    stoploss = -0.06

    # 追踪止损
    trailing_stop = True
    trailing_stop_positive = 0.02
    trailing_stop_positive_offset = 0.04
    trailing_only_offset_is_reached = True

    # 数据配置
    startup_candle_count = 100
    can_short = True

    # ========== 可优化参数 ==========
    # EMA 参数
    ema_fast = IntParameter(low=10, high=20, default=12, space="buy")
    ema_slow = IntParameter(low=20, high=50, default=26, space="buy")

    # RSI 参数
    rsi_period = IntParameter(low=10, high=21, default=14, space="buy")

    # ATR 参数
    atr_period = IntParameter(low=10, high=20, default=14, space="buy")

    # 成交量参数
    volume_ma_period = IntParameter(low=10, high=30, default=20, space="buy")

    # 因子阈值
    score_threshold = DecimalParameter(low=2.0, high=4.0, default=3.0, space="buy")

    def populate_indicators(
        self,
        dataframe: "DataFrame",
        metadata: dict[str, Any],
    ) -> "DataFrame":
        """计算技术指标"""
        # EMA
        dataframe["ema_fast"] = ta.ema(dataframe["close"], length=self.ema_fast.value)
        dataframe["ema_slow"] = ta.ema(dataframe["close"], length=self.ema_slow.value)

        # RSI
        dataframe["rsi"] = ta.rsi(dataframe["close"], length=self.rsi_period.value)

        # ATR
        dataframe["atr"] = ta.atr(
            dataframe["high"],
            dataframe["low"],
            dataframe["close"],
            length=self.atr_period.value,
        )
        dataframe["atr_pct"] = dataframe["atr"] / dataframe["close"] * 100

        # 成交量
        dataframe["volume_ma"] = ta.sma(dataframe["volume"], length=self.volume_ma_period.value)
        dataframe["volume_ratio"] = dataframe["volume"] / dataframe["volume_ma"]

        # 计算因子得分
        dataframe = self._calculate_factor_scores(dataframe)

        return dataframe

    def _calculate_factor_scores(self, dataframe: "DataFrame") -> "DataFrame":
        """计算因子得分"""
        # 趋势因子 (-1 到 1)
        dataframe["trend_score"] = 0.0
        dataframe.loc[dataframe["ema_fast"] > dataframe["ema_slow"], "trend_score"] = 1.0
        dataframe.loc[dataframe["ema_fast"] < dataframe["ema_slow"], "trend_score"] = -1.0

        # 动量因子 (-1 到 1)
        dataframe["momentum_score"] = (dataframe["rsi"] - 50) / 50

        # 波动因子 (0 到 1)
        atr_mean = dataframe["atr_pct"].rolling(20).mean()
        dataframe["volatility_score"] = (dataframe["atr_pct"] / atr_mean).clip(0, 2) - 1

        # 成交量因子 (0 到 1)
        dataframe["volume_score"] = (dataframe["volume_ratio"] - 1).clip(-1, 1)

        # 综合得分
        dataframe["long_score"] = (
            dataframe["trend_score"] +
            dataframe["momentum_score"].clip(0, 1) +
            dataframe["volatility_score"].clip(0, 1) +
            dataframe["volume_score"].clip(0, 1)
        )
        dataframe["short_score"] = (
            -dataframe["trend_score"] +
            (-dataframe["momentum_score"]).clip(0, 1) +
            dataframe["volatility_score"].clip(0, 1) +
            dataframe["volume_score"].clip(0, 1)
        )

        return dataframe

    def populate_entry_trend(
        self,
        dataframe: "DataFrame",
        metadata: dict[str, Any],
    ) -> "DataFrame":
        """生成入场信号"""
        dataframe.loc[:, "enter_long"] = 0
        dataframe.loc[:, "enter_short"] = 0
        dataframe.loc[:, "enter_tag"] = ""

        # 做多：综合得分超过阈值
        long_conditions = (
            (dataframe["long_score"] >= self.score_threshold.value) &
            (dataframe["volume"] > 0)
        )
        dataframe.loc[long_conditions, "enter_long"] = 1
        dataframe.loc[long_conditions, "enter_tag"] = "multi_factor_long"

        # 做空：综合得分超过阈值
        short_conditions = (
            (dataframe["short_score"] >= self.score_threshold.value) &
            (dataframe["volume"] > 0)
        )
        dataframe.loc[short_conditions, "enter_short"] = 1
        dataframe.loc[short_conditions, "enter_tag"] = "multi_factor_short"

        return dataframe

    def populate_exit_trend(
        self,
        dataframe: "DataFrame",
        metadata: dict[str, Any],
    ) -> "DataFrame":
        """生成出场信号"""
        dataframe.loc[:, "exit_long"] = 0
        dataframe.loc[:, "exit_short"] = 0
        dataframe.loc[:, "exit_tag"] = ""

        # 多头出场：趋势反转或得分下降
        exit_long_conditions = (
            (dataframe["trend_score"] < 0) |
            (dataframe["long_score"] < 1.0)
        )
        dataframe.loc[exit_long_conditions, "exit_long"] = 1
        dataframe.loc[exit_long_conditions, "exit_tag"] = "multi_factor_exit"

        # 空头出场：趋势反转或得分下降
        exit_short_conditions = (
            (dataframe["trend_score"] > 0) |
            (dataframe["short_score"] < 1.0)
        )
        dataframe.loc[exit_short_conditions, "exit_short"] = 1
        dataframe.loc[exit_short_conditions, "exit_tag"] = "multi_factor_exit"

        return dataframe
