#!/usr/bin/env python3
"""Build the Evidence Lab chart from the elementary-primary numeric artifact.

This is a chart renderer, not another statistical model. Run the joint-model
generator first if its CSV needs rebuilding, then run this file directly.
IN_CSV and OUT_PNG specify locations (the publisher adapts them for site use).
The saved PNG replaces any previous file at OUT_PNG.

Only elementary-primary rows excluding charter, virtual, and special-enrollment
schools are plotted. Each bar is extra IN-SAMPLE R-squared when its predictor
is added after the other; overlapping predictors mean the bars are not parts
of a total and should not be summed."""

from __future__ import annotations

import csv
import os
import tempfile
from pathlib import Path


LOCAL_CACHE_DIR = Path(tempfile.gettempdir()) / "orschool_evidence_lab_mplcache"
LOCAL_CACHE_DIR.mkdir(parents=True, exist_ok=True)
os.environ.setdefault("MPLCONFIGDIR", str(LOCAL_CACHE_DIR))
os.environ.setdefault("XDG_CACHE_HOME", str(LOCAL_CACHE_DIR))

import matplotlib.pyplot as plt
import numpy as np


ROOT = Path(__file__).resolve().parents[3]
RUN_DATE = "2026-05-19"
IN_CSV = ROOT / "evidence-lab" / "artifacts" / "reports" / f"oregon_ba_school_poverty_joint_model_{RUN_DATE}.csv"
OUT_PNG = ROOT / "evidence-lab" / "artifacts" / "reports" / f"oregon_ba_school_poverty_two_factor_delta_r2_{RUN_DATE}.png"


def load_primary_rows() -> list[dict[str, str]]:
    """Select the declared primary model rows, leaving sensitivity rows unplotted."""
    with IN_CSV.open(newline="", encoding="utf-8") as handle:
        rows = list(csv.DictReader(handle))
    return [
        row
        for row in rows
        if row["scope"] == "elementary"
        and row["estimand_role"] == "primary_stronger_geographic_linkage"
        and row["exclusion"] == "noncharter_nonvirtual"
    ]


def main() -> None:
    rows = load_primary_rows()
    subjects = [row["subject"] for row in rows]
    ba_delta = np.array([float(row["delta_r2_ba_after_poverty"]) for row in rows])
    poverty_delta = np.array([float(row["delta_r2_poverty_after_ba"]) for row in rows])

    y = np.arange(len(subjects))
    height = 0.32

    plt.rcParams.update(
        {
            "font.family": "DejaVu Sans",
            "axes.titlesize": 18,
            "axes.labelsize": 13,
            "xtick.labelsize": 11,
            "ytick.labelsize": 13,
        }
    )

    fig, ax = plt.subplots(figsize=(11.2, 6.4))
    fig.patch.set_facecolor("#f7f4ec")
    ax.set_facecolor("#f7f4ec")

    ba_color = "#1b9e77"
    poverty_color = "#bf5a36"

    ax.barh(y - height / 2, ba_delta, height, color=ba_color, label="BA+ added after poverty")
    ax.barh(y + height / 2, poverty_delta, height, color=poverty_color, label="Poverty added after BA+")

    for values, offset in ((ba_delta, -height / 2), (poverty_delta, height / 2)):
        for yi, value in zip(y, values):
            ax.text(value + 0.004, yi + offset, f"{value:.3f}", va="center", ha="left", fontsize=11, color="#243126")

    ax.set_yticks(y, subjects)
    ax.invert_yaxis()
    ax.set_xlim(0, max(poverty_delta.max(), ba_delta.max()) + 0.06)
    ax.set_xlabel("Added explanatory power (delta R²)")
    ax.set_title("Elementary Schools: Poverty Adds More; BA+ Still Adds Some")
    ax.grid(True, axis="x", color="#d8d1c2", linewidth=0.8)
    ax.spines[["top", "right", "left"]].set_visible(False)
    ax.spines["bottom"].set_color("#b8ad99")
    ax.legend(loc="lower right", frameon=False, fontsize=12)

    fig.text(
        0.02,
        0.02,
        "Oregon 2024-25 elementary-only schools; official Total Population All Grades outcomes; scored-student-weighted models.",
        ha="left",
        fontsize=10.5,
        color="#536255",
    )

    OUT_PNG.parent.mkdir(parents=True, exist_ok=True)
    fig.tight_layout(rect=[0, 0.06, 1, 0.96])
    fig.savefig(OUT_PNG, dpi=170)
    plt.close(fig)
    print(f"Wrote {OUT_PNG}")


if __name__ == "__main__":
    main()
