Skip to content
Merged
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
103 changes: 54 additions & 49 deletions raven/agent/tools/mcp.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
"""MCP client: connects to MCP servers and wraps their tools as native Raven tools."""

import asyncio
from contextlib import AsyncExitStack
from contextlib import AsyncExitStack, asynccontextmanager
from typing import TYPE_CHECKING, Any

import httpx
Expand All @@ -16,6 +16,54 @@
from raven.sandbox import SandboxExecutor


@asynccontextmanager
async def _mcp_server_connection(cfg, transport_type: str, executor: "SandboxExecutor | None"):
async with AsyncExitStack() as stack:
from mcp import ClientSession, StdioServerParameters
from mcp.client.sse import sse_client
from mcp.client.stdio import stdio_client
from mcp.client.streamable_http import streamable_http_client

if transport_type == "stdio":
if executor is not None and executor.supports_process_spawning:
read, write = await executor.start_process(cfg.command, cfg.args, env=cfg.env or None)
else:
params = StdioServerParameters(command=cfg.command, args=cfg.args, env=cfg.env or None)
read, write = await stack.enter_async_context(stdio_client(params))
elif transport_type == "sse":

def httpx_client_factory(
Comment thread
gloryfromca marked this conversation as resolved.
headers: dict[str, str] | None = None,
timeout: httpx.Timeout | None = None,
auth: httpx.Auth | None = None,
) -> httpx.AsyncClient:
merged_headers = {**(cfg.headers or {}), **(headers or {})}
return httpx.AsyncClient(
headers=merged_headers or None,
follow_redirects=True,
timeout=timeout,
auth=auth,
)

read, write = await stack.enter_async_context(
sse_client(cfg.url, httpx_client_factory=httpx_client_factory)
)
else:
http_client = await stack.enter_async_context(
httpx.AsyncClient(
headers=cfg.headers or None,
follow_redirects=True,
timeout=None,
)
)
read, write, _ = await stack.enter_async_context(streamable_http_client(cfg.url, http_client=http_client))

session = await stack.enter_async_context(ClientSession(read, write))
await session.initialize()
tools = await session.list_tools()
yield session, tools


class MCPToolWrapper(Tool):
"""Wraps a single MCP server tool as an Raven Tool."""

Expand Down Expand Up @@ -110,55 +158,12 @@ async def connect_mcp_servers(
"Either switch to an HTTP/SSE MCP server or set sandbox.backend='none'."
)

try:
from mcp import ClientSession, StdioServerParameters
from mcp.client.sse import sse_client
from mcp.client.stdio import stdio_client
from mcp.client.streamable_http import streamable_http_client

if transport_type == "stdio":
if executor is not None and executor.supports_process_spawning:
read, write = await executor.start_process(cfg.command, cfg.args, env=cfg.env or None)
else:
params = StdioServerParameters(command=cfg.command, args=cfg.args, env=cfg.env or None)
read, write = await stack.enter_async_context(stdio_client(params))
elif transport_type == "sse":

def httpx_client_factory(
headers: dict[str, str] | None = None,
timeout: httpx.Timeout | None = None,
auth: httpx.Auth | None = None,
) -> httpx.AsyncClient:
merged_headers = {**(cfg.headers or {}), **(headers or {})}
return httpx.AsyncClient(
headers=merged_headers or None,
follow_redirects=True,
timeout=timeout,
auth=auth,
)

read, write = await stack.enter_async_context(
sse_client(cfg.url, httpx_client_factory=httpx_client_factory)
)
elif transport_type == "streamableHttp":
http_client = await stack.enter_async_context(
httpx.AsyncClient(
headers=cfg.headers or None,
follow_redirects=True,
timeout=None,
)
)
read, write, _ = await stack.enter_async_context(
streamable_http_client(cfg.url, http_client=http_client)
)
else:
logger.warning("MCP server '{}': unknown transport type '{}'", name, transport_type)
continue
if transport_type not in {"stdio", "sse", "streamableHttp"}:
logger.warning("MCP server '{}': unknown transport type '{}'", name, transport_type)
continue

session = await stack.enter_async_context(ClientSession(read, write))
await session.initialize()

tools = await session.list_tools()
try:
session, tools = await stack.enter_async_context(_mcp_server_connection(cfg, transport_type, executor))
for tool_def in tools.tools:
wrapper = MCPToolWrapper(session, name, tool_def, tool_timeout=cfg.tool_timeout)
registry.register(wrapper)
Expand Down
215 changes: 215 additions & 0 deletions tests/test_sandbox_unit.py
Original file line number Diff line number Diff line change
Expand Up @@ -1052,6 +1052,221 @@ def fake_stdio_client(params):

assert reached == ["mcp-server"]

async def test_streamable_failure_does_not_cancel_following_server(self, monkeypatch):
"""A transport task failure is isolated to the server being initialized."""
from contextlib import AsyncExitStack, asynccontextmanager
from types import SimpleNamespace

import anyio
import mcp
import mcp.client.streamable_http

from raven.agent.tools.mcp import connect_mcp_servers
from raven.agent.tools.registry import ToolRegistry

attempted = []

@asynccontextmanager
async def fake_streamable_http_client(url, http_client):
attempted.append(url)
if url == "https://bad.example/mcp":
async with anyio.create_task_group() as group:

