"""Nuclide clocks and matched-ground dry/wet deposition in ECOSYS.

Install ecosys from the checkout and matplotlib, then run this file.
All outputs are written beside this script; the model data are not modified.
"""

import argparse
from datetime import date
import json
from pathlib import Path

import astropy.units as u
import matplotlib
import numpy as np

matplotlib.use("Agg")
import matplotlib.pyplot as plt

from ecosys import EcosysEngine
from ecosys.data import load_model_data
from ecosys.domain.events import DepositionEvent, IodineFractions
from ecosys.domain.landscape import Landscape
from ecosys.domain.populations import PopulationCohorts
from ecosys.domain.request import SimulationRequest
from ecosys.domain.units import DOSE_UNIT
from ecosys.reporting import serialize_run_provenance

HERE = Path(__file__).resolve().parent
ORIGIN = date(2001, 7, 15)
HORIZON = 1095
GROUND = 1000.0  # Bq/m², for each pure-route run
NUCLIDES = {"i_131": "I-131", "cs_134": "Cs-134", "cs_137": "Cs-137", "sr_90": "Sr-90"}
COLORS = dict(zip(NUCLIDES, ("#d06424", "#9650a2", "#1769a3", "#258455")))
MODES = {"dry": 0, "wet_1": 1, "wet_5": 5, "wet_20": 20}
PATHWAYS = ("cloud_inhalation", "resuspension_inhalation", "ground_external")


def air_for_ground(model, nuclide):
    """All iodine is particulate here, so all four ground velocities coincide."""
    element = nuclide.split("_")[0]
    matches = [v for v in model.deposition_velocities
               if (v.element_id, v.target_id, v.chemical_form) == (element, "ground", "particle")]
    if len(matches) != 1:
        raise ValueError("Expected one particulate ground velocity")
    return GROUND / matches[0].velocity.to_value(u.m / u.s)


def run(engine, model, nuclide, mode, step=1.0):
    mixed = mode == "mixed"
    air = air_for_ground(model, nuclide) if mode == "dry" or mixed else 0.0
    rain = 5 if mixed else MODES[mode]
    wet = 0.0 if mode == "dry" else GROUND
    event = DepositionEvent(
        nuclide_id=nuclide, date=ORIGIN,
        integrated_air_activity=air * u.Bq * u.s / u.m**3,
        wet_ground_deposition=wet * u.Bq / u.m**2,
        rainfall=rain * u.mm,
        iodine_fractions=(IodineFractions(*(x * u.one for x in (1, 0, 0)))
                          if nuclide == "i_131" else None),
    )
    days = np.arange(round(HORIZON / step) + 1) * step
    request = SimulationRequest(
        events=(event,), elapsed_times=days * u.day,
        population=PopulationCohorts(np.array([30.0]) * u.year, np.array([1.0]) * u.one),
        landscape=Landscape("arable"), output_start_date=ORIGIN,
        pathway_ids=PATHWAYS,
    )
    result = engine.run(request)
    single = result.event_results[0]
    products, plants, materials = result.animal_products, result.plants, result.materials
    data = {
        "days": days,
        "milk_Bq_kg": products.concentration[:, products.product_ids.index("cow_milk")].to_value(u.Bq / u.kg),
        "grass_Bq_kg": plants.concentration[:, plants.plant_ids.index("intensive_grass")].to_value(u.Bq / u.kg),
        "flour_Bq_kg": materials.concentration[:, materials.material_ids.index("winter_wheat_flour")].to_value(u.Bq / u.kg),
        "total_uSv": result.per_capita.cumulative[:, 0].to_value(DOSE_UNIT) * 1e6,
        "ingestion_uSv": single.ingestion.cumulative.total[:, 0].to_value(DOSE_UNIT) * 1e6,
    }
    for j, name in enumerate(single.pathways.pathway_ids):
        data[name + "_uSv"] = single.pathways.cumulative[:, 0, j].to_value(DOSE_UNIT) * 1e6
    # Scientific bookkeeping checks, not just successful plotting.
    for values in data.values():
        assert np.all(np.isfinite(values)) and np.all(values >= 0)
    assert np.all(np.diff(data["total_uSv"]) >= -1e-10)
    np.testing.assert_allclose(data["total_uSv"], data["ingestion_uSv"] +
                               sum(data[p + "_uSv"] for p in PATHWAYS), rtol=1e-11, atol=1e-10)
    np.testing.assert_allclose(result.per_capita.cumulative.value,
                               np.cumsum(result.per_capita.interval.value, axis=0), rtol=1e-11, atol=1e-15)
    dep = single.concentrations.deposition
    deposited = {target: float(dep.total_deposition[j].to_value(u.Bq / u.m**2))
                 for j, target in enumerate(dep.target_ids)}
    dry = {target: float(dep.dry_deposition[j].to_value(u.Bq / u.m**2))
           for j, target in enumerate(dep.target_ids)}
    np.testing.assert_allclose(deposited["ground"], GROUND * (2 if mixed else 1))
    if mode == "dry":
        assert np.all(dep.wet_deposition.value == 0)
    elif not mixed:
        assert np.all(dep.dry_deposition.value == 0)
        assert data["cloud_inhalation_uSv"][-1] == 0
    peak = int(np.argmax(data["milk_Bq_kg"]))
    doses = single.ingestion.cumulative.per_food[-1, 0].to_value(DOSE_UNIT) * 1e6
    np.testing.assert_allclose(np.sum(doses), data["ingestion_uSv"][-1], rtol=1e-11)
    summary = {
        "input": {"nuclide": nuclide, "date": str(ORIGIN), "step_days": step,
                  "horizon_days": HORIZON, "air_Bq_s_m3": air,
                  "wet_ground_Bq_m2": wet, "rain_mm": rain},
        "half_life_days": float(model.nuclides[nuclide].half_life.to_value(u.day)),
        "target_deposition_Bq_m2": deposited,
        "pasture_source_Bq_m2": deposited["ground"] + dry["intensive_grass"],
        "ground_shine_source_Bq_m2": deposited["ground"] + dry["turf"],
        "milk_peak_Bq_kg": float(data["milk_Bq_kg"][peak]),
        "milk_peak_day": float(days[peak]),
        "ingestion_half_of_horizon_dose_day": float(days[np.searchsorted(
            data["ingestion_uSv"], data["ingestion_uSv"][-1] / 2)]),
        "checkpoints": {
            str(day): {k: float(v[np.searchsorted(days, day)]) for k, v in data.items() if k != "days"}
            for day in (0, 7, 30, 90, 365, 730, HORIZON)
        },
        "ingestion_by_food_uSv": {single.ingestion.food_ids[j]: float(doses[j])
                                  for j in np.argsort(doses)[::-1]},
    }
    return data, summary, serialize_run_provenance(result.provenance)


def save(fig, name):
    fig.savefig(HERE / (name + ".png"), dpi=180)
    fig.savefig(HERE / (name + ".svg"))
    plt.close(fig)


