"""
Comprehensive verification tests for the experimental architecture.
Tests numerical stability, gradient flow, and correctness.
"""

import torch
import torch.nn as nn
import sys
import os
import numpy as np

sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from core.formula import ScoringFormula
from core.layers import FormulaAttentionLayer, FormulaSelectiveLayer, FormulaTransformerBlock
from model.architecture import NovelAIModel


def test_numerical_stability():
    """Test that the formula doesn't produce NaN or Inf values."""
    print("\n" + "=" * 70)
    print("TEST: Numerical Stability")
    print("=" * 70)
    
    embedding_dim = 128
    batch_size = 4
    seq_len = 16
    
    formula = ScoringFormula(embedding_dim=embedding_dim)
    
    # Test with various input ranges
    test_cases = [
        ("Normal range", torch.randn(batch_size, seq_len, embedding_dim)),
        ("Small values", torch.randn(batch_size, seq_len, embedding_dim) * 1e-6),
        ("Large values", torch.randn(batch_size, seq_len, embedding_dim) * 1e6),
        ("Extreme values", torch.randn(batch_size, seq_len, embedding_dim) * 100),
    ]
    
    all_passed = True
    for test_name, current in test_cases:
        context = current.mean(dim=1, keepdim=True).expand_as(current)
        time_steps = torch.arange(seq_len, dtype=torch.float).unsqueeze(0).expand(batch_size, -1)
        
        try:
            scores, components = formula(current, context, time_steps=time_steps)
            
            has_nan = torch.isnan(scores).any()
            has_inf = torch.isinf(scores).any()
            
            if has_nan or has_inf:
                print(f"[FAIL] {test_name}: Found NaN={has_nan}, Inf={has_inf}")
                all_passed = False
            else:
                print(f"[PASS] {test_name}: No NaN/Inf. Score range: [{scores.min():.4f}, {scores.max():.4f}]")
                
        except Exception as e:
            print(f"[FAIL] {test_name}: Exception {e}")
            all_passed = False
    
    return all_passed


def test_gradient_flow():
    """Test that gradients flow properly through the formula."""
    print("\n" + "=" * 70)
    print("TEST: Gradient Flow")
    print("=" * 70)
    
    embedding_dim = 64
    batch_size = 2
    seq_len = 8
    
    formula = ScoringFormula(embedding_dim=embedding_dim)
    
    current = torch.randn(batch_size, seq_len, embedding_dim, requires_grad=True)
    context = torch.randn(batch_size, seq_len, embedding_dim, requires_grad=True)
    time_steps = torch.arange(seq_len, dtype=torch.float).unsqueeze(0).expand(batch_size, -1)
    
    scores, _ = formula(current, context, time_steps=time_steps)
    
    # Compute loss (mean of scores)
    loss = scores.mean()
    loss.backward()
    
    # Check gradients
    current_grad = current.grad
    context_grad = context.grad
    
    all_passed = True
    
    if current_grad is None:
        print("[FAIL] No gradient for current input")
        all_passed = False
    elif torch.isnan(current_grad).any() or torch.isinf(current_grad).any():
        print("[FAIL] NaN or Inf in current gradients")
        all_passed = False
    else:
        print(f"[PASS] Current gradients flow. Mean: {current_grad.mean():.6f}, Max: {current_grad.abs().max():.6f}")
    
    if context_grad is None:
        print("[FAIL] No gradient for context input")
        all_passed = False
    elif torch.isnan(context_grad).any() or torch.isinf(context_grad).any():
        print("[FAIL] NaN or Inf in context gradients")
        all_passed = False
    else:
        print(f"[PASS] Context gradients flow. Mean: {context_grad.mean():.6f}, Max: {context_grad.abs().max():.6f}")
    
    # Check learnable parameters (some may not have gradients if not used, e.g., fatigue_net when no memory buffer)
    for name, param in formula.named_parameters():
        if param.grad is None:
            # Fatigue net doesn't get gradients when memory buffer is empty
            if 'fatigue_net' in name:
                print(f"[INFO] {name} has no gradient (expected when memory_buffer is None)")
            else:
                print(f"[FAIL] No gradient for parameter {name}")
                all_passed = False
        elif torch.isnan(param.grad).any() or torch.isinf(param.grad).any():
            print(f"[FAIL] NaN or Inf in {name} gradients")
            all_passed = False
        else:
            print(f"[PASS] {name} gradient flows: mean={param.grad.mean():.6f}")
    
    return all_passed


