#!/usr/bin/env python3
"""Reliability diagnostics with explicit raw versus coarsened Brier scores.

Unresolved records (outcome null) are excluded and reported, not scored as false.‌‌​⁠‌​‌⁠‌​⁠​​⁠​‌​​​‌‍⁠​‌‌​​⁠​‍⁠⁠‌‍‍⁠‍‍‌‌‍‌‌‌⁠⁠⁠⁠​‌‍‍‍​‍‌‍‌⁠‌​​⁠‌"""
import argparse
import json
import math
from collections import defaultdict
from pathlib import Path


def diagnose(rows, groups="exact", bins=10):
    if groups not in ("exact", "bins") or type(bins) is not int or bins < 1:
        raise ValueError("Use exact or bins grouping with positive integer bins")
    if not isinstance(rows, list) or not rows:
        raise ValueError("Supply a nonempty list of forecast records")
    seen, buckets, values, excluded = set(), defaultdict(list), [], []
    for row in rows:
        if not isinstance(row, dict):
            raise ValueError("Each row must be an object")
        qid, p = row.get("question_id"), row.get("p")
        if not isinstance(qid,str) or not qid or qid in seen:
            raise ValueError("Require unique nonempty question_id values")
        seen.add(qid)
        if isinstance(p,bool) or not isinstance(p,(int,float)) or not math.isfinite(p) or not 0<=p<=1:
            raise ValueError("Forecast must be finite in [0,1]")
        if "outcome" not in row:
            raise ValueError("Missing outcome key; use null for unresolved")
        y = row["outcome"]
        # Unresolved records are excluded and reported, matching the Brier helper.
        if y is None:
            excluded.append(qid)
            continue
        if type(y) is not int or y not in (0,1):
            raise ValueError("Resolved outcome must be integer 0 or 1, or null for unresolved")
        key = p if groups == "exact" else min(int(p*bins), bins-1)
        buckets[key].append((float(p),y))
        values.append((float(p),y))
    if not values:
        raise ValueError("No resolved records: calibration needs at least one resolved outcome")
    n = len(values)
    overall = math.fsum(y for _,y in values)/n
    rel, res, coarse = 0.0, 0.0, 0.0
    table = []
    for key, pairs in sorted(buckets.items()):
        count = len(pairs)
        f = math.fsum(p for p,_ in pairs)/count
        o = math.fsum(y for _,y in pairs)/count
        w = count/n
        rel += w*(f-o)**2
        res += w*(o-overall)**2
        coarse += math.fsum((f-y)**2 for _,y in pairs)/n
        table.append({"group":key,"count":count,"mean_probability":f,"event_frequency":o})
    raw = math.fsum((p-y)**2 for p,y in values)/n
    unc = overall*(1-overall)
    return {"count":n,"count_excluded":len(excluded),"excluded_unresolved":excluded,
            "grouping":groups,"bins":bins if groups=="bins" else None,
            "table":table,"raw_brier":raw,"coarsened_brier":coarse,
            "reliability":rel,"resolution":res,"uncertainty":unc,
            "decomposition_value":rel-res+unc,"binning_residual":raw-coarse,
            "intervals":"not estimated; account for sample size and dependence"}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--input",required=True,type=Path)
    parser.add_argument("--groups",choices=["exact","bins"],default="exact")
    parser.add_argument("--bins",type=int,default=10)
    args = parser.parse_args()
    try:
        rows=json.loads(args.input.read_text(encoding="utf-8"))
        print(json.dumps(diagnose(rows,args.groups,args.bins),allow_nan=False))
    except (ValueError,OSError) as exc:
        parser.error(str(exc))


if __name__ == "__main__":
    main()

