"""Tests for Triple Barrier labeling module."""

import numpy as np
import pandas as pd

from ai.labeling import (
    BarrierConfig,
    BarrierResult,
    compute_triple_barrier_labels,
    filter_quality_labels,
    create_training_targets,
)


def make_ohlcv_df(n_bars: int = 200, trend: str = "up", volatility: float = 0.02) -> pd.DataFrame:
    """Create synthetic OHLCV data for testing."""
    np.random.seed(42)
    
    base_price = 100.0
    prices = [base_price]
    
    for i in range(1, n_bars):
        if trend == "up":
            drift = 0.001
        elif trend == "down":
            drift = -0.001
        else:
            drift = 0.0
        
        change = drift + np.random.normal(0, volatility)
        new_price = prices[-1] * (1 + change)
        prices.append(new_price)
    
    data = []
    for i, close in enumerate(prices):
        spread = close * volatility * 2
        high = close + spread * np.random.uniform(0.3, 1.0)
        low = close - spread * np.random.uniform(0.3, 1.0)
        open_price = close * (1 + np.random.normal(0, volatility / 2))
        volume = int(np.random.uniform(100000, 1000000))
        
        data.append({
            "timestamp": 1700000000 + i * 3600,
            "Open": open_price,
            "High": max(open_price, close, high),
            "Low": min(open_price, close, low),
            "Close": close,
            "Volume": volume,
        })
    
    return pd.DataFrame(data)


class TestBarrierConfig:
    def test_default_config(self):
        config = BarrierConfig()
        assert config.tp_atr_mult == 1.5
        assert config.sl_atr_mult == 1.0
        assert config.max_holding_bars == 48
        assert config.atr_period == 14

    def test_custom_config(self):
        config = BarrierConfig(tp_atr_mult=3.0, sl_atr_mult=1.5, max_holding_bars=24)
        assert config.tp_atr_mult == 3.0
        assert config.sl_atr_mult == 1.5
        assert config.max_holding_bars == 24


class TestTripleBarrierLabels:
    def test_basic_labeling(self):
        df = make_ohlcv_df(n_bars=200, trend="up")
        labels = compute_triple_barrier_labels(df)
        
        assert not labels.empty
        assert "entry_idx" in labels.columns
        assert "outcome" in labels.columns
        assert "pnl_pct" in labels.columns
        assert "side" in labels.columns

    def test_label_outcomes(self):
        df = make_ohlcv_df(n_bars=200)
        labels = compute_triple_barrier_labels(df)
        
        assert set(labels["outcome"].unique()).issubset({-1, 0, 1})

    def test_both_sides_with_trend_filter(self):
        """With trend filter (default), each bar gets ONE side.
        LONG in uptrend, SHORT in downtrend. Both sides can appear
        across different bars in a trending dataset."""
        df = make_ohlcv_df(n_bars=500, trend="up")  # uptrend → mostly LONG
        labels = compute_triple_barrier_labels(df, sides=["LONG", "SHORT"])
        
        assert not labels.empty
        assert "side" in labels.columns
        # In uptrend, most labels should be LONG
        long_ratio = (labels["side"] == "LONG").mean()
        assert long_ratio >= 0.5, f"Expected mostly LONG in uptrend, got {long_ratio:.1%} LONG"
        
        # SHORT labels can appear near SMA or downtrend segments
        df_down = make_ohlcv_df(n_bars=500, trend="down")  # downtrend → mostly SHORT
        labels_down = compute_triple_barrier_labels(df_down, sides=["LONG", "SHORT"])
        if not labels_down.empty:
            short_ratio = (labels_down["side"] == "SHORT").mean()
            assert short_ratio >= 0.5, f"Expected mostly SHORT in downtrend, got {short_ratio:.1%} SHORT"

    def test_long_only(self):
        df = make_ohlcv_df(n_bars=200)
        labels = compute_triple_barrier_labels(df, sides=["LONG"])
        
        assert all(labels["side"] == "LONG")

    def test_short_only(self):
        df = make_ohlcv_df(n_bars=200)
        labels = compute_triple_barrier_labels(df, sides=["SHORT"])
        
        assert all(labels["side"] == "SHORT")

    def test_empty_dataframe(self):
        df = pd.DataFrame()
        labels = compute_triple_barrier_labels(df)
        assert labels.empty

    def test_insufficient_data(self):
        df = make_ohlcv_df(n_bars=30)
        labels = compute_triple_barrier_labels(df)
        assert labels.empty

    def test_pnl_calculation_long(self):
        df = make_ohlcv_df(n_bars=200, trend="up")
        config = BarrierConfig(tp_atr_mult=1.0, sl_atr_mult=1.0, max_holding_bars=10)
        labels = compute_triple_barrier_labels(df, config, sides=["LONG"])
        
        tp_labels = labels[labels["outcome"] == 1]
        assert not tp_labels.empty, "Expected at least one TP label"
        assert all(tp_labels["pnl_pct"] > 0)
        
        sl_labels = labels[labels["outcome"] == -1]
        assert not sl_labels.empty, "Expected at least one SL label"
        assert all(sl_labels["pnl_pct"] < 0)

    def test_pnl_calculation_short(self):
        df = make_ohlcv_df(n_bars=200, trend="down")
        config = BarrierConfig(tp_atr_mult=1.0, sl_atr_mult=1.0, max_holding_bars=10)
        labels = compute_triple_barrier_labels(df, config, sides=["SHORT"])
        
        tp_labels = labels[labels["outcome"] == 1]
        assert not tp_labels.empty, "Expected at least one TP label"
        assert all(tp_labels["pnl_pct"] > 0)

    def test_bars_to_exit_bounded(self):
        df = make_ohlcv_df(n_bars=200)
        config = BarrierConfig(max_holding_bars=20)
        labels = compute_triple_barrier_labels(df, config)
        
        assert all(labels["bars_to_exit"] <= 20)
        assert all(labels["bars_to_exit"] >= 1)

    def test_mfe_mae_populated(self):
        df = make_ohlcv_df(n_bars=200)
        labels = compute_triple_barrier_labels(df)
        
        assert "mfe" in labels.columns
        assert "mae" in labels.columns
        assert all(labels["mfe"] >= 0)
        assert all(labels["mae"] <= 0)


class TestFilterQualityLabels:
    def test_filter_basic(self):
        df = make_ohlcv_df(n_bars=200)
        labels = compute_triple_barrier_labels(df)
        
        filtered = filter_quality_labels(labels)
        assert len(filtered) <= len(labels)

    def test_filter_empty(self):
        labels = pd.DataFrame()
        filtered = filter_quality_labels(labels)
        assert filtered.empty


class TestCreateTrainingTargets:
    def test_targets_created(self):
        df = make_ohlcv_df(n_bars=200)
        labels = compute_triple_barrier_labels(df)
        targets = create_training_targets(labels)
        
        assert "entry_signal" in targets.columns
        assert "sl_distance_norm" in targets.columns
        assert "tp_distance_norm" in targets.columns
        assert "expected_pnl" in targets.columns
        assert "target_sl_atr" in targets.columns
        assert "target_tp_atr" in targets.columns

    def test_entry_signal_binary(self):
        df = make_ohlcv_df(n_bars=200)
        labels = compute_triple_barrier_labels(df)
        targets = create_training_targets(labels)
        
        assert set(targets["entry_signal"].unique()).issubset({0.0, 1.0})

    def test_entry_signal_only_tp(self):
        """entry_signal = 1 только для TP (outcome=1), не для SL"""
        df = make_ohlcv_df(n_bars=200)
        labels = compute_triple_barrier_labels(df)
        targets = create_training_targets(labels)
        
        # Все TP должны быть entry_signal = 1
        tp_mask = targets["outcome"] == 1
        assert tp_mask.any(), "Expected at least one TP outcome"
        assert all(targets.loc[tp_mask, "entry_signal"] == 1.0)
        
        # Все SL должны быть entry_signal = 0
        sl_mask = targets["outcome"] == -1
        assert sl_mask.any(), "Expected at least one SL outcome"
        assert all(targets.loc[sl_mask, "entry_signal"] == 0.0)

    def test_entry_signal_positive_rate_reasonable(self):
        """Positive rate должен быть в разумных пределах (не 99%)"""
        df = make_ohlcv_df(n_bars=300)
        labels = compute_triple_barrier_labels(df)
        targets = create_training_targets(labels)
        
        positive_rate = targets["entry_signal"].mean()
        # WR обычно 30-60% для разумных барьеров
        assert positive_rate < 0.8, f"Positive rate too high: {positive_rate}"
        assert positive_rate > 0.05, f"Positive rate too low: {positive_rate}"

    def test_normalized_distances_bounded(self):
        df = make_ohlcv_df(n_bars=200)
        labels = compute_triple_barrier_labels(df)
        targets = create_training_targets(labels)
        
        assert all(targets["target_sl_atr"] >= 0.3)
        assert all(targets["target_sl_atr"] <= 3.0)
        assert all(targets["target_tp_atr"] >= 0.5)
        assert all(targets["target_tp_atr"] <= 5.0)

    def test_empty_input(self):
        labels = pd.DataFrame()
        targets = create_training_targets(labels)
        assert targets.empty

    def test_min_pnl_threshold(self):
        """Можно задать минимальный PnL для positive entry"""
        df = make_ohlcv_df(n_bars=200)
        labels = compute_triple_barrier_labels(df)
        
        targets_default = create_training_targets(labels, min_pnl_for_entry=0.0)
        targets_strict = create_training_targets(labels, min_pnl_for_entry=0.01)
        
        # С более строгим порогом positive rate должен быть ниже или равен
        assert targets_strict["entry_signal"].mean() <= targets_default["entry_signal"].mean()


class TestBarrierResult:
    def test_dataclass_creation(self):
        result = BarrierResult(
            entry_idx=10,
            entry_price=100.0,
            entry_time=1700000000,
            side="LONG",
            outcome=1,
            exit_price=102.0,
            exit_idx=15,
            bars_to_exit=5,
            pnl_pct=0.02,
            max_favorable_excursion=0.025,
            max_adverse_excursion=-0.005,
            atr_at_entry=1.5,
            suggested_sl_distance=1.5,
            suggested_tp_distance=3.0,
        )
        
        assert result.entry_idx == 10
        assert result.side == "LONG"
        assert result.outcome == 1
        assert result.pnl_pct == 0.02
