"""
Comprehensive tests for MK3 implementation.

Tests cover:
1. Core components (salience formula, autoencoder, continuous model)
2. Numerical stability and gradient flow
3. Integration tests for full pipeline
4. Performance and correctness validation
"""

import torch
import torch.nn as nn
import sys
import os
import time
import numpy as np

# Add parent to path
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from core.salience_formula import GibbsSalienceFormula
from core.salience_layers import SalienceAttentionLayer, SalienceTransformerBlock
from core.continuous_embeddings import ContinuousEmbedding
from calm.autoencoder import CALMAutoencoder
from calm.continuous_model import ContinuousAutoregressiveModel
from calm.likelihood_free import ContinuousLoss
from utils.tokenizer import SimpleTokenizer


class TestSalienceFormula:
    """Test complete salience formula."""

    def __init__(self):
        self.embedding_dim = 64
        self.batch_size = 4
        self.seq_len = 8

    def test_forward_pass(self):
        """Test basic forward pass."""
        print("\n=== Testing Salience Formula Forward Pass ===")

        formula = GibbsSalienceFormula(
            embedding_dim=self.embedding_dim,
            normalization='softmax'
        )

        # Create inputs
        current = torch.randn(self.batch_size, self.seq_len, self.embedding_dim)
        context = torch.randn(self.batch_size, self.seq_len, self.embedding_dim)

        # Forward pass
        scores, components = formula(current, context, return_components=True)

        # Verify shapes
        assert scores.shape == (self.batch_size, self.seq_len), \
            f"Expected shape {(self.batch_size, self.seq_len)}, got {scores.shape}"

        # Verify normalization (scores should sum to ~1 with softmax)
        score_sums = scores.sum(dim=-1)
        assert torch.allclose(score_sums, torch.ones_like(score_sums), atol=1e-5), \
            f"Scores don't sum to 1 with softmax: {score_sums}"

        # Verify no NaN or Inf
        assert not torch.isnan(scores).any(), "Scores contain NaN"
        assert not torch.isinf(scores).any(), "Scores contain Inf"

        print("✓ Forward pass successful")
        print(f"  Score range: [{scores.min():.4f}, {scores.max():.4f}]")
        print(f"  Score mean: {scores.mean():.4f}")

        return True

    def test_gradient_flow(self):
        """Test gradient flow through salience formula."""
        print("\n=== Testing Salience Formula Gradient Flow ===")

        formula = GibbsSalienceFormula(
            embedding_dim=self.embedding_dim,
            learnable_weights=True,
            learnable_temperature=True
        )

        current = torch.randn(self.batch_size, self.seq_len, self.embedding_dim, requires_grad=True)
        context = torch.randn(self.batch_size, self.seq_len, self.embedding_dim)

        scores, _ = formula(current, context)
        loss = scores.mean()
        loss.backward()

        # Check gradients exist
        assert current.grad is not None, "No gradient for input"
        assert formula.w1.grad is not None, "No gradient for w1"
        assert formula.temperature.grad is not None, "No gradient for temperature"

        # Check gradients are not zero
        assert current.grad.abs().sum() > 0, "Input gradient is zero"

        # Check for NaN in gradients
        assert not torch.isnan(current.grad).any(), "NaN in input gradient"

        print("✓ Gradient flow successful")
        print(f"  Input grad norm: {current.grad.norm():.4f}")
        print(f"  w1 grad: {formula.w1.grad.item():.6f}")
        print(f"  temperature grad: {formula.temperature.grad.item():.6f}")

        return True

    def test_dimensional_scaling(self):
        """Test dimensional scaling invariance."""
        print("\n=== Testing Dimensional Scaling ===")

        # Test with different embedding dimensions
        dims = [32, 64, 128, 256]
        results = []

        for dim in dims:
            formula = GibbsSalienceFormula(
                embedding_dim=dim,
                use_dimensional_scaling=True,
                normalization='none'  # Test raw scores
            )

            current = torch.randn(2, 4, dim)
            context = torch.randn(2, 4, dim)
            scores, _ = formula(current, context, apply_norm=False)
            results.append(scores.mean().item())

        # Scores should be in similar range despite different dimensions
        score_std = np.std(results)
        print(f"  Score means across dimensions: {results}")
        print(f"  Standard deviation: {score_std:.4f}")

        # With dimensional scaling, variance should be low
        assert score_std < 0.5, f"High variance ({score_std}) suggests scaling not working"

        print("✓ Dimensional scaling working")

        return True


