From cb4341b3907af6245a5679bbd8b1cb836bbc77ca Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:32:23 -0700 Subject: [PATCH] fix(mcp): isolate transport test admission state and exhaust OAuth outcomes --- litellm/proxy/_experimental/mcp_server/mcp_server_manager.py | 4 +++- tests/mcp_tests/test_proxy_mcp_e2e.py | 2 ++ 2 files changed, 5 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 490190204bc..48325a34ef9 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -48,7 +48,7 @@ from mcp.types import ( ) from mcp.types import Tool as MCPTool from pydantic import AnyUrl, BaseModel, TypeAdapter -from typing_extensions import ReadOnly +from typing_extensions import ReadOnly, assert_never import litellm from litellm._logging import verbose_logger @@ -2213,6 +2213,8 @@ class MCPServerManager: detail=f"OAuth metadata discovery {reason} for MCP server {server_ref!r}", ) + return assert_never(outcome) + async def _rejoin_oauth_metadata_discovery(self, server: MCPServer, *, retry_stale: bool) -> MCPServer: if retry_stale: return await self.ensure_oauth_metadata_discovered(server, _retry_stale=False) diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index a730f6c10ee..c23a3db9d27 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -62,9 +62,11 @@ def _clear_proxy_database_env() -> typing.Iterator[None]: async def _initialize_proxy(config_path: str) -> None: + from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager cleanup_router_config_variables() + global_mcp_server_manager.catalog = TargetCatalog(global_mcp_server_manager) await initialize(config=config_path, debug=True) for server_id, upstream in tuple(global_mcp_server_manager.registry.items()): if upstream.server_name != "math_restricted":