Files
esh-pfi-infrastructure/services/intern-decision-serve/src/intern_decision_serve/app.py
T
vh 1866c003e8 fix(intern-decision-serve): 0.1.2 treats empty/null images as absent
Audit finding on 0.1.1: a Jev client that always sends an images array with no
images was rejected for nothing. Only a non-empty value is 422 now. Also aligns
the contract's response example with the wire (model is the name@revision
string, not an object).

Re-accepted live: images []/null -> 200, ["a.png"] -> 422; JevBench all 202/231,
hard 83/111, 0 changed rows across the bench's r1..r4 (924).
2026-09-30 13:04:31 -07:00

428 lines
20 KiB
Python

"""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 <decision> 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