"""Finite-cell deposition qualification with separate error allocations."""

import argparse
import json
import math
import sys

import field
import radiation
import source_sampling

STATUS = "finite-cell deposition audit only; not a primary PF-001 result"
CASES = (
    ("cell_coarse", 0.25, -math.pi / 2),
    ("cell_fine", 0.25, -math.pi / 2),
    ("cell_fine_half_dt", 0.25, -math.pi / 2),
    ("cell_finest", 0.25, -math.pi / 2),
    ("cell_finest", 0.2, -math.pi / 2),
    ("cell_fine", 0.25, math.pi / 2),
    ("cell_fine", 0.25, 0.0),
    ("cell_fine_outer_wide", 0.25, -math.pi / 2),
) + tuple(case for width in (0.25, 0.2) for case in (
    ("cell_finest", width, math.pi / 2),
    ("cell_finest", width, 0.0),
    ("cell_finest_half_dt", width, -math.pi / 2),
    ("cell_finest_outer_wide", width, -math.pi / 2),
))
PHASES = {"forward": -math.pi / 2, "reverse": math.pi / 2, "null": 0.0}
MEASUREMENTS = {"source_nodes", "energy_change", "momentum_change", "integrated_energy_flux",
                "integrated_momentum_flux", "integrated_source_work", "integrated_source_force",
                "energy_residual", "momentum_residual", "peak_abs_stored_momentum"}


def reference(width, phase, nodes):
    xy = source_sampling.axis(width, "simpson", nodes)
    axes = (xy, xy, source_sampling.axis(width, "simpson", nodes, -math.pi / 4),
            source_sampling.axis(width, "simpson", nodes, math.pi / 4))
    return radiation.integrated(axes, phase)


def evaluate():
    references = {width: {label: {str(nodes): reference(width, phase, nodes) for nodes in (33, 65)}
                          for label, phase in PHASES.items()} for width in (0.25, 0.2)}
    rows = []
    for index, (name, width, phase) in enumerate(CASES, 1):
        print(f"finite-cell case {index}/{len(CASES)}: {name}, width {width}, phase {phase}", file=sys.stderr, flush=True)
        rows.append({"field": field.simulate(name, width, phase),
                     "radiation": radiation.integrated(radiation.axes(field.CONFIGS[name][0], width), phase)})
        print(f"finite-cell case {index}/{len(CASES)} complete", file=sys.stderr, flush=True)
    return {"status": STATUS, "references": {str(key): value for key, value in references.items()}, "rows": rows}


def valid_numbers(record, keys):
    return isinstance(record, dict) and set(record) == keys and all(
        type(value) in (int, float) and math.isfinite(value) for value in record.values())