def figures(runs, summaries):
    plt.rcParams.update({"font.size": 10, "axes.spines.top": False,
                         "axes.spines.right": False, "lines.linewidth": 2,
                         "axes.titleweight": "bold"})
    fig, axes = plt.subplots(2, 2, figsize=(12, 8), layout="constrained")
    for n, label in NUCLIDES.items():
        d = runs[n + "__wet_5"]
        t = d["days"]
        style = "--" if n == "cs_134" else "-"
        milk = np.where(d["milk_Bq_kg"] > 1e-8, d["milk_Bq_kg"], np.nan)
        axes[0, 0].plot(t, milk, label=label, color=COLORS[n], linestyle=style, zorder=4 if n == "cs_134" else 3)
        axes[0, 1].plot(t / 365.25, milk, label=label, color=COLORS[n], linestyle=style)
        axes[1, 0].plot(t / 365.25, d["total_uSv"], label=label, color=COLORS[n], linestyle=style)
        axes[1, 1].plot(t / 365.25, d["ingestion_uSv"] / d["ingestion_uSv"][-1],
                        label=label, color=COLORS[n], linestyle=style)
    axes[0, 0].set(xlim=(0, 60), yscale="log", ylim=(1e-3, 100),
                   title="A  Early cow-milk pulse", xlabel="Days after deposition", ylabel="Cow milk [Bq/kg]")
    axes[0, 1].set(xlim=(0, 3), yscale="log", ylim=(1e-3, 100),
                   title="B  Later milk: feed stocks and persistent uptake",
                   xlabel="Years after deposition", ylabel="Cow milk [Bq/kg]")
    axes[1, 0].set(xlim=(0, 3), title="C  Adult cumulative effective dose",
                   xlabel="Years after deposition", ylabel="Cumulative dose [µSv]")
    axes[1, 1].set(xlim=(0, 3), ylim=(0, 1.04), title="D  When does ingestion exposure accumulate?",
                   xlabel="Years after deposition", ylabel="Fraction of each nuclide's 1,095-day ingestion dose")
    for ax in axes.flat:
        ax.grid(alpha=0.2)
        ax.legend(fontsize=9)
    fig.suptitle("Different nuclides, one summer deposition\n15 July 2001 · wet-only 1,000 Bq/m² · 5 mm rain · valley adult", fontsize=14)
    save(fig, "nuclide_clocks")

    fig, axes = plt.subplots(2, 2, figsize=(12, 8), layout="constrained")
    routes = (("ingestion_uSv", "Ingestion", "#268a65"),
              ("ground_external_uSv", "Ground external", "#ca962e"),
              ("cloud_inhalation_uSv", "Cloud inhalation", "#7c59a6"),
              ("resuspension_inhalation_uSv", "Resuspension", "#666666"))
    for ax, (n, label) in zip(axes.flat, NUCLIDES.items()):
        positions = [("dry", 30), ("wet_5", 30), ("dry", HORIZON), ("wet_5", HORIZON)]
        bottom = np.zeros(4)
        for key, name, color in routes:
            values = np.array([runs[n + "__" + mode][key][day] for mode, day in positions])
            ax.bar(np.arange(4), values, bottom=bottom, color=color, label=name, width=0.7)
            bottom += values
        for j, total in enumerate(bottom):
            ax.text(j, total, f"{total:.3g}", ha="center", va="bottom", fontsize=10)
        ax.set(xticks=np.arange(4), xticklabels=["Dry\n30 d", "Wet\n30 d", "Dry\n1,095 d", "Wet\n1,095 d"],
               title=label, ylabel="Cumulative adult effective dose [µSv]", ylim=(0, max(bottom) * 1.2))
        ax.grid(axis="y", alpha=0.2)
    handles, labels = axes[0, 0].get_legend_handles_labels()
    fig.legend(handles, labels, loc="outside lower center", ncol=4, frameon=False)
    fig.suptitle("Dry versus wet: same ground target, different pathway sources\n1,000 Bq/m² on ground · wet case: 5 mm, zero air input · independent vertical scales", fontsize=14)
    save(fig, "dry_wet_budgets")

    fig, axes = plt.subplots(1, 2, figsize=(12, 4.8), layout="constrained")
    rainfall = np.array([1, 5, 20])
    # Caesium isotopes share interception; draw one element curve for clarity.
    for n in ("i_131", "cs_137", "sr_90"):
        fractions = [summaries[n + "__wet_" + str(r)]["target_deposition_Bq_m2"]["intensive_grass"] / GROUND
                     for r in rainfall]
        axes[0].plot(rainfall, fractions, "o-", color=COLORS[n],
                     label="Cs (both isotopes)" if n == "cs_137" else NUCLIDES[n])
    for n, label in NUCLIDES.items():
        values = [runs[n + "__wet_" + str(r)]["ingestion_uSv"][-1] for r in rainfall]
        axes[1].plot(rainfall, np.array(values) / values[1], marker="o", color=COLORS[n],
                     label=label, linestyle="--" if n == "cs_134" else "-",
                     zorder=4 if n == "cs_134" else 3)
    axes[0].set(title="A  Less of the prescribed deposit stays on grass",
                ylabel="Grass wet deposition / ground wet deposition")
    axes[1].set(title="B  Resulting change in cumulative ingestion",
                ylabel="1,095-day ingestion dose / dose at 5 mm")
    for ax in axes:
        ax.set(xscale="log", xticks=rainfall, xticklabels=["1", "5", "20"], xlabel="Rainfall depth [mm]")
        ax.grid(alpha=0.2)
        ax.legend(fontsize=9)
    fig.suptitle("Rainfall sensitivity at FIXED wet ground deposition (1,000 Bq/m²)\nWet-only cases: rainfall changes interception, not the prescribed ground input", fontsize=13)
    save(fig, "rainfall_sensitivity")


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--plots-only", action="store_true", help="Redraw figures from saved results")
    args = parser.parse_args()
    if args.plots_only:
        payload = json.loads((HERE / "results.json").read_text(encoding="utf-8"))
        runs = {}
        with np.load(HERE / "trajectories.npz", allow_pickle=False) as arrays:
            for name in arrays.files:
                scenario, quantity = name.rsplit("__", 1)
                runs.setdefault(scenario, {})[quantity] = arrays[name]
        figures(runs, payload["scenarios"])
        return
    model = load_model_data(parameter_set="valley", verify=True)
    engine = EcosysEngine(parameter_set="valley")
    runs, summaries, provenance, refinement = {}, {}, {}, {}
    for n in NUCLIDES:
        for mode in MODES:
            key = n + "__" + mode
            print("Running", key, flush=True)
            runs[key], summaries[key], provenance[key] = run(engine, model, n, mode)
            if mode in ("dry", "wet_5"):
                fine, fine_summary, fine_provenance = run(engine, model, n, mode, step=0.5)
                provenance[key + "__half_daily"] = fine_provenance
                differences = {}
                for quantity in ("milk_Bq_kg", "ingestion_uSv", "total_uSv"):
                    scale = float(np.max(fine[quantity]))
                    differences[quantity] = float(np.max(np.abs(runs[key][quantity] - fine[quantity][::2])) / scale * 100)
                refinement[key] = {"common_node_max_difference_percent_of_refined_peak": differences,
                                   "half_daily_peak_milk_Bq_kg": fine_summary["milk_peak_Bq_kg"],
                                   "half_daily_final_total_uSv": float(fine["total_uSv"][-1])}
    # Same-air comparison: adding wet fallout is a different experiment from
    # replacing dry fallout at fixed ground deposition. Verify additivity here.
    key = "cs_137__mixed"
    print("Running", key, flush=True)
    runs[key], summaries[key], provenance[key] = run(engine, model, "cs_137", "mixed")
    mixed_checks = {}
    for quantity in runs[key]:
        if quantity == "days":
            continue
        combined = runs["cs_137__dry"][quantity] + runs["cs_137__wet_5"][quantity]
        residual = float(np.max(np.abs(runs[key][quantity] - combined)))
        # The ingestion logarithmic-mean quadrature need not commute exactly
        # with adding sources; report that numerical residual rather than assume it.
        mixed_checks[quantity] = {"maximum_absolute_residual": residual,
                                  "percent_of_mixed_peak": residual / float(np.max(runs[key][quantity])) * 100}
        if not quantity.endswith("_uSv"):
            np.testing.assert_allclose(runs[key][quantity], combined, rtol=1e-10, atol=1e-10)
    figures(runs, summaries)
    payload = {"scenarios": summaries, "refinement": refinement, "mixed_additivity": mixed_checks}
    (HERE / "results.json").write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")
    (HERE / "provenance.json").write_text(json.dumps(provenance, indent=2) + "\n", encoding="utf-8")
    np.savez_compressed(HERE / "trajectories.npz", **{
        name + "__" + quantity: values for name, data in runs.items() for quantity, values in data.items()
    })
    for key, summary in summaries.items():
        end = summary["checkpoints"][str(HORIZON)]
        print(f"{key:18s} milk peak={summary['milk_peak_Bq_kg']:.4g} Bq/kg; "
              f"ingestion={end['ingestion_uSv']:.4g}, total={end['total_uSv']:.4g} µSv", flush=True)
    print("Refinement:", json.dumps(refinement, indent=2))
    print("Mixed-case additivity:", json.dumps(mixed_checks, indent=2))


if __name__ == "__main__":
    main()
