"""Inference pipeline for neural network entry/SL/TP prediction.

Loads trained model and generates trading signals with SL/TP levels.
"""

from __future__ import annotations

from pathlib import Path
from typing import Dict, List, Optional

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

from ai.features import (
    CONTEXT_WINDOW,
    add_technicals,
    calculate_atr,
    extract_feature_vector,
)
from ai.trainer_v2 import MultiTaskTrainer
from ai.barrier_methods import validate_barrier_levels
from domain import NeuralSignal


class NeuralPredictor:
    """Predictor that generates entry/SL/TP signals using trained neural network."""
    
    def __init__(
        self,
        model_path: Optional[Path] = None,
        entry_threshold: float = 0.5,
        min_confidence: float = 0.3,
        context_window: int = CONTEXT_WINDOW,
        ticker_id: int = 0,
        n_tickers: int = 1,
    ):
        self.entry_threshold = entry_threshold
        self.min_confidence = min_confidence
        self.context_window = context_window
        self._ticker_id = ticker_id
        self._n_tickers = n_tickers
        self.trainer: Optional[MultiTaskTrainer] = None
        
        self._model_path = model_path
        if model_path and model_path.exists():
            self.load_model(model_path)
        else:
            logger.warning(f"Модель не найдена: {model_path}")
    
    def load_model(self, model_path: Path):
        """Load trained model from disk.
        
        Args:
            model_path: Path to saved model
        """
        self.trainer = MultiTaskTrainer()
        self.trainer.load(model_path)
        logger.debug(f"Predictor model loaded: {model_path}")
    
    def predict_signals(
        self,
        df: pd.DataFrame,
        side: Optional[str] = None,
    ) -> List[NeuralSignal]:
        """Generate trading signals for all bars in DataFrame.
        
        NEW ARCHITECTURE: Uses two separate entry heads (LONG/SHORT).
        When called without `side`, returns BOTH LONG and SHORT signals.
        When called with a specific side, returns only that side's signals.
        
        Args:
            df: OHLCV DataFrame with columns [timestamp, Open, High, Low, Close, Volume]
            side: Optional filter ("LONG" or "SHORT"). If None, returns both.
            
        Returns:
            List of NeuralSignal objects
        """
        if self.trainer is None:
            logger.warning("Модель не загружена")
            return []
        
        if df.empty or len(df) < self.context_window + 10:
            logger.warning(f"Недостаточно данных: {len(df)} баров")
            return []
        
        df = df.sort_values("timestamp").reset_index(drop=True)
        df = add_technicals(df)
        df["ATR"] = calculate_atr(df, 14)
        
        all_signals = []
        
        for idx in range(self.context_window, len(df) - 1):
            signals = self._predict_single(df, idx)
            if signals:
                all_signals.extend(signals)
        
        if side:
            all_signals = [s for s in all_signals if s.side == side]
        
        logger.debug(f"Generated {len(all_signals)} signals (total)")
        return all_signals
    
    def _predict_single(
        self,
        df: pd.DataFrame,
        idx: int,
    ) -> List[NeuralSignal]:
        """Generate signals for a single bar using both LONG and SHORT heads.
        
        ONE forward pass gives both LONG and SHORT probabilities.
        Returns signals for sides that pass the entry threshold.
        
        Args:
            df: Full OHLCV DataFrame with technicals
            idx: Index of current bar
            
        Returns:
            List of NeuralSignal objects (0, 1, or 2 signals)
        """
        start_idx = idx - self.context_window
        window_df = df.iloc[start_idx:idx].copy()
        
        if len(window_df) < self.context_window:
            return []
        
        atr_val = df["ATR"].iloc[idx]
        if pd.isna(atr_val) or atr_val <= 0:
            return []
        
        features = self._extract_features(window_df, df.iloc[idx], "LONG", atr_val)
        if features is None:
            return []
        
        features = features.reshape(1, -1)
        
        try:
            predictions = self.trainer.predict(features, entry_threshold=self.entry_threshold)
        except Exception as e:
            logger.debug(f"Prediction failed: {e}")
            return []
        
        sl_distance_atr = predictions["sl_distance"][0]
        tp_distance_atr = predictions["tp_distance"][0]
        confidence = predictions["confidence"][0]
        
        entry_long_proba = predictions["entry_long_proba"][0]
        entry_short_proba = predictions["entry_short_proba"][0]
        entry_long_signal = predictions["entry_long_signal"][0]
        entry_short_signal = predictions["entry_short_signal"][0]
        
        if confidence < self.min_confidence:
            return []
        
        entry_price = df["Close"].iloc[idx]
        sl_distance_price = sl_distance_atr * atr_val
        tp_distance_price = tp_distance_atr * atr_val
        timestamp = int(df["timestamp"].iloc[idx])
        
        signals = []
        
        if entry_long_signal == 1:
            raw_sl = entry_price - sl_distance_price
            raw_tp = entry_price + tp_distance_price
            # Validate barriers: ensure minimum distances, enforce RR
            sl_price, tp_price = validate_barrier_levels(
                sl_price=raw_sl, tp_price=raw_tp,
                entry_price=entry_price, side="LONG", atr_val=atr_val,
            )
            # Report effective (post-validation) ATR multipliers
            effective_sl_atr = (entry_price - sl_price) / (atr_val + 1e-10)
            effective_tp_atr = (tp_price - entry_price) / (atr_val + 1e-10)
            signals.append(NeuralSignal(
                timestamp=timestamp,
                side="LONG",
                entry_price=entry_price,
                sl_price=sl_price,
                tp_price=tp_price,
                sl_distance_atr=effective_sl_atr,
                tp_distance_atr=effective_tp_atr,
                entry_probability=entry_long_proba,
                confidence=confidence,
                atr_at_signal=atr_val,
            ))
        
        if entry_short_signal == 1:
            raw_sl = entry_price + sl_distance_price
            raw_tp = entry_price - tp_distance_price
            # Validate barriers: ensure minimum distances, enforce RR
            sl_price, tp_price = validate_barrier_levels(
                sl_price=raw_sl, tp_price=raw_tp,
                entry_price=entry_price, side="SHORT", atr_val=atr_val,
            )
            # Report effective (post-validation) ATR multipliers
            effective_sl_atr = (sl_price - entry_price) / (atr_val + 1e-10)
            effective_tp_atr = (entry_price - tp_price) / (atr_val + 1e-10)
            signals.append(NeuralSignal(
                timestamp=timestamp,
                side="SHORT",
                entry_price=entry_price,
                sl_price=sl_price,
                tp_price=tp_price,
                sl_distance_atr=effective_sl_atr,
                tp_distance_atr=effective_tp_atr,
                entry_probability=entry_short_proba,
                confidence=confidence,
                atr_at_signal=atr_val,
            ))
        
        return signals
    
    def _extract_features(
        self,
        window_df: pd.DataFrame,
        current_bar: pd.Series,
        side: str,
        atr_val: float,
    ) -> Optional[np.ndarray]:
        """Extract feature vector for prediction.
        
        SINGLE SOURCE OF TRUTH: delegates to ai.features.extract_feature_vector().
        
        Args:
            window_df: Context window of OHLCV data
            current_bar: Current bar data
            side: Direction
            atr_val: Current ATR value
            
        Returns:
            Feature vector or None
        """
        entry_time = current_bar.get("timestamp", 0)
        return extract_feature_vector(
            window_df=window_df,
            side=side,
            atr_val=atr_val,
            ticker_id=self._ticker_id,
            n_tickers=self._n_tickers,
            entry_time=int(entry_time) if not pd.isna(entry_time) else None,
            context_window=self.context_window,
        )
    
    def get_latest_prediction(
        self,
        df: pd.DataFrame,
        side: Optional[str] = None,
    ) -> Optional[Dict]:
        """Get raw model prediction for the latest bar without threshold filtering.

        Returns the model's raw outputs from BOTH LONG and SHORT entry heads,
        plus confidence, SL/TP distances — even if they don't meet thresholds.

        When `side` is specified, returns only that side's prediction dict.
        When `side` is None, returns a combined dict with both LONG and SHORT.

        Args:
            df: OHLCV DataFrame
            side: Optional filter ("LONG" or "SHORT"). If None, returns both.

        Returns:
            Dict with raw prediction(s), or None if prediction fails
        """
        if self.trainer is None:
            return None

        if df.empty or len(df) < self.context_window + 5:
            return None

        df_sorted = df.sort_values("timestamp").reset_index(drop=True)
        df_sorted = add_technicals(df_sorted)
        df_sorted["ATR"] = calculate_atr(df_sorted, 14)

        # Use second-to-last bar (last closed bar), matching predict_signals loop
        idx = len(df_sorted) - 2
        if idx < self.context_window:
            return None

        start_idx = idx - self.context_window
        window_df = df_sorted.iloc[start_idx:idx].copy()

        atr_val = df_sorted["ATR"].iloc[idx]
        if pd.isna(atr_val) or atr_val <= 0:
            return None

        # Single feature extraction (uses side="LONG" but side bit is now secondary)
        features = self._extract_features(window_df, df_sorted.iloc[idx], "LONG", atr_val)
        if features is None:
            return None

        features = features.reshape(1, -1)

        try:
            predictions = self.trainer.predict(features, entry_threshold=0.0)
        except Exception:
            return None

        result = {
            "confidence": float(predictions["confidence"][0]),
            "sl_distance_atr": float(predictions["sl_distance"][0]),
            "tp_distance_atr": float(predictions["tp_distance"][0]),
            "entry_price": float(df_sorted["Close"].iloc[idx]),
            "atr": float(atr_val),
            "entry_long_proba": float(predictions["entry_long_proba"][0]),
            "entry_short_proba": float(predictions["entry_short_proba"][0]),
        }
        
        # Backward compatibility
        result["entry_proba"] = result["entry_long_proba"]

        if side == "LONG":
            return {k: v for k, v in result.items() if k != "entry_short_proba"}
        elif side == "SHORT":
            return {k: v for k, v in result.items() if k != "entry_long_proba"}
        
        return result

    def predict_latest(
        self,
        df: pd.DataFrame,
        side: Optional[str] = None,
    ) -> Optional[NeuralSignal]:
        """Generate signals for the latest bar only.
        
        Args:
            df: OHLCV DataFrame
            side: Optional filter ("LONG" or "SHORT"). If None, returns the
                  higher-probability signal.
            
        Returns:
            NeuralSignal for latest bar, or the best one if both, or None
        """
        signals = self.predict_signals(df.tail(self.context_window + 5), side=None)
        if not signals:
            return None
        if side:
            side_signals = [s for s in signals if s.side == side]
            return side_signals[-1] if side_signals else None
        # Return the highest probability signal
        return max(signals, key=lambda s: s.entry_probability)
    
    def evaluate_strategy_signal(
        self,
        df: pd.DataFrame,
        idx: int,
        strategy_side: str,
    ) -> Optional[Dict]:
        """Evaluate a MACD+RSI strategy candidate bar with NN.

        Runs inference on the context window ending at `idx` and returns
        raw NN outputs WITHOUT the entry threshold filter — the strategy
        already decided this is a candidate. The NN provides:
          - quality score (entry_proba): probability of hitting TP
          - optimal SL/TP distances (in ATR)
          - model confidence

        Args:
            df: Full OHLCV DataFrame with at least context_window bars before idx
            idx: Index of the strategy signal bar (where MACD+RSI fired)
            strategy_side: "LONG" or "SHORT" (from strategy, NN doesn't change it)

        Returns:
            Dict with keys:
              - entry_proba (float): NN quality score for this bar [0-1]
              - confidence (float): NN confidence [0-1]
              - sl_distance_atr (float): optimal SL distance in ATR
              - tp_distance_atr (float): optimal TP distance in ATR
            or None if evaluation fails (insufficient data, model not loaded)
        """
        if self.trainer is None:
            return None

        if df.empty or len(df) < self.context_window + 5:
            return None

        if idx < self.context_window:
            return None

        # Ensure technical indicators are available
        df_sorted = df.sort_values("timestamp").reset_index(drop=True)
        if "ATR" not in df_sorted.columns or "rsi_14" not in df_sorted.columns:
            df_sorted = add_technicals(df_sorted)
        if "ATR" not in df_sorted.columns:
            df_sorted["ATR"] = calculate_atr(df_sorted, 14)

        atr_val = df_sorted["ATR"].iloc[idx]
        if pd.isna(atr_val) or atr_val <= 0:
            return None

        # Extract feature window
        start_idx = idx - self.context_window
        window_df = df_sorted.iloc[start_idx:idx].copy()

        if len(window_df) < self.context_window:
            return None

        features = self._extract_features(
            window_df, df_sorted.iloc[idx], strategy_side, atr_val
        )
        if features is None:
            return None

        features = features.reshape(1, -1)

        try:
            # Use threshold=0.0 so NO signal is filtered — we get raw probabilities
            predictions = self.trainer.predict(features, entry_threshold=0.0)
        except Exception as e:
            logger.debug(f"NN evaluation failed at idx {idx}: {e}")
            return None

        entry_long_proba = float(predictions["entry_long_proba"][0])
        entry_short_proba = float(predictions["entry_short_proba"][0])
        confidence = float(predictions["confidence"][0])
        sl_distance_atr = float(predictions["sl_distance"][0])
        tp_distance_atr = float(predictions["tp_distance"][0])

        # Select the probability matching the strategy side
        if strategy_side == "LONG":
            entry_proba = entry_long_proba
        else:
            entry_proba = entry_short_proba

        logger.debug(
            f"NN evaluation: side={strategy_side}, "
            f"entry_proba={entry_proba:.2%}, confidence={confidence:.2%}, "
            f"SL={sl_distance_atr:.2f}xATR, TP={tp_distance_atr:.2f}xATR"
        )

        return {
            "entry_proba": entry_proba,
            "confidence": confidence,
            "sl_distance_atr": sl_distance_atr,
            "tp_distance_atr": tp_distance_atr,
        }

    def signals_to_dataframe(self, signals: List[NeuralSignal]) -> pd.DataFrame:
        """Convert signals to DataFrame format compatible with backtest engine.
        
        Args:
            signals: List of NeuralSignal objects
            
        Returns:
            DataFrame with signal columns
        """
        if not signals:
            return pd.DataFrame()
        
        data = []
        for s in signals:
            data.append({
                "timestamp": s.timestamp,
                "signal_type": s.side,
                "entry_price": s.entry_price,
                "sl_price": s.sl_price,
                "tp_price": s.tp_price,
                "sl_distance_atr": s.sl_distance_atr,
                "tp_distance_atr": s.tp_distance_atr,
                "ai_probability": s.entry_probability,
                "ai_confidence": s.confidence,
                "atr": s.atr_at_signal,
            })
        
        return pd.DataFrame(data)


def predict_for_ticker(
    ticker: str,
    df: pd.DataFrame,
    model_path: Path,
    sides: Optional[List[str]] = None,
    entry_threshold: float = 0.5,
    min_confidence: float = 0.3,
) -> pd.DataFrame:
    """Convenience function to generate signals for a ticker.
    
    Uses dual-head architecture: single forward pass gives both LONG and SHORT.
    
    Args:
        ticker: Ticker symbol
        df: OHLCV DataFrame
        model_path: Path to trained model
        sides: List of sides to filter (default: ["LONG", "SHORT"])
        entry_threshold: Minimum entry probability
        min_confidence: Minimum confidence score
        
    Returns:
        DataFrame with signals
    """
    if sides is None:
        sides = ["LONG", "SHORT"]
    
    predictor = NeuralPredictor(
        model_path=model_path,
        entry_threshold=entry_threshold,
        min_confidence=min_confidence,
    )
    
    # Single call returns all signals (both LONG and SHORT)
    all_signals = predictor.predict_signals(df, side=None)
    
    if sides:
        all_signals = [s for s in all_signals if s.side in sides]
    
    if not all_signals:
        logger.info(f"{ticker}: сигналов не найдено")
        return pd.DataFrame()
    
    result = predictor.signals_to_dataframe(all_signals)
    result["ticker"] = ticker
    
    result = result.sort_values("timestamp").reset_index(drop=True)
    
    logger.info(f"{ticker}: {len(result)} сигналов сгенерировано")
    return result
