"""Tests for multi-task trainer."""

import numpy as np
import pytest
import torch

from ai.trainer_v2 import (
    MultiTaskLoss,
    MultiTaskTrainer,
    TrainerConfig,
    TrainingMetrics,
)


def make_synthetic_data(n_samples: int = 200, n_features: int = 600) -> tuple:
    """Create synthetic training data with dual-head side codes."""
    np.random.seed(42)
    X = np.random.randn(n_samples, n_features).astype(np.float32)
    
    entry_signal = np.random.randint(0, 2, n_samples).astype(np.float32)
    sl_atr = np.random.uniform(0.5, 2.5, n_samples).astype(np.float32)
    tp_atr = np.random.uniform(1.0, 4.0, n_samples).astype(np.float32)
    pnl = np.random.uniform(-0.05, 0.10, n_samples).astype(np.float32)
    side_code = np.random.randint(0, 2, n_samples).astype(np.float32)  # 1=LONG, 0=SHORT
    
    y = np.column_stack([entry_signal, sl_atr, tp_atr, pnl, side_code])
    
    return X, y


class TestTrainerConfig:
    def test_default_config(self):
        config = TrainerConfig()
        assert config.model_type == "lstm"
        assert config.context_window == 24
        assert config.learning_rate == 0.001
        assert config.epochs == 100

    def test_custom_config(self):
        config = TrainerConfig(
            model_type="transformer",
            epochs=50,
            batch_size=64,
        )
        assert config.model_type == "transformer"
        assert config.epochs == 50
        assert config.batch_size == 64


class TestTrainingMetrics:
    def test_default_metrics(self):
        metrics = TrainingMetrics()
        assert metrics.train_loss == 0.0
        assert metrics.entry_auc == 0.0
        assert metrics.best_epoch == 0

    def test_custom_metrics(self):
        metrics = TrainingMetrics(
            train_loss=0.5,
            val_loss=0.6,
            entry_auc=0.85,
            best_epoch=15,
        )
        assert metrics.train_loss == 0.5
        assert metrics.entry_auc == 0.85
        assert metrics.best_epoch == 15


class TestMultiTaskLoss:
    def test_loss_computation(self):
        criterion = MultiTaskLoss(
            entry_weight=1.0,
            sl_weight=0.5,
            tp_weight=0.5,
            conf_weight=0.3,
        )
        
        entry_logits = torch.randn(8)
        sl_pred = torch.rand(8) + 0.5
        tp_pred = torch.rand(8) + 1.0
        conf_pred = torch.rand(8)
        entry_target = torch.randint(0, 2, (8,)).float()
        sl_target = torch.rand(8) + 0.5
        tp_target = torch.rand(8) + 1.0
        side = torch.ones(8)  # All LONG
        
        loss, components = criterion(
            entry_logits, entry_logits, sl_pred, tp_pred, conf_pred,
            entry_target, sl_target, tp_target, side,
        )
        
        assert loss.item() > 0
        assert "entry_loss" in components
        assert "sl_loss" in components
        assert "tp_loss" in components
        assert "conf_loss" in components
        assert "total_loss" in components

    def test_loss_with_no_positive_entries(self):
        criterion = MultiTaskLoss()
        
        entry_logits = torch.randn(8)
        sl_pred = torch.rand(8) + 0.5
        tp_pred = torch.rand(8) + 1.0
        conf_pred = torch.rand(8)
        entry_target = torch.zeros(8)
        sl_target = torch.rand(8) + 0.5
        tp_target = torch.rand(8) + 1.0
        side = torch.ones(8)  # All LONG
        
        loss, components = criterion(
            entry_logits, entry_logits, sl_pred, tp_pred, conf_pred,
            entry_target, sl_target, tp_target, side,
        )
        
        assert loss.item() > 0
        # SL/TP loss weighted by 0.3 even for non-entry examples
        assert components["sl_loss"] > 0.0
        assert components["tp_loss"] > 0.0
        assert components["conf_loss"] > 0.0

    def test_loss_with_pos_weight(self):
        criterion = MultiTaskLoss(pos_weight=2.0)
        
        entry_logits = torch.randn(8)
        sl_pred = torch.rand(8) + 0.5
        tp_pred = torch.rand(8) + 1.0
        conf_pred = torch.rand(8)
        entry_target = torch.randint(0, 2, (8,)).float()
        sl_target = torch.rand(8) + 0.5
        tp_target = torch.rand(8) + 1.0
        side = torch.ones(8)  # All LONG
        
        loss, _ = criterion(
            entry_logits, entry_logits, sl_pred, tp_pred, conf_pred,
            entry_target, sl_target, tp_target, side,
        )
        
        assert loss.item() > 0


