diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 6cc2920c413..1dba4bb7c6c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -5,6 +5,7 @@ import json import logging import os import sys +from collections.abc import Callable from datetime import datetime from pathlib import Path from typing import Any, Dict, Final, Literal, Optional @@ -4074,15 +4075,23 @@ class TestMCPServerManager: @pytest.mark.asyncio @pytest.mark.parametrize( - ("manager_method", "client_method"), + ("manager_method", "client_method", "make_item"), [ - ("get_prompts_from_server", "list_prompts"), - ("get_resources_from_server", "list_resources"), - ("get_resource_templates_from_server", "list_resource_templates"), + ("get_prompts_from_server", "list_prompts", lambda name: Prompt(name=name)), + ("get_resources_from_server", "list_resources", lambda name: Resource(uri=f"demo://{name}", name=name)), + ( + "get_resource_templates_from_server", + "list_resource_templates", + lambda name: ResourceTemplate(uriTemplate=f"demo://{name}/{{id}}", name=name), + ), ], ) async def test_catalog_discovery_cache_survives_jwt_signer_and_isolates_users( - self, manager_method: str, client_method: str, monkeypatch: pytest.MonkeyPatch + self, + manager_method: str, + client_method: str, + make_item: Callable[[str], Prompt | Resource | ResourceTemplate], + monkeypatch: pytest.MonkeyPatch, ) -> None: """The MCPJWTSigner mints a fresh token on every listing, so the token itself must not be part of the discovery cache key; the user it was signed for must be, since the upstream sees that identity.""" @@ -4112,18 +4121,18 @@ class TestMCPServerManager: transport=MCPTransport.http, ) mock_client = AsyncMock() - upstream_call = AsyncMock(return_value=[]) - setattr(mock_client, client_method, upstream_call) + setattr(mock_client, client_method, AsyncMock(side_effect=[[make_item("first")], [make_item("second")]])) mock_client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash") alice = UserAPIKeyAuth(api_key="sk-alice", user_id="alice") bob = UserAPIKeyAuth(api_key="sk-bob", user_id="bob") with patch.object(manager, "_create_mcp_client", AsyncMock(return_value=mock_client)): - await getattr(manager, manager_method)(server, user_api_key_auth=alice) - await getattr(manager, manager_method)(server, user_api_key_auth=alice) - assert upstream_call.await_count == 1 - await getattr(manager, manager_method)(server, user_api_key_auth=bob) - assert upstream_call.await_count == 2 + alice_first = await getattr(manager, manager_method)(server, user_api_key_auth=alice, add_prefix=False) + alice_second = await getattr(manager, manager_method)(server, user_api_key_auth=alice, add_prefix=False) + bob_first = await getattr(manager, manager_method)(server, user_api_key_auth=bob, add_prefix=False) + assert [item.name for item in alice_first] == ["first"] + assert alice_second == alice_first + assert [item.name for item in bob_first] == ["second"] @pytest.mark.asyncio async def test_read_resource_from_server_success(self):