Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 3 additions & 28 deletions src/openenv/core/mcp_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand All @@ -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):
Expand Down
186 changes: 26 additions & 160 deletions tests/core/test_mode_selection.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"


# ============================================================================
Expand Down Expand Up @@ -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."""
Expand Down