"""Run the real finite-cell campaign and preserve coarse-source failures."""

import copy
import json
from pathlib import Path
import subprocess
import sys
import unittest

import audit

EXPECTED = [f"{name}/0.25: source {key} sampling exceeds allocation"
            for name in ("cell_coarse", "cell_fine", "cell_fine_half_dt", "cell_fine_outer_wide")
            for key in ("energy_flux", "source_reaction")]


class AuditTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.process = subprocess.run([sys.executable, str(Path(__file__).with_name("audit.py"))],
                                     stdout=subprocess.PIPE, text=True, timeout=7200, check=False)
        cls.report = json.loads(cls.process.stdout)
        cls.cohort = {key: value for key, value in cls.report.items() if key != "qualification_failures"}
        print("finite-cell campaign measurements: " + json.dumps(cls.report), flush=True)

    def test_real_cli_preserves_coarse_failures_and_finest_passes(self):
        self.assertEqual(self.process.returncode, 1, self.process.stderr)
        self.assertEqual(self.report["qualification_failures"], EXPECTED)
        self.assertEqual(audit.assess(self.cohort), EXPECTED)

    def test_separated_errors_refine_without_relying_on_cancellation(self):
        rows = self.cohort["rows"]
        for key, field_key in (("energy_flux", "integrated_energy_flux"),
                               ("source_reaction", "integrated_source_force")):
            reference = self.cohort["references"]["0.25"]["forward"]["65"][key]
            sampling = [abs(rows[i]["radiation"][key] - reference) for i in (0, 1, 3)]
            self.assertTrue(sampling[2] < sampling[1] < sampling[0])
            self.assertLess(sampling[2] / abs(reference), 0.001)
        # Compare the last spatial refinement at a fixed time step.
        fine, finest = rows[2], rows[3]
        self.assertEqual(fine["field"]["time_step"], finest["field"]["time_step"])
        self.assertLess(abs(finest["field"]["integrated_energy_flux"] - finest["radiation"]["energy_flux"]),
                        abs(fine["field"]["integrated_energy_flux"] - fine["radiation"]["energy_flux"]))

    def test_malformed_or_inconsistent_records_do_not_hide_behind_failures(self):
        for mutate in (
            lambda run: run["rows"].pop(),
            lambda run: run["rows"][0]["field"].update(spacing=float("nan")),
            lambda run: run["rows"][0]["field"].update(steps=True),
            lambda run: run["rows"][0]["field"].update(source_rule="sampled-gaussian"),
            lambda run: run["rows"][0]["radiation"].update(energy_flux=0.0),
            lambda run: run["references"]["0.25"]["forward"]["65"].update(source_reaction=0.0),
            lambda run: run["rows"][0]["field"].update(energy_change=0.1),
            lambda run: run["rows"][6]["field"].update(integrated_momentum_flux=1e-7, momentum_residual=1e-7),
        ):
            damaged = copy.deepcopy(self.cohort)
            mutate(damaged)
            failures = audit.assess(damaged)
            self.assertTrue(set(failures) - set(EXPECTED))

    def test_selected_configuration_controls_are_real_complete_runs(self):
        rows = {case: row for case, row in zip(audit.CASES, self.cohort["rows"])}
        self.assertEqual(len(rows), 16)
        for width in (0.25, 0.2):
            forward = rows["cell_finest", width, audit.PHASES["forward"]]["field"]
            half_dt = rows["cell_finest_half_dt", width, audit.PHASES["forward"]]["field"]
            outer = rows["cell_finest_outer_wide", width, audit.PHASES["forward"]]["field"]
            self.assertEqual(half_dt["steps"], 2 * forward["steps"])
            self.assertEqual(2 * half_dt["time_step"], forward["time_step"])
            self.assertEqual(outer["outer_half_width"], 12.0)
            self.assertGreater(outer["outer_half_width"], forward["outer_half_width"])
            for phase in audit.PHASES.values():
                row = rows["cell_finest", width, phase]
                self.assertEqual(row["field"]["spacing"], 0.05)
                self.assertLess(row["field"]["energy_change"], 1e-4)
                label = next(key for key, value in audit.PHASES.items() if value == phase)
                reference = self.cohort["references"][str(width)][label]["65"]
                self.assertLess(abs(row["radiation"]["energy_flux"] - reference["energy_flux"])
                                / reference["energy_flux"], 0.001)


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