continuum-task-server-sdk 0.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: 0.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
 
@@ -230,6 +231,7 @@ class QueueApi:
230
231
  input_data: Any = None,
231
232
  priority: int | None = None,
232
233
  parent_id: UUID | str | None = None,
234
+ idempotency_token: str | None = None,
233
235
  ) -> QueueItem:
234
236
  """Add a queue item. At least one of task_name or task_item_name is required."""
235
237
  if task_name is None and task_item_name is None:
@@ -241,6 +243,7 @@ class QueueApi:
241
243
  "taskItemName": task_item_name,
242
244
  "priority": priority,
243
245
  "inputData": _encode_input(input_data),
246
+ "idempotencyToken": idempotency_token,
244
247
  }
245
248
  )
246
249
  return QueueItem.model_validate(self._http.post_json(self._MGMT, body))
@@ -262,9 +265,15 @@ class QueueApi:
262
265
  return None
263
266
  return QueueItem.model_validate(result)
264
267
 
265
- 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})
266
275
  return QueueItem.model_validate(
267
- 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)
268
277
  )
269
278
 
270
279
  def update_status(
@@ -273,16 +282,50 @@ class QueueApi:
273
282
  status: TaskStatus,
274
283
  *,
275
284
  output_data: Any = None,
285
+ claim_token: UUID | str | None = None,
276
286
  ) -> QueueItem:
277
287
  # Server expects `status` on this endpoint (not `taskStatus` on queue item JSON).
278
288
  body: dict[str, Any] = {"status": status.value}
279
289
  encoded = _encode_input(output_data)
280
290
  if encoded is not None:
281
291
  body["outputData"] = encoded
292
+ if claim_token is not None:
293
+ body["claimToken"] = str(claim_token)
282
294
  return QueueItem.model_validate(
283
295
  self._http.post_json(f"{self._QUEUE}/queue-items/{queue_item_id}/status", body)
284
296
  )
285
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
+
286
329
 
287
330
  class ContentStoreApi:
288
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,7 +13,9 @@ the claim stays alive across process restarts.
13
13
 
14
14
  from __future__ import annotations
15
15
 
16
+ import inspect
16
17
  import logging
18
+ import os
17
19
  import signal
18
20
  import threading
19
21
  import time
@@ -27,13 +29,14 @@ from typing import Any
27
29
  import httpx
28
30
 
29
31
  from .client import ContinuumClient
32
+ from .context import TaskExecutionContext
30
33
  from .exceptions import ContinuumError
31
34
  from .models import QueueEvent, QueueEventType, QueueItem, TaskStatus, TransportMode
32
35
  from .websocket import ContinuumWebSocketClient, WebSocketOptions
33
36
 
34
37
  logger = logging.getLogger("continuum_task_server")
35
38
 
36
- Handler = Callable[[QueueItem], Any]
39
+ Handler = Callable[..., Any]
37
40
 
38
41
 
39
42
  @dataclass
@@ -61,10 +64,10 @@ class TaskServer:
61
64
  With ``auto_complete=True`` (default), handler return values become ``outputData``
62
65
  and the task is marked ENDED. Raising any exception marks the task KILLED.
63
66
 
64
- With ``auto_complete=False``, the handler returns without ENDED; you are
65
- responsible for heartbeats until you call ``complete_queue_item`` or
66
- ``fail_queue_item`` (see README). Heartbeats only run while the handler is
67
- 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``.
68
71
  """
69
72
 
70
73
  def __init__(
@@ -119,6 +122,7 @@ class TaskServer:
119
122
  self._claim_loop_lock = threading.Lock()
120
123
  self._base_url = base_url
121
124
  self._api_key = api_key
125
+ self._active: dict[uuid.UUID, _ActiveClaim] = {}
122
126
 
123
127
  @property
124
128
  def client(self) -> ContinuumClient:
@@ -179,7 +183,14 @@ class TaskServer:
179
183
  same ``ContinuumClient`` (and API key) as this ``TaskServer``.
180
184
  """
181
185
  qid = self._normalize_queue_id(queue_item_id)
182
- 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()
183
194
 
184
195
  def fail_queue_item(
185
196
  self,
@@ -189,6 +200,10 @@ class TaskServer:
189
200
  ) -> None:
190
201
  """Mark a queue item KILLED."""
191
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
192
207
  payload: dict[str, Any]
193
208
  if isinstance(error, dict):
194
209
  payload = dict(error)
@@ -204,7 +219,10 @@ class TaskServer:
204
219
  }
205
220
  else:
206
221
  payload = {"error": "failed"}
207
- 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()
208
226
 
209
227
  @staticmethod
210
228
  def _normalize_queue_id(queue_item_id: uuid.UUID | str) -> uuid.UUID:
