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.
@@ -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