"""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 SystemOneBody(BaseModel): """The Jev request. `model` is accepted with any value and ignored (Jev clients send 'jev-latest'); `images` is kept only to refuse it with a clear 422. Extra keys are dropped: only state+questions reach predict, exactly like the bench's 60-line wrapper.""" model_config = ConfigDict(extra="allow") state: State questions: dict[str, dict[str, Any]] @property def images_present(self) -> bool: """A NON-EMPTY images value only: [] and null count as absent (audit 2026-09-30 — a client that always sends the field must not be rejected for nothing).""" return bool((self.model_extra or {}).get("images")) 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 on_inference(fn): """Run fn on the ONE inference thread (INV-2) and map failures to contract codes (INV-4): the shared dispatch+mapping for every surface, semif and systemone alike.""" try: return await asyncio.get_running_loop().run_in_executor(executor, fn) 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 async def score(state: State, decisions: list[Decision], waves: list[list[Slot]]): return await on_inference(lambda: run(state, decisions, waves)) @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": [], "endpoints": ["/decide", "/decide/shared", "/v1/systemone"], "systemone": {"max_questions": MAX_QUESTIONS_PER_CALL, "chunking": "none", "images": "not supported", "max_tokens": settings.max_tokens}} @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() @app.post("/v1/systemone") async def systemone(request: Request): """Jev wire shape. Straight passthrough to DecisionEngine.predict — NEVER the semif mapping (a different prompt would change the answers). Contract § POST /v1/systemone.""" admit() # before the body is read try: payload = await parse(request, SystemOneBody) if payload.images_present: raise ApiError(422, "invalid_request", "images not supported (the vision tower is " "removed to fit the GPU 1 budget, INV-7)") questions = payload.questions if not 1 <= len(questions) <= MAX_QUESTIONS_PER_CALL: raise ApiError(422, "invalid_request", f"this request has {len(questions)} questions; the limit is " f"1..{MAX_QUESTIONS_PER_CALL} in one call — questions of one call " "share one prompt, so they are never chunked") response = await run_systemone({"state": payload.state, "questions": questions}) meta = engine.metadata return {"answers": response["answers"], "usage": response["usage"], "model": f"{meta.get('name')}@{meta.get('revision')}", # string: see contract **{k: v for k, v in response.items() if k not in ("answers", "usage", "model")}} finally: leave() async def run_systemone(call: dict) -> dict: """One native call on the inference thread, under the same lock and mapping as the semif surface (on_inference). The engine's own validate_request enforces the per-question Jev rules before any GPU work; its ValueError is a 422 like the token-limit check.""" def one(): with inference: response, _sha = engine.predict(call) usage = response.get("usage") response["usage"] = usage if isinstance(usage, dict) else {} # contract: {} not absent return response response = await on_inference(one) if not isinstance(response.get("answers"), dict): raise ApiError(500, "scoring_failed", "the engine returned no answers object") return response return app