"""semif-serve HTTP layer. Contract: semif-serve.contract.md.""" from __future__ import annotations import hmac import itertools import math import threading from typing import Any, Literal from fastapi import FastAPI, Request from fastapi.concurrency import run_in_threadpool from fastapi.responses import JSONResponse from pydantic import BaseModel, ConfigDict, ValidationError from .config import SEMIF_COMMIT, Settings from .errors import OutOfMemory OPEN_PATHS = frozenset({"/health"}) State = str | dict[str, Any] | list[Any] 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" def row(self, state: State) -> dict: """The SemIf row shape: exactly id, state, question, options.""" return {"id": self.id, "state": state, "question": self.question, "options": [o.model_dump() for o in self.options]} 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')}" MAX_OPTIONS_FOR_ALL = 4 def ordering_perms(decision: Decision) -> list[tuple[int, ...]]: """Index permutations of the caller's options, the caller's own order first.""" n = len(decision.options) if decision.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"orderings 'all' allows at most {MAX_OPTIONS_FOR_ALL} options ({n} given); use 'rotations'") return list(itertools.permutations(range(n))) def expanded_rows(decision: Decision, state: State, perms: list[tuple[int, ...]]) -> list[dict]: base = decision.row(state) return [{**base, "id": f"{decision.id}#o{k}", "options": [base["options"][i] for i in perm]} for k, perm in enumerate(perms)] def combine(decision: Decision, results: list[dict]) -> dict: """Average per-ordering log-softmax by option id; the native results ride along unchanged.""" option_ids = [o.id for o in decision.options] logp: dict[str, list[float]] = {i: [] for i in option_ids} probs: dict[str, list[float]] = {i: [] for i in option_ids} tops = [] for result in results: logits = result["option_logits"] top = max(logits) lse = top + math.log(sum(math.exp(x - top) for x in logits)) for oid, x, p in zip(result["option_ids"], logits, result["probabilities"]): logp[oid].append(x - lse) probs[oid].append(p) tops.append(result["option_ids"][logits.index(top)]) 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_p = [w / sum(weights) for w in weights] winner = option_ids[combined_p.index(max(combined_p))] return { "id": decision.id, "option_ids": option_ids, "combined": { "method": decision.orderings, "orderings": len(results), "probabilities": combined_p, "top": winner, "agreement": tops.count(winner) / len(tops), "spread": {i: [min(probs[i]), max(probs[i])] for i in option_ids}, }, "orderings": results, } 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 (bug hunt C2). A declared length is trusted only as ASCII digits: `"²".isdigit()` is True but `int("²")` raises (S3).""" 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 calibrated_view(result: dict, workload: str, temperature: float) -> dict: """softmax(option_logits / T): the native fields are left exactly as SemIf returned them (INV-1).""" scaled = [x / temperature for x in result["option_logits"]] top = max(scaled) weights = [math.exp(x - top) for x in scaled] total = sum(weights) return {"workload": workload, "temperature": temperature, "probabilities": [w / total for w in weights]} def create_app(settings: Settings, engine: Any) -> FastAPI: app = FastAPI(title="semif-serve") expected = f"Bearer {settings.api_token}".encode() inference = threading.Lock() # INV-2: one scorer call at a time, off the event loop in_progress = 0 # POSTs admitted and not yet answered (bug hunt C6) 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 def build(fn, *args): """Response construction from a scorer result (calibration, averaging) maps its failures to the 500 envelope too, instead of escaping as a bare 500 (bug hunt C3).""" try: return fn(*args) except ApiError: raise except Exception as exc: # noqa: BLE001 raise ApiError(500, "scoring_failed", f"building the response failed: {type(exc).__name__}: {exc}") from exc def locked(fn, *args): with inference: return fn(*args) @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 async def score(fn, *args): """Run one scorer call in a worker thread under the lock; map its failures to contract codes.""" try: return await run_in_threadpool(locked, fn, *args) except ValueError as exc: # SemIf validation, token limit, tokenisation 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 scorer failure raise ApiError(500, "scoring_failed", f"{type(exc).__name__}: {exc}") from exc def temperature_for(workload: str | None) -> float | None: if workload is None: return None if workload not in settings.calibration: raise ApiError(422, "invalid_request", f"unknown workload {workload!r}") return settings.calibration[workload] def with_calibration(result: dict, workload: str | None, temperature: float | None) -> dict: if temperature is None: return result return {**result, "calibrated": calibrated_view(result, workload, temperature)} @app.get("/health") async def health(): return {"status": "ok", "semif_commit": SEMIF_COMMIT, "model": engine.health(), "vram_cap_gib": settings.vram_cap_gib, "max_tokens": settings.max_tokens, "max_decisions": settings.max_decisions, "workloads": sorted(settings.calibration)} async def score_batch(decisions: list[Decision], state: State, workload: str | None) -> tuple[list[dict], dict]: """One engine.shared call for every row of every decision; results in request order.""" plan = [] # (decision, perms or None, row count) rows: list[dict] = [] for d in decisions: if d.orderings == "none": plan.append((d, None, 1)) rows.append(d.row(state)) else: if workload is not None: raise ApiError(422, "invalid_request", "workload calibration is not available together with orderings") perms = ordering_perms(d) plan.append((d, perms, len(perms))) rows.extend(expanded_rows(d, state, perms)) if not 1 <= len(rows) <= settings.max_decisions: raise ApiError(422, "invalid_request", f"this request expands to {len(rows)} scored rows; the limit is 1..{settings.max_decisions}") temperature = temperature_for(workload) results, timing = await score(engine.shared, rows) out, cursor = [], 0 for d, perms, count in plan: chunk = results[cursor:cursor + count] cursor += count out.append(build(with_calibration, chunk[0], workload, temperature) if perms is None else build(combine, d, chunk)) return out, timing @app.post("/decide") async def decide(request: Request): admit() try: body = await parse(request, DecideBody) if body.orderings != "none": results, _timing = await score_batch([body], body.state, body.workload) return results[0] temperature = temperature_for(body.workload) result = await score(engine.direct, body.row(body.state)) return build(with_calibration, result, body.workload, temperature) finally: leave() @app.post("/decide/shared") async def decide_shared(request: Request): admit() try: body = await parse(request, SharedBody) if not 1 <= len(body.decisions) <= settings.max_decisions: raise ApiError(422, "invalid_request", f"decisions must hold 1..{settings.max_decisions} entries, got {len(body.decisions)}") results, timing = await score_batch(body.decisions, body.state, body.workload) return {"results": results, "timing": timing} finally: leave() return app