mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
* feat(mcp)!: disable stdio MCP servers by default stdio MCP servers now only run when the proxy is started with LITELLM_ENABLE_MCP_STDIO=true. While it is off, existing stdio servers stay registered but never start: tool listings skip them quietly, direct tool calls and health checks return a 403 naming the env var, and creating or updating a stdio server is rejected. The flag is read from the process environment only, so DB-stored environment_variables cannot turn it on. The UI reads mcp_stdio_enabled from /.well-known/litellm-ui-config to grey out the stdio transport, show a banner on stdio forms, and badge stdio server cards. BREAKING CHANGE: stdio MCP servers are off by default. Set LITELLM_ENABLE_MCP_STDIO=true in the proxy environment and restart to keep using them. * fix(mcp): ignore stdio flag from config file and read UI flag from the selected worker LITELLM_ENABLE_MCP_STDIO set under environment_variables in config.yaml is now skipped like the DB-stored value, so only the process environment can enable stdio. The dashboard reads mcp_stdio_enabled from the proxy it is managing, so a control plane shows each worker's own setting. * test(mcp): cover non-mapping payloads in the shared transport validator * fix(mcp): skip blocked stdio servers quietly in every listing and keep the UI unchanged until the flag loads Prompt, resource and resource-template listings now skip a blocked stdio server at debug level like tool listing does, instead of logging a warning per server on every call. The dashboard only treats stdio as disabled once the proxy explicitly reports mcp_stdio_enabled false, so a proxy with the flag on, or an older one without the field, renders exactly as before with no flicker while loading. * fix(mcp): route blocked stdio tool calls to the flag error and warn once per server A gateway tools/call naming a blocked stdio server's tool now returns the LITELLM_ENABLE_MCP_STDIO message instead of "Tool not found". The "will not start" warning moves out of build_mcp_server_from_table, which DB reload re-runs on every cycle for rows with a NULL updated_at and which drafts and test-connection also call. It now fires when a row first enters the registry or changes transport. * fix(ui): explain on the server detail page why a stdio server is inert The Overview and MCP Tools tabs showed "No tools available" with no reason while stdio is disabled. The detail page now shows the same warning banner as the edit form, and hands off to the form's banner once editing starts. * refactor(ui): name the stdio banner conditions on the server detail page Keeps local/no-long-condition-chain within its budget * fix(proxy): log the ignored DB-stored LITELLM_ENABLE_MCP_STDIO warning once The DB config sync re-reads environment_variables on every cycle, so a stored flag logged the warning on each sync per worker
796 lines
37 KiB
Python
796 lines
37 KiB
Python
import asyncio
|
|
import json
|
|
import os
|
|
import queue
|
|
import socket
|
|
import subprocess
|
|
import sys
|
|
import tempfile
|
|
import threading
|
|
import time
|
|
import typing
|
|
from contextlib import asynccontextmanager, contextmanager
|
|
from dataclasses import dataclass
|
|
from datetime import datetime
|
|
from pathlib import Path
|
|
|
|
import httpx
|
|
import httpx2
|
|
import pytest
|
|
import uvicorn
|
|
import yaml
|
|
from mcp import ClientSession
|
|
from mcp.client.streamable_http import streamable_http_client
|
|
from mcp.types import CallToolResult
|
|
from starlette.requests import Request
|
|
|
|
from litellm.integrations.custom_logger import CustomLogger
|
|
from litellm.proxy._experimental.mcp_server.tool_search import handle_mcp_proxy_tool
|
|
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, ProxyException, UserAPIKeyAuth
|
|
from litellm.proxy.proxy_server import (
|
|
app as proxy_app,
|
|
)
|
|
from litellm.proxy.proxy_server import (
|
|
cleanup_router_config_variables,
|
|
initialize,
|
|
)
|
|
|
|
CONFIG_TEMPLATE_PATH = Path("tests/mcp_tests/test_configs/test_config_mcp_e2e.yaml")
|
|
MCP_SERVER_SCRIPT = Path("tests/mcp_tests/mcp_server.py")
|
|
MCP_PEER_PYTHON = os.environ.get("MCP_TEST_PEER_PYTHON", sys.executable)
|
|
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
|
PROXY_START_TIMEOUT = 30
|
|
|
|
|
|
PROXY_AUTHORIZATION_HEADER = "Bearer sk-1234"
|
|
|
|
|
|
@pytest.fixture(scope="session", autouse=True)
|
|
def _clear_proxy_database_env() -> typing.Iterator[None]:
|
|
"""Ensure local proxy DB settings don't leak into tests."""
|
|
mp = pytest.MonkeyPatch()
|
|
mp.delenv("DATABASE_URL", raising=False)
|
|
# The FastAPI lifespan event (proxy_startup_event) re-reads master_key from
|
|
# the LITELLM_MASTER_KEY env var, overriding whatever initialize() set from
|
|
# the config file. We must set it here so the lifespan doesn't reset it to None.
|
|
mp.setenv("LITELLM_MASTER_KEY", "sk-1234")
|
|
mp.setenv("LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY", "true")
|
|
mp.setenv("LITELLM_ENABLE_MCP_STDIO", "true")
|
|
try:
|
|
yield
|
|
finally:
|
|
mp.undo()
|
|
|
|
|
|
async def _initialize_proxy(config_path: str) -> None:
|
|
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
|
|
|
cleanup_router_config_variables()
|
|
await initialize(config=config_path, debug=True)
|
|
for server_id, upstream in tuple(global_mcp_server_manager.registry.items()):
|
|
if upstream.server_name != "math_restricted":
|
|
continue
|
|
global_mcp_server_manager.registry[server_id] = upstream.model_copy(
|
|
update={"tool_name_to_display_name": {"add": "Add Numbers"}}
|
|
)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ProxyRig:
|
|
url: str
|
|
config_path: str
|
|
loop: asyncio.AbstractEventLoop
|
|
|
|
|
|
def _start_proxy_server(
|
|
config_path: str,
|
|
) -> tuple[ProxyRig, uvicorn.Server, threading.Thread, socket.socket]:
|
|
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
|
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
|
sock.bind(("127.0.0.1", 0))
|
|
host, port = sock.getsockname()
|
|
|
|
config = uvicorn.Config(proxy_app, host=host, port=port, log_level="warning", lifespan="off")
|
|
server = uvicorn.Server(config)
|
|
|
|
loop = asyncio.new_event_loop()
|
|
|
|
async def _serve() -> None:
|
|
from litellm.proxy._experimental.mcp_server import server as mcp_server
|
|
|
|
await _initialize_proxy(config_path)
|
|
async with proxy_app.router.lifespan_context(proxy_app), mcp_server.lifespan(proxy_app):
|
|
await server.serve(sockets=[sock])
|
|
|
|
def _run() -> None:
|
|
with asyncio.Runner(loop_factory=lambda: loop) as runner:
|
|
runner.run(_serve())
|
|
|
|
thread = threading.Thread(target=_run, daemon=True)
|
|
thread.start()
|
|
|
|
start_time = time.time()
|
|
while not server.started:
|
|
if not thread.is_alive():
|
|
raise RuntimeError("Proxy server failed to start")
|
|
if time.time() - start_time > PROXY_START_TIMEOUT:
|
|
raise TimeoutError("Proxy server did not start in time")
|
|
time.sleep(0.05)
|
|
|
|
return ProxyRig(f"http://{host}:{port}", config_path, loop), server, thread, sock
|
|
|
|
|
|
@contextmanager
|
|
def _math_http_server(offset: int) -> typing.Iterator[str]:
|
|
host = "127.0.0.1"
|
|
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
|
|
sock.bind((host, 0))
|
|
_, port = sock.getsockname()
|
|
|
|
with tempfile.TemporaryFile() as server_log:
|
|
process = subprocess.Popen(
|
|
[MCP_PEER_PYTHON, str(MCP_SERVER_SCRIPT), "--transport", "http", "--host", host, "--port", str(port)],
|
|
cwd=str(PROJECT_ROOT),
|
|
stdout=server_log,
|
|
stderr=subprocess.STDOUT,
|
|
env={**os.environ, "MCP_ADD_OFFSET": str(offset)},
|
|
)
|
|
try:
|
|
start_time = time.monotonic()
|
|
while True:
|
|
if process.poll() is not None:
|
|
server_log.seek(0)
|
|
raise RuntimeError(f"MCP upstream exited early: {server_log.read().decode()}")
|
|
try:
|
|
with socket.create_connection((host, port), timeout=0.1):
|
|
break
|
|
except OSError:
|
|
if time.monotonic() - start_time > PROXY_START_TIMEOUT:
|
|
raise TimeoutError("Streamable HTTP MCP server did not start in time")
|
|
time.sleep(0.05)
|
|
yield f"http://{host}:{port}"
|
|
finally:
|
|
process.terminate()
|
|
try:
|
|
process.wait(timeout=5)
|
|
except subprocess.TimeoutExpired:
|
|
process.kill()
|
|
process.wait(timeout=5)
|
|
|
|
|
|
@pytest.fixture(scope="session")
|
|
def math_streamable_http_server() -> typing.Iterator[str]:
|
|
with _math_http_server(100) as url:
|
|
yield url
|
|
|
|
|
|
@pytest.fixture(scope="session")
|
|
def math_restricted_server() -> typing.Iterator[str]:
|
|
with _math_http_server(200) as url:
|
|
yield url
|
|
|
|
|
|
@pytest.fixture(scope="session")
|
|
def _proxy_server(
|
|
tmp_path_factory: pytest.TempPathFactory,
|
|
math_streamable_http_server: str,
|
|
math_restricted_server: str,
|
|
):
|
|
config_dir = tmp_path_factory.mktemp("mcp_e2e")
|
|
config_path = config_dir / "config.yaml"
|
|
config = yaml.safe_load(CONFIG_TEMPLATE_PATH.read_text())
|
|
config["mcp_servers"]["math_stdio"]["command"] = MCP_PEER_PYTHON
|
|
config["mcp_servers"]["math_streamable_http"]["url"] = f"{math_streamable_http_server}/mcp"
|
|
config["mcp_servers"]["math_restricted"]["url"] = f"{math_restricted_server}/mcp"
|
|
config["general_settings"]["custom_auth"] = f"{__name__}.authorize_proxy_key"
|
|
config["litellm_settings"]["callbacks"] = [f"{__name__}.proxy_call_recorder"]
|
|
config["mcp_servers"]["math_restricted"]["mcp_info"] = {"mcp_server_cost_info": {"default_cost_per_query": 0.25}}
|
|
config_path.write_text(yaml.safe_dump(config))
|
|
|
|
rig, server, thread, sock = _start_proxy_server(str(config_path))
|
|
|
|
try:
|
|
yield rig
|
|
finally:
|
|
server.should_exit = True
|
|
thread.join(timeout=10)
|
|
sock.close()
|
|
assert not thread.is_alive(), "Proxy did not shut down"
|
|
|
|
|
|
@pytest.fixture
|
|
def proxy_server_url(_proxy_server: ProxyRig, setup_and_teardown: None) -> str:
|
|
asyncio.run_coroutine_threadsafe(_initialize_proxy(_proxy_server.config_path), _proxy_server.loop).result(
|
|
timeout=30
|
|
)
|
|
return _proxy_server.url
|
|
|
|
|
|
@asynccontextmanager
|
|
async def _http_streams(url: str, headers: dict[str, str]):
|
|
async with httpx2.AsyncClient(headers=headers) as http_client:
|
|
async with streamable_http_client(url, http_client=http_client) as streams:
|
|
yield streams
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_unchanged_sdk1_langchain_peer_can_list_and_call(proxy_server_url: str) -> None:
|
|
script = """
|
|
import asyncio, json, sys
|
|
from mcp import ClientSession
|
|
from mcp.client.streamable_http import streamablehttp_client
|
|
from langchain_mcp_adapters.tools import load_mcp_tools
|
|
|
|
async def main():
|
|
async with streamablehttp_client(sys.argv[1] + '/mcp', headers={'Authorization': 'Bearer sk-1234'}) as (read, write, _):
|
|
async with ClientSession(read, write) as session:
|
|
await session.initialize()
|
|
tools = await load_mcp_tools(session)
|
|
results = {}
|
|
for name in ('math_stdio-add', 'math_streamable_http-add'):
|
|
tool = next(tool for tool in tools if tool.name == name)
|
|
results[name] = await tool.ainvoke({'a': 3, 'b': 4})
|
|
print(json.dumps(results))
|
|
asyncio.run(main())
|
|
"""
|
|
completed = await asyncio.to_thread(
|
|
subprocess.run, [MCP_PEER_PYTHON, "-c", script, proxy_server_url],
|
|
capture_output=True, text=True, timeout=30, check=True,
|
|
)
|
|
results = json.loads(completed.stdout)
|
|
assert [(item["type"], item["text"]) for item in results["math_stdio-add"]] == [("text", "7")]
|
|
assert [(item["type"], item["text"]) for item in results["math_streamable_http-add"]] == [("text", "107")]
|
|
|
|
|
|
@pytest.mark.parametrize("requested", ["2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25", "2026-07-28"])
|
|
def test_initialize_keeps_legacy_negotiation(proxy_server_url: str, requested: str) -> None:
|
|
response = httpx.post(
|
|
proxy_server_url + "/mcp",
|
|
headers={"Authorization": PROXY_AUTHORIZATION_HEADER, "Accept": "application/json, text/event-stream"},
|
|
json={"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": {
|
|
"protocolVersion": requested, "capabilities": {}, "clientInfo": {"name": "legacy-test", "version": "1"},
|
|
}},
|
|
timeout=10,
|
|
)
|
|
assert response.status_code == 200
|
|
result = _rpc_result(response)
|
|
assert result["protocolVersion"] == ("2025-11-25" if requested == "2026-07-28" else requested)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_legacy_prompts_and_resources_round_trip(proxy_server_url: str) -> None:
|
|
async with _http_streams(
|
|
proxy_server_url + "/mcp",
|
|
{"Authorization": PROXY_AUTHORIZATION_HEADER, "x-mcp-servers": "math_streamable_http"},
|
|
) as (read, write):
|
|
async with ClientSession(read, write) as session:
|
|
await session.initialize()
|
|
prompts = await session.list_prompts()
|
|
greeting = next(prompt for prompt in prompts.prompts if prompt.name.endswith("greeting"))
|
|
prompt = await session.get_prompt(greeting.name, {"name": "Ada"})
|
|
assert prompt.messages[0].content.text == "Hello, Ada"
|
|
resources = await session.list_resources()
|
|
status = next(resource for resource in resources.resources if resource.name.endswith("status"))
|
|
contents = await session.read_resource(status.uri)
|
|
assert contents.contents[0].text == "ready"
|
|
templates = await session.list_resource_templates()
|
|
greeting_template = next(template for template in templates.resource_templates if "greeting" in template.name)
|
|
contents = await session.read_resource(greeting_template.uri_template.replace("{name}", "Ada"))
|
|
assert contents.contents[0].text == "Hello, Ada"
|
|
|
|
|
|
class TestProxyMcpSimpleConnections:
|
|
@pytest.mark.asyncio
|
|
async def test_proxy_mcp_stdio_roundtrip(self, proxy_server_url: str) -> None:
|
|
async with asyncio.timeout(20):
|
|
async with _http_streams(
|
|
url=f"{proxy_server_url}/mcp",
|
|
headers={
|
|
"Authorization": PROXY_AUTHORIZATION_HEADER,
|
|
"x-mcp-servers": "math_stdio",
|
|
},
|
|
) as (read, write):
|
|
async with ClientSession(read, write) as session:
|
|
await session.initialize()
|
|
tools_result = await session.list_tools()
|
|
assert any(tool.name.endswith("add") for tool in tools_result.tools)
|
|
|
|
result = await session.call_tool("add", arguments={"a": 3, "b": 4})
|
|
assert result.content
|
|
first_content = result.content[0]
|
|
text = getattr(first_content, "text", None)
|
|
assert text == "7"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_proxy_mcp_streamable_http_roundtrip(self, proxy_server_url: str) -> None:
|
|
async with asyncio.timeout(20):
|
|
async with _http_streams(
|
|
url=f"{proxy_server_url}/mcp",
|
|
headers={
|
|
"Authorization": PROXY_AUTHORIZATION_HEADER,
|
|
"x-mcp-servers": "math_streamable_http",
|
|
},
|
|
) as (read, write):
|
|
async with ClientSession(read, write) as session:
|
|
await session.initialize()
|
|
tools_result = await session.list_tools()
|
|
assert any(tool.name.endswith("add") for tool in tools_result.tools)
|
|
|
|
result = await session.call_tool("add", arguments={"a": 5, "b": 6})
|
|
assert result.content
|
|
first_content = result.content[0]
|
|
text = getattr(first_content, "text", None)
|
|
assert text == "111"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_proxy_mcp_lists_all_servers_without_header(self, proxy_server_url: str) -> None:
|
|
async with asyncio.timeout(20):
|
|
async with _http_streams(
|
|
url=f"{proxy_server_url}/mcp",
|
|
headers={"Authorization": PROXY_AUTHORIZATION_HEADER},
|
|
) as (read, write):
|
|
async with ClientSession(read, write) as session:
|
|
await session.initialize()
|
|
tools_result = await session.list_tools()
|
|
tool_names = {tool.name for tool in tools_result.tools}
|
|
expected_tool_names = {
|
|
"math_stdio-add",
|
|
"math_stdio-multiply",
|
|
"math_streamable_http-add",
|
|
"math_streamable_http-multiply",
|
|
}
|
|
assert expected_tool_names <= tool_names
|
|
|
|
async def _call_and_get_text(tool_name: str, *, a: int, b: int) -> str | None:
|
|
result = await session.call_tool(tool_name, arguments={"a": a, "b": b})
|
|
assert result.content
|
|
first_content = result.content[0]
|
|
return getattr(first_content, "text", None)
|
|
|
|
stdio_result = await _call_and_get_text("math_stdio-add", a=2, b=3)
|
|
streamable_result = await _call_and_get_text("math_streamable_http-add", a=4, b=5)
|
|
assert stdio_result == "5"
|
|
assert streamable_result == "109"
|
|
|
|
|
|
class TestProxyMcpStatelessBehavior:
|
|
"""
|
|
Verify that the LiteLLM MCP proxy operates in stateless mode.
|
|
|
|
When StreamableHTTPSessionManager is configured with stateless=True,
|
|
independent clients must be able to connect, list tools, and call tools
|
|
without sharing or inheriting session state from other clients.
|
|
|
|
With stateless=False this fails because the server tracks sessions and
|
|
expects clients to supply an mcp-session-id header obtained from a
|
|
prior handshake — breaking clients that don't manage session IDs.
|
|
|
|
Regression test for https://github.com/BerriAI/litellm/issues/20242
|
|
"""
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_independent_clients_no_shared_session(self, proxy_server_url: str) -> None:
|
|
"""Two independent clients connect and operate without sharing session state."""
|
|
async with asyncio.timeout(30):
|
|
# --- Client A: connect, initialize, call tool ---
|
|
async with _http_streams(
|
|
url=f"{proxy_server_url}/mcp",
|
|
headers={
|
|
"Authorization": PROXY_AUTHORIZATION_HEADER,
|
|
"x-mcp-servers": "math_stdio",
|
|
},
|
|
) as (read_a, write_a):
|
|
async with ClientSession(read_a, write_a) as session_a:
|
|
await session_a.initialize()
|
|
result_a = await session_a.call_tool("math_stdio-add", arguments={"a": 10, "b": 20})
|
|
assert result_a.content
|
|
text_a = getattr(result_a.content[0], "text", None)
|
|
assert text_a == "30"
|
|
|
|
# Allow proxy and MCP SDK to fully clean up the first connection before
|
|
# opening the second. Without this, the SDK's TaskGroup can raise
|
|
# ExceptionGroup when the server closes the connection (see MCP SDK #915).
|
|
await asyncio.sleep(0.5)
|
|
|
|
# --- Client B: completely independent connection ---
|
|
async with _http_streams(
|
|
url=f"{proxy_server_url}/mcp",
|
|
headers={
|
|
"Authorization": PROXY_AUTHORIZATION_HEADER,
|
|
"x-mcp-servers": "math_stdio",
|
|
},
|
|
) as (read_b, write_b):
|
|
async with ClientSession(read_b, write_b) as session_b:
|
|
await session_b.initialize()
|
|
tools = await session_b.list_tools()
|
|
assert any(t.name.endswith("add") for t in tools.tools)
|
|
result_b = await session_b.call_tool("math_stdio-add", arguments={"a": 100, "b": 200})
|
|
assert result_b.content
|
|
text_b = getattr(result_b.content[0], "text", None)
|
|
assert text_b == "300"
|
|
|
|
|
|
PROXY_MODE_TOOLS = frozenset({"search_tools", "get_tool_schema", "call_tool"})
|
|
|
|
|
|
def _payload(result: typing.Any) -> typing.Any:
|
|
assert result.content, f"empty tool result: {result}"
|
|
return json.loads(result.content[0].text)
|
|
|
|
|
|
def _proxy_session(proxy_server_url: str, **extra_headers: str):
|
|
return _http_streams(
|
|
url=f"{proxy_server_url}/mcp/proxy",
|
|
headers={"Authorization": PROXY_AUTHORIZATION_HEADER, **extra_headers},
|
|
)
|
|
|
|
|
|
class TestProxyMcpSchemaDiscoveryMode:
|
|
"""Drive /mcp/proxy over the real streamable-HTTP transport with the MCP SDK client:
|
|
the fixed three-tool surface, opaque-id discovery, schema-validated execution against
|
|
two upstreams that expose the same tool name, and the operations the surface refuses."""
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_initialize_and_list_expose_only_discovery_tools(self, proxy_server_url: str) -> None:
|
|
async with asyncio.timeout(20):
|
|
async with _proxy_session(proxy_server_url) as (read, write):
|
|
async with ClientSession(read, write) as session:
|
|
init = await session.initialize()
|
|
assert init.capabilities.tools is not None
|
|
assert init.capabilities.prompts is None
|
|
assert init.capabilities.resources is None
|
|
|
|
listed = await session.list_tools()
|
|
assert {tool.name for tool in listed.tools} == PROXY_MODE_TOOLS
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_search_schema_and_call_round_trip_keeps_server_identity(self, proxy_server_url: str) -> None:
|
|
async with asyncio.timeout(30):
|
|
async with _proxy_session(proxy_server_url) as (read, write):
|
|
async with ClientSession(read, write) as session:
|
|
await session.initialize()
|
|
|
|
hits = _payload(await session.call_tool("search_tools", arguments={"query": "add"}))
|
|
by_name = {hit["name"]: hit for hit in hits}
|
|
assert {"math_stdio-add", "math_streamable_http-add"} <= set(by_name)
|
|
assert all("inputSchema" not in hit for hit in hits)
|
|
assert by_name["math_stdio-add"]["tool_id"] != by_name["math_streamable_http-add"]["tool_id"]
|
|
|
|
schema = _payload(
|
|
await session.call_tool(
|
|
"get_tool_schema", arguments={"tool_id": by_name["math_stdio-add"]["tool_id"]}
|
|
)
|
|
)
|
|
assert schema["name"] == "math_stdio-add"
|
|
assert set(schema["inputSchema"]["required"]) == {"a", "b"}
|
|
assert schema["outputSchema"]["properties"]["result"]["type"] == "integer"
|
|
|
|
stdio = await session.call_tool(
|
|
"call_tool",
|
|
arguments={"tool_id": by_name["math_stdio-add"]["tool_id"], "arguments": {"a": 3, "b": 4}},
|
|
)
|
|
http = await session.call_tool(
|
|
"call_tool",
|
|
arguments={
|
|
"tool_id": by_name["math_streamable_http-add"]["tool_id"],
|
|
"arguments": {"a": 5, "b": 6},
|
|
},
|
|
)
|
|
assert stdio.is_error is False and stdio.content[0].text == "7"
|
|
assert http.is_error is False and http.content[0].text == "111"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_server_scope_header_narrows_discovery(self, proxy_server_url: str) -> None:
|
|
async with asyncio.timeout(20):
|
|
async with _proxy_session(proxy_server_url, **{"x-mcp-servers": "math_streamable_http"}) as (
|
|
read,
|
|
write,
|
|
):
|
|
async with ClientSession(read, write) as session:
|
|
await session.initialize()
|
|
hits = _payload(await session.call_tool("search_tools", arguments={"query": "add"}))
|
|
assert {hit["name"] for hit in hits} == {"math_streamable_http-add"}
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_rejections_never_reach_upstream(self, proxy_server_url: str) -> None:
|
|
from mcp.shared.exceptions import MCPError
|
|
from mcp.types import METHOD_NOT_FOUND
|
|
|
|
async with asyncio.timeout(30):
|
|
async with _proxy_session(proxy_server_url) as (read, write):
|
|
async with ClientSession(read, write) as session:
|
|
await session.initialize()
|
|
hits = _payload(await session.call_tool("search_tools", arguments={"query": "add"}))
|
|
tool_id = next(hit["tool_id"] for hit in hits if hit["name"] == "math_stdio-add")
|
|
|
|
bad_args = await session.call_tool(
|
|
"call_tool", arguments={"tool_id": tool_id, "arguments": {"a": "three", "b": 4}}
|
|
)
|
|
assert bad_args.is_error is True and "Invalid arguments" in bad_args.content[0].text
|
|
|
|
stale = await session.call_tool("get_tool_schema", arguments={"tool_id": "0" * 32})
|
|
assert stale.is_error is True and "unauthorized tool_id" in stale.content[0].text
|
|
|
|
for not_an_object in ("wrong", False):
|
|
refused_args = await session.call_tool(
|
|
"call_tool", arguments={"tool_id": tool_id, "arguments": not_an_object}
|
|
)
|
|
assert refused_args.is_error is True and "object" in refused_args.content[0].text
|
|
|
|
direct = await session.call_tool("math_stdio-add", arguments={"a": 1, "b": 2})
|
|
assert direct.is_error is True and "unavailable on /mcp/proxy" in direct.content[0].text
|
|
|
|
for operation in (session.list_prompts, session.list_resources):
|
|
with pytest.raises(MCPError) as refused:
|
|
await operation()
|
|
assert refused.value.error.code == METHOD_NOT_FOUND
|
|
|
|
|
|
async def authorize_proxy_key(request: Request, api_key: str) -> UserAPIKeyAuth:
|
|
permissions = {
|
|
"sk-1234": LiteLLM_ObjectPermissionTable(object_permission_id="open", mcp_servers=["math_stdio"]),
|
|
"sk-restricted": LiteLLM_ObjectPermissionTable(
|
|
object_permission_id="restricted", mcp_servers=["math_restricted"]
|
|
),
|
|
"sk-none": LiteLLM_ObjectPermissionTable(object_permission_id="none", mcp_servers=["no-mcp-servers"]),
|
|
"sk-add-only": LiteLLM_ObjectPermissionTable(
|
|
object_permission_id="add-only", mcp_servers=["math_stdio"], mcp_tool_permissions={"math_stdio": ["add"]}
|
|
),
|
|
}
|
|
permission = permissions.get(api_key)
|
|
if permission is None:
|
|
raise ProxyException(message="Unknown test key", type="authentication_error", param=None, code=401)
|
|
return UserAPIKeyAuth(api_key=api_key, user_id=api_key, object_permission=permission)
|
|
|
|
|
|
class ProxyCallRecorder(CustomLogger):
|
|
def __init__(self) -> None:
|
|
super().__init__()
|
|
self.events: queue.Queue[str] = queue.Queue()
|
|
self.failures: queue.Queue[str] = queue.Queue()
|
|
|
|
async def async_log_success_event(
|
|
self, kwargs: dict[str, object], response_obj: object, start_time: datetime, end_time: datetime
|
|
) -> None:
|
|
payload = kwargs.get("standard_logging_object")
|
|
if isinstance(payload, dict) and payload.get("call_type") == "call_mcp_tool":
|
|
self.events.put(json.dumps(payload, default=str))
|
|
|
|
async def async_log_failure_event(
|
|
self, kwargs: dict[str, object], response_obj: object, start_time: datetime, end_time: datetime
|
|
) -> None:
|
|
payload = kwargs.get("standard_logging_object")
|
|
if isinstance(payload, dict) and payload.get("call_type") == "call_mcp_tool":
|
|
self.failures.put(json.dumps(payload, default=str))
|
|
|
|
|
|
proxy_call_recorder = ProxyCallRecorder()
|
|
|
|
|
|
@asynccontextmanager
|
|
async def _scoped_session(url: str, key: str = "sk-1234", **headers: str) -> typing.AsyncIterator[ClientSession]:
|
|
async with asyncio.timeout(30):
|
|
async with _proxy_session(url, Authorization=f"Bearer {key}", **headers) as (read, write):
|
|
async with ClientSession(read, write) as session:
|
|
await session.initialize()
|
|
yield session
|
|
|
|
|
|
async def _search(session: ClientSession, query: str) -> dict[str, str]:
|
|
result = await session.call_tool("search_tools", arguments={"query": query})
|
|
assert result.is_error is False, result
|
|
return {hit["name"]: hit["tool_id"] for hit in _payload(result)}
|
|
|
|
|
|
async def _call(session: ClientSession, tool_id: str, a: int = 3, b: int = 4) -> CallToolResult:
|
|
return await session.call_tool("call_tool", arguments={"tool_id": tool_id, "arguments": {"a": a, "b": b}})
|
|
|
|
|
|
async def _raw_rpc(
|
|
proxy_server_url: str, key: str | None, method: str, params: dict[str, object], **headers: str
|
|
) -> httpx.Response:
|
|
async with httpx.AsyncClient() as client:
|
|
return await client.post(
|
|
f"{proxy_server_url}/mcp/proxy",
|
|
headers={
|
|
"Accept": "application/json, text/event-stream",
|
|
**({"Authorization": f"Bearer {key}"} if key else {}),
|
|
**headers,
|
|
},
|
|
json={"jsonrpc": "2.0", "id": 1, "method": method, "params": params},
|
|
)
|
|
|
|
|
|
async def _raw_initialize(proxy_server_url: str, key: str | None) -> httpx.Response:
|
|
return await _raw_rpc(
|
|
proxy_server_url,
|
|
key,
|
|
"initialize",
|
|
{"protocolVersion": "2025-03-26", "capabilities": {}, "clientInfo": {"name": "auth-test", "version": "1"}},
|
|
)
|
|
|
|
|
|
def _rpc_result(response: httpx.Response) -> dict[str, typing.Any]:
|
|
if response.headers["content-type"].startswith("text/event-stream"):
|
|
data_line = next(line for line in response.text.splitlines() if line.startswith("data:"))
|
|
return json.loads(data_line.removeprefix("data:"))["result"]
|
|
return response.json()["result"]
|
|
|
|
|
|
def _assert_unauthorized(result: CallToolResult) -> None:
|
|
assert result.is_error is True
|
|
assert result.content[0].text == "Unknown or unauthorized tool_id"
|
|
|
|
|
|
class TestProxyMcpAuthorizationScope:
|
|
@pytest.mark.asyncio
|
|
async def test_server_grant_bounds_search_and_blocks_foreign_ids(self, proxy_server_url: str) -> None:
|
|
async with _scoped_session(proxy_server_url, "sk-restricted") as granted:
|
|
restricted_id = (await _search(granted, "add"))["math_restricted-add"]
|
|
assert (await _call(granted, restricted_id)).content[0].text == "207"
|
|
async with _scoped_session(proxy_server_url) as ungranted:
|
|
assert set(await _search(ungranted, "add")) == {"math_stdio-add", "math_streamable_http-add"}
|
|
_assert_unauthorized(await ungranted.call_tool("get_tool_schema", {"tool_id": restricted_id}))
|
|
_assert_unauthorized(await _call(ungranted, restricted_id))
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_no_mcp_servers_sentinel_rejects_initialize_and_hides_every_tool(self, proxy_server_url: str) -> None:
|
|
async with _scoped_session(proxy_server_url) as granted:
|
|
tool_id = (await _search(granted, "add"))["math_stdio-add"]
|
|
response = await _raw_initialize(proxy_server_url, "sk-none")
|
|
assert response.status_code == 403, response.text
|
|
assert "no MCP servers granted" in response.json()["detail"]["error"]
|
|
|
|
async def raw_call(name: str, arguments: dict[str, object]) -> dict[str, typing.Any]:
|
|
call = await _raw_rpc(proxy_server_url, "sk-none", "tools/call", {"name": name, "arguments": arguments})
|
|
assert call.status_code == 200, call.text
|
|
return _rpc_result(call)
|
|
|
|
listed = await _raw_rpc(proxy_server_url, "sk-none", "tools/list", {})
|
|
assert listed.status_code == 200, listed.text
|
|
assert {tool["name"] for tool in _rpc_result(listed)["tools"]} == {"search_tools", "get_tool_schema", "call_tool"}
|
|
search = await raw_call("search_tools", {"query": "add"})
|
|
assert search["isError"] is False, search
|
|
assert json.loads(search["content"][0]["text"]) == []
|
|
for name, arguments in (
|
|
("get_tool_schema", {"tool_id": tool_id}),
|
|
("call_tool", {"tool_id": tool_id, "arguments": {"a": 3, "b": 4}}),
|
|
):
|
|
denied = await raw_call(name, arguments)
|
|
assert denied["isError"] is True, denied
|
|
assert denied["content"][0]["text"] == "Unknown or unauthorized tool_id"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_tool_grant_hides_ungranted_tools_and_blocks_their_ids(self, proxy_server_url: str) -> None:
|
|
async with _scoped_session(proxy_server_url) as granted:
|
|
multiply_id = (await _search(granted, "multiply"))["math_stdio-multiply"]
|
|
async with _scoped_session(proxy_server_url, "sk-add-only", **{"x-mcp-servers": "math_stdio"}) as session:
|
|
ids = await _search(session, "add multiply request_headers")
|
|
assert set(ids) == {"math_stdio-add"}
|
|
assert (await _call(session, ids["math_stdio-add"])).content[0].text == "7"
|
|
_assert_unauthorized(await session.call_tool("get_tool_schema", {"tool_id": multiply_id}))
|
|
_assert_unauthorized(await _call(session, multiply_id))
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_same_named_tools_keep_distinct_ids_and_reach_their_own_upstream(self, proxy_server_url: str) -> None:
|
|
async with _scoped_session(proxy_server_url, "sk-restricted") as session:
|
|
ids = await _search(session, "add")
|
|
assert set(ids) == {"math_stdio-add", "math_streamable_http-add", "math_restricted-add"}
|
|
assert len(set(ids.values())) == 3
|
|
assert all(len(tool_id) == 32 for tool_id in ids.values())
|
|
for name, expected in (
|
|
("math_stdio-add", "7"),
|
|
("math_streamable_http-add", "107"),
|
|
("math_restricted-add", "207"),
|
|
):
|
|
schema = _payload(await session.call_tool("get_tool_schema", {"tool_id": ids[name]}))
|
|
assert schema["name"] == name
|
|
assert schema["tool_id"] == ids[name]
|
|
result = await _call(session, ids[name])
|
|
assert result.is_error is False
|
|
assert result.content[0].text == expected
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_server_scope_header_narrows_grants_and_blocks_out_of_scope_ids(self, proxy_server_url: str) -> None:
|
|
async with _scoped_session(proxy_server_url, "sk-restricted") as unscoped:
|
|
other_id = (await _search(unscoped, "add"))["math_stdio-add"]
|
|
async with _scoped_session(
|
|
proxy_server_url, "sk-restricted", **{"x-mcp-servers": "math_restricted"}
|
|
) as session:
|
|
ids = await _search(session, "add")
|
|
assert set(ids) == {"math_restricted-add"}
|
|
_assert_unauthorized(await session.call_tool("get_tool_schema", {"tool_id": other_id}))
|
|
_assert_unauthorized(await _call(session, other_id))
|
|
assert (await _call(session, ids["math_restricted-add"])).content[0].text == "207"
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("key", [None, "sk-invalid"])
|
|
async def test_missing_or_invalid_key_cannot_initialize(self, proxy_server_url: str, key: str | None) -> None:
|
|
response = await _raw_initialize(proxy_server_url, key)
|
|
assert response.status_code == 401, response.text
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_server_headers_are_forwarded_only_to_the_named_upstream(self, proxy_server_url: str) -> None:
|
|
for tag in ("first-request", "second-request"):
|
|
async with _scoped_session(
|
|
proxy_server_url,
|
|
"sk-restricted",
|
|
**{
|
|
"x-mcp-math_restricted-authorization": f"Bearer {tag}",
|
|
"x-mcp-math_restricted-x-request-tag": tag,
|
|
},
|
|
) as session:
|
|
ids = await _search(session, "request_headers")
|
|
for name, expected in (
|
|
("math_restricted", {"authorization": f"Bearer {tag}", "x-request-tag": tag}),
|
|
("math_streamable_http", {"authorization": "", "x-request-tag": ""}),
|
|
):
|
|
result = await session.call_tool(
|
|
"call_tool", {"tool_id": ids[f"{name}-request_headers"], "arguments": {}}
|
|
)
|
|
assert result.is_error is False
|
|
assert _payload(result) == expected
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_proxy_call_emits_spend_log(self, proxy_server_url: str) -> None:
|
|
async with _scoped_session(proxy_server_url, "sk-restricted") as session:
|
|
tool_id = (await _search(session, "add"))["math_restricted-add"]
|
|
result = await _call(session, tool_id, 123, 456)
|
|
assert result.is_error is False and result.content[0].text == "779"
|
|
async with asyncio.timeout(10):
|
|
while True:
|
|
payload = json.loads(await asyncio.to_thread(proxy_call_recorder.events.get, True, 5))
|
|
if payload.get("metadata", {}).get("mcp_tool_call_metadata", {}).get("arguments") == {
|
|
"a": 123,
|
|
"b": 456,
|
|
}:
|
|
break
|
|
assert payload["call_type"] == "call_mcp_tool"
|
|
assert payload["response_cost"] == 0.25
|
|
assert payload["status"] == "success"
|
|
assert payload["metadata"]["mcp_tool_call_metadata"]["mcp_server_name"] == "math_restricted"
|
|
assert payload["metadata"]["mcp_tool_call_metadata"]["name"] == "add"
|
|
assert payload["metadata"]["mcp_tool_call_metadata"]["namespaced_tool_name"] == "math_restricted/add"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_proxy_scope_exception_returns_iserror_and_emits_failure_log(self, proxy_server_url: str) -> None:
|
|
response = await _raw_rpc(
|
|
proxy_server_url,
|
|
"sk-none",
|
|
"tools/call",
|
|
{"name": "call_tool", "arguments": {"tool_id": "denied-scope", "arguments": {}}},
|
|
**{"x-mcp-servers": "math_restricted", "x-litellm-call-id": "proxy-scope-denial"},
|
|
)
|
|
assert response.status_code == 200, response.text
|
|
result = _rpc_result(response)
|
|
assert result["isError"] is True
|
|
assert result["content"][0]["text"] == (
|
|
"Error: The key is not allowed to access the requested MCP servers: math_restricted"
|
|
)
|
|
async with asyncio.timeout(10):
|
|
while True:
|
|
payload = json.loads(await asyncio.to_thread(proxy_call_recorder.failures.get, True, 5))
|
|
if payload["id"] == "proxy-scope-denial":
|
|
break
|
|
assert payload["call_type"] == "call_mcp_tool"
|
|
assert payload["status"] == "failure"
|
|
assert payload["response_cost"] == 0
|
|
assert "math_restricted" in payload["error_str"]
|
|
|
|
@pytest.mark.parametrize("arguments", ["wrong", False, None, [], 0])
|
|
def test_handler_rejects_non_object_arguments(
|
|
self, proxy_server_url: str, _proxy_server: ProxyRig, arguments: object
|
|
) -> None:
|
|
async def check() -> None:
|
|
auth = UserAPIKeyAuth(
|
|
object_permission=LiteLLM_ObjectPermissionTable(
|
|
object_permission_id="validation", mcp_servers=["math_stdio"]
|
|
)
|
|
)
|
|
hits = _payload(await handle_mcp_proxy_tool("search_tools", {"query": "add"}, auth))
|
|
tool_id = next(hit["tool_id"] for hit in hits if hit["name"] == "math_stdio-add")
|
|
result = await handle_mcp_proxy_tool("call_tool", {"tool_id": tool_id, "arguments": arguments}, auth)
|
|
assert result.is_error is True
|
|
assert result.content[0].text == "arguments must be an object"
|
|
|
|
asyncio.run_coroutine_threadsafe(check(), _proxy_server.loop).result(timeout=30)
|