"""Turn calibrated conditional probabilities into versioned cost decisions.""" from __future__ import annotations from dataclasses import dataclass import math from typing import Iterable, Sequence MAX_CALIBRATION_BINS = 1_000 MAX_TIE_TOLERANCE = 1e-6 def _require_text(value: object, label: str) -> str: if not isinstance(value, str) or not value.strip(): raise ValueError(f"{label} must be a non-empty string") return value def _finite_number(value: object, label: str) -> float: if isinstance(value, bool) or not isinstance(value, (int, float)): raise ValueError(f"{label} must be a real number") converted = float(value) if not math.isfinite(converted): raise ValueError(f"{label} must be finite") return converted def _probability(value: object, label: str = "probability") -> float: checked = _finite_number(value, label) if checked < 0.0 or checked > 1.0: raise ValueError(f"{label} must be between 0 and 1") return checked @dataclass(frozen=True) class ActionCost: action: str cost_if_event: float cost_if_no_event: float def __post_init__(self) -> None: _require_text(self.action, "action") event_cost = _finite_number(self.cost_if_event, "cost_if_event") no_event_cost = _finite_number(self.cost_if_no_event, "cost_if_no_event") if event_cost < 0.0 or no_event_cost < 0.0: raise ValueError("decision costs must be non-negative") @classmethod def capture( cls, action: str, *, cost_if_event: float, cost_if_no_event: float ) -> "ActionCost": event_cost = _finite_number(cost_if_event, "cost_if_event") no_event_cost = _finite_number(cost_if_no_event, "cost_if_no_event") if event_cost < 0.0 or no_event_cost < 0.0: raise ValueError("decision costs must be non-negative") return cls( action=_require_text(action, "action"), cost_if_event=event_cost, cost_if_no_event=no_event_cost, ) @dataclass(frozen=True) class ProbabilityEvidence: estimate_id: str event_definition: str population: str horizon: str model_version: str calibration_version: str data_version: str evaluation_id: str evaluation_window: str probability: float def __post_init__(self) -> None: for field in ( "estimate_id", "event_definition", "population", "horizon", "model_version", "calibration_version", "data_version", "evaluation_id", "evaluation_window", ): _require_text(getattr(self, field), field) _probability(self.probability) @classmethod def capture( cls, *, estimate_id: str, event_definition: str, population: str, horizon: str, model_version: str, calibration_version: str, data_version: str, evaluation_id: str, evaluation_window: str, probability: float, ) -> "ProbabilityEvidence": return cls( estimate_id=_require_text(estimate_id, "estimate_id"), event_definition=_require_text(event_definition, "event_definition"), population=_require_text(population, "population"), horizon=_require_text(horizon, "horizon"), model_version=_require_text(model_version, "model_version"), calibration_version=_require_text( calibration_version, "calibration_version" ), data_version=_require_text(data_version, "data_version"), evaluation_id=_require_text(evaluation_id, "evaluation_id"), evaluation_window=_require_text( evaluation_window, "evaluation_window" ), probability=_probability(probability), ) @dataclass(frozen=True) class DecisionPolicy: policy_id: str policy_owner: str cost_table_version: str event_definition: str population: str horizon: str model_version: str calibration_version: str data_version: str evaluation_id: str evaluation_window: str actions: tuple[ActionCost, ...] calibration_bin_count: int minimum_calibration_count: int maximum_brier_score: float maximum_calibration_gap: float tie_tolerance: float tie_break_action: str | None def __post_init__(self) -> None: for field in ( "policy_id", "policy_owner", "cost_table_version", "event_definition", "population", "horizon", "model_version", "calibration_version", "data_version", "evaluation_id", "evaluation_window", ): _require_text(getattr(self, field), field) if not isinstance(self.actions, tuple) or len(self.actions) < 2: raise ValueError("actions must be an immutable tuple with at least two items") if not all(isinstance(action, ActionCost) for action in self.actions): raise ValueError("actions must contain ActionCost values") names = tuple(action.action for action in self.actions) if len(names) != len(set(names)): raise ValueError("action names must be unique") if ( isinstance(self.calibration_bin_count, bool) or not isinstance(self.calibration_bin_count, int) or self.calibration_bin_count < 2 or self.calibration_bin_count > MAX_CALIBRATION_BINS ): raise ValueError( "calibration_bin_count must be an integer between 2 and " f"{MAX_CALIBRATION_BINS}" ) if ( isinstance(self.minimum_calibration_count, bool) or not isinstance(self.minimum_calibration_count, int) or self.minimum_calibration_count <= 0 ): raise ValueError("minimum_calibration_count must be a positive integer") _probability(self.maximum_brier_score, "maximum_brier_score") _probability(self.maximum_calibration_gap, "maximum_calibration_gap") tolerance = _finite_number(self.tie_tolerance, "tie_tolerance") if tolerance < 0.0: raise ValueError("tie_tolerance must be non-negative") if tolerance > MAX_TIE_TOLERANCE: raise ValueError( f"tie_tolerance must not exceed {MAX_TIE_TOLERANCE}" ) if self.tie_break_action is not None: _require_text(self.tie_break_action, "tie_break_action") if self.tie_break_action not in names: raise ValueError("tie_break_action must name a declared action") @classmethod def capture( cls, *, policy_id: str, policy_owner: str, cost_table_version: str, event_definition: str, population: str, horizon: str, model_version: str, calibration_version: str, data_version: str, evaluation_id: str, evaluation_window: str, actions: Iterable[ActionCost], calibration_bin_count: int, minimum_calibration_count: int, maximum_brier_score: float, maximum_calibration_gap: float, tie_tolerance: float = 1e-12, tie_break_action: str | None = None, ) -> "DecisionPolicy": frozen_actions = tuple(actions) if len(frozen_actions) < 2: raise ValueError("a decision policy requires at least two actions") if not all(isinstance(action, ActionCost) for action in frozen_actions): raise ValueError("actions must contain ActionCost values") names = tuple(action.action for action in frozen_actions) if len(names) != len(set(names)): raise ValueError("action names must be unique") if ( isinstance(calibration_bin_count, bool) or not isinstance(calibration_bin_count, int) or calibration_bin_count < 2 or calibration_bin_count > MAX_CALIBRATION_BINS ): raise ValueError( "calibration_bin_count must be an integer between 2 and " f"{MAX_CALIBRATION_BINS}" ) if ( isinstance(minimum_calibration_count, bool) or not isinstance(minimum_calibration_count, int) or minimum_calibration_count <= 0 ): raise ValueError("minimum_calibration_count must be a positive integer") checked_maximum_brier = _probability( maximum_brier_score, "maximum_brier_score" ) checked_maximum_gap = _probability( maximum_calibration_gap, "maximum_calibration_gap" ) checked_tolerance = _finite_number(tie_tolerance, "tie_tolerance") if checked_tolerance < 0.0: raise ValueError("tie_tolerance must be non-negative") if checked_tolerance > MAX_TIE_TOLERANCE: raise ValueError( f"tie_tolerance must not exceed {MAX_TIE_TOLERANCE}" ) checked_tie_action = None if tie_break_action is not None: checked_tie_action = _require_text(tie_break_action, "tie_break_action") if checked_tie_action not in names: raise ValueError("tie_break_action must name a declared action") return cls( policy_id=_require_text(policy_id, "policy_id"), policy_owner=_require_text(policy_owner, "policy_owner"), cost_table_version=_require_text( cost_table_version, "cost_table_version" ), event_definition=_require_text(event_definition, "event_definition"), population=_require_text(population, "population"), horizon=_require_text(horizon, "horizon"), model_version=_require_text(model_version, "model_version"), calibration_version=_require_text( calibration_version, "calibration_version" ), data_version=_require_text(data_version, "data_version"), evaluation_id=_require_text(evaluation_id, "evaluation_id"), evaluation_window=_require_text( evaluation_window, "evaluation_window" ), actions=frozen_actions, calibration_bin_count=calibration_bin_count, minimum_calibration_count=minimum_calibration_count, maximum_brier_score=checked_maximum_brier, maximum_calibration_gap=checked_maximum_gap, tie_tolerance=checked_tolerance, tie_break_action=checked_tie_action, ) @dataclass(frozen=True) class DecisionTrace: estimate_id: str policy_id: str policy_owner: str cost_table_version: str event_definition: str population: str horizon: str model_version: str calibration_version: str data_version: str evaluation_id: str evaluation_window: str event_probability: float expected_costs: tuple[tuple[str, float], ...] calibration_bin_count: int required_calibration_bin_count: int calibration_count: int minimum_calibration_count: int brier_score: float maximum_brier_score: float calibration_gap: float maximum_calibration_gap: float tie_tolerance: float action: str def expected_cost(action: ActionCost, event_probability: float) -> float: if not isinstance(action, ActionCost): raise ValueError("action must be an ActionCost") probability = _probability(event_probability, "event_probability") terms = ( probability * action.cost_if_event, (1.0 - probability) * action.cost_if_no_event, ) if any(not math.isfinite(term) for term in terms): raise ValueError("expected cost produced a non-finite intermediate") return _finite_number(math.fsum(terms), "expected cost") def decide( policy: DecisionPolicy, evidence: ProbabilityEvidence, calibration: CalibrationReport, ) -> DecisionTrace: if not isinstance(policy, DecisionPolicy): raise ValueError("policy must be a DecisionPolicy") if not isinstance(evidence, ProbabilityEvidence): raise ValueError("evidence must be a ProbabilityEvidence record") if not isinstance(calibration, CalibrationReport): raise ValueError("calibration must be a CalibrationReport") for field in ( "event_definition", "population", "horizon", "model_version", "calibration_version", "data_version", "evaluation_id", "evaluation_window", ): if getattr(evidence, field) != getattr(policy, field): raise ValueError(f"probability evidence {field} does not match policy") if getattr(calibration.cohort, field) != getattr(evidence, field): raise ValueError( f"calibration report {field} does not match probability evidence" ) if calibration.bin_count != policy.calibration_bin_count: raise ValueError( "calibration report bin count does not match policy: " f"{calibration.bin_count} != {policy.calibration_bin_count}" ) if calibration.count < policy.minimum_calibration_count: raise ValueError( "calibration report count is below the policy minimum: " f"{calibration.count} < {policy.minimum_calibration_count}" ) if calibration.brier_score > policy.maximum_brier_score: raise ValueError( "calibration report Brier score exceeds the policy maximum: " f"{calibration.brier_score} > {policy.maximum_brier_score}" ) observed_gap = calibration.max_calibration_gap if observed_gap > policy.maximum_calibration_gap: raise ValueError( "calibration report gap exceeds the policy maximum: " f"{observed_gap} > {policy.maximum_calibration_gap}" ) probability = evidence.probability tolerance = policy.tie_tolerance costs = tuple( (action.action, expected_cost(action, probability)) for action in policy.actions ) minimum = min(cost for _, cost in costs) tied = tuple( action for action, cost in costs if abs(cost - minimum) <= tolerance ) if len(tied) > 1: if policy.tie_break_action is None: raise ValueError( "expected-cost tie requires an explicit tie_break_action" ) if policy.tie_break_action not in tied: raise ValueError("configured tie_break_action is not among minimum-cost actions") selected = policy.tie_break_action else: selected = tied[0] return DecisionTrace( estimate_id=evidence.estimate_id, policy_id=policy.policy_id, policy_owner=policy.policy_owner, cost_table_version=policy.cost_table_version, event_definition=evidence.event_definition, population=evidence.population, horizon=evidence.horizon, model_version=evidence.model_version, calibration_version=evidence.calibration_version, data_version=evidence.data_version, evaluation_id=evidence.evaluation_id, evaluation_window=evidence.evaluation_window, event_probability=probability, expected_costs=costs, calibration_bin_count=calibration.bin_count, required_calibration_bin_count=policy.calibration_bin_count, calibration_count=calibration.count, minimum_calibration_count=policy.minimum_calibration_count, brier_score=calibration.brier_score, maximum_brier_score=policy.maximum_brier_score, calibration_gap=observed_gap, maximum_calibration_gap=policy.maximum_calibration_gap, tie_tolerance=tolerance, action=selected, ) def bayes_update(prior_probability: float, likelihood_ratio: float) -> float: """Update event probability using posterior odds = prior odds × LR.""" prior = _probability(prior_probability, "prior_probability") ratio = _finite_number(likelihood_ratio, "likelihood_ratio") if ratio <= 0.0: raise ValueError("likelihood_ratio must be positive") if prior in (0.0, 1.0): return prior posterior_odds = (prior / (1.0 - prior)) * ratio if not math.isfinite(posterior_odds): return 1.0 return posterior_odds / (1.0 + posterior_odds) @dataclass(frozen=True) class CalibrationCohort: event_definition: str population: str horizon: str model_version: str calibration_version: str data_version: str evaluation_id: str evaluation_window: str def __post_init__(self) -> None: for field in ( "event_definition", "population", "horizon", "model_version", "calibration_version", "data_version", "evaluation_id", "evaluation_window", ): _require_text(getattr(self, field), field) @classmethod def capture( cls, *, event_definition: str, population: str, horizon: str, model_version: str, calibration_version: str, data_version: str, evaluation_id: str, evaluation_window: str, ) -> "CalibrationCohort": return cls( event_definition=_require_text(event_definition, "event_definition"), population=_require_text(population, "population"), horizon=_require_text(horizon, "horizon"), model_version=_require_text(model_version, "model_version"), calibration_version=_require_text( calibration_version, "calibration_version" ), data_version=_require_text(data_version, "data_version"), evaluation_id=_require_text(evaluation_id, "evaluation_id"), evaluation_window=_require_text( evaluation_window, "evaluation_window" ), ) @dataclass(frozen=True) class CalibrationObservation: observation_id: str cohort: CalibrationCohort probability: float event_occurred: bool def __post_init__(self) -> None: _require_text(self.observation_id, "observation_id") if not isinstance(self.cohort, CalibrationCohort): raise ValueError("cohort must be a CalibrationCohort") _probability(self.probability) if not isinstance(self.event_occurred, bool): raise ValueError("event_occurred must be boolean") @classmethod def capture( cls, *, observation_id: str, cohort: CalibrationCohort, probability: float, event_occurred: bool, ) -> "CalibrationObservation": if not isinstance(cohort, CalibrationCohort): raise ValueError("cohort must be a CalibrationCohort") if not isinstance(event_occurred, bool): raise ValueError("event_occurred must be boolean") return cls( observation_id=_require_text(observation_id, "observation_id"), cohort=cohort, probability=_probability(probability), event_occurred=event_occurred, ) @dataclass(frozen=True) class ReliabilityBin: lower: float upper: float count: int mean_probability: float event_rate: float def __post_init__(self) -> None: lower = _probability(self.lower, "bin lower") upper = _probability(self.upper, "bin upper") if lower >= upper: raise ValueError("calibration bin lower bound must precede upper bound") if isinstance(self.count, bool) or not isinstance(self.count, int) or self.count <= 0: raise ValueError("calibration bin count must be a positive integer") _probability(self.mean_probability, "mean_probability") _probability(self.event_rate, "event_rate") @property def gap(self) -> float: return abs(self.mean_probability - self.event_rate) @dataclass(frozen=True) class CalibrationReport: cohort: CalibrationCohort bin_count: int count: int brier_score: float bins: tuple[ReliabilityBin, ...] def __post_init__(self) -> None: if not isinstance(self.cohort, CalibrationCohort): raise ValueError("cohort must be a CalibrationCohort") if ( isinstance(self.bin_count, bool) or not isinstance(self.bin_count, int) or self.bin_count < 2 or self.bin_count > MAX_CALIBRATION_BINS ): raise ValueError( "calibration report bin_count must be an integer between 2 and " f"{MAX_CALIBRATION_BINS}" ) if isinstance(self.count, bool) or not isinstance(self.count, int) or self.count <= 0: raise ValueError("calibration report count must be a positive integer") score = _finite_number(self.brier_score, "brier_score") if score < 0.0 or score > 1.0: raise ValueError("brier_score must be between 0 and 1") if not isinstance(self.bins, tuple) or not self.bins: raise ValueError("calibration report bins must be a non-empty immutable tuple") if not all(isinstance(bin_, ReliabilityBin) for bin_ in self.bins): raise ValueError("calibration report bins must contain ReliabilityBin values") if len(self.bins) > self.bin_count: raise ValueError("occupied calibration bins exceed requested bin_count") if sum(bin_.count for bin_ in self.bins) != self.count: raise ValueError("calibration report bin counts must equal report count") @property def max_calibration_gap(self) -> float: return max(bin_.gap for bin_ in self.bins) def calibration_report( observations: Iterable[CalibrationObservation], *, bins: int = 10 ) -> CalibrationReport: if isinstance(bins, bool) or not isinstance(bins, int) or bins < 2: raise ValueError("bins must be an integer of at least 2") if bins > MAX_CALIBRATION_BINS: raise ValueError( f"bins must not exceed the workload bound of {MAX_CALIBRATION_BINS}" ) frozen = tuple(observations) if not frozen: raise ValueError("calibration observations must not be empty") if not all(isinstance(item, CalibrationObservation) for item in frozen): raise ValueError("observations must contain CalibrationObservation values") cohort = frozen[0].cohort if any(item.cohort != cohort for item in frozen): raise ValueError("calibration observations must share one cohort contract") observation_ids = tuple(item.observation_id for item in frozen) if len(observation_ids) != len(set(observation_ids)): raise ValueError("calibration observation IDs must be unique") grouped: list[list[CalibrationObservation]] = [[] for _ in range(bins)] for observation in frozen: index = min(int(observation.probability * bins), bins - 1) grouped[index].append(observation) reliability: list[ReliabilityBin] = [] for index, group in enumerate(grouped): if not group: continue count = len(group) reliability.append( ReliabilityBin( lower=index / bins, upper=(index + 1) / bins, count=count, mean_probability=math.fsum(item.probability for item in group) / count, event_rate=math.fsum( 1.0 if item.event_occurred else 0.0 for item in group ) / count, ) ) brier = math.fsum( (item.probability - (1.0 if item.event_occurred else 0.0)) ** 2 for item in frozen ) / len(frozen) return CalibrationReport( cohort=cohort, bin_count=bins, count=len(frozen), brier_score=brier, bins=tuple(reliability), ) EXAMPLE_POLICY = DecisionPolicy.capture( policy_id="refund-escalation-v1", policy_owner="refund-risk-council", cost_table_version="refund-costs-2026-08-14", event_definition="confirmed-abusive-refund", population="authenticated-consumer-refunds", horizon="30-days-after-request", model_version="refund-risk-v3", calibration_version="temperature-2026-08-14", data_version="refund-outcomes-2026-07", evaluation_id="refund-calibration-eval-2026-08-14", evaluation_window="2026-07-01/2026-07-31", actions=( ActionCost.capture( "auto-clear", cost_if_event=12.0, cost_if_no_event=0.0 ), ActionCost.capture( "manual-review", cost_if_event=2.0, cost_if_no_event=2.0 ), ), calibration_bin_count=2, minimum_calibration_count=4, maximum_brier_score=0.05, maximum_calibration_gap=0.20, tie_tolerance=1e-12, tie_break_action="manual-review", ) EXAMPLE_EVIDENCE = ProbabilityEvidence.capture( estimate_id="estimate-1042", event_definition="confirmed-abusive-refund", population="authenticated-consumer-refunds", horizon="30-days-after-request", model_version="refund-risk-v3", calibration_version="temperature-2026-08-14", data_version="refund-outcomes-2026-07", evaluation_id="refund-calibration-eval-2026-08-14", evaluation_window="2026-07-01/2026-07-31", probability=0.2, ) EXAMPLE_CALIBRATION_COHORT = CalibrationCohort.capture( event_definition="confirmed-abusive-refund", population="authenticated-consumer-refunds", horizon="30-days-after-request", model_version="refund-risk-v3", calibration_version="temperature-2026-08-14", data_version="refund-outcomes-2026-07", evaluation_id="refund-calibration-eval-2026-08-14", evaluation_window="2026-07-01/2026-07-31", ) EXAMPLE_OBSERVATIONS = tuple( CalibrationObservation.capture( observation_id=observation_id, cohort=EXAMPLE_CALIBRATION_COHORT, probability=probability, event_occurred=outcome, ) for observation_id, probability, outcome in ( ("case-1", 0.1, False), ("case-2", 0.2, False), ("case-3", 0.8, True), ("case-4", 0.9, True), ) ) if __name__ == "__main__": report = calibration_report(EXAMPLE_OBSERVATIONS, bins=2) trace = decide(EXAMPLE_POLICY, EXAMPLE_EVIDENCE, report) print(f"policy={trace.policy_id}") print(f"event_probability={trace.event_probability:.3f}") print( "expected_costs=" + ",".join(f"{name}:{cost:.3f}" for name, cost in trace.expected_costs) ) print(f"decision={trace.action}") print(f"brier={report.brier_score:.3f}") print(f"max_calibration_gap={report.max_calibration_gap:.3f}") print( f"scope=event:{trace.event_definition} population:{trace.population} " f"horizon:{trace.horizon}" ) print( f"versions=model:{trace.model_version} " f"calibration:{trace.calibration_version} " f"data:{trace.data_version} evaluation:{trace.evaluation_id} " f"costs:{trace.cost_table_version}" ) print(f"evaluation_window={trace.evaluation_window}") print( f"calibration_quality=count:{trace.calibration_count}/" f"{trace.minimum_calibration_count} " f"bins:{trace.calibration_bin_count}/" f"{trace.required_calibration_bin_count} " f"brier:{trace.brier_score:.3f}/{trace.maximum_brier_score:.3f} " f"gap:{trace.calibration_gap:.3f}/{trace.maximum_calibration_gap:.3f}" ) print(f"tie_tolerance={trace.tie_tolerance:.12f}") print(f"owner={trace.policy_owner} estimate={trace.estimate_id}")