class TestAutoencoder:
    """Test CALM autoencoder."""

    def __init__(self):
        self.vocab_size = 1000
        self.embedding_dim = 128
        self.vector_dim = 256
        self.chunk_size = 8
        self.batch_size = 4

    def test_forward_pass(self):
        """Test autoencoder forward pass."""
        print("\n=== Testing Autoencoder Forward Pass ===")

        autoencoder = CALMAutoencoder(
            vocab_size=self.vocab_size,
            embedding_dim=self.embedding_dim,
            vector_dim=self.vector_dim,
            chunk_size=self.chunk_size,
            num_encoder_layers=2,
            num_decoder_layers=2
        )

        # Create input
        token_ids = torch.randint(0, self.vocab_size, (self.batch_size, self.chunk_size))

        # Forward pass
        logits, vectors, metrics = autoencoder(token_ids)

        # Verify shapes
        assert logits.shape == (self.batch_size, self.chunk_size, self.vocab_size), \
            f"Wrong logits shape: {logits.shape}"
        assert vectors.shape == (self.batch_size, self.vector_dim), \
            f"Wrong vectors shape: {vectors.shape}"

        # Verify no NaN
        assert not torch.isnan(logits).any(), "NaN in logits"
        assert not torch.isnan(vectors).any(), "NaN in vectors"

        print("✓ Autoencoder forward pass successful")
        print(f"  Reconstruction accuracy: {metrics['reconstruction_accuracy']:.4f}")
        print(f"  Vector norm: {metrics['vector_norm']:.4f}")

        return True

    def test_reconstruction_quality(self):
        """Test reconstruction quality with training."""
        print("\n=== Testing Autoencoder Reconstruction ===")

        autoencoder = CALMAutoencoder(
            vocab_size=self.vocab_size,
            embedding_dim=self.embedding_dim,
            vector_dim=self.vector_dim,
            chunk_size=self.chunk_size,
            num_encoder_layers=4,
            num_decoder_layers=4
        )

        optimizer = torch.optim.AdamW(autoencoder.parameters(), lr=1e-3)

        # Create fixed training batch
        token_ids = torch.randint(0, self.vocab_size, (8, self.chunk_size))

        print("  Training autoencoder...")
        accuracies = []

        for step in range(200):
            loss, metrics = autoencoder.compute_loss(token_ids)

            optimizer.zero_grad()
            loss.backward()
            torch.nn.utils.clip_grad_norm_(autoencoder.parameters(), 1.0)
            optimizer.step()

            accuracies.append(metrics['reconstruction_accuracy'])

            if step % 50 == 0:
                print(f"  Step {step}: Loss={metrics['loss']:.4f}, Acc={metrics['reconstruction_accuracy']:.4f}")

        final_accuracy = accuracies[-1]
        print(f"  Final accuracy: {final_accuracy:.4f}")

        # After training, should achieve reasonable accuracy
        assert final_accuracy > 0.5, f"Low final accuracy: {final_accuracy}"

        # Accuracy should improve over time
        assert accuracies[-1] > accuracies[0], "Accuracy did not improve"

        print("✓ Autoencoder learns to reconstruct")

        return True

    def test_encode_decode_consistency(self):
        """Test that encode->decode is consistent."""
        print("\n=== Testing Encode-Decode Consistency ===")

        autoencoder = CALMAutoencoder(
            vocab_size=self.vocab_size,
            embedding_dim=self.embedding_dim,
            vector_dim=self.vector_dim,
            chunk_size=self.chunk_size
        )

        token_ids = torch.randint(0, self.vocab_size, (self.batch_size, self.chunk_size))

        # Encode
        vectors = autoencoder.encode(token_ids)

        # Decode
        logits = autoencoder.decode(vectors)
        reconstructed = logits.argmax(dim=-1)

        # Check determinism
        vectors2 = autoencoder.encode(token_ids)
        assert torch.allclose(vectors, vectors2), "Encoding not deterministic"

        print("✓ Encode-decode is consistent")

        return True


