fix(mcp): challenge missing upstream OAuth during initialization

This commit is contained in:
Joshua Valluru 2026-09-28 12:18:26 -07:00
parent 97a662f2b2
commit cb6050b521
3 changed files with 226 additions and 11 deletions

View file

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

View file

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

View file

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