"""Tests for multi-task model architectures."""

import numpy as np
import torch

from ai.model_v2 import (
    LSTMMultiTaskModel,
    MLPMultiTaskModel,
    MultiTaskOutput,
    TransformerMultiTaskModel,
    get_multi_task_model,
)


class TestMultiTaskOutput:
    def test_output_creation(self):
        output = MultiTaskOutput(
            entry_long_logits=torch.tensor([0.5, -0.5]),
            entry_short_logits=torch.tensor([-0.5, 0.5]),
            sl_distance=torch.tensor([1.0, 1.5]),
            tp_distance=torch.tensor([2.0, 2.5]),
            confidence=torch.tensor([0.8, 0.6]),
        )
        
        assert output.entry_long_logits.shape == (2,)
        assert output.entry_short_logits.shape == (2,)
        assert output.sl_distance.shape == (2,)
        assert output.tp_distance.shape == (2,)
        assert output.confidence.shape == (2,)

    def test_entry_proba(self):
        output = MultiTaskOutput(
            entry_long_logits=torch.tensor([0.0, 2.0, -2.0]),
            entry_short_logits=torch.tensor([-2.0, 0.0, 2.0]),
            sl_distance=torch.tensor([1.0, 1.0, 1.0]),
            tp_distance=torch.tensor([2.0, 2.0, 2.0]),
            confidence=torch.tensor([0.5, 0.5, 0.5]),
        )
        
        proba = output.entry_long_proba
        assert proba.shape == (3,)
        assert abs(proba[0].item() - 0.5) < 0.01
        assert proba[1].item() > 0.8
        assert proba[2].item() < 0.2
        
        # entry_proba backward compat = entry_long_proba
        assert torch.allclose(output.entry_proba, output.entry_long_proba)

    def test_to_dict(self):
        output = MultiTaskOutput(
            entry_long_logits=torch.tensor([0.5]),
            entry_short_logits=torch.tensor([-0.5]),
            sl_distance=torch.tensor([1.0]),
            tp_distance=torch.tensor([2.0]),
            confidence=torch.tensor([0.8]),
        )
        
        d = output.to_dict()
        assert "entry_long_logits" in d
        assert "entry_short_logits" in d
        assert "entry_long_proba" in d
        assert "entry_short_proba" in d
        assert "sl_distance" in d
        assert "tp_distance" in d
        assert "confidence" in d

    def test_to_numpy(self):
        output = MultiTaskOutput(
            entry_long_logits=torch.tensor([0.5]),
            entry_short_logits=torch.tensor([-0.5]),
            sl_distance=torch.tensor([1.0]),
            tp_distance=torch.tensor([2.0]),
            confidence=torch.tensor([0.8]),
        )
        
        arr = output.to_numpy()
        assert isinstance(arr["entry_long_proba"], np.ndarray)
        assert isinstance(arr["entry_short_proba"], np.ndarray)
        assert isinstance(arr["sl_distance"], np.ndarray)


class TestLSTMMultiTaskModel:
    def test_model_creation(self):
        model = LSTMMultiTaskModel(
            input_size=600,
            context_window=24,
            hidden_size=64,
            num_layers=2,
        )
        
        assert model.input_size == 600
        assert model.context_window == 24
        assert model.hidden_size == 64
        assert model.n_per_bar == 7
        assert model.n_seq == 168
        assert model.n_global == 432  # 600 - 168

    def test_forward_pass(self):
        model = LSTMMultiTaskModel(
            input_size=600,
            context_window=24,
            hidden_size=64,
        )
        
        x = torch.randn(8, 600)
        output = model(x)
        
        assert output.entry_long_logits.shape == (8,)
        assert output.entry_short_logits.shape == (8,)
        assert output.sl_distance.shape == (8,)
        assert output.tp_distance.shape == (8,)
        assert output.confidence.shape == (8,)

    def test_sl_tp_clamping(self):
        model = LSTMMultiTaskModel(
            input_size=600,
            context_window=24,
        )
        
        x = torch.randn(16, 600) * 10
        output = model(x)
        
        assert torch.all(output.sl_distance >= model.SL_MIN)
        assert torch.all(output.sl_distance <= model.SL_MAX)
        assert torch.all(output.tp_distance >= model.TP_MIN)
        assert torch.all(output.tp_distance <= model.TP_MAX)

    def test_confidence_range(self):
        model = LSTMMultiTaskModel(
            input_size=600,
            context_window=24,
        )
        
        x = torch.randn(16, 600)
        output = model(x)
        
        assert torch.all(output.confidence >= 0)
        assert torch.all(output.confidence <= 1)

    def test_predict_numpy(self):
        model = LSTMMultiTaskModel(
            input_size=600,
            context_window=24,
        )
        
        x = np.random.randn(8, 600).astype(np.float32)
        result = model.predict(x, entry_threshold=0.5)
        
        assert "entry_long_proba" in result
        assert "entry_short_proba" in result
        assert "entry_long_signal" in result
        assert "entry_short_signal" in result
        assert "sl_distance" in result
        assert "tp_distance" in result
        assert "confidence" in result
        
        assert result["entry_long_proba"].shape == (8,)
        assert result["entry_long_signal"].dtype == int


