"""intern-decision-serve HTTP layer. Contract: intern-decision-serve.contract.md. The request surface is semif-serve's. Each semif decision becomes one Jev `choice` question; the questions of a request are asked through the engine (the checkpoint's own DecisionEngine.predict) in calls of at most 16, and the answers are re-keyed into semif's result shape. """ from __future__ import annotations import asyncio import hmac import itertools import json import math import threading import time from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from typing import Any, Literal from fastapi import FastAPI, Request from fastapi.responses import JSONResponse from pydantic import BaseModel, ConfigDict, ValidationError from .config import INFERENCE_PY_SHA256, MAX_QUESTIONS_PER_CALL, Settings from .errors import OutOfMemory, ScoringFailed OPEN_PATHS = frozenset({"/health"}) State = str | dict[str, Any] | list[Any] MIN_OPTIONS, MAX_OPTIONS = 2, 16 # SemIf's rule (its 16 answer letters), kept for the surface MAX_OPTIONS_FOR_ALL = 4 # "orderings": "all" asks n! orderings LOG_FLOOR = 1e-300 # a probability of exactly 0 before the log (combine) PROMPT_VERSION = f"intern-decision-jev/inference.py@{INFERENCE_PY_SHA256[:12]}" READOUT = ("logits at the position before each marker, softmax over the field's answer " "symbols (the checkpoint's own inference.py, DecisionEngine.predict)") PROBABILITY_STATUS = ("temperature-scaled by the checkpoint's own shipped calibration (see calibration); " "vendor-fitted, not fitted on our workloads") CHUNKING = (f"/decide/shared questions are packed greedily, in request order, into calls of at most " f"{MAX_QUESTIONS_PER_CALL} (1-{MAX_QUESTIONS_PER_CALL}, {MAX_QUESTIONS_PER_CALL + 1}-" f"{2 * MAX_QUESTIONS_PER_CALL}, ...); each call is one prompt, so the questions in a call are " f"asked together. With orderings, ordering k of every decision forms wave k, packed the same way.") class Option(BaseModel): model_config = ConfigDict(extra="ignore") id: str description: str class Decision(BaseModel): model_config = ConfigDict(extra="ignore") id: str question: str options: list[Option] orderings: Literal["none", "rotations", "all"] = "none" class DecideBody(Decision): state: State workload: str | None = None class SharedBody(BaseModel): model_config = ConfigDict(extra="ignore") state: State decisions: list[Decision] workload: str | None = None class ApiError(Exception): def __init__(self, status: int, code: str, message: str): super().__init__(message) self.status, self.code, self.message = status, code, message def error(status: int, code: str, message: str) -> JSONResponse: return JSONResponse(status_code=status, content={"error": {"code": code, "message": message}}) def _first_error(exc: ValidationError) -> str: first = exc.errors()[0] where = ".".join(str(p) for p in first.get("loc", ())) or "body" return f"{where}: {first.get('msg', 'invalid')}" async def read_limited(stream, declared: str | None, limit: int) -> bytes: """Read a request body, refusing it once it would exceed `limit` bytes. The check runs BEFORE a chunk is kept, and nothing after the crossing chunk is read. A declared length is trusted only as ASCII digits: `"²".isdigit()` is True but `int("²")` raises.""" too_large = ApiError(413, "request_too_large", f"request body exceeds {limit} bytes") if declared is not None and declared.isascii() and declared.isdigit() and int(declared) > limit: raise too_large body = bytearray() async for chunk in stream: if len(body) + len(chunk) > limit: raise too_large body.extend(chunk) return bytes(body) def check_request(state: State, decisions: list[Decision], workload: str | None) -> None: """SemIf's row validation, re-stated because SemIf is gone; a violation is a 422.""" def bad(message: str) -> ApiError: return ApiError(422, "invalid_request", message) if workload is not None: raise bad(f"unknown workload {workload!r}: no per-workload calibration is configured; " "the model's own calibration always applies") if not state: raise bad("state must be a nonempty string, object, or array") try: json.dumps(state, ensure_ascii=False, allow_nan=False) except (TypeError, ValueError): raise bad("state must be finite JSON-compatible data") from None if len({d.id for d in decisions}) != len(decisions): raise bad("decision ids must be unique") for d in decisions: if not d.id or not d.question: raise bad("id and question must be nonempty strings") if not MIN_OPTIONS <= len(d.options) <= MAX_OPTIONS: raise bad(f"decision {d.id!r}: options must contain {MIN_OPTIONS}-{MAX_OPTIONS} entries") if len({o.id for o in d.options}) != len(d.options): raise bad(f"decision {d.id!r}: option ids must be unique") def ordering_count(d: Decision) -> int: """How many orderings `d` asks, without building them (the row cap is checked first).""" n = len(d.options) if d.orderings == "none": return 1 if d.orderings == "rotations": return n if n > MAX_OPTIONS_FOR_ALL: raise ApiError(422, "invalid_request", f"decision {d.id!r}: orderings 'all' allows at most " f"{MAX_OPTIONS_FOR_ALL} options ({n} given); use 'rotations'") return math.factorial(n) def ordering_perms(d: Decision) -> list[tuple[int, ...]]: """Index permutations of the caller's options, the caller's own order first.""" n = len(d.options) if d.orderings == "none": return [tuple(range(n))] if d.orderings == "rotations": return [tuple((start + k) % n for k in range(n)) for start in range(n)] if n > MAX_OPTIONS_FOR_ALL: raise ApiError(422, "invalid_request", f"decision {d.id!r}: orderings 'all' allows at most " f"{MAX_OPTIONS_FOR_ALL} options ({n} given); use 'rotations'") return list(itertools.permutations(range(n))) @dataclass(frozen=True) class Slot: """One question to ask: ordering `k` of decision number `decision`.""" decision: int k: int row_id: str option_ids: list[str] question: dict def plan_waves(decisions: list[Decision]) -> list[list[Slot]]: """Wave k holds ordering k of every decision that has more than k orderings, in request order. Wave 0 is the request as written; no wave holds two orderings of one decision.""" perms = [ordering_perms(d) for d in decisions] waves = [] for k in range(max(map(len, perms), default=0)): wave = [] for i, (d, ps) in enumerate(zip(decisions, perms)): if k < len(ps): options = [d.options[j] for j in ps[k]] wave.append(Slot(i, k, d.id if d.orderings == "none" else f"{d.id}#o{k}", [o.id for o in options], {"type": "choice", "instructions": d.question, "criteria": {o.id: o.description for o in options}})) waves.append(wave) return waves def pack(slots: list[Slot]) -> list[list[Slot]]: """Greedy, in request order: at most MAX_QUESTIONS_PER_CALL questions per call.""" return [slots[k:k + MAX_QUESTIONS_PER_CALL] for k in range(0, len(slots), MAX_QUESTIONS_PER_CALL)] def field_names(count: int) -> list[str]: """Positional: `q` for a one-question call, `q1`..`qN` otherwise. Decision ids never reach the prompt.""" return ["q"] if count == 1 else [f"q{i + 1}" for i in range(count)] def row_result(slot: Slot, response: dict, sha: str, field: str, questions: int, index: int, model: dict) -> dict: """INV-1: re-key the model's answer for `field`; never fill in a number it did not return.""" answer = response["answers"][field] probs = answer["probabilities"] missing = [i for i in slot.option_ids if i not in probs] if missing: raise ScoringFailed(f"the model's answer for field {field!r} lacks option ids {missing}") return {"id": slot.row_id, "option_ids": slot.option_ids, "probabilities": [probs[i] for i in slot.option_ids], "top": answer["choice"], "confidence": answer["confidence"], "calibration": response["calibration"], "native": answer, "input_tokens": response["usage"]["input_tokens"], "prompt_sha256": sha, "prompt_version": PROMPT_VERSION, "model": model, "readout": READOUT, "probability_status": PROBABILITY_STATUS, "call": {"index": index, "field": field, "questions": questions}} def built(fn, *args): """Building the response from the model's output is server-side work: any failure there is a 500 scoring_failed, never a 422 the caller would read as their own bad input (review 2026-09-30).""" try: out = fn(*args) json.dumps(out, allow_nan=False) # a NaN/inf would otherwise crash rendering OUTSIDE the envelope return out except ScoringFailed: raise except Exception as exc: # noqa: BLE001 raise ScoringFailed(f"building the response failed: {type(exc).__name__}: {exc}") from None def combine(d: Decision, results: list[dict]) -> dict: """semif-serve's averaging over log p (the model returns probabilities, not logits): the mean per option id, renormalised. The per-ordering results ride along unchanged.""" option_ids = [o.id for o in d.options] logp: dict[str, list[float]] = {i: [] for i in option_ids} probs: dict[str, list[float]] = {i: [] for i in option_ids} for result in results: for oid, p in zip(result["option_ids"], result["probabilities"]): logp[oid].append(math.log(max(p, LOG_FLOOR))) probs[oid].append(p) means = [sum(logp[i]) / len(logp[i]) for i in option_ids] peak = max(means) weights = [math.exp(m - peak) for m in means] combined = [w / sum(weights) for w in weights] winner = option_ids[combined.index(max(combined))] tops = [r["top"] for r in results] return {"id": d.id, "option_ids": option_ids, "combined": {"method": d.orderings, "orderings": len(results), "probabilities": combined, "top": winner, "agreement": tops.count(winner) / len(tops), "spread": {i: [min(probs[i]), max(probs[i])] for i in option_ids}}, "orderings": results} def inference_thread() -> ThreadPoolExecutor: """INV-2: the ONE host thread that ever touches the model. torch keeps CUDA state per host thread (cuBLAS handles and workspaces), partly outside the VRAM cap: measured 2026-09-30, 40 worker threads added 252 MiB outside the cap and 326 MiB inside it.""" return ThreadPoolExecutor(max_workers=1, thread_name_prefix="inference") def create_app(settings: Settings, engine: Any, executor: ThreadPoolExecutor | None = None) -> FastAPI: """`executor` must be the single thread the engine was loaded on (main.py); tests get a fresh one.""" app = FastAPI(title="intern-decision-serve") expected = f"Bearer {settings.api_token}".encode() executor = executor or inference_thread() inference = threading.Lock() # INV-2: one request's calls at a time (the executor has one thread too) in_progress = 0 # POSTs admitted and not yet answered def admit(): nonlocal in_progress if in_progress >= settings.max_queue: raise ApiError(429, "busy", f"{in_progress} requests already in progress (limit {settings.max_queue})") in_progress += 1 def leave(): nonlocal in_progress in_progress -= 1 @app.middleware("http") async def require_bearer(request: Request, call_next): if request.url.path not in OPEN_PATHS: supplied = request.headers.get("authorization", "").encode() if not hmac.compare_digest(supplied, expected): # INV-6 return error(401, "unauthorized", "missing or wrong bearer token") return await call_next(request) @app.exception_handler(ApiError) async def api_error(_request: Request, exc: ApiError): return error(exc.status, exc.code, exc.message) async def parse(request: Request, model: type[BaseModel]): try: return model.model_validate_json( await read_limited(request.stream(), request.headers.get("content-length"), settings.max_body_bytes)) except ValidationError as exc: raise ApiError(422, "invalid_request", _first_error(exc)) from exc def planned(state: State, decisions: list[Decision], workload: str | None) -> list[list[Slot]]: check_request(state, decisions, workload) rows = sum(map(ordering_count, decisions)) # counted before anything is built if not 1 <= rows <= settings.max_decisions: raise ApiError(422, "invalid_request", f"this request scores {rows} rows; the limit is 1..{settings.max_decisions}") return plan_waves(decisions) def run(state: State, decisions: list[Decision], waves: list[list[Slot]]) -> tuple[list[dict], dict]: """Every call of one request, back to back under the lock (INV-2); results in request order.""" with inference: return _run(state, decisions, waves) def _run(state: State, decisions: list[Decision], waves: list[list[Slot]]) -> tuple[list[dict], dict]: started = time.perf_counter() model = engine.metadata # static: results never carry live memory numbers by_slot: dict[tuple[int, int], dict] = {} sizes, tokens, inference_ms = [], [], 0.0 for chunk in (chunk for wave in waves for chunk in pack(wave)): fields = field_names(len(chunk)) response, sha = engine.predict({"state": state, "questions": {f: s.question for f, s in zip(fields, chunk)}}) index = len(sizes) sizes.append(len(chunk)) tokens.append(response["usage"]["input_tokens"]) inference_ms += response["timing"]["inference_ms"] for f, s in zip(fields, chunk): by_slot[(s.decision, s.k)] = built(row_result, s, response, sha, f, len(chunk), index, model) out = [] for i, d in enumerate(decisions): if d.orderings == "none": out.append(by_slot[(i, 0)]) else: out.append(built(combine, d, [by_slot[(i, k)] for k in range(ordering_count(d))])) timing = {"total_seconds": time.perf_counter() - started, "batch_size": len(by_slot), "calls": len(sizes), "questions_per_call": sizes, "input_tokens": tokens, "inference_seconds": inference_ms / 1000} return out, timing async def score(state: State, decisions: list[Decision], waves: list[list[Slot]]): """Run one request's calls on the inference thread; map their failures to contract codes.""" try: return await asyncio.get_running_loop().run_in_executor(executor, run, state, decisions, waves) except ValueError as exc: # the model's own validation, token limit, ... raise ApiError(422, "invalid_request", str(exc)) from exc except OutOfMemory as exc: raise ApiError(503, "out_of_memory", str(exc)) from exc except Exception as exc: # noqa: BLE001 — any other failure, building the response included raise ApiError(500, "scoring_failed", f"{type(exc).__name__}: {exc}") from exc @app.get("/health") async def health(): return {"status": "ok", "model": engine.health(), "vram_cap_gib": settings.vram_cap_gib, "max_tokens": settings.max_tokens, "max_decisions": settings.max_decisions, "max_questions_per_call": MAX_QUESTIONS_PER_CALL, "chunking": CHUNKING, "workloads": []} @app.post("/decide") async def decide(request: Request): admit() # before the body is read try: body = await parse(request, DecideBody) results, timing = await score(body.state, [body], planned(body.state, [body], body.workload)) if body.orderings != "none": return results[0] return {**results[0], "total_seconds": timing["total_seconds"], "forward_seconds": timing["inference_seconds"]} finally: leave() @app.post("/decide/shared") async def decide_shared(request: Request): admit() try: body = await parse(request, SharedBody) results, timing = await score(body.state, body.decisions, planned(body.state, body.decisions, body.workload)) return {"results": results, "timing": timing} finally: leave() return app