Merge pull request #42072 from BerriAI/litellm_mcp_cold_worker_tools_call

fix(mcp): tools/call no longer 404s on a worker that has not served tools/list
This commit is contained in:
Mateo Wang 2026-09-19 20:04:33 -07:00 committed by GitHub
commit ef7da9b49f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 306 additions and 12 deletions

View file

@ -2867,7 +2867,7 @@ class MCPServerManager:
normalize_server_name(value) for value in (*iter_known_server_prefixes(server), server.name) if value
)
def _server_exposes_tool(self, server: MCPServer, tool_name: str) -> bool:
def server_exposes_tool(self, server: MCPServer, tool_name: str) -> bool:
owned: Final = self._owned_mapping_values(server)
mapped_owners: Final = (
self.tool_name_to_mcp_server_name_mapping.get(spelling)
@ -2875,6 +2875,20 @@ class MCPServerManager:
)
return any(owner is not None and normalize_server_name(owner) in owned for owner in mapped_owners)
def _known_prefix_to_server(self) -> Mapping[str, MCPServer]:
"""Every prefix form a tool name may carry, keyed to its server; a form two servers share
stays with the one registered first."""
return {
normalize_server_name(known_prefix): server
for server in reversed(tuple(self.get_registry().values()))
for known_prefix in iter_known_server_prefixes(server)
}
def server_owning_tool_name_prefix(self, tool_name: str) -> MCPServer | None:
prefix_to_server: Final = self._known_prefix_to_server()
matched: Final = match_known_server_prefix(tool_name, prefix_to_server.keys())
return None if matched is None else prefix_to_server.get(matched[0])
def remove_server(self, mcp_server: LiteLLM_MCPServerTable):
"""
Remove a server from the registry
@ -6114,7 +6128,7 @@ class MCPServerManager:
if mcp_server is None:
raise ValueError(f"Tool {name} not found")
if resolved_by_server_name_only and not self._server_exposes_tool(mcp_server, name):
if resolved_by_server_name_only and not self.server_exposes_tool(mcp_server, name):
raise ValueError(f"Tool {name} not found")
return mcp_server
@ -6475,15 +6489,7 @@ class MCPServerManager:
MCPServer if found, None otherwise
"""
registry_servers: Final = list(self.get_registry().values())
# Build prefix → server lookup covering every known form a tool name
# may take (alias / server_name / server_id / short ID). This is what
# makes the short-prefix mode work without breaking historical names.
prefix_to_server: Final[dict[str, MCPServer]] = {}
for server in registry_servers:
for known_prefix in iter_known_server_prefixes(server):
normalised = normalize_server_name(known_prefix)
prefix_to_server.setdefault(normalised, server)
prefix_to_server: Final = self._known_prefix_to_server()
# First try with the original tool name
if tool_name in self.tool_name_to_mcp_server_name_mapping:
@ -6501,7 +6507,7 @@ class MCPServerManager:
if matched is not None:
matched_prefix, original_tool_name = matched
matched_server: Final = prefix_to_server.get(matched_prefix)
if matched_server is not None and self._server_exposes_tool(matched_server, original_tool_name):
if matched_server is not None and self.server_exposes_tool(matched_server, original_tool_name):
return matched_server
return None

View file

