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.
@@ -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'")