decision-studio / pull_worker.py
Xunzhuo's picture
Add outbound HTTP pull workers for Decision runtimes
8827155
Raw History Blame Contribute Delete
7.73 kB
"""Resident queue worker; only outbound HTTPS. Never retries inference for a job."""
import json
import logging
import os
import secrets
import signal
import threading
import time
import urllib.error
import urllib.parse
import urllib.request
from contract import to_records
from engine import Engine
LOG = logging.getLogger("decision.worker")
class NoRedirect(urllib.request.HTTPRedirectHandler):
def redirect_request(self, req, fp, code, msg, headers, newurl):
# Do not forward the worker bearer or user input to a redirected origin.
return None
class GatewayError(Exception):
def __init__(self, code):
self.code = code
class Gateway:
def __init__(self, url, token, *, allow_loopback=False):
parsed = urllib.parse.urlsplit(url)
loopback = allow_loopback and parsed.hostname in {"localhost", "127.0.0.1", "::1"}
if (parsed.scheme != "https" and not (loopback and parsed.scheme == "http")
or not parsed.hostname or parsed.username or parsed.password or parsed.query or parsed.fragment
or parsed.path not in {"", "/"}):
raise ValueError("Use a fixed HTTPS gateway origin; no credentials, path, query, or fragment")
if not isinstance(token, str) or len(token) < 32:
raise ValueError("A worker token of at least 32 characters is required")
self.url, self.token = url.rstrip("/"), token
handlers = [NoRedirect()]
if loopback:
# A developer's system proxy must not receive loopback test credentials.
handlers.append(urllib.request.ProxyHandler({}))
self.opener = urllib.request.build_opener(*handlers)
def post(self, action, payload):
body = json.dumps(payload, ensure_ascii=False, separators=(",", ":"), allow_nan=False).encode()
request = urllib.request.Request(self.url + "/internal/worker/" + action, data=body, method="POST",
headers={"Content-Type": "application/json", "Authorization": "Bearer " + self.token})
try:
with self.opener.open(request, timeout=30) as response:
limit = 2 * 1024 * 1024 + 64 * 1024 if action == "claim" else 64 * 1024
data = response.read(limit + 1)
if len(data) > limit:
raise ValueError("Gateway response is too large")
return json.loads(data)
except urllib.error.HTTPError as exc:
# Do not log HTTP body, URL, or request headers.
raise GatewayError(exc.code) from None
class Worker:
heartbeat_join_timeout = 2
def __init__(self, gateway, runtime=None):
self.gateway, self.runtime = gateway, runtime or Engine()
self.worker_id, self.stop = secrets.token_hex(16), threading.Event()
def post(self, action, payload):
payload = dict(payload, model=self.runtime.model, manifest_sha256=self.runtime.manifest)
return self.gateway.post(action, payload)
def heartbeat(self):
phase = "running" if self.runtime.phase == "running" else "ready"
return self.post("heartbeat", {"worker_id": self.worker_id,
"manifest_sha256": self.runtime.manifest, "phase": phase,
"capabilities": ["context_batch_v1"]})
def heartbeat_loop(self):
while not self.stop.wait(10):
try:
self.heartbeat()
except GatewayError as exc:
LOG.warning("Heartbeat HTTP status: %d", exc.code)
if exc.code in {401, 403, 409}:
self.stop.set()
except Exception as exc:
LOG.warning("Heartbeat unavailable: %s", type(exc).__name__)
def execute(self, job):
"""One evaluate, followed only by bounded identical result retransmissions."""
started = time.monotonic()
if not isinstance(job, dict) or set(job) != {"id", "lease_token", "body", "seconds_remaining"}:
raise ValueError("Invalid gateway job")
seconds = job["seconds_remaining"]
if type(seconds) not in (int, float) or not 0 < seconds <= 180:
raise ValueError("Invalid lease duration")
reply = {"worker_id": self.worker_id, "id": job["id"], "lease_token": job["lease_token"]}
try:
records = to_records(job["body"], model=self.runtime.model)
reply["result"] = self.runtime.evaluate(job["body"], records)
except ValueError:
reply["error_code"] = "input_rejected"
except Exception as exc:
LOG.warning("Inference failed: %s", type(exc).__name__)
reply["error_code"] = "inference_failed"
# The same reply can be acknowledged twice; the model is never called twice.
for attempt in range(4):
if time.monotonic() - started >= seconds or self.stop.is_set():
LOG.warning("Result discarded after lease deadline or worker stop")
return
try:
self.post("result", reply)
LOG.info("Job finished: %s", "prediction" if "result" in reply else reply["error_code"])
return
except GatewayError as exc:
if exc.code in {409, 410}:
LOG.warning("Result no longer accepted: HTTP %d", exc.code)
return
if exc.code in {401, 403, 422}:
self.stop.set()
LOG.error("Worker result rejected: HTTP %d", exc.code)
return
LOG.warning("Result delivery HTTP status: %d", exc.code)
except Exception as exc:
LOG.warning("Result delivery unavailable: %s", type(exc).__name__)
self.stop.wait(2 ** attempt)
LOG.warning("Result delivery unconfirmed; no inference retry")
def run(self):
# A runtime must validate its artifact before advertising readiness.
self.runtime._load()
self.heartbeat()
heartbeat = threading.Thread(target=self.heartbeat_loop, daemon=True)
heartbeat.start()
LOG.info("Resident worker ready")
try:
while not self.stop.is_set():
try:
response = self.post("claim", {"worker_id": self.worker_id, "wait_seconds": 20})
if response.get("job") is not None:
self.execute(response["job"])
except GatewayError as exc:
if exc.code in {401, 403, 422}:
raise
# An uncertain claim may remain leased. Let it expire; do not rerun.
LOG.warning("Claim unavailable: HTTP %d", exc.code)
self.stop.wait(3)
except (OSError, ValueError) as exc:
LOG.warning("Claim unavailable: %s", type(exc).__name__)
self.stop.wait(3)
finally:
self.stop.set()
heartbeat.join(timeout=self.heartbeat_join_timeout)
def main():
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
gateway = Gateway(os.getenv("DECISION_STUDIO_URL", ""), os.getenv("DECISION_WORKER_TOKEN", ""),
allow_loopback=os.getenv("DECISION_ALLOW_HTTP_LOOPBACK") == "1")
worker = Worker(gateway)
for sig in (signal.SIGINT, signal.SIGTERM):
signal.signal(sig, lambda *_: worker.stop.set())
worker.run()
if __name__ == "__main__":
try:
main()
except Exception as exc:
LOG.error("Worker stopped: %s", type(exc).__name__)
raise SystemExit(1) from None