from __future__ import annotations

import os
from pathlib import Path
import unittest

import numpy as np

from ap1_r2c_implicit_dae_solver import ImplicitDAEResidual, build_residual_report


STATE = Path(os.environ["APEIRON_AP1_STATE_NPZ"])


class ImplicitDAEResidualTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.system = ImplicitDAEResidual(STATE)

    def test_pack_unpack_roundtrip(self):
        X, Q, N_seed = self.system.unpack(self.system.initial)
        np.testing.assert_allclose(self.system.pack(X, Q, N_seed), self.system.initial)

    def test_full_residual_is_square_and_finite(self):
        residual = self.system.evaluate(self.system.initial)
        self.assertEqual(len(residual), len(self.system.initial))
        self.assertEqual(len(residual), 120)
        self.assertTrue(np.all(np.isfinite(residual)))

    def test_assembly_report_preserves_claim_boundary(self):
        report = build_residual_report(STATE)
        self.assertEqual(report["classification"], "FULL_DAE_RESIDUAL_ASSEMBLED_SOLVE_NOT_RUN")
        self.assertIn("no DAE solution", report["claim_boundary"])
        self.assertFalse(report["old_solver_or_physical_map_called"])


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