from __future__ import annotations

from types import SimpleNamespace
import unittest

import numpy as np

from ap1_r2c_candidate_anderson import AndersonConfig, KEYS, SpectralCandidateMap


class FakeGrid:
    def __init__(self):
        self.cfg = SimpleNamespace(candidate_nodes=33)
        self.scales = {"rho": 2.0, "pressure": 3.0, "chi2": 5.0}
        q = {key: np.full(33, self.scales[key] * (i + 1.0))
             for i, key in enumerate(KEYS)}
        self.initial = self.pack(q)

    def pack(self, q):
        return np.concatenate([q[key] / self.scales[key] for key in KEYS])

    def unpack(self, x):
        return {key: x[i * 33:(i + 1) * 33] * self.scales[key]
                for i, key in enumerate(KEYS)}


class CandidateAndersonTests(unittest.TestCase):
    def test_default_solver_preserves_frozen_pointwise_gate(self):
        cfg = AndersonConfig()
        self.assertEqual(cfg.residual_tolerance, 1.0e-4)
        self.assertEqual(cfg.candidate.candidate_nodes, 33)
        self.assertEqual(cfg.basis_degree, 8)

    def test_spectral_constant_roundtrip(self):
        model = SpectralCandidateMap(FakeGrid(), degree=8)
        recovered = model.expand(model.initial)
        expected = model.grid.unpack(model.grid.initial)
        for key in KEYS:
            np.testing.assert_allclose(recovered[key], expected[key], rtol=0.0, atol=1.0e-13)

    def test_invalid_basis_fails_closed(self):
        with self.assertRaises(ValueError):
            SpectralCandidateMap(FakeGrid(), degree=33)


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