from __future__ import annotations

import os
from pathlib import Path
import tempfile
import unittest

import numpy as np

from chi_background_closure import (
    CANONICAL_SOURCE_HASHES,
    CanonicalSourceError,
    ChiParameters,
    ChiTrajectory,
    IncompleteRenormalizationError,
    TrajectoryError,
    audit_canonical_sources,
    analytic_canary_report,
    build_preflight_report,
    completed_quantum_sources_n,
    cosmic_derivative,
    curvature_uv_tail_n,
    flat_calibration_audit,
    frozen_renormalization_constants,
    portal_terms,
    physical_shells,
    propagate_modes_n,
    pv_moment_residuals,
    raw_pv_stress_n,
    require_production_renormalization,
    order0_uv_tail_n,
)


SOURCE_DIR = Path(os.environ["APEIRON_CANONICAL_SOURCE_DIR"])


class ChiBackgroundClosureTests(unittest.TestCase):
    def setUp(self):
        self.p = ChiParameters()

    def test_canonical_source_hashes(self):
        self.assertEqual(audit_canonical_sources(SOURCE_DIR), CANONICAL_SOURCE_HASHES)

    def test_tampered_canonical_source_fails_closed(self):
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            for name in CANONICAL_SOURCE_HASHES:
                (root / name).write_text("tampered", encoding="utf-8")
            with self.assertRaises(CanonicalSourceError):
                audit_canonical_sources(root)

    def test_portal_derivatives(self):
        sigma = 0.02
        theta = -0.69
        h = 1.0e-7
        mass2, d_sigma, d_theta = portal_terms(sigma, theta, self.p)
        numeric_sigma = (
            portal_terms(sigma + h, theta, self.p)[0]
            - portal_terms(sigma - h, theta, self.p)[0]
        ) / (2.0 * h)
        numeric_theta = (
            portal_terms(sigma, theta + h, self.p)[0]
            - portal_terms(sigma, theta - h, self.p)[0]
        ) / (2.0 * h)
        self.assertTrue(np.isfinite(mass2))
        self.assertAlmostEqual(float(d_sigma), float(numeric_sigma), delta=1.0e-15)
        self.assertAlmostEqual(float(d_theta), float(numeric_theta), delta=1.0e-15)

    def test_pv_moments_cancel_through_quadratic_order(self):
        residuals = pv_moment_residuals()
        for order in (0, 1, 2):
            self.assertEqual(residuals[order], 0.0)
        self.assertNotEqual(residuals[3], 0.0)

    def test_flat_renormalization_conditions(self):
        audit = flat_calibration_audit()
        self.assertTrue(all(audit["gates"].values()))
        self.assertEqual(audit["rho_ren_flat"], 0.0)
        self.assertEqual(audit["pressure_ren_flat"], 0.0)
        self.assertEqual(audit["chi2_ren_flat"], 0.0)

    def test_finite_constants_are_trajectory_independent(self):
        first = frozen_renormalization_constants()
        second = frozen_renormalization_constants()
        self.assertEqual(first, second)
        self.assertEqual(first["R2"], 0.0)

    def test_nonmonotonic_N_rejected(self):
        tr = ChiTrajectory(
            N=np.array([0.0, 0.1, 0.05]),
            H=np.ones(3),
            Hdot=np.zeros(3),
            sigma=np.zeros(3),
            theta=np.zeros(3),
        )
        with self.assertRaises(TrajectoryError):
            tr.validated()

    def test_nonpositive_H_rejected(self):
        tr = ChiTrajectory(
            N=np.array([0.0, 0.1]),
            H=np.array([1.0, 0.0]),
            Hdot=np.zeros(2),
            sigma=np.zeros(2),
            theta=np.zeros(2),
        )
        with self.assertRaises(TrajectoryError):
            tr.validated()

    def test_mode_core_is_finite_and_wronskian_preserving(self):
        N = np.linspace(0.0, 5.0e-4, 33)
        tr = ChiTrajectory(
            N=N,
            H=np.full_like(N, 1.0e-4),
            Hdot=np.zeros_like(N),
            sigma=np.full_like(N, 0.02),
            theta=np.full_like(N, -0.69),
        )
        result = propagate_modes_n(tr, np.array([0.0, 1.0e-4, 2.0e-4]), 1.0e-3)
        self.assertTrue(np.all(np.isfinite(result["u"])))
        self.assertTrue(np.all(np.isfinite(result["v"])))
        self.assertLess(result["wronskian_relative_error"], 2.0e-12)
        self.assertGreater(result["min_heavy_physical_omega2"], 0.0)

    def test_resolved_stress_and_order0_tail_are_finite(self):
        N = np.linspace(0.0, 2.0e-4, 17)
        tr = ChiTrajectory(
            N=N,
            H=np.full_like(N, 1.0e-4),
            Hdot=np.zeros_like(N),
            sigma=np.full_like(N, 0.02),
            theta=np.full_like(N, -0.69),
        )
        k, weights = physical_shells(self.p.Lambda, nodes=24)
        modes = propagate_modes_n(tr, k, self.p.Lambda, self.p)
        resolved = raw_pv_stress_n(tr, k, weights, modes, self.p)
        tail = order0_uv_tail_n(tr, self.p, quadrature_nodes=48)
        for value in resolved.values():
            self.assertTrue(np.all(np.isfinite(value)))
        for key in ("rho", "pressure", "chi2"):
            self.assertTrue(np.all(np.isfinite(tail[key])))
        self.assertGreater(tail["min_tail_omega2"], 0.0)

    def test_order0_tail_quadrature_converges(self):
        N = np.linspace(0.0, 1.0e-4, 5)
        tr = ChiTrajectory(
            N=N,
            H=np.full_like(N, 1.0e-4),
            Hdot=np.zeros_like(N),
            sigma=np.full_like(N, 0.02),
            theta=np.full_like(N, -0.69),
        )
        coarse = order0_uv_tail_n(tr, self.p, quadrature_nodes=48)
        fine = order0_uv_tail_n(tr, self.p, quadrature_nodes=72)
        for key in ("rho", "pressure", "chi2"):
            scale = max(float(np.max(np.abs(fine[key]))), 1.0e-30)
            self.assertLess(
                float(np.max(np.abs(coarse[key] - fine[key]))) / scale,
                2.0e-6,
            )

    def test_cosmic_derivative_on_nonuniform_N(self):
        N = np.array([0.0, 0.01, 0.025, 0.05, 0.08, 0.12, 0.17])
        H = np.full_like(N, 2.0)
        tr = ChiTrajectory(
            N=N,
            H=H,
            Hdot=np.zeros_like(N),
            sigma=np.zeros_like(N),
            theta=np.zeros_like(N),
        )
        derivative = cosmic_derivative(N**2, tr)
        np.testing.assert_allclose(derivative, 4.0 * N, atol=2.0e-15)

    def test_curvature_tail_is_finite_and_hierarchical(self):
        N = np.linspace(0.0, 0.05, 33)
        tr = ChiTrajectory(
            N=N,
            H=np.full_like(N, 1.0e-4),
            Hdot=np.zeros_like(N),
            sigma=np.full_like(N, 0.02),
            theta=np.full_like(N, -0.69),
        )
        tail = curvature_uv_tail_n(tr, self.p, nodes_per_octave=8)
        for key in ("rho", "pressure", "chi2"):
            self.assertTrue(np.all(np.isfinite(tail[key])))
        for item in tail["hierarchy"]:
            self.assertGreater(item["min_tail_omega2"], 0.0)
            self.assertLess(item["max_abs_W4_over_W0"], item["max_abs_W2_over_W0"])

    def test_completed_synthetic_sources_are_finite_but_not_production(self):
        N = np.linspace(0.0, 1.0e-4, 17)
        tr = ChiTrajectory(
            N=N,
            H=np.full_like(N, 1.0e-4),
            Hdot=np.zeros_like(N),
            sigma=np.full_like(N, 0.02),
            theta=np.full_like(N, -0.69),
        )
        result = completed_quantum_sources_n(
            tr, self.p, physical_nodes=16, curvature_nodes_per_octave=6
        )
        for key in ("rho", "pressure", "chi2", "ward_residual"):
            self.assertTrue(np.all(np.isfinite(result[key])))
        self.assertTrue(np.isfinite(result["ward_normalized"]))

    def test_analytic_canaries_pass(self):
        report = analytic_canary_report(self.p)
        self.assertTrue(report["pass"])
        self.assertTrue(all(report["gates"].values()))

    def test_production_sources_fail_closed(self):
        with self.assertRaises(IncompleteRenormalizationError):
            require_production_renormalization(None, None)

    def test_report_never_claims_production_sources(self):
        report = build_preflight_report(SOURCE_DIR)
        self.assertEqual(
            report["status"],
            "LOCAL_CHI_CLOSURE_ANALYTIC_CANARIES_PASS_PRODUCTION_TRAJECTORY_NOT_RUN",
        )
        self.assertFalse(report["N_mode_core"]["old_solver_or_physical_map_called"])
        self.assertFalse(report["production_background_sources"]["runnable"])
        self.assertTrue(
            report["production_background_sources"]["finite_counterterms_ported"]
        )
        self.assertTrue(report["production_background_sources"]["order0_uv_tail_ported"])
        self.assertTrue(
            report["production_background_sources"][
                "curvature_order2_order4_code_ported"
            ]
        )
        self.assertTrue(
            report["production_background_sources"]["analytic_canaries_pass"]
        )
        self.assertFalse(report["production_background_sources"]["rho_q_released"])


if __name__ == "__main__":
    unittest.main()
