"""Tests for fce.frenet_serret -- Quantum Frenet-Serret apparatus."""

import numpy as np
import pytest
from scipy import linalg

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

from fce.frenet_serret import QuantumFrenetSerret
from fce.lindblad import unitary_evolve
from tests.conftest import random_density_matrix, random_hermitian


@pytest.fixture
def fs():
    return QuantumFrenetSerret()


class TestEnergyMoment:
    def test_second_moment_is_variance(self, fs, qubit_H):
        """<(Delta_H)^2> = <H^2> - <H>^2."""
        rho = np.diag([0.7, 0.3]).astype(complex)
        mu2 = fs.energy_moment(rho, qubit_H, 2)

        # Manual: <H> = 0*0.7 + 1*0.3 = 0.3
        # <H^2> = 0*0.7 + 1*0.3 = 0.3
        # Var = 0.3 - 0.09 = 0.21
        expected = 0.21
        assert abs(mu2 - expected) < 1e-10

    def test_eigenstate_has_zero_variance(self, fs, qubit_H):
        """Eigenstate of H: all moments beyond 1st should factor trivially."""
        psi = np.array([0.0, 1.0], dtype=complex)
        rho = np.outer(psi, psi.conj())
        mu2 = fs.energy_moment(rho, qubit_H, 2)
        assert abs(mu2) < 1e-10

    def test_moments_are_real(self, fs, rng):
        """Central moments should be real for Hermitian H."""
        H = random_hermitian(4, rng)
        rho = random_density_matrix(4, rng)
        for order in [2, 3, 4]:
            mu = fs.energy_moment(rho, H, order)
            assert abs(np.imag(mu)) < 1e-10, f"Moment {order} has imaginary part"


class TestEvolutionSpeed:
    def test_eigenstate_zero_speed(self, fs, qubit_H):
        """Eigenstate of H should have zero evolution speed (stationary)."""
        psi = np.array([1.0, 0.0], dtype=complex)
        rho = np.outer(psi, psi.conj())
        v = fs.evolution_speed(rho, qubit_H)
        assert v < 1e-10

    def test_superposition_nonzero_speed(self, fs, qubit_H):
        """Superposition should evolve with nonzero speed."""
        psi = np.array([1.0, 1.0], dtype=complex) / np.sqrt(2)
        rho = np.outer(psi, psi.conj())
        v = fs.evolution_speed(rho, qubit_H)
        assert v > 0.1  # Should be Delta_E = 0.5 -> v = 0.5

    def test_speed_nonnegative(self, fs, rng):
        """Speed is always non-negative."""
        H = random_hermitian(4, rng)
        rho = random_density_matrix(4, rng)
        assert fs.evolution_speed(rho, H) >= 0.0


class TestCurvature:
    def test_eigenstate_zero_curvature(self, fs, qubit_H):
        """Eigenstate of H: geodesic evolution, kappa = 0."""
        psi = np.array([1.0, 0.0], dtype=complex)
        rho = np.outer(psi, psi.conj())
        kappa = fs.compute_curvature(rho, qubit_H)
        assert kappa < 1e-10, \
            f"Eigenstate should have zero curvature, got {kappa}"

    def test_equal_superposition_qubit_is_geodesic(self, fs, qubit_H):
        """For a 2-level system, equal superposition traces a great circle
        on the Bloch sphere, which IS a geodesic (kappa = 0).
        The energy distribution is Bernoulli(0.5) with kurtosis alpha_4 = 1."""
        psi = np.array([1.0, 1.0], dtype=complex) / np.sqrt(2)
        rho = np.outer(psi, psi.conj())
        kappa = fs.compute_curvature(rho, qubit_H)
        assert kappa < 0.01, \
            f"Equal qubit superposition should be geodesic, got kappa={kappa}"

    def test_unequal_superposition_nonzero_curvature(self, fs):
        """Unequal superposition in a 3-level system should have nonzero curvature.
        With 3+ energy levels and unequal weights, alpha_4 > 1."""
        H = np.diag([0.0, 1.0, 3.0]).astype(complex)
        # Unequal superposition across 3 levels
        psi = np.array([0.8, 0.5, 0.3317], dtype=complex)
        psi = psi / np.linalg.norm(psi)
        rho = np.outer(psi, psi.conj())
        kappa = fs.compute_curvature(rho, H)
        assert kappa > 0.01, \
            f"3-level unequal superposition should have nonzero curvature, got {kappa}"

    def test_curvature_nonnegative(self, fs, rng):
        """Curvature is always >= 0."""
        for _ in range(10):
            H = random_hermitian(4, rng)
            rho = random_density_matrix(4, rng)
            assert fs.compute_curvature(rho, H) >= 0.0

    def test_curvature_is_real(self, fs, rng):
        """Curvature must be a real number."""
        H = random_hermitian(8, rng)
        rho = random_density_matrix(8, rng)
        kappa = fs.compute_curvature(rho, H)
        assert np.isreal(kappa)

    def test_maximally_mixed_has_zero_curvature(self, fs, qubit_H):
        """Maximally mixed state I/d has zero speed -> zero curvature."""
        rho = np.eye(2, dtype=complex) / 2
        # For I/d under any H: <Delta_H^2> = (1/d)*Tr(H^2) - ((1/d)*Tr(H))^2
        # For H = diag(0,1): <H>=0.5, <H^2>=0.5, Var=0.25 > 0
        # But kurtosis: <Delta_H^4> / Var^2
        # This one should have nonzero curvature actually
        kappa = fs.compute_curvature(rho, qubit_H)
        assert isinstance(kappa, float)


