Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 22 additions & 3 deletions robosystems_client/clients/operation_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@

import asyncio
import logging
import threading
from dataclasses import dataclass
from typing import Dict, Any, Optional, Callable, List, cast
from datetime import datetime
Expand Down Expand Up @@ -67,7 +68,7 @@ class MonitorOptions:

on_progress: Optional[Callable[[OperationProgress], None]] = None
on_queue_update: Optional[Callable[[int, int], None]] = None
timeout: Optional[int] = None
timeout: Optional[float] = None # seconds
poll_interval: Optional[int] = None


Expand Down Expand Up @@ -132,8 +133,6 @@ def __init__(self, config: Dict[str, Any]):
self.headers["X-API-Key"] = self.token
self.active_operations: Dict[str, SSEClient] = {}
# Thread safety for operations tracking
import threading

self._lock = threading.Lock()

def monitor_operation(
Expand Down Expand Up @@ -225,13 +224,31 @@ def on_connection_error(err):
sse_client.on("error", on_connection_error)
sse_client.on("max_retries_exceeded", on_connection_error)

# connect() blocks, so the timeout closes the stream from a timer thread,
# which ends the read and returns control here.
timed_out = threading.Event()

def on_timeout():
timed_out.set()
sse_client.close()

timer = threading.Timer(options.timeout, on_timeout) if options.timeout else None

# Connect and monitor. Registered first so cancel_operation() can close
# the stream; connect() blocks until the stream ends (or never opens).
try:
with self._lock:
self.active_operations[operation_id] = sse_client
if timer:
timer.daemon = True
timer.start()
sse_client.connect(operation_id)

if not completed and timed_out.is_set():
raise TimeoutError(
f"Operation {operation_id} timed out after {options.timeout}s"
)

if not completed:
# The stream ended without a terminal event: no verdict to report,
# and nothing left to wait on.
Expand All @@ -240,6 +257,8 @@ def on_connection_error(err):
)

finally:
if timer:
timer.cancel()
# Clean up with thread safety
with self._lock:
if operation_id in self.active_operations:
Expand Down
4 changes: 4 additions & 0 deletions robosystems_client/clients/sse_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -268,6 +268,10 @@ def _handle_error(

time.sleep(delay_seconds)

# A close() during the backoff ends the reconnect.
if self.closed:
return

# Resume from last event if available
resume_from = 0
if self.last_event_id:
Expand Down
59 changes: 58 additions & 1 deletion tests/test_operation_client_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,9 @@
Dataclass and enum tests already exist in tests/test_operation_client.py.
"""

import threading
import time

import pytest
from unittest.mock import Mock, patch, MagicMock
from robosystems_client.clients.operation_client import (
Expand All @@ -15,7 +18,7 @@
OperationProgress,
MonitorOptions,
)
from robosystems_client.clients.sse_client import SSEClient
from robosystems_client.clients.sse_client import SSEClient, SSEConfig
from robosystems_client.models.cancel_operation_response_canceloperation import (
CancelOperationResponseCanceloperation,
)
Expand Down Expand Up @@ -238,6 +241,60 @@ def fake_connect(op_id):
# SSE client should have been closed
fake_sse.close.assert_called()

@patch("robosystems_client.clients.operation_client.SSEClient")
def test_monitor_timeout_closes_a_silent_stream(self, MockSSE, mock_config):
"""A stream with no terminal event is closed at `timeout` and raises."""
fake_sse = MagicMock(spec=SSEClient)
closed = threading.Event()

# connect() blocks until the stream is closed, like a live stream would.
fake_sse.connect.side_effect = lambda op_id: closed.wait(5)
fake_sse.close.side_effect = closed.set
MockSSE.return_value = fake_sse

client = OperationClient(mock_config)
started = time.monotonic()
with pytest.raises(TimeoutError, match="timed out after 0.1s"):
client.monitor_operation("op-silent", MonitorOptions(timeout=0.1))

assert time.monotonic() - started < 2
assert "op-silent" not in client.active_operations

@patch("robosystems_client.clients.operation_client.SSEClient")
def test_monitor_completion_beats_timeout(self, MockSSE, mock_config):
"""A run that finishes inside `timeout` returns its result."""
fake_sse = MagicMock(spec=SSEClient)
listeners = {}
fake_sse.on.side_effect = lambda event, handler: listeners.__setitem__(
event, handler
)
fake_sse.connect.side_effect = lambda op_id: listeners["operation_completed"](
{"result": {"ok": True}}
)
MockSSE.return_value = fake_sse

client = OperationClient(mock_config)
result = client.monitor_operation("op-fast", MonitorOptions(timeout=5))

assert result.status == OperationStatus.COMPLETED
assert result.result == {"ok": True}


@pytest.mark.unit
class TestSSEReconnect:
"""Test SSEClient's reconnect backoff."""

@patch("robosystems_client.clients.sse_client.time.sleep")
def test_close_during_backoff_ends_the_reconnect(self, mock_sleep):
"""A close() while the reconnect sleeps must not reopen the stream."""
sse = SSEClient(SSEConfig(base_url="http://localhost:8000", max_retries=3))
mock_sleep.side_effect = lambda _s: setattr(sse, "closed", True)

with patch.object(sse, "connect") as mock_connect:
sse._handle_error(RuntimeError("dropped"), "op-1", 0)

mock_connect.assert_not_called()


# ── get_operation_status ─────────────────────────────────────────────

Expand Down
Loading