class TestContinuousModel:
    """Test continuous autoregressive model."""

    def __init__(self):
        self.vocab_size = 500
        self.embedding_dim = 128
        self.vector_dim = 256
        self.chunk_size = 8
        self.num_layers = 2
        self.batch_size = 2

    def test_forward_pass(self):
        """Test continuous model forward pass."""
        print("\n=== Testing Continuous Model Forward Pass ===")

        model = ContinuousAutoregressiveModel(
            vocab_size=self.vocab_size,
            embedding_dim=self.embedding_dim,
            vector_dim=self.vector_dim,
            chunk_size=self.chunk_size,
            num_layers=self.num_layers,
            num_heads=4
        )

        # Create continuous vector input
        num_chunks = 4
        continuous_vectors = torch.randn(self.batch_size, num_chunks, self.vector_dim)

        # Forward pass
        output = model(continuous_vectors, use_cache=False)

        predicted_vectors = output['predicted_vectors']
        hidden_states = output['hidden_states']

        # Verify shapes
        assert predicted_vectors.shape == (self.batch_size, num_chunks, self.vector_dim), \
            f"Wrong predicted shape: {predicted_vectors.shape}"
        assert hidden_states.shape == (self.batch_size, num_chunks, self.embedding_dim), \
            f"Wrong hidden shape: {hidden_states.shape}"

        # Verify no NaN
        assert not torch.isnan(predicted_vectors).any(), "NaN in predictions"

        print("✓ Continuous model forward pass successful")
        print(f"  Predicted vector norm: {predicted_vectors.norm(dim=-1).mean():.4f}")

        return True

    def test_token_vector_conversion(self):
        """Test token to vector and back conversion."""
        print("\n=== Testing Token-Vector Conversion ===")

        model = ContinuousAutoregressiveModel(
            vocab_size=self.vocab_size,
            embedding_dim=self.embedding_dim,
            vector_dim=self.vector_dim,
            chunk_size=self.chunk_size,
            num_layers=self.num_layers
        )

        # Create token sequence
        seq_len = 32  # Must be multiple of chunk_size
        token_ids = torch.randint(0, self.vocab_size, (self.batch_size, seq_len))

        # Convert to vectors
        continuous_vectors, _ = model.tokenize_to_vectors(token_ids)

        expected_chunks = seq_len // self.chunk_size
        assert continuous_vectors.shape == (self.batch_size, expected_chunks, self.vector_dim), \
            f"Wrong shape after tokenize_to_vectors: {continuous_vectors.shape}"

        # Convert back to tokens
        reconstructed_tokens = model.vectors_to_tokens(continuous_vectors, return_logits=False)

        assert reconstructed_tokens.shape == (self.batch_size, seq_len), \
            f"Wrong shape after vectors_to_tokens: {reconstructed_tokens.shape}"

        print("✓ Token-vector conversion works")
        print(f"  Input shape: {token_ids.shape}")
        print(f"  Vector shape: {continuous_vectors.shape}")
        print(f"  Output shape: {reconstructed_tokens.shape}")

        return True

    def test_loss_computation(self):
        """Test loss computation."""
        print("\n=== Testing Loss Computation ===")

        model = ContinuousAutoregressiveModel(
            vocab_size=self.vocab_size,
            embedding_dim=self.embedding_dim,
            vector_dim=self.vector_dim,
            chunk_size=self.chunk_size,
            num_layers=self.num_layers
        )

        num_chunks = 4
        continuous_vectors = torch.randn(self.batch_size, num_chunks, self.vector_dim)

        loss, metrics = model.compute_loss(continuous_vectors)

        # Verify loss is valid
        assert not torch.isnan(loss), "Loss is NaN"
        assert not torch.isinf(loss), "Loss is Inf"
        assert loss.item() > 0, "Loss is not positive"

        # Test backward pass
        loss.backward()

        # Check gradients exist
        has_grads = False
        for param in model.parameters():
            if param.grad is not None and param.grad.abs().sum() > 0:
                has_grads = True
                break

        assert has_grads, "No gradients computed"

        print("✓ Loss computation successful")
        print(f"  Loss: {loss.item():.4f}")
        print(f"  Cosine similarity: {metrics['cosine_similarity']:.4f}")

        return True


