continuum-task-server-sdk 1.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- continuum_task_server/__init__.py +74 -0
- continuum_task_server/_http.py +142 -0
- continuum_task_server/_stomp.py +137 -0
- continuum_task_server/client.py +351 -0
- continuum_task_server/exceptions.py +94 -0
- continuum_task_server/models.py +197 -0
- continuum_task_server/server.py +511 -0
- continuum_task_server/websocket.py +852 -0
- continuum_task_server_sdk-1.1.0.dist-info/METADATA +258 -0
- continuum_task_server_sdk-1.1.0.dist-info/RECORD +11 -0
- continuum_task_server_sdk-1.1.0.dist-info/WHEEL +4 -0
|
@@ -0,0 +1,511 @@
|
|
|
1
|
+
"""TaskServer - worker-loop abstraction for handling Continuum queue items.
|
|
2
|
+
|
|
3
|
+
Register handlers per task-type name with ``@server.task("name")`` and call
|
|
4
|
+
``server.run()``. The server polls for OPEN items, claims them, heartbeats while
|
|
5
|
+
the handler runs, then reports ENDED (return value) or KILLED (exception).
|
|
6
|
+
|
|
7
|
+
Use ``@server.task(..., auto_complete=False)`` when the handler records work
|
|
8
|
+
elsewhere and you will call ``complete_queue_item`` / ``fail_queue_item`` later.
|
|
9
|
+
The SDK does **not** heartbeat after the handler returns; you must call
|
|
10
|
+
``client.queue.heartbeat`` yourself (for example on each pass of your DB poll) so
|
|
11
|
+
the claim stays alive across process restarts.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
import logging
|
|
17
|
+
import os
|
|
18
|
+
import signal
|
|
19
|
+
import threading
|
|
20
|
+
import time
|
|
21
|
+
import traceback
|
|
22
|
+
import uuid
|
|
23
|
+
from collections.abc import Callable
|
|
24
|
+
from concurrent.futures import Future, ThreadPoolExecutor
|
|
25
|
+
from dataclasses import dataclass, field
|
|
26
|
+
from typing import Any
|
|
27
|
+
|
|
28
|
+
import httpx
|
|
29
|
+
|
|
30
|
+
from .client import ContinuumClient
|
|
31
|
+
from .exceptions import ContinuumError
|
|
32
|
+
from .models import QueueEvent, QueueEventType, QueueItem, TaskStatus, TransportMode
|
|
33
|
+
from .websocket import ContinuumWebSocketClient, WebSocketOptions
|
|
34
|
+
|
|
35
|
+
logger = logging.getLogger("continuum_task_server")
|
|
36
|
+
|
|
37
|
+
Handler = Callable[[QueueItem], Any]
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
@dataclass
|
|
41
|
+
class _Registration:
|
|
42
|
+
name: str
|
|
43
|
+
handler: Handler
|
|
44
|
+
concurrency: int
|
|
45
|
+
auto_complete: bool
|
|
46
|
+
semaphore: threading.Semaphore = field(init=False)
|
|
47
|
+
|
|
48
|
+
def __post_init__(self) -> None:
|
|
49
|
+
self.semaphore = threading.Semaphore(self.concurrency)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
class TaskServer:
|
|
53
|
+
"""Worker-first task server.
|
|
54
|
+
|
|
55
|
+
Example:
|
|
56
|
+
>>> server = TaskServer(base_url="http://localhost:8080", api_key="...")
|
|
57
|
+
>>> @server.task("echo")
|
|
58
|
+
... def echo(item):
|
|
59
|
+
... return {"echoed": item.input_data_json}
|
|
60
|
+
>>> server.run() # blocks until SIGINT/SIGTERM
|
|
61
|
+
|
|
62
|
+
With ``auto_complete=True`` (default), handler return values become ``outputData``
|
|
63
|
+
and the task is marked ENDED. Raising any exception marks the task KILLED.
|
|
64
|
+
|
|
65
|
+
With ``auto_complete=False``, the handler returns without ENDED; you are
|
|
66
|
+
responsible for heartbeats until you call ``complete_queue_item`` or
|
|
67
|
+
``fail_queue_item`` (see README). Heartbeats only run while the handler is
|
|
68
|
+
executing, not after it returns.
|
|
69
|
+
"""
|
|
70
|
+
|
|
71
|
+
def __init__(
|
|
72
|
+
self,
|
|
73
|
+
base_url: str,
|
|
74
|
+
api_key: str,
|
|
75
|
+
*,
|
|
76
|
+
max_workers: int = 4,
|
|
77
|
+
poll_interval: float = 1.0,
|
|
78
|
+
max_poll_interval: float = 5.0,
|
|
79
|
+
heartbeat_interval: float = 15.0,
|
|
80
|
+
shutdown_timeout: float = 30.0,
|
|
81
|
+
client: ContinuumClient | None = None,
|
|
82
|
+
connect_timeout: float = 10.0,
|
|
83
|
+
request_timeout: float = 30.0,
|
|
84
|
+
transport: TransportMode | str | httpx.BaseTransport | None = None,
|
|
85
|
+
websocket_options: WebSocketOptions | None = None,
|
|
86
|
+
websocket_client: Any = None,
|
|
87
|
+
) -> None:
|
|
88
|
+
http_transport: httpx.BaseTransport | None = None
|
|
89
|
+
if isinstance(transport, httpx.BaseTransport):
|
|
90
|
+
http_transport = transport
|
|
91
|
+
self._mode = TransportMode.HTTP
|
|
92
|
+
else:
|
|
93
|
+
self._mode = _coerce_transport(transport)
|
|
94
|
+
self._client = client or ContinuumClient(
|
|
95
|
+
base_url=base_url,
|
|
96
|
+
api_key=api_key,
|
|
97
|
+
connect_timeout=connect_timeout,
|
|
98
|
+
request_timeout=request_timeout,
|
|
99
|
+
transport=http_transport,
|
|
100
|
+
)
|
|
101
|
+
self._owns_client = client is None
|
|
102
|
+
self._max_workers = max_workers
|
|
103
|
+
self._poll_interval = poll_interval
|
|
104
|
+
self._max_poll_interval = max_poll_interval
|
|
105
|
+
self._heartbeat_interval = heartbeat_interval
|
|
106
|
+
self._shutdown_timeout = shutdown_timeout
|
|
107
|
+
self._handlers: dict[str, _Registration] = {}
|
|
108
|
+
self._stop_event = threading.Event()
|
|
109
|
+
self._executor: ThreadPoolExecutor | None = None
|
|
110
|
+
self._inflight: set[Future[Any]] = set()
|
|
111
|
+
self._inflight_lock = threading.Lock()
|
|
112
|
+
self._ws = websocket_client
|
|
113
|
+
self._ws_options = websocket_options or WebSocketOptions(
|
|
114
|
+
connect_timeout=connect_timeout,
|
|
115
|
+
request_timeout=request_timeout,
|
|
116
|
+
)
|
|
117
|
+
self._owns_ws = websocket_client is None
|
|
118
|
+
self._hint_lock = threading.Lock()
|
|
119
|
+
self._pending_hints: set[str] = set()
|
|
120
|
+
self._claim_loop_lock = threading.Lock()
|
|
121
|
+
self._base_url = base_url
|
|
122
|
+
self._api_key = api_key
|
|
123
|
+
|
|
124
|
+
@property
|
|
125
|
+
def client(self) -> ContinuumClient:
|
|
126
|
+
"""Underlying ContinuumClient. Useful for management operations from handlers."""
|
|
127
|
+
return self._client
|
|
128
|
+
|
|
129
|
+
def task(
|
|
130
|
+
self,
|
|
131
|
+
name: str,
|
|
132
|
+
*,
|
|
133
|
+
concurrency: int = 1,
|
|
134
|
+
auto_complete: bool = True,
|
|
135
|
+
) -> Callable[[Handler], Handler]:
|
|
136
|
+
"""Decorator registering ``handler`` for queue items of task-type ``name``.
|
|
137
|
+
|
|
138
|
+
``concurrency`` is the max number of in-flight items for this task name.
|
|
139
|
+
|
|
140
|
+
If ``auto_complete`` is False, the handler returns without the server sending
|
|
141
|
+
ENDED; call ``complete_queue_item`` or ``fail_queue_item`` when done. You must
|
|
142
|
+
heartbeat the queue item yourself until then (see README).
|
|
143
|
+
"""
|
|
144
|
+
if concurrency < 1:
|
|
145
|
+
raise ValueError("concurrency must be >= 1")
|
|
146
|
+
|
|
147
|
+
def decorator(handler: Handler) -> Handler:
|
|
148
|
+
if name in self._handlers:
|
|
149
|
+
raise ValueError(f"a handler for task {name!r} is already registered")
|
|
150
|
+
self._handlers[name] = _Registration(
|
|
151
|
+
name=name,
|
|
152
|
+
handler=handler,
|
|
153
|
+
concurrency=concurrency,
|
|
154
|
+
auto_complete=auto_complete,
|
|
155
|
+
)
|
|
156
|
+
return handler
|
|
157
|
+
|
|
158
|
+
return decorator
|
|
159
|
+
|
|
160
|
+
def register(
|
|
161
|
+
self,
|
|
162
|
+
name: str,
|
|
163
|
+
handler: Handler,
|
|
164
|
+
*,
|
|
165
|
+
concurrency: int = 1,
|
|
166
|
+
auto_complete: bool = True,
|
|
167
|
+
) -> None:
|
|
168
|
+
"""Imperative alternative to ``@task``."""
|
|
169
|
+
self.task(name, concurrency=concurrency, auto_complete=auto_complete)(handler)
|
|
170
|
+
|
|
171
|
+
def complete_queue_item(
|
|
172
|
+
self,
|
|
173
|
+
queue_item_id: uuid.UUID | str,
|
|
174
|
+
*,
|
|
175
|
+
output_data: Any = None,
|
|
176
|
+
) -> None:
|
|
177
|
+
"""Mark a queue item ENDED.
|
|
178
|
+
|
|
179
|
+
Use after ``auto_complete=False`` when work finished successfully. Uses the
|
|
180
|
+
same ``ContinuumClient`` (and API key) as this ``TaskServer``.
|
|
181
|
+
"""
|
|
182
|
+
qid = self._normalize_queue_id(queue_item_id)
|
|
183
|
+
self._safe_update_status_by_id(qid, TaskStatus.ENDED, output=output_data)
|
|
184
|
+
|
|
185
|
+
def fail_queue_item(
|
|
186
|
+
self,
|
|
187
|
+
queue_item_id: uuid.UUID | str,
|
|
188
|
+
*,
|
|
189
|
+
error: BaseException | str | dict[str, Any] | None = None,
|
|
190
|
+
) -> None:
|
|
191
|
+
"""Mark a queue item KILLED."""
|
|
192
|
+
qid = self._normalize_queue_id(queue_item_id)
|
|
193
|
+
payload: dict[str, Any]
|
|
194
|
+
if isinstance(error, dict):
|
|
195
|
+
payload = dict(error)
|
|
196
|
+
elif isinstance(error, str):
|
|
197
|
+
payload = {"error": error}
|
|
198
|
+
elif isinstance(error, BaseException):
|
|
199
|
+
payload = {
|
|
200
|
+
"error": str(error),
|
|
201
|
+
"type": type(error).__name__,
|
|
202
|
+
"traceback": "".join(
|
|
203
|
+
traceback.format_exception(type(error), error, error.__traceback__)
|
|
204
|
+
),
|
|
205
|
+
}
|
|
206
|
+
else:
|
|
207
|
+
payload = {"error": "failed"}
|
|
208
|
+
self._safe_update_status_by_id(qid, TaskStatus.KILLED, output=payload)
|
|
209
|
+
|
|
210
|
+
@staticmethod
|
|
211
|
+
def _normalize_queue_id(queue_item_id: uuid.UUID | str) -> uuid.UUID:
|
|
212
|
+
if isinstance(queue_item_id, uuid.UUID):
|
|
213
|
+
return queue_item_id
|
|
214
|
+
return uuid.UUID(str(queue_item_id))
|
|
215
|
+
|
|
216
|
+
def run(self, *, install_signal_handlers: bool = True) -> None:
|
|
217
|
+
"""Block and poll until ``stop()`` is called or a shutdown signal arrives."""
|
|
218
|
+
if not self._handlers:
|
|
219
|
+
raise RuntimeError("no task handlers registered; use @server.task(...) first")
|
|
220
|
+
|
|
221
|
+
if install_signal_handlers:
|
|
222
|
+
self._install_signal_handlers()
|
|
223
|
+
|
|
224
|
+
self._stop_event.clear()
|
|
225
|
+
self._executor = ThreadPoolExecutor(max_workers=self._max_workers)
|
|
226
|
+
|
|
227
|
+
logger.info(
|
|
228
|
+
"Continuum task server started: tasks=%s, max_workers=%d, transport=%s",
|
|
229
|
+
list(self._handlers),
|
|
230
|
+
self._max_workers,
|
|
231
|
+
self._mode.value,
|
|
232
|
+
)
|
|
233
|
+
|
|
234
|
+
try:
|
|
235
|
+
if self._mode is TransportMode.WEBSOCKET:
|
|
236
|
+
self._websocket_loop()
|
|
237
|
+
else:
|
|
238
|
+
self._poll_loop()
|
|
239
|
+
finally:
|
|
240
|
+
self._drain()
|
|
241
|
+
|
|
242
|
+
def stop(self) -> None:
|
|
243
|
+
"""Signal the server to stop polling and drain in-flight handlers."""
|
|
244
|
+
if not self._stop_event.is_set():
|
|
245
|
+
logger.info("Continuum task server shutdown requested")
|
|
246
|
+
self._stop_event.set()
|
|
247
|
+
|
|
248
|
+
def _install_signal_handlers(self) -> None:
|
|
249
|
+
def _handler(signum: int, _frame: object) -> None:
|
|
250
|
+
logger.info("Received signal %d, stopping", signum)
|
|
251
|
+
self.stop()
|
|
252
|
+
|
|
253
|
+
try:
|
|
254
|
+
signal.signal(signal.SIGINT, _handler)
|
|
255
|
+
signal.signal(signal.SIGTERM, _handler)
|
|
256
|
+
except ValueError:
|
|
257
|
+
logger.debug("signal handlers not installed (not on main thread)")
|
|
258
|
+
|
|
259
|
+
def _poll_loop(self) -> None:
|
|
260
|
+
names = list(self._handlers)
|
|
261
|
+
index = 0
|
|
262
|
+
backoff = self._poll_interval
|
|
263
|
+
|
|
264
|
+
while not self._stop_event.is_set():
|
|
265
|
+
registration = self._handlers[names[index % len(names)]]
|
|
266
|
+
index += 1
|
|
267
|
+
|
|
268
|
+
if not registration.semaphore.acquire(blocking=False):
|
|
269
|
+
if index % len(names) == 0:
|
|
270
|
+
self._stop_event.wait(timeout=self._poll_interval)
|
|
271
|
+
continue
|
|
272
|
+
|
|
273
|
+
try:
|
|
274
|
+
item = self._client.queue.claim(registration.name)
|
|
275
|
+
except ContinuumError as e:
|
|
276
|
+
registration.semaphore.release()
|
|
277
|
+
logger.error("claim(%s) failed: %s", registration.name, e)
|
|
278
|
+
self._stop_event.wait(timeout=min(backoff, self._max_poll_interval))
|
|
279
|
+
backoff = min(backoff * 2, self._max_poll_interval)
|
|
280
|
+
continue
|
|
281
|
+
|
|
282
|
+
if item is None:
|
|
283
|
+
registration.semaphore.release()
|
|
284
|
+
if index % len(names) == 0:
|
|
285
|
+
self._stop_event.wait(timeout=backoff)
|
|
286
|
+
backoff = min(backoff * 2, self._max_poll_interval)
|
|
287
|
+
continue
|
|
288
|
+
|
|
289
|
+
backoff = self._poll_interval
|
|
290
|
+
self._dispatch(registration, item)
|
|
291
|
+
|
|
292
|
+
def _ensure_websocket(self) -> ContinuumWebSocketClient:
|
|
293
|
+
if self._ws is None:
|
|
294
|
+
self._ws = ContinuumWebSocketClient(
|
|
295
|
+
self._base_url,
|
|
296
|
+
self._api_key,
|
|
297
|
+
options=self._ws_options,
|
|
298
|
+
)
|
|
299
|
+
self._owns_ws = True
|
|
300
|
+
return self._ws
|
|
301
|
+
|
|
302
|
+
def _websocket_loop(self) -> None:
|
|
303
|
+
ws = self._ensure_websocket()
|
|
304
|
+
ws.connect()
|
|
305
|
+
for name in self._handlers:
|
|
306
|
+
ws.subscribe_work(name)
|
|
307
|
+
logger.info("WebSocket subscribed to work: %s", list(self._handlers))
|
|
308
|
+
while not self._stop_event.is_set():
|
|
309
|
+
event = ws.next_event(timeout=0.25)
|
|
310
|
+
if event is None:
|
|
311
|
+
continue
|
|
312
|
+
self._handle_queue_event(event)
|
|
313
|
+
|
|
314
|
+
def _handle_queue_event(self, event: QueueEvent) -> None:
|
|
315
|
+
if event.event_type is not QueueEventType.WORK_AVAILABLE:
|
|
316
|
+
return
|
|
317
|
+
name = event.task_name
|
|
318
|
+
if name is None or name not in self._handlers:
|
|
319
|
+
return
|
|
320
|
+
with self._hint_lock:
|
|
321
|
+
self._pending_hints.add(name)
|
|
322
|
+
self._try_claim_pending()
|
|
323
|
+
|
|
324
|
+
def _try_claim_pending(self) -> None:
|
|
325
|
+
while not self._stop_event.is_set():
|
|
326
|
+
if not self._claim_loop_lock.acquire(blocking=False):
|
|
327
|
+
return
|
|
328
|
+
try:
|
|
329
|
+
self._drain_claims()
|
|
330
|
+
finally:
|
|
331
|
+
self._claim_loop_lock.release()
|
|
332
|
+
if not self._has_pending_with_capacity():
|
|
333
|
+
return
|
|
334
|
+
|
|
335
|
+
def _has_pending_with_capacity(self) -> bool:
|
|
336
|
+
with self._hint_lock:
|
|
337
|
+
names = list(self._pending_hints)
|
|
338
|
+
for name in names:
|
|
339
|
+
registration = self._handlers.get(name)
|
|
340
|
+
if registration is None:
|
|
341
|
+
continue
|
|
342
|
+
if registration.semaphore.acquire(blocking=False):
|
|
343
|
+
registration.semaphore.release()
|
|
344
|
+
return True
|
|
345
|
+
return False
|
|
346
|
+
|
|
347
|
+
def _drain_claims(self) -> None:
|
|
348
|
+
ws = self._ws
|
|
349
|
+
if ws is None:
|
|
350
|
+
return
|
|
351
|
+
while not self._stop_event.is_set():
|
|
352
|
+
registration = self._pop_pending_with_capacity()
|
|
353
|
+
if registration is None:
|
|
354
|
+
return
|
|
355
|
+
try:
|
|
356
|
+
item = ws.claim(registration.name)
|
|
357
|
+
except ContinuumError as e:
|
|
358
|
+
registration.semaphore.release()
|
|
359
|
+
logger.error("claim(%s) failed: %s", registration.name, e)
|
|
360
|
+
with self._hint_lock:
|
|
361
|
+
self._pending_hints.add(registration.name)
|
|
362
|
+
return
|
|
363
|
+
if item is None:
|
|
364
|
+
registration.semaphore.release()
|
|
365
|
+
continue
|
|
366
|
+
with self._hint_lock:
|
|
367
|
+
self._pending_hints.add(registration.name)
|
|
368
|
+
self._dispatch(registration, item)
|
|
369
|
+
|
|
370
|
+
def _pop_pending_with_capacity(self) -> _Registration | None:
|
|
371
|
+
with self._hint_lock:
|
|
372
|
+
names = list(self._pending_hints)
|
|
373
|
+
for name in names:
|
|
374
|
+
registration = self._handlers.get(name)
|
|
375
|
+
if registration is None:
|
|
376
|
+
with self._hint_lock:
|
|
377
|
+
self._pending_hints.discard(name)
|
|
378
|
+
continue
|
|
379
|
+
if not registration.semaphore.acquire(blocking=False):
|
|
380
|
+
continue
|
|
381
|
+
with self._hint_lock:
|
|
382
|
+
self._pending_hints.discard(name)
|
|
383
|
+
return registration
|
|
384
|
+
return None
|
|
385
|
+
|
|
386
|
+
def _dispatch(self, registration: _Registration, item: QueueItem) -> None:
|
|
387
|
+
assert self._executor is not None
|
|
388
|
+
future = self._executor.submit(self._run_item, registration, item)
|
|
389
|
+
with self._inflight_lock:
|
|
390
|
+
self._inflight.add(future)
|
|
391
|
+
future.add_done_callback(self._on_done)
|
|
392
|
+
|
|
393
|
+
def _on_done(self, future: Future[Any]) -> None:
|
|
394
|
+
with self._inflight_lock:
|
|
395
|
+
self._inflight.discard(future)
|
|
396
|
+
|
|
397
|
+
def _run_item(self, registration: _Registration, item: QueueItem) -> None:
|
|
398
|
+
try:
|
|
399
|
+
logger.info("Claimed queue item %s (task=%s)", item.id, registration.name)
|
|
400
|
+
stop_heartbeat = threading.Event()
|
|
401
|
+
heartbeat_thread = threading.Thread(
|
|
402
|
+
target=self._heartbeat_loop,
|
|
403
|
+
args=(item, stop_heartbeat),
|
|
404
|
+
name=f"hb-{item.id}",
|
|
405
|
+
daemon=True,
|
|
406
|
+
)
|
|
407
|
+
heartbeat_thread.start()
|
|
408
|
+
|
|
409
|
+
try:
|
|
410
|
+
self._safe_update_status(item, TaskStatus.STARTED)
|
|
411
|
+
try:
|
|
412
|
+
result = registration.handler(item)
|
|
413
|
+
except Exception as e:
|
|
414
|
+
logger.exception("Handler raised for queue item %s", item.id)
|
|
415
|
+
error_payload = {
|
|
416
|
+
"error": str(e),
|
|
417
|
+
"type": type(e).__name__,
|
|
418
|
+
"traceback": traceback.format_exc(),
|
|
419
|
+
}
|
|
420
|
+
self._safe_update_status(item, TaskStatus.KILLED, output=error_payload)
|
|
421
|
+
return
|
|
422
|
+
|
|
423
|
+
if registration.auto_complete:
|
|
424
|
+
self._safe_update_status(item, TaskStatus.ENDED, output=result)
|
|
425
|
+
logger.info("Completed queue item %s", item.id)
|
|
426
|
+
else:
|
|
427
|
+
logger.info(
|
|
428
|
+
"Handler returned for queue item %s without ENDED "
|
|
429
|
+
"(auto_complete=False); caller must heartbeat and complete",
|
|
430
|
+
item.id,
|
|
431
|
+
)
|
|
432
|
+
finally:
|
|
433
|
+
stop_heartbeat.set()
|
|
434
|
+
heartbeat_thread.join(timeout=1.0)
|
|
435
|
+
finally:
|
|
436
|
+
registration.semaphore.release()
|
|
437
|
+
if self._mode is TransportMode.WEBSOCKET:
|
|
438
|
+
self._try_claim_pending()
|
|
439
|
+
|
|
440
|
+
def _heartbeat_loop(self, item: QueueItem, stop_event: threading.Event) -> None:
|
|
441
|
+
while not stop_event.wait(timeout=self._heartbeat_interval):
|
|
442
|
+
try:
|
|
443
|
+
if self._ws is not None:
|
|
444
|
+
self._ws.heartbeat(item.id)
|
|
445
|
+
else:
|
|
446
|
+
self._client.queue.heartbeat(item.id)
|
|
447
|
+
except ContinuumError as e:
|
|
448
|
+
logger.warning("heartbeat for %s failed: %s", item.id, e)
|
|
449
|
+
|
|
450
|
+
def _safe_update_status(
|
|
451
|
+
self,
|
|
452
|
+
item: QueueItem,
|
|
453
|
+
status: TaskStatus,
|
|
454
|
+
*,
|
|
455
|
+
output: Any = None,
|
|
456
|
+
) -> None:
|
|
457
|
+
self._safe_update_status_by_id(item.id, status, output=output)
|
|
458
|
+
|
|
459
|
+
def _safe_update_status_by_id(
|
|
460
|
+
self,
|
|
461
|
+
queue_item_id: uuid.UUID,
|
|
462
|
+
status: TaskStatus,
|
|
463
|
+
*,
|
|
464
|
+
output: Any = None,
|
|
465
|
+
) -> None:
|
|
466
|
+
try:
|
|
467
|
+
if self._ws is not None:
|
|
468
|
+
self._ws.update_status(queue_item_id, status, output_data=output)
|
|
469
|
+
else:
|
|
470
|
+
self._client.queue.update_status(queue_item_id, status, output_data=output)
|
|
471
|
+
except ContinuumError as e:
|
|
472
|
+
logger.error("update_status(%s, %s) failed: %s", queue_item_id, status.value, e)
|
|
473
|
+
|
|
474
|
+
def _drain(self) -> None:
|
|
475
|
+
if self._executor is None:
|
|
476
|
+
return
|
|
477
|
+
logger.info("Draining in-flight tasks (timeout=%.1fs)", self._shutdown_timeout)
|
|
478
|
+
deadline = time.monotonic() + self._shutdown_timeout
|
|
479
|
+
with self._inflight_lock:
|
|
480
|
+
futures = list(self._inflight)
|
|
481
|
+
for future in futures:
|
|
482
|
+
remaining = max(0.0, deadline - time.monotonic())
|
|
483
|
+
try:
|
|
484
|
+
future.result(timeout=remaining)
|
|
485
|
+
except Exception:
|
|
486
|
+
logger.debug("drain: handler future raised", exc_info=True)
|
|
487
|
+
self._executor.shutdown(wait=False, cancel_futures=True)
|
|
488
|
+
self._executor = None
|
|
489
|
+
if self._ws is not None:
|
|
490
|
+
self._ws.close()
|
|
491
|
+
if self._owns_ws:
|
|
492
|
+
self._ws = None
|
|
493
|
+
if self._owns_client:
|
|
494
|
+
self._client.close()
|
|
495
|
+
logger.info("Continuum task server stopped")
|
|
496
|
+
|
|
497
|
+
|
|
498
|
+
def _coerce_transport(value: TransportMode | str | None) -> TransportMode:
|
|
499
|
+
if value is None:
|
|
500
|
+
env = os.environ.get("CONTINUUM_TRANSPORT")
|
|
501
|
+
if env is None or not str(env).strip():
|
|
502
|
+
return TransportMode.HTTP
|
|
503
|
+
value = env
|
|
504
|
+
if isinstance(value, TransportMode):
|
|
505
|
+
return value
|
|
506
|
+
normalized = str(value).strip().lower()
|
|
507
|
+
if normalized in {"http", "https", "poll"}:
|
|
508
|
+
return TransportMode.HTTP
|
|
509
|
+
if normalized in {"websocket", "ws", "wss", "stomp"}:
|
|
510
|
+
return TransportMode.WEBSOCKET
|
|
511
|
+
raise ValueError(f"unknown transport {value!r}; expected 'http' or 'websocket'")
|