"""
Comprehensive test suite for the novel AI model.
"""

import torch
import sys
import os

sys.path.append(os.path.dirname(os.path.abspath(__file__)))

def test_imports():
    """Test all imports."""
    print("=" * 70)
    print("TEST 1: Import Tests")
    print("=" * 70)
    
    try:
        from core.formula import ScoringFormula
        print("[PASS] Core formula import")
    except Exception as e:
        print(f"[FAIL] Core formula import: {e}")
        return False
    
    try:
        from core.layers import FormulaAttentionLayer, FormulaSelectiveLayer
        print("[PASS] Core layers import")
    except Exception as e:
        print(f"[FAIL] Core layers import: {e}")
        return False
    
    try:
        from model.architecture import NovelAIModel
        print("[PASS] Model architecture import")
    except Exception as e:
        print(f"[FAIL] Model architecture import: {e}")
        return False
    
    try:
        from config.config import Config
        print("[PASS] Config import")
    except Exception as e:
        print(f"[FAIL] Config import: {e}")
        return False
    
    try:
        from training.trainer import Trainer
        print("[PASS] Trainer import")
    except Exception as e:
        print(f"[FAIL] Trainer import: {e}")
        return False
    
    try:
        from utils.tokenizer import SimpleTokenizer
        from utils.data_loader import create_data_loaders
        print("[PASS] Utils import")
    except Exception as e:
        print(f"[FAIL] Utils import: {e}")
        return False
    
    return True


def test_formula():
    """Test the scoring formula."""
    print("\n" + "=" * 70)
    print("TEST 2: Scoring Formula Functionality")
    print("=" * 70)
    
    from core.formula import ScoringFormula
    
    try:
        formula = ScoringFormula(
            embedding_dim=128,
            w1=0.4,
            w2=0.3,
            w3=0.3,
            lambda_decay=0.1,
            k_fatigue=0.2
        )
        print("[PASS] Formula initialization")
    except Exception as e:
        print(f"[FAIL] Formula initialization: {e}")
        return False
    
    try:
        # Test forward pass
        batch_size = 4
        current = torch.randn(batch_size, 128)
        context = torch.randn(batch_size, 128)
        time_steps = torch.arange(batch_size, dtype=torch.float)
        
        scores, components = formula(current, context, time_steps=time_steps)
        
        assert scores.shape == (batch_size,), f"Expected shape ({batch_size},), got {scores.shape}"
        assert 'novelty' in components
        assert 'retention' in components
        assert 'payoff' in components
        assert 'continuity' in components
        assert 'fatigue' in components
        
        print("[PASS] Formula forward pass")
        print(f"      Output shape: {scores.shape}")
        print(f"      Score range: [{scores.min():.4f}, {scores.max():.4f}]")
        
    except Exception as e:
        print(f"[FAIL] Formula forward pass: {e}")
        import traceback
        traceback.print_exc()
        return False
    
    try:
        # Test with sequence input
        current_seq = torch.randn(2, 5, 128)  # [batch, seq_len, dim]
        context_seq = torch.randn(2, 5, 128)
        time_steps_seq = torch.arange(5, dtype=torch.float).unsqueeze(0).expand(2, -1)
        
        scores_seq, _ = formula(current_seq, context_seq, time_steps=time_steps_seq)
        
        assert scores_seq.shape == (2, 5), f"Expected shape (2, 5), got {scores_seq.shape}"
        print("[PASS] Formula sequence input")
        
    except Exception as e:
        print(f"[FAIL] Formula sequence input: {e}")
        import traceback
        traceback.print_exc()
        return False
    
    try:
        # Test with memory buffer
        memory_buffer = torch.randn(10, 128)
        scores_mem, _ = formula(current, context, time_steps=time_steps, memory_buffer=memory_buffer)
        assert scores_mem.shape == (batch_size,)
        print("[PASS] Formula with memory buffer")
        
    except Exception as e:
        print(f"[FAIL] Formula with memory buffer: {e}")
        return False
    
    return True


