"""Run a reproducible probability-calibration example on an open dataset."""

from __future__ import annotations

import csv
import json
import platform
from pathlib import Path

import matplotlib

matplotlib.use("Agg")
matplotlib.rcParams["svg.hashsalt"] = "baristalabs-calibration-example"
import matplotlib.pyplot as plt
import numpy as np
import scipy
import sklearn
from sklearn.calibration import CalibratedClassifierCV
from sklearn.datasets import load_breast_cancer
from sklearn.frozen import FrozenEstimator
from sklearn.metrics import brier_score_loss, log_loss, roc_auc_score
from sklearn.model_selection import train_test_split
from sklearn.naive_bayes import GaussianNB

SEED = 20260803
N_BINS = 10
AUTO_APPROVE_AT = 0.95
HOLD_BELOW = 0.05
OUTPUT_DIR = Path(__file__).resolve().parent / "evidence"


def reliability_rows(
    y_true: np.ndarray, probabilities: np.ndarray, variant: str
) -> list[dict[str, object]]:
    """Return equal-width calibration bins, including empty bins."""
    edges = np.linspace(0.0, 1.0, N_BINS + 1)
    # Put p=1.0 in the last bin instead of creating an eleventh bin.
    assignments = np.minimum(np.digitize(probabilities, edges[1:-1]), N_BINS - 1)
    rows: list[dict[str, object]] = []

    for index in range(N_BINS):
        selected = assignments == index
        count = int(np.sum(selected))
        mean_probability = float(np.mean(probabilities[selected])) if count else None
        positive_rate = float(np.mean(y_true[selected])) if count else None
        gap = (
            abs(mean_probability - positive_rate)
            if mean_probability is not None and positive_rate is not None
            else None
        )
        rows.append(
            {
                "variant": variant,
                "bin": index + 1,
                "lower": round(float(edges[index]), 2),
                "upper": round(float(edges[index + 1]), 2),
                "count": count,
                "mean_probability": mean_probability,
                "positive_rate": positive_rate,
                "absolute_gap": gap,
            }
        )

    return rows


def expected_calibration_error(rows: list[dict[str, object]]) -> float:
    total = sum(int(row["count"]) for row in rows)
    return float(
        sum(
            (int(row["count"]) / total) * float(row["absolute_gap"])
            for row in rows
            if int(row["count"]) > 0
        )
    )


def metric_record(
    y_true: np.ndarray, probabilities: np.ndarray, rows: list[dict[str, object]]
) -> dict[str, float]:
    return {
        # ROC AUC is a ranking metric. The other values use probabilities.
        "roc_auc": float(roc_auc_score(y_true, probabilities)),
        "brier_loss": float(brier_score_loss(y_true, probabilities)),
        "log_loss": float(log_loss(y_true, probabilities)),
        "ece_10_equal_width_bins": expected_calibration_error(rows),
    }


def workflow_record(
    y_true: np.ndarray, probabilities: np.ndarray
) -> dict[str, int | float]:
    auto = probabilities >= AUTO_APPROVE_AT
    hold = probabilities < HOLD_BELOW
    review = ~(auto | hold)

    return {
        "total_test_items": int(y_true.size),
        "auto_approve_items": int(np.sum(auto)),
        "manual_review_items": int(np.sum(review)),
        "hold_items": int(np.sum(hold)),
        "review_share": float(np.mean(review)),
        "false_approvals_in_auto_lane": int(np.sum(auto & (y_true == 0))),
        "positive_items_held": int(np.sum(hold & (y_true == 1))),
        "correct_positive_auto_approvals": int(np.sum(auto & (y_true == 1))),
        "correct_negative_holds": int(np.sum(hold & (y_true == 0))),
    }


def workflow_transitions(
    uncalibrated_probability: np.ndarray, calibrated_probability: np.ndarray
) -> dict[str, int]:
    def lanes(probabilities: np.ndarray) -> np.ndarray:
        assigned = np.full(probabilities.shape, "manual_review", dtype=object)
        assigned[probabilities >= AUTO_APPROVE_AT] = "auto_approve"
        assigned[probabilities < HOLD_BELOW] = "hold"
        return assigned

    before = lanes(uncalibrated_probability)
    after = lanes(calibrated_probability)
    transitions: dict[str, int] = {}
    for before_lane in ("auto_approve", "manual_review", "hold"):
        for after_lane in ("auto_approve", "manual_review", "hold"):
            transitions[f"{before_lane}_to_{after_lane}"] = int(
                np.sum((before == before_lane) & (after == after_lane))
            )
    return transitions


def write_bins(rows: list[dict[str, object]]) -> None:
    output_path = OUTPUT_DIR / "reliability-bins.csv"
    with output_path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(
            handle, fieldnames=list(rows[0].keys()), lineterminator="\n"
        )
        writer.writeheader()
        writer.writerows(rows)


def plot_reliability(
    rows_by_variant: dict[str, list[dict[str, object]]], test_size: int
) -> None:
    colors = {"Uncalibrated": "#8B5A2B", "Sigmoid calibrated": "#236A62"}
    figure, (curve_axis, count_axis) = plt.subplots(
        2,
        1,
        figsize=(10, 8),
        sharex=True,
        gridspec_kw={"height_ratios": [2.1, 1]},
        constrained_layout=True,
    )

    curve_axis.plot([0, 1], [0, 1], color="#666666", linestyle="--", label="Ideal")
    centers = np.linspace(0.05, 0.95, N_BINS)
    bar_width = 0.038

    for position, (variant, rows) in enumerate(rows_by_variant.items()):
        nonempty = [row for row in rows if int(row["count"]) > 0]
        x_values = [float(row["mean_probability"]) for row in nonempty]
        y_values = [float(row["positive_rate"]) for row in nonempty]
        curve_axis.plot(
            x_values,
            y_values,
            marker="o",
            linewidth=2,
            color=colors[variant],
            label=variant,
        )

        offset = (position - 0.5) * bar_width
        bin_counts = [int(row["count"]) for row in rows]
        bars = count_axis.bar(
            centers + offset,
            bin_counts,
            width=bar_width,
            color=colors[variant],
            alpha=0.85,
            label=variant,
        )
        count_axis.bar_label(
            bars,
            labels=[str(value) if value else "" for value in bin_counts],
            fontsize=7,
        )

    for axis in (curve_axis, count_axis):
        axis.axvline(HOLD_BELOW, color="#6B7280", linestyle=":", linewidth=1.5)
        axis.axvline(AUTO_APPROVE_AT, color="#6B7280", linestyle=":", linewidth=1.5)
        axis.grid(axis="y", color="#D1D5DB", alpha=0.6)

    curve_axis.set_title(
        f"Reliability on the untouched test split (n={test_size})\n"
        "The lower panel shows the number of samples in each equal-width bin"
    )
    curve_axis.set_ylabel("Observed positive rate")
    curve_axis.set_ylim(-0.03, 1.08)
    curve_axis.legend(loc="upper center", ncol=3)
    curve_axis.text(
        HOLD_BELOW,
        0.03,
        "hold boundary",
        rotation=90,
        va="bottom",
        ha="right",
        fontsize=8,
    )
    curve_axis.text(
        AUTO_APPROVE_AT,
        0.03,
        "auto boundary",
        rotation=90,
        va="bottom",
        ha="right",
        fontsize=8,
    )

    count_axis.set_xlabel("Predicted probability of the positive class")
    count_axis.set_ylabel("Bin count")
    count_axis.set_xlim(-0.02, 1.02)
    count_axis.set_xticks(np.linspace(0, 1, 11))
    count_axis.legend(loc="upper left")

    for extension in ("png", "svg"):
        metadata: dict[str, str | None] = {"Creator": "BaristaLabs calibration example"}
        if extension == "svg":
            # Matplotlib otherwise writes the current timestamp into the SVG.
            metadata["Date"] = None
        output_path = (
            OUTPUT_DIR / f"calibration-reliability-with-bin-counts.{extension}"
        )
        figure.savefig(
            output_path,
            dpi=180,
            metadata=metadata,
        )
        if extension == "svg":
            # Matplotlib's multiline path data otherwise fails git's whitespace check.
            svg = output_path.read_text(encoding="utf-8")
            output_path.write_text(
                "\n".join(line.rstrip() for line in svg.splitlines()) + "\n",
                encoding="utf-8",
            )
    plt.close(figure)


def rounded(value: object) -> object:
    if isinstance(value, float):
        return round(value, 6)
    if isinstance(value, dict):
        return {key: rounded(item) for key, item in value.items()}
    if isinstance(value, list):
        return [rounded(item) for item in value]
    return value


def main() -> None:
    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)

    dataset = load_breast_cancer()
    features = np.asarray(dataset.data)
    labels = np.asarray(dataset.target)

    training_features, remainder_features, training_labels, remainder_labels = (
        train_test_split(
            features,
            labels,
            test_size=0.40,
            random_state=SEED,
            stratify=labels,
        )
    )
    calibration_features, test_features, calibration_labels, test_labels = (
        train_test_split(
            remainder_features,
            remainder_labels,
            test_size=0.50,
            random_state=SEED,
            stratify=remainder_labels,
        )
    )

    base_model = GaussianNB()
    base_model.fit(training_features, training_labels)

    calibrator = CalibratedClassifierCV(
        estimator=FrozenEstimator(base_model),
        method="sigmoid",
    )
    calibrator.fit(calibration_features, calibration_labels)

    uncalibrated_probability = base_model.predict_proba(test_features)[:, 1]
    calibrated_probability = calibrator.predict_proba(test_features)[:, 1]

    rows_by_variant = {
        "Uncalibrated": reliability_rows(
            test_labels, uncalibrated_probability, "Uncalibrated"
        ),
        "Sigmoid calibrated": reliability_rows(
            test_labels, calibrated_probability, "Sigmoid calibrated"
        ),
    }
    all_rows = [row for rows in rows_by_variant.values() for row in rows]

    results = {
        "run": {
            "random_seed": SEED,
            "python": platform.python_version(),
            "numpy": np.__version__,
            "scipy": scipy.__version__,
            "matplotlib": matplotlib.__version__,
            "scikit_learn": sklearn.__version__,
        },
        "dataset": {
            "name": "Breast Cancer Wisconsin (Diagnostic)",
            "source": "UCI ML repository copy bundled with scikit-learn",
            "samples": int(features.shape[0]),
            "features": int(features.shape[1]),
            "target_names_in_sklearn_order": [
                str(name) for name in dataset.target_names
            ],
            "positive_class": "benign (target=1)",
            "negative_class": "malignant (target=0)",
            "label_provenance": (
                "Each row is a digitized fine-needle-aspirate image record; "
                "the diagnosis field supplies the malignant or benign label."
            ),
        },
        "split": {
            "method": "two stratified random splits with random_state=20260803",
            "training": int(training_labels.size),
            "calibration": int(calibration_labels.size),
            "test": int(test_labels.size),
            "training_positive": int(np.sum(training_labels == 1)),
            "calibration_positive": int(np.sum(calibration_labels == 1)),
            "test_positive": int(np.sum(test_labels == 1)),
        },
        "model": {
            "base": "GaussianNB()",
            "calibration": "CalibratedClassifierCV(FrozenEstimator(base), method='sigmoid')",
            "fit_boundary": (
                "The base classifier fits only the training split; the sigmoid calibrator "
                "fits only the calibration split; all reported metrics use the untouched test split."
            ),
        },
        "metrics": {
            "uncalibrated": metric_record(
                test_labels,
                uncalibrated_probability,
                rows_by_variant["Uncalibrated"],
            ),
            "sigmoid_calibrated": metric_record(
                test_labels,
                calibrated_probability,
                rows_by_variant["Sigmoid calibrated"],
            ),
        },
        "workflow_thresholds": {
            "constructed_policy": (
                f"auto-approve positive-class classification at p>={AUTO_APPROVE_AT:.2f}; "
                f"manual review at {HOLD_BELOW:.2f}<=p<{AUTO_APPROVE_AT:.2f}; "
                f"hold at p<{HOLD_BELOW:.2f}"
            ),
            "uncalibrated": workflow_record(test_labels, uncalibrated_probability),
            "sigmoid_calibrated": workflow_record(test_labels, calibrated_probability),
            "transitions_uncalibrated_to_calibrated": workflow_transitions(
                uncalibrated_probability, calibrated_probability
            ),
        },
        "reliability_bins": all_rows,
    }

    write_bins(all_rows)
    plot_reliability(rows_by_variant, int(test_labels.size))

    rounded_results = rounded(results)
    results_json = json.dumps(rounded_results, indent=2)
    results_json = results_json.replace(
        '    "target_names_in_sklearn_order": [\n'
        '      "malignant",\n'
        '      "benign"\n'
        "    ]",
        '    "target_names_in_sklearn_order": ["malignant", "benign"]',
    )
    (OUTPUT_DIR / "results.json").write_text(results_json + "\n", encoding="utf-8")
    print(results_json)


if __name__ == "__main__":
    main()
