diff --git a/src/openenv/core/env_server/http_server.py b/src/openenv/core/env_server/http_server.py index 01954c5bd5..14d8d4fb3a 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,52 @@ 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_pending_closes: + raise RuntimeError( + f"Session {requested_session_id} is closing" + ) + 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 +1572,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 +1751,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 1afc6254b9..f9d73ab262 100644 --- a/src/openenv/core/mcp_client.py +++ b/src/openenv/core/mcp_client.py @@ -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 @@ -156,6 +157,7 @@ def __init__( self._tools_cache: Optional[List[Tool]] = None 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 @@ -166,11 +168,19 @@ def _next_request_id(self) -> int: 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.""" @@ -202,19 +212,35 @@ 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. + 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): - try: - await super()._connect_async() - await self._ensure_production_session() - except Exception: - await self.close() - raise + 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: + # CancelledError is a BaseException: cleanup must still run + # after session create so capacity is not leaked. + try: + await asyncio.shield(self.close()) + except Exception: + pass + raise return self return await super()._connect_async() @@ -365,33 +391,44 @@ 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. 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: + # 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: - 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 + 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: - self._production_session_id = None - - if self._http_client is not None: - try: - await self._http_client.aclose() - except Exception: - pass - finally: - self._http_client = None - - await super()._close_async() + try: + if self._http_client is not None: + try: + await self._http_client.aclose() + except Exception: + # Best effort — do not mask provider/websocket teardown + 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 04aa7ef490..fa0a8854a6 100644 --- a/tests/core/test_mode_selection.py +++ b/tests/core/test_mode_selection.py @@ -21,6 +21,7 @@ - Environment: Code mode with mode-aware tool registration """ +import asyncio import os from unittest.mock import AsyncMock, MagicMock, patch @@ -166,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.""" @@ -256,15 +271,14 @@ async def test_production_mode_call_tool_uses_jsonrpc_protocol(self, clean_env): ) @pytest.mark.asyncio - async def test_production_mode_connect_opens_websocket_and_http_session( + async def test_production_mode_connect_creates_single_session_with_websocket( self, clean_env ): - """Production connect must open WebSocket (reset/step/state) and HTTP MCP session.""" + """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 - - mock_ws = MagicMock() - mock_ws.closed = False + client._ws_url = f"{client._ws_url}?some_session_id=keep" + original_ws_url = client._ws_url with patch.object( client, @@ -275,17 +289,21 @@ async def test_production_mode_connect_opens_websocket_and_http_session( ], ) as mock_mcp_request: with patch( - "openenv.core.env_client.ws_connect", - new_callable=AsyncMock, - return_value=mock_ws, + "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() - assert client._ws is mock_ws + 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 @@ -304,61 +322,182 @@ async def test_production_mode_connect_failure_cleans_up_resources(self, clean_e 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, + "_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_close_disconnects_websocket_before_session_close( + self, clean_env + ): + """Close must detach /ws before openenv/session/close (attachment guard).""" + client = MCPToolClient(base_url="http://localhost:8000", mode="production") + client._production_session_id = "sess-attach" + client._ws = MagicMock() + client._ws_loop = None + + calls = [] + + async def fake_disconnect(): + calls.append("disconnect") + client._ws = None + + async def fake_mcp(method, params=None): + calls.append((method, params)) + return {"result": {"closed": True}} + + with patch.object(client, "_disconnect_async", side_effect=fake_disconnect): + with patch.object(client, "_production_mcp_request", side_effect=fake_mcp): + await client.close() + + assert calls[0] == "disconnect" + assert calls[1] == ("openenv/session/close", {"session_id": "sess-attach"}) + assert client._production_session_id is None - with patch( - "openenv.core.env_client.ws_connect", + 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" + + mock_http_client = AsyncMock() + client._async._http_client = mock_http_client + + with patch.object( + client._async, + "_production_mcp_request", new_callable=AsyncMock, - return_value=mock_ws, - ): - 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() + 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 - mock_close.assert_called_once() + 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 + + 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._async._production_session_id is None + mock_http_client.aclose.assert_awaited_once() + assert client._async._http_client is None @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`.""" + 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") - assert client.use_production_mode is True + client._production_session_id = "test-session-async" - mock_ws = MagicMock() - mock_ws.closed = False - mock_ws.close = AsyncMock() + mock_http_client = AsyncMock() + client._http_client = mock_http_client 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, - ): - sync_client = client.sync() - sync_client.connect() - assert client._production_session_id == "test-session" + new_callable=AsyncMock, + return_value={"result": {"status": "closed"}}, + ) as mock_mcp_req: + await client.close() - sync_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 - 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"}, - ) + @pytest.mark.asyncio + async def test_production_close_detaches_websocket_before_session_close( + self, clean_env + ): + """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 disconnect(): + teardown_events.append("websocket") + + async def request(method, params=None): + teardown_events.append("session") + return {"result": {"closed": True}} + + with ( + patch.object(client, "_disconnect_async", side_effect=disconnect), + patch.object(client, "_production_mcp_request", side_effect=request), + ): + await client.close() + + assert teardown_events[:2] == ["websocket", "session"] + + @pytest.mark.asyncio + async def test_production_connect_cancellation_closes_allocated_session( + self, clean_env + ): + """Cancelled WebSocket connect still releases the HTTP MCP session.""" + client = MCPToolClient(base_url="http://localhost:8000", mode="production") + client._ws_url = "ws://localhost:8000/ws" + close_calls = [] + + async def ensure_session(): + client._production_session_id = "cancelled-session" + return "cancelled-session" + + async def connect(*_args, **_kwargs): + raise asyncio.CancelledError() + + async def close(*_args, **_kwargs): + close_calls.append(client._production_session_id) + + with ( + patch.object(client, "_start_provider_if_needed"), + patch.object( + client, "_ensure_production_session", side_effect=ensure_session + ), + patch( + "openenv.core.env_client.EnvClient._connect_async", + side_effect=connect, + ), + patch.object(client, "close", new=close), + ): + with pytest.raises(asyncio.CancelledError): + await client._connect_async() + + assert close_calls == ["cancelled-session"] # ============================================================================ diff --git a/tests/core/test_production_mode_routes.py b/tests/core/test_production_mode_routes.py index b78df5680a..a528806e74 100644 --- a/tests/core/test_production_mode_routes.py +++ b/tests/core/test_production_mode_routes.py @@ -786,6 +786,173 @@ 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_cannot_attach_while_session_close_is_pending(self, app): + """Pending closes refuse reattach so delayed destroy cannot race a new WS.""" + 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: + 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"]["closing"] is True + + with client.websocket_connect( + f"/ws?session_id={session_id}" + ) as second_socket: + error_response = second_socket.receive_json() + assert "is closing" in error_response["data"]["message"] + + websocket.send_json({"type": "close"}) + + 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