# source: https://raw.githubusercontent.com/Mateusz-Nalezinski/CryptoBOT/ff191ca57447739040d01675b834056eefbc1e6f/services/freqtrade/user_data/strategies/RegimeHybridFreqAIV0.py
from datetime import datetime
from typing import Dict, List, Tuple

import numpy as np
import pandas as pd
from freqtrade.strategy import IStrategy
from pandas import DataFrame

from mods import exits, filters, indicators, logging as decision_log, protections as prot, regime, signals_range, signals_trend


class Github_Mateusz_Nalezinski_CryptoBOT__RegimeHybridFreqAIV0__20251220_160211(IStrategy):
    """
    RegimeHybrid z filtrem ML przez FreqAI (LightGBMClassifier).
    - MTF i wskaźniki wyliczane w feature_engineering_*.
    - populate_indicators wywołuje self.freqai.start, które dodaje predykcje/pola &-quality.
    - Wejścia jak w RegimeHybridV0 + dodatkowy filtr ML (good & do_predict==1).
    """

    timeframe = "30m"
    informative_timeframes: List[str] = ["1h", "4h"]
    process_only_new_candles = True
    startup_candle_count: int = 210
    use_exit_signal = True
    exit_profit_only = False
    ignore_buying_expired_candle_after = 20
    minimal_roi = {"0": 0.04}
    stoploss = -0.08
    entry_context: Dict[str, Dict] = {}
    protections = prot.DEFAULT_PROTECTIONS
    # Parametry ML label
    ml_tp: float = 0.015
    ml_sl: float = 0.01

    def informative_pairs(self) -> List[Tuple[str, str]]:
        pairs = self.dp.current_whitelist() if self.dp else []
        return [(pair, tf) for pair in pairs for tf in self.informative_timeframes]

    def populate_indicators(self, dataframe: DataFrame, metadata: Dict) -> DataFrame:
        # FreqAI pipeline obsługuje feature_engineering_* i targets
        return self.freqai.start(dataframe, metadata, self)

    # Feature Engineering (FreqAI contract)
    def feature_engineering_expand_all(self, df: DataFrame, period: int, **kwargs) -> DataFrame:
        # cechy zależne od period
        df[f"%ema_{period}"] = df["close"].ewm(span=period, adjust=False).mean()
        df[f"%rsi_{period}"] = indicators._rsi(df["close"], period)
        df[f"%atr_{period}"] = indicators._atr(df, period)
        df[f"%roc_{period}"] = df["close"].pct_change(periods=period)
        return df

    def feature_engineering_expand_basic(self, df: DataFrame, **kwargs) -> DataFrame:
        df["%pct_change_1"] = df["close"].pct_change()
        df["%volume"] = df["volume"]
        df["%bb_width"] = (df["close"].rolling(20).mean() + 2 * df["close"].rolling(20).std()) - (
            df["close"].rolling(20).mean() - 2 * df["close"].rolling(20).std()
        )
        df["%bb_pos"] = 0.0
        upper = df["close"].rolling(20).mean() + 2 * df["close"].rolling(20).std()
        lower = df["close"].rolling(20).mean() - 2 * df["close"].rolling(20).std()
        df["%bb_pos"] = (df["close"] - lower) / (upper - lower)
        vol_sma = df["volume"].rolling(20, min_periods=5).mean()
        df["%volume_rel"] = df["volume"] / vol_sma.replace(0, np.nan)
        return df

    def feature_engineering_standard(self, df: DataFrame, **kwargs) -> DataFrame:
        # cechy czasowe
        if "date" not in df.columns:
            df["date"] = df.index
        df["%-hour_of_day"] = df["date"].dt.hour
        df["%-day_of_week"] = df["date"].dt.dayofweek
        # kluczowe wskaźniki klasyczne do reżimu/sygnałów
        df = indicators.add_base_indicators(df)
        df["ema50_slope"] = df["ema50"].diff(3) / df["ema50"].shift(3).replace(0, np.nan)
        df["dist_ema50"] = (df["close"] - df["ema50"]) / df["ema50"].replace(0, np.nan)
        if "regime" not in df.columns:
            df["regime"] = regime.detect_regime(df)
        return df

    def set_freqai_targets(self, df: DataFrame, **kwargs) -> DataFrame:
        # Label na podstawie przyszłego ruchu (forward window = label_period_candles)
        period = int(self.freqai_info.get("feature_parameters", {}).get("label_period_candles", 24))
        future_max = df["high"].shift(-period).rolling(period, min_periods=period).max()
        future_min = df["low"].shift(-period).rolling(period, min_periods=period).min()
        entry_price = df["close"]
        tp_hit = (future_max - entry_price) / entry_price >= self.ml_tp
        sl_hit = (future_min - entry_price) / entry_price <= -self.ml_sl
        labels = pd.Series("bad", index=df.index)
        labels = labels.mask(tp_hit & ~sl_hit, "good")
        labels = labels.mask(sl_hit & ~tp_hit, "bad")
        df["&-quality"] = labels
        return df

    def populate_entry_trend(self, dataframe: DataFrame, metadata: Dict) -> DataFrame:
        df = dataframe.copy()
        vol_ok = filters.volume_ok(df, min_rel=1.1)
        calm = filters.no_high_vol(df)
        signals_tr = signals_trend.trend_long_signals(df)
        signals_rg = signals_range.range_long_signals(df)

        trend_mask = df["regime"].isin([regime.TREND_UP, regime.TREND_DOWN]) & signals_tr
        range_mask = (df["regime"] == regime.RANGE) & signals_rg
        ml_good = (df.get("&-quality") == "good") & (df.get("do_predict", 0) == 1)
        ml_enabled = bool(self.freqai_info.get("enabled", False))
        ml_filter = ml_good if ml_enabled else True

        entry_mask = (trend_mask | range_mask) & vol_ok & calm & filters.daily_trade_cap(df) & ml_filter

        df.loc[:, "enter_long"] = 0
        df.loc[entry_mask, "enter_long"] = 1
        tag_status = "ml_ok" if ml_enabled else "ml_off"
        df.loc[entry_mask, "enter_tag"] = decision_log.short_tag(df.loc[entry_mask, "regime"], "hybrid_ml", tag_status)

        if metadata and metadata.get("pair"):
            pair = metadata["pair"]
            idx = entry_mask[entry_mask].index
            if not idx.empty:
                row = df.loc[idx[-1]]
                engine = "trend" if trend_mask.loc[idx[-1]] else "range"
                self.entry_context[pair] = {
                    "regime": row.get("regime"),
                    "engine": engine,
                    "filters": {"volume_ok": bool(vol_ok.iloc[idx[-1]]), "calm": bool(calm.iloc[idx[-1]]), "ml": bool(ml_filter.iloc[idx[-1]]) if hasattr(ml_filter, "iloc") else bool(ml_filter)},
                    "indicators": {
                        "rsi14": float(row.get("rsi14", 0)),
                        "bb_pos": float(row.get("bb_position", 0)),
                        "volume_rel": float(row.get("volume_rel", 0)),
                        "ml_quality": row.get("&-quality"),
                        "do_predict": row.get("do_predict", 0),
                        "di_values": row.get("DI_values"),
                    },
                    "freqai_identifier": getattr(self.freqai, "model_identifier", None),
                    "ml_quality_pred": row.get("&-quality"),
                }
        decision_log.decision_log(
            event="entry",
            regime="hybrid_ml",
            reason="dispatcher_ml",
            extra={
                "accepted": int(entry_mask.sum()),
                "ts": datetime.utcnow().isoformat(),
                "ml_enabled": ml_enabled,
            },
        )
        return df

    def populate_exit_trend(self, dataframe: DataFrame, metadata: Dict) -> DataFrame:
        df = dataframe.copy()
        trend_exit_mask = df["regime"].isin([regime.TREND_UP, regime.TREND_DOWN]) & exits.trend_exit(df)
        range_exit_mask = (df["regime"] == regime.RANGE) & exits.range_exit(df)
        exit_mask = trend_exit_mask | range_exit_mask
        df.loc[:, "exit_long"] = 0
        df.loc[exit_mask, "exit_long"] = 1
        decision_log.decision_log(
            event="exit",
            regime="hybrid_ml",
            reason="dispatcher_exit",
            extra={"trend_exits": int(trend_exit_mask.sum()), "range_exits": int(range_exit_mask.sum())},
        )
        return df

    def custom_stoploss(self, pair: str, trade, current_time: datetime, current_rate: float, current_profit: float, **kwargs) -> float:
        return self.stoploss

    def confirm_trade_entry(self, pair: str, order_type: str, amount: float, rate: float, time_in_force: str, **kwargs):
        # zapis custom_data dla wejścia
        try:
            if pair in self.entry_context:
                kwargs.get("trade", None).set_custom_data("decision_entry", self.entry_context[pair])
        except Exception:
            pass
        return True

    def custom_exit(self, pair: str, trade, current_time: datetime, current_rate: float, current_profit: float, **kwargs):
        try:
            trade.set_custom_data("decision_exit", {"exit_at": current_time.isoformat(), "profit": current_profit, "regime": trade.enter_tag if hasattr(trade, "enter_tag") else None, "freqai_identifier": getattr(self.freqai, "model_identifier", None)})
        except Exception:
            pass
        return None
