"""Tests for MultiTickerDataset module."""

import numpy as np
import pandas as pd
from unittest.mock import patch
import pytest

from ai.dataset_v2 import (
    ALL_TICKERS,
    DatasetConfig,
    MultiTickerDataset,
    create_dataset_from_db,
)
from ai.labeling import BarrierConfig


def make_ohlcv_df(n_bars: int = 600, base_price: float = 100.0) -> pd.DataFrame:
    """Create synthetic OHLCV data for testing."""
    np.random.seed(42)
    
    data = []
    price = base_price
    
    for i in range(n_bars):
        change = np.random.normal(0.0005, 0.015)
        price = price * (1 + change)
        
        spread = price * 0.02
        high = price + spread * np.random.uniform(0.3, 1.0)
        low = price - spread * np.random.uniform(0.3, 1.0)
        open_price = price * (1 + np.random.normal(0, 0.005))
        volume = int(np.random.uniform(100000, 1000000))
        
        data.append({
            "timestamp": 1700000000 + i * 3600,
            "Open": open_price,
            "High": max(open_price, price, high),
            "Low": min(open_price, price, low),
            "Close": price,
            "Volume": volume,
        })
    
    return pd.DataFrame(data)


class TestDatasetConfig:
    def test_default_config(self):
        config = DatasetConfig()
        assert config.tickers == ALL_TICKERS
        assert config.context_window == 24
        assert config.min_bars_per_ticker == 500

    def test_custom_config(self):
        config = DatasetConfig(
            tickers=["SBER", "GAZP"],
            start_date="2024-01-01",
            context_window=24,
        )
        assert config.tickers == ["SBER", "GAZP"]
        assert config.start_date == "2024-01-01"
        assert config.context_window == 24

    def test_barrier_config_default(self):
        config = DatasetConfig()
        assert isinstance(config.barrier_config, BarrierConfig)


class TestMultiTickerDataset:
    def test_init_default(self):
        dataset = MultiTickerDataset()
        assert len(dataset.ticker_to_id) == len(ALL_TICKERS)
        assert "SBER" in dataset.ticker_to_id

    def test_init_custom_tickers(self):
        config = DatasetConfig(tickers=["SBER", "GAZP", "VTBR"])
        dataset = MultiTickerDataset(config)
        assert len(dataset.ticker_to_id) == 3
        assert dataset.ticker_to_id["SBER"] == 0
        assert dataset.ticker_to_id["GAZP"] == 1
        assert dataset.ticker_to_id["VTBR"] == 2

    def test_generate_labels_for_ticker(self):
        config = DatasetConfig(tickers=["TEST"])
        dataset = MultiTickerDataset(config)
        
        df = make_ohlcv_df(n_bars=600)
        dataset._data_cache["TEST"] = df
        
        labels = dataset.generate_labels_for_ticker("TEST", df)
        
        assert not labels.empty
        assert "ticker" in labels.columns
        assert all(labels["ticker"] == "TEST")
        assert "ticker_id" in labels.columns

    def test_generate_labels_empty_df(self):
        config = DatasetConfig(tickers=["TEST"])
        dataset = MultiTickerDataset(config)
        
        labels = dataset.generate_labels_for_ticker("TEST", pd.DataFrame())
        assert labels.empty

    def test_extract_features_for_label(self):
        config = DatasetConfig(tickers=["TEST"], context_window=24)
        dataset = MultiTickerDataset(config)
        
        df = make_ohlcv_df(n_bars=600)
        dataset._data_cache["TEST"] = df
        
        label_row = pd.Series({
            "entry_idx": 100,
            "entry_time": 1700100000,
            "side": "LONG",
            "atr_at_entry": 1.5,
            "sl_distance_norm": 1.0,
            "tp_distance_norm": 2.0,
            "ticker_id": 0,
        })
        
        features = dataset.extract_features_for_label("TEST", label_row)
        
        assert features is not None
        assert isinstance(features, np.ndarray)
        assert len(features) > 0

    def test_extract_features_invalid_idx(self):
        config = DatasetConfig(tickers=["TEST"], context_window=24)
        dataset = MultiTickerDataset(config)
        
        df = make_ohlcv_df(n_bars=100)
        dataset._data_cache["TEST"] = df
        
        label_row = pd.Series({
            "entry_idx": 10,
            "entry_time": 1700100000,
            "side": "LONG",
        })
        
        features = dataset.extract_features_for_label("TEST", label_row)
        assert features is None

    def test_extract_features_missing_ticker(self):
        config = DatasetConfig(tickers=["TEST"])
        dataset = MultiTickerDataset(config)
        
        label_row = pd.Series({"entry_idx": 100})
        features = dataset.extract_features_for_label("UNKNOWN", label_row)
        assert features is None

    @patch("db.fetch_ohlcv")
    def test_load_all_tickers_mock(self, mock_fetch):
        mock_fetch.return_value = make_ohlcv_df(n_bars=600)
        
        config = DatasetConfig(tickers=["SBER", "GAZP"])
        dataset = MultiTickerDataset(config)
        
        loaded = dataset.load_all_tickers()
        
        assert len(loaded) == 2
        assert "SBER" in loaded
        assert "GAZP" in loaded

    @patch("db.fetch_ohlcv")
    def test_load_all_tickers_empty_response(self, mock_fetch):
        mock_fetch.return_value = pd.DataFrame()
        
        config = DatasetConfig(tickers=["SBER"])
        dataset = MultiTickerDataset(config)
        
        loaded = dataset.load_all_tickers()
        assert len(loaded) == 0

    @patch("db.fetch_ohlcv")
    def test_load_all_tickers_insufficient_data(self, mock_fetch):
        mock_fetch.return_value = make_ohlcv_df(n_bars=100)
        
        config = DatasetConfig(tickers=["SBER"], min_bars_per_ticker=500)
        dataset = MultiTickerDataset(config)
        
        loaded = dataset.load_all_tickers()
        assert len(loaded) == 0


