continuum-task-server-sdk 1.4.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.4.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
@@ -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
 
@@ -408,6 +408,16 @@ class ContinuumClient:
408
408
  def base_url(self) -> str:
409
409
  return self._http.base_url
410
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
+
411
421
  def close(self) -> None:
412
422
  self._http.close()
413
423
 
@@ -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"
@@ -34,8 +34,9 @@ import httpx
34
34
 
35
35
  from .client import ContinuumClient
36
36
  from .context import TaskExecutionContext
37
- from .exceptions import ContinuumError
37
+ from .exceptions import ConflictError, ContinuumError
38
38
  from .models import QueueEvent, QueueEventType, QueueItem, TaskStatus, TransportMode
39
+ from .presence import TaskServerPresence
39
40
  from .websocket import ContinuumWebSocketClient, WebSocketOptions
40
41
 
41
42
  logger = logging.getLogger("continuum_task_server")
@@ -92,6 +93,10 @@ class TaskServer:
92
93
  transport: TransportMode | str | httpx.BaseTransport | None = None,
93
94
  websocket_options: WebSocketOptions | None = None,
94
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,
95
100
  ) -> None:
96
101
  http_transport: httpx.BaseTransport | None = None
97
102
  if isinstance(transport, httpx.BaseTransport):
@@ -129,6 +134,19 @@ class TaskServer:
129
134
  self._base_url = base_url
130
135
  self._api_key = api_key
131
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)
132
150
 
133
151
  @property
134
152
  def client(self) -> ContinuumClient:
@@ -269,6 +287,8 @@ class TaskServer:
269
287
  )
270
288
 
271
289
  try:
290
+ if self._presence is not None:
291
+ self._presence.start()
272
292
  if self._mode is TransportMode.WEBSOCKET:
273
293
  self._websocket_loop()
274
294
  else:
@@ -310,7 +330,7 @@ class TaskServer:
310
330
  continue
311
331
 
312
332
  try:
313
- item = self._client.queue.claim(registration.name)
333
+ item = self._claim_http(registration.name)
314
334
  except ContinuumError as e:
315
335
  registration.semaphore.release()
316
336
  logger.error("claim(%s) failed: %s", registration.name, e)
@@ -336,6 +356,8 @@ class TaskServer:
336
356
  options=self._ws_options,
337
357
  )
338
358
  self._owns_ws = True
359
+ if self._presence is not None:
360
+ self._presence.attach(self._ws)
339
361
  return self._ws
340
362
 
341
363
  def _websocket_loop(self) -> None:
@@ -406,7 +428,7 @@ class TaskServer:
406
428
  if registration is None:
407
429
  return
408
430
  try:
409
- item = ws.claim(registration.name)
431
+ item = self._claim_socket(ws, registration.name)
410
432
  except ContinuumError as e:
411
433
  registration.semaphore.release()
412
434
  logger.error("claim(%s) failed: %s", registration.name, e)
@@ -630,6 +652,22 @@ class TaskServer:
630
652
  logger.error("update_status(%s, %s) failed: %s", queue_item_id, status.value, e)
631
653
  return None
632
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
+
633
671
  def _drain(self) -> None:
634
672
  if self._executor is None:
635
673
  return
@@ -645,6 +683,8 @@ class TaskServer:
645
683
  logger.debug("drain: handler future raised", exc_info=True)
646
684
  self._executor.shutdown(wait=False, cancel_futures=True)
647
685
  self._executor = None
686
+ if self._presence is not None:
687
+ self._presence.close()
648
688
  if self._ws is not None:
649
689
  self._ws.close()
650
690
  if self._owns_ws:
@@ -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.4.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"