@@ -243,6 +261,8 @@ class TaskServer:
243
261
  if not self._stop_event.is_set():
244
262
  logger.info("Continuum task server shutdown requested")
245
263
  self._stop_event.set()
264
+ for claim in list(self._active.values()):
265
+ claim.deferred_gate.set()
246
266
 
247
267
  def _install_signal_handlers(self) -> None:
248
268
  def _handler(signum: int, _frame: object) -> None:
@@ -311,6 +331,9 @@ class TaskServer:
311
331
  self._handle_queue_event(event)
312
332
 
313
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
314
337
  if event.event_type is not QueueEventType.WORK_AVAILABLE:
315
338
  return
316
339
  name = event.task_name
@@ -320,6 +343,17 @@ class TaskServer:
320
343
  self._pending_hints.add(name)
321
344
  self._try_claim_pending()
322
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
+
323
357
  def _try_claim_pending(self) -> None:
324
358
  while not self._stop_event.is_set():
325
359
  if not self._claim_loop_lock.acquire(blocking=False):
@@ -394,57 +428,132 @@ class TaskServer:
394
428
  self._inflight.discard(future)
395
429
 
396
430
  def _run_item(self, registration: _Registration, item: QueueItem) -> None:
431
+ claim = _ActiveClaim(item=item, server=self)
432
+ self._active[item.id] = claim
397
433
  try:
398
434
  logger.info("Claimed queue item %s (task=%s)", item.id, registration.name)
399
- stop_heartbeat = threading.Event()
400
435
  heartbeat_thread = threading.Thread(
401
436
  target=self._heartbeat_loop,
402
- args=(item, stop_heartbeat),
437
+ args=(claim,),
403
438
  name=f"hb-{item.id}",
404
439
  daemon=True,
405
440
  )
406
441
  heartbeat_thread.start()
407
442
 
408
443
  try:
409
- 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
+ )
410
450
  try:
411
- result = registration.handler(item)
451
+ result = self._invoke_handler(registration.handler, item, claim.context)
412
452
  except Exception as e:
453
+ if claim.control_seen.is_set() or claim.context.is_cancelled():
454
+ claim.acknowledge_control()
455
+ return
413
456
  logger.exception("Handler raised for queue item %s", item.id)
414
457
  error_payload = {
415
458
  "error": str(e),
416
459
  "type": type(e).__name__,
417
460
  "traceback": traceback.format_exc(),
418
461
  }
419
- 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()
420
470
  return
421
471
 
422
472
  if registration.auto_complete:
423
- 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
+ )
424
476
  logger.info("Completed queue item %s", item.id)
477
+ claim.mark_completed()
425
478
  else:
426
479
  logger.info(
427
480
  "Handler returned for queue item %s without ENDED "
428
- "(auto_complete=False); caller must heartbeat and complete",
481
+ "(auto_complete=False); remaining subscribed until complete",
429
482
  item.id,
430
483
  )
484
+ self._stop_heartbeat(claim, heartbeat_thread)
485
+ self._wait_for_deferred(claim)
431
486
  finally:
432
- stop_heartbeat.set()
433
- heartbeat_thread.join(timeout=1.0)
487
+ self._stop_heartbeat(claim, heartbeat_thread)
488
+ self._unsubscribe_control(claim)
434
489
  finally:
490
+ self._active.pop(item.id, None)
435
491
  registration.semaphore.release()
436
492
  if self._mode is TransportMode.WEBSOCKET:
437
493
  self._try_claim_pending()
438
494
 
439
- def _heartbeat_loop(self, item: QueueItem, stop_event: threading.Event) -> None:
440
- 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):
441
521
  try:
522
+ latest: QueueItem | None
442
523
  if self._ws is not None:
443
- self._ws.heartbeat(item.id)
524
+ latest = self._ws.heartbeat(claim.item.id, claim_token=claim.claim_token)
444
525
  else:
445
- 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)
446
530
  except ContinuumError as e:
447
- 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)
448
557
 
449
558
  def _safe_update_status(
450
559
  self,
@@ -452,8 +561,11 @@ class TaskServer:
452
561
  status: TaskStatus,
453
562
  *,
454
563
  output: Any = None,
455
- ) -> None:
456
- 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
+ )
457
569
 
458
570
  def _safe_update_status_by_id(
459
571
  self,
@@ -461,14 +573,23 @@ class TaskServer:
461
573
  status: TaskStatus,
462
574
  *,
463
575
  output: Any = None,
464
- ) -> None:
576
+ claim_token: uuid.UUID | None = None,
577
+ ) -> QueueItem | None:
465
578
  try:
579
+ latest: QueueItem | None
466
580
  if self._ws is not None:
467
- 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
+ )
468
584
  else:
