# source: https://raw.githubusercontent.com/ma-pony/crypto_trade_service/9d13f8d58e31fb8c08bd8a964172dc42d5578bc9/backend/app/strategies/production/trend_following.py
"""
趋势跟踪策略

基于均线系统和 ADX 趋势强度指标的趋势跟踪策略。
适用于趋势明显的市场环境。

核心逻辑：
- 使用 EMA 快慢线交叉判断趋势方向
- 使用 ADX 过滤震荡行情，只在趋势明确时入场
- 使用 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__trend_following__20260201_072408(IStrategy):
    """趋势跟踪策略

    使用 EMA 均线系统 + ADX 趋势过滤器。
    在趋势明确时顺势入场，趋势减弱时出场。
    """

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

    # 时间周期
    timeframe = "4h"

    # 风控配置
    minimal_roi = {
        "0": 0.15,    # 立即 15% 止盈
        "120": 0.10,  # 2小时后 10%
        "240": 0.05,  # 4小时后 5%
        "480": 0.02,  # 8小时后 2%
    }
    stoploss = -0.08

    # 追踪止损
    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_period = IntParameter(low=8, high=21, default=12, space="buy")
    ema_slow_period = IntParameter(low=21, high=55, default=26, space="buy")
    ema_trend_period = IntParameter(low=100, high=250, default=200, space="buy")

    # ADX 参数
    adx_period = IntParameter(low=10, high=20, default=14, space="buy")
    adx_threshold = DecimalParameter(low=20.0, high=35.0, default=25.0, space="buy")

    # ATR 止损参数
    atr_period = IntParameter(low=10, high=20, default=14, space="sell")
    atr_multiplier = DecimalParameter(low=1.5, high=3.0, default=2.0, space="sell")

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

        # ADX 趋势强度
        adx_val, di_plus, di_minus = ta.adx(
            dataframe["high"],
            dataframe["low"],
            dataframe["close"],
            length=self.adx_period.value,
        )
        dataframe["adx"] = adx_val
        dataframe["di_plus"] = di_plus
        dataframe["di_minus"] = di_minus

        # ATR 波动率
        dataframe["atr"] = ta.atr(
            dataframe["high"],
            dataframe["low"],
            dataframe["close"],
            length=self.atr_period.value,
        )

        # 趋势判断
        dataframe["uptrend"] = (
            (dataframe["ema_fast"] > dataframe["ema_slow"]) &
            (dataframe["close"] > dataframe["ema_trend"])
        )
        dataframe["downtrend"] = (
            (dataframe["ema_fast"] < dataframe["ema_slow"]) &
            (dataframe["close"] < dataframe["ema_trend"])
        )

        # 金叉死叉信号
        dataframe["ema_cross_up"] = (
            (dataframe["ema_fast"] > dataframe["ema_slow"]) &
            (dataframe["ema_fast"].shift(1) <= dataframe["ema_slow"].shift(1))
        )
        dataframe["ema_cross_down"] = (
            (dataframe["ema_fast"] < dataframe["ema_slow"]) &
            (dataframe["ema_fast"].shift(1) >= dataframe["ema_slow"].shift(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"] = ""

        # 做多条件：金叉 + 上升趋势 + ADX 确认趋势强度
        long_conditions = (
            dataframe["ema_cross_up"] &
            dataframe["uptrend"] &
            (dataframe["adx"] > self.adx_threshold.value) &
            (dataframe["di_plus"] > dataframe["di_minus"]) &
            (dataframe["volume"] > 0)
        )
        dataframe.loc[long_conditions, "enter_long"] = 1
        dataframe.loc[long_conditions, "enter_tag"] = "trend_long"

        # 做空条件：死叉 + 下降趋势 + ADX 确认趋势强度
        short_conditions = (
            dataframe["ema_cross_down"] &
            dataframe["downtrend"] &
            (dataframe["adx"] > self.adx_threshold.value) &
            (dataframe["di_minus"] > dataframe["di_plus"]) &
            (dataframe["volume"] > 0)
        )
        dataframe.loc[short_conditions, "enter_short"] = 1
        dataframe.loc[short_conditions, "enter_tag"] = "trend_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["ema_cross_down"] |
            (dataframe["adx"] < 20) |
            (dataframe["close"] < dataframe["ema_trend"])
        )
        dataframe.loc[exit_long_conditions, "exit_long"] = 1
        dataframe.loc[exit_long_conditions, "exit_tag"] = "trend_exit"

        # 空头出场：趋势反转或趋势减弱
        exit_short_conditions = (
            dataframe["ema_cross_up"] |
            (dataframe["adx"] < 20) |
            (dataframe["close"] > dataframe["ema_trend"])
        )
        dataframe.loc[exit_short_conditions, "exit_short"] = 1
        dataframe.loc[exit_short_conditions, "exit_tag"] = "trend_exit"

        return dataframe

    def custom_stoploss(
        self,
        pair: str,
        trade_data: dict[str, Any],
        current_time: Any,
        current_rate: float,
        current_profit: float,
        **kwargs: Any,
    ) -> float:
        """基于 ATR 的动态止损"""
        dataframe = kwargs.get("dataframe")
        if dataframe is None or dataframe.empty:
            return self.stoploss

        last_candle = dataframe.iloc[-1]
        atr = last_candle.get("atr", 0)

        if atr > 0:
            # ATR 动态止损
            atr_stop = (atr * self.atr_multiplier.value) / current_rate
            return max(-atr_stop, self.stoploss)

        return self.stoploss
