mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(mcp): preserve healthy peers during authentication preflight
This commit is contained in:
parent
0891aa4b35
commit
d9239e94b3
2 changed files with 66 additions and 3 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue