continuum-task-server-sdk 1.2.0__tar.gz → 1.5.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.2.0
3
+ Version: 1.5.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
@@ -22,6 +22,7 @@ from .models import (
22
22
  Content,
23
23
  EnqueueAndSubscribeWaitResult,
24
24
  EventAckResult,
25
+ QueueDescendantsCancelResult,
25
26
  QueueEvent,
26
27
  QueueEventType,
27
28
  QueueItem,
@@ -53,6 +54,7 @@ __all__ = [
53
54
  "ForbiddenError",
54
55
  "NotFoundError",
55
56
  "ProtocolError",
57
+ "QueueDescendantsCancelResult",
56
58
  "QueueEvent",
57
59
  "QueueEventType",
58
60
  "QueueItem",
@@ -75,4 +77,4 @@ __all__ = [
75
77
  "WorkSubscriptionResult",
76
78
  ]
77
79
 
78
- __version__ = "0.1.0"
80
+ __version__ = "1.4.0"
@@ -46,6 +46,21 @@ class HttpClient:
46
46
  def close(self) -> None:
47
47
  self._client.close()
48
48
 
49
+ def set_task_server_token(self, token: str | None) -> None:
50
+ if token:
51
+ self._client.headers["Task-Server-Token"] = token
52
+ else:
53
+ self._client.headers.pop("Task-Server-Token", None)
54
+
55
+ def delete(self, path: str) -> None:
56
+ try:
57
+ response = self._client.delete(path)
58
+ except httpx.HTTPError as e:
59
+ raise ContinuumError(f"Request failed: {e}") from e
60
+ if response.status_code == 204:
61
+ return
62
+ self._raise_for_status(response)
63
+
49
64
  def __enter__(self) -> HttpClient:
50
65
  return self
51
66
 
@@ -12,7 +12,15 @@ import httpx
12
12
 
13
13
  from ._http import HttpClient, drop_none
14
14
  from .exceptions import ContinuumError
15
- from .models import Content, QueueItem, TaskItem, TaskItemVersion, TaskStatus, TaskType
15
+ from .models import (
16
+ Content,
17
+ QueueDescendantsCancelResult,
18
+ QueueItem,
19
+ TaskItem,
20
+ TaskItemVersion,
21
+ TaskStatus,
22
+ TaskType,
23
+ )
16
24
 
17
25
 
18
26
  def _encode_input(data: Any) -> str | None:
@@ -326,6 +334,24 @@ class QueueApi:
326
334
  raise ContinuumError(message or "control-complete failed")
327
335
  return QueueItem.model_validate(result["item"])
328
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
+
329
355
 
330
356
  class ContentStoreApi:
331
357
  _PATH = "/api/management/content-store"
@@ -382,6 +408,16 @@ class ContinuumClient:
382
408
  def base_url(self) -> str:
383
409
  return self._http.base_url
384
410
 
411
+ def set_task_server_token(self, token: str | None) -> None:
412
+ """Attach the secret from task-server registration. Blank clears it."""
413
+ self._http.set_task_server_token(token)
414
+
415
+ def post_raw(self, path: str, body: Any) -> Any:
416
+ return self._http.post_json(path, body)
417
+
418
+ def delete(self, path: str) -> None:
419
+ self._http.delete(path)
420
+
385
421
  def close(self) -> None:
386
422
  self._http.close()
387
423
 
@@ -125,6 +125,14 @@ class QueueItem(_Base):
125
125
  return self.output_data
126
126
 
127
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
+
128
136
  class Content(_Base):
129
137
  """Raw content returned from a content-store or queue-item content endpoint."""
130
138
 
@@ -0,0 +1,175 @@
1
+ """Registers this process with the queue API and refreshes the lease.
2
+
3
+ A missing or failing registration never stops the task server from claiming work.
4
+ """
5
+
6
+ from __future__ import annotations
7
+
8
+ import logging
9
+ import socket
10
+ import threading
11
+ import uuid
12
+ from collections.abc import Callable, Iterable
13
+ from importlib.metadata import PackageNotFoundError, version
14
+
15
+ from .client import ContinuumClient
16
+ from .exceptions import ContinuumError, NotFoundError
17
+ from .websocket import ContinuumWebSocketClient
18
+
19
+ logger = logging.getLogger("continuum_task_server")
20
+
21
+ _PATH = "/api/queue/task-servers"
22
+
23
+
24
+ def sdk_version() -> str:
25
+ try:
26
+ return version("continuum-task-server-sdk")
27
+ except PackageNotFoundError:
28
+ return "unknown"
29
+
30
+
31
+ class TaskServerPresence:
32
+ def __init__(
33
+ self,
34
+ client: ContinuumClient,
35
+ *,
36
+ name: str | None,
37
+ instance_id: str | None,
38
+ labels: dict[str, str] | None,
39
+ task_names: Callable[[], Iterable[str]],
40
+ ) -> None:
41
+ self._client = client
42
+ self._name = name or None
43
+ self._instance_id = instance_id or str(uuid.uuid4())
44
+ self._labels = _labels(labels)
45
+ self._task_names = task_names
46
+ self._enabled = True
47
+ self._stopped = threading.Event()
48
+ self._token: str | None = None
49
+ self._heartbeat_seconds = 10.0
50
+ self._websocket: ContinuumWebSocketClient | None = None
51
+ self._thread: threading.Thread | None = None
52
+
53
+ def attach(self, websocket: ContinuumWebSocketClient | None) -> None:
54
+ self._websocket = websocket
55
+ if websocket is not None and self._token:
56
+ websocket.set_task_server_token(self._token)
57
+
58
+ def start(self) -> None:
59
+ if self._thread is not None:
60
+ return
61
+ self._thread = threading.Thread(target=self._loop, name="continuum-presence", daemon=True)
62
+ self._thread.start()
63
+
64
+ def close(self) -> None:
65
+ self._stopped.set()
66
+ if self._thread is not None:
67
+ self._thread.join(timeout=1)
68
+ current = self._token
69
+ self._token = None
70
+ if not current:
71
+ self._client.set_task_server_token(None)
72
+ return
73
+ self._client.set_task_server_token(current)
74
+ try:
75
+ self._client.delete(_PATH)
76
+ except ContinuumError as exc:
77
+ logger.debug("task server deregister failed: %s", exc)
78
+ finally:
79
+ self._client.set_task_server_token(None)
80
+ if self._websocket is not None:
81
+ self._websocket.set_task_server_token(None)
82
+
83
+ def reregister(self) -> bool:
84
+ if not self._enabled or self._stopped.is_set():
85
+ return False
86
+ self._token = None
87
+ self._client.set_task_server_token(None)
88
+ return self._register()
89
+
90
+ def _loop(self) -> None:
91
+ backoff = 1.0
92
+ while not self._stopped.is_set() and self._enabled:
93
+ if self._token is None:
94
+ if not self._register():
95
+ if not self._enabled or self._stopped.is_set():
96
+ return
97
+ if self._stopped.wait(backoff):
98
+ return
99
+ backoff = min(backoff * 2, 30.0)
100
+ continue
101
+ backoff = 1.0
102
+ if self._stopped.wait(self._heartbeat_seconds) or self._token is None:
103
+ continue
104
+ if not self._heartbeat():
105
+ self._token = None
106
+ self._client.set_task_server_token(None)
107
+
108
+ def _register(self) -> bool:
109
+ body = {
110
+ "instanceId": self._instance_id,
111
+ "name": self._name,
112
+ "taskNames": list(self._task_names()),
113
+ "labels": self._labels,
114
+ }
115
+ try:
116
+ payload = self._client.post_raw(_PATH, body) or {}
117
+ if payload.get("status") != 0:
118
+ logger.warning("task server register rejected: %s", payload.get("message"))
119
+ return False
120
+ issued = payload.get("token")
121
+ if not isinstance(issued, str) or not issued:
122
+ logger.warning("task server register returned no token")
123
+ return False
124
+ seconds = payload.get("heartbeatSeconds")
125
+ self._heartbeat_seconds = (
126
+ float(seconds) if isinstance(seconds, (int, float)) and seconds > 0 else 10.0
127
+ )
128
+ self._token = issued
129
+ self._client.set_task_server_token(issued)
130
+ if self._websocket is not None:
131
+ self._websocket.set_task_server_token(issued)
132
+ server = (
133
+ payload.get("taskServer") if isinstance(payload.get("taskServer"), dict) else {}
134
+ )
135
+ logger.info("Registered task server %s", server.get("id"))
136
+ return True
137
+ except NotFoundError:
138
+ self._enabled = False
139
+ self._token = None
140
+ self._client.set_task_server_token(None)
141
+ logger.info("Queue API does not support task server presence; claiming without it")
142
+ return False
143
+ except Exception as exc:
144
+ logger.warning("task server register failed: %s", exc)
145
+ return False
146
+
147
+ def _heartbeat(self) -> bool:
148
+ try:
149
+ self._client.post_raw(_PATH + "/heartbeat", {})
150
+ return True
151
+ except NotFoundError:
152
+ self._enabled = False
153
+ logger.info("Queue API does not support task server presence; claiming without it")
154
+ return False
155
+ except Exception as exc:
156
+ logger.warning("task server heartbeat failed: %s", exc)
157
+ return False
158
+
159
+
160
+ def _labels(user: dict[str, str] | None) -> dict[str, str]:
161
+ merged: dict[str, str] = {}
162
+ if user:
163
+ merged.update(user)
164
+ merged["continuum.sdk"] = "python"
165
+ merged["continuum.sdkVersion"] = sdk_version()
166
+ merged["continuum.hostname"] = _hostname()
167
+ return merged
168
+
169
+
170
+ def _hostname() -> str:
171
+ try:
172
+ host = socket.gethostname()
173
+ except OSError:
174
+ return "unknown"
175
+ return host or "unknown"
@@ -9,6 +9,10 @@ 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
@@ -30,8 +34,9 @@ import httpx
30
34
 
31
35
  from .client import ContinuumClient
32
36
  from .context import TaskExecutionContext
33
- from .exceptions import ContinuumError
37
+ from .exceptions import ConflictError, ContinuumError
34
38
  from .models import QueueEvent, QueueEventType, QueueItem, TaskStatus, TransportMode
39
+ from .presence import TaskServerPresence
35
40
  from .websocket import ContinuumWebSocketClient, WebSocketOptions
36
41
 
37
42
  logger = logging.getLogger("continuum_task_server")
@@ -45,6 +50,7 @@ class _Registration:
45
50
  handler: Handler
46
51
  concurrency: int
47
52
  auto_complete: bool
53
+ occupy_slot_until_complete: bool
48
54
  semaphore: threading.Semaphore = field(init=False)
49
55
 
50
56
  def __post_init__(self) -> None:
@@ -67,7 +73,8 @@ class TaskServer:
67
73
  With ``auto_complete=False``, the handler returns without ENDED. Automatic
68
74
  heartbeats stop when the handler returns. Heartbeat from the completing
69
75
  process with ``client.queue.heartbeat`` until ``complete_queue_item`` or
70
- ``fail_queue_item``.
76
+ ``fail_queue_item``. The concurrency slot stays occupied until complete
77
+ unless ``occupy_slot_until_complete=False``.
71
78
  """
72
79
 
73
80
  def __init__(
@@ -86,6 +93,10 @@ class TaskServer:
86
93
  transport: TransportMode | str | httpx.BaseTransport | None = None,
87
94
  websocket_options: WebSocketOptions | None = None,
88
95
  websocket_client: Any = None,
96
+ name: str | None = None,
97
+ labels: dict[str, str] | None = None,
98
+ instance_id: str | None = None,
99
+ presence: bool = True,
89
100
  ) -> None:
90
101
  http_transport: httpx.BaseTransport | None = None
91
102
  if isinstance(transport, httpx.BaseTransport):
@@ -123,6 +134,19 @@ class TaskServer:
123
134
  self._base_url = base_url
124
135
  self._api_key = api_key
125
136
  self._active: dict[uuid.UUID, _ActiveClaim] = {}
137
+ self._presence = (
138
+ TaskServerPresence(
139
+ self._client,
140
+ name=name,
141
+ instance_id=instance_id,
142
+ labels=labels,
143
+ task_names=self._handlers.keys,
144
+ )
145
+ if presence
146
+ else None
147
+ )
148
+ if self._presence is not None and self._ws is not None:
149
+ self._presence.attach(self._ws)
126
150
 
127
151
  @property
128
152
  def client(self) -> ContinuumClient:
@@ -135,6 +159,7 @@ class TaskServer:
135
159
  *,
136
160
  concurrency: int = 1,
137
161
  auto_complete: bool = True,
162
+ occupy_slot_until_complete: bool = True,
138
163
  ) -> Callable[[Handler], Handler]:
139
164
  """Decorator registering ``handler`` for queue items of task-type ``name``.
140
165
 
@@ -143,6 +168,10 @@ class TaskServer:
143
168
  If ``auto_complete`` is False, the handler returns without the server sending
144
169
  ENDED; call ``complete_queue_item`` or ``fail_queue_item`` when done. You must
145
170
  heartbeat the queue item yourself until then (see README).
171
+
172
+ ``occupy_slot_until_complete`` (default True) keeps the concurrency slot until
173
+ that complete/fail. Set False so the slot is freed when the handler returns;
174
+ the claim stays in ``_active`` for complete/cancel.
146
175
  """
147
176
  if concurrency < 1:
148
177
  raise ValueError("concurrency must be >= 1")
@@ -155,6 +184,7 @@ class TaskServer:
155
184
  handler=handler,
156
185
  concurrency=concurrency,
157
186
  auto_complete=auto_complete,
187
+ occupy_slot_until_complete=occupy_slot_until_complete,
158
188
  )
159
189
  return handler
160
190
 
@@ -167,9 +197,15 @@ class TaskServer:
167
197
  *,
168
198
  concurrency: int = 1,
169
199
  auto_complete: bool = True,
200
+ occupy_slot_until_complete: bool = True,
170
201
  ) -> None:
171
202
  """Imperative alternative to ``@task``."""
172
- self.task(name, concurrency=concurrency, auto_complete=auto_complete)(handler)
203
+ self.task(
204
+ name,
205
+ concurrency=concurrency,
206
+ auto_complete=auto_complete,
207
+ occupy_slot_until_complete=occupy_slot_until_complete,
208
+ )(handler)
173
209
 
174
210
  def complete_queue_item(
175
211
  self,
@@ -191,6 +227,7 @@ class TaskServer:
191
227
  self._safe_update_status_by_id(qid, TaskStatus.ENDED, output=output_data, claim_token=token)
192
228
  if claim is not None:
193
229
  claim.mark_completed()
230
+ self._finish_claim(claim)
194
231
 
195
232
  def fail_queue_item(
196
233
  self,
@@ -223,6 +260,7 @@ class TaskServer:
223
260
  self._safe_update_status_by_id(qid, TaskStatus.KILLED, output=payload, claim_token=token)
224
261
  if claim is not None:
225
262
  claim.mark_completed()
263
+ self._finish_claim(claim)
226
264
 
227
265
  @staticmethod
228
266
  def _normalize_queue_id(queue_item_id: uuid.UUID | str) -> uuid.UUID:
@@ -249,6 +287,8 @@ class TaskServer:
249
287
  )
250
288
 
251
289
  try:
290
+ if self._presence is not None:
291
+ self._presence.start()
252
292
  if self._mode is TransportMode.WEBSOCKET:
253
293
  self._websocket_loop()
254
294
  else:
@@ -290,7 +330,7 @@ class TaskServer:
290
330
  continue
291
331
 
292
332
  try:
293
- item = self._client.queue.claim(registration.name)
333
+ item = self._claim_http(registration.name)
294
334
  except ContinuumError as e:
295
335
  registration.semaphore.release()
296
336
  logger.error("claim(%s) failed: %s", registration.name, e)
@@ -316,6 +356,8 @@ class TaskServer:
316
356
  options=self._ws_options,
317
357
  )
318
358
  self._owns_ws = True
359
+ if self._presence is not None:
360
+ self._presence.attach(self._ws)
319
361
  return self._ws
320
362
 
321
363
  def _websocket_loop(self) -> None:
@@ -386,7 +428,7 @@ class TaskServer:
386
428
  if registration is None:
387
429
  return
388
430
  try:
389
- item = ws.claim(registration.name)
431
+ item = self._claim_socket(ws, registration.name)
390
432
  except ContinuumError as e:
391
433
  registration.semaphore.release()
392
434
  logger.error("claim(%s) failed: %s", registration.name, e)
@@ -428,8 +470,13 @@ class TaskServer:
428
470
  self._inflight.discard(future)
429
471
 
430
472
  def _run_item(self, registration: _Registration, item: QueueItem) -> None:
431
- claim = _ActiveClaim(item=item, server=self)
473
+ claim = _ActiveClaim(
474
+ item=item,
475
+ server=self,
476
+ occupy_slot_until_complete=registration.occupy_slot_until_complete,
477
+ )
432
478
  self._active[item.id] = claim
479
+ slot_held = True
433
480
  try:
434
481
  logger.info("Claimed queue item %s (task=%s)", item.id, registration.name)
435
482
  heartbeat_thread = threading.Thread(
@@ -482,15 +529,29 @@ class TaskServer:
482
529
  item.id,
483
530
  )
484
531
  self._stop_heartbeat(claim, heartbeat_thread)
485
- self._wait_for_deferred(claim)
532
+ claim.handler_returned.set()
533
+ if registration.occupy_slot_until_complete:
534
+ self._wait_for_deferred(claim)
535
+ else:
536
+ slot_held = False
537
+ self._release_slot(registration)
486
538
  finally:
487
539
  self._stop_heartbeat(claim, heartbeat_thread)
488
- self._unsubscribe_control(claim)
540
+ if slot_held:
541
+ self._unsubscribe_control(claim)
489
542
  finally:
490
- self._active.pop(item.id, None)
491
- registration.semaphore.release()
492
- if self._mode is TransportMode.WEBSOCKET:
493
- self._try_claim_pending()
543
+ if slot_held:
544
+ self._active.pop(item.id, None)
545
+ self._release_slot(registration)
546
+
547
+ def _release_slot(self, registration: _Registration) -> None:
548
+ registration.semaphore.release()
549
+ if self._mode is TransportMode.WEBSOCKET:
550
+ self._try_claim_pending()
551
+
552
+ def _finish_claim(self, claim: _ActiveClaim) -> None:
553
+ self._unsubscribe_control(claim)
554
+ self._active.pop(claim.item.id, None)
494
555
 
495
556
  def _wait_for_deferred(self, claim: _ActiveClaim) -> None:
496
557
  while (
@@ -591,6 +652,22 @@ class TaskServer:
591
652
  logger.error("update_status(%s, %s) failed: %s", queue_item_id, status.value, e)
592
653
  return None
593
654
 
655
+ def _claim_http(self, task_name: str) -> QueueItem | None:
656
+ try:
657
+ return self._client.queue.claim(task_name)
658
+ except ConflictError:
659
+ if self._presence is None or not self._presence.reregister():
660
+ raise
661
+ return self._client.queue.claim(task_name)
662
+
663
+ def _claim_socket(self, ws: ContinuumWebSocketClient, task_name: str) -> QueueItem | None:
664
+ try:
665
+ return ws.claim(task_name)
666
+ except ConflictError:
667
+ if self._presence is None or not self._presence.reregister():
668
+ raise
669
+ return ws.claim(task_name)
670
+
594
671
  def _drain(self) -> None:
595
672
  if self._executor is None:
596
673
  return
@@ -606,19 +683,30 @@ class TaskServer:
606
683
  logger.debug("drain: handler future raised", exc_info=True)
607
684
  self._executor.shutdown(wait=False, cancel_futures=True)
608
685
  self._executor = None
686
+ if self._presence is not None:
687
+ self._presence.close()
609
688
  if self._ws is not None:
610
689
  self._ws.close()
611
690
  if self._owns_ws:
612
691
  self._ws = None
613
692
  if self._owns_client:
614
693
  self._client.close()
694
+ for leftover in list(self._active.values()):
695
+ self._finish_claim(leftover)
615
696
  logger.info("Continuum task server stopped")
616
697
 
617
698
 
618
699
  class _ActiveClaim:
619
- def __init__(self, *, item: QueueItem, server: TaskServer) -> None:
700
+ def __init__(
701
+ self,
702
+ *,
703
+ item: QueueItem,
704
+ server: TaskServer,
705
+ occupy_slot_until_complete: bool = True,
706
+ ) -> None:
620
707
  self.item = item
621
708
  self.server = server
709
+ self.occupy_slot_until_complete = occupy_slot_until_complete
622
710
  self.context = TaskExecutionContext(item)
623
711
  self.stop_heartbeat = threading.Event()
624
712
  self.control_seen = threading.Event()
@@ -626,6 +714,7 @@ class _ActiveClaim:
626
714
  self.deferred_gate = threading.Event()
627
715
  self.control_subscribed = False
628
716
  self._control_acked = threading.Event()
717
+ self.handler_returned = threading.Event()
629
718
 
630
719
  @property
631
720
  def claim_token(self) -> uuid.UUID | None:
@@ -637,6 +726,12 @@ class _ActiveClaim:
637
726
  self.control_seen.set()
638
727
  self.context.request_cancel()
639
728
  self.deferred_gate.set()
729
+ if self.handler_returned.is_set() and not self.occupy_slot_until_complete:
730
+ threading.Thread(
731
+ target=self.acknowledge_control,
732
+ name=f"control-ack-{self.item.id}",
733
+ daemon=True,
734
+ ).start()
640
735
 
641
736
  def mark_completed(self) -> None:
642
737
  self.completed.set()
@@ -659,6 +754,7 @@ class _ActiveClaim:
659
754
  except ContinuumError as e:
660
755
  logger.warning("completeControl(%s) failed: %s", self.item.id, e)
661
756
  self.mark_completed()
757
+ self.server._finish_claim(self)
662
758
 
663
759
 
664
760
  def _accepts_context(handler: Handler) -> bool:
@@ -137,6 +137,7 @@ class ContinuumWebSocketClient:
137
137
  self._heartbeat_stop = threading.Event()
138
138
  self._decoder = _stomp.StompDecoder()
139
139
  self._reconnect_attempt = 0
140
+ self._task_server_token: str | None = None
140
141
  self._client_heartbeat_ms = self._options.heartbeat_ms
141
142
  self._generation = 0
142
143
 
@@ -603,24 +604,35 @@ class ContinuumWebSocketClient:
603
604
  return
604
605
  backoff = min(backoff * 2, self._options.max_reconnect_backoff)
605
606
 
607
+ def set_task_server_token(self, token: str | None) -> None:
608
+ """Secret from registration. The next CONNECT sends it, including after reconnect."""
609
+ self._task_server_token = token or None
610
+
606
611
  def _open_and_handshake(self) -> None:
607
612
  url = derive_websocket_url(self._base_url, self._options.endpoint)
613
+ headers = {"Api-Key": self._api_key}
614
+ token = self._task_server_token
615
+ if token:
616
+ headers["Task-Server-Token"] = token
608
617
  socket = self._connector(
609
618
  url,
610
- {"Api-Key": self._api_key},
619
+ headers,
611
620
  self._options.connect_timeout,
612
621
  )
613
622
  self._socket = socket
614
623
  self._decoder = _stomp.StompDecoder()
615
624
  self._generation += 1
616
625
  cx = max(0, int(self._options.heartbeat_ms))
626
+ connect_headers = {
627
+ "accept-version": "1.2,1.1,1.0",
628
+ "heart-beat": f"{cx},{cx}",
629
+ "host": urlparse(url).hostname or "localhost",
630
+ }
631
+ if token:
632
+ connect_headers["Task-Server-Token"] = token
617
633
  connect = _stomp.StompFrame(
618
634
  command="CONNECT",
619
- headers={
620
- "accept-version": "1.2,1.1,1.0",
621
- "heart-beat": f"{cx},{cx}",
622
- "host": urlparse(url).hostname or "localhost",
623
- },
635
+ headers=connect_headers,
624
636
  )
625
637
  self._send_raw(_stomp.encode_frame(connect))
626
638
  connected = self._recv_until_command("CONNECTED", timeout=self._options.connect_timeout)
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
4
4
 
5
5
  [project]
6
6
  name = "continuum-task-server-sdk"
7
- version = "1.2.0"
7
+ version = "1.5.0"
8
8
  description = "Python SDK for the Continuum Task Server"
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.10"