From cb6050b5215be82d2f0379d861d65da46f528ed1 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 12:18:26 -0700 Subject: [PATCH] fix(mcp): challenge missing upstream OAuth during initialization --- .../proxy/_experimental/mcp_server/server.py | 71 +++++++++++--- .../test_mcp_oauth_passthrough_tools.py | 68 +++++++++++++ .../mcp_server/test_mcp_server.py | 98 +++++++++++++++++++ 3 files changed, 226 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 926aead9b7a..ea4799d9712 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -34,6 +34,7 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, + _gateway_dcr_challenge, _is_mcp_admitted_user_subject, ) from litellm.proxy._experimental.mcp_server.client_allowlist import ( @@ -1698,7 +1699,42 @@ if MCP_AVAILABLE: excludes a passthrough server is not pushed into an OAuth flow for a server it will be 403'd on immediately after authentication. """ - for server_name in mcp_servers or []: + if mcp_servers is None: + allowed: Final = await operations._get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, mcp_servers=None, client_ip=client_ip + ) + eligible: Final = tuple( + server for server in allowed if allowed_server_ids is None or server.server_id in allowed_server_ids + ) + results: Final = await asyncio.gather( + *( + _raise_preemptive_401_for_unauthenticated_servers( + scope=scope, + mcp_servers=[server.alias or server.server_name or server.name], + oauth2_headers=oauth2_headers, + mcp_server_auth_headers=mcp_server_auth_headers, + user_api_key_auth=user_api_key_auth, + client_ip=client_ip, + allowed_server_ids=allowed_server_ids, + raw_headers=raw_headers, + ) + for server in eligible + ), + return_exceptions=True, + ) + for result in results: + if isinstance(result, asyncio.CancelledError): + raise result + if results and all(isinstance(result, HTTPException) and result.status_code == 401 for result in results): + 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 + ) + first: Final = results[0] + if isinstance(first, HTTPException): + raise first + return + for server_name in mcp_servers: server = operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip) if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids: # Caller's narrowed scope excludes this server — skip the @@ -2114,16 +2150,17 @@ if MCP_AVAILABLE: # from the fully-authorized server set: a passthrough server that # the active toolset excludes should not trigger an OAuth flow # for a server the caller will be 403'd on after authentication. - await _raise_preemptive_401_for_unauthenticated_servers( - scope=scope, - mcp_servers=mcp_servers, - oauth2_headers=oauth2_headers, - mcp_server_auth_headers=mcp_server_auth_headers, - user_api_key_auth=user_api_key_auth, - client_ip=_client_ip, - allowed_server_ids=toolset_allowed_server_ids, - raw_headers=raw_headers, - ) + if mcp_servers is not None: + await _raise_preemptive_401_for_unauthenticated_servers( + scope=scope, + mcp_servers=mcp_servers, + oauth2_headers=oauth2_headers, + mcp_server_auth_headers=mcp_server_auth_headers, + user_api_key_auth=user_api_key_auth, + client_ip=_client_ip, + allowed_server_ids=toolset_allowed_server_ids, + raw_headers=raw_headers, + ) # Pre-flight auth check for pass-through servers. Must run after # toolset scoping so the probe list is derived from the fully-authorized @@ -2200,6 +2237,18 @@ if MCP_AVAILABLE: consumed_messages, body = await _read_request_body_for_routing(receive) is_initialize = _is_initialize_request(body) + if is_initialize and mcp_servers is None: + await _raise_preemptive_401_for_unauthenticated_servers( + scope=scope, + mcp_servers=None, + oauth2_headers=oauth2_headers, + mcp_server_auth_headers=mcp_server_auth_headers, + user_api_key_auth=user_api_key_auth, + client_ip=_client_ip, + allowed_server_ids=toolset_allowed_server_ids, + raw_headers=raw_headers, + ) + use_stateful: Final = bool(session_id or is_initialize) target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py index 15aa089637a..a7c7e1853a0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py @@ -854,3 +854,71 @@ async def test_tool_handler_preserves_unrelated_protocol_errors( await server.mcp_server_tool_call(_mcp_request_ctx(), CallToolRequestParams(name="example", arguments={})) assert caught.value is failure execute.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("path", ("/mcp", "/github/mcp")) +@pytest.mark.parametrize("has_token", (False, True)) +@pytest.mark.parametrize("healthy_companion", (False, True)) +async def test_initialize_challenges_missing_upstream_credentials_before_creating_session( + monkeypatch: pytest.MonkeyPatch, path: str, has_token: bool, healthy_companion: bool +) -> None: + from mcp.server.streamable_http_manager import StreamableHTTPSessionManager + from litellm.proxy._experimental.mcp_server import server + from litellm.proxy._types import UserAPIKeyAuth + + github: Final = MCPServer( + server_id="github-id", name="github", alias="github", server_name="github", + url="https://github.example/mcp", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", + authorization_url="https://github.example/authorize", token_url="https://github.example/token", + client_id="registered-client", + ) + auth: Final = UserAPIKeyAuth(api_key="test-owner", user_id="test-user") + selected: Final = ["github"] if path != "/mcp" else None + monkeypatch.setattr( + server, "extract_mcp_auth_context", AsyncMock(return_value=(auth, None, selected, None, None, None)) + ) + manager: Final = server.operations.global_mcp_server_manager + public: Final = MCPServer( + server_id="public-id", name="public", alias="public", server_name="public", + url="https://public.example/mcp", transport=MCPTransport.http, auth_type=MCPAuth.none, + ) + eligible: Final = [github, public] if healthy_companion and selected is None else [github] + monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda name, **kwargs: next(s for s in eligible if s.alias == name)) + monkeypatch.setattr(manager, "has_user_oauth_token", AsyncMock(return_value=has_token)) + monkeypatch.setattr(manager, "_ensure_upstream_initialize_instructions_cached", AsyncMock()) + monkeypatch.setattr(server.operations, "_get_allowed_mcp_servers", AsyncMock(return_value=eligible)) + monkeypatch.setattr(server, "_check_passthrough_upstream_auth", AsyncMock()) + monkeypatch.setattr( + server, "session_manager_stateful", + StreamableHTTPSessionManager(app=server.server, stateless=False, json_response=True), + ) + monkeypatch.setattr( + server, "session_manager_stateless", + StreamableHTTPSessionManager(app=server.server, stateless=True, json_response=True), + ) + await server.initialize_session_managers() + try: + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://gateway") as client: + response: Final = await client.post( + path, headers={"accept": "application/json, text/event-stream"}, + json={"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": { + "protocolVersion": "2025-06-18", "capabilities": {}, + "clientInfo": {"name": "test-client", "version": "1"}, + }}, + ) + can_initialize: Final = has_token or (healthy_companion and selected is None) + assert response.status_code == (200 if can_initialize else 401), response.text + if can_initialize: + assert response.json()["result"]["serverInfo"]["name"] + assert response.headers["mcp-session-id"] + else: + assert response.headers["www-authenticate"].startswith("Bearer ") + if path == "/mcp": + assert response.headers["www-authenticate"] == ( + 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp"' + ) + assert "mcp-session-id" not in response.headers + finally: + await server.shutdown_session_managers() 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 da11b290a1d..0ba8b3f2a40 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 @@ -10690,3 +10690,101 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct context = dispatched.await_args.args[1] assert context.user_api_key_auth.user_id == "discover-caller" assert context.mcp_servers == ("allowed",) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "token_states, allowed_ids, expected_status", + ( + ((False,), None, 401), + ((False, False), None, 401), + ((False, True), None, None), + ((True, False), None, None), + ((False, False), {"server-1"}, 401), + ((False, True), {"server-1"}, None), + ((False,), set(), None), + ((), None, None), + ), +) +async def test_unified_preflight_challenges_only_when_all_authorized_servers_need_oauth( + monkeypatch: pytest.MonkeyPatch, + token_states: tuple[bool, ...], + allowed_ids: set[str] | None, + expected_status: int | None, +) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + servers: Final = tuple( + _make_oauth2_server(f"server-{index}").model_copy(update={"server_id": f"server-{index}"}) + for index in range(len(token_states)) + ) + auth: Final = UserAPIKeyAuth(api_key="test-key", user_id="reader") + lookup: Final = AsyncMock(return_value=servers) + monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", lookup) + manager: Final = mcp_operations.global_mcp_server_manager + monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda name, **kwargs: next(s for s in servers if s.alias == name)) + tokens: Final = AsyncMock(side_effect=lambda s, user: token_states[int(s.server_id.rsplit("-", 1)[1])]) + monkeypatch.setattr(manager, "has_user_oauth_token", tokens) + scope: Final[Scope] = {"type": "http", "method": "POST", "path": "/mcp", "headers": [(b"host", b"gateway")]} + request: Final = server_module._raise_preemptive_401_for_unauthenticated_servers( + scope=scope, mcp_servers=None, oauth2_headers=None, mcp_server_auth_headers=None, + user_api_key_auth=auth, client_ip="127.0.0.1", allowed_server_ids=allowed_ids, + ) + if expected_status is not None: + with pytest.raises(HTTPException) as caught: + await request + assert caught.value.status_code == expected_status + assert "www-authenticate" in {key.lower() for key in (caught.value.headers or {})} + else: + await request + lookup.assert_awaited_once_with(user_api_key_auth=auth, mcp_servers=None, client_ip="127.0.0.1") + assert {call.args[0].server_id for call in tokens.await_args_list} == { + s.server_id for s in servers if allowed_ids is None or s.server_id in allowed_ids + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", (HTTPException(status_code=503, detail="unavailable"), asyncio.CancelledError())) +async def test_unified_preflight_does_not_misclassify_discovery_failure_as_oauth( + monkeypatch: pytest.MonkeyPatch, failure: BaseException +) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + upstream: Final = _make_oauth2_server("unavailable") + monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream])) + manager: Final = mcp_operations.global_mcp_server_manager + monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream) + discovery: Final = AsyncMock(side_effect=failure) + monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", discovery) + request: Final = server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "path": "/mcp", "method": "POST", "headers": []}, + mcp_servers=None, oauth2_headers=None, mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(user_id="reader"), client_ip=None, + ) + if isinstance(failure, asyncio.CancelledError): + with pytest.raises(asyncio.CancelledError): + await request + else: + await request + discovery.assert_awaited_once_with(upstream) + + +@pytest.mark.asyncio +async def test_unified_preflight_preserves_delegated_oauth_challenge(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + upstream: Final = _make_oauth2_server("delegated", delegate_auth_to_upstream=True) + monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream])) + monkeypatch.setattr( + mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream + ) + with pytest.raises(HTTPException) as caught: + await 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, + ) + assert caught.value.status_code == 401 + assert (caught.value.headers or {})["www-authenticate"] == ( + 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/delegated"' + )