From f438c4df2a7a48901a2d32549b5b7078a6a009f1 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 16 Sep 2026 06:15:43 +0000 Subject: [PATCH 01/11] fix: clear 0.4.3 release blockers Co-authored-by: benjamin.burtenshaw --- docs/source/guides/catalog-discovery.md | 3 + src/openenv/auto/_discovery.py | 4 +- src/openenv/core/env_client.py | 91 ++++++--- src/openenv/discovery/models.py | 10 +- .../schemas/0.1-draft/catalog.schema.json | 35 +++- .../schemas/0.1-draft/declaration.schema.json | 14 +- .../0.1-draft/environment-card.schema.json | 14 +- tests/discovery/test_catalog_contract.py | 17 +- tests/envs/test_discovery.py | 19 ++ tests/test_core/test_generic_client.py | 177 +++++++++++++++--- 10 files changed, 317 insertions(+), 67 deletions(-) diff --git a/docs/source/guides/catalog-discovery.md b/docs/source/guides/catalog-discovery.md index f5b4e34431..d5a8862231 100644 --- a/docs/source/guides/catalog-discovery.md +++ b/docs/source/guides/catalog-discovery.md @@ -165,6 +165,9 @@ JSON Schema validation is necessary but is not the whole profile contract. The packaged schemas enforce object shape, required fields, relative-path safety, supported literals, conditional artifact/license-evidence presence, and exactly one orchestration interface. Other rules require semantic validation: +The relative-path profile permits printable UTF-8 (including spaces), but +rejects C0, DEL, and C1 control characters so untrusted metadata cannot forge +CLI or log lines. | Subject | Additional rule | |---------|-----------------| diff --git a/src/openenv/auto/_discovery.py b/src/openenv/auto/_discovery.py index 7b53fc6bcc..2f7a57dbdb 100644 --- a/src/openenv/auto/_discovery.py +++ b/src/openenv/auto/_discovery.py @@ -345,9 +345,11 @@ def _default_cache_file() -> Path: shared, world-writable temporary directory. A fixed path under the shared temp dir lets another local user pre-create the cache file and redirect discovery to attacker-controlled import paths (`import_module` on a cached `client_module_path`). + Per the XDG Base Directory specification, relative `XDG_CACHE_HOME` values are + ignored so an untrusted working tree cannot supply a victim-owned cache file. """ base = os.environ.get("XDG_CACHE_HOME") - root = Path(base) if base else Path.home() / ".cache" + root = Path(base) if base and Path(base).is_absolute() else Path.home() / ".cache" return root / "openenv" / "discovery_cache.json" diff --git a/src/openenv/core/env_client.py b/src/openenv/core/env_client.py index f15df6cb2f..744f160b0c 100644 --- a/src/openenv/core/env_client.py +++ b/src/openenv/core/env_client.py @@ -255,6 +255,15 @@ async def _best_effort_close(ws: ClientConnection) -> None: pass # Best effort +async def _best_effort_graceful_close(ws: ClientConnection) -> None: + """Notify the server, then close a socket without propagating failures.""" + try: + await ws.send(json.dumps({"type": "close"})) + except (Exception, asyncio.CancelledError): + pass # Best effort + await _best_effort_close(ws) + + class EnvClient(ABC, Generic[ActT, ObsT, StateT]): """ Async environment client for persistent sessions. @@ -538,6 +547,13 @@ async def _connect_async(self) -> "EnvClient": self._ws = None self._ws_loop = None + # A timed-out request drops its socket immediately but closes it in the + # background so the timeout itself remains prompt. Wait for that close + # before opening a replacement: the old server-side session continues + # occupying a capacity slot until the close handshake finishes, and + # many environments allow only one session. + await self._drain_pending_close_tasks() + try: self._start_provider_if_needed() except Exception: @@ -573,24 +589,53 @@ async def _connect_async(self) -> "EnvClient": def disconnect(self) -> Any: return self._dispatch(self._disconnect_async) + def _schedule_socket_close( + self, ws: ClientConnection, *, notify_server: bool = False + ) -> asyncio.Task[None]: + """Schedule and track a socket close until its handshake finishes.""" + close_coro = ( + _best_effort_graceful_close(ws) + if notify_server + else _best_effort_close(ws) + ) + close_task = asyncio.create_task(close_coro) + self._pending_close_tasks.add(close_task) + close_task.add_done_callback(self._pending_close_tasks.discard) + return close_task + async def _disconnect_async(self) -> None: """Close the WebSocket connection.""" if self._ws is not None: ws = self._ws ws_loop = self._ws_loop - same_loop = ws_loop is asyncio.get_running_loop() - try: - if same_loop: - await ws.send(json.dumps({"type": "close"})) - except Exception: - pass # Best effort - try: - if same_loop: - await ws.close() - except Exception: - pass + # Detach first so cancellation during the close handshake cannot + # leave a stale socket cached for a later operation. self._ws = None self._ws_loop = None + same_loop = ws_loop is asyncio.get_running_loop() + if same_loop: + # Keep ownership of the handshake if this caller is cancelled. + # A later close/reconnect drains the task before tearing down + # the loop or opening a replacement server session. + close_task = self._schedule_socket_close(ws, notify_server=True) + await asyncio.shield(close_task) + + async def _drain_pending_close_tasks(self) -> None: + """Wait for background socket closes owned by the current event loop. + + Shielding keeps cancellation of the caller from cancelling the close + tasks themselves. This matters both before reconnecting, when the old + server session must release its capacity slot, and during explicit + client shutdown. + """ + loop = asyncio.get_running_loop() + tasks = [ + task + for task in tuple(self._pending_close_tasks) + if not task.done() and task.get_loop() is loop + ] + if tasks: + await asyncio.shield(asyncio.gather(*tasks, return_exceptions=True)) async def _ensure_connected(self) -> None: """Ensure WebSocket connection is established on the current loop. @@ -641,9 +686,7 @@ async def _receive(self) -> Dict[str, Any]: # would actually block for up to 10s before its deadline was # honored. Scheduling it lets the exception propagate # immediately while the close still happens in the background. - close_task = asyncio.ensure_future(_best_effort_close(ws)) - self._pending_close_tasks.add(close_task) - close_task.add_done_callback(self._pending_close_tasks.discard) + self._schedule_socket_close(ws) raise return json.loads(raw) @@ -963,16 +1006,16 @@ async def _close_async(self) -> None: self._child_clients.clear() try: - # Wait out any backgrounded closes from a dropped socket (see - # `_receive()` / `_best_effort_close`) so a real close() call still - # sees the handshake through. SyncEnvClient.close() waits for - # `_close_async()` before stopping its loop, so the relevant risk - # is async-context cancellation of close itself — not `_stop_loop()`. - # Keep this gather inside the provider-teardown try/finally so a - # cancelled close cannot skip container/process cleanup. - if self._pending_close_tasks: - await asyncio.gather(*self._pending_close_tasks, return_exceptions=True) - await self._disconnect_async() + try: + # A real close waits out backgrounded closes, but shield them + # from cancellation so their socket handshakes aren't + # abandoned midway. + await self._drain_pending_close_tasks() + finally: + # Run even when pending-close draining is cancelled. A client + # may already have reconnected, and that current socket must + # not remain cached or open during teardown. + await self._disconnect_async() finally: try: if self._provider is not None: diff --git a/src/openenv/discovery/models.py b/src/openenv/discovery/models.py index 39d47d6489..4f25ff85e5 100644 --- a/src/openenv/discovery/models.py +++ b/src/openenv/discovery/models.py @@ -53,7 +53,10 @@ def relative_path(value: str) -> str: return value if ( not value - or "\x00" in value + or any( + ord(character) < 0x20 or 0x7F <= ord(character) <= 0x9F + for character in value + ) or "\\" in value or PurePosixPath(value).is_absolute() or any(part in ("", ".", "..") for part in value.split("/")) @@ -64,8 +67,10 @@ def relative_path(value: str) -> str: # Positive components exclude "." and ".." without lookaround, which some # JSON Schema regex engines do not support. +_CONTROL_CHARACTER_PATTERN = r"[\x00-\x1f\x7f-\x9f]" _PATH_COMPONENT_PATTERN = ( - r"(?:[^./\\\x00]|\.[^./\\\x00]|\.\.[^./\\\x00]|\.\.\.)[^/\\\x00]*" + r"(?:[^./\\\x00-\x1f\x7f-\x9f]|\.[^./\\\x00-\x1f\x7f-\x9f]" + r"|\.\.[^./\\\x00-\x1f\x7f-\x9f]|\.\.\.)[^/\\\x00-\x1f\x7f-\x9f]*" ) RelativePath = Annotated[ NonEmpty, @@ -79,6 +84,7 @@ def relative_path(value: str) -> str: rf"(?:/{_PATH_COMPONENT_PATTERN})*)$" ), }, + {"not": {"pattern": _CONTROL_CHARACTER_PATTERN}}, ] } ), diff --git a/src/openenv/discovery/schemas/0.1-draft/catalog.schema.json b/src/openenv/discovery/schemas/0.1-draft/catalog.schema.json index 3b4491bd06..93114d0d77 100644 --- a/src/openenv/discovery/schemas/0.1-draft/catalog.schema.json +++ b/src/openenv/discovery/schemas/0.1-draft/catalog.schema.json @@ -55,7 +55,12 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, @@ -405,7 +410,12 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, @@ -472,7 +482,12 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, @@ -510,7 +525,12 @@ "items": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, @@ -524,7 +544,12 @@ "root": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, diff --git a/src/openenv/discovery/schemas/0.1-draft/declaration.schema.json b/src/openenv/discovery/schemas/0.1-draft/declaration.schema.json index 91576021ae..5c81448374 100644 --- a/src/openenv/discovery/schemas/0.1-draft/declaration.schema.json +++ b/src/openenv/discovery/schemas/0.1-draft/declaration.schema.json @@ -25,7 +25,12 @@ "source": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, @@ -104,7 +109,12 @@ { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, diff --git a/src/openenv/discovery/schemas/0.1-draft/environment-card.schema.json b/src/openenv/discovery/schemas/0.1-draft/environment-card.schema.json index 0f49575919..b0102491c4 100644 --- a/src/openenv/discovery/schemas/0.1-draft/environment-card.schema.json +++ b/src/openenv/discovery/schemas/0.1-draft/environment-card.schema.json @@ -49,7 +49,12 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, @@ -97,7 +102,12 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*(?:/(?:[^./\\\\\\x00]|\\.[^./\\\\\\x00]|\\.\\.[^./\\\\\\x00]|\\.\\.\\.)[^/\\\\\\x00]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + }, + { + "not": { + "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + } } ], "maxLength": 8192, diff --git a/tests/discovery/test_catalog_contract.py b/tests/discovery/test_catalog_contract.py index f975775264..d5ff0289c1 100644 --- a/tests/discovery/test_catalog_contract.py +++ b/tests/discovery/test_catalog_contract.py @@ -15,6 +15,9 @@ REVISION = "a" * 40 +ASCII_CONTROL_PATHS = [ + f"envs/control-{chr(codepoint)}" for codepoint in [*range(0x20), *range(0x7F, 0xA0)] +] @pytest.fixture @@ -92,7 +95,13 @@ def test_tool_declaration_cannot_borrow_another_revision(card): "envs//echo", "envs/echo/", "envs/./echo", - "envs/\x00echo", + *ASCII_CONTROL_PATHS, + "envs/trailing\n", + ".\n", + "..\n", + "envs/\n", + "envs/.\n", + "envs/..\n", ], ) def test_environment_locator_is_a_safe_repository_relative_path( @@ -115,12 +124,6 @@ def test_environment_locator_is_a_safe_repository_relative_path( "envs/...", "envs/a..b", "envs/with spaces", - "envs/trailing\n", - ".\n", - "..\n", - "envs/\n", - "envs/.\n", - "envs/..\n", ], ) def test_schema_and_model_preserve_valid_relative_locators(card, path, card_schema): diff --git a/tests/envs/test_discovery.py b/tests/envs/test_discovery.py index ff1d3890bd..08552e5b0f 100644 --- a/tests/envs/test_discovery.py +++ b/tests/envs/test_discovery.py @@ -16,6 +16,7 @@ import os import stat import tempfile +from pathlib import Path from unittest.mock import Mock, patch import openenv.auto._discovery as _discovery_module @@ -344,6 +345,24 @@ def test_cache_file_is_per_user_not_shared_tmp(self): assert tempfile.gettempdir() not in str(path) assert path.parent.name == "openenv" + def test_relative_xdg_cache_home_cannot_redirect_into_working_tree( + self, tmp_path, monkeypatch + ): + """A relative XDG path must not trust a cache planted in the checkout.""" + checkout = tmp_path / "untrusted-checkout" + planted = checkout / "cache" / "openenv" / "discovery_cache.json" + planted.parent.mkdir(parents=True) + planted.write_text("{}") + + monkeypatch.chdir(checkout) + monkeypatch.setenv("XDG_CACHE_HOME", "cache") + + path = _default_cache_file() + + assert path == Path.home() / ".cache" / "openenv" / "discovery_cache.json" + assert path.is_absolute() + assert path.resolve() != planted.resolve() + def test_world_writable_cache_is_not_trusted(self, tmp_path): f = tmp_path / "cache.json" f.write_text("{}") diff --git a/tests/test_core/test_generic_client.py b/tests/test_core/test_generic_client.py index 520f0b98e0..415047dfe5 100644 --- a/tests/test_core/test_generic_client.py +++ b/tests/test_core/test_generic_client.py @@ -17,7 +17,6 @@ import asyncio import os -from contextlib import suppress from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest @@ -1663,17 +1662,125 @@ async def close(self): ) @pytest.mark.asyncio - async def test_close_async_cancelled_during_pending_gather_still_stops_provider( + async def test_reconnect_waits_for_dropped_socket_to_release_capacity(self): + """A replacement connection must wait for the old session to close. + + A timed-out socket is detached immediately and closed in the background. + Reconnecting before that handshake finishes races the server's session + accounting; with the default capacity of one, the retry is rejected. + """ + close_started = asyncio.Event() + release_close = asyncio.Event() + + class SlowClose: + state = State.OPEN + + async def send(self, _message): + pass + + async def recv(self): + await asyncio.sleep(10) + + async def close(self): + close_started.set() + await release_close.wait() + self.state = State.CLOSED + + client = GenericEnvClient( + base_url="http://localhost:8000", message_timeout_s=0.01 + ) + dropped_ws = SlowClose() + client._ws = dropped_ws + client._ws_loop = asyncio.get_running_loop() + + with pytest.raises(asyncio.TimeoutError): + await client._send_and_receive({"type": "state"}) + await close_started.wait() + + replacement_ws = AsyncMock() + replacement_ws.state = State.OPEN + replacement_ws.recv.return_value = '{"type": "state", "data": {}}' + + async def fake_ws_connect(*args, **kwargs): + return replacement_ws + + with patch( + "openenv.core.env_client.ws_connect", side_effect=fake_ws_connect + ) as mock_connect: + reconnect = asyncio.create_task(client._connect_async()) + await asyncio.sleep(0) + assert not reconnect.done() + mock_connect.assert_not_called() + + release_close.set() + await asyncio.wait_for(reconnect, timeout=1) + + mock_connect.assert_called_once() + assert dropped_ws.state == State.CLOSED + assert client._ws is replacement_ws + await client._close_async() + + @pytest.mark.asyncio + async def test_cancelled_close_tracks_current_socket_until_handshake_finishes( self, ): - """Cancelling `_close_async` while draining pending closes must still - tear down the provider and clear provider-owned URLs. + """Cancelling the current close must not abandon its server session.""" + close_started = asyncio.Event() + release_close = asyncio.Event() - Regression: the pending-close `gather` used to run before the - try/finally that stops the provider. Cancellation of `_close_async` - itself propagates from `gather` even with `return_exceptions=True`, - which skipped provider teardown and leaked the container/process. - """ + class SlowCurrentSocket: + state = State.OPEN + + async def send(self, _message): + pass + + async def close(self): + close_started.set() + await release_close.wait() + self.state = State.CLOSED + + client = GenericEnvClient(base_url="http://localhost:8000") + current_ws = SlowCurrentSocket() + client._ws = current_ws + client._ws_loop = asyncio.get_running_loop() + + close_call = asyncio.create_task(client._close_async()) + await close_started.wait() + assert client._pending_close_tasks + + close_call.cancel() + with pytest.raises(asyncio.CancelledError): + await close_call + + assert client._ws is None + assert current_ws.state == State.OPEN + assert any(not task.cancelled() for task in client._pending_close_tasks) + + replacement_ws = AsyncMock() + replacement_ws.state = State.OPEN + + async def fake_ws_connect(*args, **kwargs): + return replacement_ws + + with patch( + "openenv.core.env_client.ws_connect", side_effect=fake_ws_connect + ) as mock_connect: + reconnect = asyncio.create_task(client._connect_async()) + await asyncio.sleep(0) + assert not reconnect.done() + mock_connect.assert_not_called() + + release_close.set() + await asyncio.wait_for(reconnect, timeout=1) + + assert current_ws.state == State.CLOSED + mock_connect.assert_called_once() + assert client._ws is replacement_ws + await client._close_async() + + @pytest.mark.asyncio + async def test_cancelled_close_still_closes_current_and_pending_sockets(self): + """Cancellation while draining an old socket must not leak either one.""" class FakeRuntimeProvider: def __init__(self): @@ -1682,37 +1789,59 @@ def __init__(self): def stop(self): self.stopped = True + close_started = asyncio.Event() + release_close = asyncio.Event() + + class PendingSocket: + state = State.OPEN + + async def close(self): + close_started.set() + await release_close.wait() + self.state = State.CLOSED + + class CurrentSocket: + state = State.OPEN + + async def send(self, _message): + pass + + async def close(self): + self.state = State.CLOSED + provider = FakeRuntimeProvider() client = GenericEnvClient(provider=provider) client._base_url = "http://localhost:8000" client._ws_url = "ws://localhost:8000/ws" - hang_gate = asyncio.Event() - - async def hang_forever(): - await hang_gate.wait() - - pending = asyncio.create_task(hang_forever()) + dropped_ws = PendingSocket() + pending = asyncio.create_task(dropped_ws.close()) client._pending_close_tasks.add(pending) + pending.add_done_callback(client._pending_close_tasks.discard) + await close_started.wait() + + current_ws = CurrentSocket() + client._ws = current_ws + client._ws_loop = asyncio.get_running_loop() close_task = asyncio.create_task(client._close_async()) - await asyncio.sleep(0) # let close enter the pending-close gather + await asyncio.sleep(0) # let close enter the shielded pending-close drain assert not close_task.done() close_task.cancel() with pytest.raises(asyncio.CancelledError): await close_task - hang_gate.set() - with suppress(asyncio.CancelledError): - await pending - - assert provider.stopped, ( - "provider.stop() must run even when _close_async is cancelled " - "during the pending-close gather" - ) + assert provider.stopped assert client._base_url is None assert client._ws_url is None + assert client._ws is None + assert current_ws.state == State.CLOSED + assert not pending.cancelled() + + release_close.set() + await pending + assert dropped_ws.state == State.CLOSED # ============================================================================ From 63f43558d63c7051d699b01bd1768e20fa7c103a Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 16 Sep 2026 06:16:24 +0000 Subject: [PATCH 02/11] style: format socket close scheduler Co-authored-by: benjamin.burtenshaw --- src/openenv/core/env_client.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/src/openenv/core/env_client.py b/src/openenv/core/env_client.py index 744f160b0c..f82b1ba0b1 100644 --- a/src/openenv/core/env_client.py +++ b/src/openenv/core/env_client.py @@ -594,9 +594,7 @@ def _schedule_socket_close( ) -> asyncio.Task[None]: """Schedule and track a socket close until its handshake finishes.""" close_coro = ( - _best_effort_graceful_close(ws) - if notify_server - else _best_effort_close(ws) + _best_effort_graceful_close(ws) if notify_server else _best_effort_close(ws) ) close_task = asyncio.create_task(close_coro) self._pending_close_tasks.add(close_task) From 5d177e3bf00e54bae80f39b826deb7d24cfd588a Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 16 Sep 2026 06:19:47 +0000 Subject: [PATCH 03/11] fix(discovery): fail closed on relative home cache Co-authored-by: benjamin.burtenshaw --- src/openenv/auto/_discovery.py | 12 +++++++++++- tests/envs/test_discovery.py | 14 ++++++++++++++ 2 files changed, 25 insertions(+), 1 deletion(-) diff --git a/src/openenv/auto/_discovery.py b/src/openenv/auto/_discovery.py index 2f7a57dbdb..12c7c19f57 100644 --- a/src/openenv/auto/_discovery.py +++ b/src/openenv/auto/_discovery.py @@ -347,9 +347,19 @@ def _default_cache_file() -> Path: attacker-controlled import paths (`import_module` on a cached `client_module_path`). Per the XDG Base Directory specification, relative `XDG_CACHE_HOME` values are ignored so an untrusted working tree cannot supply a victim-owned cache file. + The fallback home must itself be absolute; otherwise discovery fails closed. """ base = os.environ.get("XDG_CACHE_HOME") - root = Path(base) if base and Path(base).is_absolute() else Path.home() / ".cache" + if base and Path(base).is_absolute(): + root = Path(base) + else: + home = Path.home() + if not home.is_absolute(): + raise RuntimeError( + "discovery cache requires an absolute home directory when " + "XDG_CACHE_HOME is unset or relative" + ) + root = home / ".cache" return root / "openenv" / "discovery_cache.json" diff --git a/tests/envs/test_discovery.py b/tests/envs/test_discovery.py index 08552e5b0f..8285f2b4ea 100644 --- a/tests/envs/test_discovery.py +++ b/tests/envs/test_discovery.py @@ -363,6 +363,20 @@ def test_relative_xdg_cache_home_cannot_redirect_into_working_tree( assert path.is_absolute() assert path.resolve() != planted.resolve() + @pytest.mark.parametrize("xdg_cache_home", ["cache", ""]) + def test_relative_home_cannot_restore_a_relative_cache_path( + self, tmp_path, monkeypatch, xdg_cache_home + ): + """Invalid XDG and home paths must fail closed, not trust the checkout.""" + checkout = tmp_path / "untrusted-checkout" + checkout.mkdir() + monkeypatch.chdir(checkout) + monkeypatch.setenv("XDG_CACHE_HOME", xdg_cache_home) + monkeypatch.setenv("HOME", "relative-home") + + with pytest.raises(RuntimeError, match="absolute home directory"): + _default_cache_file() + def test_world_writable_cache_is_not_trusted(self, tmp_path): f = tmp_path / "cache.json" f.write_text("{}") From 366941de505991507c603f27f48acfcbbecd9865 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 16 Sep 2026 06:19:52 +0000 Subject: [PATCH 04/11] fix(client): preserve parent teardown on child cancellation Co-authored-by: benjamin.burtenshaw --- src/openenv/core/env_client.py | 23 ++++++------ tests/test_core/test_generic_client.py | 49 ++++++++++++++++++++++++++ 2 files changed, 61 insertions(+), 11 deletions(-) diff --git a/src/openenv/core/env_client.py b/src/openenv/core/env_client.py index f82b1ba0b1..36a3380fe2 100644 --- a/src/openenv/core/env_client.py +++ b/src/openenv/core/env_client.py @@ -998,21 +998,22 @@ async def _close_async(self) -> None: If this client was created via from_docker_image() or from_env(), this will also stop and remove the associated container/process. """ - for child in list(self._child_clients): - with suppress(Exception): - await child.close() - self._child_clients.clear() - try: try: - # A real close waits out backgrounded closes, but shield them - # from cancellation so their socket handshakes aren't - # abandoned midway. + try: + for child in list(self._child_clients): + with suppress(Exception): + await child.close() + finally: + self._child_clients.clear() + + # A real close waits out backgrounded closes, while shielding + # their socket handshakes from cancellation. await self._drain_pending_close_tasks() finally: - # Run even when pending-close draining is cancelled. A client - # may already have reconnected, and that current socket must - # not remain cached or open during teardown. + # Run even when child or pending-close cleanup is cancelled. A + # client may already have reconnected, and that current socket + # must not remain cached or open during teardown. await self._disconnect_async() finally: try: diff --git a/tests/test_core/test_generic_client.py b/tests/test_core/test_generic_client.py index 415047dfe5..a2ee2d3c31 100644 --- a/tests/test_core/test_generic_client.py +++ b/tests/test_core/test_generic_client.py @@ -174,6 +174,55 @@ async def test_close_stops_provider_when_child_close_raises(self, mock_provider) assert client._child_clients == [] mock_provider.stop_container.assert_called_once_with() + @pytest.mark.asyncio + async def test_cancelled_child_close_still_tears_down_parent(self): + """Child-close cancellation must not bypass parent socket/provider cleanup.""" + + class FakeRuntimeProvider: + def __init__(self): + self.stopped = False + + def stop(self): + self.stopped = True + + child_close_started = asyncio.Event() + + class SlowChild: + async def close(self): + child_close_started.set() + await asyncio.sleep(10) + + class ParentSocket: + state = State.OPEN + + async def send(self, _message): + pass + + async def close(self): + self.state = State.CLOSED + + provider = FakeRuntimeProvider() + client = GenericEnvClient(provider=provider) + client._base_url = "http://localhost:8000" + client._ws_url = "ws://localhost:8000/ws" + parent_ws = ParentSocket() + client._ws = parent_ws + client._ws_loop = asyncio.get_running_loop() + client._child_clients.append(SlowChild()) + + close_call = asyncio.create_task(client._close_async()) + await child_close_started.wait() + close_call.cancel() + with pytest.raises(asyncio.CancelledError): + await close_call + + assert client._child_clients == [] + assert client._ws is None + assert parent_ws.state == State.CLOSED + assert provider.stopped + assert client._base_url is None + assert client._ws_url is None + def test_session_client_filters_constructor_kwargs(self): """Child creation respects subclasses with narrower __init__ signatures.""" From 3b06765e3295051cc3f19bbcb023ef7e460b63c3 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 16 Sep 2026 06:23:13 +0000 Subject: [PATCH 05/11] fix(discovery): reject Unicode line separators Co-authored-by: benjamin.burtenshaw --- docs/source/guides/catalog-discovery.md | 6 +++--- src/openenv/discovery/models.py | 12 +++++++---- .../schemas/0.1-draft/catalog.schema.json | 20 +++++++++---------- .../schemas/0.1-draft/declaration.schema.json | 8 ++++---- .../0.1-draft/environment-card.schema.json | 8 ++++---- tests/discovery/test_catalog_contract.py | 8 +++++--- 6 files changed, 34 insertions(+), 28 deletions(-) diff --git a/docs/source/guides/catalog-discovery.md b/docs/source/guides/catalog-discovery.md index d5a8862231..d640bbdc44 100644 --- a/docs/source/guides/catalog-discovery.md +++ b/docs/source/guides/catalog-discovery.md @@ -164,10 +164,10 @@ does not require fetching candidate resource URLs. JSON Schema validation is necessary but is not the whole profile contract. The packaged schemas enforce object shape, required fields, relative-path safety, supported literals, conditional artifact/license-evidence presence, and -exactly one orchestration interface. Other rules require semantic validation: +exactly one orchestration interface. The relative-path profile permits printable UTF-8 (including spaces), but -rejects C0, DEL, and C1 control characters so untrusted metadata cannot forge -CLI or log lines. +rejects C0, DEL, C1, and Unicode line/paragraph separators so untrusted metadata +cannot forge CLI or log lines. Other rules require semantic validation: | Subject | Additional rule | |---------|-----------------| diff --git a/src/openenv/discovery/models.py b/src/openenv/discovery/models.py index 4f25ff85e5..4ddacb77c2 100644 --- a/src/openenv/discovery/models.py +++ b/src/openenv/discovery/models.py @@ -54,7 +54,9 @@ def relative_path(value: str) -> str: if ( not value or any( - ord(character) < 0x20 or 0x7F <= ord(character) <= 0x9F + ord(character) < 0x20 + or 0x7F <= ord(character) <= 0x9F + or ord(character) in (0x2028, 0x2029) for character in value ) or "\\" in value @@ -67,10 +69,12 @@ def relative_path(value: str) -> str: # Positive components exclude "." and ".." without lookaround, which some # JSON Schema regex engines do not support. -_CONTROL_CHARACTER_PATTERN = r"[\x00-\x1f\x7f-\x9f]" +_CONTROL_CHARACTER_PATTERN = r"[\x00-\x1f\x7f-\x9f\u2028\u2029]" _PATH_COMPONENT_PATTERN = ( - r"(?:[^./\\\x00-\x1f\x7f-\x9f]|\.[^./\\\x00-\x1f\x7f-\x9f]" - r"|\.\.[^./\\\x00-\x1f\x7f-\x9f]|\.\.\.)[^/\\\x00-\x1f\x7f-\x9f]*" + r"(?:[^./\\\x00-\x1f\x7f-\x9f\u2028\u2029]" + r"|\.[^./\\\x00-\x1f\x7f-\x9f\u2028\u2029]" + r"|\.\.[^./\\\x00-\x1f\x7f-\x9f\u2028\u2029]" + r"|\.\.\.)[^/\\\x00-\x1f\x7f-\x9f\u2028\u2029]*" ) RelativePath = Annotated[ NonEmpty, diff --git a/src/openenv/discovery/schemas/0.1-draft/catalog.schema.json b/src/openenv/discovery/schemas/0.1-draft/catalog.schema.json index 93114d0d77..40d4b35686 100644 --- a/src/openenv/discovery/schemas/0.1-draft/catalog.schema.json +++ b/src/openenv/discovery/schemas/0.1-draft/catalog.schema.json @@ -55,11 +55,11 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*)*)$" }, { "not": { - "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + "pattern": "[\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]" } } ], @@ -410,11 +410,11 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*)*)$" }, { "not": { - "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + "pattern": "[\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]" } } ], @@ -482,11 +482,11 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*)*)$" }, { "not": { - "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + "pattern": "[\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]" } } ], @@ -525,11 +525,11 @@ "items": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*)*)$" }, { "not": { - "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + "pattern": "[\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]" } } ], @@ -544,11 +544,11 @@ "root": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*)*)$" }, { "not": { - "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + "pattern": "[\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]" } } ], diff --git a/src/openenv/discovery/schemas/0.1-draft/declaration.schema.json b/src/openenv/discovery/schemas/0.1-draft/declaration.schema.json index 5c81448374..fb7422a1c8 100644 --- a/src/openenv/discovery/schemas/0.1-draft/declaration.schema.json +++ b/src/openenv/discovery/schemas/0.1-draft/declaration.schema.json @@ -25,11 +25,11 @@ "source": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*)*)$" }, { "not": { - "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + "pattern": "[\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]" } } ], @@ -109,11 +109,11 @@ { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*)*)$" }, { "not": { - "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + "pattern": "[\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]" } } ], diff --git a/src/openenv/discovery/schemas/0.1-draft/environment-card.schema.json b/src/openenv/discovery/schemas/0.1-draft/environment-card.schema.json index b0102491c4..f8980e004e 100644 --- a/src/openenv/discovery/schemas/0.1-draft/environment-card.schema.json +++ b/src/openenv/discovery/schemas/0.1-draft/environment-card.schema.json @@ -49,11 +49,11 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*)*)$" }, { "not": { - "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + "pattern": "[\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]" } } ], @@ -102,11 +102,11 @@ "path": { "allOf": [ { - "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f]*)*)$" + "pattern": "^(?:\\.|(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*(?:/(?:[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.[^./\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]|\\.\\.\\.)[^/\\\\\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]*)*)$" }, { "not": { - "pattern": "[\\x00-\\x1f\\x7f-\\x9f]" + "pattern": "[\\x00-\\x1f\\x7f-\\x9f\\u2028\\u2029]" } } ], diff --git a/tests/discovery/test_catalog_contract.py b/tests/discovery/test_catalog_contract.py index d5ff0289c1..d484a96dab 100644 --- a/tests/discovery/test_catalog_contract.py +++ b/tests/discovery/test_catalog_contract.py @@ -15,9 +15,9 @@ REVISION = "a" * 40 -ASCII_CONTROL_PATHS = [ +UNSAFE_LINE_PATHS = [ f"envs/control-{chr(codepoint)}" for codepoint in [*range(0x20), *range(0x7F, 0xA0)] -] +] + ["envs/line-\u2028separator", "envs/paragraph-\u2029separator"] @pytest.fixture @@ -95,7 +95,7 @@ def test_tool_declaration_cannot_borrow_another_revision(card): "envs//echo", "envs/echo/", "envs/./echo", - *ASCII_CONTROL_PATHS, + *UNSAFE_LINE_PATHS, "envs/trailing\n", ".\n", "..\n", @@ -124,6 +124,8 @@ def test_environment_locator_is_a_safe_repository_relative_path( "envs/...", "envs/a..b", "envs/with spaces", + "envs/café", + "envs/東京", ], ) def test_schema_and_model_preserve_valid_relative_locators(card, path, card_schema): From f1069e8f2d080e3aeeb6f7926c75af171bf13c3c Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 16 Sep 2026 06:36:21 +0000 Subject: [PATCH 06/11] fix(client): close every child despite cancellation Co-authored-by: benjamin.burtenshaw --- src/openenv/core/env_client.py | 60 +++++++++++++++++++++++--- tests/test_core/test_generic_client.py | 48 ++++++++++++++++----- 2 files changed, 91 insertions(+), 17 deletions(-) diff --git a/src/openenv/core/env_client.py b/src/openenv/core/env_client.py index 36a3380fe2..8e7f58ddc6 100644 --- a/src/openenv/core/env_client.py +++ b/src/openenv/core/env_client.py @@ -991,6 +991,40 @@ async def _state_async(self) -> StateT: def close(self) -> Any: return self._dispatch(self._close_async) + async def _close_child_clients(self) -> asyncio.CancelledError | None: + """Close every captured child while deferring caller cancellation.""" + children = list(self._child_clients) + self._child_clients.clear() + close_tasks = [] + deferred_cancellation: asyncio.CancelledError | None = None + + for child in children: + try: + close_tasks.append(asyncio.ensure_future(child.close())) + except asyncio.CancelledError as exc: + if deferred_cancellation is None: + deferred_cancellation = exc + except Exception: + pass # Best effort + + if close_tasks: + close_group = asyncio.gather(*close_tasks, return_exceptions=True) + while not close_group.done(): + try: + await asyncio.shield(close_group) + except asyncio.CancelledError as exc: + if deferred_cancellation is None: + deferred_cancellation = exc + + for result in close_group.result(): + if ( + isinstance(result, asyncio.CancelledError) + and deferred_cancellation is None + ): + deferred_cancellation = result + + return deferred_cancellation + async def _close_async(self) -> None: """ Close the WebSocket connection and clean up resources. @@ -998,23 +1032,32 @@ async def _close_async(self) -> None: If this client was created via from_docker_image() or from_env(), this will also stop and remove the associated container/process. """ + deferred_cancellation: asyncio.CancelledError | None = None try: try: try: - for child in list(self._child_clients): - with suppress(Exception): - await child.close() - finally: - self._child_clients.clear() + child_cancellation = await self._close_child_clients() + except asyncio.CancelledError as exc: + child_cancellation = exc + if child_cancellation is not None: + deferred_cancellation = child_cancellation # A real close waits out backgrounded closes, while shielding # their socket handshakes from cancellation. - await self._drain_pending_close_tasks() + try: + await self._drain_pending_close_tasks() + except asyncio.CancelledError as exc: + if deferred_cancellation is None: + deferred_cancellation = exc finally: # Run even when child or pending-close cleanup is cancelled. A # client may already have reconnected, and that current socket # must not remain cached or open during teardown. - await self._disconnect_async() + try: + await self._disconnect_async() + except asyncio.CancelledError as exc: + if deferred_cancellation is None: + deferred_cancellation = exc finally: try: if self._provider is not None: @@ -1028,6 +1071,9 @@ async def _close_async(self) -> None: self._base_url = None self._ws_url = None + if deferred_cancellation is not None: + raise deferred_cancellation + def _stop_provider_best_effort(self) -> None: """Stop the underlying provider directly, ignoring any errors. diff --git a/tests/test_core/test_generic_client.py b/tests/test_core/test_generic_client.py index a2ee2d3c31..7cec264a70 100644 --- a/tests/test_core/test_generic_client.py +++ b/tests/test_core/test_generic_client.py @@ -175,8 +175,8 @@ async def test_close_stops_provider_when_child_close_raises(self, mock_provider) mock_provider.stop_container.assert_called_once_with() @pytest.mark.asyncio - async def test_cancelled_child_close_still_tears_down_parent(self): - """Child-close cancellation must not bypass parent socket/provider cleanup.""" + async def test_cancelled_child_close_still_closes_all_children_and_parent(self): + """Cancellation waits for every child before parent teardown.""" class FakeRuntimeProvider: def __init__(self): @@ -185,12 +185,26 @@ def __init__(self): def stop(self): self.stopped = True - child_close_started = asyncio.Event() + first_close_started = asyncio.Event() + release_first_close = asyncio.Event() + second_close_completed = asyncio.Event() + + class FirstChild: + def __init__(self): + self.closed = False - class SlowChild: async def close(self): - child_close_started.set() - await asyncio.sleep(10) + first_close_started.set() + await release_first_close.wait() + self.closed = True + + class SecondChild: + def __init__(self): + self.closed = False + + async def close(self): + self.closed = True + second_close_completed.set() class ParentSocket: state = State.OPEN @@ -208,14 +222,28 @@ async def close(self): parent_ws = ParentSocket() client._ws = parent_ws client._ws_loop = asyncio.get_running_loop() - client._child_clients.append(SlowChild()) + first_child = FirstChild() + second_child = SecondChild() + client._child_clients.extend([first_child, second_child]) close_call = asyncio.create_task(client._close_async()) - await child_close_started.wait() + await first_close_started.wait() close_call.cancel() - with pytest.raises(asyncio.CancelledError): - await close_call + await asyncio.wait_for(second_close_completed.wait(), timeout=1) + assert not close_call.done() + assert parent_ws.state != State.CLOSED + assert not provider.stopped + + release_first_close.set() + close_results = await asyncio.wait_for( + asyncio.gather(close_call, return_exceptions=True), timeout=1 + ) + + assert len(close_results) == 1 + assert isinstance(close_results[0], asyncio.CancelledError) + assert first_child.closed + assert second_child.closed assert client._child_clients == [] assert client._ws is None assert parent_ws.state == State.CLOSED From 2575ab4c8aeb7891d8f05475b3de08ac268a2783 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 16 Sep 2026 09:35:56 +0000 Subject: [PATCH 07/11] Revert "fix(mcp): production JSON-RPC routing with sync-safe teardown (from #1169) (#1175)" This reverts commit e3eb3fa5bc11ff0a019c30c7a7b1b7039eb8b615. --- src/openenv/core/mcp_client.py | 31 +---- tests/core/test_mode_selection.py | 186 +++++------------------------- 2 files changed, 29 insertions(+), 188 deletions(-) diff --git a/src/openenv/core/mcp_client.py b/src/openenv/core/mcp_client.py index 1afc6254b9..7634bb0be9 100644 --- a/src/openenv/core/mcp_client.py +++ b/src/openenv/core/mcp_client.py @@ -154,7 +154,7 @@ def __init__( mode=mode, ) self._tools_cache: Optional[List[Tool]] = None - self.use_production_mode = self._mode == "production" + self.use_production_mode = False self._production_session_id: Optional[str] = None self._production_session_lock = asyncio.Lock() self._jsonrpc_request_id = 0 @@ -198,27 +198,6 @@ async def _production_mcp_request( response.raise_for_status() return response.json() - async def _connect_async(self) -> EnvClient: - """ - Establish connection to the server. - - In production mode (`use_production_mode=True`), open the WebSocket used - by `reset` / `step` / `state` and create a persistent HTTP MCP session - for `list_tools` / `call_tool`. Tool calls bypass `step()` over `/mcp`, - but the Gym lifecycle still requires `/ws` until production routing - covers those methods end-to-end. - """ - if getattr(self, "use_production_mode", False): - try: - await super()._connect_async() - await self._ensure_production_session() - except Exception: - await self.close() - raise - return self - - return await super()._connect_async() - async def _ensure_production_session(self) -> str: """Create and cache a persistent HTTP MCP session id if needed.""" async with self._production_session_lock: @@ -360,16 +339,12 @@ def _parse_state(self, payload: Dict[str, Any]) -> State: step_count=payload.get("step_count", 0), ) - async def _close_async(self) -> None: + async def close(self) -> None: """ Close client resources. In production MCP mode, this also closes the server-side persistent MCP session (best effort) before closing websocket/provider resources. - - Override `_close_async` rather than `close` so sync teardown - (`SyncEnvClient.close`, sync `__exit__`, and `_dispatch`) still cleans - up the HTTP MCP session. """ if self._production_session_id is not None: try: @@ -391,7 +366,7 @@ async def _close_async(self) -> None: finally: self._http_client = None - await super()._close_async() + await super().close() class MCPToolClient(MCPClientBase): diff --git a/tests/core/test_mode_selection.py b/tests/core/test_mode_selection.py index 04aa7ef490..cbbf543d48 100644 --- a/tests/core/test_mode_selection.py +++ b/tests/core/test_mode_selection.py @@ -22,7 +22,7 @@ """ import os -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import MagicMock, patch import pytest from fastmcp import FastMCP @@ -193,172 +193,39 @@ async def test_simulation_mode_uses_gym_protocol(self, clean_env, mock_websocket ) @pytest.mark.asyncio - async def test_production_mode_uses_jsonrpc_protocol(self, clean_env): - """Test that production mode uses HTTP JSON-RPC format for tool listing.""" - client = MCPToolClient(base_url="http://localhost:8000", mode="production") - assert client.use_production_mode is True - - with patch.object( - client, - "_production_mcp_request", - side_effect=[ - {"result": {"session_id": "test-session"}}, - { - "result": { - "tools": [ - { - "name": "echo", - "description": "Echo message", - "inputSchema": {}, - } - ] - } - }, - ], - ) as mock_mcp_request: - with patch.object(client, "step") as mock_step: - tools = await client.list_tools() - - mock_step.assert_not_called() - assert len(tools) == 1 - assert tools[0].name == "echo" - assert mock_mcp_request.call_count == 2 - mock_mcp_request.assert_called_with( - "tools/list", {"session_id": "test-session"} - ) - - @pytest.mark.asyncio - async def test_production_mode_call_tool_uses_jsonrpc_protocol(self, clean_env): - """Test that call_tool in production mode uses HTTP JSON-RPC transport.""" - client = MCPToolClient(base_url="http://localhost:8000", mode="production") - assert client.use_production_mode is True - - with patch.object( - client, - "_production_mcp_request", - side_effect=[ - {"result": {"session_id": "test-session"}}, - {"result": {"data": "hello world"}}, - ], - ) as mock_mcp_request: - with patch.object(client, "step") as mock_step: - result = await client.call_tool("echo", message="hello world") - - mock_step.assert_not_called() - assert result == "hello world" - mock_mcp_request.assert_called_with( - "tools/call", - { - "name": "echo", - "arguments": {"message": "hello world"}, - "session_id": "test-session", - }, - ) - - @pytest.mark.asyncio - async def test_production_mode_connect_opens_websocket_and_http_session( - self, clean_env + async def test_production_mode_uses_jsonrpc_protocol( + self, clean_env, mock_websocket ): - """Production connect must open WebSocket (reset/step/state) and HTTP MCP session.""" + """Test that production mode uses JSON-RPC format for tool calls.""" client = MCPToolClient(base_url="http://localhost:8000", mode="production") - assert client.use_production_mode is True - - mock_ws = MagicMock() - mock_ws.closed = False - - with patch.object( - client, - "_production_mcp_request", - side_effect=[ - {"result": {"session_id": "test-session"}}, - {"result": {"data": "hello world"}}, - ], - ) as mock_mcp_request: - with patch( - "openenv.core.env_client.ws_connect", - new_callable=AsyncMock, - return_value=mock_ws, - ) as mock_ws_connect: - await client.connect() - - mock_ws_connect.assert_called_once() - assert client._ws is mock_ws - assert client._production_session_id == "test-session" - mock_mcp_request.assert_called_once_with("openenv/session/create") - - result = await client.call_tool("echo", message="hello world") - assert result == "hello world" - assert mock_mcp_request.call_count == 2 - mock_mcp_request.assert_called_with( - "tools/call", - { - "name": "echo", - "arguments": {"message": "hello world"}, - "session_id": "test-session", - }, - ) - @pytest.mark.asyncio - async def test_production_mode_connect_failure_cleans_up_resources(self, clean_env): - """Test that failure during production mode connect() triggers client.close() cleanup.""" - client = MCPToolClient(base_url="http://localhost:8000", mode="production") - assert client.use_production_mode is True - - mock_ws = MagicMock() - mock_ws.closed = False - mock_ws.close = AsyncMock() - - with patch( - "openenv.core.env_client.ws_connect", - new_callable=AsyncMock, - return_value=mock_ws, - ): + with patch.object(client, "_send") as mock_send: with patch.object( client, - "_ensure_production_session", - side_effect=RuntimeError("Session creation failed"), - ): - with patch.object(client, "close", wraps=client.close) as mock_close: - with pytest.raises(RuntimeError, match="Session creation failed"): - await client.connect() - - mock_close.assert_called_once() - - @pytest.mark.asyncio - async def test_production_mode_sync_close_closes_mcp_session(self, clean_env): - """Sync close must tear down the HTTP MCP session via `_close_async`.""" - client = MCPToolClient(base_url="http://localhost:8000", mode="production") - assert client.use_production_mode is True - - mock_ws = MagicMock() - mock_ws.closed = False - mock_ws.close = AsyncMock() - - with patch.object( - client, - "_production_mcp_request", - side_effect=[ - {"result": {"session_id": "test-session"}}, - {"result": {}}, - ], - ) as mock_mcp_request: - with patch( - "openenv.core.env_client.ws_connect", - new_callable=AsyncMock, - return_value=mock_ws, + "_receive", + return_value={ + "type": "response", + "data": { + "observation": {"tools": []}, + "reward": None, + "done": False, + }, + }, ): - sync_client = client.sync() - sync_client.connect() - assert client._production_session_id == "test-session" + with patch.object(client, "_ws", mock_websocket): + await client.list_tools() - sync_client.close() + # Should send step message with list_tools action + call_args = mock_send.call_args_list + step_call = [ + call for call in call_args if call[0][0].get("type") == "step" + ] + assert len(step_call) > 0, "Should send message with type='step'" - assert client._production_session_id is None - assert mock_mcp_request.call_count == 2 - mock_mcp_request.assert_any_call( - "openenv/session/close", - {"session_id": "test-session"}, - ) + # Check that the action payload is list_tools + step_message = step_call[0][0][0] + assert "data" in step_message + assert step_message["data"].get("type") == "list_tools" # ============================================================================ @@ -416,7 +283,6 @@ def test_mcp_client_defaults_to_production_mode(self, clean_env): # MCPToolClient should default to production mode assert client._mode == "production" - assert client.use_production_mode is True def test_mcp_client_cannot_use_simulation_mode(self, clean_env): """Test that MCPToolClient raises error if simulation mode is requested.""" From 97b7688fa601afc993abfecc35dd942928987584 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 16 Sep 2026 09:44:53 +0000 Subject: [PATCH 08/11] fix(mcp): preserve sync-safe teardown after revert Co-authored-by: benjamin.burtenshaw --- src/openenv/core/mcp_client.py | 45 ++++++++------- tests/core/test_mode_selection.py | 93 ++++++++++++++++++++++++++++++- 2 files changed, 117 insertions(+), 21 deletions(-) diff --git a/src/openenv/core/mcp_client.py b/src/openenv/core/mcp_client.py index 7634bb0be9..d899bbf393 100644 --- a/src/openenv/core/mcp_client.py +++ b/src/openenv/core/mcp_client.py @@ -339,34 +339,39 @@ def _parse_state(self, payload: Dict[str, Any]) -> State: step_count=payload.get("step_count", 0), ) - async def close(self) -> None: + async def _close_async(self) -> None: """ Close client resources. In production MCP mode, this also closes the server-side persistent MCP session (best effort) before closing websocket/provider resources. - """ - if self._production_session_id is not None: - try: - await self._production_mcp_request( - "openenv/session/close", - {"session_id": self._production_session_id}, - ) - except Exception: - # Best effort cleanup - do not mask normal close behavior - pass - finally: - self._production_session_id = None - if self._http_client is not None: + This overrides the internal coroutine so sync and async dispatch paths + share the same cleanup. + """ + try: + if self._production_session_id is not None: + try: + await self._production_mcp_request( + "openenv/session/close", + {"session_id": self._production_session_id}, + ) + except Exception: + # Best effort cleanup - do not mask normal close behavior + pass + finally: + self._production_session_id = None + finally: try: - await self._http_client.aclose() - except Exception: - pass + if self._http_client is not None: + try: + await self._http_client.aclose() + except Exception: + pass + finally: + self._http_client = None finally: - self._http_client = None - - await super().close() + await super()._close_async() class MCPToolClient(MCPClientBase): diff --git a/tests/core/test_mode_selection.py b/tests/core/test_mode_selection.py index cbbf543d48..bc6567d725 100644 --- a/tests/core/test_mode_selection.py +++ b/tests/core/test_mode_selection.py @@ -21,8 +21,9 @@ - Environment: Code mode with mode-aware tool registration """ +import asyncio import os -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastmcp import FastMCP @@ -228,6 +229,96 @@ async def test_production_mode_uses_jsonrpc_protocol( assert step_message["data"].get("type") == "list_tools" +class TestMCPClientCleanup: + """MCP-specific resources must use the common close dispatch path.""" + + def test_bare_close_preserves_sync_dispatch(self, clean_env): + """Calling close outside an event loop must resolve cleanup synchronously.""" + client = MCPToolClient(base_url="http://localhost:8000") + client._production_session_id = "test-session" + + try: + with patch.object( + client, + "_production_mcp_request", + new=AsyncMock(return_value={"result": {}}), + ) as request: + result = client.close() + + assert result is None + request.assert_awaited_once_with( + "openenv/session/close", {"session_id": "test-session"} + ) + assert client._production_session_id is None + finally: + if client._sync_client is not None: + client._sync_client._stop_loop() + + def test_sync_wrapper_close_releases_http_resources(self, clean_env): + """Sync wrapper close must release the MCP session and HTTP client.""" + client = MCPToolClient(base_url="http://localhost:8000") + client._production_session_id = "test-session" + http_client = AsyncMock() + client._http_client = http_client + + with patch.object( + client, + "_production_mcp_request", + new=AsyncMock(return_value={"result": {}}), + ) as request: + client.sync().close() + + request.assert_awaited_once_with( + "openenv/session/close", {"session_id": "test-session"} + ) + http_client.aclose.assert_awaited_once_with() + assert client._production_session_id is None + assert client._http_client is None + + @pytest.mark.asyncio + async def test_cancelled_session_close_still_releases_http_and_provider( + self, clean_env + ): + """Cancellation must not bypass later HTTP and provider cleanup.""" + + class Provider: + def __init__(self): + self.stopped = False + + def stop(self): + self.stopped = True + + close_started = asyncio.Event() + + async def slow_session_close(*_args, **_kwargs): + close_started.set() + await asyncio.sleep(10) + + provider = Provider() + client = MCPToolClient( + base_url="http://localhost:8000", + provider=provider, + ) + client._production_session_id = "test-session" + http_client = AsyncMock() + client._http_client = http_client + + with patch.object( + client, "_production_mcp_request", side_effect=slow_session_close + ): + close_call = asyncio.create_task(client._close_async()) + await close_started.wait() + close_call.cancel() + close_results = await asyncio.gather(close_call, return_exceptions=True) + + assert len(close_results) == 1 + assert isinstance(close_results[0], asyncio.CancelledError) + http_client.aclose.assert_awaited_once_with() + assert client._production_session_id is None + assert client._http_client is None + assert provider.stopped + + # ============================================================================ # Mode Immutability Tests # ============================================================================ From 157e125f8f3ad95528d6ade8c2170ffd75a2687c Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 16 Sep 2026 09:48:26 +0000 Subject: [PATCH 09/11] style(mcp): document best-effort HTTP cleanup Co-authored-by: benjamin.burtenshaw --- src/openenv/core/mcp_client.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/openenv/core/mcp_client.py b/src/openenv/core/mcp_client.py index d899bbf393..590142056b 100644 --- a/src/openenv/core/mcp_client.py +++ b/src/openenv/core/mcp_client.py @@ -367,7 +367,7 @@ async def _close_async(self) -> None: try: await self._http_client.aclose() except Exception: - pass + pass # Best effort; continue to websocket/provider teardown finally: self._http_client = None finally: From de95e72b7bd90adada24760fb21879f1d8ad0fab Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 16 Sep 2026 11:00:52 +0000 Subject: [PATCH 10/11] fix(mcp): isolate explicit HTTP tools mode Co-authored-by: benjamin.burtenshaw --- src/openenv/core/mcp_client.py | 99 ++++++++++++++++++----- tests/core/test_mode_selection.py | 62 ++++++++++++++ tests/core/test_production_mode_routes.py | 38 ++++++--- 3 files changed, 166 insertions(+), 33 deletions(-) diff --git a/src/openenv/core/mcp_client.py b/src/openenv/core/mcp_client.py index 590142056b..6fa7b56ea3 100644 --- a/src/openenv/core/mcp_client.py +++ b/src/openenv/core/mcp_client.py @@ -20,14 +20,14 @@ │ /mcp → MCP JSON-RPC (tools/list, tools/call) │ │ /reset, /step, /state → HTTP endpoints │ ├─────────────────────────────────────────────────────────┤ - │ Production Mode (use_production_mode=True): │ + │ Explicit direct mode (use_production_mode=True): │ │ /mcp → MCP JSON-RPC (tools/list, tools/call) │ - │ Bypasses step() for direct tool access │ + │ Tools only; Gym reset/step/state are unavailable │ └─────────────────────────────────────────────────────────┘ Client Usage: MCPToolClient (default) → /ws (step-based, with rewards) - MCPToolClient (production) → /mcp (direct tool access, no rewards) + MCPToolClient (direct opt-in) → /mcp (tools only, no rewards) Examples: @@ -154,12 +154,49 @@ def __init__( mode=mode, ) self._tools_cache: Optional[List[Tool]] = None - self.use_production_mode = False + self._use_production_mode = False self._production_session_id: Optional[str] = None self._production_session_lock = asyncio.Lock() self._jsonrpc_request_id = 0 self._http_client: Optional[Any] = None # lazily-created httpx.AsyncClient + @property + def use_production_mode(self) -> bool: + """Whether explicit tools-only HTTP MCP routing is enabled.""" + return self._use_production_mode + + @use_production_mode.setter + def use_production_mode(self, value: bool) -> None: + """Enable or disable tools-only HTTP MCP routing before connecting.""" + if not isinstance(value, bool): + raise TypeError("use_production_mode must be a bool") + current = getattr(self, "_use_production_mode", False) + has_live_transport = ( + getattr(self, "_ws", None) is not None + or getattr(self, "_production_session_id", None) is not None + or getattr(self, "_http_client", None) is not None + ) + if value != current and has_live_transport: + raise RuntimeError( + "use_production_mode cannot change while a client transport is active" + ) + self._use_production_mode = value + + async def _connect_async(self) -> EnvClient: + """Connect the Gym WebSocket, or prepare explicit tools-only mode.""" + if not self.use_production_mode: + return await super()._connect_async() + + try: + self._start_provider_if_needed() + except BaseException: + try: + await asyncio.shield(self._close_async()) + except BaseException: + pass # Preserve the original startup failure + raise + return self + def _next_request_id(self) -> int: """Generate a monotonically increasing JSON-RPC request id.""" self._jsonrpc_request_id += 1 @@ -216,6 +253,27 @@ async def _ensure_production_session(self) -> str: self._production_session_id = session_id return session_id + def _tools_only_error(self, operation: str) -> RuntimeError: + return RuntimeError( + f"{operation} is unavailable while use_production_mode=True; " + "direct MCP mode supports only list_tools() and call_tool()" + ) + + async def _reset_async(self, **kwargs: Any) -> StepResult[Observation]: + if self.use_production_mode: + raise self._tools_only_error("reset()") + return await super()._reset_async(**kwargs) + + async def _step_async(self, action: Any, **kwargs: Any) -> StepResult[Observation]: + if self.use_production_mode: + raise self._tools_only_error("step()") + return await super()._step_async(action, **kwargs) + + async def _state_async(self) -> State: + if self.use_production_mode: + raise self._tools_only_error("state()") + return await super()._state_async() + async def list_tools(self, use_cache: bool = True) -> List[Tool]: """ Discover available tools from the environment. @@ -241,23 +299,22 @@ async def list_tools(self, use_cache: bool = True) -> List[Tool]: # Use production mode HTTP endpoint if enabled. # Some tests instantiate with __new__ and skip __init__, so default missing flag to False. if getattr(self, "use_production_mode", False): - try: - session_id = await self._ensure_production_session() - data = await self._production_mcp_request( - "tools/list", - {"session_id": session_id}, - ) - if "error" in data: - message = data.get("error", {}).get("message", "unknown error") - raise RuntimeError(f"list_tools failed: {message}") - if "result" in data and "tools" in data["result"]: - tools = [_tool_from_payload(t) for t in data["result"]["tools"]] - self._tools_cache = tools - return tools - except Exception: - # If HTTP request fails, return empty list - pass - return [] + session_id = await self._ensure_production_session() + data = await self._production_mcp_request( + "tools/list", + {"session_id": session_id}, + ) + if "error" in data: + message = data.get("error", {}).get("message", "unknown error") + raise RuntimeError(f"list_tools failed: {message}") + result = data.get("result") + if not isinstance(result, dict) or not isinstance( + result.get("tools"), list + ): + raise RuntimeError("list_tools failed: malformed JSON-RPC result") + tools = [_tool_from_payload(t) for t in result["tools"]] + self._tools_cache = tools + return tools result = await self.step(ListToolsAction()) if isinstance(result.observation, ListToolsObservation): diff --git a/tests/core/test_mode_selection.py b/tests/core/test_mode_selection.py index bc6567d725..a4aab83b0d 100644 --- a/tests/core/test_mode_selection.py +++ b/tests/core/test_mode_selection.py @@ -229,6 +229,68 @@ async def test_production_mode_uses_jsonrpc_protocol( assert step_message["data"].get("type") == "list_tools" +class TestDirectMCPMode: + """Explicit HTTP MCP mode must remain tools-only and single-transport.""" + + @pytest.mark.asyncio + async def test_direct_mode_connect_does_not_open_websocket(self, clean_env): + client = MCPToolClient(base_url="http://localhost:8000") + client.use_production_mode = True + + with patch( + "openenv.core.env_client.ws_connect", new_callable=AsyncMock + ) as ws_connect: + await client.connect() + + ws_connect.assert_not_awaited() + assert client._ws is None + assert client._production_session_id is None + await client.close() + + @pytest.mark.asyncio + async def test_direct_mode_rejects_gym_lifecycle_methods(self, clean_env): + client = MCPToolClient(base_url="http://localhost:8000") + client.use_production_mode = True + + with pytest.raises(RuntimeError, match="supports only"): + await client.reset() + with pytest.raises(RuntimeError, match="supports only"): + await client.step(ListToolsAction()) + with pytest.raises(RuntimeError, match="supports only"): + await client.state() + + await client.close() + + def test_direct_mode_cannot_change_with_live_websocket(self, clean_env): + client = MCPToolClient(base_url="http://localhost:8000") + client._ws = MagicMock() + + with pytest.raises(RuntimeError, match="transport is active"): + client.use_production_mode = True + + @pytest.mark.asyncio + async def test_direct_mode_list_tools_propagates_jsonrpc_errors(self, clean_env): + client = MCPToolClient(base_url="http://localhost:8000") + client.use_production_mode = True + + with ( + patch.object( + client, + "_ensure_production_session", + new=AsyncMock(return_value="test-session"), + ), + patch.object( + client, + "_production_mcp_request", + new=AsyncMock(return_value={"error": {"message": "transport failed"}}), + ), + ): + with pytest.raises(RuntimeError, match="transport failed"): + await client.list_tools() + + await client.close() + + class TestMCPClientCleanup: """MCP-specific resources must use the common close dispatch path.""" diff --git a/tests/core/test_production_mode_routes.py b/tests/core/test_production_mode_routes.py index b78df5680a..56da30b5b3 100644 --- a/tests/core/test_production_mode_routes.py +++ b/tests/core/test_production_mode_routes.py @@ -27,7 +27,7 @@ import json import sys from pathlib import Path -from unittest.mock import patch +from unittest.mock import AsyncMock, patch import pytest from fastapi import FastAPI @@ -1759,23 +1759,37 @@ class TestMCPClientProductionMode: """Tests for MCP client using production mode.""" async def test_mcp_client_can_use_production_endpoints(self): - """Test MCPToolClient can use production MCP endpoints directly.""" + """Explicit direct mode uses MCP endpoints without opening Gym transport.""" from openenv.core.mcp_client import MCPToolClient client = MCPToolClient(base_url="http://localhost:8000") - - # Client should have option to use production mode (bypasses step()) - assert hasattr(client, "use_production_mode") - client.use_production_mode = True - # Calling list_tools() should use /mcp endpoint, not step() - with patch.object(client, "step") as mock_step: - tools = await client.list_tools() - - # step() should NOT be called in production mode + with ( + patch.object( + client, + "_production_mcp_request", + new=AsyncMock( + side_effect=[ + {"result": {"session_id": "test-session"}}, + {"result": {"tools": []}}, + {"result": {"closed": True}}, + ] + ), + ) as mcp_request, + patch.object(client, "step") as mock_step, + ): + await client.connect() + assert client._ws is None + assert await client.list_tools() == [] mock_step.assert_not_called() - assert len(tools) >= 0 + await client.close() + + assert [call.args[0] for call in mcp_request.await_args_list] == [ + "openenv/session/create", + "tools/list", + "openenv/session/close", + ] @pytest.mark.skip(reason="Implementation detail - httpx is now imported locally") async def test_client_production_mode_uses_http_mcp_endpoint(self): From 5de7a2dc25e307399e7e50665c78e36a7eae32e2 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 16 Sep 2026 11:03:58 +0000 Subject: [PATCH 11/11] chore: separate shared MCP follow-up Co-authored-by: benjamin.burtenshaw --- src/openenv/core/env_server/http_server.py | 90 ++++- src/openenv/core/mcp_client.py | 201 ++++++----- tests/core/test_mode_selection.py | 380 +++++++++++++-------- tests/core/test_production_mode_routes.py | 171 ++++++++-- 4 files changed, 558 insertions(+), 284 deletions(-) diff --git a/src/openenv/core/env_server/http_server.py b/src/openenv/core/env_server/http_server.py index 01954c5bd5..4ddd96a794 100644 --- a/src/openenv/core/env_server/http_server.py +++ b/src/openenv/core/env_server/http_server.py @@ -248,6 +248,8 @@ def __init__( self._session_executors: Dict[str, ThreadPoolExecutor] = {} self._session_stacks: Dict[str, AsyncExitStack] = {} self._session_info: Dict[str, SessionInfo] = {} + self._session_websocket_attachments: set[str] = set() + self._session_pending_closes: set[str] = set() self._session_lock = asyncio.Lock() # Create thread pool for running sync code in async context @@ -451,6 +453,8 @@ async def _destroy_session(self, session_id: str) -> None: executor = self._session_executors.pop(session_id, None) stack = self._session_stacks.pop(session_id, None) self._session_info.pop(session_id, None) + self._session_websocket_attachments.discard(session_id) + self._session_pending_closes.discard(session_id) await self._cleanup_session_resources(env, executor, stack) @@ -522,7 +526,10 @@ async def _reap_idle_sessions(self) -> None: stale_ids: list[str] = [] async with self._session_lock: for sid, info in self._session_info.items(): - if now - info.last_activity_at > timeout: + if ( + sid not in self._session_websocket_attachments + and now - info.last_activity_at > timeout + ): stale_ids.append(sid) for sid in stale_ids: # Re-check under lock: activity may have arrived since @@ -532,7 +539,11 @@ async def _reap_idle_sessions(self) -> None: now = time.time() async with self._session_lock: info = self._session_info.get(sid) - if info is None or (now - info.last_activity_at) <= timeout: + if ( + info is None + or sid in self._session_websocket_attachments + or (now - info.last_activity_at) <= timeout + ): continue await self._destroy_session(sid) except asyncio.CancelledError: @@ -804,15 +815,33 @@ async def mcp_handler( ) async with self._session_lock: - env = self._sessions.pop(target_session_id, _MISSING) - if env is not _MISSING: + if target_session_id in self._session_websocket_attachments: + env = _MISSING + attached = True + self._session_pending_closes.add(target_session_id) + executor = None + stack = None + else: + attached = False + env = self._sessions.pop(target_session_id, _MISSING) + if not attached and env is not _MISSING: executor = self._session_executors.pop(target_session_id, None) stack = self._session_stacks.pop(target_session_id, None) self._session_info.pop(target_session_id, None) - else: + elif not attached: executor = None stack = None + if attached: + return JsonRpcResponse.success( + result={ + "session_id": target_session_id, + "closed": False, + "closing": True, + }, + request_id=request_id, + ) + if env is _MISSING: return JsonRpcResponse.error_response( JsonRpcErrorCode.INVALID_PARAMS, @@ -1487,18 +1516,48 @@ async def websocket_endpoint(websocket: WebSocket): """ WebSocket endpoint for persistent environment sessions. - Each WebSocket connection gets its own environment instance. The client sends - WSResetMessage, WSStepMessage, WSStateMessage, or WSCloseMessage; the server - responds with WSObservationResponse, WSStateResponse, or WSErrorResponse. + By default, each WebSocket connection gets its own environment instance. + A client can instead attach to an existing HTTP MCP session by passing its + session ID in the query string. The client sends WSResetMessage, + WSStepMessage, WSStateMessage, or WSCloseMessage; the server responds with + WSObservationResponse, WSStateResponse, or WSErrorResponse. """ await websocket.accept() session_id = None session_env = None + owns_session = False + attached_session = False try: - # Create session with dedicated environment - session_id, session_env = await self._create_session() + requested_session_id = websocket.query_params.get("session_id") + if requested_session_id: + async with self._session_lock: + attached_env = self._sessions.get( + requested_session_id, _MISSING + ) + if attached_env is _MISSING: + raise RuntimeError( + f"Unknown session_id: {requested_session_id}" + ) + if attached_env is None: + raise RuntimeError( + f"Session {requested_session_id} is still initializing" + ) + if requested_session_id in self._session_websocket_attachments: + raise RuntimeError( + f"Session {requested_session_id} already has " + "an attached WebSocket" + ) + self._session_websocket_attachments.add(requested_session_id) + session_id = requested_session_id + session_env = attached_env + attached_session = True + self._update_session_activity(session_id) + else: + session_id, session_env = await self._create_session() + owns_session = True + if session_env is None: raise RuntimeError( "Session environment not initialized for websocket" @@ -1509,7 +1568,7 @@ async def websocket_endpoint(websocket: WebSocket): async with AsyncExitStack() as stack: mcp_session_factory = getattr(session_env, "mcp_session", None) - if callable(mcp_session_factory): + if owns_session and callable(mcp_session_factory): mcp_session_cm = cast( AsyncContextManager[Any], mcp_session_factory() ) @@ -1688,7 +1747,14 @@ async def websocket_endpoint(websocket: WebSocket): ) await websocket.send_text(error_resp.model_dump_json()) finally: - if session_id: + if attached_session and session_id: + # Do not await the lock before releasing ownership: a task + # cancelled while waiting would leave this session + # permanently exempt from idle reaping. + self._session_websocket_attachments.discard(session_id) + if session_id in self._session_pending_closes: + await asyncio.shield(self._destroy_session(session_id)) + elif owns_session and session_id: await self._destroy_session(session_id) try: await websocket.close() diff --git a/src/openenv/core/mcp_client.py b/src/openenv/core/mcp_client.py index 6fa7b56ea3..e172f32eef 100644 --- a/src/openenv/core/mcp_client.py +++ b/src/openenv/core/mcp_client.py @@ -20,14 +20,14 @@ │ /mcp → MCP JSON-RPC (tools/list, tools/call) │ │ /reset, /step, /state → HTTP endpoints │ ├─────────────────────────────────────────────────────────┤ - │ Explicit direct mode (use_production_mode=True): │ + │ Production Mode (use_production_mode=True): │ │ /mcp → MCP JSON-RPC (tools/list, tools/call) │ - │ Tools only; Gym reset/step/state are unavailable │ + │ Bypasses step() for direct tool access │ └─────────────────────────────────────────────────────────┘ Client Usage: MCPToolClient (default) → /ws (step-based, with rewards) - MCPToolClient (direct opt-in) → /mcp (tools only, no rewards) + MCPToolClient (production) → /mcp (direct tool access, no rewards) Examples: @@ -56,6 +56,7 @@ import asyncio from typing import Any, Dict, List, Optional +from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit from pydantic import ConfigDict @@ -154,60 +155,32 @@ def __init__( mode=mode, ) self._tools_cache: Optional[List[Tool]] = None - self._use_production_mode = False + self.use_production_mode = self._mode == "production" self._production_session_id: Optional[str] = None + self._production_connect_lock = asyncio.Lock() self._production_session_lock = asyncio.Lock() self._jsonrpc_request_id = 0 self._http_client: Optional[Any] = None # lazily-created httpx.AsyncClient - @property - def use_production_mode(self) -> bool: - """Whether explicit tools-only HTTP MCP routing is enabled.""" - return self._use_production_mode - - @use_production_mode.setter - def use_production_mode(self, value: bool) -> None: - """Enable or disable tools-only HTTP MCP routing before connecting.""" - if not isinstance(value, bool): - raise TypeError("use_production_mode must be a bool") - current = getattr(self, "_use_production_mode", False) - has_live_transport = ( - getattr(self, "_ws", None) is not None - or getattr(self, "_production_session_id", None) is not None - or getattr(self, "_http_client", None) is not None - ) - if value != current and has_live_transport: - raise RuntimeError( - "use_production_mode cannot change while a client transport is active" - ) - self._use_production_mode = value - - async def _connect_async(self) -> EnvClient: - """Connect the Gym WebSocket, or prepare explicit tools-only mode.""" - if not self.use_production_mode: - return await super()._connect_async() - - try: - self._start_provider_if_needed() - except BaseException: - try: - await asyncio.shield(self._close_async()) - except BaseException: - pass # Preserve the original startup failure - raise - return self - def _next_request_id(self) -> int: """Generate a monotonically increasing JSON-RPC request id.""" self._jsonrpc_request_id += 1 return self._jsonrpc_request_id def _production_mcp_url(self) -> str: - """Build HTTP MCP endpoint URL from the client's websocket URL.""" - url = self._ws_url.replace("ws://", "http://").replace("wss://", "https://") - if url.endswith("/ws"): - url = url[: -len("/ws")] - return url.rstrip("/") + "/mcp" + """Build the HTTP MCP endpoint URL from the stable base URL.""" + if self._base_url is None: + raise RuntimeError("MCP client is not connected to a server.") + parts = urlsplit(self._base_url) + scheme = {"ws": "http", "wss": "https"}.get(parts.scheme, parts.scheme) + return urlunsplit( + parts._replace( + scheme=scheme, + path=parts.path.rstrip("/") + "/mcp", + query="", + fragment="", + ) + ) async def _get_http_client(self) -> Any: """Return a shared httpx.AsyncClient, creating one lazily.""" @@ -235,6 +208,41 @@ async def _production_mcp_request( response.raise_for_status() return response.json() + async def _connect_async(self) -> EnvClient: + """ + Establish connection to the server. + + In production mode (use_production_mode=True), creates an HTTP MCP session + and connects the WebSocket using that session ID so that WebSocket (reset/step/state) + and HTTP MCP (list_tools/call_tool) share the exact same server-side environment session. + """ + if getattr(self, "use_production_mode", False): + async with self._production_connect_lock: + try: + self._start_provider_if_needed() + session_id = await self._ensure_production_session() + original_ws_url = self._ws_url + if original_ws_url is None: + raise RuntimeError("MCP client has no WebSocket URL.") + + parts = urlsplit(original_ws_url) + query = dict(parse_qsl(parts.query, keep_blank_values=True)) + query["session_id"] = session_id + self._ws_url = urlunsplit(parts._replace(query=urlencode(query))) + try: + await super()._connect_async() + finally: + self._ws_url = original_ws_url + except BaseException: + # Cancellation after the HTTP session is allocated must + # release that session and any started provider before the + # cancellation propagates. + await self.close() + raise + return self + + return await super()._connect_async() + async def _ensure_production_session(self) -> str: """Create and cache a persistent HTTP MCP session id if needed.""" async with self._production_session_lock: @@ -253,27 +261,6 @@ async def _ensure_production_session(self) -> str: self._production_session_id = session_id return session_id - def _tools_only_error(self, operation: str) -> RuntimeError: - return RuntimeError( - f"{operation} is unavailable while use_production_mode=True; " - "direct MCP mode supports only list_tools() and call_tool()" - ) - - async def _reset_async(self, **kwargs: Any) -> StepResult[Observation]: - if self.use_production_mode: - raise self._tools_only_error("reset()") - return await super()._reset_async(**kwargs) - - async def _step_async(self, action: Any, **kwargs: Any) -> StepResult[Observation]: - if self.use_production_mode: - raise self._tools_only_error("step()") - return await super()._step_async(action, **kwargs) - - async def _state_async(self) -> State: - if self.use_production_mode: - raise self._tools_only_error("state()") - return await super()._state_async() - async def list_tools(self, use_cache: bool = True) -> List[Tool]: """ Discover available tools from the environment. @@ -299,22 +286,23 @@ async def list_tools(self, use_cache: bool = True) -> List[Tool]: # Use production mode HTTP endpoint if enabled. # Some tests instantiate with __new__ and skip __init__, so default missing flag to False. if getattr(self, "use_production_mode", False): - session_id = await self._ensure_production_session() - data = await self._production_mcp_request( - "tools/list", - {"session_id": session_id}, - ) - if "error" in data: - message = data.get("error", {}).get("message", "unknown error") - raise RuntimeError(f"list_tools failed: {message}") - result = data.get("result") - if not isinstance(result, dict) or not isinstance( - result.get("tools"), list - ): - raise RuntimeError("list_tools failed: malformed JSON-RPC result") - tools = [_tool_from_payload(t) for t in result["tools"]] - self._tools_cache = tools - return tools + try: + session_id = await self._ensure_production_session() + data = await self._production_mcp_request( + "tools/list", + {"session_id": session_id}, + ) + if "error" in data: + message = data.get("error", {}).get("message", "unknown error") + raise RuntimeError(f"list_tools failed: {message}") + if "result" in data and "tools" in data["result"]: + tools = [_tool_from_payload(t) for t in data["result"]["tools"]] + self._tools_cache = tools + return tools + except Exception: + # If HTTP request fails, return empty list + pass + return [] result = await self.step(ListToolsAction()) if isinstance(result.observation, ListToolsObservation): @@ -401,34 +389,43 @@ async def _close_async(self) -> None: Close client resources. In production MCP mode, this also closes the server-side persistent - MCP session (best effort) before closing websocket/provider resources. + MCP session (best effort) after detaching the WebSocket and before + closing HTTP/provider resources. - This overrides the internal coroutine so sync and async dispatch paths - share the same cleanup. + Override `_close_async` rather than `close` so sync teardown + (`SyncEnvClient.close`, sync `__exit__`, and `_dispatch`) still cleans + up the HTTP MCP session. """ try: - if self._production_session_id is not None: - try: - await self._production_mcp_request( - "openenv/session/close", - {"session_id": self._production_session_id}, - ) - except Exception: - # Best effort cleanup - do not mask normal close behavior - pass - finally: - self._production_session_id = None + # The WebSocket shares the HTTP-created session. Detach it first so + # the server's ownership guard permits the explicit session close. + await self._disconnect_async() finally: try: - if self._http_client is not None: + if self._production_session_id is not None: try: - await self._http_client.aclose() + await self._production_mcp_request( + "openenv/session/close", + {"session_id": self._production_session_id}, + ) except Exception: - pass # Best effort; continue to websocket/provider teardown + # Best effort cleanup - do not mask normal close behavior + pass finally: - self._http_client = None + self._production_session_id = None finally: - await super()._close_async() + try: + if self._http_client is not None: + try: + await self._http_client.aclose() + except Exception: + pass + finally: + self._http_client = None + finally: + # This is intentionally inside the outer finally so + # cancellation cannot skip provider teardown. + await super()._close_async() class MCPToolClient(MCPClientBase): diff --git a/tests/core/test_mode_selection.py b/tests/core/test_mode_selection.py index a4aab83b0d..ad9465f171 100644 --- a/tests/core/test_mode_selection.py +++ b/tests/core/test_mode_selection.py @@ -167,6 +167,20 @@ def test_invalid_env_var_raises_error(self): class TestModeBehavior: """Test that different modes result in different client behavior.""" + @pytest.mark.parametrize( + ("base_url", "expected_url"), + [ + ("http://localhost:8000", "http://localhost:8000/mcp"), + ("https://example.com/env", "https://example.com/env/mcp"), + ("ws://localhost:8000", "http://localhost:8000/mcp"), + ("wss://example.com/env", "https://example.com/env/mcp"), + ], + ) + def test_production_mcp_url_uses_http_scheme(self, base_url, expected_url): + """HTTP MCP requests normalize WebSocket base URL schemes.""" + client = MCPToolClient(base_url=base_url, mode="production") + assert client._production_mcp_url() == expected_url + @pytest.mark.asyncio async def test_simulation_mode_uses_gym_protocol(self, clean_env, mock_websocket): """Test that simulation mode uses Gym-style WebSocket messages.""" @@ -194,191 +208,268 @@ async def test_simulation_mode_uses_gym_protocol(self, clean_env, mock_websocket ) @pytest.mark.asyncio - async def test_production_mode_uses_jsonrpc_protocol( - self, clean_env, mock_websocket - ): - """Test that production mode uses JSON-RPC format for tool calls.""" + async def test_production_mode_uses_jsonrpc_protocol(self, clean_env): + """Test that production mode uses HTTP JSON-RPC format for tool listing.""" client = MCPToolClient(base_url="http://localhost:8000", mode="production") + assert client.use_production_mode is True - with patch.object(client, "_send") as mock_send: - with patch.object( - client, - "_receive", - return_value={ - "type": "response", - "data": { - "observation": {"tools": []}, - "reward": None, - "done": False, - }, + with patch.object( + client, + "_production_mcp_request", + side_effect=[ + {"result": {"session_id": "test-session"}}, + { + "result": { + "tools": [ + { + "name": "echo", + "description": "Echo message", + "inputSchema": {}, + } + ] + } }, - ): - with patch.object(client, "_ws", mock_websocket): - await client.list_tools() - - # Should send step message with list_tools action - call_args = mock_send.call_args_list - step_call = [ - call for call in call_args if call[0][0].get("type") == "step" - ] - assert len(step_call) > 0, "Should send message with type='step'" - - # Check that the action payload is list_tools - step_message = step_call[0][0][0] - assert "data" in step_message - assert step_message["data"].get("type") == "list_tools" - - -class TestDirectMCPMode: - """Explicit HTTP MCP mode must remain tools-only and single-transport.""" + ], + ) as mock_mcp_request: + with patch.object(client, "step") as mock_step: + tools = await client.list_tools() + + mock_step.assert_not_called() + assert len(tools) == 1 + assert tools[0].name == "echo" + assert mock_mcp_request.call_count == 2 + mock_mcp_request.assert_called_with( + "tools/list", {"session_id": "test-session"} + ) @pytest.mark.asyncio - async def test_direct_mode_connect_does_not_open_websocket(self, clean_env): - client = MCPToolClient(base_url="http://localhost:8000") - client.use_production_mode = True - - with patch( - "openenv.core.env_client.ws_connect", new_callable=AsyncMock - ) as ws_connect: - await client.connect() + async def test_production_mode_call_tool_uses_jsonrpc_protocol(self, clean_env): + """Test that call_tool in production mode uses HTTP JSON-RPC transport.""" + client = MCPToolClient(base_url="http://localhost:8000", mode="production") + assert client.use_production_mode is True - ws_connect.assert_not_awaited() - assert client._ws is None - assert client._production_session_id is None - await client.close() + with patch.object( + client, + "_production_mcp_request", + side_effect=[ + {"result": {"session_id": "test-session"}}, + {"result": {"data": "hello world"}}, + ], + ) as mock_mcp_request: + with patch.object(client, "step") as mock_step: + result = await client.call_tool("echo", message="hello world") + + mock_step.assert_not_called() + assert result == "hello world" + mock_mcp_request.assert_called_with( + "tools/call", + { + "name": "echo", + "arguments": {"message": "hello world"}, + "session_id": "test-session", + }, + ) @pytest.mark.asyncio - async def test_direct_mode_rejects_gym_lifecycle_methods(self, clean_env): - client = MCPToolClient(base_url="http://localhost:8000") - client.use_production_mode = True + async def test_production_mode_connect_creates_single_session_with_websocket( + self, clean_env + ): + """Test that connect() in production mode initializes the HTTP MCP session AND connects WebSocket using the same session ID.""" + client = MCPToolClient(base_url="http://localhost:8000", mode="production") + assert client.use_production_mode is True + client._ws_url = f"{client._ws_url}?some_session_id=keep" + original_ws_url = client._ws_url - with pytest.raises(RuntimeError, match="supports only"): - await client.reset() - with pytest.raises(RuntimeError, match="supports only"): - await client.step(ListToolsAction()) - with pytest.raises(RuntimeError, match="supports only"): - await client.state() + with patch.object( + client, + "_production_mcp_request", + side_effect=[ + {"result": {"session_id": "test-session"}}, + {"result": {"data": "hello world"}}, + ], + ) as mock_mcp_request: + with patch( + "openenv.core.env_client.ws_connect", new_callable=AsyncMock + ) as mock_ws_connect: + # Explicit connect (e.g. from async with client:) + await client.connect() + + # Should create HTTP session and connect WS with session_id query param + mock_ws_connect.assert_called_once() + connected_url = mock_ws_connect.call_args[0][0] + assert "session_id=test-session" in connected_url + assert "some_session_id=keep" in connected_url + assert client._ws_url == original_ws_url + assert client._production_session_id == "test-session" + mock_mcp_request.assert_called_once_with("openenv/session/create") + + # Subsequent call_tool should reuse the same session + result = await client.call_tool("echo", message="hello world") + assert result == "hello world" + assert mock_mcp_request.call_count == 2 + mock_mcp_request.assert_called_with( + "tools/call", + { + "name": "echo", + "arguments": {"message": "hello world"}, + "session_id": "test-session", + }, + ) - await client.close() + @pytest.mark.asyncio + async def test_production_mode_connect_failure_cleans_up_resources(self, clean_env): + """Test that failure during production mode connect() triggers client.close() cleanup.""" + client = MCPToolClient(base_url="http://localhost:8000", mode="production") + assert client.use_production_mode is True - def test_direct_mode_cannot_change_with_live_websocket(self, clean_env): - client = MCPToolClient(base_url="http://localhost:8000") - client._ws = MagicMock() + with patch.object( + client, + "_ensure_production_session", + side_effect=RuntimeError("Session creation failed"), + ): + with patch.object(client, "close", wraps=client.close) as mock_close: + with pytest.raises(RuntimeError, match="Session creation failed"): + await client.connect() - with pytest.raises(RuntimeError, match="transport is active"): - client.use_production_mode = True + mock_close.assert_called_once() @pytest.mark.asyncio - async def test_direct_mode_list_tools_propagates_jsonrpc_errors(self, clean_env): - client = MCPToolClient(base_url="http://localhost:8000") - client.use_production_mode = True + async def test_cancelled_production_connect_closes_allocated_session( + self, clean_env + ): + """Cancellation during WebSocket connect releases the HTTP session.""" + client = MCPToolClient(base_url="http://localhost:8000", mode="production") + mock_http_client = AsyncMock() + client._http_client = mock_http_client with ( - patch.object( - client, - "_ensure_production_session", - new=AsyncMock(return_value="test-session"), - ), patch.object( client, "_production_mcp_request", - new=AsyncMock(return_value={"error": {"message": "transport failed"}}), + side_effect=[ + {"result": {"session_id": "test-session"}}, + {"result": {"session_id": "test-session", "closed": True}}, + ], + ) as mock_mcp_request, + patch( + "openenv.core.env_client.ws_connect", + new_callable=AsyncMock, + side_effect=asyncio.CancelledError, ), ): - with pytest.raises(RuntimeError, match="transport failed"): - await client.list_tools() + with pytest.raises(asyncio.CancelledError): + await client.connect() - await client.close() + mock_mcp_request.assert_any_call( + "openenv/session/close", + {"session_id": "test-session"}, + ) + assert client._production_session_id is None + mock_http_client.aclose.assert_awaited_once() + assert client._http_client is None + def test_production_mode_sync_close_closes_mcp_session(self, clean_env): + """Test that production sync close() closes the MCP session and releases HTTP client.""" + client = MCPToolClient( + base_url="http://localhost:8000", mode="production" + ).sync() + client._async._production_session_id = "test-session-sync" -class TestMCPClientCleanup: - """MCP-specific resources must use the common close dispatch path.""" + mock_http_client = AsyncMock() + client._async._http_client = mock_http_client - def test_bare_close_preserves_sync_dispatch(self, clean_env): - """Calling close outside an event loop must resolve cleanup synchronously.""" - client = MCPToolClient(base_url="http://localhost:8000") - client._production_session_id = "test-session" + with patch.object( + client._async, + "_production_mcp_request", + new_callable=AsyncMock, + return_value={"result": {"status": "closed"}}, + ) as mock_mcp_req: + client.close() + + mock_mcp_req.assert_awaited_once_with( + "openenv/session/close", + {"session_id": "test-session-sync"}, + ) + assert client._async._production_session_id is None + mock_http_client.aclose.assert_awaited_once() + assert client._async._http_client is None - try: - with patch.object( - client, - "_production_mcp_request", - new=AsyncMock(return_value={"result": {}}), - ) as request: - result = client.close() + def test_production_mode_sync_context_manager_closes_mcp_session(self, clean_env): + """Test that production sync context-manager exit closes the MCP session and releases HTTP client.""" + client = MCPToolClient( + base_url="http://localhost:8000", mode="production" + ).sync() + client._async._production_session_id = "test-session-context" + + mock_http_client = AsyncMock() + client._async._http_client = mock_http_client - assert result is None - request.assert_awaited_once_with( - "openenv/session/close", {"session_id": "test-session"} + with patch.object( + client._async, + "_production_mcp_request", + new_callable=AsyncMock, + return_value={"result": {"status": "closed"}}, + ) as mock_mcp_req: + with patch.object(client._async, "_connect_async", new_callable=AsyncMock): + with client: + pass + + mock_mcp_req.assert_awaited_once_with( + "openenv/session/close", + {"session_id": "test-session-context"}, ) - assert client._production_session_id is None - finally: - if client._sync_client is not None: - client._sync_client._stop_loop() + assert client._async._production_session_id is None + mock_http_client.aclose.assert_awaited_once() + assert client._async._http_client is None - def test_sync_wrapper_close_releases_http_resources(self, clean_env): - """Sync wrapper close must release the MCP session and HTTP client.""" - client = MCPToolClient(base_url="http://localhost:8000") - client._production_session_id = "test-session" - http_client = AsyncMock() - client._http_client = http_client + @pytest.mark.asyncio + async def test_production_mode_async_close_closes_mcp_session(self, clean_env): + """Test that production async close() closes the MCP session and releases HTTP client.""" + client = MCPToolClient(base_url="http://localhost:8000", mode="production") + client._production_session_id = "test-session-async" + + mock_http_client = AsyncMock() + client._http_client = mock_http_client with patch.object( client, "_production_mcp_request", - new=AsyncMock(return_value={"result": {}}), - ) as request: - client.sync().close() - - request.assert_awaited_once_with( - "openenv/session/close", {"session_id": "test-session"} - ) - http_client.aclose.assert_awaited_once_with() - assert client._production_session_id is None - assert client._http_client is None + new_callable=AsyncMock, + return_value={"result": {"status": "closed"}}, + ) as mock_mcp_req: + await client.close() + + mock_mcp_req.assert_awaited_once_with( + "openenv/session/close", + {"session_id": "test-session-async"}, + ) + assert client._production_session_id is None + mock_http_client.aclose.assert_awaited_once() + assert client._http_client is None @pytest.mark.asyncio - async def test_cancelled_session_close_still_releases_http_and_provider( + async def test_production_close_detaches_websocket_before_session_close( self, clean_env ): - """Cancellation must not bypass later HTTP and provider cleanup.""" - - class Provider: - def __init__(self): - self.stopped = False - - def stop(self): - self.stopped = True - - close_started = asyncio.Event() + """Shared WebSocket ownership is released before HTTP session teardown.""" + client = MCPToolClient(base_url="http://localhost:8000", mode="production") + client._production_session_id = "test-session" + teardown_events = [] - async def slow_session_close(*_args, **_kwargs): - close_started.set() - await asyncio.sleep(10) + async def disconnect(): + teardown_events.append("websocket") - provider = Provider() - client = MCPToolClient( - base_url="http://localhost:8000", - provider=provider, - ) - client._production_session_id = "test-session" - http_client = AsyncMock() - client._http_client = http_client + async def request(method, params=None): + teardown_events.append("session") + return {"result": {"closed": True}} - with patch.object( - client, "_production_mcp_request", side_effect=slow_session_close + with ( + patch.object(client, "_disconnect_async", side_effect=disconnect), + patch.object(client, "_production_mcp_request", side_effect=request), ): - close_call = asyncio.create_task(client._close_async()) - await close_started.wait() - close_call.cancel() - close_results = await asyncio.gather(close_call, return_exceptions=True) - - assert len(close_results) == 1 - assert isinstance(close_results[0], asyncio.CancelledError) - http_client.aclose.assert_awaited_once_with() - assert client._production_session_id is None - assert client._http_client is None - assert provider.stopped + await client.close() + + assert teardown_events[:2] == ["websocket", "session"] # ============================================================================ @@ -436,6 +527,7 @@ def test_mcp_client_defaults_to_production_mode(self, clean_env): # MCPToolClient should default to production mode assert client._mode == "production" + assert client.use_production_mode is True def test_mcp_client_cannot_use_simulation_mode(self, clean_env): """Test that MCPToolClient raises error if simulation mode is requested.""" diff --git a/tests/core/test_production_mode_routes.py b/tests/core/test_production_mode_routes.py index 56da30b5b3..737defd188 100644 --- a/tests/core/test_production_mode_routes.py +++ b/tests/core/test_production_mode_routes.py @@ -27,7 +27,7 @@ import json import sys from pathlib import Path -from unittest.mock import AsyncMock, patch +from unittest.mock import patch import pytest from fastapi import FastAPI @@ -786,6 +786,139 @@ def test_session_create_from_websocket_is_idempotent(self, app): response2 = json.loads(response_text2) assert response2["data"]["result"]["session_id"] == ws_session_id + def test_websocket_can_attach_to_http_session_without_destroying_it(self, app): + """An attached WebSocket shares and preserves an HTTP-created session.""" + client = TestClient(app) + create_response = client.post( + "/mcp", + json={ + "jsonrpc": "2.0", + "method": "openenv/session/create", + "params": {}, + "id": 1, + }, + ) + session_id = create_response.json()["result"]["session_id"] + + with client.websocket_connect(f"/ws?session_id={session_id}") as websocket: + websocket.send_json({"type": "state"}) + state_response = websocket.receive_json() + assert state_response["type"] == "state" + + active_close_response = client.post( + "/mcp", + json={ + "jsonrpc": "2.0", + "method": "openenv/session/close", + "params": {"session_id": session_id}, + "id": 2, + }, + ) + close_result = active_close_response.json()["result"] + assert close_result == { + "session_id": session_id, + "closed": False, + "closing": True, + } + websocket.send_json({"type": "close"}) + + tools_response = client.post( + "/mcp", + json={ + "jsonrpc": "2.0", + "method": "tools/list", + "params": {"session_id": session_id}, + "id": 3, + }, + ) + assert tools_response.json()["error"]["code"] == -32602 + + replacement_response = client.post( + "/mcp", + json={ + "jsonrpc": "2.0", + "method": "openenv/session/create", + "params": {}, + "id": 4, + }, + ) + replacement_id = replacement_response.json()["result"]["session_id"] + client.post( + "/mcp", + json={ + "jsonrpc": "2.0", + "method": "openenv/session/close", + "params": {"session_id": replacement_id}, + "id": 5, + }, + ) + + def test_http_session_allows_only_one_attached_websocket(self, app): + """A second WebSocket cannot concurrently mutate the same session.""" + client = TestClient(app) + create_response = client.post( + "/mcp", + json={ + "jsonrpc": "2.0", + "method": "openenv/session/create", + "params": {}, + "id": 1, + }, + ) + session_id = create_response.json()["result"]["session_id"] + + with client.websocket_connect(f"/ws?session_id={session_id}") as first_socket: + with client.websocket_connect( + f"/ws?session_id={session_id}" + ) as second_socket: + error_response = second_socket.receive_json() + assert ( + "already has an attached WebSocket" + in (error_response["data"]["message"]) + ) + first_socket.send_json({"type": "close"}) + + close_response = client.post( + "/mcp", + json={ + "jsonrpc": "2.0", + "method": "openenv/session/close", + "params": {"session_id": session_id}, + "id": 2, + }, + ) + assert close_response.json()["result"]["closed"] is True + + def test_websocket_still_destroys_its_own_session(self, app): + """A WebSocket-created session is destroyed when the socket closes.""" + client = TestClient(app) + + with client.websocket_connect("/ws") as websocket: + websocket.send_json( + { + "type": "mcp", + "data": { + "jsonrpc": "2.0", + "method": "openenv/session/create", + "params": {}, + "id": 1, + }, + } + ) + session_id = websocket.receive_json()["data"]["result"]["session_id"] + websocket.send_json({"type": "close"}) + + tools_response = client.post( + "/mcp", + json={ + "jsonrpc": "2.0", + "method": "tools/list", + "params": {"session_id": session_id}, + "id": 2, + }, + ) + assert tools_response.json()["error"]["code"] == -32602 + def test_session_close_missing_session_id_param(self, app): """Test openenv/session/close without session_id returns INVALID_PARAMS.""" from starlette.testclient import TestClient @@ -1759,37 +1892,23 @@ class TestMCPClientProductionMode: """Tests for MCP client using production mode.""" async def test_mcp_client_can_use_production_endpoints(self): - """Explicit direct mode uses MCP endpoints without opening Gym transport.""" + """Test MCPToolClient can use production MCP endpoints directly.""" from openenv.core.mcp_client import MCPToolClient client = MCPToolClient(base_url="http://localhost:8000") + + # Client should have option to use production mode (bypasses step()) + assert hasattr(client, "use_production_mode") + client.use_production_mode = True - with ( - patch.object( - client, - "_production_mcp_request", - new=AsyncMock( - side_effect=[ - {"result": {"session_id": "test-session"}}, - {"result": {"tools": []}}, - {"result": {"closed": True}}, - ] - ), - ) as mcp_request, - patch.object(client, "step") as mock_step, - ): - await client.connect() - assert client._ws is None - assert await client.list_tools() == [] - mock_step.assert_not_called() - await client.close() + # Calling list_tools() should use /mcp endpoint, not step() + with patch.object(client, "step") as mock_step: + tools = await client.list_tools() - assert [call.args[0] for call in mcp_request.await_args_list] == [ - "openenv/session/create", - "tools/list", - "openenv/session/close", - ] + # step() should NOT be called in production mode + mock_step.assert_not_called() + assert len(tools) >= 0 @pytest.mark.skip(reason="Implementation detail - httpx is now imported locally") async def test_client_production_mode_uses_http_mcp_endpoint(self):