@ -2888,6 +2888,40 @@ if MCP_AVAILABLE:
headers={"WWW-Authenticate": get_byok_www_authenticate()},
)
async def _list_tools_before_first_call(
server: MCPServer | None,
tool_name: str,
allowed_mcp_servers: list[MCPServer],
user_api_key_auth: UserAPIKeyAuth | None,
mcp_auth_header: str | None,
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
oauth2_headers: dict[str, str] | None,
raw_headers: dict[str, str] | None,
) -> None:
"""List ``server`` with the caller's own credentials when it does not yet expose ``tool_name`` here.
The startup fill skips a server whose upstream wants the caller's token, and mcp 2 no
longer lists before an uncached tools/call, so a worker that has not served tools/list
for this caller would otherwise answer 404 for a tool the caller can see. Gating on the
requested tool, not on any prior listing, keeps callers with different upstream catalogs
from masking each other.
"""
if server is None or global_mcp_server_manager.server_exposes_tool(server, tool_name):
return
if all(allowed.server_id != server.server_id for allowed in allowed_mcp_servers):
return
try:
await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=[server.server_id],
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
except Exception as e: # noqa: BLE001 # best effort: resolution below answers as it did before
verbose_logger.debug("MCP tools/call: listing %s before its first call failed: %s", server.name, e)
async def execute_mcp_tool(
name: str,
arguments: dict[str, object],
@ -2948,6 +2982,27 @@ if MCP_AVAILABLE:
all_registry_prefixes.add(normalize_server_name(known_prefix))
name_is_prefixed = is_tool_name_prefixed(name, known_server_prefixes=all_registry_prefixes)
first_call_target: Final = (
requested_server
if requested_server is not None and not name_is_prefixed
else global_mcp_server_manager.server_owning_tool_name_prefix(name)
)
first_call_tool_name: Final = (
name
if first_call_target is None or (requested_server is not None and not name_is_prefixed)
else strip_known_server_prefix(name, first_call_target)
)
await _list_tools_before_first_call(
server=first_call_target,
tool_name=first_call_tool_name,
allowed_mcp_servers=allowed_mcp_servers,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
if requested_server is not None and not name_is_prefixed:
# REST callers may pass server_id with the upstream tool name (no
# LiteLLM prefix). The first segment is not a registered server

View file

@ -7437,6 +7437,186 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool
assert captured["name"] == "echo"
def _never_listed_passthrough_server() -> MCPServer:
return MCPServer(
server_id="lazy-map-1",
name="lazy_map",
server_name="lazy_map",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.true_passthrough,
)
@contextlib.contextmanager
def _worker_that_never_listed(server: MCPServer, upstream_tools: tuple[str, ...]):
"""A worker whose tool rows for ``server`` are empty, in front of an upstream that answers
tools/list with ``upstream_tools`` and a managed dispatch that records what reaches it."""
from mcp.types import Tool as MCPTool
from litellm.proxy._experimental.mcp_server import server as mcp_module
mcp_module.global_mcp_server_manager.registry[server.server_id] = server
dispatched: dict[str, object] = {}
async def fake_handle_managed_mcp_tool(**kwargs):
dispatched.update(kwargs)
return CallToolResult(content=[TextContent(type="text", text="ok")], is_error=False)
async def fake_fetch_tools(client, server_name):
return [MCPTool(name=tool_name, inputSchema={}) for tool_name in upstream_tools]
with (
patch.object( # test-quality-ok: the upstream MCP session is the boundary; a real one needs an initialize handshake over a live server
mcp_module.global_mcp_server_manager,
"_create_mcp_client",
new=AsyncMock(return_value=MagicMock()),
) as create_client,
patch.object( # test-quality-ok: same boundary, this is the tools/list answer the upstream would give
mcp_module.global_mcp_server_manager,
"_fetch_tools_with_timeout",
side_effect=fake_fetch_tools,
) as fetch_tools,
patch.object( # test-quality-ok: records the resolved server and bare name the managed call would forward upstream
mcp_module, "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool
),
):
yield SimpleNamespace(create_client=create_client, fetch_tools=fetch_tools, dispatched=dispatched)
@pytest.mark.asyncio
async def test_execute_mcp_tool_lists_never_listed_passthrough_server_with_caller_token_first():
"""A prefixed tools/call on a worker that has not served tools/list must list that server once
with the caller's own credentials and then dispatch, instead of answering 404."""
from litellm.proxy._experimental.mcp_server import server as mcp_module
server = _never_listed_passthrough_server()
with _worker_that_never_listed(server, upstream_tools=("add",)) as worker:
await mcp_module.execute_mcp_tool(
name="lazy_map-add",
arguments={"a": 1, "b": 2},
allowed_mcp_servers=[server],
start_time=datetime.now(),
mcp_auth_header="Bearer caller-token",
raw_headers={"authorization": "Bearer caller-token"},
)
assert worker.fetch_tools.await_count == 1
assert "caller-token" in str(worker.create_client.await_args.kwargs.get("mcp_auth_header"))
assert worker.dispatched["server_name"] == "lazy_map"
assert worker.dispatched["name"] == "add"
@pytest.mark.asyncio
async def test_execute_mcp_tool_rest_server_id_lists_never_listed_server_first():
from litellm.proxy._experimental.mcp_server import server as mcp_module
server = _never_listed_passthrough_server()
with _worker_that_never_listed(server, upstream_tools=("add",)) as worker:
await mcp_module.execute_mcp_tool(
name="add",
arguments={"a": 1, "b": 2},
allowed_mcp_servers=[server],
start_time=datetime.now(),
mcp_auth_header="Bearer caller-token",
requested_server_id=server.server_id,
)
assert worker.fetch_tools.await_count == 1
assert worker.dispatched["server_name"] == "lazy_map"
assert worker.dispatched["name"] == "add"
@pytest.mark.asyncio
async def test_execute_mcp_tool_unknown_tool_on_never_listed_server_lists_once_then_404s():
from litellm.proxy._experimental.mcp_server import server as mcp_module
server = _never_listed_passthrough_server()
with (
_worker_that_never_listed(server, upstream_tools=("add",)) as worker,
pytest.raises(HTTPException) as exc_info,
):
await mcp_module.execute_mcp_tool(
name="lazy_map-nope",
arguments={},
allowed_mcp_servers=[server],
start_time=datetime.now(),
mcp_auth_header="Bearer caller-token",
)
assert exc_info.value.status_code == 404
assert worker.fetch_tools.await_count == 1
assert worker.dispatched == {}
@pytest.mark.asyncio
async def test_execute_mcp_tool_does_not_relist_a_server_this_worker_already_listed():
from mcp.types import Tool as MCPTool
from litellm.proxy._experimental.mcp_server import server as mcp_module
server = _never_listed_passthrough_server()
with _worker_that_never_listed(server, upstream_tools=("add",)) as worker:
mcp_module.global_mcp_server_manager._create_prefixed_tools([MCPTool(name="add", inputSchema={})], server)
await mcp_module.execute_mcp_tool(
name="lazy_map-add",
arguments={"a": 1, "b": 2},
allowed_mcp_servers=[server],
start_time=datetime.now(),
mcp_auth_header="Bearer caller-token",
)
assert worker.fetch_tools.await_count == 0
assert worker.dispatched["name"] == "add"
@pytest.mark.asyncio
async def test_execute_mcp_tool_lists_a_tool_this_worker_has_not_yet_seen_on_a_listed_server():
"""A worker that already holds one of the server's tools must still list when a caller asks
for a different tool it has not cached, so callers with wider upstream catalogs are not 404ed."""
from mcp.types import Tool as MCPTool
from litellm.proxy._experimental.mcp_server import server as mcp_module
server = _never_listed_passthrough_server()
with _worker_that_never_listed(server, upstream_tools=("add", "multiply")) as worker:
mcp_module.global_mcp_server_manager._create_prefixed_tools([MCPTool(name="add", inputSchema={})], server)
await mcp_module.execute_mcp_tool(
name="lazy_map-multiply",
arguments={"a": 1, "b": 2},
allowed_mcp_servers=[server],
start_time=datetime.now(),
mcp_auth_header="Bearer caller-token",
)
assert worker.fetch_tools.await_count == 1
assert worker.dispatched["server_name"] == "lazy_map"
assert worker.dispatched["name"] == "multiply"
@pytest.mark.asyncio
async def test_execute_mcp_tool_never_lists_a_server_the_caller_cannot_access():
from litellm.proxy._experimental.mcp_server import server as mcp_module
server = _never_listed_passthrough_server()
other_server = MCPServer(server_id="other-1", name="other", transport=MCPTransport.http)
with (
_worker_that_never_listed(server, upstream_tools=("add",)) as worker,
pytest.raises(HTTPException) as exc_info,
):
await mcp_module.execute_mcp_tool(
name="lazy_map-add",
arguments={"a": 1, "b": 2},
allowed_mcp_servers=[other_server],
start_time=datetime.now(),
mcp_auth_header="Bearer caller-token",
)
assert exc_info.value.status_code == 403
assert worker.fetch_tools.await_count == 0
assert worker.dispatched == {}
@pytest.mark.asyncio
async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator():
"""A server with no alias publishes its UUID server_id as the tool prefix.

View file

@ -5556,6 +5556,59 @@ class TestMCPServerManager:
)
mock_inject.assert_awaited_once()
def test_server_owning_tool_name_prefix_is_known_before_the_server_is_ever_listed(self):
manager = MCPServerManager()
server = MCPServer(
server_id="lazy-map-1",
name="lazy_map",
server_name="lazy_map",
transport=MCPTransport.http,
auth_type=MCPAuth.true_passthrough,
)
manager.registry = {server.server_id: server}
assert manager._get_mcp_server_from_tool_name("lazy_map-add") is None
assert manager.server_owning_tool_name_prefix("lazy_map-add") is server
assert manager.server_owning_tool_name_prefix("someone_else-add") is None
manager._create_prefixed_tools([MCPTool(name="add", inputSchema={})], server)
assert manager._get_mcp_server_from_tool_name("lazy_map-add") is server
def test_server_exposes_tool_is_per_tool_not_per_server(self):
"""A tool listed for the server does not make its unlisted siblings look exposed."""
manager = MCPServerManager()
server = MCPServer(
server_id="lazy-map-2",
name="lazy_map",
server_name="lazy_map",
transport=MCPTransport.http,
auth_type=MCPAuth.true_passthrough,
)
manager.registry = {server.server_id: server}
assert manager.server_exposes_tool(server, "add") is False
manager._create_prefixed_tools([MCPTool(name="add", inputSchema={})], server)
assert manager.server_exposes_tool(server, "add") is True
assert manager.server_exposes_tool(server, "lazy_map-add") is True
assert manager.server_exposes_tool(server, "multiply") is False
def test_known_prefix_to_server_keeps_the_first_registered_owner_of_a_shared_prefix(self):
manager = MCPServerManager()
first = MCPServer(server_id="first-id", name="first", server_name="first", transport=MCPTransport.http)
second = MCPServer(
server_id="second-id", name="second", server_name="second", alias="first", transport=MCPTransport.http
)
manager.registry = {"first-id": first, "second-id": second}
prefix_to_server = manager._known_prefix_to_server()
assert prefix_to_server["first"] is first
assert prefix_to_server["second"] is second
assert manager.server_owning_tool_name_prefix("first-add") is first
def test_resolve_mcp_server_for_tool_call_via_prefixed_name(self):
"""Resolution succeeds when the prefixed tool name is in the mapping."""
manager = MCPServerManager()