469
- 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
470
590
  except ContinuumError as e:
471
591
  logger.error("update_status(%s, %s) failed: %s", queue_item_id, status.value, e)
592
+ return None
472
593
 
473
594
  def _drain(self) -> None:
474
595
  if self._executor is None:
@@ -494,9 +615,78 @@ class TaskServer:
494
615
  logger.info("Continuum task server stopped")
495
616
 
496
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
+
497
684
  def _coerce_transport(value: TransportMode | str | None) -> TransportMode:
498
685
  if value is None:
499
- return TransportMode.HTTP
686
+ env = os.environ.get("CONTINUUM_TRANSPORT")
687
+ if env is None or not str(env).strip():
688
+ return TransportMode.HTTP
689
+ value = env
500
690
  if isinstance(value, TransportMode):
501
691
  return value
502
692
  normalized = str(value).strip().lower()
@@ -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()
@@ -225,6 +226,7 @@ class ContinuumWebSocketClient:
225
226
  priority: int | None = None,
226
227
  parent_id: UUID | str | None = None,
227
228
  request_id: str | None = None,
229
+ idempotency_token: str | None = None,
228
230
  ) -> QueueItem:
229
231
  if task_name is None and task_item_name is None:
230
232
  raise ValueError("at least one of task_name or task_item_name is required")
@@ -235,6 +237,7 @@ class ContinuumWebSocketClient:
235
237
  "taskItemName": task_item_name,
236
238
  "priority": priority,
237
239
  "inputData": _encode_input(input_data),
240
+ "idempotencyToken": idempotency_token,
238
241
  }
239
242
  )
240
243
  payload = self._command("/app/enqueue", body, request_id=request_id, retry=False)
@@ -255,6 +258,7 @@ class ContinuumWebSocketClient:
255
258
  mode: WaitMode | str | None = None,
256
259
  correlation: str | None = None,
257
260
  request_id: str | None = None,
261
+ idempotency_token: str | None = None,
258
262
  ) -> EnqueueAndSubscribeWaitResult:
259
263
  if task_name is None and task_item_name is None:
260
264
  raise ValueError("at least one of task_name or task_item_name is required")
@@ -271,6 +275,7 @@ class ContinuumWebSocketClient:
271
275
  "subscriberId": subscriber_id,
272
276
  "mode": mode_value,
273
277
  "correlation": correlation,
278
+ "idempotencyToken": idempotency_token,
274
279
  }
275
280
  )
276
281
  payload = self._command(
@@ -295,10 +300,21 @@ class ContinuumWebSocketClient:
295
300
  return None
296
301
  return QueueItem.model_validate(item)
297
302
 
298
- 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:
299
310
  payload = self._command(
300
311
  "/app/heartbeat",
301
- {"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
+ ),
302
318
  request_id=request_id,
303
319
  retry=True,
304
320
  )
@@ -310,6 +326,7 @@ class ContinuumWebSocketClient:
310
326
  status: TaskStatus,
311
327
  *,
312
328
  output_data: Any = None,
329
+ claim_token: UUID | str | None = None,
313
330
  request_id: str | None = None,
314
331
  ) -> QueueItem:
315
332
  body: dict[str, Any] = {
@@ -319,6 +336,8 @@ class ContinuumWebSocketClient:
319
336
  encoded = _encode_input(output_data)
320
337
  if encoded is not None:
321
338
  body["outputData"] = encoded
339
+ if claim_token is not None:
340
+ body["claimToken"] = str(claim_token)
322
341
  payload = self._command("/app/status", body, request_id=request_id, retry=True)
323
342
  return QueueItem.model_validate(payload["item"])
324
343
 
@@ -348,6 +367,74 @@ class ContinuumWebSocketClient:
348
367
  self._work_subscriptions.discard(task_name)
349
368
  return WorkSubscriptionResult.model_validate(payload)
350
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
+
351
438
  def subscribe_wait(
352
439
  self,
353
440
  subscriber_id: str,
@@ -593,6 +680,7 @@ class ContinuumWebSocketClient:
593
680
  with self._state_lock:
594
681
  work = list(self._work_subscriptions)
595
682
  waits = list(self._wait_subscribers)
683
+ controls = list(self._control_subscriptions)
596
684
  for task_name in work:
597
685
  try:
598
686
  self._command_inline("/app/subscribeWork", {"taskName": task_name})
@@ -603,6 +691,21 @@ class ContinuumWebSocketClient:
603
691
  self._command_inline("/app/subscribeWait", {"subscriberId": subscriber_id})
604
692
  except ContinuumError:
605
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
+ )
606
709
 
607
710
  def _command_inline(self, destination: str, body: dict[str, Any]) -> dict[str, Any]:
608
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 = "0.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"