feat(mcp/v2): UpstreamConnection sse + stdio transports

Adds transport selection: a _session_streams context manager opens the SDK client over the
server's transport (streamable-http, sse, or stdio) and normalizes to a (read, write) stream pair.
Auth/headers ride the httpx client per SDK 1.26 (not the transport kwargs); stdio takes
command/args/env with no auth; only the http client needs explicit cleanup. Tested live against
in-process http + sse FastMCP servers and a stdio subprocess. Cross-server aggregation/namespacing
lands with the v2 manager (step 6).
This commit is contained in:
Tin Chi Lo 2026-06-19 12:07:39 -07:00
parent 68b2142fd7
commit b7aa83140f
2 changed files with 135 additions and 18 deletions

View file

@ -31,15 +31,25 @@ from typing import (
)
import httpx
from mcp import ClientSession
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
from pydantic import BaseModel, ConfigDict
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
from litellm.proxy._types import MCPTransport
from litellm.proxy.gateway.mcp.result import Error, Ok, Result
if TYPE_CHECKING:
from collections.abc import AsyncGenerator
from anyio.streams.memory import (
MemoryObjectReceiveStream,
MemoryObjectSendStream,
)
from mcp import ReadResourceResult, Resource
from mcp.shared.message import SessionMessage
from mcp.types import CallToolResult, GetPromptResult, Prompt
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
@ -47,6 +57,11 @@ if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
_Streams = tuple[
MemoryObjectReceiveStream[SessionMessage | Exception],
MemoryObjectSendStream[SessionMessage],
]
_T = TypeVar("_T")
_V2_EGRESS_ENV_FLAG = "LITELLM_USE_V2_MCP_EGRESS"
@ -160,29 +175,71 @@ class UpstreamConnection:
The v2 egress transport: attaches ``resolve()``'s ``httpx.Auth`` (and any static/env-var
headers) to the connection through litellm's httpx client (SSL/proxy config), opens a
streamable-http ``ClientSession`` per request, and returns typed results (errors-as-values).
Replaces v1's ``MCPClient`` for the modes routed through the v2 manager. (sse/stdio transports
and the prompt/resource ops land in later steps.)
``ClientSession`` per request over the server's transport (streamable-http, sse, or stdio),
and returns typed results (errors-as-values). Replaces v1's ``MCPClient`` for the modes routed
through the v2 manager.
"""
def __init__(
self,
server_url: str,
server_url: Optional[str] = None,
*,
transport: MCPTransport = MCPTransport.http,
auth: Optional[httpx.Auth] = None,
extra_headers: Optional[Dict[str, str]] = None,
timeout: float = 60.0,
command: Optional[str] = None,
args: Optional[List[str]] = None,
env: Optional[Dict[str, str]] = None,
) -> None:
self._server_url = server_url
self._transport = transport
self._auth = auth
self._extra_headers = extra_headers
self._timeout = timeout
self._command = command
self._args = args
self._env = env
async def _run(
self, operation: Callable[[ClientSession], Awaitable[_T]]
) -> Result[_T, ConnError]:
# The resolved auth (and any static/env-var headers) ride on the httpx client; SSL/proxy
# config comes from litellm's get_ssl_configuration, matching v1's connection behavior.
def _http_client_factory(self) -> Callable[..., httpx.AsyncClient]:
def factory(
*,
headers: Optional[Dict[str, str]] = None,
timeout: Optional[httpx.Timeout] = None,
auth: Optional[httpx.Auth] = None,
) -> httpx.AsyncClient:
return httpx.AsyncClient(
headers=headers,
timeout=timeout,
auth=auth,
verify=get_ssl_configuration(None),
follow_redirects=True,
)
return factory
@contextlib.asynccontextmanager
async def _session_streams(self) -> AsyncGenerator[_Streams, None]:
# The resolved auth and static/env-var headers ride on the httpx client (SDK 1.26: not the
# transport kwargs); SSL/proxy config comes from litellm's get_ssl_configuration. Normalizes
# all three transports to a (read, write) stream pair (streamable-http yields a third value).
if self._transport == MCPTransport.stdio:
params = StdioServerParameters(
command=self._command or "", args=self._args or [], env=self._env
)
async with stdio_client(params) as (read_stream, write_stream, *_):
yield read_stream, write_stream
return
if self._transport == MCPTransport.sse:
async with sse_client(
url=self._server_url or "",
timeout=self._timeout,
headers=self._extra_headers,
auth=self._auth,
httpx_client_factory=self._http_client_factory(),
) as (read_stream, write_stream, *_):
yield read_stream, write_stream
return
http_client = httpx.AsyncClient(
headers=self._extra_headers,
timeout=httpx.Timeout(self._timeout),
@ -192,17 +249,24 @@ class UpstreamConnection:
)
try:
async with streamable_http_client(
url=self._server_url, http_client=http_client
) as (read_stream, write_stream, _):
url=self._server_url or "", http_client=http_client
) as (read_stream, write_stream, *_):
yield read_stream, write_stream
finally:
with contextlib.suppress(Exception):
await http_client.aclose()
async def _run(
self, operation: Callable[[ClientSession], Awaitable[_T]]
) -> Result[_T, ConnError]:
try:
async with self._session_streams() as (read_stream, write_stream):
async with ClientSession(read_stream, write_stream) as session:
await session.initialize()
result = await operation(session)
return Ok(result)
except Exception as e: # transport / protocol failures -> ConnError
return Error(_classify_conn_error(e))
finally:
with contextlib.suppress(Exception):
await http_client.aclose()
async def list_tools(self) -> Result[List[MCPTool], ConnError]:
async def op(session: ClientSession) -> List[MCPTool]:

View file

@ -2,6 +2,7 @@
import contextlib
import socket
import sys
import threading
import time
@ -32,13 +33,13 @@ def test_egress_flag_falsey_values(monkeypatch, value):
@contextlib.contextmanager
def _serve(app):
"""Serve an ASGI app on a free port in a background thread; yield its /mcp url."""
def _serve(app, path="/mcp"):
"""Serve an ASGI app on a free port in a background thread; yield its endpoint url."""
sock = socket.socket()
sock.bind(("127.0.0.1", 0))
port = sock.getsockname()[1]
sock.close()
url = f"http://127.0.0.1:{port}/mcp"
url = f"http://127.0.0.1:{port}{path}"
server = uvicorn.Server(
uvicorn.Config(app, host="127.0.0.1", port=port, log_level="error")
)
@ -105,6 +106,20 @@ def protected_server_url():
yield url, "secret-token"
@pytest.fixture
def sse_server_url():
from mcp.server.fastmcp import FastMCP
mcp = FastMCP("egress-sse-test")
@mcp.tool()
def echo(text: str) -> str:
return f"echo: {text}"
with _serve(mcp.sse_app(), path="/sse") as url:
yield url
@pytest.mark.asyncio
async def test_upstream_connection_lists_and_calls(echo_server_url):
from litellm.proxy._experimental.mcp_server.v2_egress import UpstreamConnection
@ -178,3 +193,41 @@ async def test_upstream_connection_prompts_and_resources(echo_server_url):
read = await conn.read_resource(target.uri)
assert isinstance(read, Ok)
assert read.ok.contents # at least one content block
@pytest.mark.asyncio
async def test_upstream_connection_stdio(tmp_path):
from litellm.proxy._experimental.mcp_server.v2_egress import UpstreamConnection
from litellm.proxy._types import MCPTransport
from litellm.proxy.gateway.mcp.result import Ok
script = tmp_path / "stdio_server.py"
script.write_text(
"from mcp.server.fastmcp import FastMCP\n"
"mcp = FastMCP('stdio-test')\n"
"@mcp.tool()\n"
"def ping() -> str:\n"
" return 'pong'\n"
"mcp.run(transport='stdio')\n"
)
conn = UpstreamConnection(
transport=MCPTransport.stdio, command=sys.executable, args=[str(script)]
)
tools = await conn.list_tools()
assert isinstance(tools, Ok)
assert any(t.name == "ping" for t in tools.ok)
@pytest.mark.asyncio
async def test_upstream_connection_sse(sse_server_url):
from litellm.proxy._experimental.mcp_server.v2_egress import UpstreamConnection
from litellm.proxy._types import MCPTransport
from litellm.proxy.gateway.mcp.outbound_credentials.httpx_auth import NoOpAuth
from litellm.proxy.gateway.mcp.result import Ok
conn = UpstreamConnection(
sse_server_url, transport=MCPTransport.sse, auth=NoOpAuth()
)
tools = await conn.list_tools()
assert isinstance(tools, Ok)
assert any(t.name == "echo" for t in tools.ok)