class TestNumericalStability:
    """Test numerical stability under various conditions."""

    def test_large_values(self):
        """Test stability with large input values."""
        print("\n=== Testing Numerical Stability (Large Values) ===")

        formula = GibbsSalienceFormula(
            embedding_dim=64,
            normalization='softmax'
        )

        # Large values
        current = torch.randn(2, 4, 64) * 100
        context = torch.randn(2, 4, 64) * 100

        scores, _ = formula(current, context)

        assert not torch.isnan(scores).any(), "NaN with large values"
        assert not torch.isinf(scores).any(), "Inf with large values"

        print("✓ Stable with large values")

        return True

    def test_small_values(self):
        """Test stability with small input values."""
        print("\n=== Testing Numerical Stability (Small Values) ===")

        formula = GibbsSalienceFormula(
            embedding_dim=64,
            normalization='softmax'
        )

        # Small values
        current = torch.randn(2, 4, 64) * 0.001
        context = torch.randn(2, 4, 64) * 0.001

        scores, _ = formula(current, context)

        assert not torch.isnan(scores).any(), "NaN with small values"
        assert not torch.isinf(scores).any(), "Inf with small values"

        print("✓ Stable with small values")

        return True

    def test_gradient_explosion_check(self):
        """Check for gradient explosion."""
        print("\n=== Testing Gradient Magnitude ===")

        model = ContinuousAutoregressiveModel(
            vocab_size=500,
            embedding_dim=128,
            vector_dim=256,
            chunk_size=8,
            num_layers=4
        )

        continuous_vectors = torch.randn(2, 4, 256)

        loss, _ = model.compute_loss(continuous_vectors)
        loss.backward()

        # Check gradient magnitudes
        max_grad_norm = 0.0
        for param in model.parameters():
            if param.grad is not None:
                grad_norm = param.grad.norm().item()
                max_grad_norm = max(max_grad_norm, grad_norm)

                assert not np.isnan(grad_norm), "NaN gradient detected"
                assert not np.isinf(grad_norm), "Inf gradient detected"

        print(f"  Max gradient norm: {max_grad_norm:.4f}")

        # Gradient shouldn't explode (threshold is somewhat arbitrary)
        assert max_grad_norm < 1000, f"Gradient explosion detected: {max_grad_norm}"

        print("✓ No gradient explosion")

        return True


class TestPerformance:
    """Test performance characteristics."""

    def test_salience_attention_performance(self):
        """Test attention layer performance (identify bottlenecks)."""
        print("\n=== Testing Attention Performance ===")

        # Small test to check if it completes in reasonable time
        attention = SalienceAttentionLayer(
            embedding_dim=64,
            num_heads=2
        )

        # Small inputs
        query = torch.randn(1, 4, 64)  # Only 4 positions to keep it fast

        start_time = time.time()
        output, attn_weights, _ = attention(query, query, query)
        elapsed = time.time() - start_time

        print(f"  Attention forward pass time: {elapsed:.4f}s")

        # Warn if too slow even for tiny inputs
        if elapsed > 1.0:
            print("  ⚠ Warning: Attention is slow even for 4 positions!")
            print("  ⚠ This suggests O(n²) loops instead of batched operations")
        else:
            print("✓ Attention performance acceptable for small inputs")

        return True