def test_formula_component_consistency():
    """Test that formula components are computed consistently."""
    print("\n" + "=" * 70)
    print("TEST: Formula Component Consistency")
    print("=" * 70)
    
    embedding_dim = 64
    batch_size = 4
    
    formula = ScoringFormula(embedding_dim=embedding_dim)
    
    current = torch.randn(batch_size, embedding_dim)
    context = torch.randn(batch_size, embedding_dim)
    time_steps = torch.arange(batch_size, dtype=torch.float)
    
    scores, components = formula(current, context, time_steps=time_steps)
    
    all_passed = True
    
    # Check all expected components exist
    expected_components = ['novelty', 'retention', 'payoff', 'continuity', 'fatigue', 'weighted_sum']
    for comp_name in expected_components:
        if comp_name not in components:
            print(f"[FAIL] Missing component: {comp_name}")
            all_passed = False
        else:
            comp = components[comp_name]
            if comp is not None:
                if comp.dim() == 0:
                    comp = comp.unsqueeze(0)
                if torch.isnan(comp).any() or torch.isinf(comp).any():
                    print(f"[FAIL] {comp_name} has NaN/Inf")
                    all_passed = False
                else:
                    print(f"[PASS] {comp_name}: range [{comp.min():.4f}, {comp.max():.4f}]")
    
    # Verify formula computation manually
    w1 = abs(formula.w1) / (abs(formula.w1) + abs(formula.w2) + abs(formula.w3) + 1e-8)
    w2 = abs(formula.w2) / (abs(formula.w1) + abs(formula.w2) + abs(formula.w3) + 1e-8)
    w3 = abs(formula.w3) / (abs(formula.w1) + abs(formula.w2) + abs(formula.w3) + 1e-8)
    
    weighted_sum = (w1 * components['novelty'] + 
                    w2 * components['retention'] + 
                    w3 * components['payoff'])
    
    time_decay = torch.exp(-formula.lambda_decay * time_steps)
    fatigue_term = 1 - formula.k_fatigue * components['fatigue']
    
    manual_score = weighted_sum * components['continuity'] * time_decay * fatigue_term
    
    # Allow small numerical differences
    diff = (scores - manual_score).abs().max()
    if diff > 1e-5:
        print(f"[FAIL] Formula computation mismatch. Max diff: {diff:.8f}")
        all_passed = False
    else:
        print(f"[PASS] Formula computation consistent. Max diff: {diff:.8f}")
    
    return all_passed


def test_vectorized_attention_performance():
    """Verify that vectorized attention produces same results as loop version (if old version exists)."""
    print("\n" + "=" * 70)
    print("TEST: Vectorized Attention Correctness")
    print("=" * 70)
    
    embedding_dim = 64
    batch_size = 2
    seq_len = 8
    
    layer = FormulaAttentionLayer(embedding_dim=embedding_dim)
    x = torch.randn(batch_size, seq_len, embedding_dim)
    time_steps = torch.arange(seq_len, dtype=torch.float).unsqueeze(0).expand(batch_size, -1)
    
    # Test self-attention
    output, attn_weights = layer(x, x, x, time_steps=time_steps)
    
    all_passed = True
    
    # Check output shape
    expected_shape = (batch_size, seq_len, embedding_dim)
    if output.shape != expected_shape:
        print(f"[FAIL] Output shape mismatch: {output.shape} vs {expected_shape}")
        all_passed = False
    else:
        print(f"[PASS] Output shape correct: {output.shape}")
    
    # Check attention weights shape
    expected_attn_shape = (batch_size, seq_len, seq_len)
    if attn_weights.shape != expected_attn_shape:
        print(f"[FAIL] Attention weights shape mismatch: {attn_weights.shape} vs {expected_attn_shape}")
        all_passed = False
    else:
        print(f"[PASS] Attention weights shape correct: {attn_weights.shape}")
    
    # Check that attention weights sum approximately to 1 (after dropout they won't sum exactly to 1)
    # Dropout zeros some values, but we can check the mean is close to 1
    attn_sums = attn_weights.sum(dim=-1)
    # Allow some tolerance since dropout modifies the weights
    attn_mean = attn_sums.mean()
    if attn_mean < 0.5 or attn_mean > 1.5:
        print(f"[FAIL] Attention weights sum is abnormal. Mean: {attn_mean:.6f}, Range: [{attn_sums.min():.6f}, {attn_sums.max():.6f}]")
        all_passed = False
    else:
        print(f"[PASS] Attention weights sum is reasonable (after dropout). Mean: {attn_mean:.6f}, Range: [{attn_sums.min():.6f}, {attn_sums.max():.6f}]")
    
    # Check for NaN/Inf
    if torch.isnan(output).any() or torch.isinf(output).any():
        print("[FAIL] NaN/Inf in attention output")
        all_passed = False
    else:
        print(f"[PASS] No NaN/Inf in output. Range: [{output.min():.4f}, {output.max():.4f}]")
    
    if torch.isnan(attn_weights).any() or torch.isinf(attn_weights).any():
        print("[FAIL] NaN/Inf in attention weights")
        all_passed = False
    else:
        print(f"[PASS] No NaN/Inf in attention weights. Range: [{attn_weights.min():.4f}, {attn_weights.max():.4f}]")
    
    return all_passed


