Spaces:
Running
Running
Download pull_worker.py from vllm-sr/decision-studio: direct link, hf CLI and curl.
- Browser
- Download file 7.73 kB
-
https://huggingface.co/spaces/vllm-sr/decision-studio/resolve/main/pull_worker.py
- Command line
-
hf download hf://spaces/vllm-sr/decision-studio/pull_worker.py
-
curl -L -o pull_worker.py https://huggingface.co/spaces/vllm-sr/decision-studio/resolve/main/pull_worker.py
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 | |