class TestIntegration:
    """Integration tests for full pipeline."""

    def test_end_to_end_training_step(self):
        """Test one complete training step."""
        print("\n=== Testing End-to-End Training Step ===")

        # Small model
        model = ContinuousAutoregressiveModel(
            vocab_size=500,
            embedding_dim=64,
            vector_dim=128,
            chunk_size=8,
            num_layers=2,
            num_heads=2
        )

        optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)

        # Create synthetic token data
        token_ids = torch.randint(0, 500, (2, 32))  # 2 batches, 32 tokens

        # Convert to vectors
        continuous_vectors, _ = model.tokenize_to_vectors(token_ids)

        # Forward pass
        loss, metrics = model.compute_loss(continuous_vectors)

        # Backward pass
        optimizer.zero_grad()
        loss.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        optimizer.step()

        print("✓ End-to-end training step successful")
        print(f"  Loss: {loss.item():.4f}")

        return True


def run_all_tests():
    """Run all test suites."""
    print("\n" + "="*70)
    print("MK3 COMPREHENSIVE TEST SUITE")
    print("="*70)

    all_passed = True
    failed_tests = []

    # Test Salience Formula
    print("\n" + "="*70)
    print("TEST SUITE 1: Salience Formula")
    print("="*70)
    try:
        salience_tests = TestSalienceFormula()
        salience_tests.test_forward_pass()
        salience_tests.test_gradient_flow()
        salience_tests.test_dimensional_scaling()
    except Exception as e:
        print(f"✗ Salience Formula tests failed: {e}")
        failed_tests.append(("Salience Formula", str(e)))
        all_passed = False

    # Test Autoencoder
    print("\n" + "="*70)
    print("TEST SUITE 2: Autoencoder")
    print("="*70)
    try:
        ae_tests = TestAutoencoder()
        ae_tests.test_forward_pass()
        ae_tests.test_reconstruction_quality()
        ae_tests.test_encode_decode_consistency()
    except Exception as e:
        print(f"✗ Autoencoder tests failed: {e}")
        failed_tests.append(("Autoencoder", str(e)))
        all_passed = False

    # Test Continuous Model
    print("\n" + "="*70)
    print("TEST SUITE 3: Continuous Model")
    print("="*70)
    try:
        model_tests = TestContinuousModel()
        model_tests.test_forward_pass()
        model_tests.test_token_vector_conversion()
        model_tests.test_loss_computation()
    except Exception as e:
        print(f"✗ Continuous Model tests failed: {e}")
        failed_tests.append(("Continuous Model", str(e)))
        all_passed = False

    # Test Numerical Stability
    print("\n" + "="*70)
    print("TEST SUITE 4: Numerical Stability")
    print("="*70)
    try:
        stability_tests = TestNumericalStability()
        stability_tests.test_large_values()
        stability_tests.test_small_values()
        stability_tests.test_gradient_explosion_check()
    except Exception as e:
        print(f"✗ Numerical Stability tests failed: {e}")
        failed_tests.append(("Numerical Stability", str(e)))
        all_passed = False

    # Test Performance
    print("\n" + "="*70)
    print("TEST SUITE 5: Performance")
    print("="*70)
    try:
        perf_tests = TestPerformance()
        perf_tests.test_salience_attention_performance()
    except Exception as e:
        print(f"✗ Performance tests failed: {e}")
        failed_tests.append(("Performance", str(e)))
        all_passed = False

    # Test Integration
    print("\n" + "="*70)
    print("TEST SUITE 6: Integration")
    print("="*70)
    try:
        integration_tests = TestIntegration()
        integration_tests.test_end_to_end_training_step()
    except Exception as e:
        print(f"✗ Integration tests failed: {e}")
        failed_tests.append(("Integration", str(e)))
        all_passed = False

    # Summary
    print("\n" + "="*70)
    print("TEST SUMMARY")
    print("="*70)

    if all_passed:
        print("✓ ALL TESTS PASSED")
    else:
        print("✗ SOME TESTS FAILED:")
        for test_name, error in failed_tests:
            print(f"\n  {test_name}:")
            print(f"    {error}")

    return all_passed


if __name__ == '__main__':
    success = run_all_tests()
    exit(0 if success else 1)
