continuum-task-server-sdk 1.1.0__tar.gz → 1.2.0__tar.gz

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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: continuum-task-server-sdk
3
- Version: 1.1.0
3
+ Version: 1.2.0
4
4
  Summary: Python SDK for the Continuum Task Server
5
5
  Project-URL: Homepage, https://github.com/ContinuumWorkflow/continuum-task-server-sdk-python
6
6
  Project-URL: Issues, https://github.com/ContinuumWorkflow/continuum-task-server-sdk-python/issues
@@ -3,6 +3,7 @@
3
3
  from __future__ import annotations
4
4
 
5
5
  from .client import ContinuumClient
6
+ from .context import TaskExecutionContext
6
7
  from .exceptions import (
7
8
  AmbiguousCommandError,
8
9
  BadRequestError,
@@ -24,6 +25,7 @@ from .models import (
24
25
  QueueEvent,
25
26
  QueueEventType,
26
27
  QueueItem,
28
+ QueueItemTrigger,
27
29
  TaskItem,
28
30
  TaskItemVersion,
29
31
  TaskStatus,
@@ -54,8 +56,10 @@ __all__ = [
54
56
  "QueueEvent",
55
57
  "QueueEventType",
56
58
  "QueueItem",
59
+ "QueueItemTrigger",
57
60
  "RateLimitError",
58
61
  "ServerError",
62
+ "TaskExecutionContext",
59
63
  "TaskItem",
60
64
  "TaskItemVersion",
61
65
  "TaskServer",
@@ -11,6 +11,7 @@ from uuid import UUID
11
11
  import httpx
12
12
 
13
13
  from ._http import HttpClient, drop_none
14
+ from .exceptions import ContinuumError
14
15
  from .models import Content, QueueItem, TaskItem, TaskItemVersion, TaskStatus, TaskType
15
16
 
16
17
 
@@ -264,9 +265,15 @@ class QueueApi:
264
265
  return None
265
266
  return QueueItem.model_validate(result)
266
267
 
267
- def heartbeat(self, queue_item_id: UUID | str) -> QueueItem:
268
+ def heartbeat(
269
+ self,
270
+ queue_item_id: UUID | str,
271
+ *,
272
+ claim_token: UUID | str | None = None,
273
+ ) -> QueueItem:
274
+ body = drop_none({"claimToken": str(claim_token) if claim_token else None})
268
275
  return QueueItem.model_validate(
269
- self._http.post_json(f"{self._QUEUE}/queue-items/{queue_item_id}/heartbeat", {})
276
+ self._http.post_json(f"{self._QUEUE}/queue-items/{queue_item_id}/heartbeat", body)
270
277
  )
271
278
 
272
279
  def update_status(
@@ -275,16 +282,50 @@ class QueueApi:
275
282
  status: TaskStatus,
276
283
  *,
277
284
  output_data: Any = None,
285
+ claim_token: UUID | str | None = None,
278
286
  ) -> QueueItem:
279
287
  # Server expects `status` on this endpoint (not `taskStatus` on queue item JSON).
280
288
  body: dict[str, Any] = {"status": status.value}
281
289
  encoded = _encode_input(output_data)
282
290
  if encoded is not None:
283
291
  body["outputData"] = encoded
292
+ if claim_token is not None:
293
+ body["claimToken"] = str(claim_token)
284
294
  return QueueItem.model_validate(
285
295
  self._http.post_json(f"{self._QUEUE}/queue-items/{queue_item_id}/status", body)
286
296
  )
287
297
 
298
+ def cancel(self, queue_item_id: UUID | str) -> QueueItem:
299
+ result = self._http.post_json(f"{self._MGMT}/{queue_item_id}/cancel", {})
300
+ status = 0 if result is None else int(result.get("status", 0))
301
+ if result is None or status != 0:
302
+ message = None if result is None else result.get("message")
303
+ raise ContinuumError(message or "cancel failed")
304
+ return QueueItem.model_validate(result["item"])
305
+
306
+ def complete_control(
307
+ self,
308
+ queue_item_id: UUID | str,
309
+ *,
310
+ claim_token: UUID | str | None = None,
311
+ output_data: Any = None,
312
+ ) -> QueueItem:
313
+ body = drop_none(
314
+ {
315
+ "claimToken": str(claim_token) if claim_token else None,
316
+ "outputData": _encode_input(output_data),
317
+ }
318
+ )
319
+ result = self._http.post_json(
320
+ f"{self._QUEUE}/queue-items/{queue_item_id}/control-complete",
321
+ body,
322
+ )
323
+ status = 0 if result is None else int(result.get("status", 0))
324
+ if result is None or status != 0:
325
+ message = None if result is None else result.get("message")
326
+ raise ContinuumError(message or "control-complete failed")
327
+ return QueueItem.model_validate(result["item"])
328
+
288
329
 
289
330
  class ContentStoreApi:
290
331
  _PATH = "/api/management/content-store"
@@ -0,0 +1,79 @@
1
+ """Per-claim execution context with an idempotent cancellation callback."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import threading
6
+ from collections.abc import Callable
7
+ from uuid import UUID
8
+
9
+ from .models import QueueItem
10
+
11
+ _callback_lock = threading.Lock()
12
+ _callback_pool: list[threading.Thread] = []
13
+
14
+
15
+ class TaskExecutionContext:
16
+ """Cancellation token/callback for one claimed queue item."""
17
+
18
+ def __init__(self, item: QueueItem) -> None:
19
+ self._item = item
20
+ self._cancelled = threading.Event()
21
+ self._callback: Callable[[], None] | None = None
22
+ self._callback_started = threading.Event()
23
+ self._cleanup_done = threading.Event()
24
+ self._lock = threading.Lock()
25
+
26
+ @property
27
+ def item(self) -> QueueItem:
28
+ return self._item
29
+
30
+ @property
31
+ def claim_token(self) -> UUID | None:
32
+ return self._item.claim_token
33
+
34
+ def is_cancelled(self) -> bool:
35
+ return self._cancelled.is_set()
36
+
37
+ def on_cancelled(self, handler: Callable[[], None] | None) -> None:
38
+ if handler is None:
39
+ return
40
+ with self._lock:
41
+ if self._callback is not None:
42
+ return
43
+ self._callback = handler
44
+ already = self._cancelled.is_set()
45
+ if already:
46
+ self._start_callback()
47
+
48
+ def request_cancel(self) -> None:
49
+ with self._lock:
50
+ if self._cancelled.is_set():
51
+ return
52
+ self._cancelled.set()
53
+ has_callback = self._callback is not None
54
+ if has_callback:
55
+ self._start_callback()
56
+
57
+ def await_cleanup(self, timeout: float) -> None:
58
+ if not self._cancelled.is_set() or self._callback is None:
59
+ return
60
+ self._cleanup_done.wait(timeout=max(0.001, timeout))
61
+
62
+ def _start_callback(self) -> None:
63
+ with self._lock:
64
+ if self._callback_started.is_set():
65
+ return
66
+ handler = self._callback
67
+ if handler is None:
68
+ self._cleanup_done.set()
69
+ return
70
+ self._callback_started.set()
71
+
72
+ def _run() -> None:
73
+ try:
74
+ handler()
75
+ finally:
76
+ self._cleanup_done.set()
77
+
78
+ thread = threading.Thread(target=_run, name="continuum-cancel-cb", daemon=True)
79
+ thread.start()
@@ -36,6 +36,8 @@ class TaskType(_Base):
36
36
  is_global: bool = Field(default=False, alias="isGlobal")
37
37
  max_duration_sec: int | None = Field(default=None, alias="maxDurationSec")
38
38
  heartbeat_timeout_sec: int | None = Field(default=None, alias="heartbeatTimeoutSec")
39
+ claim_start_timeout_sec: int | None = Field(default=None, alias="claimStartTimeoutSec")
40
+ control_grace_sec: int | None = Field(default=None, alias="controlGraceSec")
39
41
 
40
42
 
41
43
  class TaskItem(_Base):
@@ -69,6 +71,18 @@ class TaskItemVersion(_Base):
69
71
  return self.item_definition
70
72
 
71
73
 
74
+ class QueueItemTrigger(_Base):
75
+ type: str | None = None
76
+ source: str | None = None
77
+ schedule_id: UUID | None = Field(default=None, alias="scheduleId")
78
+ schedule_name: str | None = Field(default=None, alias="scheduleName")
79
+ occurrence_id: UUID | None = Field(default=None, alias="occurrenceId")
80
+ scheduled_for: datetime | None = Field(default=None, alias="scheduledFor")
81
+ enqueued_at: datetime | None = Field(default=None, alias="enqueuedAt")
82
+ time_zone: str | None = Field(default=None, alias="timeZone")
83
+ misfire: bool | None = None
84
+
85
+
72
86
  class QueueItem(_Base):
73
87
  id: UUID
74
88
  depth: int = 0
@@ -81,6 +95,14 @@ class QueueItem(_Base):
81
95
  created_by: UUID | None = Field(default=None, alias="createdBy")
82
96
  input_data: str | None = Field(default=None, alias="inputData")
83
97
  output_data: str | None = Field(default=None, alias="outputData")
98
+ claim_token: UUID | None = Field(default=None, alias="claimToken")
99
+ claim_attempt: int | None = Field(default=None, alias="claimAttempt")
100
+ control_request: str | None = Field(default=None, alias="controlRequest")
101
+ control_reason: str | None = Field(default=None, alias="controlReason")
102
+ heartbeat_dtm: datetime | None = Field(default=None, alias="heartbeatDtm")
103
+ start_dtm: datetime | None = Field(default=None, alias="startDtm")
104
+ end_dtm: datetime | None = Field(default=None, alias="endDtm")
105
+ trigger: QueueItemTrigger | None = None
84
106
 
85
107
  @property
86
108
  def input_data_json(self) -> Any:
@@ -132,6 +154,7 @@ class QueueEventType(str, Enum):
132
154
 
133
155
  WORK_AVAILABLE = "work.available"
134
156
  COMPLETED = "queue.completed"
157
+ CONTROL_REQUESTED = "queue.control.requested"
135
158
 
136
159
 
137
160
  class WaitTargetStatus(_Base):
@@ -157,6 +180,8 @@ class QueueEvent(_Base):
157
180
  correlation: str | None = None
158
181
  wait_mode: WaitMode | None = Field(default=None, alias="waitMode")
159
182
  targets: list[WaitTargetStatus] | None = None
183
+ control_request: str | None = Field(default=None, alias="controlRequest")
184
+ claim_token: UUID | None = Field(default=None, alias="claimToken")
160
185
 
161
186
 
162
187
  class WorkSubscriptionResult(_Base):
@@ -13,6 +13,7 @@ the claim stays alive across process restarts.
13
13
 
14
14
  from __future__ import annotations
15
15
 
16
+ import inspect
16
17
  import logging
17
18
  import os
18
19
  import signal
@@ -28,13 +29,14 @@ from typing import Any
28
29
  import httpx
29
30
 
30
31
  from .client import ContinuumClient
32
+ from .context import TaskExecutionContext
31
33
  from .exceptions import ContinuumError
32
34
  from .models import QueueEvent, QueueEventType, QueueItem, TaskStatus, TransportMode
33
35
  from .websocket import ContinuumWebSocketClient, WebSocketOptions
34
36
 
35
37
  logger = logging.getLogger("continuum_task_server")
36
38
 
37
- Handler = Callable[[QueueItem], Any]
39
+ Handler = Callable[..., Any]
38
40
 
39
41
 
40
42
  @dataclass
@@ -62,10 +64,10 @@ class TaskServer:
62
64
  With ``auto_complete=True`` (default), handler return values become ``outputData``
63
65
  and the task is marked ENDED. Raising any exception marks the task KILLED.
64
66
 
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.
67
+ With ``auto_complete=False``, the handler returns without ENDED. Automatic
68
+ heartbeats stop when the handler returns. Heartbeat from the completing
69
+ process with ``client.queue.heartbeat`` until ``complete_queue_item`` or
70
+ ``fail_queue_item``.
69
71
  """
70
72
 
71
73
  def __init__(
@@ -120,6 +122,7 @@ class TaskServer:
120
122
  self._claim_loop_lock = threading.Lock()
121
123
  self._base_url = base_url
122
124
  self._api_key = api_key
125
+ self._active: dict[uuid.UUID, _ActiveClaim] = {}
123
126
 
124
127
  @property
125
128
  def client(self) -> ContinuumClient:
@@ -180,7 +183,14 @@ class TaskServer:
180
183
  same ``ContinuumClient`` (and API key) as this ``TaskServer``.
181
184
  """
182
185
  qid = self._normalize_queue_id(queue_item_id)
183
- self._safe_update_status_by_id(qid, TaskStatus.ENDED, output=output_data)
186
+ claim = self._active.get(qid)
187
+ if claim is not None and claim.control_seen.is_set():
188
+ claim.acknowledge_control()
189
+ return
190
+ token = None if claim is None else claim.claim_token
191
+ self._safe_update_status_by_id(qid, TaskStatus.ENDED, output=output_data, claim_token=token)
192
+ if claim is not None:
193
+ claim.mark_completed()
184
194
 
185
195
  def fail_queue_item(
186
196
  self,
@@ -190,6 +200,10 @@ class TaskServer:
190
200
  ) -> None:
191
201
  """Mark a queue item KILLED."""
192
202
  qid = self._normalize_queue_id(queue_item_id)
203
+ claim = self._active.get(qid)
204
+ if claim is not None and claim.control_seen.is_set():
205
+ claim.acknowledge_control()
206
+ return
193
207
  payload: dict[str, Any]
194
208
  if isinstance(error, dict):
195
209
  payload = dict(error)
@@ -205,7 +219,10 @@ class TaskServer:
205
219
  }
206
220
  else:
207
221
  payload = {"error": "failed"}
208
- self._safe_update_status_by_id(qid, TaskStatus.KILLED, output=payload)
222
+ token = None if claim is None else claim.claim_token
223
+ self._safe_update_status_by_id(qid, TaskStatus.KILLED, output=payload, claim_token=token)
224
+ if claim is not None:
225
+ claim.mark_completed()
209
226
 
210
227
  @staticmethod
211
228
  def _normalize_queue_id(queue_item_id: uuid.UUID | str) -> uuid.UUID:
@@ -244,6 +261,8 @@ class TaskServer:
244
261
  if not self._stop_event.is_set():
245
262
  logger.info("Continuum task server shutdown requested")
246
263
  self._stop_event.set()
264
+ for claim in list(self._active.values()):
265
+ claim.deferred_gate.set()
247
266
 
248
267
  def _install_signal_handlers(self) -> None:
249
268
  def _handler(signum: int, _frame: object) -> None:
@@ -312,6 +331,9 @@ class TaskServer:
312
331
  self._handle_queue_event(event)
313
332
 
314
333
  def _handle_queue_event(self, event: QueueEvent) -> None:
334
+ if event.event_type is QueueEventType.CONTROL_REQUESTED:
335
+ self._apply_control_event(event)
336
+ return
315
337
  if event.event_type is not QueueEventType.WORK_AVAILABLE:
316
338
  return
317
339
  name = event.task_name
@@ -321,6 +343,17 @@ class TaskServer:
321
343
  self._pending_hints.add(name)
322
344
  self._try_claim_pending()
323
345
 
346
+ def _apply_control_event(self, event: QueueEvent) -> None:
347
+ if event.queue_item_id is None:
348
+ return
349
+ claim = self._active.get(event.queue_item_id)
350
+ if claim is None:
351
+ return
352
+ token = event.claim_token
353
+ if token is not None and claim.claim_token is not None and token != claim.claim_token:
354
+ return
355
+ claim.signal_control()
356
+
324
357
  def _try_claim_pending(self) -> None:
325
358
  while not self._stop_event.is_set():
326
359
  if not self._claim_loop_lock.acquire(blocking=False):
@@ -395,57 +428,132 @@ class TaskServer:
395
428
  self._inflight.discard(future)
396
429
 
397
430
  def _run_item(self, registration: _Registration, item: QueueItem) -> None:
431
+ claim = _ActiveClaim(item=item, server=self)
432
+ self._active[item.id] = claim
398
433
  try:
399
434
  logger.info("Claimed queue item %s (task=%s)", item.id, registration.name)
400
- stop_heartbeat = threading.Event()
401
435
  heartbeat_thread = threading.Thread(
402
436
  target=self._heartbeat_loop,
403
- args=(item, stop_heartbeat),
437
+ args=(claim,),
404
438
  name=f"hb-{item.id}",
405
439
  daemon=True,
406
440
  )
407
441
  heartbeat_thread.start()
408
442
 
409
443
  try:
410
- self._safe_update_status(item, TaskStatus.STARTED)
444
+ self._subscribe_control(claim)
445
+ self._inspect_control(
446
+ self._safe_update_status(
447
+ item, TaskStatus.STARTED, claim_token=claim.claim_token
448
+ )
449
+ )
411
450
  try:
412
- result = registration.handler(item)
451
+ result = self._invoke_handler(registration.handler, item, claim.context)
413
452
  except Exception as e:
453
+ if claim.control_seen.is_set() or claim.context.is_cancelled():
454
+ claim.acknowledge_control()
455
+ return
414
456
  logger.exception("Handler raised for queue item %s", item.id)
415
457
  error_payload = {
416
458
  "error": str(e),
417
459
  "type": type(e).__name__,
418
460
  "traceback": traceback.format_exc(),
419
461
  }
420
- self._safe_update_status(item, TaskStatus.KILLED, output=error_payload)
462
+ self._safe_update_status(
463
+ item, TaskStatus.KILLED, output=error_payload, claim_token=claim.claim_token
464
+ )
465
+ claim.mark_completed()
466
+ return
467
+
468
+ if claim.control_seen.is_set() or claim.context.is_cancelled():
469
+ claim.acknowledge_control()
421
470
  return
422
471
 
423
472
  if registration.auto_complete:
424
- self._safe_update_status(item, TaskStatus.ENDED, output=result)
473
+ self._safe_update_status(
474
+ item, TaskStatus.ENDED, output=result, claim_token=claim.claim_token
475
+ )
425
476
  logger.info("Completed queue item %s", item.id)
477
+ claim.mark_completed()
426
478
  else:
427
479
  logger.info(
428
480
  "Handler returned for queue item %s without ENDED "
429
- "(auto_complete=False); caller must heartbeat and complete",
481
+ "(auto_complete=False); remaining subscribed until complete",
430
482
  item.id,
431
483
  )
484
+ self._stop_heartbeat(claim, heartbeat_thread)
485
+ self._wait_for_deferred(claim)
432
486
  finally:
433
- stop_heartbeat.set()
434
- heartbeat_thread.join(timeout=1.0)
487
+ self._stop_heartbeat(claim, heartbeat_thread)
488
+ self._unsubscribe_control(claim)
435
489
  finally:
490
+ self._active.pop(item.id, None)
436
491
  registration.semaphore.release()
437
492
  if self._mode is TransportMode.WEBSOCKET:
438
493
  self._try_claim_pending()
439
494
 
440
- def _heartbeat_loop(self, item: QueueItem, stop_event: threading.Event) -> None:
441
- while not stop_event.wait(timeout=self._heartbeat_interval):
495
+ def _wait_for_deferred(self, claim: _ActiveClaim) -> None:
496
+ while (
497
+ not self._stop_event.is_set()
498
+ and not claim.completed.is_set()
499
+ and not claim.control_seen.is_set()
500
+ ):
501
+ if claim.deferred_gate.wait(timeout=0.1):
502
+ break
503
+ if claim.control_seen.is_set() or claim.context.is_cancelled():
504
+ claim.acknowledge_control()
505
+
506
+ def _stop_heartbeat(self, claim: _ActiveClaim, heartbeat_thread: threading.Thread) -> None:
507
+ claim.stop_heartbeat.set()
508
+ if heartbeat_thread is threading.current_thread():
509
+ return
510
+ heartbeat_thread.join(timeout=1.0)
511
+
512
+ def _invoke_handler(
513
+ self, handler: Handler, item: QueueItem, context: TaskExecutionContext
514
+ ) -> Any:
515
+ if _accepts_context(handler):
516
+ return handler(item, context)
517
+ return handler(item)
518
+
519
+ def _heartbeat_loop(self, claim: _ActiveClaim) -> None:
520
+ while not claim.stop_heartbeat.wait(timeout=self._heartbeat_interval):
442
521
  try:
522
+ latest: QueueItem | None
443
523
  if self._ws is not None:
444
- self._ws.heartbeat(item.id)
524
+ latest = self._ws.heartbeat(claim.item.id, claim_token=claim.claim_token)
445
525
  else:
446
- self._client.queue.heartbeat(item.id)
526
+ latest = self._client.queue.heartbeat(
527
+ claim.item.id, claim_token=claim.claim_token
528
+ )
529
+ self._inspect_control(latest)
447
530
  except ContinuumError as e:
448
- logger.warning("heartbeat for %s failed: %s", item.id, e)
531
+ logger.warning("heartbeat for %s failed: %s", claim.item.id, e)
532
+
533
+ def _inspect_control(self, latest: QueueItem | None) -> None:
534
+ if latest is None or latest.control_request in (None, ""):
535
+ return
536
+ claim = self._active.get(latest.id)
537
+ if claim is None:
538
+ return
539
+ claim.signal_control()
540
+
541
+ def _subscribe_control(self, claim: _ActiveClaim) -> None:
542
+ if self._ws is None:
543
+ return
544
+ try:
545
+ self._ws.subscribe_control(claim.item.id, claim_token=claim.claim_token)
546
+ claim.control_subscribed = True
547
+ except ContinuumError as e:
548
+ logger.warning("subscribeControl(%s) failed: %s", claim.item.id, e)
549
+
550
+ def _unsubscribe_control(self, claim: _ActiveClaim) -> None:
551
+ if self._ws is None or not claim.control_subscribed:
552
+ return
553
+ try:
554
+ self._ws.unsubscribe_control(claim.item.id, claim_token=claim.claim_token)
555
+ except ContinuumError as e:
556
+ logger.debug("unsubscribeControl(%s) failed: %s", claim.item.id, e)
449
557
 
450
558
  def _safe_update_status(
451
559
  self,
@@ -453,8 +561,11 @@ class TaskServer:
453
561
  status: TaskStatus,
454
562
  *,
455
563
  output: Any = None,
456
- ) -> None:
457
- self._safe_update_status_by_id(item.id, status, output=output)
564
+ claim_token: uuid.UUID | None = None,
565
+ ) -> QueueItem | None:
566
+ return self._safe_update_status_by_id(
567
+ item.id, status, output=output, claim_token=claim_token
568
+ )
458
569
 
459
570
  def _safe_update_status_by_id(
460
571
  self,
@@ -462,14 +573,23 @@ class TaskServer:
462
573
  status: TaskStatus,
463
574
  *,
464
575
  output: Any = None,
465
- ) -> None:
576
+ claim_token: uuid.UUID | None = None,
577
+ ) -> QueueItem | None:
466
578
  try:
579
+ latest: QueueItem | None
467
580
  if self._ws is not None:
468
- self._ws.update_status(queue_item_id, status, output_data=output)
581
+ latest = self._ws.update_status(
582
+ queue_item_id, status, output_data=output, claim_token=claim_token
583
+ )
469
584
  else:
470
- self._client.queue.update_status(queue_item_id, status, output_data=output)
585
+ latest = self._client.queue.update_status(
586
+ queue_item_id, status, output_data=output, claim_token=claim_token
587
+ )
588
+ self._inspect_control(latest)
589
+ return latest
471
590
  except ContinuumError as e:
472
591
  logger.error("update_status(%s, %s) failed: %s", queue_item_id, status.value, e)
592
+ return None
473
593
 
474
594
  def _drain(self) -> None:
475
595
  if self._executor is None:
@@ -495,6 +615,72 @@ class TaskServer:
495
615
  logger.info("Continuum task server stopped")
496
616
 
497
617
 
618
+ class _ActiveClaim:
619
+ def __init__(self, *, item: QueueItem, server: TaskServer) -> None:
620
+ self.item = item
621
+ self.server = server
622
+ self.context = TaskExecutionContext(item)
623
+ self.stop_heartbeat = threading.Event()
624
+ self.control_seen = threading.Event()
625
+ self.completed = threading.Event()
626
+ self.deferred_gate = threading.Event()
627
+ self.control_subscribed = False
628
+ self._control_acked = threading.Event()
629
+
630
+ @property
631
+ def claim_token(self) -> uuid.UUID | None:
632
+ return self.item.claim_token
633
+
634
+ def signal_control(self) -> None:
635
+ if self.control_seen.is_set():
636
+ return
637
+ self.control_seen.set()
638
+ self.context.request_cancel()
639
+ self.deferred_gate.set()
640
+
641
+ def mark_completed(self) -> None:
642
+ self.completed.set()
643
+ self.deferred_gate.set()
644
+
645
+ def acknowledge_control(self) -> None:
646
+ if self._control_acked.is_set():
647
+ return
648
+ self._control_acked.set()
649
+ self.context.request_cancel()
650
+ self.context.await_cleanup(max(0.001, self.server._shutdown_timeout))
651
+ try:
652
+ if self.server._ws is not None:
653
+ self.server._ws.complete_control(self.item.id, claim_token=self.claim_token)
654
+ else:
655
+ self.server._client.queue.complete_control(
656
+ self.item.id, claim_token=self.claim_token
657
+ )
658
+ logger.info("Acknowledged control for queue item %s", self.item.id)
659
+ except ContinuumError as e:
660
+ logger.warning("completeControl(%s) failed: %s", self.item.id, e)
661
+ self.mark_completed()
662
+
663
+
664
+ def _accepts_context(handler: Handler) -> bool:
665
+ try:
666
+ signature = inspect.signature(handler)
667
+ except (TypeError, ValueError):
668
+ return False
669
+ positional = [
670
+ param
671
+ for param in signature.parameters.values()
672
+ if param.kind
673
+ in (
674
+ inspect.Parameter.POSITIONAL_ONLY,
675
+ inspect.Parameter.POSITIONAL_OR_KEYWORD,
676
+ inspect.Parameter.VAR_POSITIONAL,
677
+ )
678
+ ]
679
+ if any(param.kind == inspect.Parameter.VAR_POSITIONAL for param in positional):
680
+ return True
681
+ return len(positional) >= 2
682
+
683
+
498
684
  def _coerce_transport(value: TransportMode | str | None) -> TransportMode:
499
685
  if value is None:
500
686
  env = os.environ.get("CONTINUUM_TRANSPORT")
@@ -127,6 +127,7 @@ class ContinuumWebSocketClient:
127
127
  self._events_cv = threading.Condition()
128
128
  self._event_callbacks: list[EventCallback] = []
129
129
  self._work_subscriptions: set[str] = set()
130
+ self._control_subscriptions: set[tuple[str, str | None]] = set()
130
131
  self._wait_subscribers: set[str] = set()
131
132
  self._seen_completed: OrderedDict[str, None] = OrderedDict()
132
133
  self._connected = threading.Event()
@@ -299,10 +300,21 @@ class ContinuumWebSocketClient:
299
300
  return None
300
301
  return QueueItem.model_validate(item)
301
302
 
302
- def heartbeat(self, queue_item_id: UUID | str, *, request_id: str | None = None) -> QueueItem:
303
+ def heartbeat(
304
+ self,
305
+ queue_item_id: UUID | str,
306
+ *,
307
+ claim_token: UUID | str | None = None,
308
+ request_id: str | None = None,
309
+ ) -> QueueItem:
303
310
  payload = self._command(
304
311
  "/app/heartbeat",
305
- {"queueItemId": str(queue_item_id)},
312
+ drop_none(
313
+ {
314
+ "queueItemId": str(queue_item_id),
315
+ "claimToken": str(claim_token) if claim_token else None,
316
+ }
317
+ ),
306
318
  request_id=request_id,
307
319
  retry=True,
308
320
  )
@@ -314,6 +326,7 @@ class ContinuumWebSocketClient:
314
326
  status: TaskStatus,
315
327
  *,
316
328
  output_data: Any = None,
329
+ claim_token: UUID | str | None = None,
317
330
  request_id: str | None = None,
318
331
  ) -> QueueItem:
319
332
  body: dict[str, Any] = {
@@ -323,6 +336,8 @@ class ContinuumWebSocketClient:
323
336
  encoded = _encode_input(output_data)
324
337
  if encoded is not None:
325
338
  body["outputData"] = encoded
339
+ if claim_token is not None:
340
+ body["claimToken"] = str(claim_token)
326
341
  payload = self._command("/app/status", body, request_id=request_id, retry=True)
327
342
  return QueueItem.model_validate(payload["item"])
328
343
 
@@ -352,6 +367,74 @@ class ContinuumWebSocketClient:
352
367
  self._work_subscriptions.discard(task_name)
353
368
  return WorkSubscriptionResult.model_validate(payload)
354
369
 
370
+ def subscribe_control(
371
+ self,
372
+ queue_item_id: UUID | str,
373
+ *,
374
+ claim_token: UUID | str | None = None,
375
+ request_id: str | None = None,
376
+ ) -> WorkSubscriptionResult:
377
+ payload = self._command(
378
+ "/app/subscribeControl",
379
+ drop_none(
380
+ {
381
+ "queueItemId": str(queue_item_id),
382
+ "claimToken": str(claim_token) if claim_token else None,
383
+ }
384
+ ),
385
+ request_id=request_id,
386
+ retry=True,
387
+ )
388
+ with self._state_lock:
389
+ self._control_subscriptions.add(
390
+ (str(queue_item_id), str(claim_token) if claim_token else None)
391
+ )
392
+ return WorkSubscriptionResult.model_validate(payload)
393
+
394
+ def unsubscribe_control(
395
+ self,
396
+ queue_item_id: UUID | str,
397
+ *,
398
+ claim_token: UUID | str | None = None,
399
+ request_id: str | None = None,
400
+ ) -> WorkSubscriptionResult:
401
+ payload = self._command(
402
+ "/app/unsubscribeControl",
403
+ drop_none(
404
+ {
405
+ "queueItemId": str(queue_item_id),
406
+ "claimToken": str(claim_token) if claim_token else None,
407
+ }
408
+ ),
409
+ request_id=request_id,
410
+ retry=True,
411
+ )
412
+ with self._state_lock:
413
+ self._control_subscriptions.discard(
414
+ (str(queue_item_id), str(claim_token) if claim_token else None)
415
+ )
416
+ return WorkSubscriptionResult.model_validate(payload)
417
+
418
+ def complete_control(
419
+ self,
420
+ queue_item_id: UUID | str,
421
+ *,
422
+ claim_token: UUID | str | None = None,
423
+ request_id: str | None = None,
424
+ ) -> QueueItem:
425
+ payload = self._command(
426
+ "/app/completeControl",
427
+ drop_none(
428
+ {
429
+ "queueItemId": str(queue_item_id),
430
+ "claimToken": str(claim_token) if claim_token else None,
431
+ }
432
+ ),
433
+ request_id=request_id,
434
+ retry=True,
435
+ )
436
+ return QueueItem.model_validate(payload["item"])
437
+
355
438
  def subscribe_wait(
356
439
  self,
357
440
  subscriber_id: str,
@@ -597,6 +680,7 @@ class ContinuumWebSocketClient:
597
680
  with self._state_lock:
598
681
  work = list(self._work_subscriptions)
599
682
  waits = list(self._wait_subscribers)
683
+ controls = list(self._control_subscriptions)
600
684
  for task_name in work:
601
685
  try:
602
686
  self._command_inline("/app/subscribeWork", {"taskName": task_name})
@@ -607,6 +691,21 @@ class ContinuumWebSocketClient:
607
691
  self._command_inline("/app/subscribeWait", {"subscriberId": subscriber_id})
608
692
  except ContinuumError:
609
693
  logger.error("Failed to restore wait subscriber %s", subscriber_id, exc_info=True)
694
+ for queue_item_id, claim_token in controls:
695
+ try:
696
+ self._command_inline(
697
+ "/app/subscribeControl",
698
+ drop_none(
699
+ {
700
+ "queueItemId": queue_item_id,
701
+ "claimToken": claim_token,
702
+ }
703
+ ),
704
+ )
705
+ except ContinuumError:
706
+ logger.error(
707
+ "Failed to restore control subscription %s", queue_item_id, exc_info=True
708
+ )
610
709
 
611
710
  def _command_inline(self, destination: str, body: dict[str, Any]) -> dict[str, Any]:
612
711
  rid = str(uuid.uuid4())
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
4
4
 
5
5
  [project]
6
6
  name = "continuum-task-server-sdk"
7
- version = "1.1.0"
7
+ version = "1.2.0"
8
8
  description = "Python SDK for the Continuum Task Server"
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.10"