mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): record a listed-tool catalog only for a listing the caller is served
_get_tools_from_server now records the catalog into the caller's listed-tools slot only when asked (record_listing=True), which the served listings pass: the /mcp and Responses API tools/list handlers via _get_tools_from_mcp_servers, MCPServerManager.list_tools, and the REST listing via _list_server_tools. Four internal listings stop recording, so a later tools/call hands pre_mcp_call hooks name and arguments only, as on main: - _list_tools_before_first_call, the implicit listing inside tools/call when this worker does not yet expose the tool - fetch_pinnable_tool_catalog, the admin pin snapshot listed without the catalog guard and without description overrides - _initialize_tool_name_to_mcp_server_name_mapping, the startup fill - get_tools_for_server, used by the semantic tool filter _create_prefixed_tools returns to its tool-name mapping job only; the record follows it in _get_tools_from_server.
This commit is contained in:
parent
ed75c1aef7
commit
ad746dadb7
6 changed files with 227 additions and 55 deletions
|
|
@ -3870,6 +3870,7 @@ class MCPServerManager:
|
|||
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:
|
||||
|
|
@ -4482,6 +4483,7 @@ class MCPServerManager:
|
|||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
*,
|
||||
catalog_auth_header: str | dict[str, str] | None | EllipsisType = ...,
|
||||
record_listing: bool = False,
|
||||
) -> Sequence[MCPTool]:
|
||||
"""
|
||||
Helper method to get tools from a single MCP server with prefixed names.
|
||||
|
|
@ -4491,6 +4493,8 @@ class MCPServerManager:
|
|||
mcp_auth_header: Optional auth header for MCP server
|
||||
catalog_auth_header: The header the client supplied, keying the caller's catalog slot;
|
||||
defaults to ``mcp_auth_header``
|
||||
record_listing: Record the served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
|
||||
Returns:
|
||||
List[MCPTool]: List of tools available on the server with prefixed names
|
||||
|
|
@ -4608,7 +4612,8 @@ class MCPServerManager:
|
|||
# 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)
|
||||
if record_listing:
|
||||
self._record_listed_tools(server, unprefixed_tools, listed_caller, listed_generation)
|
||||
if not add_prefix:
|
||||
return unprefixed_tools
|
||||
return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi]
|
||||
|
|
@ -4624,8 +4629,10 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
)
|
||||
prefixed_or_original_tools: Final = self._create_prefixed_tools(
|
||||
guarded_tools, server, add_prefix=add_prefix, caller=listed_caller, generation=listed_generation
|
||||
guarded_tools, server, add_prefix=add_prefix
|
||||
)
|
||||
if record_listing:
|
||||
self._record_listed_tools(server, guarded_tools, listed_caller, listed_generation)
|
||||
|
||||
return prefixed_or_original_tools
|
||||
|
||||
|
|
@ -5655,8 +5662,6 @@ class MCPServerManager:
|
|||
tools: Sequence[MCPTool],
|
||||
server: MCPServer,
|
||||
add_prefix: bool = True,
|
||||
caller: ListedToolsCaller | None = None,
|
||||
generation: int | None = None,
|
||||
) -> list[MCPTool]:
|
||||
"""
|
||||
Create prefixed tools and update tool mapping.
|
||||
|
|
@ -5682,7 +5687,6 @@ class MCPServerManager:
|
|||
for spelling in iter_known_tool_name_spellings(original_name, server):
|
||||
self.tool_name_to_mcp_server_name_mapping[spelling] = prefix
|
||||
|
||||
self._record_listed_tools(server, tools, caller, generation)
|
||||
verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name)
|
||||
return prefixed_tools
|
||||
|
||||
|
|
|
|||
|
|
@ -957,6 +957,8 @@ async def _get_tools_from_mcp_servers(
|
|||
request_tags: list[str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
*,
|
||||
record_listing: bool = True,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
Helper method to fetch tools from MCP servers based on server filtering criteria.
|
||||
|
|
@ -967,6 +969,8 @@ async def _get_tools_from_mcp_servers(
|
|||
mcp_servers: Optional list of server names/aliases to filter by
|
||||
mcp_server_auth_headers: Optional dict of server-specific auth headers
|
||||
oauth2_headers: Optional dict of oauth2 headers
|
||||
record_listing: Record each served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
|
||||
Returns:
|
||||
AggregateToolListing: Combined tools from filtered servers plus each server's
|
||||
|
|
@ -1132,6 +1136,7 @@ async def _get_tools_from_mcp_servers(
|
|||
oauth2_headers=oauth2_headers,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
catalog_auth_header=catalog_auth_header,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
|
|
@ -1805,6 +1810,7 @@ async def _list_tools_before_first_call(
|
|||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
record_listing=False,
|
||||
)
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -703,6 +703,8 @@ if MCP_AVAILABLE:
|
|||
extra_headers: dict[str, str] | None,
|
||||
client_ip: str | None,
|
||||
proxy_logging_obj: "ProxyLogging | None",
|
||||
*,
|
||||
record_listing: bool,
|
||||
) -> list[MCPTool]:
|
||||
return await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
|
|
@ -713,6 +715,7 @@ if MCP_AVAILABLE:
|
|||
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(
|
||||
|
|
@ -734,7 +737,14 @@ if MCP_AVAILABLE:
|
|||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
tools = await _list_server_tools(
|
||||
server, server_auth_header, raw_headers, user_api_key_auth, extra_headers, client_ip, proxy_logging_obj
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers,
|
||||
user_api_key_auth,
|
||||
extra_headers,
|
||||
client_ip,
|
||||
proxy_logging_obj,
|
||||
record_listing=True,
|
||||
)
|
||||
|
||||
if not apply_tool_filters:
|
||||
|
|
@ -776,6 +786,7 @@ if MCP_AVAILABLE:
|
|||
await _get_user_oauth_extra_headers(server, user_api_key_dict),
|
||||
IPAddressUtils.get_mcp_client_ip(request),
|
||||
None,
|
||||
record_listing=False,
|
||||
)
|
||||
scan: Final = await scan_tool_descriptions(
|
||||
apply_description_overrides(upstream, server), server, proxy_logging_obj, user_api_key_dict, raw_headers
|
||||
|
|
|
|||
|
|
@ -964,6 +964,7 @@ async def test_get_tools_from_mcp_servers():
|
|||
oauth2_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
catalog_auth_header=None,
|
||||
record_listing=True,
|
||||
):
|
||||
if server.server_id == "server1_id":
|
||||
return [mock_tool_1]
|
||||
|
|
@ -1999,6 +2000,7 @@ async def test_get_tools_for_single_server():
|
|||
client_ip=None,
|
||||
user_api_key_auth=None,
|
||||
proxy_logging_obj=ANY,
|
||||
record_listing=True,
|
||||
)
|
||||
|
||||
# Verify the result
|
||||
|
|
|
|||
|
|
@ -7050,7 +7050,8 @@ class TestMCPServerManager:
|
|||
manager.registry = {"test-server": server}
|
||||
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, caller=caller)
|
||||
manager._create_prefixed_tools(listed_tools, server)
|
||||
manager._record_listed_tools(server, listed_tools, caller)
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False)
|
||||
|
|
@ -7134,8 +7135,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._create_prefixed_tools([MCPTool(name="echo", description="v1", inputSchema={})], server)
|
||||
manager._create_prefixed_tools([MCPTool(name="echo", description="v2", inputSchema={})], server)
|
||||
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"
|
||||
|
|
@ -7147,13 +7148,13 @@ 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._create_prefixed_tools(
|
||||
manager._record_listed_tools(
|
||||
server,
|
||||
[
|
||||
MCPTool(name="foo", description="Fetches foo records", inputSchema={"type": "object"}),
|
||||
MCPTool(name="bar", description="Fetches bar records", inputSchema={"type": "object"}),
|
||||
],
|
||||
server,
|
||||
caller=caller,
|
||||
caller,
|
||||
)
|
||||
|
||||
assert manager.get_listed_tool(server, "srv-foo", caller) is None
|
||||
|
|
@ -7174,7 +7175,7 @@ class TestMCPServerManager:
|
|||
url="http://srv",
|
||||
tool_name_to_description={"echo": "Admin wording"},
|
||||
)
|
||||
await manager._get_tools_from_server(server, add_prefix=True)
|
||||
await manager._get_tools_from_server(server, add_prefix=True, record_listing=True)
|
||||
|
||||
overridden = manager.get_listed_tool(server, "echo")
|
||||
assert overridden is not None
|
||||
|
|
@ -7194,7 +7195,9 @@ class TestMCPServerManager:
|
|||
transport=MCPTransport.http,
|
||||
tool_name_to_description={"read_note": "Read a SECRET note"},
|
||||
)
|
||||
served = await manager._get_tools_from_server(server, add_prefix=True, proxy_logging_obj=proxy_logging_obj)
|
||||
served = await manager._get_tools_from_server(
|
||||
server, add_prefix=True, proxy_logging_obj=proxy_logging_obj, record_listing=True
|
||||
)
|
||||
assert [tool.description for tool in served] == ["Read a [MASKED] note"]
|
||||
|
||||
listed = manager.get_listed_tool(server, "read_note")
|
||||
|
|
@ -7206,8 +7209,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._create_prefixed_tools([MCPTool(name="echo", description="old", inputSchema={})], server)
|
||||
manager._create_prefixed_tools([MCPTool(name="ping", description="kept", inputSchema={})], 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._invalidate_server_definition_caches(server.server_id)
|
||||
|
||||
|
|
@ -7236,7 +7239,7 @@ class TestMCPServerManager:
|
|||
caller = ListedToolsCaller(user_api_key_auth=user)
|
||||
|
||||
async def list_tools() -> None:
|
||||
await manager._get_tools_from_server(server=server, user_api_key_auth=user)
|
||||
await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True)
|
||||
|
||||
listing = asyncio.create_task(list_tools())
|
||||
await fetch_started.wait()
|
||||
|
|
@ -7249,7 +7252,7 @@ class TestMCPServerManager:
|
|||
manager._fetch_tools_with_timeout = AsyncMock(
|
||||
return_value=[MCPTool(name="turn", description="after save", inputSchema={})]
|
||||
)
|
||||
await manager._get_tools_from_server(server=server, user_api_key_auth=user)
|
||||
await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True)
|
||||
listed = manager.get_listed_tool(server, "turn", caller)
|
||||
assert listed is not None and listed.description == "after save"
|
||||
|
||||
|
|
@ -7348,7 +7351,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._create_prefixed_tools([MCPTool(name="echo", description="shared", inputSchema={})], server)
|
||||
manager._record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None)
|
||||
|
||||
await manager.invalidate_user_oauth_token_cache("alice", server.server_id)
|
||||
|
||||
|
|
@ -7368,15 +7371,15 @@ 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._create_prefixed_tools(
|
||||
manager._record_listed_tools(
|
||||
server,
|
||||
[MCPTool(name="read", description="alice view", inputSchema=alice_schema)],
|
||||
server,
|
||||
caller=ListedToolsCaller(user_api_key_auth=alice),
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
manager._create_prefixed_tools(
|
||||
[MCPTool(name="read", description="bob view", inputSchema=bob_schema)],
|
||||
manager._record_listed_tools(
|
||||
server,
|
||||
caller=ListedToolsCaller(user_api_key_auth=bob),
|
||||
[MCPTool(name="read", description="bob view", inputSchema=bob_schema)],
|
||||
ListedToolsCaller(user_api_key_auth=bob),
|
||||
)
|
||||
|
||||
alice_tool = manager.get_listed_tool(server, "read", ListedToolsCaller(user_api_key_auth=alice))
|
||||
|
|
@ -7390,10 +7393,10 @@ 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._create_prefixed_tools(
|
||||
[MCPTool(name="echo", description="everyone", inputSchema={})],
|
||||
manager._record_listed_tools(
|
||||
shared,
|
||||
caller=ListedToolsCaller(user_api_key_auth=alice),
|
||||
[MCPTool(name="echo", description="everyone", inputSchema={})],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
for_bob = manager.get_listed_tool(shared, "echo", ListedToolsCaller(user_api_key_auth=bob))
|
||||
assert for_bob is None, "keyed callers get their own slot even on servers without upstream per-user auth"
|
||||
|
|
@ -7446,12 +7449,8 @@ class TestMCPServerManager:
|
|||
server = MCPServer(
|
||||
**{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs}
|
||||
)
|
||||
manager._create_prefixed_tools(
|
||||
[MCPTool(name="turn", description="Catalog A", inputSchema={})], server, caller=caller_a
|
||||
)
|
||||
manager._create_prefixed_tools(
|
||||
[MCPTool(name="turn", description="Catalog B", inputSchema={})], server, caller=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)
|
||||
|
|
@ -7462,10 +7461,10 @@ 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._create_prefixed_tools(
|
||||
[MCPTool(name="turn", description="everyone", inputSchema={})],
|
||||
manager._record_listed_tools(
|
||||
server,
|
||||
caller=ListedToolsCaller(raw_headers={"authorization": "Bearer sk-litellm", "x-workspace": "A"}),
|
||||
[MCPTool(name="turn", description="everyone", inputSchema={})],
|
||||
ListedToolsCaller(raw_headers={"authorization": "Bearer sk-litellm", "x-workspace": "A"}),
|
||||
)
|
||||
|
||||
other = ListedToolsCaller(raw_headers={"authorization": "Bearer sk-other", "x-workspace": "B"})
|
||||
|
|
@ -7499,7 +7498,7 @@ class TestMCPServerManager:
|
|||
AsyncMock(side_effect=RuntimeError("DB DOWN")),
|
||||
),
|
||||
):
|
||||
await manager._get_tools_from_server(server=server, user_api_key_auth=user)
|
||||
await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True)
|
||||
|
||||
listed = manager.get_listed_tool(server, "turn", listed_tools_caller_for(server, user, None, None, None, None))
|
||||
assert listed is not None and listed.description == "listed while db down"
|
||||
|
|
@ -7550,7 +7549,9 @@ class TestMCPServerManager:
|
|||
proxy_logging_obj.during_call_hook = AsyncMock(return_value=None)
|
||||
cache_byok_credential("byok-user", "byok-catalog", "stored-secret")
|
||||
try:
|
||||
await manager._get_tools_from_server(server=server, mcp_auth_header=list_header, user_api_key_auth=user)
|
||||
await manager._get_tools_from_server(
|
||||
server=server, mcp_auth_header=list_header, user_api_key_auth=user, record_listing=True
|
||||
)
|
||||
listed = manager.get_listed_tool(
|
||||
server, "turn", listed_tools_caller_for(server, user, list_header, None, None, None)
|
||||
)
|
||||
|
|
@ -7590,6 +7591,7 @@ class TestMCPServerManager:
|
|||
server=server,
|
||||
mcp_auth_header="Bearer hdr",
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"),
|
||||
record_listing=True,
|
||||
)
|
||||
|
||||
caller: Final = ListedToolsCaller(
|
||||
|
|
@ -7659,7 +7661,7 @@ class TestMCPServerManager:
|
|||
signer_headers,
|
||||
),
|
||||
):
|
||||
await manager._get_tools_from_server(server=server, user_api_key_auth=alice)
|
||||
await manager._get_tools_from_server(server=server, user_api_key_auth=alice, record_listing=True)
|
||||
finally:
|
||||
byok_credential_cache.delete_cache(byok_credential_cache_key("alice", "cc1"))
|
||||
|
||||
|
|
@ -7692,8 +7694,8 @@ class TestMCPServerManager:
|
|||
"litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer",
|
||||
return_value=signer,
|
||||
):
|
||||
manager._create_prefixed_tools(
|
||||
[MCPTool(name="turn", description="alice view", inputSchema={})], server, caller=alice
|
||||
manager._record_listed_tools(
|
||||
server, [MCPTool(name="turn", description="alice view", inputSchema={})], alice
|
||||
)
|
||||
assert manager.get_listed_tool(server, "turn", bob) is None
|
||||
|
||||
|
|
@ -7712,9 +7714,7 @@ class TestMCPServerManager:
|
|||
"litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer",
|
||||
return_value=MagicMock(),
|
||||
):
|
||||
manager._create_prefixed_tools(
|
||||
[MCPTool(name="turn", description="slot a", inputSchema={})], server, caller=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"))
|
||||
|
|
@ -7756,6 +7756,7 @@ class TestMCPServerManager:
|
|||
extra_headers={"X-Workspace": workspace},
|
||||
raw_headers={"x-workspace": workspace, "authorization": "Bearer sk-litellm"},
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="shared-key"),
|
||||
record_listing=True,
|
||||
)
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
|
|
@ -7789,20 +7790,18 @@ class TestMCPServerManager:
|
|||
url="http://srv",
|
||||
auth_type=MCPAuth.oauth2_token_exchange,
|
||||
)
|
||||
manager._create_prefixed_tools([MCPTool(name="read", description="shared", inputSchema={})], server)
|
||||
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._create_prefixed_tools(
|
||||
[MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})],
|
||||
manager._record_listed_tools(
|
||||
server,
|
||||
caller=caller,
|
||||
[MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})],
|
||||
caller,
|
||||
)
|
||||
manager._create_prefixed_tools(
|
||||
[MCPTool(name="read", description="u1 again", inputSchema={})], server, caller=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])
|
||||
|
|
@ -7840,7 +7839,7 @@ class TestMCPServerManager:
|
|||
handler=_handler,
|
||||
)
|
||||
try:
|
||||
listed = await manager._get_tools_from_server(server=server, add_prefix=add_prefix)
|
||||
listed = await manager._get_tools_from_server(server=server, add_prefix=add_prefix, record_listing=True)
|
||||
finally:
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
|
||||
|
||||
|
|
@ -7882,7 +7881,7 @@ class TestMCPServerManager:
|
|||
handler=_handler,
|
||||
)
|
||||
try:
|
||||
listed = await manager._get_tools_from_server(server=server, add_prefix=True)
|
||||
listed = await manager._get_tools_from_server(server=server, add_prefix=True, record_listing=True)
|
||||
finally:
|
||||
for prefix in ("pet-", "petstore-"):
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix(prefix)
|
||||
|
|
@ -7892,6 +7891,77 @@ class TestMCPServerManager:
|
|||
assert tool is not None and tool.description == "Local pet tool"
|
||||
assert tool.input_schema["properties"] == {"limit": {"type": "integer"}}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("openapi", [False, True], ids=["remote", "openapi"])
|
||||
async def test_get_tools_from_server_records_the_catalog_only_when_asked_to(self, openapi):
|
||||
"""The startup fill, the implicit pre-call listing and the pin snapshot reuse this fetch without
|
||||
serving its result, so only a listing that asks to be recorded sets what tools/call hooks see."""
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
if openapi:
|
||||
server = MCPServer(
|
||||
server_id="srv", name="srv", alias="srv", transport=MCPTransport.http, url=None, spec_path="/spec.yaml"
|
||||
)
|
||||
manager = MCPServerManager()
|
||||
manager._create_mcp_client = AsyncMock(return_value=AsyncMock())
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix("srv-")
|
||||
global_mcp_tool_registry.register_tool(
|
||||
name="srv-echo", description="Echoes", input_schema={"type": "object"}, handler=lambda **kwargs: None
|
||||
)
|
||||
else:
|
||||
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"}))
|
||||
user = UserAPIKeyAuth(api_key="sk-litellm", user_id="lister")
|
||||
|
||||
try:
|
||||
listed = await manager._get_tools_from_server(server=server, user_api_key_auth=user)
|
||||
assert [t.name for t in listed] == ["srv-echo"]
|
||||
assert server.server_id not in manager._listed_tools_by_server_id
|
||||
|
||||
await manager._get_tools_from_server(server=server, user_api_key_auth=user, record_listing=True)
|
||||
finally:
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix("srv-")
|
||||
|
||||
recorded = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user))
|
||||
assert recorded is not None and recorded.description == "Echoes"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools_records_the_served_catalog(self):
|
||||
manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"}))
|
||||
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
manager.registry = {"srv": server}
|
||||
manager.get_allowed_mcp_servers = AsyncMock(return_value=["srv"])
|
||||
user = UserAPIKeyAuth(api_key="sk-litellm", user_id="lister")
|
||||
|
||||
listed = await manager.list_tools(user_api_key_auth=user)
|
||||
|
||||
assert [t.name for t in listed] == ["srv-echo"]
|
||||
recorded = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user))
|
||||
assert recorded is not None and recorded.description == "Echoes"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_startup_tool_name_mapping_records_no_listed_catalog(self):
|
||||
manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"}))
|
||||
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
manager.registry = {"srv": server}
|
||||
|
||||
await manager._initialize_tool_name_to_mcp_server_name_mapping()
|
||||
|
||||
assert manager.server_exposes_tool(server, "echo") is True
|
||||
assert server.server_id not in manager._listed_tools_by_server_id
|
||||
assert manager.get_listed_tool(server, "echo") is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tools_for_server_records_no_listed_catalog(self):
|
||||
manager = _catalog_manager(MCPTool(name="echo", description="Echoes", inputSchema={"type": "object"}))
|
||||
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
|
||||
manager.registry = {"srv": server}
|
||||
|
||||
listed = await manager.get_tools_for_server("srv")
|
||||
|
||||
assert [t.name for t in listed] == ["srv-echo"]
|
||||
assert server.server_id not in manager._listed_tools_by_server_id
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_allowed_mcp_servers_with_user_api_key_auth(self):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -37,7 +37,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer, PinnedMCPTool
|
||||
|
||||
|
||||
def test_mcp_available_on_sdk2():
|
||||
|
|
@ -8332,6 +8332,85 @@ async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation
|
|||
assert result.content[0].text == "long"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hooks_no_description():
|
||||
"""The listing tools/call runs on its own when this worker does not yet expose the tool is never served
|
||||
to the caller, so it leaves the caller's listed slot empty and the pre-call hooks still get name and
|
||||
arguments only, as on main."""
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = _never_listed_passthrough_server()
|
||||
manager.registry[server.server_id] = server
|
||||
manager._listed_tools_by_server_id.pop(server.server_id, None)
|
||||
upstream = AsyncMock()
|
||||
upstream.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False)
|
||||
proxy_logging = _mock_mcp_proxy_logging()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
proxy_logging.during_call_hook = AsyncMock(return_value=None)
|
||||
fetch_tools = AsyncMock(
|
||||
return_value=[MCPTool(name="add", description="Adds. FLAGWORD", inputSchema={"type": "object"})]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=upstream)),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging),
|
||||
):
|
||||
result = await mcp_operations.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 fetch_tools.await_count == 1
|
||||
assert upstream.call_tool.await_count == 1
|
||||
assert result.content[0].text == "ok"
|
||||
hook_kwargs = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0]
|
||||
assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None)
|
||||
assert server.server_id not in manager._listed_tools_by_server_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_pinnable_tool_catalog_records_no_listed_catalog_for_the_admin():
|
||||
"""The pin snapshot lists the raw upstream catalog, without the catalog guard or the admin's description
|
||||
overrides, so it must not become what the admin's own later tools/call is evaluated against."""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server.rest_endpoints import fetch_pinnable_tool_catalog
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = MCPServer(
|
||||
server_id="pin-srv",
|
||||
name="pin_srv",
|
||||
transport=MCPTransport.http,
|
||||
url="https://up.example.com/mcp",
|
||||
tool_name_to_description={"add": "Admin wording"},
|
||||
)
|
||||
manager._listed_tools_by_server_id.pop(server.server_id, None)
|
||||
admin = UserAPIKeyAuth(api_key="sk-admin", user_id="admin")
|
||||
request = MagicMock()
|
||||
request.client.host = "10.1.2.3"
|
||||
request.headers = {"x-litellm-api-key": "sk-admin"}
|
||||
fetch_tools = AsyncMock(
|
||||
return_value=[MCPTool(name="add", description="Upstream wording", inputSchema={"type": "object"})]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=MagicMock())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", ProxyLogging(user_api_key_cache=DualCache())),
|
||||
):
|
||||
snapshot = await fetch_pinnable_tool_catalog(server, request, admin)
|
||||
|
||||
assert snapshot == {"add": PinnedMCPTool(description="Upstream wording", input_schema={"type": "object"})}
|
||||
assert server.server_id not in manager._listed_tools_by_server_id
|
||||
assert manager.get_listed_tool(server, "add", ListedToolsCaller(user_api_key_auth=admin)) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requested_server():
|
||||
"""A prefixed REST name that resolves to no tool must still dispatch to the server_id.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue