File size: 7,157 Bytes
8827155
 
 
 
 
 
 
 
 
 
 
 
 
 
c646299
8827155
 
 
c646299
 
 
 
 
 
8827155
 
 
 
 
c646299
 
8827155
 
c646299
 
8827155
 
 
 
 
 
 
 
59ba696
8827155
 
c646299
8827155
 
 
 
 
 
c646299
8827155
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c646299
 
8827155
 
 
 
82cafc8
 
8827155
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
"""Outbound queue worker for a pinned private Decision HTTP runtime."""

import asyncio
import concurrent.futures
import json
import logging
import math
import os
import signal
import threading

from contract import to_records
from direct_gateway import DirectGateway, DirectGatewayError
from model_registry import PROFILES, model_registry
from model_runtime_client import ModelRuntimeGateway
from pull_worker import Gateway, Worker

LOG = logging.getLogger("decision.http_worker")
# `decision_serve_v1`: `vllm-sr decision serve` (artifact headers, native batches).
# `vllm_sr_runtime_v1`: the built-in model runtime (`vllm-sr-runtime serve`).
RUNTIME_APIS = {
    "decision_serve_v1": DirectGateway,
    "vllm_sr_runtime_v1": ModelRuntimeGateway,
}


class HTTPRuntime:
    """Run one attested private client on its own event loop for probe and inference."""

    def __init__(self, model, origin, artifact, *, timeout_seconds=40, client=None,
                 api="decision_serve_v1"):
        if model not in PROFILES or artifact.get("repo_id") != PROFILES[model]["repo_id"]:
            raise ValueError("Select one exact Decision model and canonical artifact")
        if api not in RUNTIME_APIS:
            raise ValueError("Select a supported Decision runtime API")
        if (
            isinstance(timeout_seconds, bool)
            or not isinstance(timeout_seconds, (int, float))
            or not math.isfinite(timeout_seconds)
            or not 0.1 <= timeout_seconds <= 90
        ):
            raise ValueError("Runtime timeout must be between 0.1 and 90 seconds")
        self.model = model
        self.canonical_model = PROFILES[model].get("runtime_model", artifact["repo_id"])
        self.manifest = artifact["manifest_sha256"]
        self.phase = "not_loaded"
        self._gateway = RUNTIME_APIS[api](
            {self.canonical_model: origin},
            expected_artifacts={self.canonical_model: {
                "revision": artifact["revision"],
                "manifest_sha256": self.manifest,
                **({"content_sha256": artifact["content_sha256"]}
                   if "content_sha256" in artifact else {}),
                "confidence_definition": PROFILES[model]["confidence"],
            }},
            timeout_seconds=timeout_seconds,
            client=client,
        )
        self._loop = None
        self._thread = None

    def _submit(self, coroutine, timeout):
        if self._loop is None:
            coroutine.close()
            raise RuntimeError("Runtime client is not started")
        future = asyncio.run_coroutine_threadsafe(coroutine, self._loop)
        try:
            return future.result(timeout=timeout)
        except concurrent.futures.TimeoutError:
            future.cancel()
            raise TimeoutError("Runtime call exceeded its deadline") from None

    def _load(self):
        if self._loop is not None:
            return
        self._loop = asyncio.new_event_loop()

        def serve():
            asyncio.set_event_loop(self._loop)
            self._loop.run_forever()

        self._thread = threading.Thread(target=serve, name="decision-runtime-http", daemon=True)
        self._thread.start()
        if not self.ready():
            raise RuntimeError("Selected Decision runtime is not ready with its pinned artifact")
        self.phase = "ready"

    def ready(self):
        observed = self._submit(self._gateway.probe(self.canonical_model), 5)
        return observed["loaded"] is True

    def evaluate(self, body, records):
        if body.get("model") != self.model or records != to_records(body, model=self.model):
            raise ValueError("Request does not match this worker")
        request = dict(body, model=self.canonical_model)
        self.phase = "running"
        try:
            try:
                result = self._submit(
                    self._gateway.evaluate(request, batch="states" in body),
                    self._gateway.timeout_seconds + 5,
                )
            except DirectGatewayError as exc:
                if exc.code == 413:
                    raise ValueError("Selected model rejected the complete input") from None
                raise
            return {"kind": "http_runtime_v1", "response": result}
        finally:
            self.phase = "ready"

    def close(self):
        if self._loop is None:
            return
        try:
            self._submit(self._gateway.aclose(), 5)
        finally:
            self._loop.call_soon_threadsafe(self._loop.stop)
            self._thread.join(timeout=5)
            if self._thread.is_alive():
                raise RuntimeError("Runtime client did not stop")
            self._loop.close()
            self._loop = None
            self._thread = None


class HTTPWorker(Worker):
    heartbeat_join_timeout = 38

    def heartbeat(self):
        if self.runtime.phase != "running" and not self.runtime.ready():
            raise RuntimeError("Selected Decision runtime is unavailable")
        return super().heartbeat()

    def run(self):
        try:
            super().run()
        finally:
            self.runtime.close()


def runtime_from_environment():
    raw_registry = os.getenv("DECISION_MODEL_REGISTRY_V2", "")
    if not raw_registry:
        raise ValueError("DECISION_MODEL_REGISTRY_V2 is required")
    try:
        registry = model_registry(json.loads(raw_registry))
    except (TypeError, ValueError, KeyError, RecursionError) as exc:
        raise ValueError("Invalid Decision model registry") from exc
    model = os.getenv("DECISION_WORKER_MODEL", "")
    if model not in registry:
        raise ValueError("DECISION_WORKER_MODEL must be an exact configured wire ID")
    item = registry[model]
    if "revision" not in item:
        raise ValueError("The selected model requires a pinned Hub revision")
    try:
        timeout = float(os.getenv("DECISION_RUNTIME_TIMEOUT_SECONDS", "40"))
    except ValueError as exc:
        raise ValueError("DECISION_RUNTIME_TIMEOUT_SECONDS must be numeric") from exc
    return HTTPRuntime(model, os.getenv("DECISION_RUNTIME_URL", ""), item,
                       timeout_seconds=timeout,
                       api=os.getenv("DECISION_RUNTIME_API", "decision_serve_v1"))


def main():
    logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
    # Every heartbeat probes the runtime; keep per-request transport lines out of the log.
    logging.getLogger("httpx").setLevel(logging.WARNING)
    runtime = runtime_from_environment()
    gateway = Gateway(os.getenv("DECISION_STUDIO_URL", ""),
                      os.getenv("DECISION_WORKER_TOKEN", ""),
                      allow_loopback=os.getenv("DECISION_ALLOW_HTTP_LOOPBACK") == "1")
    worker = HTTPWorker(gateway, runtime)
    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:  # noqa: BLE001 - never print secret-bearing exception text
        LOG.error("HTTP worker stopped: %s", type(exc).__name__)
        raise SystemExit(1) from None