# source: https://raw.githubusercontent.com/CrocoPruki/GalaxyWatch_BitGetApp_WearOs/a09a4c01078d3dd5ffc96388a812e64006284702/BrooksPA_V1.py
from __future__ import annotations

from datetime import datetime
from typing import Any

import numpy as np
import talib.abstract as ta
from pandas import DataFrame
from sklearn.cluster import KMeans

from freqtrade.persistence import Trade
from freqtrade.strategy import IStrategy, stoploss_from_absolute


class Github_CrocoPruki_GalaxyWatch_BitGetApp_WearOs__BrooksPA_V1__20260719_235030(IStrategy):
    """
    Simplified price-action strategy built around:
    - H2 / H3 and L2 / L3 continuation logic
    - cluster reversals
    - EMA return impulse setup
    - trend channels and range environments
    - simple wedge compression reversals
    """

    timeframe = "5m"
    startup_candle_count: int = 320

    stoploss = -0.12
    trailing_stop = False
    process_only_new_candles = True
    use_custom_stoploss = True
    use_exit_signal = True
    can_short = True

    minimal_roi = {
        "0": 0.25,
    }

    risk_per_trade = 0.03
    atr_period = 14
    cluster_lookback = 36
    kmeans_clusters = 3
    stop_lookback_candles = 3
    long_stop_buffer = 0.999
    short_stop_buffer = 1.001

    def _is_live_like(self) -> bool:
        runmode: Any = self.config.get("runmode") if self.config else None
        if hasattr(runmode, "value"):
            runmode = runmode.value
        return runmode in ("live", "dry_run")

    def _calculate_clusters(self, dataframe: DataFrame) -> DataFrame:
        if self._is_live_like():
            for column in ("cluster_support", "cluster_mid", "cluster_resistance"):
                if column not in dataframe.columns:
                    dataframe[column] = float("nan")

            if len(dataframe) < max(self.cluster_lookback, 10):
                dataframe["cluster_support"] = dataframe["low"].rolling(10, min_periods=1).min()
                dataframe["cluster_mid"] = dataframe["close"].rolling(10, min_periods=1).median()
                dataframe["cluster_resistance"] = dataframe["high"].rolling(10, min_periods=1).max()
                return dataframe

            try:
                price_points = (
                    dataframe[["high", "low", "close"]]
                    .tail(self.cluster_lookback)
                    .to_numpy()
                    .reshape(-1, 1)
                )
                model = KMeans(n_clusters=self.kmeans_clusters, n_init=5, random_state=42)
                model.fit(price_points)
                centers = sorted(model.cluster_centers_.flatten())

                dataframe.loc[dataframe.index[-1], "cluster_support"] = centers[0]
                dataframe.loc[dataframe.index[-1], "cluster_mid"] = centers[len(centers) // 2]
                dataframe.loc[dataframe.index[-1], "cluster_resistance"] = centers[-1]
            except Exception:
                dataframe.loc[dataframe.index[-1], "cluster_support"] = dataframe["low"].tail(10).min()
                dataframe.loc[dataframe.index[-1], "cluster_mid"] = dataframe["close"].tail(10).median()
                dataframe.loc[dataframe.index[-1], "cluster_resistance"] = dataframe["high"].tail(10).max()
        else:
            lb = self.cluster_lookback
            dataframe["cluster_support"] = dataframe["low"].rolling(lb, min_periods=10).quantile(0.2)
            dataframe["cluster_mid"] = dataframe["close"].rolling(lb, min_periods=10).median()
            dataframe["cluster_resistance"] = dataframe["high"].rolling(lb, min_periods=10).quantile(0.8)

        dataframe["cluster_support"] = dataframe["cluster_support"].ffill().bfill()
        dataframe["cluster_mid"] = dataframe["cluster_mid"].ffill().bfill()
        dataframe["cluster_resistance"] = dataframe["cluster_resistance"].ffill().bfill()
        return dataframe

    def _annotate_swings(self, dataframe: DataFrame) -> DataFrame:
        confirmed_high = np.where(
            (dataframe["high"].shift(2) > dataframe["high"].shift(3))
            & (dataframe["high"].shift(2) > dataframe["high"].shift(4))
            & (dataframe["high"].shift(2) >= dataframe["high"].shift(1))
            & (dataframe["high"].shift(2) >= dataframe["high"]),
            dataframe["high"].shift(2),
            np.nan,
        )
        confirmed_low = np.where(
            (dataframe["low"].shift(2) < dataframe["low"].shift(3))
            & (dataframe["low"].shift(2) < dataframe["low"].shift(4))
            & (dataframe["low"].shift(2) <= dataframe["low"].shift(1))
            & (dataframe["low"].shift(2) <= dataframe["low"]),
            dataframe["low"].shift(2),
            np.nan,
        )

        last_highs: list[float] = []
        prev_highs: list[float] = []
        last_lows: list[float] = []
        prev_lows: list[float] = []

        last_high = np.nan
        prev_high = np.nan
        last_low = np.nan
        prev_low = np.nan

        for high_val, low_val in zip(confirmed_high, confirmed_low):
            if not np.isnan(high_val):
                prev_high = last_high
                last_high = float(high_val)
            if not np.isnan(low_val):
                prev_low = last_low
                last_low = float(low_val)

            last_highs.append(last_high)
            prev_highs.append(prev_high)
            last_lows.append(last_low)
            prev_lows.append(prev_low)

        dataframe["confirmed_swing_high"] = confirmed_high
        dataframe["confirmed_swing_low"] = confirmed_low
        dataframe["last_swing_high"] = last_highs
        dataframe["prev_swing_high"] = prev_highs
        dataframe["last_swing_low"] = last_lows
        dataframe["prev_swing_low"] = prev_lows
        return dataframe

    def populate_indicators(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        dataframe["ema20"] = ta.EMA(dataframe, timeperiod=20)
        dataframe["atr14"] = ta.ATR(dataframe, timeperiod=self.atr_period)

        dataframe = self._calculate_clusters(dataframe)
        dataframe = self._annotate_swings(dataframe)

        candle_range = (dataframe["high"] - dataframe["low"]).replace(0, np.nan)
        dataframe["body"] = (dataframe["close"] - dataframe["open"]).abs()
        dataframe["lower_wick"] = np.minimum(dataframe["open"], dataframe["close"]) - dataframe["low"]
        dataframe["upper_wick"] = dataframe["high"] - np.maximum(dataframe["open"], dataframe["close"])
        dataframe["close_in_range"] = ((dataframe["close"] - dataframe["low"]) / candle_range).clip(0, 1).fillna(0.5)
        dataframe["ema_slope_5"] = dataframe["ema20"] - dataframe["ema20"].shift(5)

        dataframe["strong_bull"] = (
            (dataframe["close"] > dataframe["open"])
            & (
                (
                    (dataframe["body"] >= dataframe["atr14"] * 0.60)
                    & (dataframe["close_in_range"] >= 0.75)
                )
                |
                (
                    (dataframe["body"] >= dataframe["atr14"] * 0.45)
                    & (dataframe["lower_wick"] >= candle_range * 0.33)
                    & (dataframe["lower_wick"] <= candle_range * 0.50)
                    & (dataframe["close_in_range"] >= 0.50)
                )
            )
            & (dataframe["close"] > dataframe["high"].shift(1))
        )
        dataframe["strong_bear"] = (
            (dataframe["close"] < dataframe["open"])
            & (
                (
                    (dataframe["body"] >= dataframe["atr14"] * 0.60)
                    & (dataframe["close_in_range"] <= 0.25)
                )
                |
                (
                    (dataframe["body"] >= dataframe["atr14"] * 0.45)
                    & (dataframe["upper_wick"] >= candle_range * 0.33)
                    & (dataframe["upper_wick"] <= candle_range * 0.50)
                    & (dataframe["close_in_range"] <= 0.50)
                )
            )
            & (dataframe["close"] < dataframe["low"].shift(1))
        )

        dataframe["high_break"] = dataframe["high"] > dataframe["high"].shift(1)
        dataframe["low_break"] = dataframe["low"] < dataframe["low"].shift(1)

        dataframe["h1_strong"] = dataframe["high_break"] & dataframe["strong_bull"]
        dataframe["l1_strong"] = dataframe["low_break"] & dataframe["strong_bear"]

        prior_h_count = dataframe["h1_strong"].shift(1).rolling(10, min_periods=1).sum().fillna(0)
        prior_l_count = dataframe["l1_strong"].shift(1).rolling(10, min_periods=1).sum().fillna(0)
        dataframe["h2_signal"] = dataframe["high_break"] & dataframe["strong_bull"] & (prior_h_count >= 1)
        dataframe["h3_signal"] = dataframe["high_break"] & dataframe["strong_bull"] & (prior_h_count >= 2)
        dataframe["l2_signal"] = dataframe["low_break"] & dataframe["strong_bear"] & (prior_l_count >= 1)
        dataframe["l3_signal"] = dataframe["low_break"] & dataframe["strong_bear"] & (prior_l_count >= 2)

        dataframe["failed_low"] = (dataframe["low"] < dataframe["low"].shift(1)) & (
            dataframe["close"] >= dataframe["low"].shift(1)
        )
        dataframe["failed_high"] = (dataframe["high"] > dataframe["high"].shift(1)) & (
            dataframe["close"] <= dataframe["high"].shift(1)
        )
        dataframe["failed_low2"] = dataframe["failed_low"] & (
            dataframe["failed_low"].shift(1).rolling(8, min_periods=1).sum().fillna(0) >= 1
        )
        dataframe["failed_high2"] = dataframe["failed_high"] & (
            dataframe["failed_high"].shift(1).rolling(8, min_periods=1).sum().fillna(0) >= 1
        )

        atr_pad = dataframe["atr14"] * 0.35
        dataframe["in_support_zone"] = dataframe["low"] <= (dataframe["cluster_support"] + atr_pad)
        dataframe["below_support"] = dataframe["low"] < (dataframe["cluster_support"] - dataframe["atr14"] * 0.10)
        dataframe["in_resistance_zone"] = dataframe["high"] >= (dataframe["cluster_resistance"] - atr_pad)
        dataframe["above_resistance"] = dataframe["high"] > (dataframe["cluster_resistance"] + dataframe["atr14"] * 0.10)

        swing_pad = dataframe["atr14"] * 0.10
        dataframe["higher_highs"] = dataframe["last_swing_high"] > (dataframe["prev_swing_high"] + swing_pad)
        dataframe["higher_lows"] = dataframe["last_swing_low"] > (dataframe["prev_swing_low"] + swing_pad)
        dataframe["lower_highs"] = dataframe["last_swing_high"] < (dataframe["prev_swing_high"] - swing_pad)
        dataframe["lower_lows"] = dataframe["last_swing_low"] < (dataframe["prev_swing_low"] - swing_pad)

        dataframe["bull_channel_env"] = (
            dataframe["higher_highs"]
            & dataframe["higher_lows"]
            & (dataframe["ema_slope_5"] > dataframe["atr14"] * 0.05)
        )
        dataframe["bear_channel_env"] = (
            dataframe["lower_highs"]
            & dataframe["lower_lows"]
            & (dataframe["ema_slope_5"] < -dataframe["atr14"] * 0.05)
        )

        dataframe["channel_upper"] = dataframe["high"].rolling(20, min_periods=10).max()
        dataframe["channel_lower"] = dataframe["low"].rolling(20, min_periods=10).min()
        dataframe["touch_lower_channel"] = dataframe["low"] <= (
            dataframe["channel_lower"].shift(1) + dataframe["atr14"] * 0.25
        )
        dataframe["touch_upper_channel"] = dataframe["high"] >= (
            dataframe["channel_upper"].shift(1) - dataframe["atr14"] * 0.25
        )

        dataframe["flat_highs"] = (
            (dataframe["last_swing_high"] - dataframe["prev_swing_high"]).abs() <= dataframe["atr14"] * 0.60
        )
        dataframe["flat_lows"] = (
            (dataframe["last_swing_low"] - dataframe["prev_swing_low"]).abs() <= dataframe["atr14"] * 0.60
        )
        dataframe["range_env"] = (
            dataframe["flat_highs"]
            & dataframe["flat_lows"]
            & (dataframe["ema_slope_5"].abs() <= dataframe["atr14"] * 0.05)
        )
        dataframe["range_high"] = dataframe["high"].rolling(30, min_periods=15).max()
        dataframe["range_low"] = dataframe["low"].rolling(30, min_periods=15).min()
        range_span = (dataframe["range_high"] - dataframe["range_low"]).replace(0, np.nan)
        dataframe["range_position"] = ((dataframe["close"] - dataframe["range_low"]) / range_span).clip(0, 1).fillna(0.5)

        prior_above_ema_10 = (dataframe["low"].shift(1) > dataframe["ema20"].shift(1)).rolling(10).sum() == 10
        prior_below_ema_10 = (dataframe["high"].shift(1) < dataframe["ema20"].shift(1)).rolling(10).sum() == 10
        ema_touch_prev = (
            (dataframe["low"].shift(1) <= dataframe["ema20"].shift(1))
            | (dataframe["close"].shift(1) <= dataframe["ema20"].shift(1))
        )
        ema_touch_prev_short = (
            (dataframe["high"].shift(1) >= dataframe["ema20"].shift(1))
            | (dataframe["close"].shift(1) >= dataframe["ema20"].shift(1))
        )
        dataframe["ema_return_long"] = (
            prior_above_ema_10
            & ema_touch_prev
            & dataframe["strong_bull"].shift(1)
            & (dataframe["body"].shift(1) > dataframe["atr14"].shift(1))
            & (dataframe["close"] > dataframe["high"].shift(1))
        )
        dataframe["ema_return_short"] = (
            prior_below_ema_10
            & ema_touch_prev_short
            & dataframe["strong_bear"].shift(1)
            & (dataframe["body"].shift(1) > dataframe["atr14"].shift(1))
            & (dataframe["close"] < dataframe["low"].shift(1))
        )

        rolling_span = dataframe["high"].rolling(20, min_periods=10).max() - dataframe["low"].rolling(20, min_periods=10).min()
        dataframe["compressed_env"] = rolling_span < (rolling_span.shift(5) * 0.90)
        dataframe["double_top_or_lower_high"] = (
            (dataframe["high"] <= dataframe["high"].shift(1))
            | ((dataframe["high"] - dataframe["high"].shift(1)).abs() <= dataframe["atr14"] * 0.20)
        )
        dataframe["double_bottom_or_higher_low"] = (
            (dataframe["low"] >= dataframe["low"].shift(1))
            | ((dataframe["low"] - dataframe["low"].shift(1)).abs() <= dataframe["atr14"] * 0.20)
        )

        dataframe["cluster_long_setup"] = (
            dataframe["in_support_zone"]
            & (dataframe["h2_signal"] | dataframe["h3_signal"])
            & dataframe["strong_bull"]
            & ((~dataframe["below_support"]) | dataframe["failed_low2"])
        )
        dataframe["cluster_short_setup"] = (
            dataframe["in_resistance_zone"]
            & (dataframe["l2_signal"] | dataframe["l3_signal"])
            & dataframe["strong_bear"]
            & ((~dataframe["above_resistance"]) | dataframe["failed_high2"])
        )

        dataframe["channel_long_setup"] = (
            dataframe["bull_channel_env"]
            & dataframe["touch_lower_channel"].shift(1).rolling(4, min_periods=1).max().fillna(0).astype(bool)
            & (dataframe["h2_signal"] | dataframe["h3_signal"])
            & dataframe["strong_bull"]
        )
        dataframe["channel_short_setup"] = (
            dataframe["bear_channel_env"]
            & dataframe["touch_upper_channel"].shift(1).rolling(4, min_periods=1).max().fillna(0).astype(bool)
            & (dataframe["l2_signal"] | dataframe["l3_signal"])
            & dataframe["strong_bear"]
        )

        dataframe["range_long_setup"] = (
            dataframe["range_env"]
            & (dataframe["range_position"] <= 0.35)
            & (dataframe["h2_signal"] | dataframe["h3_signal"])
            & dataframe["strong_bull"]
        )
        dataframe["range_short_setup"] = (
            dataframe["range_env"]
            & (dataframe["range_position"] >= 0.65)
            & (dataframe["l2_signal"] | dataframe["l3_signal"])
            & dataframe["strong_bear"]
        )

        dataframe["wedge_long_setup"] = (
            dataframe["bear_channel_env"]
            & dataframe["compressed_env"]
            & dataframe["touch_lower_channel"]
            & dataframe["double_bottom_or_higher_low"]
            & dataframe["strong_bull"]
        )
        dataframe["wedge_short_setup"] = (
            dataframe["bull_channel_env"]
            & dataframe["compressed_env"]
            & dataframe["touch_upper_channel"]
            & dataframe["double_top_or_lower_high"]
            & dataframe["strong_bear"]
        )

        dataframe["impulse_up"] = (dataframe["close"] - dataframe["low"].rolling(8, min_periods=4).min()).clip(lower=dataframe["atr14"])
        dataframe["impulse_down"] = (dataframe["high"].rolling(8, min_periods=4).max() - dataframe["close"]).clip(lower=dataframe["atr14"])
        dataframe["ema_long_target"] = dataframe["close"] + dataframe["impulse_up"]
        dataframe["ema_short_target"] = dataframe["close"] - dataframe["impulse_down"]

        return dataframe

    def _absolute_stop_from_dataframe(self, dataframe: DataFrame, is_short: bool) -> float | None:
        if dataframe is None or dataframe.empty:
            return None

        if is_short:
            recent_high = dataframe["high"].tail(self.stop_lookback_candles).max()
            return float(recent_high) * self.short_stop_buffer

        recent_low = dataframe["low"].tail(self.stop_lookback_candles).min()
        return float(recent_low) * self.long_stop_buffer

    def _resolve_total_stake_capital(self, fallback_value: float) -> float:
        wallets = getattr(self, "wallets", None)
        if wallets is None:
            return fallback_value

        for attr_name in ("get_total_stake_amount", "get_available_stake_amount"):
            getter = getattr(wallets, attr_name, None)
            if callable(getter):
                try:
                    resolved = float(getter())
                except Exception:
                    continue
                if resolved > 0:
                    return resolved

        return fallback_value

    def custom_stake_amount(
        self,
        pair: str,
        current_time: datetime,
        current_rate: float,
        proposed_stake: float,
        min_stake: float | None,
        max_stake: float,
        leverage: float,
        entry_tag: str | None,
        side: str,
        **kwargs,
    ) -> float:
        dataframe, _ = self.dp.get_analyzed_dataframe(pair, self.timeframe)
        if dataframe is None or dataframe.empty or current_rate <= 0:
            return proposed_stake

        stop_abs = self._absolute_stop_from_dataframe(dataframe, side == "short")
        if stop_abs is None or stop_abs <= 0:
            return proposed_stake

        stop_distance_ratio = abs(current_rate - stop_abs) / current_rate
        leveraged_loss_ratio = stop_distance_ratio * max(1.0, float(leverage or 1.0))
        if leveraged_loss_ratio <= 0:
            return proposed_stake

        capital_base = self._resolve_total_stake_capital(float(proposed_stake))
        risk_budget = capital_base * self.risk_per_trade
        stake = risk_budget / leveraged_loss_ratio

        if min_stake is not None:
            stake = max(stake, float(min_stake))
        stake = min(stake, float(max_stake))
        return max(0.0, float(stake))

    def populate_entry_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        dataframe["enter_long"] = 0
        dataframe["enter_short"] = 0
        dataframe["enter_tag"] = None

        long_priority = [
            (dataframe["wedge_long_setup"], "wedge_long"),
            (dataframe["ema_return_long"], "ema_return_long"),
            (dataframe["channel_long_setup"], "channel_h2_long"),
            (dataframe["range_long_setup"], "range_h2_long"),
            (dataframe["cluster_long_setup"], "cluster_h2_long"),
        ]
        short_priority = [
            (dataframe["wedge_short_setup"], "wedge_short"),
            (dataframe["ema_return_short"], "ema_return_short"),
            (dataframe["channel_short_setup"], "channel_l2_short"),
            (dataframe["range_short_setup"], "range_l2_short"),
            (dataframe["cluster_short_setup"], "cluster_l2_short"),
        ]

        for condition, tag in long_priority:
            dataframe.loc[condition, ["enter_long", "enter_tag"]] = (1, tag)
        for condition, tag in short_priority:
            dataframe.loc[condition, ["enter_short", "enter_tag"]] = (1, tag)

        return dataframe

    def custom_stoploss(
        self,
        pair: str,
        trade: Trade,
        current_time: datetime,
        current_rate: float,
        current_profit: float,
        **kwargs,
    ) -> float:
        dataframe, _ = self.dp.get_analyzed_dataframe(pair, self.timeframe)
        if dataframe is None or dataframe.empty:
            return self.stoploss

        sl_abs = self._absolute_stop_from_dataframe(dataframe, trade.is_short)
        if sl_abs is None:
            return self.stoploss

        dyn_sl = stoploss_from_absolute(
            sl_abs,
            current_rate=current_rate,
            is_short=trade.is_short,
            leverage=max(1.0, float(getattr(trade, "leverage", 1.0))),
        )
        if dyn_sl is None:
            return self.stoploss

        return max(float(dyn_sl), self.stoploss)

    def _target_from_last_row(self, last: Any, is_short: bool, entry_tag: str) -> float | None:
        tag = (entry_tag or "").lower()
        if is_short:
            if "wedge" in tag or "channel" in tag:
                return float(last["channel_lower"])
            if "range" in tag:
                return float(last["range_low"])
            if "ema" in tag:
                return float(last["ema_short_target"])
            return float(last["cluster_support"])

        if "wedge" in tag or "channel" in tag:
            return float(last["channel_upper"])
        if "range" in tag:
            return float(last["range_high"])
        if "ema" in tag:
            return float(last["ema_long_target"])
        return float(last["cluster_resistance"])

    def custom_exit(
        self,
        pair: str,
        trade: Trade,
        current_time: datetime,
        current_rate: float,
        current_profit: float,
        **kwargs,
    ) -> str | None:
        dataframe, _ = self.dp.get_analyzed_dataframe(pair, self.timeframe)
        if dataframe is None or dataframe.empty:
            return None

        last = dataframe.iloc[-1]
        target = self._target_from_last_row(last, trade.is_short, getattr(trade, "enter_tag", ""))
        if target is None:
            return None

        if trade.is_short and current_rate <= target:
            return "target_short"
        if not trade.is_short and current_rate >= target:
            return "target_long"

        return None

    def populate_exit_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        dataframe["exit_long"] = 0
        dataframe["exit_short"] = 0
        return dataframe

    def confirm_trade_entry(
        self,
        pair: str,
        order_type: str,
        amount: float,
        rate: float,
        time_in_force: str,
        current_time: datetime,
        entry_tag: str | None,
        side: str,
        **kwargs,
    ) -> bool:
        dataframe, _ = self.dp.get_analyzed_dataframe(pair, self.timeframe)
        if dataframe is None or dataframe.empty:
            return False

        last = dataframe.iloc[-1]
        stop_abs = self._absolute_stop_from_dataframe(dataframe, side == "short")
        if stop_abs is None:
            return False

        target = self._target_from_last_row(last, side == "short", entry_tag or "")
        if target is None:
            return False

        risk = abs(rate - stop_abs)
        reward = abs(target - rate)
        if risk <= 0:
            return False

        return reward >= (risk * 1.2)