atlas-iac/services/hermes/scripts/lan_generate.py

176 lines
8.1 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""Stateless, pinned RTX 3080 inference without routing, proxies, or persistence."""
import json
import threading
import time
from urllib.error import HTTPError, URLError
from urllib.request import HTTPRedirectHandler, ProxyHandler, Request, build_opener
UPSTREAM = "http://ollama-gpu.ai.svc.cluster.local:11434"
MODEL = "qwen2.5:14b-instruct-q4_0"
DIGEST = "5449194ff8035ccb13a6409a5814de6c8f9c39f555f429e383ae0fb7137001bd"
RUNTIME = "0.13.5"
MAX_BODY = 131072
CONTEXT = 8192
TIMEOUT = 1200
_inference = threading.BoundedSemaphore(1)
class NoRedirect(HTTPRedirectHandler):
"""Reject redirects before another host can receive a request."""
def redirect_request(self, req, fp, code, msg, headers, newurl):
return None
_http = build_opener(ProxyHandler({}), NoRedirect())
class InputTooLarge(ValueError):
"""Distinguish context admission failures from malformed requests."""
def _finite_json(raw):
"""Reject nonstandard constants without including input in diagnostics."""
def invalid_constant(_value):
raise ValueError("nonfinite JSON value")
return json.loads(raw, parse_constant=invalid_constant)
def _local_schema(value):
"""Allow bounded schemas without external references or dynamic resolution."""
if isinstance(value, dict):
for key, child in value.items():
if key in ("$ref", "$dynamicRef") and (
not isinstance(child, str) or not child.startswith("#")
):
raise ValueError("external schema references are unsupported")
_local_schema(child)
elif isinstance(value, list):
for child in value:
_local_schema(child)
def normalize(body):
"""Return a bounded native generate payload; never trim source material."""
if not body or len(body) > MAX_BODY:
raise InputTooLarge("request body exceeds limit")
payload = _finite_json(body)
if not isinstance(payload, dict) or set(payload) - {"model", "prompt", "stream", "format", "options"}:
raise ValueError("unsupported fields")
if payload.get("model") != MODEL or payload.get("stream") is not False:
raise ValueError("exact model and stream=false are required")
prompt = payload.get("prompt")
if not isinstance(prompt, str) or not prompt.strip():
raise ValueError("nonempty prompt is required")
if "format" in payload:
shape = payload["format"]
if shape != "json":
if not isinstance(shape, dict) or shape.get("type") != "object":
raise ValueError("format must be json or an object JSON schema")
if len(json.dumps(shape).encode()) > 32768:
raise InputTooLarge("schema exceeds limit")
_local_schema(shape)
options = payload.get("options", {})
allowed = {"num_ctx", "num_predict", "temperature", "top_p", "top_k", "seed"}
if not isinstance(options, dict) or set(options) - allowed:
raise ValueError("unsupported options")
options = {"num_ctx": CONTEXT, "num_predict": 256, "temperature": 0,
"top_p": 1, "top_k": 40, "seed": 0, **options}
for key, lower, upper in (("num_ctx", CONTEXT, CONTEXT), ("num_predict", 1, 2048),
("top_k", 1, 100), ("seed", 0, 2147483647)):
if type(options[key]) is not int or not lower <= options[key] <= upper:
raise ValueError("unsupported integer option")
for key, upper in (("temperature", 2), ("top_p", 1)):
if type(options[key]) not in (int, float) or not 0 <= options[key] <= upper:
raise ValueError("unsupported sampling option")
# UTF-8 bytes upper-bound this model's byte-level tokenizer. Reserve the
# full output budget plus 1024 template tokens before calling Ollama.
if len(prompt.encode("utf-8")) + options["num_predict"] + 1024 > CONTEXT:
raise InputTooLarge("prompt exceeds conservative context budget")
return {**payload, "options": options}
def _request(path, payload=None, timeout=10):
"""Call only the fixed in-cluster backend and bound the response body."""
data = None if payload is None else json.dumps(payload, allow_nan=False).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 exceeds limit")
return _finite_json(raw)
def verify_model():
"""Check the runtime and local manifest before sending any prompt text."""
if _request("/api/version").get("version") != RUNTIME:
raise ValueError("runtime changed")
matches = [item for item in _request("/api/tags")["models"] if item.get("name") == MODEL]
if len(matches) != 1 or matches[0].get("digest", "").removeprefix("sha256:") != DIGEST:
raise ValueError("pinned model unavailable")
return {"model": MODEL, "model_digest": DIGEST, "runtime": RUNTIME,
"placement": "titan-24/RTX-3080-10GB", "context_tokens": CONTEXT,
"max_output_tokens": 2048, "timeout_seconds": TIMEOUT,
"concurrency": 1, "fallback": None, "protocol_version": 1,
"serving_configuration": {"parallel_requests": 1, "max_queue": 1,
"flash_attention": True, "kv_cache_type": "q8_0"}}
def health():
"""Report authenticated model readiness, not merely process liveness."""
try:
return 200, {"status": "ok", **verify_model()}
except (OSError, URLError, ValueError, KeyError, TypeError, AttributeError):
return 503, {"error": "pinned local model unavailable; no fallback"}
def generate(body):
"""Reject busy requests and return native output with reproducible metadata."""
try:
payload = normalize(body)
except InputTooLarge:
return 413, {"error": "request exceeds body, schema, or context budget"}
except (ValueError, TypeError, RecursionError):
return 400, {"error": "invalid inference request; consult endpoint contract"}
if not _inference.acquire(blocking=False):
return 429, {"error": "local inference busy; retry explicitly"}
started = time.monotonic()
try:
provenance = verify_model()
result = _request("/api/generate", payload, timeout=max(1, TIMEOUT - (time.monotonic() - started)))
if result.get("model") != MODEL or result.get("done") is not True:
raise ValueError("model mismatch or incomplete response")
if result.get("done_reason") == "length":
return 422, {"error": "output token budget exhausted; result is incomplete"}
if not isinstance(result.get("response"), str):
raise ValueError("missing output")
if "format" in payload:
try:
_finite_json(result["response"])
except ValueError:
return 422, {"error": "model returned invalid JSON"}
# Project the response so context token IDs, debug fields, or upstream
# error text can never become an accidental history or logging channel.
fields = ("model", "created_at", "response", "done", "done_reason",
"total_duration", "load_duration", "prompt_eval_count",
"prompt_eval_duration", "eval_count", "eval_duration")
output = {key: result[key] for key in fields if key in result}
output["inference_provenance"] = {
**provenance, "options": payload["options"], "backend_api": "/api/generate",
"wall_seconds": round(time.monotonic() - started, 3),
}
return 200, output
except TimeoutError:
return 504, {"error": "local inference timed out; no fallback"}
except HTTPError as exc:
if exc.code in (429, 503):
return 503, {"error": "local inference capacity unavailable; no fallback"}
return 502, {"error": "local inference failed; no fallback"}
except (OSError, URLError, ValueError, KeyError, TypeError, AttributeError, RecursionError):
return 503, {"error": "pinned local inference unavailable; no fallback"}
finally:
_inference.release()