class TestTransformerMultiTaskModel:
    def test_model_creation(self):
        model = TransformerMultiTaskModel(
            input_size=600,
            context_window=24,
            d_model=64,
            n_heads=4,
            n_layers=2,
        )
        
        assert model.input_size == 600
        assert model.d_model == 64
        assert model.n_heads == 4

    def test_forward_pass(self):
        model = TransformerMultiTaskModel(
            input_size=600,
            context_window=24,
            d_model=64,
            n_heads=4,
        )
        
        x = torch.randn(8, 600)
        output = model(x)
        
        assert output.entry_long_logits.shape == (8,)
        assert output.entry_short_logits.shape == (8,)
        assert output.sl_distance.shape == (8,)
        assert output.tp_distance.shape == (8,)
        assert output.confidence.shape == (8,)

    def test_sl_tp_clamping(self):
        model = TransformerMultiTaskModel(
            input_size=600,
            context_window=24,
        )
        
        x = torch.randn(16, 600) * 10
        output = model(x)
        
        assert torch.all(output.sl_distance >= model.SL_MIN)
        assert torch.all(output.sl_distance <= model.SL_MAX)
        assert torch.all(output.tp_distance >= model.TP_MIN)
        assert torch.all(output.tp_distance <= model.TP_MAX)


class TestMLPMultiTaskModel:
    def test_model_creation(self):
        model = MLPMultiTaskModel(
            input_size=600,
            context_window=24,
            hidden_sizes=[128, 64],
        )
        
        assert model.input_size == 600

    def test_forward_pass(self):
        model = MLPMultiTaskModel(
            input_size=600,
            context_window=24,
        )
        
        x = torch.randn(8, 600)
        output = model(x)
        
        assert output.entry_long_logits.shape == (8,)
        assert output.entry_short_logits.shape == (8,)
        assert output.sl_distance.shape == (8,)
        assert output.tp_distance.shape == (8,)
        assert output.confidence.shape == (8,)

    def test_predict_numpy(self):
        model = MLPMultiTaskModel(
            input_size=600,
            context_window=24,
        )
        
        x = np.random.randn(8, 600).astype(np.float32)
        result = model.predict(x)
        
        assert "entry_long_proba" in result
        assert "entry_short_proba" in result
        assert "sl_distance" in result


class TestGetMultiTaskModel:
    def test_get_lstm_model(self):
        model = get_multi_task_model(
            "lstm",
            input_size=600,
            context_window=24,
        )
        assert isinstance(model, LSTMMultiTaskModel)

    def test_get_transformer_model(self):
        model = get_multi_task_model(
            "transformer",
            input_size=600,
            context_window=24,
        )
        assert isinstance(model, TransformerMultiTaskModel)

    def test_get_mlp_model(self):
        model = get_multi_task_model(
            "mlp",
            input_size=600,
            context_window=24,
        )
        assert isinstance(model, MLPMultiTaskModel)

    def test_unknown_model_defaults_to_lstm(self):
        model = get_multi_task_model(
            "unknown",
            input_size=600,
            context_window=24,
        )
        assert isinstance(model, LSTMMultiTaskModel)

    def test_case_insensitive(self):
        model = get_multi_task_model("LSTM", input_size=600, context_window=24)
        assert isinstance(model, LSTMMultiTaskModel)
        
        model = get_multi_task_model("Transformer", input_size=600, context_window=24)
        assert isinstance(model, TransformerMultiTaskModel)


class TestModelGradients:
    def test_lstm_gradients_flow(self):
        model = LSTMMultiTaskModel(input_size=600, context_window=24)
        
        x = torch.randn(4, 600, requires_grad=True)
        output = model(x)
        
        loss = output.entry_long_logits.sum() + output.entry_short_logits.sum() + output.sl_distance.sum()
        loss.backward()
        
        assert x.grad is not None
        assert x.grad.shape == x.shape

    def test_transformer_gradients_flow(self):
        model = TransformerMultiTaskModel(input_size=600, context_window=24)
        
        x = torch.randn(4, 600, requires_grad=True)
        output = model(x)
        
        loss = output.entry_long_logits.sum() + output.entry_short_logits.sum() + output.tp_distance.sum()
        loss.backward()
        
        assert x.grad is not None


class TestModelSaveLoad:
    def test_model_state_dict(self):
        model = LSTMMultiTaskModel(input_size=600, context_window=24)
        
        state_dict = model.state_dict()
        assert len(state_dict) > 0
        
        model2 = LSTMMultiTaskModel(input_size=600, context_window=24)
        model2.load_state_dict(state_dict)
        
        model.eval()
        model2.eval()
        
        x = torch.randn(4, 600)
        out1 = model(x)
        out2 = model2(x)
        
        torch.testing.assert_close(out1.entry_long_logits, out2.entry_long_logits)
        torch.testing.assert_close(out1.entry_short_logits, out2.entry_short_logits)
