"""
Quick test to verify Gibbs-normalized salience formula works correctly.
"""

import torch
import sys
import os

# Add parent to path
sys.path.append(os.path.dirname(os.path.abspath(__file__)))

from core.salience_formula import GibbsSalienceFormula


def test_basic_forward():
    """Test basic forward pass."""
    print("\n" + "="*60)
    print("Testing Gibbs Salience Formula - Basic Forward Pass")
    print("="*60)

    formula = GibbsSalienceFormula(
        embedding_dim=64,
        learnable_coefficients=True,
        learnable_temperature=True,
        learnable_budget=False,
        capacity_budget=1.0
    )

    # Create test inputs
    batch_size, seq_len, dim = 2, 8, 64
    current = torch.randn(batch_size, seq_len, dim)
    context = torch.randn(batch_size, seq_len, dim)
    time_steps = torch.arange(seq_len).float().unsqueeze(0).repeat(batch_size, 1)

    # Forward pass
    scores, components = formula(
        current=current,
        context=context,
        time_steps=time_steps,
        return_components=True
    )

    print(f"\nInput shape: {current.shape}")
    print(f"Output shape: {scores.shape}")
    print(f"Output range: [{scores.min().item():.6f}, {scores.max().item():.6f}]")
    print(f"Sum per scope: {scores.sum(dim=-1).tolist()}")
    print(f"Expected budget: {formula.capacity_budget.item()}")

    # Check components
    print("\nComponent statistics:")
    for name, values in components.items():
        if isinstance(values, torch.Tensor):
            print(f"  {name}: mean={values.mean().item():.4f}, std={values.std().item():.4f}")

    print("\n✓ Basic forward pass successful!")
    return True


def test_sanity_checks():
    """Test all sanity checks."""
    print("\n" + "="*60)
    print("Testing Gibbs Salience Formula - Sanity Checks")
    print("="*60)

    formula = GibbsSalienceFormula(
        embedding_dim=64,
        learnable_coefficients=False,  # Use fixed coefficients for testing
        learnable_temperature=False,
        learnable_budget=False,
        capacity_budget=1.0
    )

    # Create test inputs
    batch_size, seq_len, dim = 2, 8, 64
    current = torch.randn(batch_size, seq_len, dim)
    context = torch.randn(batch_size, seq_len, dim)
    time_steps = torch.arange(seq_len).float().unsqueeze(0).repeat(batch_size, 1)

    # Run all sanity checks
    results = formula.run_all_sanity_checks(
        current=current,
        context=context,
        time_steps=time_steps,
        verbose=True
    )

    return results


def test_gradient_flow():
    """Test gradient flow through the formula."""
    print("\n" + "="*60)
    print("Testing Gibbs Salience Formula - Gradient Flow")
    print("="*60)

    formula = GibbsSalienceFormula(
        embedding_dim=64,
        learnable_coefficients=True,
        learnable_temperature=True,
        learnable_budget=True
    )

    # Create test inputs
    batch_size, seq_len, dim = 2, 8, 64
    current = torch.randn(batch_size, seq_len, dim, requires_grad=True)
    context = torch.randn(batch_size, seq_len, dim, requires_grad=True)

    # Forward pass
    scores, _ = formula(current, context)

    # Compute loss
    loss = scores.mean()

    # Backward pass
    loss.backward()

    print(f"\nLoss: {loss.item():.6f}")
    print(f"Current gradient norm: {current.grad.norm().item():.6f}")
    print(f"Context gradient norm: {context.grad.norm().item():.6f}")

    # Check parameter gradients
    print("\nParameter gradients:")
    for name, param in formula.named_parameters():
        if param.grad is not None:
            print(f"  {name}: {param.grad.norm().item():.6f}")

    print("\n✓ Gradient flow successful!")
    return True


def test_component_ranges():
    """Test that components output appropriate ranges."""
    print("\n" + "="*60)
    print("Testing Gibbs Salience Formula - Component Ranges")
    print("="*60)

    formula = GibbsSalienceFormula(embedding_dim=64)

    batch_size, seq_len, dim = 2, 8, 64
    current = torch.randn(batch_size, seq_len, dim)
    context = torch.randn(batch_size, seq_len, dim)

    current_flat = current.view(-1, dim)
    context_flat = context.view(-1, dim)

    # Test novelty (unbounded)
    novelty = formula.compute_novelty(current_flat, context_flat)
    print(f"\nNovelty (unbounded logit space):")
    print(f"  Range: [{novelty.min().item():.4f}, {novelty.max().item():.4f}]")
    print(f"  Mean: {novelty.mean().item():.4f}, Std: {novelty.std().item():.4f}")

    # Test retention (unbounded)
    retention = formula.compute_retention(current_flat)
    print(f"\nRetention (unbounded logit space):")
    print(f"  Range: [{retention.min().item():.4f}, {retention.max().item():.4f}]")
    print(f"  Mean: {retention.mean().item():.4f}, Std: {retention.std().item():.4f}")

    # Test payoff (unbounded)
    payoff = formula.compute_payoff(current_flat)
    print(f"\nPayoff (unbounded logit space):")
    print(f"  Range: [{payoff.min().item():.4f}, {payoff.max().item():.4f}]")
    print(f"  Mean: {payoff.mean().item():.4f}, Std: {payoff.std().item():.4f}")

    # Test continuity (positive for log)
    continuity = formula.compute_continuity(current_flat, context_flat)
    print(f"\nContinuity (positive, for ln(C+ε)):")
    print(f"  Range: [{continuity.min().item():.4f}, {continuity.max().item():.4f}]")
    print(f"  All positive: {torch.all(continuity > 0).item()}")

    # Test fatigue (positive for subtraction)
    memory_buffer = torch.randn(10, dim)
    fatigue = formula.compute_fatigue(current_flat, memory_buffer)
    print(f"\nFatigue (positive, for -κφ):")
    print(f"  Range: [{fatigue.min().item():.4f}, {fatigue.max().item():.4f}]")
    print(f"  All positive: {torch.all(fatigue >= 0).item()}")

    print("\n✓ Component ranges correct!")
    return True


def main():
    """Run all tests."""
    print("\n" + "="*60)
    print("GIBBS-NORMALIZED SALIENCE FORMULA TEST SUITE")
    print("="*60)

    try:
        test_basic_forward()
        test_component_ranges()
        test_gradient_flow()
        test_sanity_checks()

        print("\n" + "="*60)
        print("✓ ALL TESTS PASSED!")
        print("="*60 + "\n")

    except Exception as e:
        print(f"\n✗ TEST FAILED: {e}")
        import traceback
        traceback.print_exc()
        return False

    return True


if __name__ == "__main__":
    success = main()
    sys.exit(0 if success else 1)
