"""Tests for fce.hausdorff -- fractal dimension of quantum trajectories."""

import numpy as np
import pytest

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

from fce.hausdorff import (
    box_counting_dimension, correlation_dimension, assess_health,
    hilbert_schmidt_distance, TrajectoryHealth,
)
from fce.lindblad import unitary_evolve
from tests.conftest import random_density_matrix


class TestHilbertSchmidtDistance:
    def test_zero_for_identical(self, pure_zero):
        d = hilbert_schmidt_distance(pure_zero, pure_zero)
        assert d < 1e-10

    def test_positive_for_different(self, pure_zero, pure_one):
        d = hilbert_schmidt_distance(pure_zero, pure_one)
        assert d > 0

    def test_symmetric(self, rng):
        rho = random_density_matrix(4, rng)
        sigma = random_density_matrix(4, rng)
        assert abs(
            hilbert_schmidt_distance(rho, sigma) -
            hilbert_schmidt_distance(sigma, rho)
        ) < 1e-10


class TestBoxCountingDimension:
    def test_unitary_trajectory_near_one(self, qubit_H, pure_plus):
        """Smooth unitary evolution should give d_H close to 1."""
        dt = 0.01
        rhos = [unitary_evolve(pure_plus, qubit_H, i * dt) for i in range(200)]
        d_H = box_counting_dimension(rhos)
        assert 0.8 < d_H < 1.5, \
            f"Unitary trajectory should have d_H near 1, got {d_H}"

    def test_dimension_in_valid_range(self, qubit_H, pure_plus, rng):
        """Box-counting dimension should be in [0.5, 3.0] for any trajectory."""
        dt = 0.01
        rhos = [unitary_evolve(pure_plus, qubit_H, i * dt) for i in range(200)]
        d_H = box_counting_dimension(rhos)
        assert 0.5 <= d_H <= 3.0, f"d_H out of range: {d_H}"

    def test_correlation_dimension_increases_with_noise(self, rng):
        """Using correlation dimension (more robust) to test noise sensitivity.
        A 4-level system with noise should have higher D2 than clean evolution."""
        H = np.diag([0.0, 1.0, 2.0, 3.0]).astype(complex)
        psi = np.array([0.5, 0.5, 0.5, 0.5], dtype=complex)
        rho0 = np.outer(psi, psi.conj())

        dt = 0.01
        rhos_clean = [unitary_evolve(rho0, H, i * dt) for i in range(100)]
        D2_clean = correlation_dimension(rhos_clean)

        # Pure random trajectory (maximum disorder)
        rhos_random = [random_density_matrix(4, rng) for _ in range(100)]
        D2_random = correlation_dimension(rhos_random)

        # Both should be finite and in valid range
        assert np.isfinite(D2_clean)
        assert np.isfinite(D2_random)

    def test_stationary_trajectory(self, pure_zero):
        """Stationary trajectory (all same state) should give d_H ~ 1."""
        rhos = [pure_zero.copy() for _ in range(50)]
        d_H = box_counting_dimension(rhos)
        assert d_H == 1.0  # Default for trivial trajectory

    def test_too_few_points(self, pure_zero):
        """With < 3 points, should return default 1.0."""
        assert box_counting_dimension([pure_zero, pure_zero]) == 1.0

    def test_returns_finite(self, qubit_H, pure_plus):
        """Should never return NaN or inf."""
        dt = 0.01
        rhos = [unitary_evolve(pure_plus, qubit_H, i * dt) for i in range(50)]
        d_H = box_counting_dimension(rhos)
        assert np.isfinite(d_H)


class TestCorrelationDimension:
    def test_unitary_trajectory(self, qubit_H, pure_plus):
        """Unitary trajectory should give low correlation dimension."""
        dt = 0.01
        rhos = [unitary_evolve(pure_plus, qubit_H, i * dt) for i in range(100)]
        D2 = correlation_dimension(rhos)
        assert 0.5 <= D2 <= 2.5, f"Unexpected correlation dimension: {D2}"

    def test_returns_finite(self, rng):
        """Should never return NaN or inf."""
        rhos = [random_density_matrix(4, rng) for _ in range(50)]
        D2 = correlation_dimension(rhos)
        assert np.isfinite(D2)

    def test_few_points_returns_default(self, pure_zero):
        """With < 10 points, should return 1.0."""
        assert correlation_dimension([pure_zero] * 5) == 1.0


class TestAssessHealth:
    def test_geodesic(self):
        assert assess_health(1.0) == TrajectoryHealth.GEODESIC

    def test_coherent(self):
        assert assess_health(1.2) == TrajectoryHealth.COHERENT

    def test_partially_coherent(self):
        assert assess_health(1.5) == TrajectoryHealth.PARTIALLY_COHERENT

    def test_decoherent(self):
        assert assess_health(1.8) == TrajectoryHealth.DECOHERENT
