mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(mcp): preserve local tool ownership across discovery
Some checks failed
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run
Terraform Modules / fmt, validate, test (aws) (push) Has been cancelled
Terraform Modules / fmt, validate, test (gcp) (push) Has been cancelled
Some checks failed
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run
Terraform Modules / fmt, validate, test (aws) (push) Has been cancelled
Terraform Modules / fmt, validate, test (gcp) (push) Has been cancelled
This commit is contained in:
parent
57dfed6e39
commit
b63edce2dc
3 changed files with 101 additions and 0 deletions
|
|
@ -5292,6 +5292,8 @@ class MCPServerManager:
|
|||
Returns:
|
||||
List of tools with prefixed names
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
prefixed_tools: Final = []
|
||||
prefix: Final = get_server_prefix(server)
|
||||
|
||||
|
|
@ -5304,6 +5306,11 @@ class MCPServerManager:
|
|||
# short ID) so call_tool can resolve regardless of which form a
|
||||
# caller / cached client is using.
|
||||
for spelling in iter_known_tool_name_spellings(original_name, server):
|
||||
namespace_owner = self.server_owning_tool_name_prefix(spelling)
|
||||
if namespace_owner is not None and namespace_owner.server_id != server.server_id:
|
||||
continue
|
||||
if namespace_owner is None and global_mcp_tool_registry.get_tool(spelling) is not None:
|
||||
continue
|
||||
self.tool_name_to_mcp_server_name_mapping[spelling] = prefix
|
||||
|
||||
verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name)
|
||||
|
|
@ -6466,6 +6473,13 @@ class MCPServerManager:
|
|||
Returns:
|
||||
MCPServer if found, None otherwise
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
# Local handlers belong to their registered namespace, even if a native
|
||||
# discovery result or an older cached route claims the same spelling.
|
||||
if global_mcp_tool_registry.get_tool(tool_name) is not None:
|
||||
return self.server_owning_tool_name_prefix(tool_name)
|
||||
|
||||
registry_servers: Final = list(self.get_registry().values())
|
||||
prefix_to_server: Final = self._known_prefix_to_server()
|
||||
|
||||
|
|
|
|||
|
|
@ -17031,3 +17031,89 @@ async def test_catalog_cached_revision_retains_newly_discovered_tool_routes(monk
|
|||
if reader is not None:
|
||||
reader.cancel()
|
||||
await asyncio.gather(reader, return_exceptions=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("poisoned_route", [False, True])
|
||||
async def test_cached_discovery_cannot_authorize_another_servers_local_handler(monkeypatch, poisoned_route):
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import HTTPException
|
||||
from mcp.types import Tool
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
_catalog_database(monkeypatch, AsyncMock(return_value=[]), AsyncMock(return_value=_revision_row(7)))
|
||||
manager = MCPServerManager()
|
||||
allowed = MCPServer(server_id="allowed", name="allowed", transport=MCPTransport.http)
|
||||
private = MCPServer(server_id="private", name="private", transport=MCPTransport.http)
|
||||
manager.config_mcp_servers = {server.server_id: server for server in (allowed, private)}
|
||||
manager.published_tool_routes = {"getsecret": "private", "private-getsecret": "private"}
|
||||
handler = AsyncMock(return_value="private result")
|
||||
monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {})
|
||||
global_mcp_tool_registry.register_tool("private-getsecret", "Private tool", {}, handler)
|
||||
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
|
||||
check = AsyncMock(return_value={})
|
||||
monkeypatch.setattr(manager, "pre_call_tool_check", check)
|
||||
async with manager.catalog.operation():
|
||||
manager._create_prefixed_tools([Tool(name="private-getsecret", input_schema={})], allowed)
|
||||
if poisoned_route:
|
||||
manager.published_tool_routes["private-getsecret"] = "allowed"
|
||||
async with manager.catalog.operation():
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await operations._execute_mcp_tool(
|
||||
name="private-getsecret", arguments={}, allowed_mcp_servers=[allowed], start_time=datetime.now()
|
||||
)
|
||||
assert denied.value.status_code == 403
|
||||
handler.assert_not_awaited()
|
||||
check.assert_not_awaited()
|
||||
result = await operations._execute_mcp_tool(
|
||||
name="private-getsecret", arguments={}, allowed_mcp_servers=[private], start_time=datetime.now()
|
||||
)
|
||||
assert result.is_error is False
|
||||
handler.assert_awaited_once_with()
|
||||
assert check.await_args.kwargs["server"].server_id == private.server_id
|
||||
assert manager._get_mcp_server_from_tool_name("allowed-private-getsecret").server_id == allowed.server_id
|
||||
|
||||
|
||||
@pytest.mark.parametrize("local_handler", [False, True])
|
||||
def test_discovery_preserves_registered_tool_namespaces(monkeypatch, local_handler):
|
||||
from mcp.types import Tool
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
manager = MCPServerManager()
|
||||
allowed = MCPServer(server_id="allowed", name="allowed", transport=MCPTransport.http)
|
||||
private = MCPServer(server_id="private", name="private", alias="private-alias", transport=MCPTransport.http)
|
||||
manager.registry = {server.server_id: server for server in (allowed, private)}
|
||||
manager.published_tool_routes = {"private-getsecret": "private", "private-alias-getsecret": "private"}
|
||||
monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {})
|
||||
if local_handler:
|
||||
global_mcp_tool_registry.register_tool("orphan", "Orphan local tool", {}, AsyncMock())
|
||||
tools = [Tool(name=name, input_schema={}) for name in ("private-getsecret", "private-alias-getsecret", "orphan")]
|
||||
listed = manager._create_prefixed_tools(tools, allowed)
|
||||
assert [tool.name for tool in listed] == ["allowed-" + tool.name for tool in tools]
|
||||
assert manager.published_tool_routes["private-getsecret"] == "private"
|
||||
assert manager.published_tool_routes["private-alias-getsecret"] == "private"
|
||||
assert manager._get_mcp_server_from_tool_name("allowed-private-getsecret").server_id == "allowed"
|
||||
assert ("orphan" in manager.published_tool_routes) is not local_handler
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cached_catalog_excludes_routes_for_servers_outside_its_snapshot(monkeypatch):
|
||||
read_rows = AsyncMock(return_value=[])
|
||||
_catalog_database(monkeypatch, read_rows, AsyncMock(return_value=_revision_row(7)))
|
||||
manager = MCPServerManager()
|
||||
pinned = MCPServer(server_id="pinned", name="pinned", transport=MCPTransport.http)
|
||||
manager.config_mcp_servers = {pinned.server_id: pinned}
|
||||
async with manager.catalog.operation():
|
||||
assert manager.get_mcp_server_by_id("pinned") is not None
|
||||
published = MCPServer(server_id="published", name="published", transport=MCPTransport.http)
|
||||
manager.registry[published.server_id] = published
|
||||
manager.published_tool_routes = {"search": "published", "known": "pinned"}
|
||||
async with manager.catalog.operation():
|
||||
assert manager.get_mcp_server_by_id("published") is None
|
||||
assert "search" not in manager.tool_name_to_mcp_server_name_mapping
|
||||
assert manager._get_mcp_server_from_tool_name("known").server_id == "pinned"
|
||||
read_rows.assert_awaited_once()
|
||||
|
|
|
|||
|
|
@ -8632,6 +8632,7 @@ async def test_stateful_mcp_tool_call_uses_current_requests_otel_destinations(_m
|
|||
server = MCPServer(
|
||||
server_id="otel-context-test",
|
||||
name="otelcontext",
|
||||
server_name="otelcontext",
|
||||
transport=MCPTransport.http,
|
||||
allow_all_keys=True,
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue