diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 17c64c0f5d3..c6bd22d1729 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1865,15 +1865,29 @@ if MCP_AVAILABLE: raw_headers=raw_headers, ) for server in eligible - ) + ), + return_exceptions=True, ) - if not results or (first := results[0]) is None or any(result is None for result in results): + failures: Final = tuple( + (server, result) for server, result in zip(eligible, results) if isinstance(result, BaseException) + ) + for server, failure in failures: + if not isinstance(failure, Exception): + raise failure + if not isinstance(failure, HTTPException) or failure.status_code != 401: + verbose_logger.warning( + "MCP authentication preflight failed for %s (%s)", server.name, type(failure).__name__ + ) + if not failures or len(failures) != len(results): return + for _, failure in failures: + if not isinstance(failure, HTTPException) or failure.status_code != 401: + raise failure if all(server.is_gateway_managed_oauth2 for server in eligible): raise _gateway_dcr_challenge( StarletteRequest(scope), get_route_relative_request_path(scope), None, invalid_token=False ) - raise first + raise failures[0][1] for server_name in mcp_servers: if ( server := operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip) 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 4117e5c9ea2..063491f802d 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 @@ -10892,3 +10892,52 @@ async def test_preflight_does_not_request_oauth_for_excluded_server(monkeypatch: ) discovery.assert_not_awaited() tokens.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("healthy_peer", (False, True)) +@pytest.mark.parametrize("failure_first", (False, True)) +@pytest.mark.parametrize("failure", (HTTPException(503, "unavailable"), RuntimeError("discovery failed"), asyncio.CancelledError())) +async def test_unified_preflight_preserves_usable_peer_during_discovery_failure( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + healthy_peer: bool, + failure_first: bool, + failure: BaseException, +) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + unavailable: Final = _make_oauth2_server("unavailable") + peer: Final = _make_oauth2_server("peer") + servers: Final = [unavailable, peer] if failure_first else [peer, unavailable] + monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=servers)) + manager: Final = mcp_operations.global_mcp_server_manager + + async def discover(server: MCPServer) -> MCPServer: + if server.server_id == unavailable.server_id: + raise failure + return server + + monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", discover) + tokens: Final = AsyncMock(return_value=healthy_peer) + monkeypatch.setattr(manager, "has_user_oauth_token", tokens) + request: Final = server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "path": "/mcp", "method": "POST", "headers": [(b"host", b"gateway")]}, + mcp_servers=None, + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(user_id="reader"), + client_ip=None, + ) + if healthy_peer and not isinstance(failure, asyncio.CancelledError): + with caplog.at_level("WARNING", logger="LiteLLM"): + await request + assert any("unavailable" in record.getMessage() for record in caplog.records) + assert str(failure) not in caplog.text + else: + with pytest.raises(type(failure)) as caught: + await request + if not isinstance(failure, asyncio.CancelledError): + assert caught.value is failure + tokens.assert_awaited_once() + assert tokens.await_args.args[0].server_id == peer.server_id