"""Reproducible, content-addressed temperature and nucleus decoding policy.""" from __future__ import annotations from dataclasses import asdict, dataclass import hashlib import json import math import re from typing import Any, Iterable MAX_VOCABULARY = 100_000 MAX_PROMPT_TOKENS = 65_536 MAX_NEW_TOKENS = 16_384 MAX_TOKEN_ID = 10_000_000 MAX_SEED = 2**63 - 1 MIN_BINARY64_RESOLUTION = 1e-12 MAX_BINARY64_RESOLUTION = 1e-4 _IDENTIFIER = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:/@=+-]{0,159}$") def _identifier(name: str, value: object) -> str: if type(value) is not str or not _IDENTIFIER.fullmatch(value): raise ValueError(f"{name} must be a stable identifier") return value def _owner(name: str, value: object) -> str: checked = _identifier(name, value) if not checked.startswith("team:"): raise ValueError(f"{name} must identify an accountable team") return checked def _integer(name: str, value: object, minimum: int, maximum: int) -> int: if type(value) is not int or not minimum <= value <= maximum: raise ValueError(f"{name} must be an integer in [{minimum}, {maximum}]") return value def _finite(name: str, value: object) -> float: if isinstance(value, bool) or not isinstance(value, (int, float)): raise ValueError(f"{name} must be a real number") checked = float(value) if not math.isfinite(checked): raise ValueError(f"{name} must be finite") return checked def _canonical(value: Any) -> Any: if isinstance(value, dict): return {key: _canonical(item) for key, item in sorted(value.items())} if isinstance(value, tuple): return [_canonical(item) for item in value] return value def _content_id(prefix: str, payload: dict[str, Any]) -> str: encoded = json.dumps( _canonical(payload), sort_keys=True, separators=(",", ":"), allow_nan=False ).encode("utf-8") return f"{prefix}@sha256:{hashlib.sha256(encoded).hexdigest()}" def _token_tuple( name: str, values: object, *, maximum_count: int, unique: bool = True ) -> tuple[int, ...]: if type(values) is not tuple or not 1 <= len(values) <= maximum_count: raise TypeError(f"{name} must be a bounded immutable tuple") checked = tuple(_integer(name, item, 0, MAX_TOKEN_ID) for item in values) if unique and len(set(checked)) != len(checked): raise ValueError(f"{name} must not contain duplicates") return checked @dataclass(frozen=True) class DecodingPolicy: policy_id: str policy_version: str model_id: str model_version: str tokenizer_id: str tokenizer_version: str logits_schema_version: str temperature: float top_p: float seed: int maximum_new_tokens: int stop_token_ids: tuple[int, ...] constraint_id: str constraint_version: str allowed_token_ids: tuple[int, ...] maximum_vocabulary: int minimum_numeric_resolution: float maximum_absolute_logit: float numeric_convention: str policy_owner: str def __post_init__(self) -> None: for field in ( "policy_id", "policy_version", "model_id", "model_version", "tokenizer_id", "tokenizer_version", "logits_schema_version", "constraint_id", "constraint_version", ): _identifier(field, getattr(self, field)) _owner("policy_owner", self.policy_owner) temperature = _finite("temperature", self.temperature) top_p = _finite("top_p", self.top_p) resolution = _finite( "minimum_numeric_resolution", self.minimum_numeric_resolution ) logit_limit = _finite("maximum_absolute_logit", self.maximum_absolute_logit) if not MIN_BINARY64_RESOLUTION <= resolution <= MAX_BINARY64_RESOLUTION: raise ValueError("minimum_numeric_resolution is outside the safe range") if not resolution <= temperature <= 100.0: raise ValueError("temperature is outside the safe range") if not resolution <= top_p <= 1.0: raise ValueError("top_p is outside the safe range") if not resolution <= logit_limit <= 1.0 / resolution: raise ValueError("maximum_absolute_logit is outside the safe range") _integer("seed", self.seed, 0, MAX_SEED) _integer("maximum_new_tokens", self.maximum_new_tokens, 1, MAX_NEW_TOKENS) _integer("maximum_vocabulary", self.maximum_vocabulary, 1, MAX_VOCABULARY) stops = _token_tuple( "stop_token_ids", self.stop_token_ids, maximum_count=self.maximum_vocabulary ) allowed = _token_tuple( "allowed_token_ids", self.allowed_token_ids, maximum_count=self.maximum_vocabulary, ) if not set(stops).issubset(allowed): raise ValueError("every stop token must be allowed by the constraint") if self.numeric_convention != "binary64-stable-softmax-sha256-draw-v1": raise ValueError("unsupported numeric_convention") @property def content_id(self) -> str: return _content_id("decoding-policy", asdict(self)) @dataclass(frozen=True) class DecodingRequest: request_id: str request_version: str submitted_at: str tenant_scope: str model_id: str model_version: str tokenizer_id: str tokenizer_version: str policy_content_id: str prompt_token_ids: tuple[int, ...] requested_new_tokens: int request_owner: str request_content_id: str @classmethod def capture( cls, *, policy: DecodingPolicy, request_id: str, request_version: str, submitted_at: str, tenant_scope: str, prompt_token_ids: Iterable[int], requested_new_tokens: int, request_owner: str, ) -> "DecodingRequest": if type(policy) is not DecodingPolicy: raise TypeError("policy must be a concrete DecodingPolicy") frozen_prompt = tuple(prompt_token_ids) _token_tuple( "prompt_token_ids", frozen_prompt, maximum_count=MAX_PROMPT_TOKENS, unique=False, ) _integer( "requested_new_tokens", requested_new_tokens, 1, policy.maximum_new_tokens, ) payload = { "request_id": request_id, "request_version": request_version, "submitted_at": submitted_at, "tenant_scope": tenant_scope, "model_id": policy.model_id, "model_version": policy.model_version, "tokenizer_id": policy.tokenizer_id, "tokenizer_version": policy.tokenizer_version, "policy_content_id": policy.content_id, "prompt_token_ids": frozen_prompt, "requested_new_tokens": requested_new_tokens, "request_owner": request_owner, } return cls( **payload, request_content_id=_content_id("decoding-request", payload), ) def __post_init__(self) -> None: for field in ( "request_id", "request_version", "submitted_at", "tenant_scope", "model_id", "model_version", "tokenizer_id", "tokenizer_version", "policy_content_id", "request_content_id", ): _identifier(field, getattr(self, field)) _owner("request_owner", self.request_owner) if type(self.prompt_token_ids) is not tuple: raise TypeError("prompt_token_ids must be an immutable tuple") @dataclass(frozen=True) class LogitsEvidence: logits_id: str logits_version: str captured_at: str model_id: str model_version: str tokenizer_id: str tokenizer_version: str logits_schema_version: str policy_content_id: str request_content_id: str position: int token_ids: tuple[int, ...] logits: tuple[float, ...] evidence_owner: str logits_content_id: str @classmethod def capture( cls, *, policy: DecodingPolicy, request: DecodingRequest, logits_id: str, logits_version: str, captured_at: str, position: int, token_ids: Iterable[int], logits: Iterable[float], evidence_owner: str, ) -> "LogitsEvidence": if type(policy) is not DecodingPolicy or type(request) is not DecodingRequest: raise TypeError("policy and request must be concrete records") frozen_tokens = tuple(token_ids) frozen_logits = tuple(logits) _validate_logits_values(policy, frozen_tokens, frozen_logits) _integer("position", position, 0, MAX_PROMPT_TOKENS + MAX_NEW_TOKENS) payload = { "logits_id": logits_id, "logits_version": logits_version, "captured_at": captured_at, "model_id": policy.model_id, "model_version": policy.model_version, "tokenizer_id": policy.tokenizer_id, "tokenizer_version": policy.tokenizer_version, "logits_schema_version": policy.logits_schema_version, "policy_content_id": policy.content_id, "request_content_id": request.request_content_id, "position": position, "token_ids": frozen_tokens, "logits": frozen_logits, "evidence_owner": evidence_owner, } return cls( **payload, logits_content_id=_content_id("decoding-logits", payload), ) def __post_init__(self) -> None: for field in ( "logits_id", "logits_version", "captured_at", "model_id", "model_version", "tokenizer_id", "tokenizer_version", "logits_schema_version", "policy_content_id", "request_content_id", "logits_content_id", ): _identifier(field, getattr(self, field)) _owner("evidence_owner", self.evidence_owner) if type(self.token_ids) is not tuple or type(self.logits) is not tuple: raise TypeError("token_ids and logits must be immutable tuples") @dataclass(frozen=True) class DecodingDecision: policy_content_id: str request_content_id: str logits_content_id: str candidate_token_ids: tuple[int, ...] selected_token_id: int selected_is_stop: bool distribution_content_id: str decision_content_id: str status: str claim: str def _validate_logits_values( policy: DecodingPolicy, token_ids: object, logits: object ) -> None: checked_tokens = _token_tuple( "token_ids", token_ids, maximum_count=policy.maximum_vocabulary ) if type(logits) is not tuple or len(logits) != len(checked_tokens): raise TypeError("logits must be an immutable tuple aligned to token_ids") for logit in logits: checked = _finite("logit", logit) if checked != 0.0 and abs(checked) < policy.minimum_numeric_resolution: raise ValueError("logit is below numeric resolution") if abs(checked) > policy.maximum_absolute_logit: raise ValueError("logit exceeds the contracted magnitude") def _validate_request(policy: DecodingPolicy, request: DecodingRequest) -> None: if type(request) is not DecodingRequest: raise TypeError("request must be a concrete DecodingRequest") request.__post_init__() scope = ( request.model_id, request.model_version, request.tokenizer_id, request.tokenizer_version, request.policy_content_id, ) expected = ( policy.model_id, policy.model_version, policy.tokenizer_id, policy.tokenizer_version, policy.content_id, ) if scope != expected: raise ValueError("request scope does not match the policy") _token_tuple( "prompt_token_ids", request.prompt_token_ids, maximum_count=MAX_PROMPT_TOKENS, unique=False, ) _integer( "requested_new_tokens", request.requested_new_tokens, 1, policy.maximum_new_tokens, ) payload = { field: getattr(request, field) for field in ( "request_id", "request_version", "submitted_at", "tenant_scope", "model_id", "model_version", "tokenizer_id", "tokenizer_version", "policy_content_id", "prompt_token_ids", "requested_new_tokens", "request_owner", ) } if request.request_content_id != _content_id("decoding-request", payload): raise ValueError("request content identity does not match contents") def _validate_logits( policy: DecodingPolicy, request: DecodingRequest, evidence: LogitsEvidence ) -> None: if type(evidence) is not LogitsEvidence: raise TypeError("evidence must be concrete LogitsEvidence") evidence.__post_init__() scope = ( evidence.model_id, evidence.model_version, evidence.tokenizer_id, evidence.tokenizer_version, evidence.logits_schema_version, evidence.policy_content_id, evidence.request_content_id, ) expected = ( policy.model_id, policy.model_version, policy.tokenizer_id, policy.tokenizer_version, policy.logits_schema_version, policy.content_id, request.request_content_id, ) if scope != expected: raise ValueError("logits scope does not match policy and request") _integer( "position", evidence.position, 0, MAX_PROMPT_TOKENS + MAX_NEW_TOKENS, ) if evidence.position < len(request.prompt_token_ids): raise ValueError("logits position precedes the end of the prompt") if evidence.position >= len(request.prompt_token_ids) + request.requested_new_tokens: raise ValueError("logits position exceeds the request generation bound") _validate_logits_values(policy, evidence.token_ids, evidence.logits) payload = { field: getattr(evidence, field) for field in ( "logits_id", "logits_version", "captured_at", "model_id", "model_version", "tokenizer_id", "tokenizer_version", "logits_schema_version", "policy_content_id", "request_content_id", "position", "token_ids", "logits", "evidence_owner", ) } if evidence.logits_content_id != _content_id("decoding-logits", payload): raise ValueError("logits content identity does not match contents") def _stable_distribution( policy: DecodingPolicy, evidence: LogitsEvidence ) -> tuple[tuple[int, float], ...]: allowed = set(policy.allowed_token_ids) selected = tuple( (token_id, logit / policy.temperature) for token_id, logit in zip(evidence.token_ids, evidence.logits) if token_id in allowed ) if not selected: raise ValueError("constraint removes every token from the logits") pivot = max(value for _, value in selected) exponentials = tuple((token_id, math.exp(value - pivot)) for token_id, value in selected) total = math.fsum(value for _, value in exponentials) if not math.isfinite(total) or total <= 0.0: raise ArithmeticError("stable softmax normalization failed") probabilities = tuple((token_id, value / total) for token_id, value in exponentials) if not all(math.isfinite(probability) for _, probability in probabilities): raise ArithmeticError("stable softmax produced a non-finite probability") return tuple(sorted(probabilities, key=lambda item: (-item[1], item[0]))) def apply_decoding_policy( policy: DecodingPolicy, request: DecodingRequest, evidence: LogitsEvidence, ) -> DecodingDecision: """Apply the bound policy with stable ranking and a SHA-256-derived draw.""" if type(policy) is not DecodingPolicy: raise TypeError("policy must be a concrete DecodingPolicy") policy.__post_init__() _validate_request(policy, request) _validate_logits(policy, request, evidence) ranked = _stable_distribution(policy, evidence) nucleus: list[tuple[int, float]] = [] nucleus_mass = 0.0 for item in ranked: nucleus.append(item) nucleus_mass = math.fsum((nucleus_mass, item[1])) if nucleus_mass >= policy.top_p: break if not nucleus: raise ArithmeticError("nucleus selection produced no candidates") nucleus_total = math.fsum(probability for _, probability in nucleus) normalized = tuple( (token_id, probability / nucleus_total) for token_id, probability in nucleus ) distribution_content_id = _content_id( "decoding-distribution", { "policy_content_id": policy.content_id, "request_content_id": request.request_content_id, "logits_content_id": evidence.logits_content_id, "candidates": normalized, }, ) draw_material = "|".join( ( str(policy.seed), policy.content_id, request.request_content_id, evidence.logits_content_id, str(evidence.position), ) ).encode("utf-8") draw_integer = int.from_bytes(hashlib.sha256(draw_material).digest(), "big") draw_numerator = draw_integer >> (256 - 53) draw = (draw_numerator + 0.5) / 2**53 cumulative = 0.0 selected_token_id = normalized[-1][0] for token_id, probability in normalized: cumulative = math.fsum((cumulative, probability)) if draw <= cumulative: selected_token_id = token_id break decision_payload = { "policy_content_id": policy.content_id, "request_content_id": request.request_content_id, "logits_content_id": evidence.logits_content_id, "candidate_token_ids": tuple(token_id for token_id, _ in normalized), "selected_token_id": selected_token_id, "selected_is_stop": selected_token_id in policy.stop_token_ids, "distribution_content_id": distribution_content_id, "status": "PASS", "claim": "REPRODUCIBLE_POLICY_EXECUTION_NOT_QUALITY_GUARANTEE", } return DecodingDecision( **decision_payload, decision_content_id=_content_id("decoding-decision", decision_payload), ) ILLUSTRATIVE_POLICY = DecodingPolicy( policy_id="support-assistant-policy", policy_version="policy-2026-08-25", model_id="decoder:support", model_version="model-v12", tokenizer_id="tokenizer:support", tokenizer_version="tokenizer-v4", logits_schema_version="logits-v2", temperature=0.8, top_p=0.75, seed=20260825, maximum_new_tokens=64, stop_token_ids=(2,), constraint_id="constraint:support-response", constraint_version="constraint-v3", allowed_token_ids=(2, 11, 17, 23), maximum_vocabulary=16, minimum_numeric_resolution=1e-12, maximum_absolute_logit=1_000.0, numeric_convention="binary64-stable-softmax-sha256-draw-v1", policy_owner="team:inference", ) ILLUSTRATIVE_REQUEST = DecodingRequest.capture( policy=ILLUSTRATIVE_POLICY, request_id="request-001", request_version="request-v1", submitted_at="2026-08-25T00:00:00Z", tenant_scope="tenant:academy", prompt_token_ids=(101, 205, 309), requested_new_tokens=16, request_owner="team:inference", ) ILLUSTRATIVE_LOGITS = LogitsEvidence.capture( policy=ILLUSTRATIVE_POLICY, request=ILLUSTRATIVE_REQUEST, logits_id="logits-001", logits_version="logits-snapshot-v1", captured_at="2026-08-25T00:00:01Z", position=3, token_ids=(2, 11, 17, 23), logits=(0.5, 3.0, 2.0, 1.0), evidence_owner="team:inference", ) def format_example() -> str: decision = apply_decoding_policy( ILLUSTRATIVE_POLICY, ILLUSTRATIVE_REQUEST, ILLUSTRATIVE_LOGITS ) return "\n".join( ( "example=illustrative_only", f"policy_content_id={decision.policy_content_id}", f"request_content_id={decision.request_content_id}", f"logits_content_id={decision.logits_content_id}", "temperature=0.800", "top_p=0.750", "candidate_tokens=" + ",".join(map(str, decision.candidate_token_ids)), f"selected_token={decision.selected_token_id}", f"selected_is_stop={str(decision.selected_is_stop).lower()}", f"status={decision.status}", f"claim={decision.claim}", ) ) if __name__ == "__main__": print(format_example())