fix(mcp): preserve healthy peers during authentication preflight

This commit is contained in:
Joshua Valluru 2026-09-28 16:00:40 -07:00
parent 0891aa4b35
commit d9239e94b3
2 changed files with 66 additions and 3 deletions

View file

@ -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)

View file

@ -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