"""Live scanner using MACD+RSI strategy for signal generation.

Two-stage architecture:
  Stage 1 — MACD+RSI Confluence strategy finds candidate bars
  Stage 2 — (optional) Neural Network evaluates signal quality & SL/TP

Strategy:
    LONG:  RSI(14) rising 3+ bars AND MACD histogram crosses 0 from below
    SHORT: RSI(14) falling 3+ bars AND MACD histogram crosses 0 from above
    Exit:  TP=1.5xATR, SL=1.0xATR (or NN-predicted if NN available)

NN role (when loaded):
  - Evaluates signal quality: probability of reaching TP
  - Predicts optimal SL and TP distances (in ATR)
  - Only emits signals with quality >= nn_entry_threshold
"""

from __future__ import annotations

from datetime import datetime
from pathlib import Path
from typing import Dict, List, Optional
from zoneinfo import ZoneInfo

import numpy as np
import pandas as pd
from loguru import logger

from core.data_loader import DataPreparator
from ai.features import calculate_atr
from ai.strategy_signals import get_strategy_signals, StrategyConfig
from ai.inference_v2 import NeuralPredictor
from ai.barrier_methods import validate_barrier_levels, suggest_dynamic_barriers
from scanner.repository import InstrumentRepository

_MS = ZoneInfo("Europe/Moscow")


class StrategyScanner:
    """Scanner that uses MACD+RSI Confluence strategy for signal generation.

    Primary entry signals are rule-based (MACD+RSI). Neural network is optional
    and only used to evaluate signal quality and predict optimal SL/TP levels.

    Per-ticker models: when NN is available, automatically loads ticker-specific
    model from {models_dir}/{TICKER}_strategy.pt, falling back to the default model.

    Implements the same interface as NeuralScanner so it can be used
    interchangeably with the virtual trading infrastructure.

    Two modes:
      1. Strategy-only (no NN): fixed TP=1.5xATR, SL=1.0xATR
      2. Hybrid (strategy + NN): NN evaluates quality, provides SL/TP
    """

    def __init__(
        self,
        strategy_config: Optional[StrategyConfig] = None,
        tickers: Optional[List[str]] = None,
        neural_predictor: Optional[NeuralPredictor] = None,
        nn_entry_threshold: float = 0.5,
        models_dir: Optional[Path] = None,
        live_mode: bool = False,
    ):
        self.strategy_config = strategy_config or StrategyConfig()
        self.repo = InstrumentRepository()
        self.nn_entry_threshold = nn_entry_threshold
        self.live_mode = live_mode

        # Per-ticker predictor cache: {ticker: NeuralPredictor}
        self._predictor_cache: Dict[str, NeuralPredictor] = {}

        # Default predictor (fallback when per-ticker model not found)
        self._default_predictor = neural_predictor

        # Models directory (derived from default predictor path or default)
        # Resolve relative to this file's location: scanner/ -> moex_vsa_backtester/
        _scanner_dir = Path(__file__).resolve().parent.parent

        if models_dir:
            self._models_dir = Path(models_dir)
        elif neural_predictor and neural_predictor.trainer:
            model_path = getattr(neural_predictor, '_model_path', None)
            self._models_dir = model_path.parent if model_path else _scanner_dir / "models"
        else:
            self._models_dir = _scanner_dir / "models"

        # Default: TOP7 tickers (GAZP, MTSS, SNGSP, X5 excluded — consistently unprofitable OOS)
        self.default_tickers = tickers or [
            "SBER", "NVTK", "PHOR", "ROSN", "VTBR", "LKOH", "ASTR",
        ]

        mode = "HYBRID (per-ticker NN)" if neural_predictor else "STRATEGY-ONLY"
        if live_mode:
            mode += " [LIVE]"
        logger.info(
            f"StrategyScanner [{mode}]: "
            f"MACD fast={self.strategy_config.macd_fast}, "
            f"slow={self.strategy_config.macd_slow}, "
            f"signal={self.strategy_config.macd_signal}, "
            f"RSI period={self.strategy_config.rsi_period}, "
            f"TP={self.strategy_config.tp_atr_mult}xATR, "
            f"SL={self.strategy_config.sl_atr_mult}xATR"
        )

    @property
    def has_neural_predictor(self) -> bool:
        """Whether this scanner has a neural predictor (HYBRID mode)."""
        return self._default_predictor is not None

    def _get_predictor(self, ticker: str) -> Optional[NeuralPredictor]:
        """Get or create a NeuralPredictor for a specific ticker.

        Loads per-ticker model from {models_dir}/{TICKER}_strategy.pt.
        Falls back to the default predictor if per-ticker model doesn't exist.

        Args:
            ticker: Ticker symbol (e.g. "SBER")

        Returns:
            NeuralPredictor instance or None if no model available
        """
        # Return cached predictor if available
        if ticker in self._predictor_cache:
            return self._predictor_cache[ticker]

        # Try per-ticker model first
        per_ticker_path = self._models_dir / f"{ticker}_strategy.pt"
        if per_ticker_path.exists():
            logger.info(f"[{ticker}] Loading per-ticker model: {per_ticker_path.name}")
            predictor = NeuralPredictor(
                model_path=per_ticker_path,
                entry_threshold=0.0,
                min_confidence=0.0,
            )
            self._predictor_cache[ticker] = predictor
            return predictor

        # Fall back to default predictor
        if self._default_predictor is not None:
            logger.debug(f"[{ticker}] No per-ticker model, using default")
            self._predictor_cache[ticker] = self._default_predictor
            return self._default_predictor

        return None

    def get_all_instruments(self) -> List[str]:
        """Return available instruments (filtered to strategy tickers)."""
        all_instruments = self.repo.get_all_instruments()
        # Filter to known strategy tickers
        strategy_set = set(t.upper() for t in self.default_tickers)
        return [t for t in all_instruments if t.upper() in strategy_set]

    def get_latest_timestamp(self, ticker: str) -> Optional[int]:
        return self.repo.get_latest_timestamp(ticker)

    def get_latest_price(self, ticker: str) -> Optional[float]:
        return self.repo.get_latest_price(ticker)

    def get_latest_candle_state(self, ticker: str) -> Optional[Dict]:
        """Get latest timestamp and close price for change detection."""
        return self.repo.get_latest_candle_state(ticker)

    def scan_instrument(
        self,
        ticker: str,
        lookback_hours: int = 72,
        sides: Optional[List[str]] = None,
    ) -> List[Dict]:
        """Scan a single instrument for MACD+RSI signals.

        Two-stage pipeline:
          1. MACD+RSI strategy identifies candidate bars
          2. If NN is loaded, evaluates each candidate's quality and SL/TP

        Args:
            ticker: Ticker symbol
            lookback_hours: How many hours of H1 data to load
            sides: Which sides to include (None = both LONG and SHORT)

        Returns:
            List of signal dictionaries (same format as NeuralScanner)
        """
        end_ts = int(datetime.now(_MS).timestamp())
        start_ts = end_ts - lookback_hours * 3600

        try:
            prep = DataPreparator([ticker], ["H1"])
            df_h1 = prep.load_h1(ticker, start_ts, end_ts)

            if df_h1.empty or len(df_h1) < 50:
                return []

            # --- Stage 1: MACD+RSI strategy signals ---
            signals_df = get_strategy_signals(df_h1, self.strategy_config)

            if signals_df.empty:
                return []

            # Merge signals into OHLCV data
            df_h1 = df_h1.sort_values("timestamp").reset_index(drop=True)
            signals_df = signals_df.sort_values("timestamp").reset_index(drop=True)

            df_h1["strategy_signal"] = signals_df["strategy_signal"].values
            df_h1["strategy_side"] = signals_df["side"].values
            df_h1["signal_strength"] = signals_df["signal_strength"].values

            # Calculate ATR (needed for both modes)
            atr = calculate_atr(df_h1, 14)

            all_signals = []

            # Get per-ticker predictor (loads per-ticker model on demand)
            predictor = self._get_predictor(ticker)
            nn_available = predictor is not None and predictor.trainer is not None

            # In live_mode, only process the latest completed bar
            if self.live_mode:
                # Find the latest bar with a strategy signal
                signal_indices = df_h1.index[df_h1["strategy_signal"] == 1].tolist()
                if not signal_indices:
                    return []
                # Only process the most recent signal
                indices_to_process = [signal_indices[-1]]
            else:
                # Backtest mode: process all signals
                indices_to_process = range(len(df_h1))

            for i in indices_to_process:
                if df_h1.iloc[i]["strategy_signal"] != 1:
                    continue

                side = df_h1.iloc[i]["strategy_side"]
                if sides and side not in sides:
                    continue

                entry_price = df_h1.iloc[i]["Close"]
                bar_ts = df_h1.iloc[i]["timestamp"]
                atr_val = atr.iloc[i] if i < len(atr) else atr.iloc[-1]

                # --- Stage 2: NN quality evaluation (optional) ---
                if nn_available:
                    nn_result = predictor.evaluate_strategy_signal(
                        df_h1, i, side
                    )

                    if nn_result is None:
                        # NN evaluation failed — skip this signal
                        logger.debug(
                            f"[{ticker}] NN evaluation failed at bar {i}, skipping"
                        )
                        continue

                    nn_proba = nn_result["entry_proba"]
                    nn_confidence = nn_result["confidence"]

                    # NN quality filter: skip if NN thinks this is a bad signal
                    if nn_proba < self.nn_entry_threshold:
                        logger.debug(
                            f"[{ticker}] NN REJECTED {side} signal: "
                            f"proba={nn_proba:.2%} < threshold={self.nn_entry_threshold:.2%}"
                        )
                        continue

                    # Use NN-predicted SL/TP distances
                    sl_atr = nn_result["sl_distance_atr"]
                    tp_atr = nn_result["tp_distance_atr"]
                    strength = nn_proba  # NN probability as signal strength

                    if side == "LONG":
                        raw_sl = entry_price - atr_val * sl_atr
                        raw_tp = entry_price + atr_val * tp_atr
                    else:
                        raw_sl = entry_price + atr_val * sl_atr
                        raw_tp = entry_price - atr_val * tp_atr

                    # Validate barriers: enforce RR & prevent degenerate levels
                    sl_price, tp_price = validate_barrier_levels(
                        sl_price=raw_sl, tp_price=raw_tp,
                        entry_price=entry_price, side=side, atr_val=atr_val,
                    )
                    # Report effective (post-validation) ATR multipliers
                    if side == "LONG":
                        eff_sl_atr = (entry_price - sl_price) / (atr_val + 1e-10)
                        eff_tp_atr = (tp_price - entry_price) / (atr_val + 1e-10)
                    else:
                        eff_sl_atr = (sl_price - entry_price) / (atr_val + 1e-10)
                        eff_tp_atr = (entry_price - tp_price) / (atr_val + 1e-10)

                    signal = self._format_signal(
                        ticker=ticker,
                        timestamp=bar_ts,
                        side=side,
                        entry_price=entry_price,
                        sl_price=sl_price,
                        tp_price=tp_price,
                        strength=strength,
                        sl_atr=eff_sl_atr,
                        tp_atr=eff_tp_atr,
                        nn_filtered=True,
                        confidence=nn_confidence,
                    )
                else:
                    # Mode 1: strategy-only — SL/TP from config
                    strength = df_h1.iloc[i]["signal_strength"]
                    
                    if self.strategy_config.enable_dynamic_barriers:
                        # Volatility-adaptive barriers (includes RR enforcement)
                        barrier_result = suggest_dynamic_barriers(
                            entry_price=entry_price,
                            side=side,
                            atr_val=atr_val,
                            atr_series=atr,
                            df=df_h1,
                            idx=i,
                            use_structure=True,
                        )
                        sl_price = barrier_result.sl_price
                        tp_price = barrier_result.tp_price
                        sl_atr = barrier_result.sl_multiplier
                        tp_atr = barrier_result.tp_multiplier
                    else:
                        # Fixed ATR multipliers with per-ticker TP + enforce RR
                        sl_atr = self.strategy_config.sl_atr_mult
                        tp_atr = self.strategy_config.get_tp_atr_mult(ticker)
                        if side == "LONG":
                            raw_sl = entry_price - atr_val * sl_atr
                            raw_tp = entry_price + atr_val * tp_atr
                        else:
                            raw_sl = entry_price + atr_val * sl_atr
                            raw_tp = entry_price - atr_val * tp_atr
                        sl_price, tp_price = validate_barrier_levels(
                            sl_price=raw_sl, tp_price=raw_tp,
                            entry_price=entry_price, side=side, atr_val=atr_val,
                        )
                        # Effective multipliers after RR enforcement
                        if side == "LONG":
                            sl_atr = (entry_price - sl_price) / (atr_val + 1e-10)
                            tp_atr = (tp_price - entry_price) / (atr_val + 1e-10)
                        else:
                            sl_atr = (sl_price - entry_price) / (atr_val + 1e-10)
                            tp_atr = (entry_price - tp_price) / (atr_val + 1e-10)

                    signal = self._format_signal(
                        ticker=ticker,
                        timestamp=bar_ts,
                        side=side,
                        entry_price=entry_price,
                        sl_price=sl_price,
                        tp_price=tp_price,
                        strength=strength,
                        sl_atr=sl_atr,
                        tp_atr=tp_atr,
                        nn_filtered=False,
                    )

                all_signals.append(signal)

            if all_signals:
                n_rejected = (
                    len([s for s in all_signals if s.get("nn_filtered", False)])
                    if nn_available
                    else 0
                )
                logger.debug(
                    f"[{ticker}] TRADING SIGNALS: {len(all_signals)} "
                    f"({'NN-filtered' if nn_available else 'strategy-only'}) "
                    f"({sum(1 for s in all_signals if s['signal_type']=='LONG')} LONG, "
                    f"{sum(1 for s in all_signals if s['signal_type']=='SHORT')} SHORT)"
                )

            return all_signals

        except Exception as e:
            logger.error(f"Error scanning {ticker}: {e}")
            return []

    def scan_all_instruments(
        self,
        lookback_hours: int = 72,
        instruments: Optional[List[str]] = None,
        sides: Optional[List[str]] = None,
    ) -> List[Dict]:
        """Scan all instruments for MACD+RSI signals.

        Args:
            lookback_hours: Hours of H1 data to load per ticker
            instruments: List of tickers (default: strategy tickers)
            sides: Which sides to include

        Returns:
            List of all signals found
        """
        if instruments is None:
            instruments = self.get_all_instruments()

        all_signals = []
        for ticker in instruments:
            signals = self.scan_instrument(ticker, lookback_hours, sides)
            for sig in signals:
                self.log_signal(sig)
                all_signals.append(sig)

        return all_signals

    def _format_signal(
        self,
        ticker: str,
        timestamp: int,
        side: str,
        entry_price: float,
        sl_price: float,
        tp_price: float,
        strength: float,
        sl_atr: float,
        tp_atr: float,
        nn_filtered: bool = False,
        confidence: Optional[float] = None,
    ) -> Dict:
        """Format a signal dictionary (matches NeuralScanner format).

        Args:
            confidence: Override for model confidence. If None, uses ``strength``
                        (strategy-only mode, where strength ≈ confidence).
        """
        if confidence is None:
            confidence = strength
        return {
            "ticker": ticker,
            "signal_time": datetime.fromtimestamp(timestamp, _MS).strftime(
                "%Y-%m-%d %H:%M:%S"
            ),
            "timestamp": timestamp,
            "signal_type": side,
            "entry_price": entry_price,
            "sl_price": sl_price,
            "tp_price": tp_price,
            "entry_probability": strength,
            "confidence": confidence,
            "sl_distance_atr": sl_atr,
            "tp_distance_atr": tp_atr,
            "nn_filtered": nn_filtered,
        }

    def log_signal(self, signal: Dict):
        """Log signal to console in compact single-line format."""
        nn_tag = "NN-Q" if signal.get("nn_filtered") else "STRAT"
        ts = signal['signal_time']
        logger.info(
            f"[{nn_tag}] {signal['ticker']} {signal['signal_type']} @ "
            f"{signal['entry_price']:.2f} | "
            f"SL={signal['sl_price']:.2f} ({signal.get('sl_distance_atr', 0):.2f}x) "
            f"TP={signal['tp_price']:.2f} ({signal.get('tp_distance_atr', 0):.2f}x) | "
            f"Str={signal.get('entry_probability', 0):.0%} [{ts}]"
        )


