# source: https://raw.githubusercontent.com/sadhahacker/freqtrade/f3b0bcc6236aea7c2a0c2f218007f50f0ed244ff/user_data/strategies/CompoundScalper.py
# pragma pylint: disable=missing-docstring, invalid-name, pointless-string-statement
# flake8: noqa: F401
# isort: skip_file
# --- Do not remove these imports ---
import logging
import numpy as np
import pandas as pd
from datetime import datetime, timedelta, timezone
from pandas import DataFrame
from typing import Optional

from freqtrade.strategy import IStrategy, Trade
from freqtrade.persistence import Trade as TradeModel

import talib.abstract as ta

logger = logging.getLogger(__name__)


class Github_sadhahacker_freqtrade__CompoundScalper__20260419_104106(IStrategy):
    """
    Github_sadhahacker_freqtrade__CompoundScalper__20260419_104106 v3 - Supertrend + EMA + RSI pullback strategy.

    Entry logic:
      - Supertrend (ATR 10, factor 3) as adaptive trend filter
      - EMA 9/21 crossover confirms short-term direction
      - RSI pullback provides entry timing
      - Bullish/bearish candle confirmation

    Risk math (5x leverage):
      TP = 30% account (6% price)  |  SL = 23% account (4.6% price)
      R:R = 1:1.3 — breakeven at 43% win rate.

    20-Pip Challenge: $100 -> compound 30% per level across 30 levels.
    """

    INTERFACE_VERSION = 3
    can_short: bool = True

    # ---- 30% take profit (6% price × 5x) ----
    minimal_roi = {"0": 0.30}

    # ---- 23% stoploss (4.6% price × 5x) — wider stop for crypto volatility ----
    stoploss = -0.23

    trailing_stop = False

    timeframe = "1h"
    process_only_new_candles = True

    use_exit_signal = True
    exit_profit_only = False
    ignore_roi_if_entry_signal = False

    startup_candle_count: int = 210

    order_types = {
        "entry": "market",
        "exit": "market",
        "stoploss": "market",
        "stoploss_on_exchange": True,
    }

    order_time_in_force = {"entry": "GTC", "exit": "GTC"}

    plot_config = {
        "main_plot": {
            "supertrend": {"color": "blue"},
            "ema9": {"color": "orange"},
            "ema21": {"color": "yellow"},
        },
        "subplots": {
            "RSI": {"rsi": {"color": "red"}},
        },
    }

    # ---- Risk management state ----
    peak_balance: float = 0.0
    last_loss_time: Optional[datetime] = None
    cooldown_candles: int = 2
    starting_balance: float = 2.0


    def informative_pairs(self):
        return []

    def _calc_supertrend(self, dataframe: DataFrame, period: int = 10, factor: float = 3.0):
        atr = ta.ATR(dataframe, timeperiod=period)
        hl2 = (dataframe["high"] + dataframe["low"]) / 2

        upper_basic = hl2 + factor * atr
        lower_basic = hl2 - factor * atr

        upper_band = upper_basic.values.copy()
        lower_band = lower_basic.values.copy()
        close = dataframe["close"].values
        direction = np.ones(len(dataframe), dtype=int)

        for i in range(1, len(dataframe)):
            if np.isnan(atr.iloc[i]):
                continue

            if np.isnan(lower_band[i - 1]):
                lower_band[i] = lower_basic.iloc[i]
                upper_band[i] = upper_basic.iloc[i]
                direction[i] = -1 if close[i] > upper_band[i] else 1
                continue

            if lower_basic.iloc[i] > lower_band[i - 1] or close[i - 1] < lower_band[i - 1]:
                lower_band[i] = lower_basic.iloc[i]
            else:
                lower_band[i] = lower_band[i - 1]

            if upper_basic.iloc[i] < upper_band[i - 1] or close[i - 1] > upper_band[i - 1]:
                upper_band[i] = upper_basic.iloc[i]
            else:
                upper_band[i] = upper_band[i - 1]

            if direction[i - 1] == 1:
                direction[i] = -1 if close[i] > upper_band[i] else 1
            else:
                direction[i] = 1 if close[i] < lower_band[i] else -1

        dir_series = pd.Series(direction, index=dataframe.index)
        lb_series = pd.Series(lower_band, index=dataframe.index)
        ub_series = pd.Series(upper_band, index=dataframe.index)
        return lb_series.where(dir_series == -1, ub_series), dir_series

    def populate_indicators(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        # === Trend indicators ===
        dataframe["ema9"] = ta.EMA(dataframe, timeperiod=9)
        dataframe["ema21"] = ta.EMA(dataframe, timeperiod=21)
        dataframe["ema200"] = ta.EMA(dataframe, timeperiod=200)
        dataframe["sar"] = ta.SAR(dataframe, acceleration=0.02, maximum=0.2)
        dataframe["adx"] = ta.ADX(dataframe, timeperiod=14)

        # === Momentum / Oscillators ===
        dataframe["rsi"] = ta.RSI(dataframe, timeperiod=14)
        dataframe["cci"] = ta.CCI(dataframe, timeperiod=20)
        dataframe["willr"] = ta.WILLR(dataframe, timeperiod=14)
        dataframe["mfi"] = ta.MFI(dataframe, timeperiod=14)
        macd = ta.MACD(dataframe, fastperiod=12, slowperiod=26, signalperiod=9)
        dataframe["macd"] = macd["macd"]
        dataframe["macd_signal"] = macd["macdsignal"]

        # === Volatility ===
        bb = ta.BBANDS(dataframe, timeperiod=20, nbdevup=2.0, nbdevdn=2.0)
        dataframe["bb_upper"] = bb["upperband"]
        dataframe["bb_lower"] = bb["lowerband"]
        dataframe["bb_mid"] = bb["middleband"]

        # === Candlestick patterns (TA-Lib) ===
        dataframe["cdl_engulfing"] = ta.CDLENGULFING(dataframe)
        dataframe["cdl_hammer"] = ta.CDLHAMMER(dataframe)
        dataframe["cdl_invhammer"] = ta.CDLINVERTEDHAMMER(dataframe)
        dataframe["cdl_shooting_star"] = ta.CDLSHOOTINGSTAR(dataframe)
        dataframe["cdl_morningstar"] = ta.CDLMORNINGSTAR(dataframe)
        dataframe["cdl_eveningstar"] = ta.CDLEVENINGSTAR(dataframe)

        # === Volume ===
        dataframe["vol_sma20"] = dataframe["volume"].rolling(20).mean()

        # === Confluence scores ===
        # LONG score: each indicator contributing 0 or 1
        long_score = pd.Series(0.0, index=dataframe.index)
        long_score += (dataframe["close"] > dataframe["ema200"]).astype(float) * 2      # strongest filter
        long_score += (dataframe["ema9"] > dataframe["ema21"]).astype(float)
        long_score += (dataframe["sar"] < dataframe["close"]).astype(float)
        long_score += (dataframe["adx"] > 20).astype(float)
        long_score += (dataframe["macd"] > dataframe["macd_signal"]).astype(float)
        long_score += ((dataframe["rsi"] > 30) & (dataframe["rsi"] < 50)).astype(float)
        long_score += ((dataframe["cci"] > -100) & (dataframe["cci"] < 0)).astype(float)
        long_score += ((dataframe["willr"] > -80) & (dataframe["willr"] < -30)).astype(float)
        long_score += ((dataframe["mfi"] > 20) & (dataframe["mfi"] < 50)).astype(float)
        long_score += (dataframe["close"] < dataframe["bb_mid"]).astype(float)           # price below BB middle = pullback
        long_score += (dataframe["close"] > dataframe["open"]).astype(float)
        long_score += (dataframe["cdl_engulfing"] > 0).astype(float)
        long_score += (dataframe["cdl_hammer"] > 0).astype(float)
        long_score += (dataframe["cdl_morningstar"] > 0).astype(float)
        long_score += (dataframe["volume"] > dataframe["vol_sma20"]).astype(float)
        dataframe["long_score"] = long_score

        # SHORT score
        short_score = pd.Series(0.0, index=dataframe.index)
        short_score += (dataframe["close"] < dataframe["ema200"]).astype(float) * 2
        short_score += (dataframe["ema9"] < dataframe["ema21"]).astype(float)
        short_score += (dataframe["sar"] > dataframe["close"]).astype(float)
        short_score += (dataframe["adx"] > 20).astype(float)
        short_score += (dataframe["macd"] < dataframe["macd_signal"]).astype(float)
        short_score += ((dataframe["rsi"] > 50) & (dataframe["rsi"] < 70)).astype(float)
        short_score += ((dataframe["cci"] > 0) & (dataframe["cci"] < 100)).astype(float)
        short_score += ((dataframe["willr"] > -70) & (dataframe["willr"] < -20)).astype(float)
        short_score += ((dataframe["mfi"] > 50) & (dataframe["mfi"] < 80)).astype(float)
        short_score += (dataframe["close"] > dataframe["bb_mid"]).astype(float)
        short_score += (dataframe["close"] < dataframe["open"]).astype(float)
        short_score += (dataframe["cdl_engulfing"] < 0).astype(float)
        short_score += (dataframe["cdl_shooting_star"] > 0).astype(float)
        short_score += (dataframe["cdl_eveningstar"] > 0).astype(float)
        short_score += (dataframe["volume"] > dataframe["vol_sma20"]).astype(float)
        dataframe["short_score"] = short_score

        return dataframe

    def populate_entry_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        # LONG: High confluence score (>= 9 out of 16)
        long_mask = dataframe["long_score"] >= 9
        dataframe.loc[long_mask, ["enter_long", "enter_tag"]] = (1, "confluence_long")

        # SHORT: High confluence score (>= 9 out of 16)
        short_mask = dataframe["short_score"] >= 9
        dataframe.loc[short_mask, ["enter_short", "enter_tag"]] = (1, "confluence_short")

        return dataframe

    def populate_exit_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        # Exit long only when strong bearish reversal (score >= 10)
        dataframe["exit_long"] = (dataframe["short_score"] >= 10).astype(int)
        # Exit short only when strong bullish reversal
        dataframe["exit_short"] = (dataframe["long_score"] >= 10).astype(int)
        return dataframe

    def _get_challenge_level(self, balance: float) -> int:
        if balance <= 0:
            return 0
        level = 1
        threshold = self.starting_balance
        while threshold * 1.30 <= balance:
            threshold *= 1.30
            level += 1
        return level

    def custom_exit(
        self,
        pair: str,
        trade: Trade,
        current_time: datetime,
        current_rate: float,
        current_profit: float,
        **kwargs,
    ) -> Optional[str]:
        return None

    def confirm_trade_exit(
        self,
        pair: str,
        trade: Trade,
        order_type: str,
        amount: float,
        rate: float,
        time_in_force: str,
        exit_reason: str,
        current_time: datetime,
        **kwargs,
    ) -> bool:
        profit = trade.calc_profit_ratio(rate)
        wallet_balance = self.wallets.get_free(self.config["stake_currency"])
        current_balance = wallet_balance + (amount * rate * profit)
        if current_balance > self.peak_balance:
            self.peak_balance = current_balance
        if profit < 0:
            self.last_loss_time = current_time
        level = self._get_challenge_level(current_balance)
        logger.info(
            "Github_sadhahacker_freqtrade__CompoundScalper__20260419_104106: EXIT %s | Profit: %.2f%% | Balance: $%.2f | Level: %d/30",
            pair, profit * 100, current_balance, level,
        )
        return True

    def confirm_trade_entry(
        self,
        pair: str,
        order_type: str,
        amount: float,
        rate: float,
        time_in_force: str,
        current_time: datetime,
        entry_tag: Optional[str],
        side: str,
        **kwargs,
    ) -> bool:
        all_trades = TradeModel.get_trades_proxy(is_open=None)
        today_date = current_time.date()
        trades_today = [t for t in all_trades if t.open_date_utc.date() == today_date]

        # Gate 1: Max 8 trades per day (across all pairs)
        if len(trades_today) >= 8:
            return False

        # Gate 2: Max 3 consecutive losses today — stop trading
        closed_today = sorted(
            [t for t in trades_today if not t.is_open],
            key=lambda t: t.close_date_utc,
        )
        if len(closed_today) >= 3:
            last_three = closed_today[-3:]
            if all(t.calc_profit_ratio(rate=t.close_rate) < 0 for t in last_three):
                return False

        # Gate 3: Cooldown after loss (2 candles = 2 hours)
        if self.last_loss_time is not None:
            cooldown_seconds = self.cooldown_candles * 3600
            elapsed = (current_time - self.last_loss_time).total_seconds()
            if elapsed < cooldown_seconds:
                return False

        # Gate 4: Drawdown circuit breaker (50% of peak)
        wallet_balance = self.wallets.get_free(self.config["stake_currency"])
        if self.peak_balance == 0:
            self.peak_balance = wallet_balance
        if wallet_balance < self.peak_balance * 0.50:
            return False


        level = self._get_challenge_level(wallet_balance)
        logger.info(
            "Github_sadhahacker_freqtrade__CompoundScalper__20260419_104106: ENTRY %s %s [%s] | Balance: $%.2f | Level: %d/30",
            side.upper(), pair, entry_tag or "", wallet_balance, level,
        )
        return True

    def leverage(
        self,
        pair: str,
        current_time: datetime,
        current_rate: float,
        proposed_leverage: float,
        max_leverage: float,
        entry_tag: Optional[str],
        side: str,
        **kwargs,
    ) -> float:
        return 5.0
