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
4 changes: 2 additions & 2 deletions robosystems_client/clients/ledger_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -2503,8 +2503,8 @@ def create_report(
graph_id: str,
name: str,
mapping_id: str,
period_start: str,
period_end: str,
period_start: str | datetime.date,
period_end: str | datetime.date,
taxonomy_id: str = "rs-gaap",
period_type: str = "quarterly",
comparative: bool = True,
Expand Down
44 changes: 23 additions & 21 deletions robosystems_client/clients/operation_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -150,6 +150,8 @@ def monitor_operation(
result = OperationResult(operation_id=operation_id, status=OperationStatus.PENDING)
completed = False
error = None
# Set by a terminal event, or when the stream ends without one.
settled = threading.Event()

# Set up SSE connection with event replay from the beginning
# This handles the race condition where the operation may have already completed.
Expand Down Expand Up @@ -189,6 +191,7 @@ def on_operation_completed(data):
result.completed_at = datetime.now()
result.execution_time_ms = data.get("execution_time_ms")
completed = True
settled.set()

def on_operation_error(err):
nonlocal completed, error
Expand All @@ -197,12 +200,14 @@ def on_operation_error(err):
result.completed_at = datetime.now()
error = Exception(result.error)
completed = True
settled.set()

def on_operation_cancelled(_data=None):
nonlocal completed
result.status = OperationStatus.CANCELLED
result.completed_at = datetime.now()
completed = True
settled.set()

def on_connection_error(err):
nonlocal completed, error
Expand All @@ -211,6 +216,7 @@ def on_connection_error(err):
result.completed_at = datetime.now()
error = err if isinstance(err, Exception) else Exception(str(err))
completed = True
settled.set()

# Register event handlers
sse_client.on(EventType.OPERATION_STARTED.value, on_operation_started)
Expand All @@ -224,30 +230,28 @@ 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
def read_stream():
try:
sse_client.connect(operation_id)
finally:
settled.set()

# Connect and monitor. Registered first so cancel_operation() can close
# the stream; connect() blocks until the stream ends (or never opens).
# the stream.
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 options.timeout:
# connect() blocks on the socket, and closing the stream from another
# thread does not wake that read, so the stream is read on a worker
# and the timeout is kept here.
threading.Thread(target=read_stream, daemon=True).start()
if not settled.wait(options.timeout):
raise TimeoutError(
f"Operation {operation_id} timed out after {options.timeout}s"
)
else:
read_stream()

if not completed:
# The stream ended without a terminal event: no verdict to report,
Expand All @@ -257,8 +261,6 @@ def on_timeout():
)

finally:
if timer:
timer.cancel()
# Clean up with thread safety
with self._lock:
if operation_id in self.active_operations:
Expand Down
8 changes: 8 additions & 0 deletions robosystems_client/graphql/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -90,6 +90,14 @@ def __init__(
# use cases, but we keep the routing symmetric with the TS
# client so a caller that forwards a JWT (e.g. a backend
# proxying a request-scoped token) still works.
#
# The resolved token replaces any credential the static headers carry,
# as the REST writes do, so exactly one is sent.
self._headers = {
k: v
for k, v in self._headers.items()
if k.lower() not in ("x-api-key", "authorization")
}
if token.startswith("rfs"):
self._headers["X-API-Key"] = token
else:
Expand Down
33 changes: 33 additions & 0 deletions tests/test_auth_header_resolution.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
from robosystems_client.clients.operation_client import OperationClient
from robosystems_client.clients.retry import RetryingClient
from robosystems_client.clients.sse_client import SSEClient, event_error_message
from robosystems_client.graphql.client import GraphQLClient
from robosystems_client.clients.token_utils import (
apply_auth_header,
resolve_auth_headers,
Expand Down Expand Up @@ -153,3 +154,35 @@ def test_status_call_uses_provider_credential(self, mock_config):
passed = mock_get.call_args.kwargs["client"]
assert passed.get_httpx_client().headers["X-API-Key"] == "rfs_fresh"
assert isinstance(passed.get_httpx_client(), RetryingClient)


@pytest.mark.unit
class TestGraphQLClientCredential:
"""GraphQL reads send exactly one credential, as the REST writes do."""

def test_resolved_jwt_replaces_a_static_api_key(self):
# A facade with a static `rfs` key in its headers and a token_provider
# handing out JWTs: the resolved JWT is the one sent.
client = GraphQLClient(
"http://localhost:8000",
token="eyJ.jwt",
headers={"X-API-Key": "rfs_static", "X-Trace": "1"},
)

assert client._headers == {
"Content-Type": "application/json",
"X-Trace": "1",
"Authorization": "Bearer eyJ.jwt",
}

def test_resolved_api_key_replaces_a_static_bearer(self):
client = GraphQLClient(
"http://localhost:8000",
token="rfs_key",
headers={"authorization": "Bearer stale"},
)

assert client._headers == {
"Content-Type": "application/json",
"X-API-Key": "rfs_key",
}
38 changes: 38 additions & 0 deletions tests/test_operation_client_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@

import threading
import time
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer

import pytest
from unittest.mock import Mock, patch, MagicMock
Expand Down Expand Up @@ -260,6 +261,43 @@ def test_monitor_timeout_closes_a_silent_stream(self, MockSSE, mock_config):
assert time.monotonic() - started < 2
assert "op-silent" not in client.active_operations

def test_monitor_timeout_fires_on_time_over_a_real_socket(self, mock_config):
"""A blocked socket read must not hold the timeout to the stream's own
30s read timeout: a server that keeps the stream open with keepalives
and never finishes is given up on at `timeout`."""
stop = threading.Event()

class KeepaliveStream(BaseHTTPRequestHandler):
def do_GET(self):
self.send_response(200)
self.send_header("Content-Type", "text/event-stream")
self.end_headers()
while not stop.is_set():
try:
self.wfile.write(b": keepalive\n\n")
self.wfile.flush()
except OSError:
return
stop.wait(0.5)

def log_message(self, *args):
pass

server = ThreadingHTTPServer(("127.0.0.1", 0), KeepaliveStream)
server.daemon_threads = True
threading.Thread(target=server.serve_forever, daemon=True).start()
try:
config = {**mock_config, "base_url": f"http://127.0.0.1:{server.server_port}"}
client = OperationClient(config)
started = time.monotonic()
with pytest.raises(TimeoutError, match="timed out after 1s"):
client.monitor_operation("op-open", MonitorOptions(timeout=1))
assert time.monotonic() - started < 5
finally:
stop.set()
server.shutdown()
server.server_close()

@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."""
Expand Down
Loading