144 lines
6.7 KiB
Python
Executable File

#!/usr/bin/env python3
"""Serve pinned batch models through a fixed, local-only Ollama endpoint."""
import json
import threading
import time
from urllib.error import HTTPError, URLError
from urllib.request import HTTPRedirectHandler, ProxyHandler, Request, build_opener
UPSTREAM = "http://ollama-batch.ai.svc.cluster.local:11434"
VERSION = "0.34.1"
PINS = {
"qwen3.5:9b": "6488c96fa5faab64bb65cbd30d4289e20e6130ef535a93ef9a49f42eda893ea7",
"qwen3.6:27b": "9d5803d493a991af27b9441c098aa56f2ed7bbd260877f075ec09b575c049bc3",
}
MAX_BODY = 1048576
TIMEOUT = 1800
_inference = threading.BoundedSemaphore(1)
class NoRedirect(HTTPRedirectHandler):
"""Never forward inference bodies or credentials through a redirect."""
def redirect_request(self, req, fp, code, msg, headers, newurl):
return None
_http = build_opener(ProxyHandler({}), NoRedirect())
def _request(path, payload=None, timeout=10):
"""Read a bounded local response without consulting environment proxies."""
data = None if payload is None else json.dumps(payload).encode()
request = Request(UPSTREAM + path, data=data, headers={"Content-Type": "application/json"})
with _http.open(request, timeout=timeout) as response:
raw = response.read(4 * MAX_BODY + 1)
if len(raw) > 4 * MAX_BODY:
raise ValueError("upstream response too large")
return json.loads(raw)
def catalog():
"""Fail closed if the runtime version or any approved model digest changes."""
runtime = _request("/api/version")["version"]
if runtime != VERSION:
raise ValueError("runtime version mismatch")
installed = {item["name"]: item for item in _request("/api/tags")["models"]}
models = []
for name, digest in PINS.items():
item = installed.get(name, {})
if item.get("digest", "").removeprefix("sha256:") != digest:
raise ValueError("model digest mismatch or model unavailable")
models.append({"model": name, "digest": digest, "size": item["size"],
"details": item.get("details", {})})
return {"runtime": runtime, "placement": "titan-23/cpu", "models": models,
"context_limits": [16384, 32768, 65536], "max_output_tokens": 16384,
"timeout_seconds": TIMEOUT, "fallback": None}
def normalize(body):
"""Validate explicit inference settings and conservatively avoid truncation."""
payload = json.loads(body)
allowed = {"model", "prompt", "stream", "think", "format", "options"}
if not isinstance(payload, dict) or set(payload) - allowed:
raise ValueError("unsupported batch fields")
if payload.get("model") not in PINS:
raise ValueError("an exact approved batch model is required")
if payload.get("stream") is not False or type(payload.get("think")) is not bool:
raise ValueError("stream=false and explicit boolean think are required")
prompt = payload.get("prompt")
if not isinstance(prompt, str) or not prompt.strip():
raise ValueError("nonempty text prompt is required")
shape = payload.get("format")
if shape != "json" and not isinstance(shape, dict):
raise ValueError("JSON format or a JSON schema is required")
if isinstance(shape, dict) and (shape.get("type") != "object" or len(json.dumps(shape)) > 65536):
raise ValueError("bounded object JSON schema is required")
options = payload.get("options")
required = {"num_ctx", "num_predict", "temperature", "top_p", "top_k", "seed"}
if not isinstance(options, dict) or set(options) != required:
raise ValueError("explicit num_ctx, num_predict, temperature, top_p, top_k and seed are required")
context, count = options["num_ctx"], options["num_predict"]
if type(context) is not int or context not in (16384, 32768, 65536):
raise ValueError("unsupported context size")
if type(count) is not int or not 1 <= count <= 16384:
raise ValueError("num_predict must be between 1 and 16384")
# One UTF-8 byte per token is a deliberately conservative upper bound.
# Reserve template overhead and the full output budget; never trim a suite.
if len(prompt.encode("utf-8")) + count + 1024 > context:
raise ValueError("prompt exceeds conservative context budget; use candidate batches and reconciliation")
for key, upper in (("temperature", 2), ("top_p", 1)):
value = options[key]
if type(value) not in (int, float) or not 0 <= value <= upper:
raise ValueError("invalid sampling value")
for key, low, high in (("seed", 0, 2147483647), ("top_k", 1, 100)):
if type(options[key]) is not int or not low <= options[key] <= high:
raise ValueError("invalid integer sampling value")
payload["options"] = {**options, "num_thread": 16, "num_gpu": 0}
payload["keep_alive"] = "10m"
return payload
def models_response():
"""Return verifiable configuration without leaking internal failure details."""
try:
return 200, catalog()
except (OSError, HTTPError, URLError, ValueError, KeyError, TypeError):
return 503, {"error": "pinned local batch models unavailable; no fallback"}
def generate(body):
"""Serialize a batch call, attaching measured wall time and exact provenance."""
try:
payload = normalize(body)
except (ValueError, TypeError):
return 400, {"error": "invalid batch request or context budget exceeded"}
if not _inference.acquire(blocking=False):
return 429, {"error": "local batch model busy; retry explicitly"}
started = time.monotonic()
try:
provenance = catalog()
result = _request("/api/generate", payload, timeout=TIMEOUT)
if result.get("model") != payload["model"] or not result.get("done"):
raise ValueError("unexpected model or incomplete generation")
if result.get("done_reason") == "length":
return 422, {"error": "output budget exhausted; result is incomplete"}
if not isinstance(result.get("response"), str):
raise ValueError("missing output")
json.loads(result["response"])
# Keep the requested rationale in the answer, not raw hidden thinking.
result.pop("thinking", None)
result.pop("context", None)
result["batch_provenance"] = {
"model_digest": PINS[payload["model"]], "runtime": provenance["runtime"],
"placement": provenance["placement"], "options": payload["options"],
"think": payload["think"], "wall_seconds": round(time.monotonic() - started, 3),
}
return 200, result
except (OSError, HTTPError, URLError, ValueError, KeyError, TypeError):
return 503, {"error": "pinned local batch inference failed; no fallback"}
finally:
_inference.release()