continuum-task-server-sdk 1.1.0__py3-none-any.whl
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.
- continuum_task_server/__init__.py +74 -0
- continuum_task_server/_http.py +142 -0
- continuum_task_server/_stomp.py +137 -0
- continuum_task_server/client.py +351 -0
- continuum_task_server/exceptions.py +94 -0
- continuum_task_server/models.py +197 -0
- continuum_task_server/server.py +511 -0
- continuum_task_server/websocket.py +852 -0
- continuum_task_server_sdk-1.1.0.dist-info/METADATA +258 -0
- continuum_task_server_sdk-1.1.0.dist-info/RECORD +11 -0
- continuum_task_server_sdk-1.1.0.dist-info/WHEEL +4 -0
|
@@ -0,0 +1,852 @@
|
|
|
1
|
+
"""Synchronous STOMP-over-WebSocket client for the Queue API worker protocol."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import json
|
|
6
|
+
import logging
|
|
7
|
+
import threading
|
|
8
|
+
import time
|
|
9
|
+
import uuid
|
|
10
|
+
from collections import OrderedDict
|
|
11
|
+
from collections.abc import Callable
|
|
12
|
+
from dataclasses import dataclass
|
|
13
|
+
from typing import Any, Protocol
|
|
14
|
+
from urllib.parse import urljoin, urlparse
|
|
15
|
+
from uuid import UUID
|
|
16
|
+
|
|
17
|
+
from . import _stomp
|
|
18
|
+
from ._http import drop_none
|
|
19
|
+
from .client import _encode_input
|
|
20
|
+
from .exceptions import (
|
|
21
|
+
AmbiguousCommandError,
|
|
22
|
+
CommandTimeoutError,
|
|
23
|
+
ContinuumError,
|
|
24
|
+
ProtocolError,
|
|
25
|
+
TransportError,
|
|
26
|
+
error_for_status,
|
|
27
|
+
)
|
|
28
|
+
from .models import (
|
|
29
|
+
EnqueueAndSubscribeWaitResult,
|
|
30
|
+
EventAckResult,
|
|
31
|
+
QueueEvent,
|
|
32
|
+
QueueEventType,
|
|
33
|
+
QueueItem,
|
|
34
|
+
TaskStatus,
|
|
35
|
+
WaitMode,
|
|
36
|
+
WaitSubscriptionResult,
|
|
37
|
+
WorkSubscriptionResult,
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
logger = logging.getLogger("continuum_task_server")
|
|
41
|
+
|
|
42
|
+
Connector = Callable[[str, dict[str, str], float], "SocketLike"]
|
|
43
|
+
EventCallback = Callable[[QueueEvent], None]
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class SocketLike(Protocol):
|
|
47
|
+
def send(self, data: str) -> None: ...
|
|
48
|
+
|
|
49
|
+
def recv(self) -> str | bytes: ...
|
|
50
|
+
|
|
51
|
+
def close(self) -> None: ...
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
@dataclass
|
|
55
|
+
class WebSocketOptions:
|
|
56
|
+
"""Tunables for :class:`ContinuumWebSocketClient`."""
|
|
57
|
+
|
|
58
|
+
endpoint: str = "/ws/queue"
|
|
59
|
+
connect_timeout: float = 10.0
|
|
60
|
+
request_timeout: float = 30.0
|
|
61
|
+
heartbeat_ms: int = 10_000
|
|
62
|
+
reconnect_backoff: float = 1.0
|
|
63
|
+
max_reconnect_backoff: float = 30.0
|
|
64
|
+
auto_reconnect: bool = True
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def derive_websocket_url(base_url: str, endpoint: str = "/ws/queue") -> str:
|
|
68
|
+
"""Map an HTTP(S) Queue API base URL to the native STOMP WebSocket URL.
|
|
69
|
+
|
|
70
|
+
``http://`` stays ``ws://`` and ``https://`` stays ``wss://``. The path
|
|
71
|
+
``endpoint`` is appended to the origin (not to extra HTTP path prefixes).
|
|
72
|
+
"""
|
|
73
|
+
if not base_url:
|
|
74
|
+
raise ValueError("base_url must not be empty")
|
|
75
|
+
parsed = urlparse(base_url)
|
|
76
|
+
scheme = parsed.scheme.lower()
|
|
77
|
+
if scheme in ("https", "wss"):
|
|
78
|
+
ws_scheme = "wss"
|
|
79
|
+
elif scheme in ("http", "ws"):
|
|
80
|
+
ws_scheme = "ws"
|
|
81
|
+
else:
|
|
82
|
+
raise ValueError(f"unsupported URL scheme: {parsed.scheme!r}")
|
|
83
|
+
path = endpoint if endpoint.startswith("/") else "/" + endpoint
|
|
84
|
+
netloc = parsed.netloc
|
|
85
|
+
if not netloc:
|
|
86
|
+
raise ValueError(f"base_url has no host: {base_url!r}")
|
|
87
|
+
return urljoin(f"{ws_scheme}://{netloc}", path)
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
class _Pending:
|
|
91
|
+
def __init__(self) -> None:
|
|
92
|
+
self.event = threading.Event()
|
|
93
|
+
self.payload: dict[str, Any] | None = None
|
|
94
|
+
self.error: Exception | None = None
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
class ContinuumWebSocketClient:
|
|
98
|
+
"""Synchronous Queue API STOMP session.
|
|
99
|
+
|
|
100
|
+
Authenticate the WebSocket upgrade with ``Api-Key``. Commands are correlated
|
|
101
|
+
by ``requestId``. ``work.available`` hints are not reservations; ``claim``
|
|
102
|
+
remains the arbiter. Durable ``queue.completed`` deliveries must be
|
|
103
|
+
acknowledged with :meth:`ack_event`.
|
|
104
|
+
"""
|
|
105
|
+
|
|
106
|
+
def __init__(
|
|
107
|
+
self,
|
|
108
|
+
base_url: str,
|
|
109
|
+
api_key: str,
|
|
110
|
+
*,
|
|
111
|
+
options: WebSocketOptions | None = None,
|
|
112
|
+
connector: Connector | None = None,
|
|
113
|
+
) -> None:
|
|
114
|
+
if not base_url:
|
|
115
|
+
raise ValueError("base_url must not be empty")
|
|
116
|
+
if not api_key:
|
|
117
|
+
raise ValueError("api_key must not be empty")
|
|
118
|
+
self._base_url = base_url.rstrip("/")
|
|
119
|
+
self._api_key = api_key
|
|
120
|
+
self._options = options or WebSocketOptions()
|
|
121
|
+
self._connector = connector or _default_connector
|
|
122
|
+
self._socket: SocketLike | None = None
|
|
123
|
+
self._send_lock = threading.Lock()
|
|
124
|
+
self._state_lock = threading.Lock()
|
|
125
|
+
self._pending: dict[str, _Pending] = {}
|
|
126
|
+
self._events: list[QueueEvent] = []
|
|
127
|
+
self._events_cv = threading.Condition()
|
|
128
|
+
self._event_callbacks: list[EventCallback] = []
|
|
129
|
+
self._work_subscriptions: set[str] = set()
|
|
130
|
+
self._wait_subscribers: set[str] = set()
|
|
131
|
+
self._seen_completed: OrderedDict[str, None] = OrderedDict()
|
|
132
|
+
self._connected = threading.Event()
|
|
133
|
+
self._closed = threading.Event()
|
|
134
|
+
self._receiver: threading.Thread | None = None
|
|
135
|
+
self._heartbeat_thread: threading.Thread | None = None
|
|
136
|
+
self._heartbeat_stop = threading.Event()
|
|
137
|
+
self._decoder = _stomp.StompDecoder()
|
|
138
|
+
self._reconnect_attempt = 0
|
|
139
|
+
self._client_heartbeat_ms = self._options.heartbeat_ms
|
|
140
|
+
self._generation = 0
|
|
141
|
+
|
|
142
|
+
@property
|
|
143
|
+
def connected(self) -> bool:
|
|
144
|
+
return self._connected.is_set() and not self._closed.is_set()
|
|
145
|
+
|
|
146
|
+
def connect(self) -> None:
|
|
147
|
+
"""Open the WebSocket and complete the STOMP handshake.
|
|
148
|
+
|
|
149
|
+
Idempotent while already connected. Blocks until ``CONNECTED`` or
|
|
150
|
+
``connect_timeout``.
|
|
151
|
+
"""
|
|
152
|
+
if self._closed.is_set():
|
|
153
|
+
self._closed.clear()
|
|
154
|
+
if self._connected.is_set() and self._socket is not None:
|
|
155
|
+
return
|
|
156
|
+
self._start_receiver()
|
|
157
|
+
if not self._connected.wait(timeout=self._options.connect_timeout):
|
|
158
|
+
raise TransportError("timed out connecting to Queue API WebSocket")
|
|
159
|
+
|
|
160
|
+
def close(self) -> None:
|
|
161
|
+
self._closed.set()
|
|
162
|
+
self._heartbeat_stop.set()
|
|
163
|
+
self._connected.clear()
|
|
164
|
+
self._fail_pending(TransportError("WebSocket client closed"))
|
|
165
|
+
sock = self._socket
|
|
166
|
+
self._socket = None
|
|
167
|
+
if sock is not None:
|
|
168
|
+
try:
|
|
169
|
+
sock.close()
|
|
170
|
+
except Exception:
|
|
171
|
+
logger.debug("error closing websocket", exc_info=True)
|
|
172
|
+
with self._events_cv:
|
|
173
|
+
self._events_cv.notify_all()
|
|
174
|
+
if self._receiver is not None and self._receiver is not threading.current_thread():
|
|
175
|
+
self._receiver.join(timeout=5.0)
|
|
176
|
+
if (
|
|
177
|
+
self._heartbeat_thread is not None
|
|
178
|
+
and self._heartbeat_thread is not threading.current_thread()
|
|
179
|
+
):
|
|
180
|
+
self._heartbeat_thread.join(timeout=2.0)
|
|
181
|
+
self._receiver = None
|
|
182
|
+
self._heartbeat_thread = None
|
|
183
|
+
|
|
184
|
+
def __enter__(self) -> ContinuumWebSocketClient:
|
|
185
|
+
self.connect()
|
|
186
|
+
return self
|
|
187
|
+
|
|
188
|
+
def __exit__(self, *exc: object) -> None:
|
|
189
|
+
self.close()
|
|
190
|
+
|
|
191
|
+
def on_event(self, callback: EventCallback) -> Callable[[], None]:
|
|
192
|
+
"""Register a callback invoked on the receiver thread. Returns unsubscribe."""
|
|
193
|
+
self._event_callbacks.append(callback)
|
|
194
|
+
|
|
195
|
+
def _off() -> None:
|
|
196
|
+
try:
|
|
197
|
+
self._event_callbacks.remove(callback)
|
|
198
|
+
except ValueError:
|
|
199
|
+
pass
|
|
200
|
+
|
|
201
|
+
return _off
|
|
202
|
+
|
|
203
|
+
def next_event(self, timeout: float | None = None) -> QueueEvent | None:
|
|
204
|
+
"""Block until the next event, or return ``None`` on timeout / close."""
|
|
205
|
+
deadline = None if timeout is None else time.monotonic() + timeout
|
|
206
|
+
with self._events_cv:
|
|
207
|
+
while True:
|
|
208
|
+
if self._events:
|
|
209
|
+
return self._events.pop(0)
|
|
210
|
+
if self._closed.is_set():
|
|
211
|
+
return None
|
|
212
|
+
remaining = None
|
|
213
|
+
if deadline is not None:
|
|
214
|
+
remaining = deadline - time.monotonic()
|
|
215
|
+
if remaining <= 0:
|
|
216
|
+
return None
|
|
217
|
+
self._events_cv.wait(timeout=remaining)
|
|
218
|
+
|
|
219
|
+
def enqueue(
|
|
220
|
+
self,
|
|
221
|
+
*,
|
|
222
|
+
task_name: str | None = None,
|
|
223
|
+
task_item_name: str | None = None,
|
|
224
|
+
input_data: Any = None,
|
|
225
|
+
priority: int | None = None,
|
|
226
|
+
parent_id: UUID | str | None = None,
|
|
227
|
+
request_id: str | None = None,
|
|
228
|
+
idempotency_token: str | None = None,
|
|
229
|
+
) -> QueueItem:
|
|
230
|
+
if task_name is None and task_item_name is None:
|
|
231
|
+
raise ValueError("at least one of task_name or task_item_name is required")
|
|
232
|
+
body = drop_none(
|
|
233
|
+
{
|
|
234
|
+
"parent_id": str(parent_id) if parent_id else None,
|
|
235
|
+
"taskName": task_name,
|
|
236
|
+
"taskItemName": task_item_name,
|
|
237
|
+
"priority": priority,
|
|
238
|
+
"inputData": _encode_input(input_data),
|
|
239
|
+
"idempotencyToken": idempotency_token,
|
|
240
|
+
}
|
|
241
|
+
)
|
|
242
|
+
payload = self._command("/app/enqueue", body, request_id=request_id, retry=False)
|
|
243
|
+
item = payload.get("item")
|
|
244
|
+
if item is None:
|
|
245
|
+
raise ContinuumError("enqueue succeeded without an item")
|
|
246
|
+
return QueueItem.model_validate(item)
|
|
247
|
+
|
|
248
|
+
def enqueue_and_subscribe_wait(
|
|
249
|
+
self,
|
|
250
|
+
*,
|
|
251
|
+
subscriber_id: str,
|
|
252
|
+
task_name: str | None = None,
|
|
253
|
+
task_item_name: str | None = None,
|
|
254
|
+
input_data: Any = None,
|
|
255
|
+
priority: int | None = None,
|
|
256
|
+
parent_id: UUID | str | None = None,
|
|
257
|
+
mode: WaitMode | str | None = None,
|
|
258
|
+
correlation: str | None = None,
|
|
259
|
+
request_id: str | None = None,
|
|
260
|
+
idempotency_token: str | None = None,
|
|
261
|
+
) -> EnqueueAndSubscribeWaitResult:
|
|
262
|
+
if task_name is None and task_item_name is None:
|
|
263
|
+
raise ValueError("at least one of task_name or task_item_name is required")
|
|
264
|
+
if not subscriber_id or not str(subscriber_id).strip():
|
|
265
|
+
raise ValueError("subscriber_id is required")
|
|
266
|
+
mode_value = mode.value if isinstance(mode, WaitMode) else mode
|
|
267
|
+
body = drop_none(
|
|
268
|
+
{
|
|
269
|
+
"parent_id": str(parent_id) if parent_id else None,
|
|
270
|
+
"taskName": task_name,
|
|
271
|
+
"taskItemName": task_item_name,
|
|
272
|
+
"priority": priority,
|
|
273
|
+
"inputData": _encode_input(input_data),
|
|
274
|
+
"subscriberId": subscriber_id,
|
|
275
|
+
"mode": mode_value,
|
|
276
|
+
"correlation": correlation,
|
|
277
|
+
"idempotencyToken": idempotency_token,
|
|
278
|
+
}
|
|
279
|
+
)
|
|
280
|
+
payload = self._command(
|
|
281
|
+
"/app/enqueueAndSubscribeWait", body, request_id=request_id, retry=False
|
|
282
|
+
)
|
|
283
|
+
result = EnqueueAndSubscribeWaitResult.model_validate(payload)
|
|
284
|
+
if result.item is None:
|
|
285
|
+
raise ContinuumError("enqueueAndSubscribeWait succeeded without an item")
|
|
286
|
+
with self._state_lock:
|
|
287
|
+
self._wait_subscribers.add(subscriber_id)
|
|
288
|
+
return result
|
|
289
|
+
|
|
290
|
+
def claim(self, task_name: str, *, request_id: str | None = None) -> QueueItem | None:
|
|
291
|
+
payload = self._command(
|
|
292
|
+
"/app/claim",
|
|
293
|
+
{"taskName": task_name},
|
|
294
|
+
request_id=request_id,
|
|
295
|
+
retry=False,
|
|
296
|
+
)
|
|
297
|
+
item = payload.get("item")
|
|
298
|
+
if item is None:
|
|
299
|
+
return None
|
|
300
|
+
return QueueItem.model_validate(item)
|
|
301
|
+
|
|
302
|
+
def heartbeat(self, queue_item_id: UUID | str, *, request_id: str | None = None) -> QueueItem:
|
|
303
|
+
payload = self._command(
|
|
304
|
+
"/app/heartbeat",
|
|
305
|
+
{"queueItemId": str(queue_item_id)},
|
|
306
|
+
request_id=request_id,
|
|
307
|
+
retry=True,
|
|
308
|
+
)
|
|
309
|
+
return QueueItem.model_validate(payload["item"])
|
|
310
|
+
|
|
311
|
+
def update_status(
|
|
312
|
+
self,
|
|
313
|
+
queue_item_id: UUID | str,
|
|
314
|
+
status: TaskStatus,
|
|
315
|
+
*,
|
|
316
|
+
output_data: Any = None,
|
|
317
|
+
request_id: str | None = None,
|
|
318
|
+
) -> QueueItem:
|
|
319
|
+
body: dict[str, Any] = {
|
|
320
|
+
"queueItemId": str(queue_item_id),
|
|
321
|
+
"status": status.value if isinstance(status, TaskStatus) else str(status),
|
|
322
|
+
}
|
|
323
|
+
encoded = _encode_input(output_data)
|
|
324
|
+
if encoded is not None:
|
|
325
|
+
body["outputData"] = encoded
|
|
326
|
+
payload = self._command("/app/status", body, request_id=request_id, retry=True)
|
|
327
|
+
return QueueItem.model_validate(payload["item"])
|
|
328
|
+
|
|
329
|
+
def subscribe_work(
|
|
330
|
+
self, task_name: str, *, request_id: str | None = None
|
|
331
|
+
) -> WorkSubscriptionResult:
|
|
332
|
+
payload = self._command(
|
|
333
|
+
"/app/subscribeWork",
|
|
334
|
+
{"taskName": task_name},
|
|
335
|
+
request_id=request_id,
|
|
336
|
+
retry=True,
|
|
337
|
+
)
|
|
338
|
+
with self._state_lock:
|
|
339
|
+
self._work_subscriptions.add(task_name)
|
|
340
|
+
return WorkSubscriptionResult.model_validate(payload)
|
|
341
|
+
|
|
342
|
+
def unsubscribe_work(
|
|
343
|
+
self, task_name: str, *, request_id: str | None = None
|
|
344
|
+
) -> WorkSubscriptionResult:
|
|
345
|
+
payload = self._command(
|
|
346
|
+
"/app/unsubscribeWork",
|
|
347
|
+
{"taskName": task_name},
|
|
348
|
+
request_id=request_id,
|
|
349
|
+
retry=True,
|
|
350
|
+
)
|
|
351
|
+
with self._state_lock:
|
|
352
|
+
self._work_subscriptions.discard(task_name)
|
|
353
|
+
return WorkSubscriptionResult.model_validate(payload)
|
|
354
|
+
|
|
355
|
+
def subscribe_wait(
|
|
356
|
+
self,
|
|
357
|
+
subscriber_id: str,
|
|
358
|
+
*,
|
|
359
|
+
mode: WaitMode | str | None = None,
|
|
360
|
+
correlation: str | None = None,
|
|
361
|
+
queue_item_ids: list[UUID | str] | None = None,
|
|
362
|
+
request_id: str | None = None,
|
|
363
|
+
) -> WaitSubscriptionResult:
|
|
364
|
+
ids = [str(qid) for qid in queue_item_ids] if queue_item_ids else None
|
|
365
|
+
mode_value = mode.value if isinstance(mode, WaitMode) else mode
|
|
366
|
+
payload = self._command(
|
|
367
|
+
"/app/subscribeWait",
|
|
368
|
+
drop_none(
|
|
369
|
+
{
|
|
370
|
+
"subscriberId": subscriber_id,
|
|
371
|
+
"mode": mode_value,
|
|
372
|
+
"correlation": correlation,
|
|
373
|
+
"queueItemIds": ids,
|
|
374
|
+
}
|
|
375
|
+
),
|
|
376
|
+
request_id=request_id,
|
|
377
|
+
retry=True,
|
|
378
|
+
)
|
|
379
|
+
with self._state_lock:
|
|
380
|
+
self._wait_subscribers.add(subscriber_id)
|
|
381
|
+
return WaitSubscriptionResult.model_validate(payload)
|
|
382
|
+
|
|
383
|
+
def bind_wait(
|
|
384
|
+
self, subscriber_id: str, *, request_id: str | None = None
|
|
385
|
+
) -> WaitSubscriptionResult:
|
|
386
|
+
"""Reconnect helper: bind this socket and replay unacked deliveries."""
|
|
387
|
+
return self.subscribe_wait(subscriber_id, request_id=request_id)
|
|
388
|
+
|
|
389
|
+
def ack_event(
|
|
390
|
+
self,
|
|
391
|
+
subscriber_id: str,
|
|
392
|
+
event_id: UUID | str,
|
|
393
|
+
*,
|
|
394
|
+
request_id: str | None = None,
|
|
395
|
+
) -> EventAckResult:
|
|
396
|
+
payload = self._command(
|
|
397
|
+
"/app/ackEvent",
|
|
398
|
+
{"subscriberId": subscriber_id, "eventId": str(event_id)},
|
|
399
|
+
request_id=request_id,
|
|
400
|
+
retry=True,
|
|
401
|
+
)
|
|
402
|
+
return EventAckResult.model_validate(payload)
|
|
403
|
+
|
|
404
|
+
def _command(
|
|
405
|
+
self,
|
|
406
|
+
destination: str,
|
|
407
|
+
body: dict[str, Any],
|
|
408
|
+
*,
|
|
409
|
+
request_id: str | None,
|
|
410
|
+
retry: bool,
|
|
411
|
+
) -> dict[str, Any]:
|
|
412
|
+
if self._closed.is_set():
|
|
413
|
+
raise TransportError("WebSocket client is closed")
|
|
414
|
+
deadline = time.monotonic() + self._options.request_timeout
|
|
415
|
+
last_error: Exception | None = None
|
|
416
|
+
sent = False
|
|
417
|
+
while True:
|
|
418
|
+
remaining = deadline - time.monotonic()
|
|
419
|
+
if remaining <= 0:
|
|
420
|
+
if last_error is not None:
|
|
421
|
+
raise last_error
|
|
422
|
+
raise CommandTimeoutError(f"timed out waiting to send {destination}")
|
|
423
|
+
try:
|
|
424
|
+
self._wait_connected(min(remaining, self._options.connect_timeout))
|
|
425
|
+
rid = request_id or str(uuid.uuid4())
|
|
426
|
+
payload = dict(body)
|
|
427
|
+
payload["requestId"] = rid
|
|
428
|
+
pending = _Pending()
|
|
429
|
+
with self._state_lock:
|
|
430
|
+
self._pending[rid] = pending
|
|
431
|
+
self._send_json(destination, payload)
|
|
432
|
+
sent = True
|
|
433
|
+
wait_for = max(0.01, deadline - time.monotonic())
|
|
434
|
+
if not pending.event.wait(timeout=wait_for):
|
|
435
|
+
with self._state_lock:
|
|
436
|
+
self._pending.pop(rid, None)
|
|
437
|
+
raise CommandTimeoutError(f"timed out waiting for reply to {destination}")
|
|
438
|
+
if pending.error is not None:
|
|
439
|
+
raise pending.error
|
|
440
|
+
assert pending.payload is not None
|
|
441
|
+
return pending.payload
|
|
442
|
+
except AmbiguousCommandError:
|
|
443
|
+
raise
|
|
444
|
+
except CommandTimeoutError:
|
|
445
|
+
raise
|
|
446
|
+
except TransportError as exc:
|
|
447
|
+
last_error = exc
|
|
448
|
+
if sent and not retry:
|
|
449
|
+
raise AmbiguousCommandError(
|
|
450
|
+
f"{destination} was sent but the reply was lost; outcome is unknown"
|
|
451
|
+
) from exc
|
|
452
|
+
if not self._options.auto_reconnect or self._closed.is_set():
|
|
453
|
+
raise
|
|
454
|
+
request_id = None
|
|
455
|
+
sent = False
|
|
456
|
+
time.sleep(min(0.05, max(0.0, deadline - time.monotonic())))
|
|
457
|
+
|
|
458
|
+
def _wait_connected(self, timeout: float) -> None:
|
|
459
|
+
if self._connected.wait(timeout=timeout) and self._socket is not None:
|
|
460
|
+
return
|
|
461
|
+
if self._closed.is_set():
|
|
462
|
+
raise TransportError("WebSocket client is closed")
|
|
463
|
+
raise TransportError("WebSocket is not connected")
|
|
464
|
+
|
|
465
|
+
def _send_json(self, destination: str, payload: dict[str, Any]) -> None:
|
|
466
|
+
frame = _stomp.StompFrame(
|
|
467
|
+
command="SEND",
|
|
468
|
+
headers={
|
|
469
|
+
"destination": destination,
|
|
470
|
+
"content-type": "application/json",
|
|
471
|
+
},
|
|
472
|
+
body=json.dumps(payload, default=_json_default),
|
|
473
|
+
)
|
|
474
|
+
self._send_raw(_stomp.encode_frame(frame))
|
|
475
|
+
|
|
476
|
+
def _send_raw(self, data: str) -> None:
|
|
477
|
+
sock = self._socket
|
|
478
|
+
if sock is None:
|
|
479
|
+
raise TransportError("WebSocket is not connected")
|
|
480
|
+
with self._send_lock:
|
|
481
|
+
try:
|
|
482
|
+
sock.send(data)
|
|
483
|
+
except Exception as exc:
|
|
484
|
+
raise TransportError(f"WebSocket send failed: {exc}") from exc
|
|
485
|
+
|
|
486
|
+
def _start_receiver(self) -> None:
|
|
487
|
+
if self._receiver is not None and self._receiver.is_alive():
|
|
488
|
+
return
|
|
489
|
+
self._receiver = threading.Thread(
|
|
490
|
+
target=self._receiver_loop, name="continuum-stomp-recv", daemon=True
|
|
491
|
+
)
|
|
492
|
+
self._receiver.start()
|
|
493
|
+
|
|
494
|
+
def _receiver_loop(self) -> None:
|
|
495
|
+
backoff = self._options.reconnect_backoff
|
|
496
|
+
while not self._closed.is_set():
|
|
497
|
+
try:
|
|
498
|
+
self._open_and_handshake()
|
|
499
|
+
self._reconnect_attempt = 0
|
|
500
|
+
backoff = self._options.reconnect_backoff
|
|
501
|
+
self._restore_subscriptions()
|
|
502
|
+
self._connected.set()
|
|
503
|
+
self._read_until_disconnect()
|
|
504
|
+
except Exception:
|
|
505
|
+
logger.info("WebSocket session ended", exc_info=True)
|
|
506
|
+
self._teardown_socket()
|
|
507
|
+
self._connected.clear()
|
|
508
|
+
self._fail_pending(TransportError("WebSocket disconnected"))
|
|
509
|
+
if self._closed.is_set() or not self._options.auto_reconnect:
|
|
510
|
+
with self._events_cv:
|
|
511
|
+
self._events_cv.notify_all()
|
|
512
|
+
return
|
|
513
|
+
self._reconnect_attempt += 1
|
|
514
|
+
logger.info(
|
|
515
|
+
"Reconnecting to Queue API WebSocket in %.1fs (attempt %d)",
|
|
516
|
+
backoff,
|
|
517
|
+
self._reconnect_attempt,
|
|
518
|
+
)
|
|
519
|
+
if self._closed.wait(timeout=backoff):
|
|
520
|
+
return
|
|
521
|
+
backoff = min(backoff * 2, self._options.max_reconnect_backoff)
|
|
522
|
+
|
|
523
|
+
def _open_and_handshake(self) -> None:
|
|
524
|
+
url = derive_websocket_url(self._base_url, self._options.endpoint)
|
|
525
|
+
socket = self._connector(
|
|
526
|
+
url,
|
|
527
|
+
{"Api-Key": self._api_key},
|
|
528
|
+
self._options.connect_timeout,
|
|
529
|
+
)
|
|
530
|
+
self._socket = socket
|
|
531
|
+
self._decoder = _stomp.StompDecoder()
|
|
532
|
+
self._generation += 1
|
|
533
|
+
cx = max(0, int(self._options.heartbeat_ms))
|
|
534
|
+
connect = _stomp.StompFrame(
|
|
535
|
+
command="CONNECT",
|
|
536
|
+
headers={
|
|
537
|
+
"accept-version": "1.2,1.1,1.0",
|
|
538
|
+
"heart-beat": f"{cx},{cx}",
|
|
539
|
+
"host": urlparse(url).hostname or "localhost",
|
|
540
|
+
},
|
|
541
|
+
)
|
|
542
|
+
self._send_raw(_stomp.encode_frame(connect))
|
|
543
|
+
connected = self._recv_until_command("CONNECTED", timeout=self._options.connect_timeout)
|
|
544
|
+
self._client_heartbeat_ms = _negotiate_heartbeat(
|
|
545
|
+
cx, connected.headers.get("heart-beat", "0,0")
|
|
546
|
+
)
|
|
547
|
+
self._subscribe_destinations(connected.headers.get("session"))
|
|
548
|
+
self._start_heartbeat()
|
|
549
|
+
logger.info("WebSocket connected to %s", url)
|
|
550
|
+
|
|
551
|
+
def _subscribe_destinations(self, session_id: str | None) -> None:
|
|
552
|
+
"""Subscribe to user destinations and Spring's resolved session queues."""
|
|
553
|
+
self._send_raw(
|
|
554
|
+
_stomp.encode_frame(
|
|
555
|
+
_stomp.StompFrame(
|
|
556
|
+
command="SUBSCRIBE",
|
|
557
|
+
headers={"id": "replies", "destination": "/user/queue/replies"},
|
|
558
|
+
)
|
|
559
|
+
)
|
|
560
|
+
)
|
|
561
|
+
self._send_raw(
|
|
562
|
+
_stomp.encode_frame(
|
|
563
|
+
_stomp.StompFrame(
|
|
564
|
+
command="SUBSCRIBE",
|
|
565
|
+
headers={"id": "events", "destination": "/user/queue/events"},
|
|
566
|
+
)
|
|
567
|
+
)
|
|
568
|
+
)
|
|
569
|
+
if not session_id:
|
|
570
|
+
return
|
|
571
|
+
self._send_raw(
|
|
572
|
+
_stomp.encode_frame(
|
|
573
|
+
_stomp.StompFrame(
|
|
574
|
+
command="SUBSCRIBE",
|
|
575
|
+
headers={
|
|
576
|
+
"id": "replies-session",
|
|
577
|
+
"destination": f"/queue/replies-user{session_id}",
|
|
578
|
+
},
|
|
579
|
+
)
|
|
580
|
+
)
|
|
581
|
+
)
|
|
582
|
+
self._send_raw(
|
|
583
|
+
_stomp.encode_frame(
|
|
584
|
+
_stomp.StompFrame(
|
|
585
|
+
command="SUBSCRIBE",
|
|
586
|
+
headers={
|
|
587
|
+
"id": "events-session",
|
|
588
|
+
"destination": f"/queue/events-user{session_id}",
|
|
589
|
+
},
|
|
590
|
+
)
|
|
591
|
+
)
|
|
592
|
+
)
|
|
593
|
+
|
|
594
|
+
def _restore_subscriptions(self) -> None:
|
|
595
|
+
# Runs on the receiver thread; must recv replies inline (not via
|
|
596
|
+
# ``_command``, which waits for this same thread to dispatch).
|
|
597
|
+
with self._state_lock:
|
|
598
|
+
work = list(self._work_subscriptions)
|
|
599
|
+
waits = list(self._wait_subscribers)
|
|
600
|
+
for task_name in work:
|
|
601
|
+
try:
|
|
602
|
+
self._command_inline("/app/subscribeWork", {"taskName": task_name})
|
|
603
|
+
except ContinuumError:
|
|
604
|
+
logger.error("Failed to restore work subscription %s", task_name, exc_info=True)
|
|
605
|
+
for subscriber_id in waits:
|
|
606
|
+
try:
|
|
607
|
+
self._command_inline("/app/subscribeWait", {"subscriberId": subscriber_id})
|
|
608
|
+
except ContinuumError:
|
|
609
|
+
logger.error("Failed to restore wait subscriber %s", subscriber_id, exc_info=True)
|
|
610
|
+
|
|
611
|
+
def _command_inline(self, destination: str, body: dict[str, Any]) -> dict[str, Any]:
|
|
612
|
+
rid = str(uuid.uuid4())
|
|
613
|
+
payload = dict(body)
|
|
614
|
+
payload["requestId"] = rid
|
|
615
|
+
self._send_json(destination, payload)
|
|
616
|
+
deadline = time.monotonic() + self._options.request_timeout
|
|
617
|
+
sock = self._socket
|
|
618
|
+
if sock is None:
|
|
619
|
+
raise TransportError("WebSocket is not connected")
|
|
620
|
+
while time.monotonic() < deadline:
|
|
621
|
+
try:
|
|
622
|
+
data = sock.recv()
|
|
623
|
+
except Exception as exc:
|
|
624
|
+
if _is_timeout(exc):
|
|
625
|
+
continue
|
|
626
|
+
raise TransportError(f"WebSocket recv failed: {exc}") from exc
|
|
627
|
+
if data is None or data == "":
|
|
628
|
+
raise TransportError("WebSocket closed during subscription restore")
|
|
629
|
+
for frame in self._decoder.feed(data):
|
|
630
|
+
if frame is None:
|
|
631
|
+
continue
|
|
632
|
+
if frame.command == "ERROR":
|
|
633
|
+
raise _error_frame(frame)
|
|
634
|
+
if frame.command != "MESSAGE":
|
|
635
|
+
continue
|
|
636
|
+
destination_header = frame.headers.get("destination", "")
|
|
637
|
+
try:
|
|
638
|
+
parsed = json.loads(frame.body) if frame.body else {}
|
|
639
|
+
except json.JSONDecodeError as exc:
|
|
640
|
+
raise ProtocolError(f"invalid JSON: {exc}") from exc
|
|
641
|
+
if "events" in destination_header:
|
|
642
|
+
self._handle_event(parsed)
|
|
643
|
+
continue
|
|
644
|
+
if str(parsed.get("requestId")) != rid:
|
|
645
|
+
self._handle_reply(parsed)
|
|
646
|
+
continue
|
|
647
|
+
status = int(parsed.get("status", 0))
|
|
648
|
+
if status != 0:
|
|
649
|
+
message = parsed.get("message") or f"command failed with status {status}"
|
|
650
|
+
raise error_for_status(status, message)
|
|
651
|
+
return parsed
|
|
652
|
+
raise CommandTimeoutError(f"timed out restoring {destination}")
|
|
653
|
+
|
|
654
|
+
def _read_until_disconnect(self) -> None:
|
|
655
|
+
sock = self._socket
|
|
656
|
+
if sock is None:
|
|
657
|
+
return
|
|
658
|
+
while not self._closed.is_set():
|
|
659
|
+
try:
|
|
660
|
+
data = sock.recv()
|
|
661
|
+
except Exception as exc:
|
|
662
|
+
if _is_timeout(exc):
|
|
663
|
+
continue
|
|
664
|
+
raise TransportError(f"WebSocket recv failed: {exc}") from exc
|
|
665
|
+
if data is None or data == "":
|
|
666
|
+
raise TransportError("WebSocket closed by peer")
|
|
667
|
+
self._dispatch_bytes(data)
|
|
668
|
+
|
|
669
|
+
def _recv_until_command(self, command: str, *, timeout: float) -> _stomp.StompFrame:
|
|
670
|
+
sock = self._socket
|
|
671
|
+
if sock is None:
|
|
672
|
+
raise TransportError("WebSocket is not connected")
|
|
673
|
+
deadline = time.monotonic() + timeout
|
|
674
|
+
while time.monotonic() < deadline:
|
|
675
|
+
try:
|
|
676
|
+
data = sock.recv()
|
|
677
|
+
except Exception as exc:
|
|
678
|
+
if _is_timeout(exc):
|
|
679
|
+
continue
|
|
680
|
+
raise TransportError(f"WebSocket recv failed: {exc}") from exc
|
|
681
|
+
if data is None or data == "":
|
|
682
|
+
raise TransportError("WebSocket closed during handshake")
|
|
683
|
+
for frame in self._decoder.feed(data):
|
|
684
|
+
if frame is None:
|
|
685
|
+
continue
|
|
686
|
+
if frame.command == "ERROR":
|
|
687
|
+
raise _error_frame(frame)
|
|
688
|
+
if frame.command == command:
|
|
689
|
+
return frame
|
|
690
|
+
self._handle_frame(frame)
|
|
691
|
+
raise TransportError(f"timed out waiting for STOMP {command}")
|
|
692
|
+
|
|
693
|
+
def _dispatch_bytes(self, data: str | bytes) -> None:
|
|
694
|
+
for frame in self._decoder.feed(data):
|
|
695
|
+
if frame is None:
|
|
696
|
+
continue
|
|
697
|
+
self._handle_frame(frame)
|
|
698
|
+
|
|
699
|
+
def _handle_frame(self, frame: _stomp.StompFrame) -> None:
|
|
700
|
+
if frame.command == "ERROR":
|
|
701
|
+
exc = _error_frame(frame)
|
|
702
|
+
self._fail_pending(exc)
|
|
703
|
+
raise exc
|
|
704
|
+
if frame.command != "MESSAGE":
|
|
705
|
+
return
|
|
706
|
+
destination = frame.headers.get("destination", "")
|
|
707
|
+
try:
|
|
708
|
+
payload = json.loads(frame.body) if frame.body else {}
|
|
709
|
+
except json.JSONDecodeError as exc:
|
|
710
|
+
raise ProtocolError(f"invalid JSON on {destination}: {exc}") from exc
|
|
711
|
+
if "replies" in destination:
|
|
712
|
+
self._handle_reply(payload)
|
|
713
|
+
elif "events" in destination:
|
|
714
|
+
self._handle_event(payload)
|
|
715
|
+
|
|
716
|
+
def _handle_reply(self, payload: dict[str, Any]) -> None:
|
|
717
|
+
request_id = payload.get("requestId")
|
|
718
|
+
pending: _Pending | None = None
|
|
719
|
+
if request_id is not None:
|
|
720
|
+
with self._state_lock:
|
|
721
|
+
pending = self._pending.pop(str(request_id), None)
|
|
722
|
+
status = int(payload.get("status", 0))
|
|
723
|
+
if status != 0:
|
|
724
|
+
message = payload.get("message") or f"command failed with status {status}"
|
|
725
|
+
error = error_for_status(status, message)
|
|
726
|
+
if pending is not None:
|
|
727
|
+
pending.error = error
|
|
728
|
+
pending.event.set()
|
|
729
|
+
return
|
|
730
|
+
if pending is not None:
|
|
731
|
+
pending.payload = payload
|
|
732
|
+
pending.event.set()
|
|
733
|
+
|
|
734
|
+
def _handle_event(self, payload: dict[str, Any]) -> None:
|
|
735
|
+
event_body = payload.get("event")
|
|
736
|
+
if not event_body:
|
|
737
|
+
return
|
|
738
|
+
event = QueueEvent.model_validate(event_body)
|
|
739
|
+
if event.event_type == QueueEventType.COMPLETED:
|
|
740
|
+
key = str(event.event_id)
|
|
741
|
+
if key in self._seen_completed:
|
|
742
|
+
return
|
|
743
|
+
self._seen_completed[key] = None
|
|
744
|
+
while len(self._seen_completed) > 1024:
|
|
745
|
+
self._seen_completed.popitem(last=False)
|
|
746
|
+
for callback in list(self._event_callbacks):
|
|
747
|
+
try:
|
|
748
|
+
callback(event)
|
|
749
|
+
except Exception:
|
|
750
|
+
logger.exception("event callback failed")
|
|
751
|
+
with self._events_cv:
|
|
752
|
+
self._events.append(event)
|
|
753
|
+
self._events_cv.notify_all()
|
|
754
|
+
|
|
755
|
+
def _fail_pending(self, error: Exception) -> None:
|
|
756
|
+
with self._state_lock:
|
|
757
|
+
pending = list(self._pending.values())
|
|
758
|
+
self._pending.clear()
|
|
759
|
+
for item in pending:
|
|
760
|
+
if item.error is None:
|
|
761
|
+
item.error = error
|
|
762
|
+
item.event.set()
|
|
763
|
+
|
|
764
|
+
def _teardown_socket(self) -> None:
|
|
765
|
+
self._heartbeat_stop.set()
|
|
766
|
+
sock = self._socket
|
|
767
|
+
self._socket = None
|
|
768
|
+
if sock is not None:
|
|
769
|
+
try:
|
|
770
|
+
sock.close()
|
|
771
|
+
except Exception:
|
|
772
|
+
logger.debug("error closing websocket after disconnect", exc_info=True)
|
|
773
|
+
thread = self._heartbeat_thread
|
|
774
|
+
self._heartbeat_thread = None
|
|
775
|
+
if thread is not None and thread is not threading.current_thread():
|
|
776
|
+
thread.join(timeout=1.0)
|
|
777
|
+
|
|
778
|
+
def _start_heartbeat(self) -> None:
|
|
779
|
+
self._heartbeat_stop = threading.Event()
|
|
780
|
+
if self._client_heartbeat_ms <= 0:
|
|
781
|
+
return
|
|
782
|
+
interval = self._client_heartbeat_ms / 1000.0
|
|
783
|
+
|
|
784
|
+
def _loop() -> None:
|
|
785
|
+
while not self._heartbeat_stop.wait(timeout=interval):
|
|
786
|
+
if self._closed.is_set() or self._socket is None:
|
|
787
|
+
return
|
|
788
|
+
try:
|
|
789
|
+
self._send_raw(_stomp.encode_heartbeat())
|
|
790
|
+
except Exception:
|
|
791
|
+
return
|
|
792
|
+
|
|
793
|
+
self._heartbeat_thread = threading.Thread(
|
|
794
|
+
target=_loop, name="continuum-stomp-heartbeat", daemon=True
|
|
795
|
+
)
|
|
796
|
+
self._heartbeat_thread.start()
|
|
797
|
+
|
|
798
|
+
|
|
799
|
+
def _negotiate_heartbeat(client_cx_ms: int, server_header: str) -> int:
|
|
800
|
+
"""Return the client send interval in milliseconds, or 0 to disable."""
|
|
801
|
+
parts = [p.strip() for p in server_header.split(",")]
|
|
802
|
+
server_cy = 0
|
|
803
|
+
if len(parts) >= 2:
|
|
804
|
+
try:
|
|
805
|
+
server_cy = int(parts[1])
|
|
806
|
+
except ValueError:
|
|
807
|
+
server_cy = 0
|
|
808
|
+
if client_cx_ms <= 0 or server_cy <= 0:
|
|
809
|
+
return 0
|
|
810
|
+
return max(client_cx_ms, server_cy)
|
|
811
|
+
|
|
812
|
+
|
|
813
|
+
def _is_timeout(exc: BaseException) -> bool:
|
|
814
|
+
name = type(exc).__name__.lower()
|
|
815
|
+
return "timeout" in name or isinstance(exc, TimeoutError)
|
|
816
|
+
|
|
817
|
+
|
|
818
|
+
def _error_frame(frame: _stomp.StompFrame) -> ContinuumError:
|
|
819
|
+
message = frame.headers.get("message") or frame.body or "STOMP ERROR"
|
|
820
|
+
lower = message.lower()
|
|
821
|
+
if "unauthorized" in lower or "401" in lower:
|
|
822
|
+
return error_for_status(401, message)
|
|
823
|
+
return ProtocolError(message, body=frame.body)
|
|
824
|
+
|
|
825
|
+
|
|
826
|
+
def _json_default(obj: Any) -> Any:
|
|
827
|
+
import datetime as _dt
|
|
828
|
+
import enum as _enum
|
|
829
|
+
|
|
830
|
+
if isinstance(obj, UUID):
|
|
831
|
+
return str(obj)
|
|
832
|
+
if isinstance(obj, _dt.datetime):
|
|
833
|
+
return obj.isoformat()
|
|
834
|
+
if isinstance(obj, _enum.Enum):
|
|
835
|
+
return obj.value
|
|
836
|
+
raise TypeError(f"Object of type {type(obj).__name__} is not JSON serializable")
|
|
837
|
+
|
|
838
|
+
|
|
839
|
+
def _default_connector(url: str, headers: dict[str, str], timeout: float) -> SocketLike:
|
|
840
|
+
try:
|
|
841
|
+
import websocket
|
|
842
|
+
except ImportError as exc: # pragma: no cover
|
|
843
|
+
raise TransportError(
|
|
844
|
+
"websocket-client is required for WebSocket mode; install continuum-task-server-sdk"
|
|
845
|
+
) from exc
|
|
846
|
+
header = [f"{key}: {value}" for key, value in headers.items()]
|
|
847
|
+
try:
|
|
848
|
+
sock = websocket.create_connection(url, header=header, timeout=timeout)
|
|
849
|
+
sock.settimeout(1.0)
|
|
850
|
+
return sock
|
|
851
|
+
except Exception as exc:
|
|
852
|
+
raise TransportError(f"WebSocket handshake failed for {url}: {exc}") from exc
|