mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
commit
704cc41f28
3 changed files with 91 additions and 24 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue