# source: https://raw.githubusercontent.com/xRetr00/Freqtradex/479e46c35cd0bde6dba4eb34797081ae6ca9af3c/templates/FreqaiExampleStrategy.py
import logging
from functools import reduce
from typing import Dict, List, Optional, Tuple
from datetime import datetime, timedelta
import numpy as np
from dataclasses import dataclass

import talib.abstract as ta
from pandas import DataFrame, Series
from technical import qtpylib

from freqtrade.strategy import IStrategy
from freqtrade.strategy.strategy_helper import merge_informative_pair


logger = logging.getLogger(__name__)


@dataclass
class MarketState:
    """Dataclass to store market state information"""
    state: str
    volatility: float
    trend_strength: float
    volume_profile: str
    timestamp: datetime


class Github_xRetr00_Freqtradex__FreqaiExampleStrategy__20241129_193154(IStrategy):
    """
    Example strategy showing how the user connects their own
    IFreqaiModel to the strategy.

    Warning! This is a showcase of functionality,
    which means that it is designed to show various functions of FreqAI
    and it runs on all computers. We use this showcase to help users
    understand how to build a strategy, and we use it as a benchmark
    to help debug possible problems.

    This means this is *not* meant to be run live in production.
    """

    # Add class variables for adaptive features
    MARKET_STATES = ["trending", "ranging", "volatile"]
    VOLUME_PROFILES = ["high", "normal", "low"]
    
    def __init__(self, config: dict) -> None:
        super().__init__(config)
        
        # Initialize market state tracking
        self.market_state_history: List[MarketState] = []
        self.current_market_state: Optional[MarketState] = None
        
        # Dynamic parameter sets for different market conditions
        self.parameter_sets = {
            "trending": {
                "roi": {"0": 0.15, "120": 0.08, "240": 0.04, "360": 0},
                "stoploss": -0.08,
                "entry_thresholds": {"long": 0.02, "short": -0.02},
                "exit_thresholds": {"long": -0.01, "short": 0.01},
            },
            "ranging": {
                "roi": {"0": 0.08, "60": 0.04, "120": 0.02, "180": 0},
                "stoploss": -0.05,
                "entry_thresholds": {"long": 0.015, "short": -0.015},
                "exit_thresholds": {"long": -0.008, "short": 0.008},
            },
            "volatile": {
                "roi": {"0": 0.20, "180": 0.10, "360": 0.05, "720": 0},
                "stoploss": -0.12,
                "entry_thresholds": {"long": 0.03, "short": -0.03},
                "exit_thresholds": {"long": -0.015, "short": 0.015},
            }
        }

    minimal_roi = {"0": 0.1, "240": -1}

    plot_config = {
        "main_plot": {},
        "subplots": {
            "&-s_close": {"&-s_close": {"color": "blue"}},
            "do_predict": {
                "do_predict": {"color": "brown"},
            },
        },
    }

    process_only_new_candles = True
    stoploss = -0.05
    use_exit_signal = True
    # this is the maximum period fed to talib (timeframe independent)
    startup_candle_count: int = 40
    can_short = True

    def feature_engineering_expand_all(
        self, dataframe: DataFrame, period: int, metadata: dict, **kwargs
    ) -> DataFrame:
        """
        *Only functional with FreqAI enabled strategies*
        This function will automatically expand the defined features on the config defined
        `indicator_periods_candles`, `include_timeframes`, `include_shifted_candles`, and
        `include_corr_pairs`. In other words, a single feature defined in this function
        will automatically expand to a total of
        `indicator_periods_candles` * `include_timeframes` * `include_shifted_candles` *
        `include_corr_pairs` numbers of features added to the model.

        All features must be prepended with `%` to be recognized by FreqAI internals.

        Access metadata such as the current pair/timeframe with:

        `metadata["pair"]` `metadata["tf"]`

        More details on how these config defined parameters accelerate feature engineering
        in the documentation at:

        https://www.freqtrade.io/en/latest/freqai-parameter-table/#feature-parameters

        https://www.freqtrade.io/en/latest/freqai-feature-engineering/#defining-the-features

        :param dataframe: strategy dataframe which will receive the features
        :param period: period of the indicator - usage example:
        :param metadata: metadata of current pair
        dataframe["%-ema-period"] = ta.EMA(dataframe, timeperiod=period)
        """

        dataframe["%-rsi-period"] = ta.RSI(dataframe, timeperiod=period)
        dataframe["%-mfi-period"] = ta.MFI(dataframe, timeperiod=period)
        dataframe["%-adx-period"] = ta.ADX(dataframe, timeperiod=period)
        dataframe["%-sma-period"] = ta.SMA(dataframe, timeperiod=period)
        dataframe["%-ema-period"] = ta.EMA(dataframe, timeperiod=period)

        bollinger = qtpylib.bollinger_bands(
            qtpylib.typical_price(dataframe), window=period, stds=2.2
        )
        dataframe["bb_lowerband-period"] = bollinger["lower"]
        dataframe["bb_middleband-period"] = bollinger["mid"]
        dataframe["bb_upperband-period"] = bollinger["upper"]

        dataframe["%-bb_width-period"] = (
            dataframe["bb_upperband-period"] - dataframe["bb_lowerband-period"]
        ) / dataframe["bb_middleband-period"]
        dataframe["%-close-bb_lower-period"] = dataframe["close"] / dataframe["bb_lowerband-period"]

        dataframe["%-roc-period"] = ta.ROC(dataframe, timeperiod=period)

        dataframe["%-relative_volume-period"] = (
            dataframe["volume"] / dataframe["volume"].rolling(period).mean()
        )

        # Add market state detection features
        dataframe[f"%-volatility-{period}"] = self.calculate_volatility(dataframe, period)
        dataframe[f"%-trend_strength-{period}"] = self.calculate_trend_strength(dataframe, period)
        dataframe[f"%-volume_profile-{period}"] = self.calculate_volume_profile(dataframe, period)
        
        return dataframe

    def feature_engineering_expand_basic(
        self, dataframe: DataFrame, metadata: dict, **kwargs
    ) -> DataFrame:
        """
        *Only functional with FreqAI enabled strategies*
        This function will automatically expand the defined features on the config defined
        `include_timeframes`, `include_shifted_candles`, and `include_corr_pairs`.
        In other words, a single feature defined in this function
        will automatically expand to a total of
        `include_timeframes` * `include_shifted_candles` * `include_corr_pairs`
        numbers of features added to the model.

        Features defined here will *not* be automatically duplicated on user defined
        `indicator_periods_candles`

        All features must be prepended with `%` to be recognized by FreqAI internals.

        Access metadata such as the current pair/timeframe with:

        `metadata["pair"]` `metadata["tf"]`

        More details on how these config defined parameters accelerate feature engineering
        in the documentation at:

        https://www.freqtrade.io/en/latest/freqai-parameter-table/#feature-parameters

        https://www.freqtrade.io/en/latest/freqai-feature-engineering/#defining-the-features

        :param dataframe: strategy dataframe which will receive the features
        :param metadata: metadata of current pair
        dataframe["%-pct-change"] = dataframe["close"].pct_change()
        dataframe["%-ema-200"] = ta.EMA(dataframe, timeperiod=200)
        """
        dataframe["%-pct-change"] = dataframe["close"].pct_change()
        dataframe["%-raw_volume"] = dataframe["volume"]
        dataframe["%-raw_price"] = dataframe["close"]
        return dataframe

    def feature_engineering_standard(
        self, dataframe: DataFrame, metadata: dict, **kwargs
    ) -> DataFrame:
        """
        *Only functional with FreqAI enabled strategies*
        This optional function will be called once with the dataframe of the base timeframe.
        This is the final function to be called, which means that the dataframe entering this
        function will contain all the features and columns created by all other
        freqai_feature_engineering_* functions.

        This function is a good place to do custom exotic feature extractions (e.g. tsfresh).
        This function is a good place for any feature that should not be auto-expanded upon
        (e.g. day of the week).

        All features must be prepended with `%` to be recognized by FreqAI internals.

        Access metadata such as the current pair with:

        `metadata["pair"]`

        More details about feature engineering available:

        https://www.freqtrade.io/en/latest/freqai-feature-engineering

        :param dataframe: strategy dataframe which will receive the features
        :param metadata: metadata of current pair
        usage example: dataframe["%-day_of_week"] = (dataframe["date"].dt.dayofweek + 1) / 7
        """
        dataframe["%-day_of_week"] = dataframe["date"].dt.dayofweek
        dataframe["%-hour_of_day"] = dataframe["date"].dt.hour
        return dataframe

    def set_freqai_targets(self, dataframe: DataFrame, metadata: dict, **kwargs) -> DataFrame:
        """
        *Only functional with FreqAI enabled strategies*
        Required function to set the targets for the model.
        All targets must be prepended with `&` to be recognized by the FreqAI internals.

        Access metadata such as the current pair with:

        `metadata["pair"]`

        More details about feature engineering available:

        https://www.freqtrade.io/en/latest/freqai-feature-engineering

        :param dataframe: strategy dataframe which will receive the targets
        :param metadata: metadata of current pair
        usage example: dataframe["&-target"] = dataframe["close"].shift(-1) / dataframe["close"]
        """
        dataframe["&-s_close"] = (
            dataframe["close"]
            .shift(-self.freqai_info["feature_parameters"]["label_period_candles"])
            .rolling(self.freqai_info["feature_parameters"]["label_period_candles"])
            .mean()
            / dataframe["close"]
            - 1
        )

        # Classifiers are typically set up with strings as targets:
        # df['&s-up_or_down'] = np.where( df["close"].shift(-100) >
        #                                 df["close"], 'up', 'down')

        # If user wishes to use multiple targets, they can add more by
        # appending more columns with '&'. User should keep in mind that multi targets
        # requires a multioutput prediction model such as
        # freqai/prediction_models/CatboostRegressorMultiTarget.py,
        # freqtrade trade --freqaimodel CatboostRegressorMultiTarget

        # df["&-s_range"] = (
        #     df["close"]
        #     .shift(-self.freqai_info["feature_parameters"]["label_period_candles"])
        #     .rolling(self.freqai_info["feature_parameters"]["label_period_candles"])
        #     .max()
        #     -
        #     df["close"]
        #     .shift(-self.freqai_info["feature_parameters"]["label_period_candles"])
        #     .rolling(self.freqai_info["feature_parameters"]["label_period_candles"])
        #     .min()
        # )

        return dataframe

    def calculate_volatility(self, dataframe: DataFrame, period: int) -> Series:
        """Calculate market volatility using multiple indicators"""
        bb = qtpylib.bollinger_bands(dataframe['close'], window=period)
        atr = ta.ATR(dataframe, timeperiod=period)
        
        volatility = (
            (bb['upper'] - bb['lower']) / bb['mid'] * 0.5 +
            (atr / dataframe['close']) * 0.5
        )
        return volatility

    def calculate_trend_strength(self, dataframe: DataFrame, period: int) -> Series:
        """Calculate trend strength using multiple indicators"""
        adx = ta.ADX(dataframe, timeperiod=period)
        ema_short = ta.EMA(dataframe, timeperiod=period)
        ema_long = ta.EMA(dataframe, timeperiod=period * 2)
        
        trend_strength = (
            (adx / 100.0) * 0.4 +
            abs(ema_short - ema_long) / ema_long * 0.6
        )
        return trend_strength

    def calculate_volume_profile(self, dataframe: DataFrame, period: int) -> Series:
        """Analyze volume profile"""
        volume_ma = dataframe['volume'].rolling(period).mean()
        return dataframe['volume'] / volume_ma

    def detect_market_state(self, dataframe: DataFrame) -> MarketState:
        """Enhanced market state detection"""
        current_volatility = dataframe['%-volatility-14'].iloc[-1]
        current_trend = dataframe['%-trend_strength-14'].iloc[-1]
        current_volume = dataframe['%-volume_profile-14'].iloc[-1]
        
        # Determine volume profile
        if current_volume > 1.5:
            volume_profile = "high"
        elif current_volume < 0.5:
            volume_profile = "low"
        else:
            volume_profile = "normal"
            
        # Determine market state
        if current_volatility > 0.15:
            state = "volatile"
        elif current_trend > 0.12:
            state = "trending"
        else:
            state = "ranging"
            
        return MarketState(
            state=state,
            volatility=current_volatility,
            trend_strength=current_trend,
            volume_profile=volume_profile,
            timestamp=datetime.now()
        )

    def adapt_parameters(self, market_state: MarketState) -> None:
        """Adapt strategy parameters based on market state"""
        params = self.parameter_sets[market_state.state]
        
        # Adjust parameters based on volume profile
        if market_state.volume_profile == "high":
            params = self._adjust_for_high_volume(params)
        elif market_state.volume_profile == "low":
            params = self._adjust_for_low_volume(params)
            
        # Update strategy parameters
        self.minimal_roi = params["roi"]
        self.stoploss = params["stoploss"]
        
        # Store state for reference
        self.current_market_state = market_state
        self.market_state_history.append(market_state)

    def _adjust_for_high_volume(self, params: Dict) -> Dict:
        """Adjust parameters for high volume conditions"""
        return {
            **params,
            "entry_thresholds": {
                k: v * 1.2 for k, v in params["entry_thresholds"].items()
            },
            "exit_thresholds": {
                k: v * 1.2 for k, v in params["exit_thresholds"].items()
            }
        }

    def _adjust_for_low_volume(self, params: Dict) -> Dict:
        """Adjust parameters for low volume conditions"""
        return {
            **params,
            "entry_thresholds": {
                k: v * 0.8 for k, v in params["entry_thresholds"].items()
            },
            "exit_thresholds": {
                k: v * 0.8 for k, v in params["exit_thresholds"].items()
            }
        }

    def populate_indicators(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        # Original indicator calculations
        dataframe = super().populate_indicators(dataframe, metadata)
        
        # Detect and adapt to market state
        market_state = self.detect_market_state(dataframe)
        self.adapt_parameters(market_state)
        
        return dataframe

    def populate_entry_trend(self, df: DataFrame, metadata: dict) -> DataFrame:
        params = self.parameter_sets[self.current_market_state.state]
        
        # Long entries
        long_conditions = [
            df["do_predict"] == 1,
            df["&-s_close"] > params["entry_thresholds"]["long"],
            # Add volume profile conditions
            (df["volume"] > 0),
            (
                (self.current_market_state.volume_profile == "high") |
                (df["volume"] > df["volume"].rolling(24).mean())
            )
        ]
        
        if long_conditions:
            df.loc[reduce(lambda x, y: x & y, long_conditions), ["enter_long", "enter_tag"]] = (1, "long")

        # Short entries with similar adaptations
        short_conditions = [
            df["do_predict"] == 1,
            df["&-s_close"] < params["entry_thresholds"]["short"],
            (df["volume"] > 0),
            (
                (self.current_market_state.volume_profile == "high") |
                (df["volume"] > df["volume"].rolling(24).mean())
            )
        ]
        
        if short_conditions:
            df.loc[reduce(lambda x, y: x & y, short_conditions), ["enter_short", "enter_tag"]] = (1, "short")

        return df

    def populate_exit_trend(self, df: DataFrame, metadata: dict) -> DataFrame:
        exit_long_conditions = [df["do_predict"] == 1, df["&-s_close"] < 0]
        if exit_long_conditions:
            df.loc[reduce(lambda x, y: x & y, exit_long_conditions), "exit_long"] = 1

        exit_short_conditions = [df["do_predict"] == 1, df["&-s_close"] > 0]
        if exit_short_conditions:
            df.loc[reduce(lambda x, y: x & y, exit_short_conditions), "exit_short"] = 1

        return df

    def confirm_trade_entry(
        self,
        pair: str,
        order_type: str,
        amount: float,
        rate: float,
        time_in_force: str,
        current_time,
        entry_tag,
        side: str,
        **kwargs,
    ) -> bool:
        df, _ = self.dp.get_analyzed_dataframe(pair, self.timeframe)
        last_candle = df.iloc[-1].squeeze()

        if side == "long":
            if rate > (last_candle["close"] * (1 + 0.0025)):
                return False
        else:
            if rate < (last_candle["close"] * (1 - 0.0025)):
                return False

        return True