class TestTorsion:
    def test_torsion_is_real(self, fs, rng):
        """Torsion must be a real number."""
        H = random_hermitian(4, rng)
        rho = random_density_matrix(4, rng)
        tau = fs.compute_torsion(rho, H)
        assert np.isreal(tau)

    def test_eigenstate_zero_torsion(self, fs, qubit_H):
        """Eigenstate: zero speed -> zero torsion."""
        psi = np.array([1.0, 0.0], dtype=complex)
        rho = np.outer(psi, psi.conj())
        tau = fs.compute_torsion(rho, qubit_H)
        assert abs(tau) < 1e-10

    def test_qubit_symmetric_superposition_zero_skewness(self, fs, qubit_H):
        """For equal superposition of a 2-level system, alpha_3 should be 0
        (symmetric energy distribution)."""
        psi = np.array([1.0, 1.0], dtype=complex) / np.sqrt(2)
        rho = np.outer(psi, psi.conj())
        mu3 = fs.energy_moment(rho, qubit_H, 3)
        # For equal weights: <(H-<H>)^3> = 0.5*(-0.5)^3 + 0.5*(0.5)^3 = 0
        assert abs(mu3) < 1e-10


class TestTangent:
    def test_tangent_shape(self, fs, qubit_H):
        """Tangent should have same shape as density matrix."""
        psi = np.array([1.0, 1.0], dtype=complex) / np.sqrt(2)
        rho = np.outer(psi, psi.conj())
        T = fs.compute_tangent(rho, qubit_H)
        assert T.shape == rho.shape

    def test_eigenstate_zero_tangent(self, fs, qubit_H):
        """Eigenstate has zero speed -> zero tangent."""
        psi = np.array([1.0, 0.0], dtype=complex)
        rho = np.outer(psi, psi.conj())
        T = fs.compute_tangent(rho, qubit_H)
        assert np.allclose(T, 0.0, atol=1e-10)

    def test_tangent_is_normalized(self, fs, rng):
        """Tangent should have unit Hilbert-Schmidt norm (if nonzero)."""
        H = random_hermitian(4, rng)
        rho = random_density_matrix(4, rng)
        T = fs.compute_tangent(rho, H)
        hs_norm = np.sqrt(np.real(np.trace(T.conj().T @ T)))
        if hs_norm > 1e-10:
            assert abs(hs_norm - 1.0) < 1e-8


class TestComputeFrame:
    def test_frame_length_matches_trajectory(self, fs, qubit_H, pure_plus):
        """One frame per density matrix."""
        dt = 0.01
        rhos = [unitary_evolve(pure_plus, qubit_H, i * dt) for i in range(20)]
        frames = fs.compute_frame(rhos, qubit_H, dt)
        assert len(frames) == 20

    def test_arc_length_increases(self, fs, qubit_H, pure_plus):
        """Arc length should be monotonically non-decreasing."""
        dt = 0.01
        rhos = [unitary_evolve(pure_plus, qubit_H, i * dt) for i in range(50)]
        frames = fs.compute_frame(rhos, qubit_H, dt)
        arc_lengths = [f.arc_length for f in frames]
        for i in range(1, len(arc_lengths)):
            assert arc_lengths[i] >= arc_lengths[i - 1] - 1e-15

    def test_curvature_constant_for_stationary_H(self, fs):
        """Under time-independent H, curvature should be approximately constant.
        Use a 3-level system with unequal weights so curvature is nonzero."""
        H = np.diag([0.0, 1.0, 3.0]).astype(complex)
        psi = np.array([0.8, 0.5, 0.3317], dtype=complex)
        psi = psi / np.linalg.norm(psi)
        rho0 = np.outer(psi, psi.conj())

        dt = 0.01
        rhos = [unitary_evolve(rho0, H, i * dt) for i in range(50)]
        frames = fs.compute_frame(rhos, H, dt)
        curvatures = [f.curvature for f in frames]
        mean_k = np.mean(curvatures)
        # Should all be the same (up to numerical precision)
        assert mean_k > 0.01, f"Expected nonzero curvature, got {mean_k}"
        assert np.std(curvatures) < 0.01 * mean_k + 1e-6


class TestCurvatureFromTrajectory:
    def test_numerical_curvature_nonnegative(self, fs, qubit_H, pure_plus):
        """Numerically computed curvature should be non-negative."""
        dt = 0.01
        rhos = [unitary_evolve(pure_plus, qubit_H, i * dt) for i in range(30)]
        kappas = fs.curvature_from_trajectory(rhos, dt)
        assert np.all(kappas >= -1e-10)

    def test_stationary_state_near_zero(self, fs, qubit_H):
        """Eigenstate trajectory: all points identical -> curvature ~ 0."""
        psi = np.array([1.0, 0.0], dtype=complex)
        rho = np.outer(psi, psi.conj())
        # Under H = diag(0,1), |0> just picks up phase -> rho unchanged
        rhos = [unitary_evolve(rho, qubit_H, i * 0.01) for i in range(20)]
        kappas = fs.curvature_from_trajectory(rhos, 0.01)
        assert np.all(np.abs(kappas) < 1e-5)
