continuum-task-server-sdk 1.1.0__tar.gz → 1.4.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.4.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
@@ -112,6 +112,14 @@ for row in db.outstanding_rows():
112
112
 
113
113
  Standalone process (no `TaskServer`): build a `ContinuumClient` with the worker key and call `client.queue.heartbeat` / `client.queue.update_status` the same way.
114
114
 
115
+ By default the deferred claim still occupies the per-task ``concurrency`` slot until ``complete_queue_item``. Set ``occupy_slot_until_complete=False`` so the handler is still gated (you will not insert 10k rows at once) but the slot is freed when the function returns:
116
+
117
+ ```python
118
+ @server.task("discord", auto_complete=False, occupy_slot_until_complete=False)
119
+ def discord(item):
120
+ db.insert_await(queue_item_id=str(item.id), payload=item.input_data_json)
121
+ ```
122
+
115
123
  ### TaskServer options
116
124
 
117
125
  ```python
@@ -83,6 +83,14 @@ for row in db.outstanding_rows():
83
83
 
84
84
  Standalone process (no `TaskServer`): build a `ContinuumClient` with the worker key and call `client.queue.heartbeat` / `client.queue.update_status` the same way.
85
85
 
86
+ By default the deferred claim still occupies the per-task ``concurrency`` slot until ``complete_queue_item``. Set ``occupy_slot_until_complete=False`` so the handler is still gated (you will not insert 10k rows at once) but the slot is freed when the function returns:
87
+
88
+ ```python
89
+ @server.task("discord", auto_complete=False, occupy_slot_until_complete=False)
90
+ def discord(item):
91
+ db.insert_await(queue_item_id=str(item.id), payload=item.input_data_json)
92
+ ```
93
+
86
94
  ### TaskServer options
87
95
 
88
96
  ```python
@@ -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,
@@ -21,9 +22,11 @@ from .models import (
21
22
  Content,
22
23
  EnqueueAndSubscribeWaitResult,
23
24
  EventAckResult,
25
+ QueueDescendantsCancelResult,
24
26
  QueueEvent,
25
27
  QueueEventType,
26
28
  QueueItem,
29
+ QueueItemTrigger,
27
30
  TaskItem,
28
31
  TaskItemVersion,
29
32
  TaskStatus,
@@ -51,11 +54,14 @@ __all__ = [
51
54
  "ForbiddenError",
52
55
  "NotFoundError",
53
56
  "ProtocolError",
57
+ "QueueDescendantsCancelResult",
54
58
  "QueueEvent",
55
59
  "QueueEventType",
56
60
  "QueueItem",
61
+ "QueueItemTrigger",
57
62
  "RateLimitError",
58
63
  "ServerError",
64
+ "TaskExecutionContext",
59
65
  "TaskItem",
60
66
  "TaskItemVersion",
61
67
  "TaskServer",
@@ -71,4 +77,4 @@ __all__ = [
71
77
  "WorkSubscriptionResult",
72
78
  ]
73
79
 
74
- __version__ = "0.1.0"
80
+ __version__ = "1.4.0"
@@ -11,7 +11,16 @@ from uuid import UUID
11
11
  import httpx
12
12
 
13
13
  from ._http import HttpClient, drop_none
14
- from .models import Content, QueueItem, TaskItem, TaskItemVersion, TaskStatus, TaskType
14
+ from .exceptions import ContinuumError
15
+ from .models import (
16
+ Content,
17
+ QueueDescendantsCancelResult,
18
+ QueueItem,
19
+ TaskItem,
20
+ TaskItemVersion,
21
+ TaskStatus,
22
+ TaskType,
23
+ )
15
24
 
16
25
 
17
26
  def _encode_input(data: Any) -> str | None:
@@ -264,9 +273,15 @@ class QueueApi:
264
273
  return None
265
274
  return QueueItem.model_validate(result)
266
275
 
267
- def heartbeat(self, queue_item_id: UUID | str) -> QueueItem:
276
+ def heartbeat(
277
+ self,
278
+ queue_item_id: UUID | str,
279
+ *,
280
+ claim_token: UUID | str | None = None,
281
+ ) -> QueueItem:
282
+ body = drop_none({"claimToken": str(claim_token) if claim_token else None})
268
283
  return QueueItem.model_validate(
269
- self._http.post_json(f"{self._QUEUE}/queue-items/{queue_item_id}/heartbeat", {})
284
+ self._http.post_json(f"{self._QUEUE}/queue-items/{queue_item_id}/heartbeat", body)
270
285
  )
271
286
 
272
287
  def update_status(
@@ -275,16 +290,68 @@ class QueueApi:
275
290
  status: TaskStatus,
276
291
  *,
277
292
  output_data: Any = None,
293
+ claim_token: UUID | str | None = None,
278
294
  ) -> QueueItem:
279
295
  # Server expects `status` on this endpoint (not `taskStatus` on queue item JSON).
280
296
  body: dict[str, Any] = {"status": status.value}
281
297
  encoded = _encode_input(output_data)
282
298
  if encoded is not None:
283
299
  body["outputData"] = encoded
300
+ if claim_token is not None:
301
+ body["claimToken"] = str(claim_token)
284
302
  return QueueItem.model_validate(
285
303
  self._http.post_json(f"{self._QUEUE}/queue-items/{queue_item_id}/status", body)
286
304
  )
287
305
 
306
+ def cancel(self, queue_item_id: UUID | str) -> QueueItem:
307
+ result = self._http.post_json(f"{self._MGMT}/{queue_item_id}/cancel", {})
308
+ status = 0 if result is None else int(result.get("status", 0))
309
+ if result is None or status != 0:
310
+ message = None if result is None else result.get("message")
311
+ raise ContinuumError(message or "cancel failed")
312
+ return QueueItem.model_validate(result["item"])
313
+
314
+ def complete_control(
315
+ self,
316
+ queue_item_id: UUID | str,
317
+ *,
318
+ claim_token: UUID | str | None = None,
319
+ output_data: Any = None,
320
+ ) -> QueueItem:
321
+ body = drop_none(
322
+ {
323
+ "claimToken": str(claim_token) if claim_token else None,
324
+ "outputData": _encode_input(output_data),
325
+ }
326
+ )
327
+ result = self._http.post_json(
328
+ f"{self._QUEUE}/queue-items/{queue_item_id}/control-complete",
329
+ body,
330
+ )
331
+ status = 0 if result is None else int(result.get("status", 0))
332
+ if result is None or status != 0:
333
+ message = None if result is None else result.get("message")
334
+ raise ContinuumError(message or "control-complete failed")
335
+ return QueueItem.model_validate(result["item"])
336
+
337
+ def cancel_descendants(
338
+ self,
339
+ queue_item_id: UUID | str,
340
+ *,
341
+ claim_token: UUID | str | None = None,
342
+ ) -> QueueDescendantsCancelResult:
343
+ """Cancel descendants of a claimed item without cancelling the root."""
344
+ body = drop_none({"claimToken": str(claim_token) if claim_token else None})
345
+ result = self._http.post_json(
346
+ f"{self._QUEUE}/queue-items/{queue_item_id}/cancel-descendants",
347
+ body,
348
+ )
349
+ status = 0 if result is None else int(result.get("status", 0))
350
+ if result is None or status != 0:
351
+ message = None if result is None else result.get("message")
352
+ raise ContinuumError(message or "cancel-descendants failed")
353
+ return QueueDescendantsCancelResult.model_validate(result)
354
+
288
355
 
289
356
  class ContentStoreApi:
290
357
  _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:
@@ -103,6 +125,14 @@ class QueueItem(_Base):
103
125
  return self.output_data
104
126
 
105
127
 
128
+ class QueueDescendantsCancelResult(_Base):
129
+ """Result of ``POST /api/queue/queue-items/{id}/cancel-descendants``."""
130
+
131
+ status: int = 0
132
+ message: str | None = None
133
+ cancelled_count: int | None = Field(default=None, alias="cancelledCount")
134
+
135
+
106
136
  class Content(_Base):
107
137
  """Raw content returned from a content-store or queue-item content endpoint."""
108
138
 
@@ -132,6 +162,7 @@ class QueueEventType(str, Enum):
132
162
 
133
163
  WORK_AVAILABLE = "work.available"
134
164
  COMPLETED = "queue.completed"
165
+ CONTROL_REQUESTED = "queue.control.requested"
135
166
 
136
167
 
137
168
  class WaitTargetStatus(_Base):
@@ -157,6 +188,8 @@ class QueueEvent(_Base):
157
188
  correlation: str | None = None
158
189
  wait_mode: WaitMode | None = Field(default=None, alias="waitMode")
159
190
  targets: list[WaitTargetStatus] | None = None
191
+ control_request: str | None = Field(default=None, alias="controlRequest")
192
+ claim_token: UUID | None = Field(default=None, alias="claimToken")
160
193
 
161
194
 
162
195
  class WorkSubscriptionResult(_Base):
@@ -9,10 +9,15 @@ elsewhere and you will call ``complete_queue_item`` / ``fail_queue_item`` later.
9
9
  The SDK does **not** heartbeat after the handler returns; you must call
10
10
  ``client.queue.heartbeat`` yourself (for example on each pass of your DB poll) so
11
11
  the claim stays alive across process restarts.
12
+
13
+ By default a deferred claim still occupies the per-task ``concurrency`` slot
14
+ until complete. Pass ``occupy_slot_until_complete=False`` to free the slot when
15
+ the handler returns so another item can be claimed while the first stays open.
12
16
  """
13
17
 
14
18
  from __future__ import annotations
15
19
 
20
+ import inspect
16
21
  import logging
17
22
  import os
18
23
  import signal
@@ -28,13 +33,14 @@ from typing import Any
28
33
  import httpx
29
34
 
30
35
  from .client import ContinuumClient
36
+ from .context import TaskExecutionContext
31
37
  from .exceptions import ContinuumError
32
38
  from .models import QueueEvent, QueueEventType, QueueItem, TaskStatus, TransportMode
33
39
  from .websocket import ContinuumWebSocketClient, WebSocketOptions
34
40
 
35
41
  logger = logging.getLogger("continuum_task_server")
36
42
 
37
- Handler = Callable[[QueueItem], Any]
43
+ Handler = Callable[..., Any]
38
44
 
39
45
 
40
46
  @dataclass
@@ -43,6 +49,7 @@ class _Registration:
43
49
  handler: Handler
44
50
  concurrency: int
45
51
  auto_complete: bool
52
+ occupy_slot_until_complete: bool
46
53
  semaphore: threading.Semaphore = field(init=False)
47
54
 
48
55
  def __post_init__(self) -> None:
@@ -62,10 +69,11 @@ class TaskServer:
62
69
  With ``auto_complete=True`` (default), handler return values become ``outputData``
63
70
  and the task is marked ENDED. Raising any exception marks the task KILLED.
64
71
 
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.
72
+ With ``auto_complete=False``, the handler returns without ENDED. Automatic
73
+ heartbeats stop when the handler returns. Heartbeat from the completing
74
+ process with ``client.queue.heartbeat`` until ``complete_queue_item`` or
75
+ ``fail_queue_item``. The concurrency slot stays occupied until complete
76
+ unless ``occupy_slot_until_complete=False``.
69
77
  """
70
78
 
71
79
  def __init__(
@@ -120,6 +128,7 @@ class TaskServer:
120
128
  self._claim_loop_lock = threading.Lock()
121
129
  self._base_url = base_url
122
130
  self._api_key = api_key
131
+ self._active: dict[uuid.UUID, _ActiveClaim] = {}
123
132
 
124
133
  @property
125
134
  def client(self) -> ContinuumClient:
@@ -132,6 +141,7 @@ class TaskServer:
132
141
  *,
133
142
  concurrency: int = 1,
134
143
  auto_complete: bool = True,
144
+ occupy_slot_until_complete: bool = True,
135
145
  ) -> Callable[[Handler], Handler]:
136
146
  """Decorator registering ``handler`` for queue items of task-type ``name``.
137
147
 
@@ -140,6 +150,10 @@ class TaskServer:
140
150
  If ``auto_complete`` is False, the handler returns without the server sending
141
151
  ENDED; call ``complete_queue_item`` or ``fail_queue_item`` when done. You must
142
152
  heartbeat the queue item yourself until then (see README).
153
+
154
+ ``occupy_slot_until_complete`` (default True) keeps the concurrency slot until
155
+ that complete/fail. Set False so the slot is freed when the handler returns;
156
+ the claim stays in ``_active`` for complete/cancel.
143
157
  """
144
158
  if concurrency < 1:
145
159
  raise ValueError("concurrency must be >= 1")
@@ -152,6 +166,7 @@ class TaskServer:
152
166
  handler=handler,
153
167
  concurrency=concurrency,
154
168
  auto_complete=auto_complete,
169
+ occupy_slot_until_complete=occupy_slot_until_complete,
155
170
  )
156
171
  return handler
157
172
 
@@ -164,9 +179,15 @@ class TaskServer:
164
179
  *,
165
180
  concurrency: int = 1,
166
181
  auto_complete: bool = True,
182
+ occupy_slot_until_complete: bool = True,
167
183
  ) -> None:
168
184
  """Imperative alternative to ``@task``."""
169
- self.task(name, concurrency=concurrency, auto_complete=auto_complete)(handler)
185
+ self.task(
186
+ name,
187
+ concurrency=concurrency,
188
+ auto_complete=auto_complete,
189
+ occupy_slot_until_complete=occupy_slot_until_complete,
190
+ )(handler)
170
191
 
171
192
  def complete_queue_item(
172
193
  self,
@@ -180,7 +201,15 @@ class TaskServer:
180
201
  same ``ContinuumClient`` (and API key) as this ``TaskServer``.
181
202
  """
182
203
  qid = self._normalize_queue_id(queue_item_id)
183
- self._safe_update_status_by_id(qid, TaskStatus.ENDED, output=output_data)
204
+ claim = self._active.get(qid)
205
+ if claim is not None and claim.control_seen.is_set():
206
+ claim.acknowledge_control()
207
+ return
208
+ token = None if claim is None else claim.claim_token
209
+ self._safe_update_status_by_id(qid, TaskStatus.ENDED, output=output_data, claim_token=token)
210
+ if claim is not None:
211
+ claim.mark_completed()
212
+ self._finish_claim(claim)
184
213
 
185
214
  def fail_queue_item(
186
215
  self,
@@ -190,6 +219,10 @@ class TaskServer:
190
219
  ) -> None:
191
220
  """Mark a queue item KILLED."""
192
221
  qid = self._normalize_queue_id(queue_item_id)
222
+ claim = self._active.get(qid)
223
+ if claim is not None and claim.control_seen.is_set():
224
+ claim.acknowledge_control()
225
+ return
193
226
  payload: dict[str, Any]
194
227
  if isinstance(error, dict):
195
228
  payload = dict(error)
@@ -205,7 +238,11 @@ class TaskServer:
205
238
  }
206
239
  else:
207
240
  payload = {"error": "failed"}
208
- self._safe_update_status_by_id(qid, TaskStatus.KILLED, output=payload)
241
+ token = None if claim is None else claim.claim_token
242
+ self._safe_update_status_by_id(qid, TaskStatus.KILLED, output=payload, claim_token=token)
243
+ if claim is not None:
244
+ claim.mark_completed()
245
+ self._finish_claim(claim)
209
246
 
210
247
  @staticmethod
211
248
  def _normalize_queue_id(queue_item_id: uuid.UUID | str) -> uuid.UUID:
@@ -244,6 +281,8 @@ class TaskServer:
244
281
  if not self._stop_event.is_set():
245
282
  logger.info("Continuum task server shutdown requested")
246
283
  self._stop_event.set()
284
+ for claim in list(self._active.values()):
285
+ claim.deferred_gate.set()
247
286
 
248
287
  def _install_signal_handlers(self) -> None:
249
288
  def _handler(signum: int, _frame: object) -> None:
@@ -312,6 +351,9 @@ class TaskServer:
312
351
  self._handle_queue_event(event)
313
352
 
314
353
  def _handle_queue_event(self, event: QueueEvent) -> None:
354
+ if event.event_type is QueueEventType.CONTROL_REQUESTED:
355
+ self._apply_control_event(event)
356
+ return
315
357
  if event.event_type is not QueueEventType.WORK_AVAILABLE:
316
358
  return
317
359
  name = event.task_name
@@ -321,6 +363,17 @@ class TaskServer:
321
363
  self._pending_hints.add(name)
322
364
  self._try_claim_pending()
323
365
 
366
+ def _apply_control_event(self, event: QueueEvent) -> None:
367
+ if event.queue_item_id is None:
368
+ return
369
+ claim = self._active.get(event.queue_item_id)
370
+ if claim is None:
371
+ return
372
+ token = event.claim_token
373
+ if token is not None and claim.claim_token is not None and token != claim.claim_token:
374
+ return
375
+ claim.signal_control()
376
+
324
377
  def _try_claim_pending(self) -> None:
325
378
  while not self._stop_event.is_set():
326
379
  if not self._claim_loop_lock.acquire(blocking=False):
@@ -395,57 +448,151 @@ class TaskServer:
395
448
  self._inflight.discard(future)
396
449
 
397
450
  def _run_item(self, registration: _Registration, item: QueueItem) -> None:
451
+ claim = _ActiveClaim(
452
+ item=item,
453
+ server=self,
454
+ occupy_slot_until_complete=registration.occupy_slot_until_complete,
455
+ )
456
+ self._active[item.id] = claim
457
+ slot_held = True
398
458
  try:
399
459
  logger.info("Claimed queue item %s (task=%s)", item.id, registration.name)
400
- stop_heartbeat = threading.Event()
401
460
  heartbeat_thread = threading.Thread(
402
461
  target=self._heartbeat_loop,
403
- args=(item, stop_heartbeat),
462
+ args=(claim,),
404
463
  name=f"hb-{item.id}",
405
464
  daemon=True,
406
465
  )
407
466
  heartbeat_thread.start()
408
467
 
409
468
  try:
410
- self._safe_update_status(item, TaskStatus.STARTED)
469
+ self._subscribe_control(claim)
470
+ self._inspect_control(
471
+ self._safe_update_status(
472
+ item, TaskStatus.STARTED, claim_token=claim.claim_token
473
+ )
474
+ )
411
475
  try:
412
- result = registration.handler(item)
476
+ result = self._invoke_handler(registration.handler, item, claim.context)
413
477
  except Exception as e:
478
+ if claim.control_seen.is_set() or claim.context.is_cancelled():
479
+ claim.acknowledge_control()
480
+ return
414
481
  logger.exception("Handler raised for queue item %s", item.id)
415
482
  error_payload = {
416
483
  "error": str(e),
417
484
  "type": type(e).__name__,
418
485
  "traceback": traceback.format_exc(),
419
486
  }
420
- self._safe_update_status(item, TaskStatus.KILLED, output=error_payload)
487
+ self._safe_update_status(
488
+ item, TaskStatus.KILLED, output=error_payload, claim_token=claim.claim_token
489
+ )
490
+ claim.mark_completed()
491
+ return
492
+
493
+ if claim.control_seen.is_set() or claim.context.is_cancelled():
494
+ claim.acknowledge_control()
421
495
  return
422
496
 
423
497
  if registration.auto_complete:
424
- self._safe_update_status(item, TaskStatus.ENDED, output=result)
498
+ self._safe_update_status(
499
+ item, TaskStatus.ENDED, output=result, claim_token=claim.claim_token
500
+ )
425
501
  logger.info("Completed queue item %s", item.id)
502
+ claim.mark_completed()
426
503
  else:
427
504
  logger.info(
428
505
  "Handler returned for queue item %s without ENDED "
429
- "(auto_complete=False); caller must heartbeat and complete",
506
+ "(auto_complete=False); remaining subscribed until complete",
430
507
  item.id,
431
508
  )
509
+ self._stop_heartbeat(claim, heartbeat_thread)
510
+ claim.handler_returned.set()
511
+ if registration.occupy_slot_until_complete:
512
+ self._wait_for_deferred(claim)
513
+ else:
514
+ slot_held = False
515
+ self._release_slot(registration)
432
516
  finally:
433
- stop_heartbeat.set()
434
- heartbeat_thread.join(timeout=1.0)
517
+ self._stop_heartbeat(claim, heartbeat_thread)
518
+ if slot_held:
519
+ self._unsubscribe_control(claim)
435
520
  finally:
436
- registration.semaphore.release()
437
- if self._mode is TransportMode.WEBSOCKET:
438
- self._try_claim_pending()
521
+ if slot_held:
522
+ self._active.pop(item.id, None)
523
+ self._release_slot(registration)
524
+
525
+ def _release_slot(self, registration: _Registration) -> None:
526
+ registration.semaphore.release()
527
+ if self._mode is TransportMode.WEBSOCKET:
528
+ self._try_claim_pending()
529
+
530
+ def _finish_claim(self, claim: _ActiveClaim) -> None:
531
+ self._unsubscribe_control(claim)
532
+ self._active.pop(claim.item.id, None)
533
+
534
+ def _wait_for_deferred(self, claim: _ActiveClaim) -> None:
535
+ while (
536
+ not self._stop_event.is_set()
537
+ and not claim.completed.is_set()
538
+ and not claim.control_seen.is_set()
539
+ ):
540
+ if claim.deferred_gate.wait(timeout=0.1):
541
+ break
542
+ if claim.control_seen.is_set() or claim.context.is_cancelled():
543
+ claim.acknowledge_control()
544
+
545
+ def _stop_heartbeat(self, claim: _ActiveClaim, heartbeat_thread: threading.Thread) -> None:
546
+ claim.stop_heartbeat.set()
547
+ if heartbeat_thread is threading.current_thread():
548
+ return
549
+ heartbeat_thread.join(timeout=1.0)
550
+
551
+ def _invoke_handler(
552
+ self, handler: Handler, item: QueueItem, context: TaskExecutionContext
553
+ ) -> Any:
554
+ if _accepts_context(handler):
555
+ return handler(item, context)
556
+ return handler(item)
439
557
 
440
- def _heartbeat_loop(self, item: QueueItem, stop_event: threading.Event) -> None:
441
- while not stop_event.wait(timeout=self._heartbeat_interval):
558
+ def _heartbeat_loop(self, claim: _ActiveClaim) -> None:
559
+ while not claim.stop_heartbeat.wait(timeout=self._heartbeat_interval):
442
560
  try:
561
+ latest: QueueItem | None
443
562
  if self._ws is not None:
444
- self._ws.heartbeat(item.id)
563
+ latest = self._ws.heartbeat(claim.item.id, claim_token=claim.claim_token)
445
564
  else:
446
- self._client.queue.heartbeat(item.id)
565
+ latest = self._client.queue.heartbeat(
566
+ claim.item.id, claim_token=claim.claim_token
567
+ )
568
+ self._inspect_control(latest)
447
569
  except ContinuumError as e:
448
- logger.warning("heartbeat for %s failed: %s", item.id, e)
570
+ logger.warning("heartbeat for %s failed: %s", claim.item.id, e)
571
+
572
+ def _inspect_control(self, latest: QueueItem | None) -> None:
573
+ if latest is None or latest.control_request in (None, ""):
574
+ return
575
+ claim = self._active.get(latest.id)
576
+ if claim is None:
577
+ return
578
+ claim.signal_control()
579
+
580
+ def _subscribe_control(self, claim: _ActiveClaim) -> None:
581
+ if self._ws is None:
582
+ return
583
+ try:
584
+ self._ws.subscribe_control(claim.item.id, claim_token=claim.claim_token)
585
+ claim.control_subscribed = True
586
+ except ContinuumError as e:
587
+ logger.warning("subscribeControl(%s) failed: %s", claim.item.id, e)
588
+
589
+ def _unsubscribe_control(self, claim: _ActiveClaim) -> None:
590
+ if self._ws is None or not claim.control_subscribed:
591
+ return
592
+ try:
593
+ self._ws.unsubscribe_control(claim.item.id, claim_token=claim.claim_token)
594
+ except ContinuumError as e:
595
+ logger.debug("unsubscribeControl(%s) failed: %s", claim.item.id, e)
449
596
 
450
597
  def _safe_update_status(
451
598
  self,
@@ -453,8 +600,11 @@ class TaskServer:
453
600
  status: TaskStatus,
454
601
  *,
455
602
  output: Any = None,
456
- ) -> None:
457
- self._safe_update_status_by_id(item.id, status, output=output)
603
+ claim_token: uuid.UUID | None = None,
604
+ ) -> QueueItem | None:
605
+ return self._safe_update_status_by_id(
606
+ item.id, status, output=output, claim_token=claim_token
607
+ )
458
608
 
459
609
  def _safe_update_status_by_id(
460
610
  self,
@@ -462,14 +612,23 @@ class TaskServer:
462
612
  status: TaskStatus,
463
613
  *,
464
614
  output: Any = None,
465
- ) -> None:
615
+ claim_token: uuid.UUID | None = None,
616
+ ) -> QueueItem | None:
466
617
  try:
618
+ latest: QueueItem | None
467
619
  if self._ws is not None:
468
- self._ws.update_status(queue_item_id, status, output_data=output)
620
+ latest = self._ws.update_status(
621
+ queue_item_id, status, output_data=output, claim_token=claim_token
622
+ )
469
623
  else:
470
- self._client.queue.update_status(queue_item_id, status, output_data=output)
624
+ latest = self._client.queue.update_status(
625
+ queue_item_id, status, output_data=output, claim_token=claim_token
626
+ )
627
+ self._inspect_control(latest)
628
+ return latest
471
629
  except ContinuumError as e:
472
630
  logger.error("update_status(%s, %s) failed: %s", queue_item_id, status.value, e)
631
+ return None
473
632
 
474
633
  def _drain(self) -> None:
475
634
  if self._executor is None:
@@ -492,9 +651,92 @@ class TaskServer:
492
651
  self._ws = None
493
652
  if self._owns_client:
494
653
  self._client.close()
654
+ for leftover in list(self._active.values()):
655
+ self._finish_claim(leftover)
495
656
  logger.info("Continuum task server stopped")
496
657
 
497
658
 
659
+ class _ActiveClaim:
660
+ def __init__(
661
+ self,
662
+ *,
663
+ item: QueueItem,
664
+ server: TaskServer,
665
+ occupy_slot_until_complete: bool = True,
666
+ ) -> None:
667
+ self.item = item
668
+ self.server = server
669
+ self.occupy_slot_until_complete = occupy_slot_until_complete
670
+ self.context = TaskExecutionContext(item)
671
+ self.stop_heartbeat = threading.Event()
672
+ self.control_seen = threading.Event()
673
+ self.completed = threading.Event()
674
+ self.deferred_gate = threading.Event()
675
+ self.control_subscribed = False
676
+ self._control_acked = threading.Event()
677
+ self.handler_returned = threading.Event()
678
+
679
+ @property
680
+ def claim_token(self) -> uuid.UUID | None:
681
+ return self.item.claim_token
682
+
683
+ def signal_control(self) -> None:
684
+ if self.control_seen.is_set():
685
+ return
686
+ self.control_seen.set()
687
+ self.context.request_cancel()
688
+ self.deferred_gate.set()
689
+ if self.handler_returned.is_set() and not self.occupy_slot_until_complete:
690
+ threading.Thread(
691
+ target=self.acknowledge_control,
692
+ name=f"control-ack-{self.item.id}",
693
+ daemon=True,
694
+ ).start()
695
+
696
+ def mark_completed(self) -> None:
697
+ self.completed.set()
698
+ self.deferred_gate.set()
699
+
700
+ def acknowledge_control(self) -> None:
701
+ if self._control_acked.is_set():
702
+ return
703
+ self._control_acked.set()
704
+ self.context.request_cancel()
705
+ self.context.await_cleanup(max(0.001, self.server._shutdown_timeout))
706
+ try:
707
+ if self.server._ws is not None:
708
+ self.server._ws.complete_control(self.item.id, claim_token=self.claim_token)
709
+ else:
710
+ self.server._client.queue.complete_control(
711
+ self.item.id, claim_token=self.claim_token
712
+ )
713
+ logger.info("Acknowledged control for queue item %s", self.item.id)
714
+ except ContinuumError as e:
715
+ logger.warning("completeControl(%s) failed: %s", self.item.id, e)
716
+ self.mark_completed()
717
+ self.server._finish_claim(self)
718
+
719
+
720
+ def _accepts_context(handler: Handler) -> bool:
721
+ try:
722
+ signature = inspect.signature(handler)
723
+ except (TypeError, ValueError):
724
+ return False
725
+ positional = [
726
+ param
727
+ for param in signature.parameters.values()
728
+ if param.kind
729
+ in (
730
+ inspect.Parameter.POSITIONAL_ONLY,
731
+ inspect.Parameter.POSITIONAL_OR_KEYWORD,
732
+ inspect.Parameter.VAR_POSITIONAL,
733
+ )
734
+ ]
735
+ if any(param.kind == inspect.Parameter.VAR_POSITIONAL for param in positional):
736
+ return True
737
+ return len(positional) >= 2
738
+
739
+
498
740
  def _coerce_transport(value: TransportMode | str | None) -> TransportMode:
499
741
  if value is None:
500
742
  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.4.0"
8
8
  description = "Python SDK for the Continuum Task Server"
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.10"