def test_memory_buffer_handling():
    """Test memory buffer updates and fatigue computation."""
    print("\n" + "=" * 70)
    print("TEST: Memory Buffer Handling")
    print("=" * 70)
    
    embedding_dim = 64
    batch_size = 2
    seq_len = 8
    
    formula = ScoringFormula(embedding_dim=embedding_dim)
    
    # Test with empty memory buffer
    current = torch.randn(batch_size, embedding_dim)
    context = torch.randn(batch_size, embedding_dim)
    
    scores_no_mem, comp_no_mem = formula(current, context, memory_buffer=None)
    print(f"[PASS] Formula works without memory buffer")
    
    # Test with memory buffer
    memory_buffer = torch.randn(10, embedding_dim)
    scores_with_mem, comp_with_mem = formula(current, context, memory_buffer=memory_buffer)
    print(f"[PASS] Formula works with memory buffer")
    
    # Fatigue should be higher with similar items in memory
    # Create current similar to memory buffer
    similar_current = memory_buffer[0:1].expand(batch_size, -1)
    scores_similar, comp_similar = formula(similar_current, context, memory_buffer=memory_buffer)
    
    # Fatigue should be higher for similar items
    if comp_similar['fatigue'].mean() > comp_with_mem['fatigue'].mean():
        print(f"[PASS] Fatigue correctly higher for similar items")
    else:
        print(f"[INFO] Fatigue comparison: similar={comp_similar['fatigue'].mean():.4f}, random={comp_with_mem['fatigue'].mean():.4f}")
    
    return True


def test_model_end_to_end():
    """Test full model forward pass and generation."""
    print("\n" + "=" * 70)
    print("TEST: Full Model End-to-End")
    print("=" * 70)
    
    vocab_size = 100
    embedding_dim = 64
    num_layers = 2
    max_seq_length = 32
    
    model = NovelAIModel(
        vocab_size=vocab_size,
        embedding_dim=embedding_dim,
        num_layers=num_layers,
        max_seq_length=max_seq_length,
        use_memory_buffer=True
    )
    
    batch_size = 2
    seq_len = 8
    input_ids = torch.randint(0, vocab_size, (batch_size, seq_len))
    
    all_passed = True
    
    # Test forward pass
    try:
        output = model(input_ids=input_ids)
        
        assert 'logits' in output
        assert 'hidden_states' in output
        assert output['logits'].shape == (batch_size, seq_len, vocab_size)
        
        if torch.isnan(output['logits']).any() or torch.isinf(output['logits']).any():
            print("[FAIL] NaN/Inf in model output logits")
            all_passed = False
        else:
            print(f"[PASS] Forward pass successful. Logits shape: {output['logits'].shape}")
            
    except Exception as e:
        print(f"[FAIL] Forward pass failed: {e}")
        import traceback
        traceback.print_exc()
        all_passed = False
    
    # Test generation
    try:
        model.eval()
        with torch.no_grad():
            prompt = torch.randint(0, vocab_size, (1, 3))
            generated = model.generate(prompt, max_new_tokens=5, temperature=1.0)
            
            assert generated.shape[0] == 1
            assert generated.shape[1] >= prompt.shape[1]
            print(f"[PASS] Generation successful. Output length: {generated.shape[1]}")
            
    except Exception as e:
        print(f"[FAIL] Generation failed: {e}")
        import traceback
        traceback.print_exc()
        all_passed = False
    
    return all_passed


def main():
    """Run all verification tests."""
    print("=" * 70)
    print("EXPERIMENTAL ARCHITECTURE - VERIFICATION TEST SUITE")
    print("=" * 70)
    
    tests = [
        ("Numerical Stability", test_numerical_stability),
        ("Gradient Flow", test_gradient_flow),
        ("Formula Component Consistency", test_formula_component_consistency),
        ("Vectorized Attention Correctness", test_vectorized_attention_performance),
        ("Memory Buffer Handling", test_memory_buffer_handling),
        ("Full Model End-to-End", test_model_end_to_end),
    ]
    
    results = []
    for test_name, test_func in tests:
        try:
            result = test_func()
            results.append((test_name, result))
        except Exception as e:
            print(f"\n[ERROR] {test_name} crashed: {e}")
            import traceback
            traceback.print_exc()
            results.append((test_name, False))
    
    # Summary
    print("\n" + "=" * 70)
    print("VERIFICATION SUMMARY")
    print("=" * 70)
    
    passed = sum(1 for _, result in results if result)
    total = len(results)
    
    for test_name, result in results:
        status = "PASS" if result else "FAIL"
        print(f"{test_name:40s}: {status}")
    
    print("-" * 70)
    print(f"Total: {passed}/{total} verification tests passed")
    
    if passed == total:
        print("\n[SUCCESS] All verification tests passed!")
        return 0
    else:
        print(f"\n[WARNING] {total - passed} verification test(s) failed")
        return 1


if __name__ == '__main__':
    exit_code = main()
    sys.exit(exit_code)