async def fail_transport():
await anyio.sleep(0)
raise RuntimeError("transport failed")

group.start_soon(fail_transport)
yield url, object(), None
else:
yield url, object(), None

class FakeSession:
def __init__(self, read, write):
self.read = read

async def __aenter__(self):
return self

async def __aexit__(self, exc_type, exc, traceback):
return False

async def initialize(self):
if self.read == "https://bad.example/mcp":
await asyncio.Event().wait()

async def list_tools(self):
return SimpleNamespace(tools=[])

monkeypatch.setattr(mcp, "ClientSession", FakeSession)
monkeypatch.setattr(
mcp.client.streamable_http,
"streamable_http_client",
fake_streamable_http_client,
)

def config(url):
return SimpleNamespace(type="streamableHttp", url=url, headers=None, tool_timeout=30)

async with AsyncExitStack() as stack:
await connect_mcp_servers(
{
"bad": config("https://bad.example/mcp"),
"good": config("https://good.example/mcp"),
},
ToolRegistry(),
stack,
)

assert attempted == ["https://bad.example/mcp", "https://good.example/mcp"]
assert asyncio.current_task().cancelling() == 0

async def test_streamable_external_cancellation_propagates(self, monkeypatch):
"""Cancellation of Raven's connection task is not treated as a server failure."""
from contextlib import AsyncExitStack, asynccontextmanager
from types import SimpleNamespace

import mcp
import mcp.client.streamable_http

from raven.agent.tools.mcp import connect_mcp_servers
from raven.agent.tools.registry import ToolRegistry

entered = asyncio.Event()

@asynccontextmanager
async def fake_streamable_http_client(url, http_client):
yield object(), object(), None

class FakeSession:
def __init__(self, read, write):
pass

async def __aenter__(self):
return self

async def __aexit__(self, exc_type, exc, traceback):
return False

async def initialize(self):
entered.set()
await asyncio.Event().wait()

monkeypatch.setattr(mcp, "ClientSession", FakeSession)
monkeypatch.setattr(
mcp.client.streamable_http,
"streamable_http_client",
fake_streamable_http_client,
)

cfg = SimpleNamespace(
type="streamableHttp",
url="https://wait.example/mcp",
headers=None,
tool_timeout=30,
)
task = asyncio.create_task(connect_mcp_servers({"svc": cfg}, ToolRegistry(), AsyncExitStack()))
await entered.wait()
task.cancel()

with pytest.raises(asyncio.CancelledError):
await task

async def test_sse_connection_uses_merged_headers(self, monkeypatch):
"""The SSE transport keeps config headers when the SDK adds its own."""
from contextlib import AsyncExitStack, asynccontextmanager
from types import SimpleNamespace

import httpx
import mcp
import mcp.client.sse

from raven.agent.tools.mcp import connect_mcp_servers
from raven.agent.tools.registry import ToolRegistry

clients = []

class FakeHttpClient:
async def __aenter__(self):
return self

async def __aexit__(self, exc_type, exc, traceback):
return False

def fake_http_client(**kwargs):
clients.append(kwargs)
return FakeHttpClient()

@asynccontextmanager
async def fake_sse_client(url, httpx_client_factory):
client = httpx_client_factory(
headers={"X-SDK": "sdk", "X-Shared": "sdk"},
timeout="timeout",
auth="auth",
)
async with client:
yield object(), object()

class FakeSession:
def __init__(self, read, write):
pass

async def __aenter__(self):
return self

async def __aexit__(self, exc_type, exc, traceback):
return False

async def initialize(self):
pass

async def list_tools(self):
return SimpleNamespace(tools=[])

monkeypatch.setattr(httpx, "AsyncClient", fake_http_client)
monkeypatch.setattr(mcp, "ClientSession", FakeSession)
monkeypatch.setattr(mcp.client.sse, "sse_client", fake_sse_client)

cfg = SimpleNamespace(
type="sse",
url="https://example.test/sse",
headers={"X-Config": "config", "X-Shared": "config"},
tool_timeout=30,
)
async with AsyncExitStack() as stack:
await connect_mcp_servers({"svc": cfg}, ToolRegistry(), stack)

assert clients == [
{
"headers": {"X-Config": "config", "X-Shared": "sdk", "X-SDK": "sdk"},
"follow_redirects": True,
"timeout": "timeout",
"auth": "auth",
}
]

async def test_unknown_transport_is_skipped(self, monkeypatch):
"""An unknown transport does not attempt to open an MCP connection."""
from contextlib import AsyncExitStack
from types import SimpleNamespace

from raven.agent.tools import mcp as mcp_tools
from raven.agent.tools.registry import ToolRegistry

attempted = False

def fake_connection(cfg, transport_type, executor):
nonlocal attempted
attempted = True
raise AssertionError("unknown transport attempted a connection")

monkeypatch.setattr(mcp_tools, "_mcp_server_connection", fake_connection)
cfg = SimpleNamespace(type="websocket", command=None, url="wss://example.test/mcp")

await mcp_tools.connect_mcp_servers({"svc": cfg}, ToolRegistry(), AsyncExitStack())

assert attempted is False

async def test_stdio_sandboxed_with_spawning_does_not_raise(self):
"""Sandboxed executor that supports spawning does not trigger the guard."""
from contextlib import AsyncExitStack
Expand Down
Loading