atlas-iac/scripts/ops/hermes_batch_client.py

196 lines
8.9 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""Call the private batch API using urllib, pinned LAN TLS and a local cache."""
import argparse
import hashlib
import http.client
import json
import os
from pathlib import Path
import socket
import ssl
import tempfile
import time
from urllib.error import HTTPError, URLError
from urllib.request import HTTPSHandler, HTTPRedirectHandler, ProxyHandler, Request, build_opener
HOST = "worker.bstein.dev"
ADDRESS = "192.168.22.50"
BASE = f"https://{HOST}/local-model/api/batch"
RUNTIME = "0.34.1"
MODELS = {
"qwen3.5:9b": "6488c96fa5faab64bb65cbd30d4289e20e6130ef535a93ef9a49f42eda893ea7",
"qwen3.6:27b": "9d5803d493a991af27b9441c098aa56f2ed7bbd260877f075ec09b575c049bc3",
}
def canonical(value):
"""Encode deterministic JSON for content-addressed cache keys."""
return json.dumps(value, sort_keys=True, separators=(",", ":"), ensure_ascii=False, allow_nan=False).encode()
class LanConnection(http.client.HTTPSConnection):
"""Connect to the fixed LAN IP, retaining certificate and SNI hostname checks."""
def connect(self):
if self.host != HOST or self.port != 443 or self._tunnel_host:
raise ValueError("only the approved LAN endpoint is supported")
connection = socket.create_connection((ADDRESS, 443), self.timeout)
try:
self.sock = self._context.wrap_socket(connection, server_hostname=HOST)
except Exception:
connection.close()
raise
class LanHandler(HTTPSHandler):
"""Use the pinned connection without globally changing DNS resolution."""
def https_open(self, request):
return self.do_open(LanConnection, request, context=ssl.create_default_context())
class NoRedirect(HTTPRedirectHandler):
"""Refuse redirects before any request body or token could be forwarded."""
def redirect_request(self, req, fp, code, msg, headers, newurl):
return None
class BatchClient:
"""Expose exact-model batch inference; never retry or substitute implicitly."""
def __init__(self, token_file, cache_dir):
token_path = Path(token_file)
if token_path.stat().st_mode & 0o077:
raise ValueError("token file must be private (chmod 600)")
self.token = token_path.read_text().strip()
if len(self.token) != 64 or any(c not in "0123456789abcdef" for c in self.token):
raise ValueError("invalid API credential")
self.cache_dir = Path(cache_dir)
self.cache_dir.mkdir(mode=0o700, parents=True, exist_ok=True)
self.cache_dir.chmod(0o700)
self.http = build_opener(ProxyHandler({}), LanHandler(), NoRedirect())
def _request(self, suffix, payload=None):
data = None if payload is None else canonical(payload)
request = Request(BASE + suffix, data=data, headers={
"Authorization": "Bearer " + self.token, "Content-Type": "application/json"})
with self.http.open(request, timeout=1810 if data else 15) as response:
raw = response.read(4194305)
if len(raw) > 4194304:
raise ValueError("response too large")
return json.loads(raw)
def catalog(self):
"""Confirm the actual local runtime and immutable model manifests."""
result = self._request("/models")
actual = {item["model"]: item["digest"] for item in result["models"]}
if result.get("runtime") != RUNTIME or actual != MODELS or result.get("fallback") is not None:
raise ValueError("approved runtime/model configuration mismatch")
return result
def run(self, envelope, validator=None):
"""Cache a FA01/synthetic request by source, prompt, schema and model identity.
The roster application must filter source rows and permitted fields before
constructing this envelope. This client never reads an Excel workbook.
"""
required = {"campaign_id", "suite_id", "source_sha256", "prompt_version", "request"}
if set(envelope) != required or envelope["campaign_id"] not in ("FA01", "SYNTHETIC"):
raise ValueError("a FA01 or SYNTHETIC scope envelope is required")
if not all(isinstance(envelope[k], str) and envelope[k] for k in required - {"request"}):
raise ValueError("scope and version fields must be nonempty strings")
source = envelope["source_sha256"]
if len(source) != 64 or any(c not in "0123456789abcdef" for c in source):
raise ValueError("source_sha256 must identify the filtered source content")
payload = envelope["request"]
model = payload.get("model")
if model not in MODELS:
raise ValueError("exact approved model required")
identity = {"envelope": envelope, "model_digest": MODELS[model], "runtime": RUNTIME,
"client_version": 1, "endpoint": BASE, "connection_address": ADDRESS}
key = hashlib.sha256(canonical(identity)).hexdigest()
destination = self.cache_dir / (key + ".json")
if destination.exists():
record = json.loads(destination.read_text())
self._verify(record["result"], payload)
if record.get("cache_key") != key:
raise ValueError("cache identity mismatch")
self._validate_output(destination, record, validator)
return {"cache_hit": True, "path": str(destination), "record": record}
self.catalog()
started = time.monotonic()
result = self._request("/generate", payload)
self._verify(result, payload)
record = {"cache_key": key, "source_sha256": source,
"prompt_version": envelope["prompt_version"],
"request_sha256": hashlib.sha256(canonical(payload)).hexdigest(),
"campaign_id": envelope["campaign_id"], "suite_id": envelope["suite_id"],
"client_wall_seconds": round(time.monotonic() - started, 3), "result": result}
self._validate_output(destination, record, validator)
self._save(destination, record)
return {"cache_hit": False, "path": str(destination), "record": record}
def _validate_output(self, destination, record, validator):
"""Keep failed application contracts as diagnostics, outside reusable cache."""
if validator is None:
return
try:
validator(json.loads(record["result"]["response"]))
except (ValueError, KeyError, TypeError):
self._save(destination.with_suffix(".rejected.json"), record)
destination.unlink(missing_ok=True)
raise ValueError("application validation failed; local diagnostic retained") from None
@staticmethod
def _verify(result, payload):
"""Reject changed models, settings and incomplete results before caching."""
provenance = result.get("batch_provenance", {})
expected_options = {**payload["options"], "num_thread": 16, "num_gpu": 0}
if (result.get("model") != payload["model"] or not result.get("done")
or result.get("done_reason") == "length"
or provenance.get("model_digest") != MODELS[payload["model"]]
or provenance.get("runtime") != RUNTIME
or provenance.get("options") != expected_options
or provenance.get("think") != payload["think"]):
raise ValueError("incomplete or substituted model result")
json.loads(result["response"])
@staticmethod
def _save(destination, record):
"""Atomically publish mode-600 results so interruption cannot poison cache."""
descriptor, temporary = tempfile.mkstemp(dir=destination.parent, prefix=".pending-")
try:
with os.fdopen(descriptor, "wb") as handle:
handle.write(canonical(record))
handle.flush()
os.fsync(handle.fileno())
os.replace(temporary, destination)
finally:
Path(temporary).unlink(missing_ok=True)
def main():
"""Print only transport metadata; model results stay in the local cache."""
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--token-file", required=True)
parser.add_argument("--cache-dir", required=True)
parser.add_argument("--request", help="JSON scope envelope, prepared on the laptop")
arguments = parser.parse_args()
try:
client = BatchClient(arguments.token_file, arguments.cache_dir)
if not arguments.request:
print(json.dumps(client.catalog(), indent=2))
return
outcome = client.run(json.loads(Path(arguments.request).read_text()))
print(json.dumps({"cache_hit": outcome["cache_hit"], "result_file": outcome["path"],
"client_wall_seconds": outcome["record"]["client_wall_seconds"]}))
except (OSError, HTTPError, URLError, ValueError, KeyError, TypeError):
parser.exit(1, "Local batch request failed; no fallback. Check API health, scope and configuration.\n")
if __name__ == "__main__":
main()