#!/usr/bin/env python3
"""Score adjudicated, aligned forecast snapshots. Python 3 standard library.‌‌​⁠‌​‌⁠‌​⁠​​⁠​‌​​​‌​‌‌‌‍‍⁠‌‍‌​​⁠‌‍‌‍⁠​‌​‌⁠‌⁠⁠‍‍‌‌‌​‍‌‍​‌‍⁠‍‍⁠​‍"""
import argparse
import json
import math
from pathlib import Path


def prob(x):
    if isinstance(x, bool) or not isinstance(x, (int, float)) or not math.isfinite(x) or not 0 <= x <= 1:
        raise ValueError("Probabilities must be finite numbers in [0,1]")
    return float(x)


def score(rows, mode="binary"):
    if mode not in ("binary", "categorical") or not isinstance(rows, list) or not rows:
        raise ValueError("Supply a nonempty list and binary or categorical mode")
    seen, scored, excluded = set(), [], []
    for row in rows:
        if not isinstance(row, dict):
            raise ValueError("Each row must be an object")
        qid = row.get("question_id")
        if not isinstance(qid, str) or not qid or qid in seen:
            raise ValueError("Require a unique nonempty question_id per snapshot")
        seen.add(qid)
        if "outcome" not in row:
            raise ValueError("Missing outcome key; use null for unresolved")
        y = row["outcome"]
        if mode == "binary":
            p = prob(row.get("p"))
            if y is not None and (type(y) is not int or y not in (0, 1)):
                raise ValueError("Binary outcome must be integer 0, 1 or null")
            loss = None if y is None else (p-y)**2
        else:
            raw = row.get("probabilities")
            if not isinstance(raw, list) or len(raw) < 2:
                raise ValueError("Require at least two categorical probabilities")
            ps = [prob(p) for p in raw]
            if not math.isclose(math.fsum(ps), 1.0, rel_tol=0, abs_tol=1e-10):
                raise ValueError("Categorical probabilities must sum to one")
            if y is not None and (type(y) is not int or not 0 <= y < len(ps)):
                raise ValueError("Categorical outcome must be an integer index or null")
            loss = None if y is None else math.fsum((p-int(i==y))**2 for i,p in enumerate(ps))
        if y is None:
            excluded.append(qid)
        else:
            scored.append({"question_id": qid, "brier": loss})
    return {"mode": mode, "range": [0, 1 if mode == "binary" else 2],
            "count_scored": len(scored), "count_excluded": len(excluded),
            "excluded_unresolved": excluded, "per_question": scored,
            "mean_brier": math.fsum(x["brier"] for x in scored)/len(scored) if scored else None}


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


if __name__ == "__main__":
    main()

