"""Valley versus mountain: the same Cs-137 fallout on every week of the year.

Sweeps the deposition date weekly through 2001 for the valley and mountain
parameter sets and writes ``seasonality_mountain.png`` next to this script. It
also prints the four-date table used in ``docs/examples/seasonality.md``.
Requires matplotlib, which is not an ecosys dependency; runs in about 90 s.
"""
from datetime import date, timedelta
from pathlib import Path

import astropy.units as u
import matplotlib
import matplotlib.dates as mdates
import matplotlib.pyplot as plt
import numpy as np

from ecosys import EcosysEngine
from ecosys.domain.events import DepositionEvent
from ecosys.domain.populations import PopulationCohorts
from ecosys.domain.request import SimulationRequest
from ecosys.domain.units import DOSE_UNIT

matplotlib.use("Agg")
HERE = Path(__file__).parent
SETS = ("valley", "mountain")
HORIZON_YEARS = 3
ADULT = 1  # cohort index (0 = 1-year-old, 1 = 30-year-old)
EXAMPLE_DATES = {"1 Feb": date(2001, 2, 1), "1 May": date(2001, 5, 1),
                 "15 Jul": date(2001, 7, 15), "15 Oct": date(2001, 10, 15)}
SWEEP_DATES = [date(2001, 1, 1) + timedelta(days=k) for k in range(0, 365, 7)]

population = PopulationCohorts(
    initial_age=np.array([1.0, 30.0]) * u.year,
    population=np.array([1.0, 1.0]) * u.dimensionless_unscaled,
)
engines = {name: EcosysEngine(parameter_set=name) for name in SETS}


def run(parameter_set, day):
    event = DepositionEvent.from_constant_air_concentration(
        "cs_137", day,
        100 * u.Bq / u.m**3, 1 * u.day,  # air concentration drives dry deposition
        10_000 * u.Bq / u.m**2,          # wet deposition
        5 * u.mm,                        # rainfall
    )
    request = SimulationRequest.on_default_grid(
        (event,), population, output_start_date=day, horizon_years=HORIZON_YEARS,
    )
    return engines[parameter_set].run(request)


GROUPS = {
    "Dairy": ["drinking_milk", "butter", "cream", "condensed_milk", "rennet_cheese",
              "acid_set_cheese", "goat_milk", "sheep_milk"],
    "Beef, veal, lamb": ["cow_beef", "fattened_cattle_beef", "veal", "lamb_meat", "venison"],
    "Pork, poultry": ["pork", "chicken_meat", "eggs"],
    "Leafy veg.": ["leafy_vegetables"],
    "Fruit, potatoes, field veg.": ["orchard_fruit", "berries", "fruiting_vegetables",
                                    "potatoes", "root_vegetables"],
    "Cereals, beer": ["beer", "oats", "rye_bran", "rye_flour", "rye_wholegrain",
                      "summer_wheat_bran", "summer_wheat_flour", "summer_wheat_wholegrain",
                      "winter_wheat_bran", "winter_wheat_flour", "winter_wheat_wholegrain"],
}
OTHER = "Ground shine, inhalation"


def adult_breakdown(result):
    """Three-year adult dose in mSv by food group, plus the non-ingestion pathways."""
    single = result.event_results[0]
    foods = single.ingestion.food_ids
    assert sorted(sum(GROUPS.values(), [])) == sorted(foods)
    per_food = single.ingestion.cumulative.per_food[-1, ADULT].to_value(DOSE_UNIT) * 1e3
    pathways = single.pathways.cumulative[-1, ADULT].to_value(DOSE_UNIT) * 1e3
    other = pathways.sum() - pathways[single.pathways.pathway_ids.index("ingestion")]
    groups = {name: sum(per_food[foods.index(f)] for f in members) for name, members in GROUPS.items()}
    return {**groups, OTHER: other}


# Weekly sweep of the deposition date.
sweep = {name: [adult_breakdown(run(name, day)) for day in SWEEP_DATES] for name in SETS}

# Four example dates: table for the documentation and the 15 Jul trajectories.
print("| Deposition | Valley adult | Mountain adult | Valley 1-year-old | Mountain 1-year-old |")
examples = {}
for label, day in EXAMPLE_DATES.items():
    examples[label] = {name: run(name, day) for name in SETS}
    final = {name: r.per_capita.cumulative[-1].to_value(DOSE_UNIT) * 1e3
             for name, r in examples[label].items()}
    print(f"| {label} | {final['valley'][1]:.2f} | {final['mountain'][1]:.2f} "
          f"| {final['valley'][0]:.2f} | {final['mountain'][0]:.2f} |")

# Figure
INK, INK2, GRID, SURFACE = "#0b0b0b", "#52514e", "#e4e3df", "#fcfcfb"
SET_COLORS = {"valley": "#52514e", "mountain": "#4a3aa7"}
GROUP_COLORS = ["#2a78d6", "#eb6834", "#1baf7a", "#eda100", "#e87ba4", "#008300", "#c3c2b7"]
plt.rcParams.update({
    "font.size": 9.5, "axes.edgecolor": GRID, "axes.labelcolor": INK2,
    "xtick.color": INK2, "ytick.color": INK2, "axes.titlecolor": INK,
    "axes.titlesize": 10.5, "axes.titleweight": "bold", "axes.titlelocation": "left",
    "figure.facecolor": SURFACE, "axes.facecolor": SURFACE,
    "axes.spines.top": False, "axes.spines.right": False, "lines.linewidth": 2,
})
fig = plt.figure(figsize=(12, 9.4))
gs = fig.add_gridspec(2, 2, hspace=0.38, wspace=0.2)
ax_total = fig.add_subplot(gs[0, 0])
ax_cum = fig.add_subplot(gs[0, 1])
ax_stack = {"valley": fig.add_subplot(gs[1, 0])}
ax_stack["mountain"] = fig.add_subplot(gs[1, 1], sharey=ax_stack["valley"])


def format_months(ax):
    ax.xaxis.set_major_locator(mdates.MonthLocator())
    ax.xaxis.set_major_formatter(mdates.DateFormatter("%b"))
    ax.set_xlim(SWEEP_DATES[0], SWEEP_DATES[-1])
    ax.grid(axis="y", color=GRID, lw=0.8)


# (a) Total three-year adult dose against the deposition date.
for name in SETS:
    total = [sum(row.values()) for row in sweep[name]]
    ax_total.plot(SWEEP_DATES, total, color=SET_COLORS[name])
    peak = int(np.argmax(total))
    ax_total.annotate(f"{name}\npeak {total[peak]:.1f} mSv, {SWEEP_DATES[peak]:%d %b}",
                      (SWEEP_DATES[peak], total[peak]),
                      xytext=(-12, 0) if name == "valley" else (10, 4),
                      textcoords="offset points", color=INK, fontsize=8.5,
                      ha="right" if name == "valley" else "left", va="center" if name == "valley" else "bottom")
for label, day in EXAMPLE_DATES.items():
    ax_total.axvline(day, color=GRID, lw=1, zorder=0)
    ax_total.text(day, 0.02, f" {label}", color=INK2, fontsize=7.5, rotation=90, va="bottom",
                  transform=ax_total.get_xaxis_transform())
format_months(ax_total)
ax_total.set_ylim(0, 5)
ax_total.set_ylabel("Adult dose over 3 years (mSv)")
ax_total.set_xlabel("Deposition date")
ax_total.set_title("(a) Mountain: shorter, later and higher summer window")

# (b) Time evolution after the 15 Jul deposit: cumulative adult dose.
for name in SETS:
    result = examples["15 Jul"][name]
    years = result.times.to_value(u.year)
    cumulative = result.per_capita.cumulative[:, ADULT].to_value(DOSE_UNIT) * 1e3
    ax_cum.plot(years, cumulative, color=SET_COLORS[name])
    ax_cum.annotate(f"{name}  {cumulative[-1]:.2f} mSv", (years[-1], cumulative[-1]),
                    xytext=(5, 0), textcoords="offset points", va="center", color=INK, fontsize=8.5)
ax_cum.axvspan(1.55, 2.6, color="#f1f0ec", lw=0, zorder=0)
ax_cum.text(2.07, 0.3, "stored flour\nreaches the table", ha="center", color=INK2, fontsize=8)
ax_cum.set_xlim(0, 3.7)
ax_cum.set_xticks([0, 0.5, 1, 1.5, 2, 2.5, 3])
ax_cum.set_ylim(0, 5)
ax_cum.set_xlabel("Years after 15 Jul deposition")
ax_cum.set_ylabel("Cumulative adult dose (mSv)")
ax_cum.grid(axis="y", color=GRID, lw=0.8)
ax_cum.set_title("(b) 15 Jul: mountain crops still green, 40 % more dose")

# (c, d) Composition of the dose against the deposition date.
names = list(GROUPS) + [OTHER]
for name, tag in zip(SETS, "cd"):
    ax = ax_stack[name]
    values = np.array([[row[g] for g in names] for row in sweep[name]]).T
    ax.stackplot(SWEEP_DATES, values, colors=GROUP_COLORS, labels=names,
                 edgecolor=SURFACE, linewidth=0.8)
    format_months(ax)
    ax.set_xlabel("Deposition date")
    ax.set_title(f"({tag}) {name.capitalize()}: dose by food group")
ax_stack["valley"].set_ylabel("Adult dose over 3 years (mSv)")
ax_stack["valley"].set_ylim(0, 5)
plt.setp(ax_stack["mountain"].get_yticklabels(), visible=False)
handles, legend_labels = ax_stack["valley"].get_legend_handles_labels()
fig.legend(handles[::-1], legend_labels[::-1], loc="lower center", ncol=7, frameon=False,
           fontsize=8.5, bbox_to_anchor=(0.5, 0.0), handlelength=1.2, columnspacing=1.2)

fig.suptitle("Cs-137, identical fallout (14 kBq/m² on the ground) on each week of 2001 — valley vs mountain",
             x=0.075, ha="left", fontsize=12, color=INK, y=0.975)
fig.subplots_adjust(bottom=0.1)
fig.savefig(HERE / "seasonality_mountain.png", dpi=150, bbox_inches="tight")