class TestMultiTaskTrainer:
    def test_trainer_creation(self):
        trainer = MultiTaskTrainer()
        assert trainer.model is None
        assert trainer.scaler is None

    def test_trainer_with_config(self):
        config = TrainerConfig(model_type="mlp", epochs=5)
        trainer = MultiTaskTrainer(config)
        assert trainer.config.model_type == "mlp"
        assert trainer.config.epochs == 5

    def test_fit_basic(self):
        config = TrainerConfig(
            model_type="mlp",
            epochs=5,
            batch_size=16,
            context_window=24,
        )
        trainer = MultiTaskTrainer(config)
        
        X, y = make_synthetic_data(n_samples=100, n_features=600)
        
        metrics = trainer.fit(X, y)
        
        assert trainer.model is not None
        assert metrics.train_loss > 0
        assert metrics.val_loss > 0
        assert 0 <= metrics.entry_auc <= 1

    def test_fit_lstm(self):
        config = TrainerConfig(
            model_type="lstm",
            epochs=3,
            batch_size=16,
            context_window=24,
            hidden_size=32,
        )
        trainer = MultiTaskTrainer(config)
        
        X, y = make_synthetic_data(n_samples=100, n_features=600)
        
        metrics = trainer.fit(X, y)
        
        assert trainer.model is not None
        assert metrics.entry_f1 >= 0

    def test_fit_empty_data_raises(self):
        trainer = MultiTaskTrainer()
        
        with pytest.raises(ValueError):
            trainer.fit(np.array([]), np.array([]))

    def test_predict_before_fit_raises(self):
        trainer = MultiTaskTrainer()
        X, _ = make_synthetic_data(n_samples=10)
        
        with pytest.raises(RuntimeError):
            trainer.predict(X)

    def test_predict_after_fit(self):
        config = TrainerConfig(
            model_type="mlp",
            epochs=3,
            batch_size=16,
            context_window=24,
        )
        trainer = MultiTaskTrainer(config)
        
        X, y = make_synthetic_data(n_samples=100, n_features=600)
        trainer.fit(X, y)
        
        X_test, _ = make_synthetic_data(n_samples=20, n_features=600)
        predictions = trainer.predict(X_test, entry_threshold=0.5)
        
        assert "entry_proba" in predictions
        assert "entry_signal" in predictions
        assert "sl_distance" in predictions
        assert "tp_distance" in predictions
        assert "confidence" in predictions
        
        assert predictions["entry_proba"].shape == (20,)
        assert predictions["sl_distance"].shape == (20,)


class TestTrainerSaveLoad:
    def test_save_and_load(self, tmp_path):
        config = TrainerConfig(
            model_type="mlp",
            epochs=3,
            batch_size=16,
            context_window=24,
        )
        trainer = MultiTaskTrainer(config)
        
        X, y = make_synthetic_data(n_samples=100, n_features=600)
        trainer.fit(X, y)
        
        model_path = tmp_path / "test_model.pt"
        trainer.save(model_path)
        
        assert model_path.exists()
        
        trainer2 = MultiTaskTrainer()
        trainer2.load(model_path)
        
        assert trainer2.model is not None
        assert trainer2.config.model_type == "mlp"

    def test_save_before_fit_raises(self, tmp_path):
        trainer = MultiTaskTrainer()
        
        with pytest.raises(RuntimeError):
            trainer.save(tmp_path / "model.pt")

    def test_predictions_match_after_load(self, tmp_path):
        config = TrainerConfig(
            model_type="mlp",
            epochs=3,
            batch_size=16,
            context_window=24,
        )
        trainer = MultiTaskTrainer(config)
        
        X, y = make_synthetic_data(n_samples=100, n_features=600)
        trainer.fit(X, y)
        
        X_test, _ = make_synthetic_data(n_samples=20, n_features=600)
        pred1 = trainer.predict(X_test)
        
        model_path = tmp_path / "test_model.pt"
        trainer.save(model_path)
        
        trainer2 = MultiTaskTrainer()
        trainer2.load(model_path)
        pred2 = trainer2.predict(X_test)
        
        np.testing.assert_array_almost_equal(
            pred1["entry_proba"],
            pred2["entry_proba"],
            decimal=5,
        )


class TestTrainerHistory:
    def test_history_populated(self):
        config = TrainerConfig(
            model_type="mlp",
            epochs=5,
            batch_size=16,
            context_window=24,
        )
        trainer = MultiTaskTrainer(config)
        
        X, y = make_synthetic_data(n_samples=100, n_features=600)
        trainer.fit(X, y)
        
        assert "train_loss" in trainer.history
        assert "val_loss" in trainer.history
        assert "entry_auc" in trainer.history
        assert len(trainer.history["train_loss"]) == 5