def assess(run):
    if not isinstance(run, dict) or set(run) != {"status", "references", "rows"}:
        return ["incomplete deposition audit"]
    if run["status"] != STATUS or not isinstance(run["rows"], list) or len(run["rows"]) != len(CASES):
        return ["deposition cohort drifted"]
    refs = run["references"]
    quantities = {"energy_flux", "momentum_flux", "source_reaction"}
    if not isinstance(refs, dict) or set(refs) != {"0.25", "0.2"}:
        return ["reference widths drifted"]
    failures = []
    for width, phases in refs.items():
        if not isinstance(phases, dict) or set(phases) != set(PHASES):
            return ["reference phases drifted"]
        for label, phase in PHASES.items():
            resolutions = phases[label]
            if not isinstance(resolutions, dict) or set(resolutions) != {"33", "65"}:
                return ["reference quadratures drifted"]
            for nodes, values in resolutions.items():
                if not valid_numbers(values, quantities):
                    return ["invalid radiation reference"]
                expected = reference(float(width), phase, int(nodes))
                if any(abs(values[key] - expected[key]) > 1e-12 for key in quantities):
                    return ["radiation reference drifted"]
            for key in quantities:
                difference = abs(resolutions["33"][key] - resolutions["65"][key])
                limit = 1e-8 if phase == 0 and key != "energy_flux" else abs(resolutions["65"][key]) * source_sampling.REFERENCE_LIMIT
                if difference > limit:
                    failures.append(f"width {width}/{label}: radiation reference unresolved")
    for row, (name, width, phase) in zip(run["rows"], CASES):
        if not isinstance(row, dict) or set(row) != {"field", "radiation"}:
            return ["incomplete deposition paths"]
        data = row["field"]
        h, half, cells, steps, dt = field.CONFIGS[name]
        identity = {"status": field.STATUS, "backend": "numpy-2.5.3", "configuration": name, "width": width, "phase": phase,
                    "source_rule": "clipped-cell-mass-at-node-v1", "spacing": h, "time_step": dt,
                    "steps": steps, "control_half_width": cells * h, "outer_half_width": half * h,
                    "earliest_boundary_return": 2 * half * h - cells * h - math.pi / 4 - 3 * width - h / 2}
        if not isinstance(data, dict) or set(data) != set(identity) | MEASUREMENTS:
            return ["malformed finite-cell field record"]
        for key, value in identity.items():
            if isinstance(value, float):
                valid = type(data[key]) in (int, float) and math.isfinite(data[key]) and abs(data[key] - value) <= 1e-12
            else:
                valid = type(data[key]) is type(value) and data[key] == value
            if not valid:
                return ["finite-cell field identity drifted"]
        if not valid_numbers({key: data[key] for key in MEASUREMENTS}, MEASUREMENTS):
            return ["non-finite finite-cell measurement"]
        if type(data["source_nodes"]) is not int or data["source_nodes"] != len(field.sources(h, half, width)[1]):
            return ["finite-cell source count drifted"]
        oracle = row["radiation"]
        if not valid_numbers(oracle, quantities):
            return ["invalid matched-source radiation record"]
        expected = radiation.integrated(radiation.axes(h, width), phase)
        if any(abs(oracle[key] - expected[key]) > 1e-12 for key in quantities):
            return ["matched-source radiation identity drifted"]
        label = f"{name}/{width}/{phase}"
        eb = data["energy_change"] + data["integrated_energy_flux"] - data["integrated_source_work"]
        mb = data["momentum_change"] + data["integrated_momentum_flux"] + data["integrated_source_force"]
        if abs(eb - data["energy_residual"]) > 1e-12 or abs(mb - data["momentum_residual"]) > 1e-12:
            failures.append(f"{label}: inconsistent balance")
        if abs(eb) > 0.006 or abs(mb) > 0.0005:
            failures.append(f"{label}: internal balance failed")
        if data["energy_change"] < 0 or data["integrated_energy_flux"] <= 0 or data["integrated_source_work"] <= 0 or data["peak_abs_stored_momentum"] < 0:
            failures.append(f"{label}: invalid energy or storage sign")
        if phase != 0 and data["peak_abs_stored_momentum"] < 0.01:
            failures.append(f"{label}: near-field momentum storage was not observed")
        if abs(data["energy_change"]) > 1e-4 or abs(data["momentum_change"]) > 1e-5:
            failures.append(f"{label}: terminal storage persists")
        support = cells * h + 3 * width + h / 2
        departure = 4 + math.sqrt(2 * support**2 + (support + math.pi / 4)**2)
        if not departure < steps * dt < data["earliest_boundary_return"]:
            failures.append(f"{label}: invalid observation window")
        for key, limit in (("energy", 0.004), ("momentum", 0.0022)):
            if abs(data[f"integrated_{key}_flux"] - oracle[f"{key}_flux"]) > limit:
                failures.append(f"{label}: matched-source {key} paths disagree")
        force = oracle["source_reaction"]
        if abs(data["integrated_source_force"] - force) > (1e-8 if phase == 0 else 0.005 * abs(force)):
            failures.append(f"{label}: matched-source reaction paths disagree")
        # Coarser representations remain explicit failed controls. Never use
        # cancellation with propagation error to qualify emitter sampling.
        if phase == -math.pi / 2 or h == 0.05:
            phase_label = next(label for label, value in PHASES.items() if value == phase)
            ref = refs[str(width)][phase_label]["65"]
            for key in ("energy_flux", "source_reaction"):
                limit = 1e-8 if phase == 0 and key == "source_reaction" else abs(ref[key]) * source_sampling.SAMPLING_LIMIT
                if abs(oracle[key] - ref[key]) > limit:
                    suffix = "" if phase == -math.pi / 2 else f"/{phase_label}"
                    failures.append(f"{name}/{width}{suffix}: source {key} sampling exceeds allocation")
    forward, reverse, null = run["rows"][1], run["rows"][5], run["rows"][6]
    if abs(forward["field"]["integrated_source_force"] + reverse["field"]["integrated_source_force"]) > 1e-8:
        failures.append("finite-cell phase reversal failed")
    for key in ("integrated_source_force", "integrated_momentum_flux", "momentum_change"):
        if abs(null["field"][key]) > 1e-8:
            failures.append("finite-cell null failed")
    for key in MEASUREMENTS - {"source_nodes"}:
        if abs(forward["field"][key] - run["rows"][7]["field"][key]) > 1e-7:
            failures.append("finite-cell outer-domain sensitivity failed")
    records = {case: row["field"] for case, row in zip(CASES, run["rows"])}
    for width in (0.25, 0.2):
        forward = records["cell_finest", width, -math.pi / 2]
        reverse = records["cell_finest", width, math.pi / 2]
        null = records["cell_finest", width, 0.0]
        half_dt = records["cell_finest_half_dt", width, -math.pi / 2]
        outer = records["cell_finest_outer_wide", width, -math.pi / 2]
        if abs(forward["integrated_source_force"] + reverse["integrated_source_force"]) > 1e-8:
            failures.append(f"finest/{width}: phase reversal failed")
        for key in ("integrated_source_force", "integrated_momentum_flux", "momentum_change"):
            if abs(null[key]) > 1e-8:
                failures.append(f"finest/{width}: null failed")
        for key in MEASUREMENTS - {"source_nodes"}:
            if abs(forward[key] - outer[key]) > 1e-7:
                failures.append(f"finest/{width}: outer-domain sensitivity failed")
        # Paired temporal differences use the existing flux/reaction ceilings;
        # each row must also pass the independent matched-source checks above.
        limits = {"integrated_energy_flux": 0.004, "integrated_source_work": 0.004,
                  "integrated_momentum_flux": 0.0022,
                  "integrated_source_force": 0.005 * abs(refs[str(width)]["forward"]["65"]["source_reaction"])}
        for key, limit in limits.items():
            if abs(forward[key] - half_dt[key]) > limit:
                failures.append(f"finest/{width}: time-step sensitivity failed for {key}")
    return failures


def main():
    argparse.ArgumentParser(description=__doc__).parse_args()
    run = evaluate()
    failures = assess(run)
    print(json.dumps({**run, "qualification_failures": failures}, indent=2, allow_nan=False))
    if failures:
        raise SystemExit(1)


if __name__ == "__main__":
    main()