def test_model():
    """Test the full model."""
    print("\n" + "=" * 70)
    print("TEST 3: Model Architecture")
    print("=" * 70)
    
    from model.architecture import NovelAIModel
    from config.config import Config
    
    try:
        config = Config.default()
        config.model.vocab_size = 100
        config.model.embedding_dim = 64
        config.model.num_layers = 2
        config.model.max_seq_length = 32
        
        formula_config = {
            'w1': 0.4,
            'w2': 0.3,
            'w3': 0.3,
            'lambda_decay': 0.1,
            'k_fatigue': 0.2,
        }
        
        model = NovelAIModel(
            vocab_size=config.model.vocab_size,
            embedding_dim=config.model.embedding_dim,
            num_layers=config.model.num_layers,
            max_seq_length=config.model.max_seq_length,
            formula_config=formula_config
        )
        print("[PASS] Model initialization")
        
        param_count = sum(p.numel() for p in model.parameters())
        print(f"      Parameters: {param_count:,}")
        
    except Exception as e:
        print(f"[FAIL] Model initialization: {e}")
        import traceback
        traceback.print_exc()
        return False
    
    try:
        # Test forward pass
        batch_size = 2
        seq_len = 8
        input_ids = torch.randint(0, config.model.vocab_size, (batch_size, seq_len))
        
        output = model(input_ids=input_ids, return_components=False)
        
        assert 'logits' in output
        assert 'hidden_states' in output
        assert output['logits'].shape == (batch_size, seq_len, config.model.vocab_size)
        assert output['hidden_states'].shape == (batch_size, seq_len, config.model.embedding_dim)
        
        print("[PASS] Model forward pass")
        print(f"      Logits shape: {output['logits'].shape}")
        print(f"      Hidden states shape: {output['hidden_states'].shape}")
        
    except Exception as e:
        print(f"[FAIL] Model forward pass: {e}")
        import traceback
        traceback.print_exc()
        return False
    
    try:
        # Test generation
        model.eval()
        prompt = torch.randint(0, config.model.vocab_size, (1, 3))
        
        with torch.no_grad():
            generated = model.generate(
                input_ids=prompt,
                max_new_tokens=5,
                temperature=1.0,
                do_sample=True
            )
        
        assert generated.shape[0] == 1
        assert generated.shape[1] >= prompt.shape[1]
        print("[PASS] Model generation")
        print(f"      Generated length: {generated.shape[1]}")
        
    except Exception as e:
        print(f"[FAIL] Model generation: {e}")
        import traceback
        traceback.print_exc()
        return False
    
    return True


def test_layers():
    """Test the formula-based layers."""
    print("\n" + "=" * 70)
    print("TEST 4: Formula-Based Layers")
    print("=" * 70)
    
    from core.layers import FormulaAttentionLayer, FormulaSelectiveLayer, FormulaTransformerBlock
    
    embedding_dim = 64
    batch_size = 2
    seq_len = 8
    
    try:
        # Test FormulaAttentionLayer
        attn_layer = FormulaAttentionLayer(embedding_dim=embedding_dim)
        x = torch.randn(batch_size, seq_len, embedding_dim)
        
        output, attn_weights = attn_layer(x, x, x)
        
        assert output.shape == (batch_size, seq_len, embedding_dim)
        assert attn_weights.shape == (batch_size, seq_len, seq_len)
        print("[PASS] FormulaAttentionLayer")
        print(f"      Output shape: {output.shape}")
        print(f"      Attention weights shape: {attn_weights.shape}")
        
    except Exception as e:
        print(f"[FAIL] FormulaAttentionLayer: {e}")
        import traceback
        traceback.print_exc()
        return False
    
    try:
        # Test FormulaSelectiveLayer
        selective_layer = FormulaSelectiveLayer(embedding_dim=embedding_dim)
        x = torch.randn(batch_size, seq_len, embedding_dim)
        context = torch.randn(batch_size, embedding_dim)
        
        output, mask = selective_layer(x, context)
        
        assert output.shape == (batch_size, seq_len, embedding_dim)
        assert mask.shape == (batch_size, seq_len)
        print("[PASS] FormulaSelectiveLayer")
        print(f"      Output shape: {output.shape}")
        print(f"      Selection mask shape: {mask.shape}")
        print(f"      Selected items: {mask.sum().item()}/{mask.numel()}")
        
    except Exception as e:
        print(f"[FAIL] FormulaSelectiveLayer: {e}")
        import traceback
        traceback.print_exc()
        return False
    
    try:
        # Test FormulaTransformerBlock
        block = FormulaTransformerBlock(embedding_dim=embedding_dim)
        x = torch.randn(batch_size, seq_len, embedding_dim)
        
        output, attn_weights = block(x)
        
        assert output.shape == (batch_size, seq_len, embedding_dim)
        print("[PASS] FormulaTransformerBlock")
        print(f"      Output shape: {output.shape}")
        
    except Exception as e:
        print(f"[FAIL] FormulaTransformerBlock: {e}")
        import traceback
        traceback.print_exc()
        return False
    
    return True


def test_training_setup():
    """Test training infrastructure."""
    print("\n" + "=" * 70)
    print("TEST 5: Training Infrastructure")
    print("=" * 70)
    
    from utils.tokenizer import SimpleTokenizer
    from utils.data_loader import create_data_loaders
    from training.trainer import Trainer
    from model.architecture import NovelAIModel
    from config.config import Config
    
    try:
        # Create sample data
        train_texts = [
            "This is a sample training text.",
            "The model uses a novel scoring formula.",
            "Each component has a specific meaning.",
        ]
        
        val_texts = [
            "This is a validation text.",
            "It tests generalization.",
        ]
        
        # Create data loaders
        train_loader, val_loader, tokenizer = create_data_loaders(
            train_texts=train_texts,
            val_texts=val_texts,
            max_length=16,
            batch_size=2
        )
        
        assert len(train_loader) > 0
        print(f"[PASS] Data loaders created")
        print(f"      Training batches: {len(train_loader)}")
        if val_loader:
            print(f"      Validation batches: {len(val_loader)}")
        
        # Test a batch
        sample_batch = next(iter(train_loader))
        assert 'input_ids' in sample_batch
        assert 'target_ids' in sample_batch
        assert 'attention_mask' in sample_batch
        print(f"[PASS] Batch structure correct")
        print(f"      Batch input shape: {sample_batch['input_ids'].shape}")
        
    except Exception as e:
        print(f"[FAIL] Data loading: {e}")
        import traceback
        traceback.print_exc()
        return False
    
    try:
        # Create model for training
        config = Config.default()
        config.model.vocab_size = tokenizer.vocab_size
        config.model.embedding_dim = 64
        config.model.num_layers = 2
        config.model.max_seq_length = 16
        
        formula_config = {
            'w1': 0.4,
            'w2': 0.3,
            'w3': 0.3,
        }
        
        model = NovelAIModel(
            vocab_size=config.model.vocab_size,
            embedding_dim=config.model.embedding_dim,
            num_layers=config.model.num_layers,
            max_seq_length=config.model.max_seq_length,
            formula_config=formula_config
        )
        
        # Create trainer
        trainer = Trainer(
            model=model,
            train_loader=train_loader,
            val_loader=val_loader,
            learning_rate=1e-4,
            device='cpu',  # Use CPU for testing
            save_dir='./test_checkpoints'
        )
        
        print("[PASS] Trainer initialization")
        
    except Exception as e:
        print(f"[FAIL] Trainer initialization: {e}")
        import traceback
        traceback.print_exc()
        return False
    
    try:
        # Test a single training step
        trainer.model.train()
        batch = next(iter(train_loader))
        
        input_ids = batch['input_ids'].to(trainer.device)
        target_ids = batch['target_ids'].to(trainer.device)
        attention_mask = batch['attention_mask'].to(trainer.device)
        
        output = trainer.model(input_ids=input_ids, attention_mask=attention_mask)
        logits = output['logits']
        
        logits_flat = logits.view(-1, logits.shape[-1])
        targets_flat = target_ids.view(-1)
        
        loss = trainer.criterion(logits_flat, targets_flat)
        
        # Backward pass
        trainer.optimizer.zero_grad()
        loss.backward()
        trainer.optimizer.step()
        
        print("[PASS] Training step executed")
        print(f"      Loss: {loss.item():.4f}")
        
    except Exception as e:
        print(f"[FAIL] Training step: {e}")
        import traceback
        traceback.print_exc()
        return False
    
    return True


def test_config():
    """Test configuration management."""
    print("\n" + "=" * 70)
    print("TEST 6: Configuration Management")
    print("=" * 70)
    
    from config.config import Config
    
    try:
        # Test default config
        config = Config.default()
        assert config.model is not None
        assert config.training is not None
        print("[PASS] Default config creation")
        
    except Exception as e:
        print(f"[FAIL] Default config: {e}")
        return False
    
    try:
        # Test config save/load
        import tempfile
        import os
        
        with tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False) as f:
            temp_path = f.name
        
        try:
            config.save(temp_path)
            assert os.path.exists(temp_path)
            print("[PASS] Config save")
            
            loaded_config = Config.load(temp_path)
            assert loaded_config.model.embedding_dim == config.model.embedding_dim
            print("[PASS] Config load")
            
        finally:
            if os.path.exists(temp_path):
                os.remove(temp_path)
        
    except Exception as e:
        print(f"[FAIL] Config save/load: {e}")
        import traceback
        traceback.print_exc()
        return False
    
    return True


def main():
    """Run all tests."""
    print("\n" + "=" * 70)
    print("NOVEL AI MODEL - COMPREHENSIVE TEST SUITE")
    print("=" * 70)
    
    tests = [
        ("Imports", test_imports),
        ("Scoring Formula", test_formula),
        ("Model Architecture", test_model),
        ("Formula Layers", test_layers),
        ("Training Infrastructure", test_training_setup),
        ("Configuration", test_config),
    ]
    
    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("TEST 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:30s}: {status}")
    
    print("-" * 70)
    print(f"Total: {passed}/{total} tests passed")
    
    if passed == total:
        print("\n[SUCCESS] All tests passed!")
        return 0
    else:
        print(f"\n[FAILURE] {total - passed} test(s) failed")
        return 1


if __name__ == '__main__':
    exit_code = main()
    sys.exit(exit_code)



