from __future__ import annotations

import json
from pathlib import Path

import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from scipy import stats


ROOT = Path(__file__).resolve().parent
INPUT = ROOT / "irrigation-water-2020-2026.csv"
FIGURES = ROOT / "figures"
FIGURES.mkdir(exist_ok=True)
REGION_ORDER = ["Yilan", "Hualien", "Taitung"]
REGION_LABELS = {"Yilan": "宜蘭", "Hualien": "花蓮", "Taitung": "臺東"}
PLOT_LABELS = {"Yilan": "Yilan", "Hualien": "Hualien", "Taitung": "Taitung"}
COLORS = {"Yilan": "#315c4c", "Hualien": "#c18a3d", "Taitung": "#577a9e"}
RNG = np.random.default_rng(20260712)


def q1(x):
    return x.quantile(0.25)


def q3(x):
    return x.quantile(0.75)


def bootstrap_median_ci(values: np.ndarray, iterations: int = 5000) -> list[float]:
    values = values[np.isfinite(values)]
    estimates = np.empty(iterations)
    for i in range(iterations):
        estimates[i] = np.median(RNG.choice(values, size=len(values), replace=True))
    return [float(x) for x in np.quantile(estimates, [0.025, 0.975])]


def epsilon_squared(h: float, n: int, k: int) -> float:
    return max(0.0, (h - k + 1) / (n - k))


def savefig(name: str):
    plt.tight_layout()
    plt.savefig(FIGURES / name, dpi=220, bbox_inches="tight", facecolor="white")
    plt.close()


df = pd.read_csv(INPUT, parse_dates=["sample_date"])
df["site_id"] = df["region"] + "|" + df["station"] + "|" + df["monitoring_point"]
df["year"] = df["sample_date"].dt.year
df["month"] = df["sample_date"].dt.month

# Values outside broad physical plausibility bounds are preserved in the raw file,
# flagged here, and excluded only from the sensitivity dataset.
df["plausible"] = (
    df["water_temperature_c"].between(0, 45)
    & df["ph"].between(0, 14)
    & df["ec_us_cm_25c"].between(0, 5000)
)
df["ph_outside_standard"] = ~df["ph"].between(6.0, 9.0)
df["ec_above_standard"] = df["ec_us_cm_25c"] > 750

site = (
    df.groupby(["region", "site_id"], as_index=False)
    .agg(
        station=("station", "first"),
        monitoring_point=("monitoring_point", "first"),
        n=("sample_date", "size"),
        temperature_median=("water_temperature_c", "median"),
        ph_median=("ph", "median"),
        ec_median=("ec_us_cm_25c", "median"),
        ph_outside_rate=("ph_outside_standard", "mean"),
        ec_above_rate=("ec_above_standard", "mean"),
    )
)

summary = (
    df.groupby("region")
    .agg(
        observations=("sample_date", "size"),
        sites=("site_id", "nunique"),
        start_date=("sample_date", "min"),
        end_date=("sample_date", "max"),
        temperature_median=("water_temperature_c", "median"),
        temperature_q1=("water_temperature_c", q1),
        temperature_q3=("water_temperature_c", q3),
        ph_median=("ph", "median"),
        ph_q1=("ph", q1),
        ph_q3=("ph", q3),
        ec_median=("ec_us_cm_25c", "median"),
        ec_q1=("ec_us_cm_25c", q1),
        ec_q3=("ec_us_cm_25c", q3),
        ph_outside_rate=("ph_outside_standard", "mean"),
        ec_above_rate=("ec_above_standard", "mean"),
        implausible=("plausible", lambda x: int((~x).sum())),
    )
    .reindex(REGION_ORDER)
)

tests = {}
for field, label in [
    ("temperature_median", "water_temperature"),
    ("ph_median", "ph"),
    ("ec_median", "ec"),
]:
    groups = [site.loc[site.region == region, field].to_numpy() for region in REGION_ORDER]
    h, p = stats.kruskal(*groups)
    tests[label] = {
        "method": "Kruskal-Wallis on site medians",
        "H": float(h),
        "df": 2,
        "p": float(p),
        "epsilon_squared": float(epsilon_squared(h, len(site), 3)),
    }

site_summary = {}
for region in REGION_ORDER:
    part = site[site.region == region]
    site_summary[region] = {}
    for field in ["temperature_median", "ph_median", "ec_median"]:
        values = part[field].to_numpy()
        site_summary[region][field] = {
            "median": float(np.median(values)),
            "q1": float(np.quantile(values, 0.25)),
            "q3": float(np.quantile(values, 0.75)),
            "median_ci95": bootstrap_median_ci(values),
        }

common = df[df["year"].between(2022, 2025) & df["plausible"]]
common_site = common.groupby(["region", "site_id"], as_index=False).agg(
    temperature_median=("water_temperature_c", "median"),
    ph_median=("ph", "median"),
    ec_median=("ec_us_cm_25c", "median"),
)
sensitivity = {
    region: {
        "sites": int(len(part)),
        "temperature_median": float(part.temperature_median.median()),
        "ph_median": float(part.ph_median.median()),
        "ec_median": float(part.ec_median.median()),
    }
    for region in REGION_ORDER
    for part in [common_site[common_site.region == region]]
}

monthly = (
    df.groupby(["region", "month"], as_index=False)
    .agg(temperature_median=("water_temperature_c", "median"), ec_median=("ec_us_cm_25c", "median"))
)
annual = (
    df.groupby(["region", "year"], as_index=False)
    .agg(ph_median=("ph", "median"), ec_median=("ec_us_cm_25c", "median"), n=("sample_date", "size"))
)

# Figure 1: study logic and evidence boundary.
fig, ax = plt.subplots(figsize=(11, 3.2))
ax.axis("off")
boxes = [
    (0.02, "Irrigation water\nPublic: temperature, pH, EC\nPhase 2: ions, Si, B, As, Fe, Mn"),
    (0.36, "Paddy soil\nPhase 2: pH, OM, CEC, salinity,\navailable elements"),
    (0.70, "Rice grain\nPhase 2: protein, amylose, chalk,\nRVA, taste, aroma, As/Cd"),
]
for x, text_value in boxes:
    ax.text(x, 0.5, text_value, transform=ax.transAxes, va="center", ha="left", fontsize=10,
            bbox=dict(boxstyle="round,pad=0.7", facecolor="#f3efe5", edgecolor="#315c4c", linewidth=1.4))
for x in [0.30, 0.64]:
    ax.annotate("", xy=(x + 0.045, 0.5), xytext=(x, 0.5), xycoords="axes fraction",
                arrowprops=dict(arrowstyle="->", lw=2, color="#315c4c"))
ax.text(0.5, 0.05, "Phase 1 describes water monitoring; it does not estimate water-to-grain causality.",
        transform=ax.transAxes, ha="center", fontsize=10, color="#7a4e24")
savefig("figure-1-water-soil-rice-chain.png")

# Figures 2-3: site-level distributions (equal weight per monitoring point).
for field, ylabel, filename in [
    ("ec_median", "Site median EC (μS/cm at 25°C)", "figure-2-site-ec.png"),
    ("ph_median", "Site median pH", "figure-3-site-ph.png"),
]:
    fig, ax = plt.subplots(figsize=(8.6, 5.2))
    values = [site.loc[site.region == r, field] for r in REGION_ORDER]
    bp = ax.boxplot(values, tick_labels=[PLOT_LABELS[r] for r in REGION_ORDER], patch_artist=True,
                    showfliers=True, widths=0.55)
    for box, region in zip(bp["boxes"], REGION_ORDER):
        box.set_facecolor(COLORS[region]); box.set_alpha(0.75)
    ax.set_ylabel(ylabel)
    ax.grid(axis="y", alpha=0.2)
    if field == "ec_median":
        ax.axhline(750, color="#a33b2b", linestyle="--", linewidth=1.4, label="Irrigation standard 750")
        ax.legend(frameon=False)
    savefig(filename)

# Figure 4: seasonal temperature.
fig, ax = plt.subplots(figsize=(9.2, 5.2))
for region in REGION_ORDER:
    part = monthly[monthly.region == region]
    ax.plot(part.month, part.temperature_median, marker="o", linewidth=2, color=COLORS[region], label=PLOT_LABELS[region])
ax.set_xticks(range(1, 13))
ax.set_xlabel("Month")
ax.set_ylabel("Median water temperature (°C)")
ax.grid(alpha=0.2)
ax.legend(frameon=False, ncol=3)
savefig("figure-4-seasonal-temperature.png")

# Figure 5: annual EC medians and observation counts caveat.
fig, ax = plt.subplots(figsize=(9.2, 5.2))
for region in REGION_ORDER:
    part = annual[annual.region == region]
    ax.plot(part.year, part.ec_median, marker="o", linewidth=2, color=COLORS[region], label=PLOT_LABELS[region])
ax.axhline(750, color="#a33b2b", linestyle="--", linewidth=1.2)
ax.set_xlabel("Year")
ax.set_ylabel("Observation-level median EC (μS/cm at 25°C)")
ax.grid(alpha=0.2)
ax.legend(frameon=False, ncol=3)
savefig("figure-5-annual-ec.png")

# Figure 6: observation-level standard exceedance rates.
fig, ax = plt.subplots(figsize=(8.8, 5.2))
x = np.arange(3)
width = 0.34
ph_rates = [summary.loc[r, "ph_outside_rate"] * 100 for r in REGION_ORDER]
ec_rates = [summary.loc[r, "ec_above_rate"] * 100 for r in REGION_ORDER]
ax.bar(x - width / 2, ph_rates, width, color="#795548", label="pH outside 6.0–9.0")
ax.bar(x + width / 2, ec_rates, width, color="#577a9e", label="EC above 750")
ax.set_xticks(x, [PLOT_LABELS[r] for r in REGION_ORDER])
ax.set_ylabel("Share of observations (%)")
ax.grid(axis="y", alpha=0.2)
ax.legend(frameon=False)
savefig("figure-6-standard-exceedance.png")

summary_out = summary.reset_index().copy()
summary_out["start_date"] = summary_out["start_date"].dt.strftime("%Y-%m-%d")
summary_out["end_date"] = summary_out["end_date"].dt.strftime("%Y-%m-%d")
summary_out.to_csv(ROOT / "regional-summary.csv", index=False, encoding="utf-8-sig")
site.to_csv(ROOT / "site-level-summary.csv", index=False, encoding="utf-8-sig")
monthly.to_csv(ROOT / "monthly-summary.csv", index=False, encoding="utf-8-sig")
annual.to_csv(ROOT / "annual-summary.csv", index=False, encoding="utf-8-sig")

result = {
    "generated_on": "2026-07-12",
    "input_rows": int(len(df)),
    "exact_duplicates": int(df.duplicated().sum()),
    "total_sites": int(df.site_id.nunique()),
    "physical_plausibility_flags": int((~df.plausible).sum()),
    "regional_observation_summary": summary_out.to_dict(orient="records"),
    "site_weighted_summary": site_summary,
    "kruskal_wallis": tests,
    "sensitivity_2022_2025": sensitivity,
    "standards": {"ph": "6.0–9.0", "ec_us_cm_25c_max": 750},
    "interpretation_boundary": "Public monitoring contains temperature, pH and EC only; no water-to-soil-to-grain causal effect is estimated.",
}
(ROOT / "analysis-results.json").write_text(json.dumps(result, ensure_ascii=False, indent=2), encoding="utf-8")
print(json.dumps(result, ensure_ascii=False, indent=2))
