diff --git a/CHANGELOG.md b/CHANGELOG.md index 60f9b4c..6626919 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -114,6 +114,32 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Fixed + +- **The ActionCable handshake is bounded, and a failed setup no longer leaks the + socket (#108).** `open_timeout` was passed to the WebSocket upgrade and + nothing else: the waits for `welcome` and `confirm_subscription` that follow + had no deadline at all, so a socket that upgraded and then went quiet hung the + caller indefinitely, and cancelling out of that hang left the upgraded socket + open. Connect, `welcome` and `confirm_subscription` are now one bounded setup + lifecycle governed by a new `setup_timeout` (defaulting to `open_timeout`, so + the timeout you already configure does cover protocol setup), and every socket + the stream allocates is closed on any failure, timeout or cancellation -- + including a `__aenter__` that raises, where `__aexit__` never runs. +- **A failed reconnect consumes the reconnect budget instead of escaping on the + first attempt (#108).** An `OSError` raised while reconnecting inside the + `ConnectionClosed` handler propagated straight out of the iterator, so a + stream configured with `max_reconnect_attempts=10` gave up after one. Transient + failures now spend the configured consecutive-attempt budget with backoff and + end in `ConnectionError: Stream lost after N reconnect attempts`; a permanent + refusal stops immediately with the new `StreamAuthError` (a `ConnectionError` + subclass, so existing handlers are unaffected) rather than retrying a rejected + key ten times. +- **`close()` retires the stream.** It is idempotent, closes the socket under a + bounded teardown timeout, and prevents any subsequent reconnect; `connect()` + on a closed stream raises instead of quietly opening a new socket. Reconnects + now close the socket they are replacing. + ## [1.12.6] - 2026-08-11 ### Changed diff --git a/oilpriceapi/__init__.py b/oilpriceapi/__init__.py index 0bd3ed7..8640979 100644 --- a/oilpriceapi/__init__.py +++ b/oilpriceapi/__init__.py @@ -46,6 +46,7 @@ PriceStream, PriceUpdate, RigCountUpdate, + StreamAuthError, StreamingNotInstalledError, StreamUpdate, ) @@ -82,6 +83,7 @@ "StreamUpdate", "PriceUpdate", "RigCountUpdate", + "StreamAuthError", "StreamingNotInstalledError", ] diff --git a/oilpriceapi/streaming/__init__.py b/oilpriceapi/streaming/__init__.py index 0483fe4..513f0bd 100644 --- a/oilpriceapi/streaming/__init__.py +++ b/oilpriceapi/streaming/__init__.py @@ -14,6 +14,7 @@ CHANNEL_NAME, AsyncStreamNamespace, PriceStream, + StreamAuthError, StreamingNotInstalledError, ) from .models import ( @@ -30,6 +31,7 @@ __all__ = [ "AsyncStreamNamespace", "PriceStream", + "StreamAuthError", "StreamingNotInstalledError", "CHANNEL_NAME", "StreamUpdate", diff --git a/oilpriceapi/streaming/client.py b/oilpriceapi/streaming/client.py index 16f4f32..286ed4c 100644 --- a/oilpriceapi/streaming/client.py +++ b/oilpriceapi/streaming/client.py @@ -19,6 +19,10 @@ 5. Broadcasts arrive as ``{"identifier": ..., "message": {...}}``. 6. Server periodically sends ``{"type": "ping"}`` keepalives (ignored). +The whole of steps 1-4 is bounded by ``setup_timeout`` (default: +``open_timeout``), and every socket allocated along the way is closed on any +failure or cancellation. + Requires the optional ``[stream]`` extra (``pip install oilpriceapi[stream]``). """ from __future__ import annotations @@ -29,7 +33,7 @@ import random import sys from types import TracebackType -from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, List, Optional, Type +from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, List, Optional, Tuple, Type from .models import StreamUpdate @@ -46,6 +50,24 @@ class StreamingNotInstalledError(ImportError): """Raised when the optional ``websockets`` dependency is missing.""" +class StreamAuthError(ConnectionError): + """A permanent streaming setup failure -- retrying will not fix it. + + Raised when the server refuses the connection outright (an ActionCable + ``disconnect`` frame) or rejects the subscription (``reject_subscription``): + a bad API key, a missing streaming entitlement, or an unknown channel. + Subclasses :class:`ConnectionError` so existing ``except ConnectionError`` + handlers keep working, while the reconnect loop can tell it apart from a + transient network failure and stop immediately instead of burning the + reconnect budget on a refusal that will not change. + """ + + +# Teardown is bounded too: a socket that refuses to close must not wedge the +# caller inside ``close()`` or inside the cleanup path of a failed setup. +TEARDOWN_TIMEOUT = 5.0 + + def _import_websockets() -> Any: """Import the ``websockets`` library, raising a friendly error if absent.""" try: @@ -59,6 +81,17 @@ def _import_websockets() -> Any: ) from exc +def _transient_errors() -> Tuple[Type[BaseException], ...]: + """Error types a reconnect should retry (as opposed to give up on). + + ``OSError`` covers refused/reset connections and the ``ConnectionError`` + the bounded setup raises on timeout; :class:`StreamAuthError` is a subclass + of it and is therefore matched *before* this tuple by the reconnect loop. + """ + websockets = _import_websockets() + return (OSError, asyncio.TimeoutError, websockets.exceptions.WebSocketException) + + class PriceStream: """An async-iterable handle over a single ActionCable subscription. @@ -66,6 +99,20 @@ class PriceStream: ActionCable handshake, and yields :class:`StreamUpdate` objects. On transient disconnects it reconnects with exponential backoff + jitter, transparently re-subscribing, up to ``max_reconnect_attempts``. + + **Timeouts.** ``open_timeout`` is handed to the transport and bounds the + WebSocket upgrade only. ``setup_timeout`` bounds the *whole* setup + lifecycle -- upgrade, ``welcome`` and ``confirm_subscription`` -- and + defaults to ``open_timeout``, so by default the configured timeout does + cover protocol setup and there is no unbounded wait anywhere in + :meth:`connect`. Pass ``setup_timeout`` explicitly when a slow server needs + longer for the handshake than for the upgrade. + + **Cleanup.** Every socket this stream allocates is closed on any failure or + cancellation, including a handshake that times out and a ``__aenter__`` + that raises (where ``__aexit__`` never runs). Once :meth:`close` has been + called the stream is retired: it will not reconnect and :meth:`connect` + raises. """ def __init__( @@ -81,6 +128,7 @@ def __init__( reconnect_max_delay: float = 30.0, ping_interval: Optional[float] = None, open_timeout: float = 10.0, + setup_timeout: Optional[float] = None, ) -> None: self._cable_url = cable_url self._api_key = api_key @@ -92,6 +140,8 @@ def __init__( self._reconnect_max_delay = reconnect_max_delay self._ping_interval = ping_interval self._open_timeout = open_timeout + # Default: the caller's open_timeout also bounds the handshake. + self._setup_timeout = open_timeout if setup_timeout is None else setup_timeout self._ws: Any = None self._closed = False @@ -107,10 +157,60 @@ def identifier(self) -> str: # Sort keys for a stable identifier (ActionCable matches on exact string). return json.dumps(ident, sort_keys=True) + @property + def setup_timeout(self) -> float: + """Deadline, in seconds, for connect + welcome + confirm_subscription.""" + return self._setup_timeout + # -- connection lifecycle --------------------------------------------- async def connect(self) -> None: - """Open the WebSocket and complete the ActionCable handshake.""" + """Open the WebSocket and complete the ActionCable handshake. + + The entire lifecycle is bounded by :attr:`setup_timeout`; on timeout, + failure or cancellation the socket allocated by this call is closed + before the error propagates, so no upgraded socket is ever orphaned. + """ + if self._closed: + raise ConnectionError( + "Stream is closed; open a new stream to reconnect." + ) + + # The socket lives in a box the *caller* of wait_for can reach, so the + # cleanup below runs outside the (possibly cancelled) setup coroutine. + box: Dict[str, Any] = {} + try: + ws = await asyncio.wait_for(self._setup(box), timeout=self._setup_timeout) + except asyncio.TimeoutError as exc: + await self._close_socket(box.get("ws")) + raise ConnectionError( + f"ActionCable setup timed out after {self._setup_timeout:g}s " + "(connect, welcome, confirm_subscription). Raise setup_timeout " + "if the server needs longer, or check the /cable endpoint." + ) from exc + except BaseException: + # Includes cancellation and a rejected subscription: close first. + await self._close_socket(box.get("ws")) + raise + + if self._closed: + # close() landed while the handshake was in flight. + await self._close_socket(ws) + raise ConnectionError("Stream is closed; open a new stream to reconnect.") + + self._ws = ws + self._subscribed = True + + async def _setup(self, box: Dict[str, Any]) -> Any: + """Allocate a socket and run the handshake on it. Bounded by connect().""" + ws = await self._open_socket() + box["ws"] = ws + await self._await_welcome(ws) + await self._subscribe(ws) + return ws + + async def _open_socket(self) -> Any: + """Open the raw WebSocket (the transport upgrade only).""" websockets = _import_websockets() # Auth via query param is the most portable across proxies; the server # also accepts the Authorization header (connection.rb find_verified_user). @@ -133,66 +233,79 @@ async def connect(self) -> None: "X-SDK-Version": SDK_VERSION, } - self._ws = await websockets.connect( + return await websockets.connect( url, additional_headers=headers, ping_interval=self._ping_interval, open_timeout=self._open_timeout, ) - await self._await_welcome() - await self._subscribe() - async def _await_welcome(self) -> None: + async def _await_welcome(self, ws: Any) -> None: """Wait for the ActionCable ``welcome`` frame before subscribing.""" while True: - raw = await self._ws.recv() + raw = await ws.recv() data = json.loads(raw) msg_type = data.get("type") if msg_type == "welcome": return if msg_type == "disconnect": - raise ConnectionError( + raise StreamAuthError( f"Server refused connection: {data.get('reason', 'unknown')}" ) # Ignore stray pings while waiting for welcome. - async def _subscribe(self) -> None: + async def _subscribe(self, ws: Any) -> None: """Send the subscribe command and await ``confirm_subscription``.""" - await self._ws.send( + await ws.send( json.dumps({"command": "subscribe", "identifier": self.identifier}) ) while True: - raw = await self._ws.recv() + raw = await ws.recv() data = json.loads(raw) msg_type = data.get("type") if msg_type == "confirm_subscription": - self._subscribed = True return if msg_type == "reject_subscription": - raise ConnectionError( + raise StreamAuthError( "Subscription rejected; confirm the API key and streaming entitlement at " "https://www.oilpriceapi.com/pricing." ) # Ignore pings / pre-confirmation noise. + async def _close_socket(self, ws: Any) -> None: + """Close one socket. Best-effort and bounded; never raises.""" + if ws is None: + return + try: + await asyncio.wait_for(ws.close(), timeout=TEARDOWN_TIMEOUT) + except Exception: # noqa: BLE001 - best-effort teardown + logger.debug("Failed to close websocket cleanly", exc_info=True) + async def close(self) -> None: - """Unsubscribe and close the underlying WebSocket.""" + """Unsubscribe, close the socket, and retire the stream. + + Idempotent. After this returns the stream will not reconnect, holds no + socket, and :meth:`connect` raises. + """ self._closed = True - if self._ws is not None: - try: - if self._subscribed: - await self._ws.send( - json.dumps({"command": "unsubscribe", "identifier": self.identifier}) - ) - except Exception: # noqa: BLE001 - best-effort teardown - logger.debug("Failed to send unsubscribe on close", exc_info=True) - try: - await self._ws.close() - except Exception: # noqa: BLE001 - best-effort teardown - logger.debug("Failed to close websocket cleanly", exc_info=True) - finally: - self._ws = None - self._subscribed = False + ws, self._ws = self._ws, None + subscribed, self._subscribed = self._subscribed, False + if ws is None: + return + try: + if subscribed: + await asyncio.wait_for( + ws.send( + json.dumps( + {"command": "unsubscribe", "identifier": self.identifier} + ) + ), + timeout=TEARDOWN_TIMEOUT, + ) + except Exception: # noqa: BLE001 - best-effort teardown + logger.debug("Failed to send unsubscribe on close", exc_info=True) + finally: + await self._close_socket(ws) # -- async context manager -------------------------------------------- @@ -223,29 +336,35 @@ async def _iterate(self) -> AsyncIterator[StreamUpdate]: attempts = 0 while not self._closed: + ws = self._ws + if ws is None: + break try: - raw = await self._ws.recv() + raw = await ws.recv() except connection_closed_ok: # Clean server-side close — end iteration, do not reconnect. break except connection_closed: if not self._auto_reconnect or self._closed: break - attempts += 1 - if attempts > self._max_reconnect_attempts: - raise ConnectionError( - f"Stream lost after {self._max_reconnect_attempts} reconnect attempts" - ) - await self._backoff(attempts) - await self._reconnect() + # Consume the configured budget here rather than letting a + # failed reconnect escape on the first attempt. + attempts = await self._reconnect_with_budget(attempts) + if self._closed or self._ws is None: + break continue - attempts = 0 # reset backoff after any successful receive + attempts = 0 # reset the budget after any successful receive data = json.loads(raw) update = self._dispatch(data) if update is not None: yield update + # Iteration finished for good (clean close, close() by the caller, or + # auto_reconnect disabled): release the socket rather than leaving it + # to garbage collection. + await self.close() + def _dispatch(self, data: Dict[str, Any]) -> Optional[StreamUpdate]: """Translate a raw ActionCable frame into a StreamUpdate (or None).""" msg_type = data.get("type") @@ -272,9 +391,50 @@ async def _backoff(self, attempt: int) -> None: logger.info("Reconnecting stream in %.2fs (attempt %d)", delay, attempt) await asyncio.sleep(delay) + async def _reconnect_with_budget(self, attempts: int) -> int: + """Reconnect, spending the bounded consecutive-failure budget. + + A transient failure (refused connection, a setup that timed out, a drop + mid-handshake) costs one attempt and is retried after backoff. A + permanent failure (:class:`StreamAuthError` -- bad key, no entitlement, + rejected subscription) stops at once with that error, because no number + of retries changes the answer. Raises ``ConnectionError`` once + ``max_reconnect_attempts`` consecutive attempts have been spent. + + Returns the number of consecutive attempts consumed so far, so the + caller can reset it on the next successful receive. + """ + last_exc: Optional[BaseException] = None + while not self._closed: + attempts += 1 + if attempts > self._max_reconnect_attempts: + raise ConnectionError( + f"Stream lost after {self._max_reconnect_attempts} reconnect attempts" + ) from last_exc + await self._backoff(attempts) + if self._closed: + break + try: + await self._reconnect() + except StreamAuthError: + raise + except _transient_errors() as exc: + last_exc = exc + logger.warning( + "Reconnect attempt %d/%d failed: %s", + attempts, + self._max_reconnect_attempts, + exc, + ) + continue + return attempts + return attempts + async def _reconnect(self) -> None: + """Drop the current socket (closing it) and run a fresh setup.""" + old, self._ws = self._ws, None self._subscribed = False - self._ws = None + await self._close_socket(old) await self.connect() @@ -304,6 +464,7 @@ def prices( reconnect_base_delay: float = 1.0, reconnect_max_delay: float = 30.0, open_timeout: float = 10.0, + setup_timeout: Optional[float] = None, ) -> PriceStream: """Open a price-update stream over ``EnergyPricesChannel``. @@ -316,7 +477,12 @@ def prices( max_reconnect_attempts: Give up after this many failures. reconnect_base_delay: Initial backoff delay (seconds). reconnect_max_delay: Maximum backoff delay (seconds). - open_timeout: Connection open timeout (seconds). + open_timeout: WebSocket upgrade timeout (seconds). Also the + default deadline for the whole ActionCable handshake. + setup_timeout: Deadline (seconds) for the complete setup + lifecycle -- upgrade, ``welcome`` and + ``confirm_subscription``. Defaults to ``open_timeout``; there + is no unbounded wait either way. Returns: A :class:`PriceStream` async context manager / iterator. @@ -344,4 +510,5 @@ def prices( reconnect_base_delay=reconnect_base_delay, reconnect_max_delay=reconnect_max_delay, open_timeout=open_timeout, + setup_timeout=setup_timeout, ) diff --git a/tests/unit/test_streaming_lifecycle.py b/tests/unit/test_streaming_lifecycle.py new file mode 100644 index 0000000..edcfaf1 --- /dev/null +++ b/tests/unit/test_streaming_lifecycle.py @@ -0,0 +1,516 @@ +""" +Lifecycle tests for the async WebSocket streaming client (issue #108). + +These cover the two defects the coverage review found in +``oilpriceapi/streaming/client.py``: + +1. ``open_timeout`` bounded only the WebSocket upgrade. The ActionCable + handshake that follows it -- waiting for ``welcome`` and then for + ``confirm_subscription`` -- had no deadline at all, so a socket that + upgrades and then goes quiet hangs the caller forever. Cancelling out of + that hang left the upgraded socket open. +2. A reconnect that fails inside the ``ConnectionClosed`` handler escaped the + iterator on the first attempt instead of consuming the configured + ``max_reconnect_attempts`` budget. + +Every await in this module carries its own deadline so a regression fails fast +instead of blocking CI. The ``websockets`` transport is fully mocked; nothing +here touches the network. +""" +from __future__ import annotations + +import asyncio +import json +from typing import Any, Dict, List, Optional + +import pytest + +from oilpriceapi.streaming.client import PriceStream + +# Hard ceiling for any await in this file. A correct implementation finishes +# these in milliseconds; a regression trips the deadline instead of hanging. +DEADLINE = 2.0 + + +# -------------------------------------------------------------------------- +# Fakes +# -------------------------------------------------------------------------- + +class ScriptedSocket: + """A fake ``websockets`` connection driven by a scripted frame list. + + Frame items: + * ``dict`` / ``str`` -- delivered by ``recv`` + * ``HANG`` -- ``recv`` never completes + * ``DROP`` -- ``recv`` raises ``ConnectionClosed`` + An exhausted script raises ``ConnectionClosedOK`` (graceful server close). + """ + + HANG = object() + DROP = object() + + def __init__(self, frames: Optional[List[Any]] = None) -> None: + self.frames: List[Any] = list(frames or []) + self.sent: List[Dict[str, Any]] = [] + self.closed = False + self.recv_started = asyncio.Event() + + async def recv(self) -> str: + self.recv_started.set() + import websockets + + if not self.frames: + raise websockets.exceptions.ConnectionClosedOK(None, None) + item = self.frames.pop(0) + if item is self.HANG: + # Upgraded, connected, and silent forever. + await asyncio.sleep(3600) + raise AssertionError("unreachable") # pragma: no cover + if item is self.DROP: + raise websockets.exceptions.ConnectionClosed(None, None) + if isinstance(item, dict): + return json.dumps(item) + return item + + async def send(self, message: str) -> None: + self.sent.append(json.loads(message)) + + async def close(self) -> None: + self.closed = True + + +class RecordingConnector: + """Stands in for ``websockets.connect``; records every socket handed out. + + ``results`` entries are either a ``ScriptedSocket`` to return or an + exception instance to raise (a transport-level connect failure). + """ + + def __init__(self, results: List[Any], default: Any = None) -> None: + self._results = list(results) + self._default = default + self.calls: List[Dict[str, Any]] = [] + self.allocated: List[ScriptedSocket] = [] + + async def __call__(self, url: str, **kwargs: Any) -> ScriptedSocket: + self.calls.append({"url": url, **kwargs}) + if self._results: + result = self._results.pop(0) + elif self._default is not None: + result = self._default() if callable(self._default) else self._default + else: + raise AssertionError("connector called more times than scripted") + if isinstance(result, BaseException): + raise result + self.allocated.append(result) + return result + + +def _patch_connect(monkeypatch: pytest.MonkeyPatch, connector: RecordingConnector) -> None: + import websockets + + monkeypatch.setattr(websockets, "connect", connector) + + +def _welcome() -> Dict[str, Any]: + return {"type": "welcome"} + + +def _confirm() -> Dict[str, Any]: + return {"type": "confirm_subscription"} + + +def assert_no_socket_leaked(connector: RecordingConnector) -> None: + """Every socket the connector handed out must have been closed.""" + leaked = [i for i, ws in enumerate(connector.allocated) if not ws.closed] + assert not leaked, ( + f"{len(leaked)} of {len(connector.allocated)} allocated socket(s) left " + f"open (indices {leaked})" + ) + + +def assert_no_pending_tasks() -> None: + pending = [ + t + for t in asyncio.all_tasks() + if t is not asyncio.current_task() and not t.done() + ] + assert not pending, f"{len(pending)} task(s) left pending: {pending}" + + +# -------------------------------------------------------------------------- +# 1. Bounded setup lifecycle +# -------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_upgraded_socket_with_no_welcome_is_bounded(monkeypatch, api_key): + """A socket that upgrades and never sends ``welcome`` must not hang.""" + ws = ScriptedSocket([ScriptedSocket.HANG]) + connector = RecordingConnector([ws]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream(cable_url="ws://h/cable", api_key=api_key, open_timeout=0.05) + + with pytest.raises(ConnectionError, match="timed out"): + await asyncio.wait_for(stream.connect(), timeout=DEADLINE) + + assert_no_socket_leaked(connector) + assert stream._ws is None + assert_no_pending_tasks() + + +@pytest.mark.asyncio +async def test_welcome_without_confirmation_is_bounded(monkeypatch, api_key): + """``welcome`` then silence must not hang waiting for the confirmation.""" + ws = ScriptedSocket([_welcome(), ScriptedSocket.HANG]) + connector = RecordingConnector([ws]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream(cable_url="ws://h/cable", api_key=api_key, open_timeout=0.05) + + with pytest.raises(ConnectionError, match="timed out"): + await asyncio.wait_for(stream.connect(), timeout=DEADLINE) + + # The subscribe command did go out; the server simply never confirmed it. + assert ws.sent and ws.sent[0]["command"] == "subscribe" + assert_no_socket_leaked(connector) + assert stream._ws is None + assert_no_pending_tasks() + + +@pytest.mark.asyncio +async def test_setup_timeout_overrides_open_timeout(monkeypatch, api_key): + """``setup_timeout`` bounds the whole handshake independently of open_timeout.""" + ws = ScriptedSocket([ScriptedSocket.HANG]) + connector = RecordingConnector([ws]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream( + cable_url="ws://h/cable", + api_key=api_key, + open_timeout=300.0, + setup_timeout=0.05, + ) + + with pytest.raises(ConnectionError, match="timed out"): + await asyncio.wait_for(stream.connect(), timeout=DEADLINE) + + # open_timeout is still handed to the transport for the upgrade itself. + assert connector.calls[0]["open_timeout"] == 300.0 + assert_no_socket_leaked(connector) + + +@pytest.mark.asyncio +async def test_rejected_subscription_closes_the_socket(monkeypatch, api_key): + ws = ScriptedSocket([_welcome(), {"type": "reject_subscription"}]) + connector = RecordingConnector([ws]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream(cable_url="ws://h/cable", api_key=api_key) + + with pytest.raises(ConnectionError, match="Subscription rejected"): + await asyncio.wait_for(stream.connect(), timeout=DEADLINE) + + assert_no_socket_leaked(connector) + assert stream._ws is None + + +@pytest.mark.asyncio +async def test_server_disconnect_during_handshake_closes_the_socket(monkeypatch, api_key): + ws = ScriptedSocket([{"type": "disconnect", "reason": "unauthorized"}]) + connector = RecordingConnector([ws]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream(cable_url="ws://h/cable", api_key=api_key) + + with pytest.raises(ConnectionError, match="unauthorized"): + await asyncio.wait_for(stream.connect(), timeout=DEADLINE) + + assert_no_socket_leaked(connector) + assert stream._ws is None + + +@pytest.mark.asyncio +async def test_failed_aenter_closes_the_socket(monkeypatch, api_key): + """``__aexit__`` never runs when ``__aenter__`` raises -- setup must clean up.""" + ws = ScriptedSocket([_welcome(), {"type": "reject_subscription"}]) + connector = RecordingConnector([ws]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream(cable_url="ws://h/cable", api_key=api_key) + + async def _enter() -> None: + async with stream: # pragma: no cover - body never reached + raise AssertionError("unreachable") + + with pytest.raises(ConnectionError, match="Subscription rejected"): + await asyncio.wait_for(_enter(), timeout=DEADLINE) + + assert_no_socket_leaked(connector) + assert stream._ws is None + + +@pytest.mark.asyncio +async def test_cancellation_mid_setup_closes_the_socket(monkeypatch, api_key): + """A caller who cancels during the handshake leaves no socket behind.""" + ws = ScriptedSocket([ScriptedSocket.HANG]) + connector = RecordingConnector([ws]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream(cable_url="ws://h/cable", api_key=api_key, setup_timeout=60.0) + + task = asyncio.create_task(stream.connect()) + await asyncio.wait_for(ws.recv_started.wait(), timeout=DEADLINE) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, timeout=DEADLINE) + + assert_no_socket_leaked(connector) + assert stream._ws is None + assert_no_pending_tasks() + + +# -------------------------------------------------------------------------- +# 2. Reconnect budget +# -------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_transient_reconnect_failure_consumes_the_budget(monkeypatch, api_key): + """An OSError on reconnect consumes one attempt, not the whole iterator.""" + first = ScriptedSocket([_welcome(), _confirm(), ScriptedSocket.DROP]) + connector = RecordingConnector( + [first], default=lambda: OSError("connection refused") + ) + _patch_connect(monkeypatch, connector) + + stream = PriceStream( + cable_url="ws://h/cable", + api_key=api_key, + max_reconnect_attempts=10, + reconnect_base_delay=0.0, + reconnect_max_delay=0.0, + ) + + async def _drain() -> None: + async for _ in stream: + pass + + with pytest.raises(ConnectionError, match="10 reconnect attempts"): + await asyncio.wait_for(_drain(), timeout=DEADLINE) + + # 1 initial connect + 10 reconnect attempts, the full configured budget. + assert len(connector.calls) == 11 + assert_no_socket_leaked(connector) + + +@pytest.mark.asyncio +async def test_reconnect_succeeds_within_the_budget(monkeypatch, api_key): + """Two transient failures then a success: iteration continues.""" + first = ScriptedSocket([_welcome(), _confirm(), ScriptedSocket.DROP]) + third = ScriptedSocket([_welcome(), _confirm()]) + connector = RecordingConnector( + [first, OSError("refused"), OSError("refused"), third] + ) + _patch_connect(monkeypatch, connector) + + stream = PriceStream( + cable_url="ws://h/cable", + api_key=api_key, + max_reconnect_attempts=5, + reconnect_base_delay=0.0, + reconnect_max_delay=0.0, + ) + + async def _drain() -> List[Any]: + return [u async for u in stream] + + updates = await asyncio.wait_for(_drain(), timeout=DEADLINE) + + assert updates == [] # third socket closes gracefully with no broadcasts + assert len(connector.calls) == 4 + assert_no_socket_leaked(connector) + + +@pytest.mark.asyncio +async def test_permanent_rejection_on_reconnect_stops_immediately(monkeypatch, api_key): + """An auth/entitlement rejection is permanent: stop, do not burn the budget.""" + first = ScriptedSocket([_welcome(), _confirm(), ScriptedSocket.DROP]) + second = ScriptedSocket([_welcome(), {"type": "reject_subscription"}]) + connector = RecordingConnector([first, second]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream( + cable_url="ws://h/cable", + api_key=api_key, + max_reconnect_attempts=10, + reconnect_base_delay=0.0, + reconnect_max_delay=0.0, + ) + + async def _drain() -> None: + async for _ in stream: + pass + + with pytest.raises(ConnectionError, match="Subscription rejected"): + await asyncio.wait_for(_drain(), timeout=DEADLINE) + + # One initial connect, exactly one reconnect, then a hard stop. + assert len(connector.calls) == 2 + assert_no_socket_leaked(connector) + + +@pytest.mark.asyncio +async def test_permanent_failure_is_a_distinguishable_error_type(monkeypatch, api_key): + """Callers can tell a permanent refusal from a transient network error.""" + from oilpriceapi.streaming.client import StreamAuthError + + assert issubclass(StreamAuthError, ConnectionError) + + ws = ScriptedSocket([_welcome(), {"type": "reject_subscription"}]) + connector = RecordingConnector([ws]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream(cable_url="ws://h/cable", api_key=api_key) + with pytest.raises(StreamAuthError) as excinfo: + await asyncio.wait_for(stream.connect(), timeout=DEADLINE) + + assert "pricing" in str(excinfo.value) + assert_no_socket_leaked(connector) + + +@pytest.mark.asyncio +async def test_reconnect_closes_the_previous_socket(monkeypatch, api_key): + first = ScriptedSocket([_welcome(), _confirm(), ScriptedSocket.DROP]) + second = ScriptedSocket([_welcome(), _confirm()]) + connector = RecordingConnector([first, second]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream( + cable_url="ws://h/cable", + api_key=api_key, + reconnect_base_delay=0.0, + reconnect_max_delay=0.0, + ) + + async def _drain() -> List[Any]: + return [u async for u in stream] + + await asyncio.wait_for(_drain(), timeout=DEADLINE) + + assert first.closed is True, "the dropped socket was never closed" + assert_no_socket_leaked(connector) + + +# -------------------------------------------------------------------------- +# 3. Close / cancel stop the lifecycle +# -------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_close_prevents_any_further_connect(monkeypatch, api_key): + ws = ScriptedSocket([_welcome(), _confirm()]) + connector = RecordingConnector([ws]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream(cable_url="ws://h/cable", api_key=api_key) + await asyncio.wait_for(stream.connect(), timeout=DEADLINE) + await asyncio.wait_for(stream.close(), timeout=DEADLINE) + + assert ws.closed is True + assert stream._ws is None + + # A closed stream must not silently open a new socket. + with pytest.raises(ConnectionError, match="closed"): + await asyncio.wait_for(stream.connect(), timeout=DEADLINE) + + # Iterating a closed stream yields nothing and reconnects nothing. + async def _drain() -> List[Any]: + return [u async for u in stream] + + assert await asyncio.wait_for(_drain(), timeout=DEADLINE) == [] + assert len(connector.calls) == 1 + assert_no_pending_tasks() + + +@pytest.mark.asyncio +async def test_close_is_idempotent(monkeypatch, api_key): + ws = ScriptedSocket([_welcome(), _confirm()]) + connector = RecordingConnector([ws]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream(cable_url="ws://h/cable", api_key=api_key) + await asyncio.wait_for(stream.connect(), timeout=DEADLINE) + await asyncio.wait_for(stream.close(), timeout=DEADLINE) + await asyncio.wait_for(stream.close(), timeout=DEADLINE) + + unsubscribes = [m for m in ws.sent if m.get("command") == "unsubscribe"] + assert len(unsubscribes) == 1 + assert_no_socket_leaked(connector) + + +@pytest.mark.asyncio +async def test_close_during_backoff_stops_reconnecting(monkeypatch, api_key): + """Closing while the iterator waits to reconnect must not open a new socket.""" + first = ScriptedSocket([_welcome(), _confirm(), ScriptedSocket.DROP]) + connector = RecordingConnector([first]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream( + cable_url="ws://h/cable", + api_key=api_key, + max_reconnect_attempts=10, + reconnect_base_delay=0.2, + reconnect_max_delay=0.2, + ) + + async def _drain() -> List[Any]: + return [u async for u in stream] + + task = asyncio.create_task(_drain()) + await asyncio.sleep(0.05) # let it drop and enter backoff + await asyncio.wait_for(stream.close(), timeout=DEADLINE) + + assert await asyncio.wait_for(task, timeout=DEADLINE) == [] + assert len(connector.calls) == 1, "reconnected after close()" + assert_no_socket_leaked(connector) + assert_no_pending_tasks() + + +@pytest.mark.asyncio +async def test_cancelling_an_active_stream_leaves_nothing_pending(monkeypatch, api_key): + first = ScriptedSocket([_welcome(), _confirm(), ScriptedSocket.HANG]) + connector = RecordingConnector([first]) + _patch_connect(monkeypatch, connector) + + stream = PriceStream(cable_url="ws://h/cable", api_key=api_key) + + async def _consume() -> None: + async with stream as s: + async for _ in s: + pass + + task = asyncio.create_task(_consume()) + await asyncio.wait_for(first.recv_started.wait(), timeout=DEADLINE) + await asyncio.sleep(0) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, timeout=DEADLINE) + + assert_no_socket_leaked(connector) + assert len(connector.calls) == 1 + assert_no_pending_tasks() + + +# -------------------------------------------------------------------------- +# 4. Sync/async parity +# -------------------------------------------------------------------------- + +def test_streaming_is_async_only(api_key): + """Streaming has no sync counterpart; pin that so parity cannot drift silently. + + If a sync stream is ever added, this test fails and forces the same bounded + setup + cleanup + reconnect-budget lifecycle to be applied to it. + """ + from oilpriceapi import OilPriceAPI + + client = OilPriceAPI(api_key=api_key) + assert not hasattr(client, "stream")