"""Dataset generator for neural network training on all tickers H1.

Loads OHLCV data for all tickers, generates Triple Barrier labels,
and extracts features for multi-task learning.
"""

from __future__ import annotations

from pathlib import Path
from typing import Dict, List, Optional, Tuple
from dataclasses import dataclass

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

from ai.features import (
    CONTEXT_WINDOW,
    add_technicals,
    extract_feature_vector,
)
from ai.labeling import (
    BarrierConfig,
    compute_triple_barrier_labels,
    create_training_targets,
    filter_quality_labels,
)
from ai.strategy_signals import (
    StrategyConfig as StrategySignalConfig,
    get_strategy_signals,
)


ALL_TICKERS = [
    "SBER", "GAZP", "PLZL", "VTBR", "LKOH", "ROSN",
    "NVTK", "MTSS", "PHOR", "SNGSP", "ASTR", "X5",
    "MOEX", "BITCOIN",
]


@dataclass
class DatasetConfig:
    """Configuration for dataset generation."""
    tickers: List[str] = None
    start_date: str = "2023-01-01"
    end_date: str = "2024-01-01"
    context_window: int = CONTEXT_WINDOW
    barrier_config: BarrierConfig = None
    sides: List[str] = None
    min_bars_per_ticker: int = 500
    cache_dir: Optional[Path] = None
    use_strategy_signals: bool = False
    strategy_config: Optional[StrategySignalConfig] = None

    def __post_init__(self):
        if self.tickers is None:
            self.tickers = ALL_TICKERS
        if self.barrier_config is None:
            self.barrier_config = BarrierConfig()
        if self.sides is None:
            self.sides = ["LONG", "SHORT"]
        if self.strategy_config is None:
            self.strategy_config = StrategySignalConfig()


class MultiTickerDataset:
    """Dataset generator that loads data from all tickers and creates
    a unified training set with Triple Barrier labels."""

    def __init__(self, config: Optional[DatasetConfig] = None):
        self.config = config or DatasetConfig()
        self.ticker_to_id: Dict[str, int] = {
            t: i for i, t in enumerate(self.config.tickers)
        }
        self._data_cache: Dict[str, pd.DataFrame] = {}
        self._labels_cache: Dict[str, pd.DataFrame] = {}

    def load_all_tickers(self) -> Dict[str, pd.DataFrame]:
        """Load H1 OHLCV data for all tickers.
        
        Returns:
            Dictionary mapping ticker -> OHLCV DataFrame
        """
        from db import fetch_ohlcv
        
        start_ts = int(pd.Timestamp(self.config.start_date).timestamp())
        end_ts = int(pd.Timestamp(self.config.end_date).timestamp())
        
        loaded = {}
        for ticker in self.config.tickers:
            try:
                df = fetch_ohlcv(ticker, "H1", start_ts, end_ts)
                
                if df.empty:
                    logger.warning(f"{ticker}: нет данных H1")
                    continue
                
                if len(df) < self.config.min_bars_per_ticker:
                    logger.warning(
                        f"{ticker}: недостаточно данных ({len(df)} < {self.config.min_bars_per_ticker})"
                    )
                    continue
                
                df = df.sort_values("timestamp").reset_index(drop=True)
                loaded[ticker] = df
                self._data_cache[ticker] = df
                
                logger.debug(f"{ticker}: загружено {len(df)} баров H1")
                
            except Exception as e:
                logger.error(f"{ticker}: ошибка загрузки: {e}")
        
        logger.info(f"Загружено {len(loaded)}/{len(self.config.tickers)} тикеров")
        return loaded

    def generate_labels_for_ticker(
        self, ticker: str, df: pd.DataFrame
    ) -> pd.DataFrame:
        """Generate Triple Barrier labels for a single ticker.
        
        Args:
            ticker: Ticker symbol
            df: OHLCV DataFrame
            
        Returns:
            DataFrame with labels
        """
        labels = compute_triple_barrier_labels(
            df=df,
            config=self.config.barrier_config,
            sides=self.config.sides,
        )
        
        if labels.empty:
            return labels
        
        labels["ticker"] = ticker
        labels["ticker_id"] = self.ticker_to_id.get(ticker, -1)
        
        self._labels_cache[ticker] = labels
        return labels

    def generate_all_labels(self) -> pd.DataFrame:
        """Generate labels for all loaded tickers.
        
        When use_strategy_signals=True, labels are generated ONLY for bars
        where the MACD+RSI strategy triggers. This creates a focused dataset
        of strategy-augmented samples with positive rate = TP rate of the
        strategy (~40-55%), rather than diluting with non-strategy bars.
        
        Returns:
            Combined DataFrame with labels from all tickers
        """
        if not self._data_cache:
            self.load_all_tickers()
        
        all_labels = []
        for ticker, df in self._data_cache.items():
            if self.config.use_strategy_signals:
                # Strategy-augmented mode:
                # 1. Generate MACD+RSI signals first
                # 2. Compute Triple Barrier labels ONLY for strategy-triggered bars
                # 3. entry_signal = TP rate of strategy bars
                signals_df = get_strategy_signals(df, self.config.strategy_config)
                if signals_df.empty:
                    logger.warning(f"{ticker}: нет стратегических сигналов, пропускаем")
                    continue
                
                strategy_indices = signals_df.index[signals_df["strategy_signal"] == 1].tolist()
                
                if not strategy_indices:
                    logger.warning(f"{ticker}: стратегия не дала сигналов, пропускаем")
                    continue
                
                # Generate Triple Barrier labels only for strategy bars
                full_labels = compute_triple_barrier_labels(
                    df=df,
                    config=self.config.barrier_config,
                    sides=self.config.sides,
                )
                
                if full_labels.empty:
                    continue
                
                # Filter to only bars where strategy triggered
                labels = full_labels[full_labels["entry_idx"].isin(strategy_indices)].copy()
                
                if labels.empty:
                    logger.warning(f"{ticker}: ни один стратегический сигнал не совпал с метками")
                    continue
                
                # Create training targets for these strategy-filtered labels
                labels = create_training_targets(labels, min_pnl_for_entry=0.001)
                labels["ticker"] = ticker
                labels["ticker_id"] = self.ticker_to_id.get(ticker, -1)
                
                logger.info(
                    f"{ticker}: strategy-augmented labels: "
                    f"{len(labels)} примеров, "
                    f"entry_signal={int(labels['entry_signal'].sum())}/{len(labels)} "
                    f"({labels['entry_signal'].mean():.1%})"
                )
            else:
                # Standard mode: all bars get labels
                labels = self.generate_labels_for_ticker(ticker, df)
            
            if not labels.empty:
                all_labels.append(labels)
        
        if not all_labels:
            logger.warning("Не сгенерировано ни одной метки")
            return pd.DataFrame()
        
        combined = pd.concat(all_labels, ignore_index=True)
        
        if not self.config.use_strategy_signals:
            # Standard pipeline: filter + create targets
            combined = filter_quality_labels(combined)
            combined = create_training_targets(combined)
        
        self._log_label_statistics(combined)
        return combined

    def extract_features_for_label(
        self, ticker: str, label_row: pd.Series
    ) -> Optional[np.ndarray]:
        """Extract feature vector for a single label.
        
        Args:
            ticker: Ticker symbol
            label_row: Single row from labels DataFrame
            
        Returns:
            Feature vector or None if extraction failed
        """
        df = self._data_cache.get(ticker)
        if df is None:
            return None
        
        entry_idx = int(label_row["entry_idx"])
        start_idx = entry_idx - self.config.context_window
        
        if start_idx < 0:
            return None
        
        window_df = df.iloc[start_idx:entry_idx].copy()
        if len(window_df) < self.config.context_window:
            return None
        
        window_df = add_technicals(window_df)
        
        features = self._extract_feature_vector(window_df, label_row)
        return features

    def _extract_features_fast(
        self,
        technicals_cache: Dict[str, pd.DataFrame],
        ticker: str,
        label_row: pd.Series,
    ) -> Optional[np.ndarray]:
        """Fast feature extraction using pre-computed technicals.
        
        Args:
            technicals_cache: Dict of ticker -> DataFrame with technicals
            ticker: Ticker symbol
            label_row: Single row from labels DataFrame
            
        Returns:
            Feature vector or None if extraction failed
        """
        df = technicals_cache.get(ticker)
        if df is None:
            return None
        
        entry_idx = int(label_row["entry_idx"])
        start_idx = entry_idx - self.config.context_window
        
        if start_idx < 0:
            return None
        
        if entry_idx >= len(df):
            return None
        
        window_df = df.iloc[start_idx:entry_idx]
        if len(window_df) < self.config.context_window:
            return None
        
        features = self._extract_feature_vector(window_df, label_row)
        return features

    def _extract_feature_vector(
        self, window_df: pd.DataFrame, label_row: pd.Series
    ) -> np.ndarray:
        """Extract feature vector from window DataFrame.
        
        SINGLE SOURCE OF TRUTH: delegates to ai.features.extract_feature_vector().
        
        Args:
            window_df: Window of OHLCV data with technicals
            label_row: Label information
            
        Returns:
            Feature vector as numpy array
        """
        return extract_feature_vector(
            window_df=window_df,
            side=label_row.get("side", "LONG"),
            atr_val=label_row.get("atr_at_entry", 0),
            ticker_id=label_row.get("ticker_id", 0),
            n_tickers=len(self.config.tickers),
            entry_time=label_row.get("entry_time", 0),
            context_window=self.config.context_window,
        )

    def build_dataset(
        self, labels_df: Optional[pd.DataFrame] = None
    ) -> Tuple[np.ndarray, np.ndarray, pd.DataFrame]:
        """Build complete training dataset.
        
        Args:
            labels_df: Pre-computed labels (if None, will generate)
            
        Returns:
            Tuple of (X, y, metadata_df) where:
            - X: Feature matrix (n_samples, n_features)
            - y: Target array with columns [entry_signal, sl_atr, tp_atr, pnl]
            - metadata_df: DataFrame with label metadata
        """
        if labels_df is None:
            labels_df = self.generate_all_labels()
        
        if labels_df.empty:
            return np.array([]), np.array([]), pd.DataFrame()
        
        # Оптимизация: предварительно вычислить technicals для каждого тикера
        logger.info("Предварительное вычисление технических индикаторов...")
        technicals_cache: Dict[str, pd.DataFrame] = {}
        for ticker in labels_df["ticker"].unique():
            if ticker in self._data_cache:
                df = self._data_cache[ticker].copy()
                df = add_technicals(df)
                technicals_cache[ticker] = df
                logger.debug(f"{ticker}: technicals вычислены")
        
        X_list = []
        y_list = []
        meta_list = []
        
        total = len(labels_df)
        log_interval = max(1, total // 20)  # Логировать каждые 5%
        
        for i, (idx, row) in enumerate(labels_df.iterrows()):
            if i % log_interval == 0:
                logger.info(f"Building dataset: {i}/{total} ({i/total:.1%})")
            
            ticker = row["ticker"]
            
            # Использовать кэшированные technicals
            features = self._extract_features_fast(technicals_cache, ticker, row)
            if features is None:
                continue
            
            X_list.append(features)
            
            side_code = 1.0 if row.get("side", "LONG") == "LONG" else 0.0
            y_list.append([
                row.get("entry_signal", 0),
                row.get("target_sl_atr", 1.0),
                row.get("target_tp_atr", 2.0),
                row.get("expected_pnl", 0),
                side_code,  # 5th column: 1=LONG, 0=SHORT
            ])
            
            meta_list.append({
                "ticker": ticker,
                "ticker_id": row.get("ticker_id", -1),
                "entry_idx": row.get("entry_idx", -1),
                "entry_time": row.get("entry_time", 0),
                "side": row.get("side", "LONG"),
                "outcome": row.get("outcome", 0),
            })
        
        if not X_list:
            logger.warning("Не удалось извлечь ни одного примера")
            return np.array([]), np.array([]), pd.DataFrame()
        
        X = np.array(X_list, dtype=np.float32)
        y = np.array(y_list, dtype=np.float32)
        metadata = pd.DataFrame(meta_list)
        
        logger.info(
            f"Датасет собран: {X.shape[0]} примеров, {X.shape[1]} признаков"
        )
        
        return X, y, metadata

    def save_dataset(
        self, X: np.ndarray, y: np.ndarray, metadata: pd.DataFrame, path: Path
    ):
        """Save dataset to disk.
        
        Args:
            X: Feature matrix
            y: Target array
            metadata: Metadata DataFrame
            path: Output path (without extension)
        """
        path = Path(path)
        path.parent.mkdir(parents=True, exist_ok=True)
        
        np.savez(
            path.with_suffix(".npz"),
            X=X,
            y=y,
        )
        
        metadata.to_parquet(path.with_suffix(".meta.parquet"))
        
        logger.info(f"Датасет сохранён: {path.with_suffix('.npz')}")

    def load_dataset(
        self, path: Path
    ) -> Tuple[np.ndarray, np.ndarray, pd.DataFrame]:
        """Load dataset from disk.
        
        Args:
            path: Base path (without extension)
            
        Returns:
            Tuple of (X, y, metadata_df)
        """
        path = Path(path)
        
        data = np.load(path.with_suffix(".npz"))
        X = data["X"]
        y = data["y"]
        
        metadata = pd.read_parquet(path.with_suffix(".meta.parquet"))
        
        logger.info(f"Датасет загружен: {X.shape[0]} примеров")
        return X, y, metadata

    def _log_label_statistics(self, labels_df: pd.DataFrame):
        """Log statistics about generated labels."""
        if labels_df.empty:
            return
        
        total = len(labels_df)
        tp_count = (labels_df["outcome"] == 1).sum()
        sl_count = (labels_df["outcome"] == -1).sum()
        timeout_count = (labels_df["outcome"] == 0).sum()
        
        win_rate = tp_count / total if total > 0 else 0
        avg_pnl = labels_df["pnl_pct"].mean()
        
        logger.info(
            f"Статистика меток: {total} примеров | "
            f"TP={tp_count} ({tp_count/total:.1%}), "
            f"SL={sl_count} ({sl_count/total:.1%}), "
            f"Timeout={timeout_count} ({timeout_count/total:.1%}) | "
            f"WR={win_rate:.1%}, Avg PnL={avg_pnl:.2%}"
        )
        
        if "ticker" in labels_df.columns:
            ticker_stats = labels_df.groupby("ticker").agg({
                "outcome": "count",
                "pnl_pct": "mean",
            }).rename(columns={"outcome": "count", "pnl_pct": "avg_pnl"})
            
            ticker_stats["win_rate"] = labels_df.groupby("ticker")["outcome"].apply(
                lambda x: (x == 1).mean()
            )
            
            logger.info(f"По тикерам:\n{ticker_stats.to_string()}")


def create_dataset_from_db(
    start_date: str = "2023-01-01",
    end_date: str = "2024-01-01",
    tickers: Optional[List[str]] = None,
    barrier_config: Optional[BarrierConfig] = None,
) -> Tuple[np.ndarray, np.ndarray, pd.DataFrame]:
    """Convenience function to create dataset from database.
    
    Args:
        start_date: Start date for data loading
        end_date: End date for data loading
        tickers: List of tickers (default: all)
        barrier_config: Barrier configuration
        
    Returns:
        Tuple of (X, y, metadata_df)
    """
    config = DatasetConfig(
        tickers=tickers,
        start_date=start_date,
        end_date=end_date,
        barrier_config=barrier_config,
    )
    
    dataset = MultiTickerDataset(config)
    dataset.load_all_tickers()
    labels = dataset.generate_all_labels()
    X, y, metadata = dataset.build_dataset(labels)
    
    return X, y, metadata
