"""Same Cs-137 fallout, four deposition dates: how the season shapes dose.

Runs the valley parameter set for identical deposition inputs on four calendar
dates and writes ``seasonality.png`` next to this script. Requires matplotlib,
which is not an ecosys dependency. See ``docs/examples/seasonality.md``.
"""
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
DATES = {
    "1 Feb": date(2001, 2, 1),    # winter: no growth, cows in the stall
    "1 May": date(2001, 5, 1),    # spring: first grazing, young grass
    "15 Jul": date(2001, 7, 15),  # summer: full canopy, just before cereal harvest
    "15 Oct": date(2001, 10, 15),  # autumn: cereals harvested, maize harvest starting
}
HORIZON_YEARS = 3
ADULT = 1  # cohort index (0 = 1-year-old, 1 = 30-year-old)

population = PopulationCohorts(
    initial_age=np.array([1.0, 30.0]) * u.year,
    population=np.array([1.0, 1.0]) * u.dimensionless_unscaled,
)
engine = EcosysEngine()  # valley parameter set
results = {}
for label, day in DATES.items():
    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,
    )
    results[label] = engine.run(request)

# Figure
labels = list(DATES)
COLORS = dict(zip(labels, ["#2a78d6", "#eb6834", "#1baf7a", "#eda100"]))
INK, INK2, GRID, SURFACE = "#0b0b0b", "#52514e", "#e4e3df", "#fcfcfb"
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, 8.6))
gs = fig.add_gridspec(2, 2, height_ratios=[1, 1], hspace=0.42, wspace=0.3)
ax_milk = fig.add_subplot(gs[0, :])
ax_cum = fig.add_subplot(gs[1, 0])
ax_bar = fig.add_subplot(gs[1, 1])


def calendar(label):
    return np.array([DATES[label] + timedelta(days=float(x)) for x in results[label].times.to_value(u.day)])


# (a) Cow milk on the calendar, grazing seasons shaded.
for year in range(2001, 2004):
    ax_milk.axvspan(date(year, 4, 20), date(year, 11, 1), color="#eef4e8", lw=0, zorder=0)
for year in (2001, 2002, 2003):
    ax_milk.text(date(year, 7, 27), 1.3e4, "pasture", color="#5d7a4a", fontsize=8.5, ha="center")
for year in (2002, 2003):
    ax_milk.text(date(year, 1, 27), 1.3e4, "stall feeding", color=INK2, fontsize=8.5, ha="center")
label_offset = {"1 Feb": (0, 7), "1 May": (-4, 7), "15 Jul": (0, 16), "15 Oct": (4, 7)}
for label in labels:
    products = results[label].animal_products
    milk = products.concentration[:, products.product_ids.index("cow_milk")].to_value(u.Bq / u.kg)
    t = calendar(label)
    ok = milk > 0
    ax_milk.plot(t[ok], milk[ok], color=COLORS[label], zorder=3)
    peak = int(np.argmax(milk))
    ax_milk.annotate(label, (t[peak], milk[peak]), xytext=label_offset[label],
                     textcoords="offset points", ha="center", color=INK, fontsize=8.5)
ax_milk.set_yscale("log")
ax_milk.set_ylim(0.1, 3e4)
ax_milk.set_xlim(date(2001, 1, 15), date(2003, 9, 30))
ax_milk.xaxis.set_major_locator(mdates.MonthLocator(bymonth=(1, 4, 7, 10)))
ax_milk.xaxis.set_major_formatter(mdates.DateFormatter("%b\n%Y"))
ax_milk.set_ylabel("Cs-137 in cow milk (Bq/kg)")
ax_milk.grid(axis="y", color=GRID, lw=0.8)
ax_milk.set_title("(a) Milk follows the cow's diet, not the calendar of the fallout")

# (b) Cumulative adult dose vs time since deposition.
for label in labels:
    t_years = results[label].times.to_value(u.year)
    cum = results[label].per_capita.cumulative[:, ADULT].to_value(DOSE_UNIT) * 1e3
    ax_cum.plot(t_years, cum, color=COLORS[label])
    ax_cum.annotate(f"{label}  {cum[-1]:.2f} mSv", (t_years[-1], cum[-1]),
                    xytext=(5, {"15 Oct": 5, "1 May": -5}.get(label, 0)), textcoords="offset points",
                    va="center", color=INK, fontsize=8.5)
ax_cum.set_xlim(0, 3.75)
ax_cum.set_xticks([0, 0.5, 1, 1.5, 2, 2.5, 3])
ax_cum.set_ylim(0, None)
ax_cum.set_xlabel("Years after deposition")
ax_cum.set_ylabel("Cumulative adult dose (mSv)")
ax_cum.grid(axis="y", color=GRID, lw=0.8)
ax_cum.set_title("(b) Same fallout, 13x spread in dose")

# (c) Three-year adult dose by route, grouped by deposition date.
single = results[labels[0]].event_results[0]
foods = list(single.ingestion.food_ids)
pathways = list(single.pathways.pathway_ids)
GROUPS = {
    "Dairy": ["drinking_milk", "butter", "cream", "condensed_milk", "rennet_cheese",
              "acid_set_cheese", "goat_milk", "sheep_milk"],
    "Beef, veal,\nlamb": ["cow_beef", "fattened_cattle_beef", "veal", "lamb_meat", "venison"],
    "Pork,\npoultry": ["pork", "chicken_meat", "eggs"],
    "Leafy\nveg.": ["leafy_vegetables"],
    "Fruit,\npotatoes,\nfield veg.": ["orchard_fruit", "berries", "fruiting_vegetables",
                                    "potatoes", "root_vegetables"],
    "Cereals,\nbeer": [f for f in foods if any(s in f for s in ("wheat", "rye", "oats", "beer"))],
}
assert sorted(sum(GROUPS.values(), [])) == sorted(foods)
names = list(GROUPS) + ["Ground\nshine"]
width = 0.2
x = np.arange(len(names))
for j, label in enumerate(labels):
    single = results[label].event_results[0]
    per_food = single.ingestion.cumulative.per_food[-1, ADULT].to_value(DOSE_UNIT)
    values = [sum(per_food[foods.index(f)] for f in members) for members in GROUPS.values()]
    ground = single.pathways.cumulative[-1, ADULT, pathways.index("ground_external")]
    values.append(ground.to_value(DOSE_UNIT))
    ax_bar.bar(x + (j - 1.5) * width, np.array(values) * 1e3, width * 0.9,
               color=COLORS[label], label=label, zorder=3)
ax_bar.set_xticks(x, names, fontsize=8)
ax_bar.set_ylabel("Adult dose over 3 years (mSv)")
ax_bar.grid(axis="y", color=GRID, lw=0.8)
ax_bar.legend(title="Deposition date", frameon=False, fontsize=8.5, title_fontsize=8.5)
ax_bar.set_title("(c) Each season opens a different door into the diet")

fig.suptitle("Cs-137, identical fallout (14 kBq/m² on the ground) on four dates — ECOSYS valley parameter set",
             x=0.075, ha="left", fontsize=12, color=INK, y=0.975)
fig.savefig(HERE / "seasonality.png", dpi=150, bbox_inches="tight")
