diff --git a/src/openenv/core/mcp_client.py b/src/openenv/core/mcp_client.py index 7634bb0be9..1afc6254b9 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 = False + self.use_production_mode = self._mode == "production" self._production_session_id: Optional[str] = None self._production_session_lock = asyncio.Lock() self._jsonrpc_request_id = 0 @@ -198,6 +198,27 @@ 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: @@ -339,12 +360,16 @@ 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. + + 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: @@ -366,7 +391,7 @@ async def close(self) -> 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..04aa7ef490 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 MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastmcp import FastMCP @@ -193,39 +193,172 @@ 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 + 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 ): - """Test that production mode uses JSON-RPC format for tool calls.""" + """Production connect must open WebSocket (reset/step/state) and HTTP MCP session.""" 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", + }, + ) - with patch.object(client, "_send") as mock_send: + @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, - "_receive", - return_value={ - "type": "response", - "data": { - "observation": {"tools": []}, - "reward": None, - "done": False, - }, - }, + "_ensure_production_session", + side_effect=RuntimeError("Session creation failed"), ): - with patch.object(client, "_ws", mock_websocket): - await client.list_tools() + with patch.object(client, "close", wraps=client.close) as mock_close: + with pytest.raises(RuntimeError, match="Session creation failed"): + await client.connect() - # 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'" + 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, + ): + sync_client = client.sync() + sync_client.connect() + assert client._production_session_id == "test-session" + + sync_client.close() - # 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" + 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"}, + ) # ============================================================================ @@ -283,6 +416,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."""