Merge pull request #37388 from BerriAI/litellm_lit_5718_mcp_tool_bound_to_server

fix(mcp): bind tool existence check to the selected server
This commit is contained in:
Mateo Wang 2026-08-18 18:22:00 -07:00 • committed by GitHub
commit 704cc41f28
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 91 additions and 24 deletions

View file

@ -2321,23 +2321,30 @@ class MCPServerManager:
openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR
global_mcp_tool_registry.unregister_tools_with_prefix(openapi_key_prefix)
owned_raw: Final[set[str]] = set()
for p in iter_known_server_prefixes(server):
if p:
owned_raw.add(p)
if server.name:
owned_raw.add(server.name)
owned_normalized: Final = self._owned_mapping_values(server)
owned_normalized: Final = {normalize_server_name(x) for x in owned_raw}
stale_mapping_keys: Final[list[str]] = []
for tool_name, mapped_server in list(self.tool_name_to_mcp_server_name_mapping.items()):
if mapped_server in owned_raw or normalize_server_name(str(mapped_server)) in owned_normalized:
stale_mapping_keys.append(tool_name)
stale_mapping_keys: Final = tuple(
tool_name
for tool_name, mapped_server in self.tool_name_to_mcp_server_name_mapping.items()
if normalize_server_name(str(mapped_server)) in owned_normalized
)
for key in stale_mapping_keys:
del self.tool_name_to_mcp_server_name_mapping[key]
def _owned_mapping_values(self, server: MCPServer) -> frozenset[str]:
return frozenset(
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:
owned: Final = self._owned_mapping_values(server)
mapped_owners: Final = (
self.tool_name_to_mcp_server_name_mapping.get(spelling)
for spelling in iter_known_tool_name_spellings(tool_name, server)
)
return any(owner is not None and normalize_server_name(owner) in owned for owner in mapped_owners)
def remove_server(self, mcp_server: LiteLLM_MCPServerTable):
"""
Remove a server from the registry
@ -5463,13 +5470,8 @@ class MCPServerManager:
if mcp_server is None:
raise ValueError(f"Tool {name} not found")
if resolved_by_server_name_only:
tool_known: Final = (
name in self.tool_name_to_mcp_server_name_mapping
or prefixed_tool_name in self.tool_name_to_mcp_server_name_mapping
)
if not tool_known:
raise ValueError(f"Tool {name} not found")
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
@ -5847,10 +5849,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 (
original_tool_name in self.tool_name_to_mcp_server_name_mapping
or tool_name in self.tool_name_to_mcp_server_name_mapping
):
if matched_server is not None and self._server_exposes_tool(matched_server, original_tool_name):
return matched_server
return None

View file

@ -6406,7 +6406,7 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti
with (
patch.dict(
mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_name_mapping,
{"echo": collision_server.name},
{"echo": collision_server.name, "echo_requested-echo": requested_server.name},
),
patch.object(
mcp_module.global_mcp_server_manager,

View file

@ -4833,6 +4833,74 @@ class TestMCPServerManager:
with pytest.raises(ValueError, match="Tool missing_tool not found"):
manager._resolve_mcp_server_for_tool_call("github", "missing_tool")
@staticmethod
def _manager_with_deepwiki_and_huggingface() -> MCPServerManager:
manager = MCPServerManager()
deepwiki = MCPServer(server_id="deepwiki-id", name="deepwiki", server_name="deepwiki", transport=MCPTransport.http)
huggingface = MCPServer(
server_id="huggingface-id", name="huggingface", server_name="huggingface", transport=MCPTransport.http
)
manager.registry = {"deepwiki-id": deepwiki, "huggingface-id": huggingface}
manager.tool_name_to_mcp_server_name_mapping = {
"read_wiki_structure": "deepwiki",
"deepwiki-read_wiki_structure": "deepwiki",
"hub_repo_search": "huggingface",
"huggingface-hub_repo_search": "huggingface",
}
return manager
def test_resolve_mcp_server_for_tool_call_rejects_tool_exposed_only_by_another_server(self):
manager = self._manager_with_deepwiki_and_huggingface()
with pytest.raises(ValueError, match="Tool read_wiki_structure not found"):
manager._resolve_mcp_server_for_tool_call("huggingface", "read_wiki_structure")
with pytest.raises(ValueError, match="Tool hub_repo_search not found"):
manager._resolve_mcp_server_for_tool_call("deepwiki", "hub_repo_search")
assert manager._resolve_mcp_server_for_tool_call("deepwiki", "read_wiki_structure") is manager.registry["deepwiki-id"]
assert manager._resolve_mcp_server_for_tool_call("huggingface", "hub_repo_search") is manager.registry["huggingface-id"]
def test_get_mcp_server_from_tool_name_rejects_other_servers_prefix(self):
manager = self._manager_with_deepwiki_and_huggingface()
assert manager._get_mcp_server_from_tool_name("huggingface-read_wiki_structure") is None
assert manager._get_mcp_server_from_tool_name("deepwiki-hub_repo_search") is None
assert manager._get_mcp_server_from_tool_name("deepwiki-read_wiki_structure") is manager.registry["deepwiki-id"]
assert manager._get_mcp_server_from_tool_name("huggingface-hub_repo_search") is manager.registry["huggingface-id"]
def test_resolve_mcp_server_for_tool_call_shared_bare_name_resolves_via_own_prefixed_spelling(self):
manager = MCPServerManager()
zapier = MCPServer(server_id="zapier-id", name="zapier", alias="zapier-alias", transport=MCPTransport.http)
other = MCPServer(server_id="other-id", name="other", server_name="other", transport=MCPTransport.http)
manager.registry = {"zapier-id": zapier, "other-id": other}
manager.tool_name_to_mcp_server_name_mapping = {
"create_zap": "other",
"other-create_zap": "other",
"zapier-alias-create_zap": "zapier-alias",
}
assert manager._resolve_mcp_server_for_tool_call("zapier", "create_zap") is zapier
assert manager._resolve_mcp_server_for_tool_call("other", "create_zap") is other
def test_remove_server_drops_only_its_own_tool_mapping_rows(self):
manager = self._manager_with_deepwiki_and_huggingface()
manager.remove_server(
LiteLLM_MCPServerTable(
server_id="huggingface-id",
alias="huggingface",
url="https://huggingface.co/mcp",
transport=MCPTransport.http,
created_at=datetime.now(),
updated_at=datetime.now(),
)
)
assert manager.tool_name_to_mcp_server_name_mapping == {
"read_wiki_structure": "deepwiki",
"deepwiki-read_wiki_structure": "deepwiki",
}
@pytest.mark.asyncio
async def test_resolve_oauth2_headers_skipped_when_not_user_oauth(self):
"""Returns input headers unchanged when server does not need user OAuth."""