# source: https://raw.githubusercontent.com/hankniel/freqtrade/83ba4265808f598d5e1f810daa10ebea1726a645/user_data/strategies/freqai/prod/ConsumerSwingShortStrategy4h.py
import logging
from pandas import DataFrame
from datetime import datetime
import talib.abstract as ta

from freqtrade.strategy import (IStrategy, merge_informative_pair)
from freqtrade.constants import Config
from user_data.utils.constants import DEFAULT_LABEL_NAME
from user_data.utils.helpers import clean_memory

logger = logging.getLogger(__name__)

import json
with open('configs/enter_exit_configs/Github_hankniel_freqtrade__ConsumerSwingShortStrategy4h__20251014_102948.json', 'r') as f:
    enter_exit_configs = json.load(f)

class Github_hankniel_freqtrade__ConsumerSwingShortStrategy4h__20251014_102948(IStrategy):
    """
    
    """
    timeframe = '4h'
    
    process_only_new_candles = False #! required for consumers
    stoploss = -0.99
    use_exit_signal = True
    # this is the maximum period fed to talib (timeframe independent)
    startup_candle_count: int = 400
    can_short = True
    
    def __init__(self, config: Config) -> None:
        super().__init__(config)
        # Access FreqAI config from the main config
        self.freqai_info = self.config.get('freqai', {})
    
    def populate_indicators(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        """
        Use the websocket api to get pre-populated indicators from another freqtrade instance.
        Use `self.dp.get_producer_df(pair)` to get the dataframe
        """
        label = f'&-s_{self.freqai_info.get('label_config', {}).get('label_name', DEFAULT_LABEL_NAME)}'
        pair = metadata['pair']
        timeframe = self.timeframe
        
        # producer_pairs = self.dp.get_producer_pairs()
        # You can specify which producer to get pairs from via:
        # self.dp.get_producer_pairs("my_other_producer")
        
        # This func returns the analyzed dataframe, and when it was analyzed
        producer_df, _ = self.dp.get_producer_df(pair)
        # You can get other data if the producer makes it available:
        # self.dp.get_producer_df(
        #   pair,
        #   timeframe="1h",
        #   candle_type=CandleType.SPOT,
        #   producer_name="my_other_producer"
        # )
        
        if not producer_df.empty:
            # If you plan on passing the producer's entry/exit signal directly,
            # specify ffill=False or it will have unintended results
            #* Here we use custom entry/exit signals of consumer
            columns = ['date', 'do_predict', label, 'open', 'high', 'low', 'close', 'volume']
            try:
                merged_df = merge_informative_pair(
                    dataframe, producer_df[columns],
                    timeframe, timeframe,
                    append_timeframe=False,
                    suffix='default'
                )
            except:
                merged_df = merge_informative_pair(
                    dataframe, producer_df,
                    timeframe, timeframe,
                    append_timeframe=False,
                    suffix='default'
                )
            del producer_df
            clean_memory()
            return merged_df
        else:
            raise ValueError(f"No producer dataframe found for pair {pair}")
    
    def populate_entry_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        logger.info(f"Populating entry signal for {metadata['pair']}")
        #* Primary entry short conditions
        label = f'&-s_{self.freqai_info.get('label_config', {}).get('label_name', DEFAULT_LABEL_NAME)}'
        
        dataframe.loc[:, 'enter_short'] = 0
        dataframe.loc[:, 'enter_tag'] = None
        
        # Get filter configurations
        # filters_config = self.config.get('signal_filters', {})
        filters_config = enter_exit_configs.get(metadata['pair'], {})
        
        # Base condition: ML model prediction
        base_enter_condition = (
            (dataframe['do_predict'] == 1) &
            (dataframe[label] == 'short')
        )
        
        # Indicators
        sma = ta.SMA(dataframe, timeperiod=filters_config.get('trend_filter', {}).get('entry_period', 20))
        rsi = ta.RSI(dataframe, timeperiod=filters_config.get('momentum_filter', {}).get('entry_momentum_period', 14))
        mfi = ta.MFI(dataframe, timeperiod=filters_config.get('volume_filter', {}).get('entry_volume_period', 20))
        oversold_threshold = filters_config.get('momentum_filter', {}).get('oversold_threshold', 30)
        entry_volume_threshold = filters_config.get('volume_filter', {}).get('entry_volume_threshold', 20)
        
        # Filter config
        num_entry_filters = filters_config.get('num_entry_filters', 1)
        entry_logics = filters_config.get('entry_logics', ["trend_filter & momentum_filter & volume_filter"])
        assert len(entry_logics) == num_entry_filters, f"Entry_logics length must match num_entry_filters. " \
            f"Currently having {num_entry_filters} and {len(entry_logics)}: {entry_logics}"
        
        trend_filter = (dataframe['close'] < sma)
        momentum_filter = (rsi < oversold_threshold)
        volume_filter = (mfi < entry_volume_threshold)
        
        # Create evaluation context for safe eval
        eval_context = {
            'trend_filter': trend_filter,
            'momentum_filter': momentum_filter,
            'volume_filter': volume_filter
        }
        
        # Conditions
        for f in range(num_entry_filters):
            logic = entry_logics[f]
            combined = eval(logic, {"__builtins__": {}}, eval_context)
            conditions = base_enter_condition & combined
            
            enter_tag = logic.replace('&', 'and').replace('|', 'or').replace(' ', '_')
            
            dataframe.loc[conditions, 'enter_short'] = 1
            dataframe.loc[conditions, 'enter_tag'] = enter_tag
        
        logger.info(f"Last candle signals for pair {metadata['pair']}: enter_short={dataframe['enter_short'].iloc[-1]}")
        if dataframe['enter_short'].iloc[-1] == 1:
            try:
                self.dp.send_msg(f"Enter pair {metadata['pair']} at {dataframe['close'].iloc[-1]}")
            except Exception as e:
                logger.error(f"Error sending enter message: {e}")
        
        return dataframe
    
    def populate_exit_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        """Exit logic with filtering (opposite conditions to entry)"""
        logger.info(f"Populating exit signal for {metadata['pair']}")
        label = f'&-s_{self.freqai_info.get('label_config', {}).get('label_name', DEFAULT_LABEL_NAME)}'
        
        dataframe.loc[:, 'exit_short'] = 0
        dataframe.loc[:, 'exit_tag'] = None
        
        # Get filter configurations
        # filters_config = self.config.get('signal_filters', {})
        filters_config = enter_exit_configs.get(metadata['pair'], {})
        
        # Base exit condition: ML model prediction
        base_exit_condition = (
            (dataframe['do_predict'] == 1) &
            (dataframe[label] == 'short')
        )
        
        # Indicators
        sma = ta.SMA(dataframe, timeperiod=filters_config.get('trend_filter', {}).get('exit_period', 20))
        rsi = ta.RSI(dataframe, timeperiod=filters_config.get('momentum_filter', {}).get('exit_momentum_period', 14))
        mfi = ta.MFI(dataframe, timeperiod=filters_config.get('volume_filter', {}).get('exit_volume_period', 20))
        overbought_threshold = filters_config.get('momentum_filter', {}).get('overbought_threshold', 30)
        exit_volume_threshold = filters_config.get('volume_filter', {}).get('exit_volume_threshold', 80)
        
        # Filter config
        num_exit_filters = filters_config.get('num_exit_filters', 1)
        exit_logics = filters_config.get('exit_logics', ["trend_filter & momentum_filter & volume_filter"])
        assert len(exit_logics) == num_exit_filters, f"exit_logics length must match num_exit_filters. " \
            f"Currently having {num_exit_filters} and {len(exit_logics)}: {exit_logics}"
        
        trend_filter = (dataframe['close'] > sma)
        momentum_filter = (rsi > overbought_threshold)
        volume_filter = (mfi > exit_volume_threshold)
        
        # Create evaluation context for safe eval
        eval_context = {
            'trend_filter': trend_filter,
            'momentum_filter': momentum_filter,
            'volume_filter': volume_filter
        }
        
        # Conditions
        for f in range(num_exit_filters):
            logic = exit_logics[f]
            combined = eval(logic, {"__builtins__": {}}, eval_context)
            conditions = base_exit_condition & combined
            
            exit_tag = logic.replace('&', 'and').replace('|', 'or').replace(' ', '_')
            
            dataframe.loc[conditions, 'exit_short'] = 1
            dataframe.loc[conditions, 'exit_tag'] = exit_tag
        
        logger.info(f"Last candle signals for pair {metadata['pair']}: exit_short={dataframe['exit_short'].iloc[-1]}")
        if dataframe['exit_short'].iloc[-1] == 1:
            try:
                self.dp.send_msg(f"Exit pair {metadata['pair']} at {dataframe['close'].iloc[-1]}")
            except Exception as e:
                logger.error(f"Error sending exit message: {e}")
        
        return dataframe
    
    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:
        return self.config['leverage']