mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
refactor(mcp): drop Sequence/list return-type mismatch and collapse record_listed_tools wrapper (#44556)
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
7a7d27c550
commit
254c2f6b3b
5 changed files with 67 additions and 73 deletions
|
|
@ -3801,7 +3801,7 @@ class MCPServerManager:
|
|||
if server is None:
|
||||
verbose_logger.warning("MCP Server %s not found", server_id)
|
||||
return []
|
||||
return await self._get_tools_from_server(server)
|
||||
return list(await self._get_tools_from_server(server))
|
||||
except Exception as e:
|
||||
verbose_logger.warning("Failed to get tools from server %s: %s", server_id, e)
|
||||
return []
|
||||
|
|
@ -3838,11 +3838,13 @@ class MCPServerManager:
|
|||
server_auth_header: Final = _server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header)
|
||||
|
||||
try:
|
||||
tools: Final = await self._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
record_listing=True,
|
||||
tools: Final = list(
|
||||
await self._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
record_listing=True,
|
||||
)
|
||||
)
|
||||
return tools
|
||||
except Exception as e:
|
||||
|
|
@ -4350,12 +4352,11 @@ class MCPServerManager:
|
|||
# applied (e.g. "test_petstore-getinventory"). Do NOT pass them
|
||||
# through _create_prefixed_tools — that would add the prefix a second
|
||||
# time producing "test_petstore-test_petstore-getinventory".
|
||||
unprefixed_tools: Final = guarded_openapi
|
||||
self._record_listed_tools(
|
||||
server, unprefixed_tools, listed_caller, listed_generation, record_listing=record_listing
|
||||
self.record_listed_tools(
|
||||
server, guarded_openapi, listed_caller, listed_generation, record_listing=record_listing
|
||||
)
|
||||
if not add_prefix:
|
||||
return unprefixed_tools
|
||||
return guarded_openapi
|
||||
return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi]
|
||||
else:
|
||||
tools = await self._fetch_tools_with_timeout(client, server.name)
|
||||
|
|
@ -4371,7 +4372,7 @@ class MCPServerManager:
|
|||
prefixed_or_original_tools: Final = self._create_prefixed_tools(
|
||||
guarded_tools, server, add_prefix=add_prefix
|
||||
)
|
||||
self._record_listed_tools(
|
||||
self.record_listed_tools(
|
||||
server, guarded_tools, listed_caller, listed_generation, record_listing=record_listing
|
||||
)
|
||||
|
||||
|
|
@ -4478,17 +4479,6 @@ class MCPServerManager:
|
|||
return self._listed_tools_generations.get(server_id, 0)
|
||||
|
||||
def record_listed_tools(
|
||||
self,
|
||||
server: MCPServer,
|
||||
tools: Sequence[MCPTool],
|
||||
caller: ListedToolsCaller | None,
|
||||
generation: int,
|
||||
*,
|
||||
record_listing: bool = True,
|
||||
) -> None:
|
||||
self._record_listed_tools(server, tools, caller, generation, record_listing=record_listing)
|
||||
|
||||
def _record_listed_tools(
|
||||
self,
|
||||
server: MCPServer,
|
||||
tools: Sequence[MCPTool],
|
||||
|
|
|
|||
|
|
@ -1125,18 +1125,20 @@ async def _get_tools_from_mcp_servers(
|
|||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id)
|
||||
tools: Final = await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
add_prefix=True, # Always add server prefix
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
oauth2_headers=oauth2_headers,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
catalog_auth_header=catalog_auth_header,
|
||||
record_listing=False,
|
||||
tools: Final = list(
|
||||
await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
add_prefix=True, # Always add server prefix
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
oauth2_headers=oauth2_headers,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
catalog_auth_header=catalog_auth_header,
|
||||
record_listing=False,
|
||||
)
|
||||
)
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
|
|
|
|||
|
|
@ -705,16 +705,18 @@ if MCP_AVAILABLE:
|
|||
*,
|
||||
record_listing: bool,
|
||||
) -> list[MCPTool]:
|
||||
return await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
add_prefix=False,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
record_listing=record_listing,
|
||||
return list(
|
||||
await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
add_prefix=False,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
)
|
||||
|
||||
async def _get_tools_for_single_server(
|
||||
|
|
|
|||
|
|
@ -7197,7 +7197,7 @@ class TestMCPServerManager:
|
|||
manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server"
|
||||
manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server"
|
||||
manager._create_prefixed_tools(listed_tools, server)
|
||||
manager._record_listed_tools(server, listed_tools, caller)
|
||||
manager.record_listed_tools(server, listed_tools, caller)
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False)
|
||||
|
|
@ -7281,8 +7281,8 @@ class TestMCPServerManager:
|
|||
def test_get_listed_tool_resolves_the_bare_name_from_the_latest_listing(self):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
manager._record_listed_tools(server, [MCPTool(name="echo", description="v1", inputSchema={})], None)
|
||||
manager._record_listed_tools(server, [MCPTool(name="echo", description="v2", inputSchema={})], None)
|
||||
manager.record_listed_tools(server, [MCPTool(name="echo", description="v1", inputSchema={})], None)
|
||||
manager.record_listed_tools(server, [MCPTool(name="echo", description="v2", inputSchema={})], None)
|
||||
|
||||
latest = manager.get_listed_tool(server, "echo")
|
||||
assert latest is not None and latest.description == "v2"
|
||||
|
|
@ -7294,7 +7294,7 @@ class TestMCPServerManager:
|
|||
manager = MCPServerManager()
|
||||
server = MCPServer(server_id="srv-id", name="srv", alias="srv", transport=MCPTransport.http, url="http://srv")
|
||||
caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice"))
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
server,
|
||||
[
|
||||
MCPTool(name="foo", description="Fetches foo records", inputSchema={"type": "object"}),
|
||||
|
|
@ -7355,8 +7355,8 @@ class TestMCPServerManager:
|
|||
manager = MCPServerManager()
|
||||
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
other = MCPServer(server_id="other", name="other", transport=MCPTransport.http, url="http://other")
|
||||
manager._record_listed_tools(server, [MCPTool(name="echo", description="old", inputSchema={})], None)
|
||||
manager._record_listed_tools(other, [MCPTool(name="ping", description="kept", inputSchema={})], None)
|
||||
manager.record_listed_tools(server, [MCPTool(name="echo", description="old", inputSchema={})], None)
|
||||
manager.record_listed_tools(other, [MCPTool(name="ping", description="kept", inputSchema={})], None)
|
||||
|
||||
manager._invalidate_server_definition_caches(server.server_id)
|
||||
|
||||
|
|
@ -7418,7 +7418,7 @@ class TestMCPServerManager:
|
|||
caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="lister"))
|
||||
|
||||
async def register_while_a_listing_records(server: MCPServer, *, initialize_mapping: bool = True) -> None:
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
server,
|
||||
[MCPTool(name="search", description="pre-save", inputSchema={})],
|
||||
caller,
|
||||
|
|
@ -7463,7 +7463,7 @@ class TestMCPServerManager:
|
|||
return [Prompt(name="greet")]
|
||||
|
||||
async def register_while_discovery_fills(server: MCPServer, *, initialize_mapping: bool = True) -> None:
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
server,
|
||||
[MCPTool(name="search", description="pre-save", inputSchema={})],
|
||||
caller,
|
||||
|
|
@ -7497,7 +7497,7 @@ class TestMCPServerManager:
|
|||
async def test_user_oauth_refresh_keeps_listed_tools(self):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
manager._record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None)
|
||||
manager.record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None)
|
||||
|
||||
await manager.invalidate_user_oauth_token_cache("alice", server.server_id)
|
||||
|
||||
|
|
@ -7517,12 +7517,12 @@ class TestMCPServerManager:
|
|||
bob = UserAPIKeyAuth(user_id="bob", token="hashed-bob")
|
||||
alice_schema = {"type": "object", "properties": {"path": {"type": "string"}}}
|
||||
bob_schema = {"type": "object", "properties": {"path": {"type": "string"}, "site": {"type": "string"}}}
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
server,
|
||||
[MCPTool(name="read", description="alice view", inputSchema=alice_schema)],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
server,
|
||||
[MCPTool(name="read", description="bob view", inputSchema=bob_schema)],
|
||||
ListedToolsCaller(user_api_key_auth=bob),
|
||||
|
|
@ -7539,7 +7539,7 @@ class TestMCPServerManager:
|
|||
assert manager.get_listed_tool(server, "read", carol) is None
|
||||
|
||||
shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared")
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
shared,
|
||||
[MCPTool(name="echo", description="everyone", inputSchema={})],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
|
|
@ -7595,8 +7595,8 @@ class TestMCPServerManager:
|
|||
server = MCPServer(
|
||||
**{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs}
|
||||
)
|
||||
manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog A", inputSchema={})], caller_a)
|
||||
manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog B", inputSchema={})], caller_b)
|
||||
manager.record_listed_tools(server, [MCPTool(name="turn", description="Catalog A", inputSchema={})], caller_a)
|
||||
manager.record_listed_tools(server, [MCPTool(name="turn", description="Catalog B", inputSchema={})], caller_b)
|
||||
|
||||
for_a = manager.get_listed_tool(server, "turn", caller_a)
|
||||
for_b = manager.get_listed_tool(server, "turn", caller_b)
|
||||
|
|
@ -7607,7 +7607,7 @@ class TestMCPServerManager:
|
|||
def test_shared_server_ignores_headers_it_never_forwards(self):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
server,
|
||||
[MCPTool(name="turn", description="everyone", inputSchema={})],
|
||||
ListedToolsCaller(raw_headers={"authorization": "Bearer sk-litellm", "x-workspace": "A"}),
|
||||
|
|
@ -7840,7 +7840,7 @@ class TestMCPServerManager:
|
|||
"litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer",
|
||||
return_value=signer,
|
||||
):
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
server, [MCPTool(name="turn", description="alice view", inputSchema={})], alice
|
||||
)
|
||||
assert manager.get_listed_tool(server, "turn", bob) is None
|
||||
|
|
@ -7860,7 +7860,7 @@ class TestMCPServerManager:
|
|||
"litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer",
|
||||
return_value=MagicMock(),
|
||||
):
|
||||
manager._record_listed_tools(server, [MCPTool(name="turn", description="slot a", inputSchema={})], alice)
|
||||
manager.record_listed_tools(server, [MCPTool(name="turn", description="slot a", inputSchema={})], alice)
|
||||
assert manager.get_listed_tool(server, "turn", bob) is None
|
||||
|
||||
same_key = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha"))
|
||||
|
|
@ -7878,7 +7878,7 @@ class TestMCPServerManager:
|
|||
team_two: Final = ListedToolsCaller(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-two")
|
||||
)
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], team_one
|
||||
)
|
||||
|
||||
|
|
@ -7896,7 +7896,7 @@ class TestMCPServerManager:
|
|||
alice_in_two: Final = ListedToolsCaller(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id="alice", team_id="team-two")
|
||||
)
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], alice_in_one
|
||||
)
|
||||
|
||||
|
|
@ -7917,7 +7917,7 @@ class TestMCPServerManager:
|
|||
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one"),
|
||||
raw_headers={"authorization": "Bearer jwt-bob"},
|
||||
)
|
||||
manager._record_listed_tools(server, [MCPTool(name="foo", description="alice view", inputSchema={})], alice)
|
||||
manager.record_listed_tools(server, [MCPTool(name="foo", description="alice view", inputSchema={})], alice)
|
||||
|
||||
assert manager.get_listed_tool(server, "foo", bob) is None
|
||||
listed: Final = manager.get_listed_tool(server, "foo", alice)
|
||||
|
|
@ -7976,7 +7976,7 @@ class TestMCPServerManager:
|
|||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-master"),
|
||||
raw_headers={"x-litellm-api-key": "Bearer sk-master", "authorization": "Bearer UP-B"},
|
||||
)
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
server, [MCPTool(name="lookup", description="Workspace A lookup FLAGWORD", inputSchema={})], caller_a
|
||||
)
|
||||
|
||||
|
|
@ -8053,18 +8053,18 @@ class TestMCPServerManager:
|
|||
url="http://srv",
|
||||
auth_type=MCPAuth.oauth2_token_exchange,
|
||||
)
|
||||
manager._record_listed_tools(server, [MCPTool(name="read", description="shared", inputSchema={})], None)
|
||||
manager.record_listed_tools(server, [MCPTool(name="read", description="shared", inputSchema={})], None)
|
||||
callers = [
|
||||
ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id=f"u{i}", api_key=f"k{i}"))
|
||||
for i in range(_LISTED_TOOLS_CALLERS_PER_SERVER + 1)
|
||||
]
|
||||
for caller in callers:
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
server,
|
||||
[MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})],
|
||||
caller,
|
||||
)
|
||||
manager._record_listed_tools(server, [MCPTool(name="read", description="u1 again", inputSchema={})], callers[1])
|
||||
manager.record_listed_tools(server, [MCPTool(name="read", description="u1 again", inputSchema={})], callers[1])
|
||||
|
||||
assert manager.get_listed_tool(server, "read", callers[0]) is None
|
||||
second = manager.get_listed_tool(server, "read", callers[1])
|
||||
|
|
|
|||
|
|
@ -8115,7 +8115,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_listed_entry_and_nothing
|
|||
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging),
|
||||
):
|
||||
never_listed_tool, never_listed_data = await call()
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="list_pets", description="ADMIN DESC", inputSchema=schema)],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
|
|
@ -8156,7 +8156,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_cl
|
|||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=pinned_schema)],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
|
|
@ -8206,12 +8206,12 @@ async def test_execute_mcp_tool_hands_openapi_hooks_each_callers_own_listed_entr
|
|||
manager = mcp_module.global_mcp_server_manager
|
||||
guarded = UserAPIKeyAuth(api_key="sk-guarded", user_id="alice")
|
||||
opted_out = UserAPIKeyAuth(api_key="sk-opted-out", user_id="bob")
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=schema)],
|
||||
ListedToolsCaller(user_api_key_auth=guarded),
|
||||
)
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="getpetbyid", description="Find a SECRET pet", inputSchema=schema)],
|
||||
ListedToolsCaller(user_api_key_auth=opted_out),
|
||||
|
|
@ -8305,7 +8305,7 @@ async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation
|
|||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
|
||||
manager._record_listed_tools(
|
||||
manager.record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="get_pet", description="Fetches pet records. FLAGWORD", inputSchema={"type": "object"})],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue