From ff7a59eebbddbecb2a2c322dc8a19dabce064080 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 11 Jun 2026 09:39:29 +0530 Subject: [PATCH 01/10] fix(mcp): honor server_id for REST tool calls with shared upstream URLs When multiple MCP server entries point at the same backend URL and tool name, REST /mcp-rest/tools/call now routes and applies auth from the requested server_id instead of the global unprefixed tool-name mapping. Co-authored-by: Cursor --- .../proxy/_experimental/mcp_server/server.py | 93 +++++++------ .../mcp_server/test_mcp_server.py | 131 ++++++++++++++++++ 2 files changed, 183 insertions(+), 41 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 731493b1337d..5484ed9aef2a 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -270,6 +270,8 @@ def _jsonrpc_text_has_top_level_method(text: str) -> bool: global_mcp_tool_registry, ) from litellm.proxy._experimental.mcp_server.utils import ( + is_tool_name_prefixed, + normalize_server_name, split_server_prefix_from_name, ) @@ -2483,47 +2485,56 @@ async def execute_mcp_tool( # noqa: PLR0915 None, ) - # Resolve the actual MCP server up-front so the permission check uses - # the canonical server.name even when the tool name is prefixed with a - # short ID (LITELLM_USE_SHORT_MCP_TOOL_PREFIX) that doesn't match the - # server's display name directly. - mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) - if mcp_server is None and requested_server is not None: - # REST callers may pass the raw tool name (no prefix) plus a - # ``requested_server_id``. The mapping might only contain the - # prefixed form, so retry the lookup with every known prefix of - # the requested server before treating the tool as unresolved — - # otherwise the tool_server_mismatch guard below is silently - # bypassed. - for known_prefix in iter_known_server_prefixes(requested_server): - candidate = global_mcp_server_manager._get_mcp_server_from_tool_name( - add_server_prefix_to_name(name, known_prefix) - ) - if candidate is not None: - mcp_server = candidate - break - if mcp_server is not None: - server_name = mcp_server.name - - # REST /mcp-rest/tools/call passes server_id — tool must belong to that server - if requested_server is not None: - if ( - mcp_server is not None - and mcp_server.server_id != requested_server.server_id - ): - raise HTTPException( - status_code=403, - detail={ - "error": "tool_server_mismatch", - "message": ( - f"Tool '{name}' belongs to MCP server '{mcp_server.name}' " - f"but request specified server_id for '{requested_server.name}'." - ), - }, - ) - if mcp_server is None: - mcp_server = requested_server - server_name = requested_server.name + known_server_prefixes: Set[str] = set() + for allowed_server in allowed_mcp_servers: + for known_prefix in iter_known_server_prefixes(allowed_server): + known_server_prefixes.add(normalize_server_name(known_prefix)) + + name_is_prefixed = is_tool_name_prefixed( + name, known_server_prefixes=known_server_prefixes + ) + + if requested_server is not None and not name_is_prefixed: + # REST callers may pass server_id with the upstream tool name (no + # LiteLLM prefix). Multiple MCP entries can share the same URL and + # tool name; server_id is authoritative for routing and auth. + mcp_server = requested_server + server_name = requested_server.name + else: + # Resolve from tool name (MCP JSON-RPC or prefixed REST tool names). + mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + if mcp_server is None and requested_server is not None: + for known_prefix in iter_known_server_prefixes(requested_server): + candidate = ( + global_mcp_server_manager._get_mcp_server_from_tool_name( + add_server_prefix_to_name(name, known_prefix) + ) + ) + if candidate is not None: + mcp_server = candidate + break + if mcp_server is not None: + server_name = mcp_server.name + + if requested_server is not None: + if ( + mcp_server is not None + and mcp_server.server_id != requested_server.server_id + ): + raise HTTPException( + status_code=403, + detail={ + "error": "tool_server_mismatch", + "message": ( + f"Tool '{name}' belongs to MCP server " + f"'{mcp_server.name}' but request specified " + f"server_id for '{requested_server.name}'." + ), + }, + ) + if mcp_server is None: + mcp_server = requested_server + server_name = requested_server.name # Only enforce server-level permissions when we can resolve a server if server_name: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index b6550fee6b96..9ea6aa8fa8e8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -5489,3 +5489,134 @@ async def test_create_mcp_client_sampling_enabled(): client = await manager._create_mcp_client(server=server) assert client._sampling_callback is not None + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool(): + """REST server_id + unprefixed tool name must not use global tool-name mapping.""" + from mcp.types import TextContent + + from litellm.proxy._experimental.mcp_server import server as mcp_module + + api_key_server = MCPServer( + server_id="api-key-server-id", + name="echo_api_key", + server_name="echo_api_key", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + authentication_token="abc123", + ) + oauth_server = MCPServer( + server_id="oauth-server-id", + name="echo_oauth_m2m", + server_name="echo_oauth_m2m", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + token_url="http://127.0.0.1:8080/token", + client_id="client", + client_secret="secret", + ) + + mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_name_mapping[ + "echo" + ] = oauth_server.name + + captured: dict = {} + + async def fake_handle_managed_mcp_tool(**kwargs): + captured.update(kwargs) + return mcp_module.CallToolResult( + content=[TextContent(type="text", text="ok")], + isError=False, + ) + + with ( + patch.object( + mcp_module.global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=oauth_server, + ), + patch.object( + mcp_module, + "_handle_managed_mcp_tool", + new=fake_handle_managed_mcp_tool, + ), + patch.object( + mcp_module.MCPRequestHandler, + "is_tool_allowed", + return_value=True, + ), + patch.object( + mcp_module.global_mcp_tool_registry, + "get_tool", + return_value=None, + ), + ): + await mcp_module.execute_mcp_tool( + name="echo", + arguments={"message": "hello"}, + allowed_mcp_servers=[api_key_server, oauth_server], + start_time=datetime.now(), + requested_server_id=api_key_server.server_id, + ) + + assert captured["server_name"] == "echo_api_key" + assert captured["name"] == "echo" + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_rest_prefixed_tool_still_validates_server_id(): + """Prefixed REST tool names must still match the requested server_id.""" + from litellm.proxy._experimental.mcp_server import server as mcp_module + + api_key_server = MCPServer( + server_id="api-key-server-id", + name="echo_api_key", + server_name="echo_api_key", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + authentication_token="abc123", + ) + oauth_server = MCPServer( + server_id="oauth-server-id", + name="echo_oauth_m2m", + server_name="echo_oauth_m2m", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + token_url="http://127.0.0.1:8080/token", + client_id="client", + client_secret="secret", + ) + + with ( + patch.object( + mcp_module.global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=oauth_server, + ), + patch.object( + mcp_module.MCPRequestHandler, + "is_tool_allowed", + return_value=True, + ), + patch.object( + mcp_module.global_mcp_tool_registry, + "get_tool", + return_value=None, + ), + pytest.raises(HTTPException) as exc_info, + ): + await mcp_module.execute_mcp_tool( + name="echo_oauth_m2m-echo", + arguments={"message": "hello"}, + allowed_mcp_servers=[api_key_server, oauth_server], + start_time=datetime.now(), + requested_server_id=api_key_server.server_id, + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail["error"] == "tool_server_mismatch" From 2625237a9347238f88d4a4f1f26f9015da7182fd Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 11 Jun 2026 18:47:26 +0530 Subject: [PATCH 02/10] fix(mcp): classify prefixed REST tool names against full registry Use all registered MCP server prefixes for prefix detection so unauthorized prefixed names still trigger tool_server_mismatch, and reject ambiguous hyphenated REST tool names with server_id. Co-authored-by: Cursor --- .../proxy/_experimental/mcp_server/server.py | 25 +++- .../mcp_server/test_mcp_server.py | 118 ++++++++++++++++++ 2 files changed, 138 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 5484ed9aef2a..061a0b2cd115 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -270,6 +270,7 @@ def _jsonrpc_text_has_top_level_method(text: str) -> bool: global_mcp_tool_registry, ) from litellm.proxy._experimental.mcp_server.utils import ( + MCP_TOOL_PREFIX_SEPARATOR, is_tool_name_prefixed, normalize_server_name, split_server_prefix_from_name, @@ -2485,16 +2486,30 @@ async def execute_mcp_tool( # noqa: PLR0915 None, ) - known_server_prefixes: Set[str] = set() - for allowed_server in allowed_mcp_servers: - for known_prefix in iter_known_server_prefixes(allowed_server): - known_server_prefixes.add(normalize_server_name(known_prefix)) + all_registry_prefixes: Set[str] = set() + for registry_server in global_mcp_server_manager.get_registry().values(): + for known_prefix in iter_known_server_prefixes(registry_server): + all_registry_prefixes.add(normalize_server_name(known_prefix)) name_is_prefixed = is_tool_name_prefixed( - name, known_server_prefixes=known_server_prefixes + name, known_server_prefixes=all_registry_prefixes ) if requested_server is not None and not name_is_prefixed: + if MCP_TOOL_PREFIX_SEPARATOR in name: + raise HTTPException( + status_code=400, + detail={ + "error": "ambiguous_tool_name", + "message": ( + f"Tool name '{name}' contains " + f"'{MCP_TOOL_PREFIX_SEPARATOR}' but does not match " + "a registered MCP server prefix. Pass the unprefixed " + "upstream tool name with server_id, or use the full " + "prefixed form 'server_name-tool_name'." + ), + }, + ) # REST callers may pass server_id with the upstream tool name (no # LiteLLM prefix). Multiple MCP entries can share the same URL and # tool name; server_id is authoritative for routing and auth. diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 9ea6aa8fa8e8..381d8eca46a7 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -5533,6 +5533,14 @@ async def fake_handle_managed_mcp_tool(**kwargs): ) with ( + patch.object( + mcp_module.global_mcp_server_manager, + "get_registry", + return_value={ + api_key_server.server_id: api_key_server, + oauth_server.server_id: oauth_server, + }, + ), patch.object( mcp_module.global_mcp_server_manager, "_get_mcp_server_from_tool_name", @@ -5593,6 +5601,14 @@ async def test_execute_mcp_tool_rest_prefixed_tool_still_validates_server_id(): ) with ( + patch.object( + mcp_module.global_mcp_server_manager, + "get_registry", + return_value={ + api_key_server.server_id: api_key_server, + oauth_server.server_id: oauth_server, + }, + ), patch.object( mcp_module.global_mcp_server_manager, "_get_mcp_server_from_tool_name", @@ -5620,3 +5636,105 @@ async def test_execute_mcp_tool_rest_prefixed_tool_still_validates_server_id(): assert exc_info.value.status_code == 403 assert exc_info.value.detail["error"] == "tool_server_mismatch" + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_rest_unauthorized_prefix_still_mismatches(): + """Prefixed name for a registry server the caller cannot access must 403.""" + from litellm.proxy._experimental.mcp_server import server as mcp_module + + api_key_server = MCPServer( + server_id="api-key-server-id", + name="echo_api_key", + server_name="echo_api_key", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + authentication_token="abc123", + ) + restricted_server = MCPServer( + server_id="restricted-server-id", + name="restricted_server", + server_name="restricted_server", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token="secret", + ) + + with ( + patch.object( + mcp_module.global_mcp_server_manager, + "get_registry", + return_value={ + api_key_server.server_id: api_key_server, + restricted_server.server_id: restricted_server, + }, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=restricted_server, + ), + patch.object( + mcp_module.MCPRequestHandler, + "is_tool_allowed", + return_value=True, + ), + patch.object( + mcp_module.global_mcp_tool_registry, + "get_tool", + return_value=None, + ), + pytest.raises(HTTPException) as exc_info, + ): + await mcp_module.execute_mcp_tool( + name="restricted_server-echo", + arguments={"message": "hello"}, + allowed_mcp_servers=[api_key_server], + start_time=datetime.now(), + requested_server_id=api_key_server.server_id, + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail["error"] == "tool_server_mismatch" + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_rest_ambiguous_hyphenated_name_rejected(): + """REST server_id + hyphenated name with no registry prefix must 400.""" + from litellm.proxy._experimental.mcp_server import server as mcp_module + + api_key_server = MCPServer( + server_id="api-key-server-id", + name="echo_api_key", + server_name="echo_api_key", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + authentication_token="abc123", + ) + + with ( + patch.object( + mcp_module.global_mcp_server_manager, + "get_registry", + return_value={api_key_server.server_id: api_key_server}, + ), + patch.object( + mcp_module.global_mcp_tool_registry, + "get_tool", + return_value=None, + ), + pytest.raises(HTTPException) as exc_info, + ): + await mcp_module.execute_mcp_tool( + name="text-to-speech", + arguments={"message": "hello"}, + allowed_mcp_servers=[api_key_server], + start_time=datetime.now(), + requested_server_id=api_key_server.server_id, + ) + + assert exc_info.value.status_code == 400 + assert exc_info.value.detail["error"] == "ambiguous_tool_name" From dea48f1685dbe37a546b5f4d35e04a47d0e5d697 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 12 Jun 2026 05:40:16 +0000 Subject: [PATCH 03/10] test(mcp): cover server_id fallback for unresolved prefixed REST tool names execute_mcp_tool left the prefix-retry and requested-server fallback branches uncovered, dropping diff coverage below the project target. Add a regression test for a REST call that passes server_id with a prefixed tool name that resolves to no managed tool; it must still dispatch to the server identified by server_id rather than the server named by the prefix. --- .../mcp_server/test_mcp_server.py | 77 +++++++++++++++++++ 1 file changed, 77 insertions(+) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 295f999c9fd4..ac2e6a9091c4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -5819,3 +5819,80 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details(): assert litellm_logging_obj.model_call_details["model"] == "MCP: list_pets" assert litellm_logging_obj.model == "MCP: list_pets" + + +@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.""" + from mcp.types import TextContent + + from litellm.proxy._experimental.mcp_server import server as mcp_module + + requested_server = MCPServer( + server_id="rest-target-id", + name="rest_target", + server_name="rest_target", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + authentication_token="abc123", + ) + prefix_owner = MCPServer( + server_id="prefix-owner-id", + name="known_prefix", + server_name="known_prefix", + url="http://127.0.0.1:5116/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + authentication_token="def456", + ) + + captured: dict = {} + + async def fake_handle_managed_mcp_tool(**kwargs): + captured.update(kwargs) + return mcp_module.CallToolResult( + content=[TextContent(type="text", text="ok")], + isError=False, + ) + + with ( + patch.object( + mcp_module.global_mcp_server_manager, + "get_registry", + return_value={ + requested_server.server_id: requested_server, + prefix_owner.server_id: prefix_owner, + }, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=None, + ), + patch.object( + mcp_module, + "_handle_managed_mcp_tool", + new=fake_handle_managed_mcp_tool, + ), + patch.object( + mcp_module.MCPRequestHandler, + "is_tool_allowed", + return_value=True, + ), + patch.object( + mcp_module.global_mcp_tool_registry, + "get_tool", + return_value=None, + ), + ): + await mcp_module.execute_mcp_tool( + name="known_prefix-list_things", + arguments={"message": "hello"}, + allowed_mcp_servers=[requested_server, prefix_owner], + start_time=datetime.now(), + requested_server_id=requested_server.server_id, + ) + + assert captured["server_name"] == "rest_target" + assert captured["name"] == "list_things" From 1c19e3e2150a46d6881685df162b507ee5bb4585 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 12 Jun 2026 05:45:19 +0000 Subject: [PATCH 04/10] test(mcp): scope global tool-name mapping mutation with patch.dict --- .../proxy/_experimental/mcp_server/test_mcp_server.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index ac2e6a9091c4..e2991ff03506 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -5519,10 +5519,6 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool client_secret="secret", ) - mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_name_mapping[ - "echo" - ] = oauth_server.name - captured: dict = {} async def fake_handle_managed_mcp_tool(**kwargs): @@ -5533,6 +5529,10 @@ async def fake_handle_managed_mcp_tool(**kwargs): ) with ( + patch.dict( + mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_name_mapping, + {"echo": oauth_server.name}, + ), patch.object( mcp_module.global_mcp_server_manager, "get_registry", From 8cc2390e410820ae52db208acadb7a02070c0946 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 12 Jun 2026 05:58:33 +0000 Subject: [PATCH 05/10] test(mcp): cover server_id guard on prefix-retry tool resolution The prefix-retry branch in execute_mcp_tool re-prefixes the tool name with the requested server's known prefixes when the bare lookup misses. The candidate-found path that assigns mcp_server from that lookup stayed uncovered, so codecov patch coverage remained below the diff target. Add a regression test where the re-prefixed lookup resolves a server whose server_id differs from the requested server_id; the tool_server_mismatch 403 guard must still fire instead of being silently bypassed. --- .../mcp_server/test_mcp_server.py | 67 +++++++++++++++++++ 1 file changed, 67 insertions(+) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index e2991ff03506..eb47e949f3af 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -5896,3 +5896,70 @@ async def fake_handle_managed_mcp_tool(**kwargs): assert captured["server_name"] == "rest_target" assert captured["name"] == "list_things" + + +@pytest.mark.asyncio +async def test_execute_mcp_tool_rest_prefix_retry_resolution_still_enforces_server_id(): + """A managed tool resolved via the requested server's prefix must still honor the server_id guard.""" + from litellm.proxy._experimental.mcp_server import server as mcp_module + + requested_server = MCPServer( + server_id="api-key-server-id", + name="echo_api_key", + server_name="echo_api_key", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + authentication_token="abc123", + ) + prefix_owner = MCPServer( + server_id="prefix-owner-id", + name="known_prefix", + server_name="known_prefix", + url="http://127.0.0.1:5116/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token="secret", + ) + + def resolve_only_when_requested_prefix_added(tool_name): + if tool_name == "known_prefix-echo": + return None + return prefix_owner + + with ( + patch.object( + mcp_module.global_mcp_server_manager, + "get_registry", + return_value={ + requested_server.server_id: requested_server, + prefix_owner.server_id: prefix_owner, + }, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + side_effect=resolve_only_when_requested_prefix_added, + ), + patch.object( + mcp_module.MCPRequestHandler, + "is_tool_allowed", + return_value=True, + ), + patch.object( + mcp_module.global_mcp_tool_registry, + "get_tool", + return_value=None, + ), + pytest.raises(HTTPException) as exc_info, + ): + await mcp_module.execute_mcp_tool( + name="known_prefix-echo", + arguments={"message": "hello"}, + allowed_mcp_servers=[requested_server, prefix_owner], + start_time=datetime.now(), + requested_server_id=requested_server.server_id, + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail["error"] == "tool_server_mismatch" From 40b5135abbe24b2adc505a9456790855da1fb92e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 12 Jun 2026 06:17:43 +0000 Subject: [PATCH 06/10] test(mcp): assert requested server credentials injected on cross-server REST routing --- .../mcp_server/test_mcp_server.py | 97 ++++++++++++++++++- 1 file changed, 96 insertions(+), 1 deletion(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index eb47e949f3af..a7996d96bdd0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -5574,6 +5574,92 @@ async def fake_handle_managed_mcp_tool(**kwargs): assert captured["name"] == "echo" +@pytest.mark.asyncio +async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credentials(): + """REST server_id must inject the requested server's auth, not a URL-collision peer's.""" + from mcp.types import TextContent + + from litellm.proxy._experimental.mcp_server import server as mcp_module + + requested_server = MCPServer( + server_id="requested-server-id", + name="echo_requested", + server_name="echo_requested", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + authentication_token="requested-secret", + ) + collision_server = MCPServer( + server_id="collision-server-id", + name="echo_collision", + server_name="echo_collision", + url="http://127.0.0.1:5115/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token="collision-secret", + ) + + fake_client = MagicMock() + fake_client._last_initialize_instructions = None + fake_client.call_tool = AsyncMock( + return_value=mcp_module.CallToolResult( + content=[TextContent(type="text", text="ok")], + isError=False, + ) + ) + + injected: dict = {} + + async def fake_create_mcp_client(server, **kwargs): + injected["server"] = server + return fake_client + + with ( + patch.dict( + mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_name_mapping, + {"echo": collision_server.name}, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "get_registry", + return_value={ + requested_server.server_id: requested_server, + collision_server.server_id: collision_server, + }, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "_create_mcp_client", + new=fake_create_mcp_client, + ), + patch.object( + mcp_module.MCPRequestHandler, + "is_tool_allowed", + return_value=True, + ), + patch.object( + mcp_module.global_mcp_tool_registry, + "get_tool", + return_value=None, + ), + patch("litellm.proxy.proxy_server.proxy_logging_obj", None), + ): + await mcp_module.execute_mcp_tool( + name="echo", + arguments={"message": "hello"}, + allowed_mcp_servers=[requested_server, collision_server], + start_time=datetime.now(), + requested_server_id=requested_server.server_id, + ) + + routed = injected["server"] + assert routed.server_id == requested_server.server_id + assert routed.auth_type == MCPAuth.api_key + assert routed.authentication_token == "requested-secret" + assert routed.authentication_token != collision_server.authentication_token + + @pytest.mark.asyncio async def test_execute_mcp_tool_rest_prefixed_tool_still_validates_server_id(): """Prefixed REST tool names must still match the requested server_id.""" @@ -5843,7 +5929,7 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste server_name="known_prefix", url="http://127.0.0.1:5116/mcp", transport=MCPTransport.http, - auth_type=MCPAuth.api_key, + auth_type=MCPAuth.bearer_token, authentication_token="def456", ) @@ -5897,6 +5983,15 @@ async def fake_handle_managed_mcp_tool(**kwargs): assert captured["server_name"] == "rest_target" assert captured["name"] == "list_things" + routed_server = { + requested_server.name: requested_server, + prefix_owner.name: prefix_owner, + }[captured["server_name"]] + assert routed_server.server_id == requested_server.server_id + assert routed_server.auth_type == MCPAuth.api_key + assert routed_server.authentication_token == "abc123" + assert routed_server.authentication_token != prefix_owner.authentication_token + @pytest.mark.asyncio async def test_execute_mcp_tool_rest_prefix_retry_resolution_still_enforces_server_id(): From 424a610ba3a4457ae46ec35d7e45433188062ab0 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 12 Jun 2026 06:35:03 +0000 Subject: [PATCH 07/10] perf(mcp): scan registry prefixes only when server_id is supplied --- .../proxy/_experimental/mcp_server/server.py | 17 +++++++++-------- 1 file changed, 9 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index a90201ba6a47..bdd9ff0eb99d 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2486,14 +2486,15 @@ async def execute_mcp_tool( # noqa: PLR0915 None, ) - all_registry_prefixes: Set[str] = set() - for registry_server in global_mcp_server_manager.get_registry().values(): - for known_prefix in iter_known_server_prefixes(registry_server): - all_registry_prefixes.add(normalize_server_name(known_prefix)) - - name_is_prefixed = is_tool_name_prefixed( - name, known_server_prefixes=all_registry_prefixes - ) + name_is_prefixed = False + if requested_server is not None: + all_registry_prefixes: Set[str] = set() + for registry_server in global_mcp_server_manager.get_registry().values(): + for known_prefix in iter_known_server_prefixes(registry_server): + all_registry_prefixes.add(normalize_server_name(known_prefix)) + name_is_prefixed = is_tool_name_prefixed( + name, known_server_prefixes=all_registry_prefixes + ) if requested_server is not None and not name_is_prefixed: if MCP_TOOL_PREFIX_SEPARATOR in name: From b7df649d5aba22c1fdc66a3d7f3ec881c26f97fd Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 12 Jun 2026 06:37:06 +0000 Subject: [PATCH 08/10] fix(mcp): allow hyphenated upstream tool names when REST server_id is authoritative --- .../proxy/_experimental/mcp_server/server.py | 22 +++--------- .../mcp_server/test_mcp_server.py | 35 ++++++++++++++++--- 2 files changed, 35 insertions(+), 22 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index bdd9ff0eb99d..e2065a8e1b4c 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -270,7 +270,6 @@ def _jsonrpc_text_has_top_level_method(text: str) -> bool: global_mcp_tool_registry, ) from litellm.proxy._experimental.mcp_server.utils import ( - MCP_TOOL_PREFIX_SEPARATOR, is_tool_name_prefixed, normalize_server_name, split_server_prefix_from_name, @@ -2497,25 +2496,14 @@ async def execute_mcp_tool( # noqa: PLR0915 ) if requested_server is not None and not name_is_prefixed: - if MCP_TOOL_PREFIX_SEPARATOR in name: - raise HTTPException( - status_code=400, - detail={ - "error": "ambiguous_tool_name", - "message": ( - f"Tool name '{name}' contains " - f"'{MCP_TOOL_PREFIX_SEPARATOR}' but does not match " - "a registered MCP server prefix. Pass the unprefixed " - "upstream tool name with server_id, or use the full " - "prefixed form 'server_name-tool_name'." - ), - }, - ) # REST callers may pass server_id with the upstream tool name (no - # LiteLLM prefix). Multiple MCP entries can share the same URL and - # tool name; server_id is authoritative for routing and auth. + # LiteLLM prefix). The first segment is not a registered server + # prefix, so the whole string is the upstream tool name and may + # legitimately contain the separator (e.g. "text-to-speech"). + # server_id is authoritative for routing and auth. mcp_server = requested_server server_name = requested_server.name + original_tool_name = name else: # Resolve from tool name (MCP JSON-RPC or prefixed REST tool names). mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index a7996d96bdd0..1c31f437363c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -5787,8 +5787,10 @@ async def test_execute_mcp_tool_rest_unauthorized_prefix_still_mismatches(): @pytest.mark.asyncio -async def test_execute_mcp_tool_rest_ambiguous_hyphenated_name_rejected(): - """REST server_id + hyphenated name with no registry prefix must 400.""" +async def test_execute_mcp_tool_rest_hyphenated_upstream_tool_name_routes_to_requested_server(): + """REST server_id + hyphenated upstream tool name (no registry prefix) must route, not 400.""" + from mcp.types import TextContent + from litellm.proxy._experimental.mcp_server import server as mcp_module api_key_server = MCPServer( @@ -5801,18 +5803,41 @@ async def test_execute_mcp_tool_rest_ambiguous_hyphenated_name_rejected(): authentication_token="abc123", ) + captured: dict = {} + + async def fake_handle_managed_mcp_tool(**kwargs): + captured.update(kwargs) + return mcp_module.CallToolResult( + content=[TextContent(type="text", text="ok")], + isError=False, + ) + with ( patch.object( mcp_module.global_mcp_server_manager, "get_registry", return_value={api_key_server.server_id: api_key_server}, ), + patch.object( + mcp_module.global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=None, + ), + patch.object( + mcp_module, + "_handle_managed_mcp_tool", + new=fake_handle_managed_mcp_tool, + ), + patch.object( + mcp_module.MCPRequestHandler, + "is_tool_allowed", + return_value=True, + ), patch.object( mcp_module.global_mcp_tool_registry, "get_tool", return_value=None, ), - pytest.raises(HTTPException) as exc_info, ): await mcp_module.execute_mcp_tool( name="text-to-speech", @@ -5822,8 +5847,8 @@ async def test_execute_mcp_tool_rest_ambiguous_hyphenated_name_rejected(): requested_server_id=api_key_server.server_id, ) - assert exc_info.value.status_code == 400 - assert exc_info.value.detail["error"] == "ambiguous_tool_name" + assert captured["server_name"] == "echo_api_key" + assert captured["name"] == "text-to-speech" @pytest.mark.asyncio From 162340359a6924eadd287f2529584b0c05e5b73c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 12 Jun 2026 06:43:42 +0000 Subject: [PATCH 09/10] perf(mcp): skip registry prefix scan for separator-free REST tool names --- litellm/proxy/_experimental/mcp_server/server.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index e2065a8e1b4c..746fc4e7d3fa 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -270,6 +270,7 @@ def _jsonrpc_text_has_top_level_method(text: str) -> bool: global_mcp_tool_registry, ) from litellm.proxy._experimental.mcp_server.utils import ( + MCP_TOOL_PREFIX_SEPARATOR, is_tool_name_prefixed, normalize_server_name, split_server_prefix_from_name, @@ -2486,7 +2487,7 @@ async def execute_mcp_tool( # noqa: PLR0915 ) name_is_prefixed = False - if requested_server is not None: + if requested_server is not None and MCP_TOOL_PREFIX_SEPARATOR in name: all_registry_prefixes: Set[str] = set() for registry_server in global_mcp_server_manager.get_registry().values(): for known_prefix in iter_known_server_prefixes(registry_server): From 8acfb9610a6d856f2787c770af43ac7ef185ec23 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 12 Jun 2026 07:04:21 +0000 Subject: [PATCH 10/10] test(http_handler): drop httpbin dependence from per-request timeout test The per-request timeout test posted to https://httpbin.org/delay/10 and asserted a Timeout was raised. httpbin's free /delay endpoint intermittently returns 503 even when the /get reachability guard succeeds, so local_testing_part1 flaked on that 503 instead of the expected timeout (failed identically across an initial run and a rerun-from-failed). Serve the slow response from a local ThreadingHTTPServer so the timeout fires deterministically with no third-party network dependence. --- .../test_azure_anthropic_sync_post.py | 38 ++++++++++++++----- 1 file changed, 28 insertions(+), 10 deletions(-) diff --git a/tests/local_testing/test_azure_anthropic_sync_post.py b/tests/local_testing/test_azure_anthropic_sync_post.py index 5ceb9ae3ed9e..578a309d7ed7 100644 --- a/tests/local_testing/test_azure_anthropic_sync_post.py +++ b/tests/local_testing/test_azure_anthropic_sync_post.py @@ -2,8 +2,9 @@ ``_get_httpx_client`` + ``HTTPHandler.post`` (same pattern as Azure Anthropic sync path: ``_get_httpx_client(params={"timeout": ...})`` then ``post(..., timeout=...)``). -Uses https://httpbin.org/delay/10 with ``timeout=5`` — the handler must raise :class:`~litellm.exceptions.Timeout` -before the 10s delay completes. Skips if httpbin is unreachable. +A local server stalls longer than the per-request ``timeout`` but well under the client +default, so the handler must raise :class:`~litellm.exceptions.Timeout` from the per-request +override rather than completing under the (much larger) client default. Lives under ``local_testing`` (not ``make test-unit``). """ @@ -11,8 +12,10 @@ import json import os import sys +import threading +import time +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer -import httpx import pytest sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../.."))) @@ -20,25 +23,40 @@ from litellm.exceptions import Timeout as LitellmTimeout from litellm.llms.custom_httpx.http_handler import _get_httpx_client -_HTTPBIN_DELAY_S = 10 -_PER_REQUEST_TIMEOUT_S = 5.0 +_SERVER_DELAY_S = 5 +_PER_REQUEST_TIMEOUT_S = 1.0 _CLIENT_DEFAULT_TIMEOUT_S = 60.0 +class _SlowHandler(BaseHTTPRequestHandler): + def do_POST(self): + time.sleep(_SERVER_DELAY_S) + try: + self.send_response(200) + self.end_headers() + self.wfile.write(b"{}") + except OSError: + pass + + def log_message(self, *args): + pass + + def test_post_delay_exceeds_per_request_timeout_raises(): - try: - httpx.get("https://httpbin.org/get", timeout=5.0) - except Exception as e: - pytest.skip(f"httpbin.org unreachable: {e}") + server = ThreadingHTTPServer(("127.0.0.1", 0), _SlowHandler) + threading.Thread(target=server.serve_forever, daemon=True).start() + host, port = server.server_address handler = _get_httpx_client(params={"timeout": _CLIENT_DEFAULT_TIMEOUT_S}) try: with pytest.raises(LitellmTimeout): handler.post( - f"https://httpbin.org/delay/{_HTTPBIN_DELAY_S}", + f"http://{host}:{port}/delay", headers={"content-type": "application/json"}, data=json.dumps({"model": "claude", "messages": []}), timeout=_PER_REQUEST_TIMEOUT_S, ) finally: handler.close() + server.shutdown() + server.server_close()