"""Integration tests for fce.engine and fce.navigation."""

import numpy as np
import pytest

import sys, os
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..'))

from fce.engine import FractalCorrectionEngine, EngineConfig, EvolutionResult
from fce.navigation import (
    NavigationEngine, DimensionalScaling, build_dimension_info,
    NavigationResult,
)
from fce.fidelity import validate_density_matrix, uhlmann_fidelity
from fce.lindblad import build_lindblad_operators, compute_decoherence_rates


@pytest.fixture
def qubit_engine():
    """FCE engine for a 2-level system with moderate noise."""
    H = np.diag([0.0, 1.0]).astype(complex)
    params = {'T1': 1e-3, 'T2': 5e-4, 'gate_fidelity': 0.999}
    config = EngineConfig(feedback_gain=5.0, correction_threshold=0.001)
    return FractalCorrectionEngine(
        hamiltonian=H, system_params=params, config=config,
    )


@pytest.fixture
def qubit_engine_no_qec():
    """FCE engine with QEC disabled."""
    H = np.diag([0.0, 1.0]).astype(complex)
    params = {'T1': 1e-3, 'T2': 5e-4, 'gate_fidelity': 0.999}
    config = EngineConfig(feedback_gain=0.0, qec_enabled=False)
    return FractalCorrectionEngine(
        hamiltonian=H, system_params=params, config=config,
    )


@pytest.fixture
def superposition_init():
    """Initial state: |+><+|."""
    psi = np.array([1.0, 1.0], dtype=complex) / np.sqrt(2)
    return np.outer(psi, psi.conj())


@pytest.fixture
def three_level_engine():
    """FCE engine for a 3-level system."""
    H = np.diag([0.0, 1.0, 3.0]).astype(complex)
    params = {'T1': 1e-3, 'T2': 5e-4, 'gate_fidelity': 0.999}
    config = EngineConfig(feedback_gain=5.0, correction_threshold=0.001)
    return FractalCorrectionEngine(
        hamiltonian=H, system_params=params, config=config,
    )


# === Engine Integration Tests ===

class TestEngineBasic:
    def test_evolve_returns_result(self, qubit_engine, superposition_init):
        """Engine should return an EvolutionResult."""
        result = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=20)
        assert isinstance(result, EvolutionResult)

    def test_result_array_lengths(self, qubit_engine, superposition_init):
        """All result arrays should have length n_steps + 1."""
        n_steps = 30
        result = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=n_steps)
        expected = n_steps + 1
        assert len(result.fidelities) == expected
        assert len(result.fidelities_uncorrected) == expected
        assert len(result.curvatures) == expected
        assert len(result.torsions) == expected
        assert len(result.kappa_errors) == expected
        assert len(result.tau_errors) == expected
        assert len(result.feedback_norms) == expected
        assert len(result.times) == expected
        assert len(result.density_matrices) == expected
        assert len(result.ideal_states) == expected
        assert len(result.prediction_fidelities) == expected

    def test_all_density_matrices_valid(self, qubit_engine, superposition_init):
        """Every density matrix in the trajectory should be valid."""
        result = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=30)
        for i, rho in enumerate(result.density_matrices):
            checks = validate_density_matrix(rho)
            assert checks['valid'], f"State at step {i} invalid: {checks}"

    def test_initial_fidelity_is_one(self, qubit_engine, superposition_init):
        """Fidelity at t=0 should be 1.0."""
        result = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=10)
        assert abs(result.fidelities[0] - 1.0) < 1e-6

    def test_times_monotonic(self, qubit_engine, superposition_init):
        """Times should be monotonically increasing."""
        result = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=20)
        assert np.all(np.diff(result.times) > 0)


class TestEngineQEC:
    def test_qec_improves_fidelity(self, qubit_engine, qubit_engine_no_qec,
                                    superposition_init):
        """QEC-corrected evolution should have higher fidelity than uncorrected."""
        result_on = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=50)
        result_off = qubit_engine_no_qec.evolve(superposition_init, dt=1e-6, n_steps=50)

        fid_on = result_on.fidelities[-1]
        fid_off = result_off.fidelities[-1]

        assert fid_on >= fid_off - 0.01, \
            f"QEC fidelity ({fid_on}) should be >= uncorrected ({fid_off})"

    def test_uncorrected_track_matches_no_qec(self, qubit_engine, qubit_engine_no_qec,
                                                superposition_init):
        """The uncorrected fidelity track should match running without QEC."""
        result = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=30)
        result_off = qubit_engine_no_qec.evolve(superposition_init, dt=1e-6, n_steps=30)

        # These should be similar (not identical due to independent Lindblad steps)
        # but within reasonable tolerance
        diff = abs(result.fidelities_uncorrected[-1] - result_off.fidelities[-1])
        assert diff < 0.1, \
            f"Uncorrected tracks diverged: {result.fidelities_uncorrected[-1]} vs {result_off.fidelities[-1]}"

    def test_corrections_applied(self, qubit_engine, superposition_init):
        """Some corrections should be applied during evolution."""
        result = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=50)
        # With moderate noise and low threshold, corrections should occur
        assert result.n_corrections >= 0  # At least plausible

    def test_no_hardcoded_metrics(self, qubit_engine, superposition_init):
        """Error rates should NOT be any known hardcoded value."""
        result = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=50)
        error = 1.0 - result.fidelities[-1]
        assert abs(error - 0.909) > 0.01, "Error rate should not be hardcoded 0.909"
        assert abs(error - 0.001) > 0.0005, "Error rate should not be hardcoded 0.001"

    def test_different_noise_different_results(self):
        """Different noise parameters should produce different fidelities."""
        H = np.diag([0.0, 1.0]).astype(complex)
        psi = np.array([1.0, 1.0], dtype=complex) / np.sqrt(2)
        rho = np.outer(psi, psi.conj())

        # Low noise
        engine_low = FractalCorrectionEngine(
            hamiltonian=H,
            system_params={'T1': 1.0, 'T2': 0.5, 'gate_fidelity': 0.9999},
        )
        result_low = engine_low.evolve(rho, dt=1e-6, n_steps=30)

        # High noise
        engine_high = FractalCorrectionEngine(
            hamiltonian=H,
            system_params={'T1': 1e-4, 'T2': 5e-5, 'gate_fidelity': 0.99},
        )
        result_high = engine_high.evolve(rho, dt=1e-6, n_steps=30)

        assert result_low.fidelities[-1] > result_high.fidelities[-1], \
            "Low noise should give higher fidelity than high noise"


class TestEngineHealthMetrics:
    def test_hausdorff_finite(self, qubit_engine, superposition_init):
        """Hausdorff dimension should be finite."""
        result = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=50)
        assert np.isfinite(result.hausdorff_dimension)

    def test_trajectory_health_valid(self, qubit_engine, superposition_init):
        """Trajectory health should be one of the valid categories."""
        result = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=50)
        valid_healths = {'geodesic', 'coherent', 'partially_coherent', 'decoherent'}
        assert result.trajectory_health in valid_healths, \
            f"Invalid health: {result.trajectory_health}"


class TestEngineSummary:
    def test_summary_keys(self, qubit_engine, superposition_init):
        """Summary should contain all expected keys."""
        result = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=20)
        summary = qubit_engine.summary(result)

        expected_keys = {
            'n_steps', 'total_time',
            'final_fidelity_corrected', 'final_fidelity_uncorrected',
            'error_rate_corrected', 'error_rate_uncorrected',
            'error_reduction_factor',
            'mean_curvature', 'mean_torsion',
            'n_corrections_applied',
            'hausdorff_dimension', 'trajectory_health',
            'mean_prediction_fidelity',
        }
        assert expected_keys.issubset(summary.keys()), \
            f"Missing keys: {expected_keys - summary.keys()}"

    def test_summary_values_consistent(self, qubit_engine, superposition_init):
        """Summary values should be consistent with the result arrays."""
        result = qubit_engine.evolve(superposition_init, dt=1e-6, n_steps=20)
        summary = qubit_engine.summary(result)

        assert abs(summary['final_fidelity_corrected'] - result.fidelities[-1]) < 1e-10
        assert abs(summary['error_rate_corrected'] - (1 - result.fidelities[-1])) < 1e-10
        assert summary['n_steps'] == 20


class TestEngineThreeLevel:
    def test_three_level_runs(self, three_level_engine):
        """3-level system should work end-to-end."""
        psi = np.array([0.5, 0.5, 1 / np.sqrt(2)], dtype=complex)
        psi /= np.linalg.norm(psi)
        rho = np.outer(psi, psi.conj())

        result = three_level_engine.evolve(rho, dt=1e-6, n_steps=20)
        assert isinstance(result, EvolutionResult)
        assert len(result.fidelities) == 21

    def test_three_level_nonzero_curvature(self, three_level_engine):
        """3-level system with unequal weights should have nonzero curvature."""
        psi = np.array([0.5, 0.3, np.sqrt(1 - 0.25 - 0.09)], dtype=complex)
        psi /= np.linalg.norm(psi)
        rho = np.outer(psi, psi.conj())

        result = three_level_engine.evolve(rho, dt=1e-6, n_steps=10)
        # At least some curvatures should be nonzero
        assert np.any(result.curvatures > 1e-10)


# === Navigation Integration Tests ===

class TestDimensionalScaling:
    def test_spacetime_scaling(self):
        """Spacetime dimensions should have larger scaling."""
        scaler = DimensionalScaling()
        for d in range(4):
            assert scaler.scaling_factor(d) > 0

    def test_compact_dim_suppressed(self):
        """Compact dimensions should be suppressed relative to spacetime."""
        scaler = DimensionalScaling()
        spacetime_max = max(scaler.scaling_factor(d) for d in range(4))
        for d in range(4, 11):
            assert scaler.scaling_factor(d) < spacetime_max

    def test_all_factors_finite(self):
        """All 11 scaling factors should be finite and positive."""
        factors = DimensionalScaling().all_factors()
        assert len(factors) == 11
        assert np.all(np.isfinite(factors))
        assert np.all(factors > 0)

    def test_kk_mass_increases_with_mode(self):
        """Higher KK modes should have higher mass."""
        scaler = DimensionalScaling()
        for d in range(5, 11):
            m1 = scaler.kk_mass(d, 1)
            m2 = scaler.kk_mass(d, 2)
            assert m2 > m1, f"KK mass should increase with mode number for dim {d}"


class TestDimensionInfo:
    def test_eleven_dimensions(self):
        """Should describe all 11 dimensions."""
        dims = build_dimension_info()
        assert len(dims) == 11

    def test_spacetime_types(self):
        """First 4 should be spacetime."""
        dims = build_dimension_info()
        for d in dims[:4]:
            assert d.dim_type == 'spacetime'
            assert d.compactification_radius == float('inf')
            assert d.suppression_factor == 1.0

    def test_compact_types(self):
        """Dimensions 4-10 should be compact."""
        dims = build_dimension_info()
        for d in dims[4:]:
            assert d.dim_type == 'compact'
            assert np.isfinite(d.compactification_radius)
            assert d.suppression_factor < 1.0


class TestNavigationEngine:
    def test_navigate_returns_result(self):
        """Navigation should return NavigationResult."""
        H = np.diag([0.0, 1.0]).astype(complex)
        psi = np.array([1.0, 1.0], dtype=complex) / np.sqrt(2)
        rho = np.outer(psi, psi.conj())

        nav = NavigationEngine(
            hamiltonian=H,
            system_params={'T1': 1e-3, 'T2': 5e-4, 'gate_fidelity': 0.999},
        )
        result = nav.navigate(rho, dt=1e-6, n_steps=20)
        assert isinstance(result, NavigationResult)

    def test_navigation_state_count(self):
        """Should have n_steps + 1 navigation states."""
        H = np.diag([0.0, 1.0]).astype(complex)
        psi = np.array([1.0, 0.0], dtype=complex)
        rho = np.outer(psi, psi.conj())

        nav = NavigationEngine(hamiltonian=H)
        n_steps = 15
        result = nav.navigate(rho, dt=1e-6, n_steps=n_steps)
        assert len(result.states) == n_steps + 1

    def test_coordinates_evolve(self):
        """11D coordinates should change over time (not stay at zero)."""
        H = np.diag([0.0, 1.0]).astype(complex)
        psi = np.array([1.0, 1.0], dtype=complex) / np.sqrt(2)
        rho = np.outer(psi, psi.conj())

        nav = NavigationEngine(hamiltonian=H)
        result = nav.navigate(rho, dt=1e-6, n_steps=20)

        final_coords = result.states[-1].coordinates
        assert np.any(np.abs(final_coords) > 1e-15), \
            "Coordinates should evolve from zero"

    def test_dimension_curvatures_shape(self):
        """Dimension curvatures should have shape (n_steps+1, 11)."""
        H = np.diag([0.0, 1.0]).astype(complex)
        rho = np.diag([0.5, 0.5]).astype(complex)

        nav = NavigationEngine(hamiltonian=H)
        result = nav.navigate(rho, dt=1e-6, n_steps=10)
        assert result.dimension_curvatures.shape == (11, 11)

    def test_vulnerabilities_normalized(self):
        """Dimension vulnerabilities should be in [0, 1]."""
        H = np.diag([0.0, 1.0]).astype(complex)
        psi = np.array([1.0, 1.0], dtype=complex) / np.sqrt(2)
        rho = np.outer(psi, psi.conj())

        nav = NavigationEngine(
            hamiltonian=H,
            system_params={'T1': 1e-3, 'T2': 5e-4, 'gate_fidelity': 0.999},
        )
        result = nav.navigate(rho, dt=1e-6, n_steps=20)

        assert np.all(result.dimension_vulnerabilities >= 0)
        assert np.all(result.dimension_vulnerabilities <= 1.0 + 1e-10)

    def test_summary_has_dimension_info(self):
        """Summary should include dimension vulnerabilities."""
        H = np.diag([0.0, 1.0]).astype(complex)
        rho = np.diag([0.5, 0.5]).astype(complex)

        nav = NavigationEngine(hamiltonian=H)
        result = nav.navigate(rho, dt=1e-6, n_steps=10)

        assert 'dimension_vulnerabilities' in result.summary
        assert 'compact_dim_suppression' in result.summary
        assert len(result.summary['dimension_vulnerabilities']) == 11

    def test_fidelity_tracked(self):
        """total_fidelity should match the final evolution fidelity."""
        H = np.diag([0.0, 1.0]).astype(complex)
        psi = np.array([1.0, 1.0], dtype=complex) / np.sqrt(2)
        rho = np.outer(psi, psi.conj())

        nav = NavigationEngine(hamiltonian=H)
        result = nav.navigate(rho, dt=1e-6, n_steps=10)

        assert abs(result.total_fidelity - result.evolution_result.fidelities[-1]) < 1e-10
