mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
test(mcp): align listed-tool slot tests with per-caller keying
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
7dee160b9f
commit
001fcad930
1 changed files with 36 additions and 29 deletions
|
|
@ -7007,10 +7007,8 @@ class TestMCPServerManager:
|
|||
# Mock _create_mcp_client to return our mock client
|
||||
manager._create_mcp_client = AsyncMock(return_value=mock_client)
|
||||
|
||||
# Mock user auth with no restrictions
|
||||
user_api_key_auth = MagicMock()
|
||||
user_api_key_auth.object_permission = None
|
||||
user_api_key_auth.object_permission_id = None
|
||||
# Real auth: the listed-tool slot identity is hashed from these fields
|
||||
user_api_key_auth = UserAPIKeyAuth(api_key="sk-test")
|
||||
|
||||
# Mock proxy logging
|
||||
proxy_logging_obj = MagicMock()
|
||||
|
|
@ -7037,7 +7035,9 @@ class TestMCPServerManager:
|
|||
assert mock_client.call_tool.await_count == 1
|
||||
|
||||
@staticmethod
|
||||
def _manager_ready_for_call_tool(listed_tools: list[MCPTool]) -> tuple[MCPServerManager, MagicMock]:
|
||||
def _manager_ready_for_call_tool(
|
||||
listed_tools: list[MCPTool], caller: ListedToolsCaller | None = None
|
||||
) -> tuple[MCPServerManager, MagicMock]:
|
||||
from mcp.types import CallToolResult
|
||||
|
||||
manager = MCPServerManager()
|
||||
|
|
@ -7050,7 +7050,7 @@ 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)
|
||||
manager._create_prefixed_tools(listed_tools, server, caller=caller)
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False)
|
||||
|
|
@ -7064,23 +7064,23 @@ class TestMCPServerManager:
|
|||
return manager, proxy_logging_obj
|
||||
|
||||
@staticmethod
|
||||
def _unrestricted_auth() -> MagicMock:
|
||||
user_api_key_auth = MagicMock()
|
||||
user_api_key_auth.object_permission = None
|
||||
user_api_key_auth.object_permission_id = None
|
||||
return user_api_key_auth
|
||||
def _unrestricted_auth() -> UserAPIKeyAuth:
|
||||
return UserAPIKeyAuth(api_key="sk-test")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_tool_hands_listed_tool_description_and_schema_to_pre_call_hooks(self):
|
||||
schema = {"type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"]}
|
||||
listed = [MCPTool(name="test_tool", description="Runs the test tool", inputSchema=schema)]
|
||||
manager, proxy_logging_obj = self._manager_ready_for_call_tool(listed)
|
||||
auth = self._unrestricted_auth()
|
||||
manager, proxy_logging_obj = self._manager_ready_for_call_tool(
|
||||
listed, caller=ListedToolsCaller(user_api_key_auth=auth)
|
||||
)
|
||||
|
||||
await manager.call_tool(
|
||||
server_name="test-server",
|
||||
name="test_tool",
|
||||
arguments={"param": "value"},
|
||||
user_api_key_auth=self._unrestricted_auth(),
|
||||
user_api_key_auth=auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
|
|
@ -7091,7 +7091,8 @@ class TestMCPServerManager:
|
|||
async def test_call_tool_hands_listed_tool_metadata_to_during_call_hooks_through_real_conversion(self):
|
||||
schema = {"type": "object", "properties": {"param": {"type": "string"}}}
|
||||
listed = [MCPTool(name="test_tool", description="Runs the test tool", inputSchema=schema)]
|
||||
manager, _ = self._manager_ready_for_call_tool(listed)
|
||||
auth = UserAPIKeyAuth(api_key="sk-test")
|
||||
manager, _ = self._manager_ready_for_call_tool(listed, caller=ListedToolsCaller(user_api_key_auth=auth))
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
|
||||
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
|
||||
proxy_logging_obj.during_call_hook = AsyncMock(return_value=None)
|
||||
|
|
@ -7100,7 +7101,7 @@ class TestMCPServerManager:
|
|||
server_name="test-server",
|
||||
name="test_tool",
|
||||
arguments={"param": "value"},
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-test"),
|
||||
user_api_key_auth=auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
|
|
@ -7112,15 +7113,17 @@ class TestMCPServerManager:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_tool_passes_no_tool_metadata_when_tool_was_never_listed(self):
|
||||
auth = self._unrestricted_auth()
|
||||
manager, proxy_logging_obj = self._manager_ready_for_call_tool(
|
||||
[MCPTool(name="other_tool", description="Unrelated", inputSchema={"type": "object"})]
|
||||
[MCPTool(name="other_tool", description="Unrelated", inputSchema={"type": "object"})],
|
||||
caller=ListedToolsCaller(user_api_key_auth=auth),
|
||||
)
|
||||
|
||||
await manager.call_tool(
|
||||
server_name="test-server",
|
||||
name="test_tool",
|
||||
arguments={"param": "value"},
|
||||
user_api_key_auth=self._unrestricted_auth(),
|
||||
user_api_key_auth=auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
|
|
@ -7244,7 +7247,9 @@ class TestMCPServerManager:
|
|||
caller=ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
for_bob = manager.get_listed_tool(shared, "echo", ListedToolsCaller(user_api_key_auth=bob))
|
||||
assert for_bob is not None and for_bob.description == "everyone"
|
||||
assert for_bob is None, "keyed callers get their own slot even on servers without upstream per-user auth"
|
||||
anonymous = manager.get_listed_tool(shared, "echo")
|
||||
assert anonymous is None
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("server_kwargs", "caller_a", "caller_b"),
|
||||
|
|
@ -7371,20 +7376,24 @@ class TestMCPServerManager:
|
|||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"),
|
||||
)
|
||||
|
||||
caller: Final = ListedToolsCaller(mcp_auth_header="Bearer hdr")
|
||||
caller: Final = ListedToolsCaller(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm"), mcp_auth_header="Bearer hdr"
|
||||
)
|
||||
listed = manager.get_listed_tool(server, "turn", caller)
|
||||
assert listed is not None and listed.description == "t"
|
||||
assert manager._create_mcp_client.await_args.kwargs["mcp_auth_header"] == "Bearer hdr"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("signer", "static_headers", "shared"),
|
||||
("signer", "static_headers"),
|
||||
[
|
||||
pytest.param(MagicMock(), None, False, id="signer-mints-per-caller-authorization"),
|
||||
pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, True, id="static-authorization-wins"),
|
||||
pytest.param(None, None, True, id="no-signer-stays-shared"),
|
||||
pytest.param(MagicMock(), None, id="signer"),
|
||||
pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, id="static-authorization"),
|
||||
pytest.param(None, None, id="no-signer"),
|
||||
],
|
||||
)
|
||||
def test_jwt_signer_makes_a_shared_server_list_per_caller(self, signer, static_headers, shared):
|
||||
def test_keyed_callers_always_list_into_their_own_slot(self, signer, static_headers):
|
||||
"""The catalog is guardrail-shaped per key, so a keyed caller never reads another caller's
|
||||
listing regardless of the signer or static authorization configuration."""
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", static_headers=static_headers
|
||||
|
|
@ -7399,12 +7408,10 @@ class TestMCPServerManager:
|
|||
manager._create_prefixed_tools(
|
||||
[MCPTool(name="turn", description="alice view", inputSchema={})], server, caller=alice
|
||||
)
|
||||
for_bob = manager.get_listed_tool(server, "srv-turn", bob)
|
||||
assert manager.get_listed_tool(server, "srv-turn", bob) is None
|
||||
|
||||
if shared:
|
||||
assert for_bob is not None and for_bob.description == "alice view"
|
||||
else:
|
||||
assert for_bob is None
|
||||
for_alice = manager.get_listed_tool(server, "srv-turn", alice)
|
||||
assert for_alice is not None and for_alice.description == "alice view"
|
||||
|
||||
def test_signed_server_slot_splits_on_the_callers_key_not_only_the_user(self):
|
||||
"""Two keys sharing a user_id get different signed JWTs, so they split; the same key
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue