#!/usr/bin/env python3
"""Compute T^GPT v0.8.1 observable semantic loss and recovery metrics.

Input JSON schema:
{
  "tau_support": 0.02,
  "minimum_count": 2,
  "bootstrap_reps": 10000,
  "seed": 8081,
  "relevance": {"C1":1,"C2":0},
  "reference": ["C1","C1","C2","UNMAPPED"],
  "contracted": ["C1",...],
  "recovery": ["C1","C2",...]
}

Cluster identities and relevance must already be frozen before confirmation.
"""
from __future__ import annotations
import argparse
import json
import math
import random
from collections import Counter
from pathlib import Path
from typing import Dict, Iterable, List, Optional


def k_support(n: int, tau: float, minimum: int) -> int:
    return max(minimum, math.ceil(tau * n))


def counts(samples: Iterable[str]) -> Counter:
    return Counter(str(x) for x in samples if str(x) != "UNMAPPED")


def support(samples: List[str], relevance: Dict[str, int], tau: float, minimum: int, filtered: bool = True):
    c = counts(samples)
    k = k_support(len(samples), tau, minimum)
    s = set()
    for cluster, n in c.items():
        if n >= k and (not filtered or relevance.get(cluster, 0) == 1):
            s.add(cluster)
    return s, c, k


def probs(samples: List[str]) -> Dict[str, float]:
    c = counts(samples)
    denom = sum(c.values())
    if denom == 0:
        return {}
    return {cluster: n / denom for cluster, n in c.items()}


def descriptive(samples: List[str], relevance: Dict[str, int], tau: float, minimum: int):
    s, c, k = support(samples, relevance, tau, minimum, filtered=True)
    p = probs(samples)
    mass = sum(p.get(x, 0.0) for x in s)
    q = {x: p.get(x, 0.0) / mass for x in s} if mass > 0 else {}
    h = -sum(v * math.log(v) for v in q.values() if v > 0)
    mapped = sum(c.values())
    relevant_n = sum(n for cluster, n in c.items() if relevance.get(cluster, 0) == 1)
    return {
        "N_total": len(samples),
        "N_mapped": mapped,
        "UNMAPPED_rate": (len(samples) - mapped) / len(samples) if samples else None,
        "support_k": k,
        "D_sem": len(s),
        "H_sem": h,
        "P": relevant_n / mapped if mapped else None,
        "support_relevant": sorted(s),
    }


def recovery_metrics(reference: List[str], contracted: List[str], recovery: List[str],
                     relevance: Dict[str, int], tau: float, minimum: int):
    s_a, _, _ = support(reference, relevance, tau, minimum, True)
    s_b, _, _ = support(contracted, relevance, tau, minimum, True)
    s_r, _, _ = support(recovery, relevance, tau, minimum, True)
    p_a = probs(reference)
    denom = sum(p_a.get(c, 0.0) for c in s_a)
    lost = s_a - s_b
    lost_mass = sum(p_a.get(c, 0.0) for c in lost)
    es = lost_mass / denom if denom > 0 else None
    recovered = lost & s_r
    if lost_mass > 0:
        rec = sum(p_a.get(c, 0.0) for c in recovered) / lost_mass
        irrev = 1.0 - rec
    else:
        rec = None
        irrev = None
    return {
        "ES_LOSS": es,
        "REC_GAIN": rec,
        "IRREV_OBS": irrev,
        "lost_clusters": sorted(lost),
        "recovered_clusters": sorted(recovered),
    }


def percentile(values: List[float], q: float) -> Optional[float]:
    xs = sorted(x for x in values if x is not None and not math.isnan(x))
    if not xs:
        return None
    pos = (len(xs) - 1) * q
    lo = math.floor(pos)
    hi = math.ceil(pos)
    if lo == hi:
        return xs[lo]
    return xs[lo] + (xs[hi] - xs[lo]) * (pos - lo)


def resample(xs: List[str], rng: random.Random) -> List[str]:
    return [xs[rng.randrange(len(xs))] for _ in xs] if xs else []


def bootstrap(reference, contracted, recovery, relevance, tau, minimum, reps, seed):
    rng = random.Random(seed)
    bag = {"ES_LOSS": [], "REC_GAIN": [], "IRREV_OBS": []}
    for _ in range(reps):
        m = recovery_metrics(
            resample(reference, rng),
            resample(contracted, rng),
            resample(recovery, rng),
            relevance,
            tau,
            minimum,
        )
        for key in bag:
            if m[key] is not None:
                bag[key].append(float(m[key]))
    return {
        key: {
            "n_defined": len(vals),
            "ci95_percentile": [percentile(vals, 0.025), percentile(vals, 0.975)],
        }
        for key, vals in bag.items()
    }


def main():
    ap = argparse.ArgumentParser(description="Compute T^GPT v0.8.1 semantic loss/recovery metrics.")
    ap.add_argument("input", type=Path)
    ap.add_argument("-o", "--output", type=Path)
    args = ap.parse_args()
    data = json.loads(args.input.read_text(encoding="utf-8"))
    tau = float(data.get("tau_support", 0.02))
    minimum = int(data.get("minimum_count", 2))
    reps = int(data.get("bootstrap_reps", 10000))
    seed = int(data.get("seed", 8081))
    relevance = {str(k): int(v) for k, v in data["relevance"].items()}
    reference = [str(x) for x in data["reference"]]
    contracted = [str(x) for x in data["contracted"]]
    recovery = [str(x) for x in data["recovery"]]

    result = {
        "schema": "T_GPT_RECOVERY_RESULT_V0_8_1",
        "parameters": {
            "tau_support": tau,
            "minimum_count": minimum,
            "bootstrap_reps": reps,
            "seed": seed,
        },
        "descriptive": {
            "reference": descriptive(reference, relevance, tau, minimum),
            "contracted": descriptive(contracted, relevance, tau, minimum),
            "recovery": descriptive(recovery, relevance, tau, minimum),
        },
        "primary": recovery_metrics(reference, contracted, recovery, relevance, tau, minimum),
        "uncertainty": bootstrap(reference, contracted, recovery, relevance, tau, minimum, reps, seed),
        "interpretation_guardrail": "IRREV_OBS is relative to this probe family/model/budget and is not weight irreversibility.",
    }
    payload = json.dumps(result, ensure_ascii=False, indent=2) + "\n"
    if args.output:
        args.output.write_text(payload, encoding="utf-8")
    else:
        print(payload, end="")


if __name__ == "__main__":
    main()