class TestBuildDataset:
    def test_build_dataset_from_labels(self):
        config = DatasetConfig(tickers=["TEST"], context_window=24)
        dataset = MultiTickerDataset(config)
        
        df = make_ohlcv_df(n_bars=600)
        dataset._data_cache["TEST"] = df
        
        labels = dataset.generate_labels_for_ticker("TEST", df)
        labels = labels[labels["entry_idx"] >= 50].head(10)
        
        X, y, metadata = dataset.build_dataset(labels)
        
        assert X.shape[0] > 0
        assert X.shape[0] == y.shape[0]
        assert y.shape[1] == 5
        assert len(metadata) == X.shape[0]

    def test_build_dataset_empty_labels(self):
        config = DatasetConfig(tickers=["TEST"])
        dataset = MultiTickerDataset(config)
        
        X, y, metadata = dataset.build_dataset(pd.DataFrame())
        
        assert X.size == 0
        assert y.size == 0
        assert metadata.empty


class TestSaveLoadDataset:
    def test_save_and_load(self, tmp_path):
        config = DatasetConfig(tickers=["TEST"])
        dataset = MultiTickerDataset(config)
        
        X = np.random.randn(100, 50).astype(np.float32)
        y = np.random.randn(100, 4).astype(np.float32)
        metadata = pd.DataFrame({
            "ticker": ["TEST"] * 100,
            "ticker_id": [0] * 100,
            "entry_idx": range(100),
            "entry_time": [1700000000] * 100,
            "side": ["LONG"] * 100,
            "outcome": [1] * 100,
        })
        
        base_path = tmp_path / "test_dataset"
        dataset.save_dataset(X, y, metadata, base_path)
        
        X_loaded, y_loaded, meta_loaded = dataset.load_dataset(base_path)
        
        np.testing.assert_array_almost_equal(X, X_loaded)
        np.testing.assert_array_almost_equal(y, y_loaded)
        assert len(meta_loaded) == len(metadata)


class TestAllTickers:
    def test_all_tickers_list(self):
        assert len(ALL_TICKERS) == 14
        assert "SBER" in ALL_TICKERS
        assert "GAZP" in ALL_TICKERS
        assert "BITCOIN" in ALL_TICKERS
        assert "EURUSD" not in ALL_TICKERS


class TestCreateDatasetFromDB:
    @patch("db.fetch_ohlcv")
    def test_create_dataset_convenience(self, mock_fetch):
        mock_fetch.return_value = make_ohlcv_df(n_bars=600)
        
        X, y, metadata = create_dataset_from_db(
            start_date="2023-01-01",
            end_date="2024-01-01",
            tickers=["SBER"],
        )
        
        assert X.shape[0] > 0
        assert y.shape[1] == 5
