fix(mcp): simplify OAuth checks and preserve mixed-failure challenges

This commit is contained in:
Joshua Valluru 2026-09-28 15:45:32 -07:00
parent 8c7a4372ee
commit 0891aa4b35
7 changed files with 357 additions and 259 deletions

View file

@ -73,7 +73,7 @@ def listing_auth_error(outcomes: Mapping[str, ServerOutcome]) -> MCPUpstreamAuth
for name, outcome in outcomes.items()
if isinstance(outcome, ServerListFault) and outcome.tag in ("auth_required", "forbidden")
)
if not blocked or len(blocked) != len(outcomes):
if not blocked or any(isinstance(outcome, ServerListOk) for outcome in outcomes.values()):
return None
name, outcome = next((entry for entry in blocked if entry[1].tag == "auth_required"), blocked[0])
return MCPUpstreamAuthError(
@ -132,17 +132,22 @@ def raise_classified_list_failure(
a classified fault. Every fetch site delegates here so the two channels cannot drift apart per
call site. ``suppress_challenge`` is for dcr_bridge servers, whose upstream challenge points
clients at the wrong protected-resource metadata and must never relay."""
auth: Final = upstream_auth_challenge(exc)
auth: Final = upstream_auth_error(exc, server_name, suppress_challenge=suppress_challenge)
if auth is not None:
status_code, challenge = auth
raise MCPUpstreamAuthError(
status_code=status_code,
www_authenticate=None if suppress_challenge else challenge,
server_name=server_name,
) from exc
raise auth from exc
raise MCPServerListError(classify_list_exception(exc), server_name) from exc
def upstream_auth_error(
exc: BaseException, server_name: str, *, suppress_challenge: bool = False
) -> MCPUpstreamAuthError | None:
auth: Final = upstream_auth_challenge(exc)
if auth is None:
return None
status_code, challenge = auth
return MCPUpstreamAuthError(status_code, None if suppress_challenge else challenge, server_name)
def classify_list_exception(exc: BaseException) -> ServerListFault:
"""Classify a per-server listing failure into exactly one outcome. Total: an exception this
function cannot recognize is the gateway's own fault (``internal``), never a re-raise."""

View file

@ -87,6 +87,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
ServerListFault,
raise_classified_list_failure,
upstream_auth_challenge,
upstream_auth_error,
)
from litellm.proxy._experimental.mcp_server.mcp_debug import describe_upstream_http_failure, record_auth_resolution
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
@ -4632,9 +4633,7 @@ class MCPServerManager:
return self._create_prefixed_prompts(items, server, add_prefix=add_prefix)
except Exception as error:
verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, error)
if upstream_auth_challenge(error) is not None:
raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge)
return []
raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge)
async def get_resources_from_server(
self,
@ -4680,9 +4679,7 @@ class MCPServerManager:
return self._create_prefixed_resources(items, server, add_prefix=add_prefix)
except Exception as error:
verbose_logger.warning("Failed to get resources from server %s: %s", server.name, error)
if upstream_auth_challenge(error) is not None:
raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge)
return []
raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge)
async def get_resource_templates_from_server(
self,
@ -4728,9 +4725,7 @@ class MCPServerManager:
return self._create_prefixed_resource_templates(items, server, add_prefix=add_prefix)
except Exception as error:
verbose_logger.warning("Failed to get resource_templates from server %s: %s", server.name, error)
if upstream_auth_challenge(error) is not None:
raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge)
return []
raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge)
async def read_resource_from_server(
self,
@ -4769,12 +4764,9 @@ class MCPServerManager:
return await client.read_resource(url)
except Exception as exc:
auth_failure: Final = upstream_auth_challenge(exc)
auth_failure: Final = upstream_auth_error(exc, server.name, suppress_challenge=server.is_dcr_bridge)
if auth_failure is not None:
status_code, challenge = auth_failure
raise MCPUpstreamAuthError(
status_code, None if server.is_dcr_bridge else challenge, server.name
) from exc
raise auth_failure from exc
raise
async def get_prompt_from_server(
@ -4819,12 +4811,9 @@ class MCPServerManager:
)
return await client.get_prompt(get_prompt_request_params)
except Exception as exc:
auth_failure: Final = upstream_auth_challenge(exc)
auth_failure: Final = upstream_auth_error(exc, server.name, suppress_challenge=server.is_dcr_bridge)
if auth_failure is not None:
status_code, challenge = auth_failure
raise MCPUpstreamAuthError(
status_code, None if server.is_dcr_bridge else challenge, server.name
) from exc
raise auth_failure from exc
raise
@staticmethod

View file

@ -1678,6 +1678,151 @@ if MCP_AVAILABLE:
)
return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id})
async def _server_auth_challenge(
configured_server: MCPServer,
server_name: str,
scope: Scope,
mcp_servers: list[str],
oauth2_headers: dict[str, str] | None,
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
user_api_key_auth: UserAPIKeyAuth | None,
client_ip: str | None,
raw_headers: Mapping[str, str] | None,
) -> HTTPException | None:
try:
if configured_server.auth_type == MCPAuth.oauth2 and configured_server.oauth2_flow == "client_credentials":
return None
server: Final = await operations.global_mcp_server_manager.ensure_oauth_metadata_discovered(
configured_server
)
if server.auth_type == MCPAuth.oauth2:
if MCPServerManager.effective_oauth2_flow(server) == "client_credentials":
return None
if getattr(server, "delegate_auth_to_upstream", False) is not True:
if await operations.global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth):
return None
if _is_mcp_admitted_user_subject(user_api_key_auth):
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
},
)
request: Final = StarletteRequest(scope)
base_url: Final = get_request_base_url(request)
_path: Final = get_route_relative_request_path(scope)
as_metadata_root: Final = (
f"{base_url}/.well-known/oauth-authorization-server{well_known_root_suffix()}"
)
as_url: Final = (
f"{as_metadata_root}/mcp/{server_name}"
if _path.startswith(f"/mcp/{server_name}")
else f"{as_metadata_root}/{server_name}"
)
authorization_uri: Final = f'Bearer authorization_uri="{as_url}"'
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": authorization_uri},
)
if not oauth2_headers:
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(scope=scope, server_name=server_name)
},
)
return None
if server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers:
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
raise_token_exchange_challenge,
)
from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports proxy utils
get_request_root_path,
)
raise_token_exchange_challenge(server, root_path=get_request_root_path())
if len(mcp_servers) == 1 and server.server_id in frozenset(
allowed.server_id
for allowed in await operations._get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
)
):
await operations.global_mcp_server_manager.preflight_token_exchange(
server=server,
oauth2_headers=oauth2_headers,
user_api_key_auth=user_api_key_auth,
raw_headers=raw_headers,
)
if server.is_oauth_passthrough and not operations._client_has_passthrough_authorization(
server, oauth2_headers, mcp_server_auth_headers
):
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(scope=scope, server_name=server_name)
},
)
if (
server.is_oauth_delegate
and len(mcp_servers) == 1
and _get_forwarded_auth_from_scope(scope) is None
and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers)
):
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(scope=scope, server_name=server_name)
},
)
if (
server.is_true_passthrough
and len(mcp_servers) == 1
and not _scope_has_authorization_header(scope)
and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers)
):
if server.is_dcr_bridge:
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
},
)
upstream_status, upstream_www_authenticate = await _probe_upstream_auth(server.url or "", "")
if upstream_status == 401 and upstream_www_authenticate:
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": upstream_www_authenticate},
)
except HTTPException as exc:
if exc.status_code != 401:
raise
return exc
return None
async def _raise_preemptive_401_for_unauthenticated_servers(
scope: Scope,
mcp_servers: list[str] | None,
@ -1708,234 +1853,48 @@ if MCP_AVAILABLE:
)
results: Final = await asyncio.gather(
*(
_raise_preemptive_401_for_unauthenticated_servers(
_server_auth_challenge(
configured_server=server,
server_name=server.alias or server.server_name or server.name,
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
if not results or (first := results[0]) is None or any(result is None for result in results):
return
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
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
# preemptive challenge and let downstream authorization
# return 403.
continue
if server is not None and server.auth_type == MCPAuth.oauth2 and server.oauth2_flow == "client_credentials":
# Stamped M2M: the challenge decision below never reads discovered
# metadata, so deferred-discovery failures must not 503 this loop.
# Unstamped rows stay on the discover-first path because filling
# authorization_url/token_url can change their inferred flow.
continue
if server is not None:
server = await operations.global_mcp_server_manager.ensure_oauth_metadata_discovered(server)
if server and server.auth_type == MCPAuth.oauth2:
# The challenge decision is per oauth2 sub-mode, not per header:
# gateway-managed modes (M2M and interactive authorization_code)
# never receive a client-supplied upstream token, so a bearer in
# Authorization is a LiteLLM key (surfaced here as oauth2_headers)
# and must not suppress the challenge. Only the delegate mode
# treats a present bearer as the upstream token. The sub-mode is
# resolved the same way egress resolves it, via
# effective_oauth2_flow: an unstamped (null oauth2_flow) row with
# the M2M shape resolves to client_credentials, so the bare
# has_client_credentials column is never trusted here.
if MCPServerManager.effective_oauth2_flow(server) == "client_credentials":
# M2M: the gateway mints its own token at egress from the
# stored client credentials, so there is nothing to challenge.
continue
if getattr(server, "delegate_auth_to_upstream", False) is not True:
# Gateway-managed interactive (authorization_code): the only
# thing that authorizes egress is a stored per-user token, so
# challenge whenever one is absent, regardless of any bearer.
# The v2 resolver owns the existence check, so every
# authorization_code resolution (egress and this discovery
# challenge) runs through it. A keyless admitted subject is
# challenged with the per-server resource_metadata (whose
# authorization server is the gateway itself, vaulting via the
# authorize interlude); the per-server relay advertised below
# cannot vault without a litellm key on its token request.
if await operations.global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth):
continue
if _is_mcp_admitted_user_subject(user_api_key_auth):
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
},
)
request = StarletteRequest(scope)
base_url = get_request_base_url(request)
_path = get_route_relative_request_path(scope)
# Pick the well-known AS-metadata form that matches the inbound route
# so strict RFC 9728 §3.2 clients can resolve it correctly.
as_metadata_root = f"{base_url}/.well-known/oauth-authorization-server{well_known_root_suffix()}"
if _path.startswith(f"/mcp/{server_name}"):
_as_url = f"{as_metadata_root}/mcp/{server_name}"
else:
_as_url = f"{as_metadata_root}/{server_name}"
authorization_uri = f'Bearer authorization_uri="{_as_url}"'
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": authorization_uri},
)
if not oauth2_headers:
# Delegate-auth servers run upstream PKCE: a present bearer is
# the upstream token, so only challenge when it is absent, with
# the proxied resource_metadata (RFC 9728), not the gateway
# authorization_uri above which would authorize against the
# gateway instead of the upstream IdP.
www_authenticate = get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": www_authenticate},
)
# Delegate server with a bearer present: it is the upstream token,
# so admit the session and move to the next target. Every oauth2
# sub-mode is terminal here (continue or raise) so no oauth2 server
# reaches the token_exchange / pass-through blocks below.
continue
# token_exchange (OBO): the caller supplied no subject token. Challenge at connect
# (transport level, where WWW-Authenticate survives) with the RFC 9728 resource_metadata
# so the client discovers the IdP, SSOs, and retries with a subject token, which LiteLLM
# then exchanges. A tool-call-time 401 would be wrapped into a JSON-RPC error and the
# header lost, so the discovery flow needs this pre-emptive challenge.
if server and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers:
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
raise_token_exchange_challenge,
)
from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports proxy utils
get_request_root_path,
)
raise_token_exchange_challenge(server, root_path=get_request_root_path())
# Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run
# the exchange here at the transport edge, so a rejected subject raises the RFC 9728
# challenge and any other failure its public status, instead of the session opening and
# list_tools masking it as an empty tool list. The manager owns which modes pre-flight
# and what each mints from. Gated to single-server routes the key may reach; the
# multi-server aggregate keeps absorbing per-server auth failures so one bad server
# cannot 401 the whole connect.
if (
server
and len(mcp_servers or []) == 1
and server.server_id
in frozenset(
allowed.server_id
for allowed in await operations._get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
)
)
):
await operations.global_mcp_server_manager.preflight_token_exchange(
server=server,
server := operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip)
) is None:
continue
if allowed_server_ids is not None and server.server_id not in allowed_server_ids:
continue
if (
challenge := await _server_auth_challenge(
configured_server=server,
server_name=server_name,
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,
raw_headers=raw_headers,
)
# Pass-through OAuth: when the admin has opted a server into
# forwarding the client's bearer token (is_oauth_passthrough) and
# the client hasn't supplied one, fail fast with 401 and point
# them at the gateway's oauth-protected-resource well-known URL.
# That endpoint proxies the upstream's metadata so the client
# kicks off OAuth against the real upstream IdP, not the gateway.
if (
server
and server.is_oauth_passthrough
and not operations._client_has_passthrough_authorization(
server, oauth2_headers, mcp_server_auth_headers
)
):
www_authenticate = get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": www_authenticate},
)
if (
server
and server.is_oauth_delegate
and len(mcp_servers or []) == 1
and _get_forwarded_auth_from_scope(scope) is None
and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers)
):
www_authenticate = get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": www_authenticate},
)
if (
server
and server.is_true_passthrough
and len(mcp_servers or []) == 1
and not _scope_has_authorization_header(scope)
and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers)
):
if server.is_dcr_bridge:
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
},
)
upstream_status, upstream_www_authenticate = await _probe_upstream_auth(server.url or "", "")
if upstream_status == 401 and upstream_www_authenticate:
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": upstream_www_authenticate},
)
) is not None:
raise challenge
def _get_authorization_header_from_scope(scope: Scope) -> str | None:
"""First ``Authorization`` header value in the ASGI scope, or None."""

View file

@ -261,16 +261,18 @@ def test_pure_non_auth_response_still_classifies_upstream_error():
"outcomes,expected_status",
(
({}, None),
({"timeout": ServerListFault(tag="timeout")}, None),
({"denied": ServerListFault(tag="forbidden", status_code=403), "timeout": ServerListFault(tag="timeout")}, 403),
({"empty": ServerListOk(tool_count=0)}, None),
({"auth": ServerListFault(tag="auth_required", status_code=401)}, 401),
({"denied": ServerListFault(tag="forbidden", status_code=403)}, 403),
({"denied": ServerListFault(tag="forbidden", status_code=403), "auth": ServerListFault(tag="auth_required", status_code=401)}, 401),
({"auth": ServerListFault(tag="auth_required", status_code=401), "empty": ServerListOk(tool_count=0)}, None),
({"auth": ServerListFault(tag="auth_required", status_code=401), "healthy": ServerListOk(tool_count=2)}, None),
({"auth": ServerListFault(tag="auth_required", status_code=401), "timeout": ServerListFault(tag="timeout")}, None),
({"auth": ServerListFault(tag="auth_required", status_code=401), "timeout": ServerListFault(tag="timeout")}, 401),
),
)
def test_listing_auth_failure_requires_every_server_to_be_blocked(
def test_listing_auth_failure_requires_no_successful_server(
outcomes: dict[str, ServerListOk | ServerListFault], expected_status: int | None
) -> None:
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import listing_auth_error

View file

@ -735,14 +735,15 @@ async def test_optional_listing_propagates_auth_without_discarding_healthy_serve
),
)
@pytest.mark.parametrize("status", (401, 403))
@pytest.mark.parametrize("dcr_bridge", (False, True))
async def test_manager_preserves_auth_failures_for_prompts_and_resources(
operation: str, status: int, monkeypatch: pytest.MonkeyPatch
operation: str, status: int, dcr_bridge: bool, monkeypatch: pytest.MonkeyPatch
) -> None:
from fastapi import HTTPException
from pydantic import AnyUrl
manager: Final = MCPServerManager()
upstream: Final = _http_server("upstream", "upstream")
upstream: Final = _http_server("upstream", "upstream", auth_type=MCPAuth.oauth_delegate, dcr_bridge=dcr_bridge)
create: Final = AsyncMock(side_effect=HTTPException(status, headers={"WWW-Authenticate": "Bearer"}))
monkeypatch.setattr(manager, "_create_mcp_client", create)
kwargs: Final = (
@ -755,7 +756,7 @@ async def test_manager_preserves_auth_failures_for_prompts_and_resources(
with pytest.raises(MCPUpstreamAuthError) as failure:
await getattr(manager, operation)(server=upstream, user_api_key_auth=None, **kwargs)
assert failure.value.status_code == status
assert failure.value.www_authenticate == "Bearer"
assert failure.value.www_authenticate == (None if dcr_bridge else "Bearer")
assert failure.value.server_name == "upstream"
create.assert_awaited_once()
@ -773,8 +774,9 @@ async def test_optional_listing_preserves_cancellation() -> None:
@pytest.mark.asyncio
@pytest.mark.parametrize("operation", ("get_prompt_from_server", "read_resource_from_server"))
@pytest.mark.parametrize("extra_headers", (None, {"x-forwarded": "caller"}))
async def test_prompt_and_resource_calls_preserve_static_headers_and_non_auth_failures(
operation: str, monkeypatch: pytest.MonkeyPatch
operation: str, extra_headers: dict[str, str] | None, monkeypatch: pytest.MonkeyPatch
) -> None:
from pydantic import AnyUrl
@ -789,9 +791,9 @@ async def test_prompt_and_resource_calls_preserve_static_headers_and_non_auth_fa
else {"url": AnyUrl("https://example.com/resource")}
)
with pytest.raises(RuntimeError) as caught:
await getattr(manager, operation)(server=upstream, user_api_key_auth=None, **kwargs)
await getattr(manager, operation)(server=upstream, user_api_key_auth=None, extra_headers=extra_headers, **kwargs)
assert caught.value is failure
assert create.await_args.kwargs["extra_headers"] == {"x-upstream": "configured"}
assert create.await_args.kwargs["extra_headers"] == {**(extra_headers or {}), "x-upstream": "configured"}
create.assert_awaited_once()
@ -922,3 +924,26 @@ async def test_initialize_challenges_missing_upstream_credentials_before_creatin
assert "mcp-session-id" not in response.headers
finally:
await server.shutdown_session_managers()
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ("prompts", "resources", "resource_templates"))
async def test_optional_listing_challenges_auth_when_other_server_times_out(
kind: str, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.proxy._experimental.mcp_server import operations
blocked: Final = _http_server("blocked", "blocked")
unavailable: Final = _http_server("unavailable", "unavailable")
manager: Final = MCPServerManager()
create: Final = AsyncMock(side_effect=[MCPUpstreamAuthError(401, "Bearer", "blocked"), TimeoutError()])
monkeypatch.setattr(manager, "_create_mcp_client", create)
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
monkeypatch.setattr(operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None)))
monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[blocked, unavailable]))
with pytest.raises(MCPUpstreamAuthError) as caught:
await getattr(operations, f"_list_mcp_{kind}")()
assert caught.value.status_code == 401
assert caught.value.www_authenticate == "Bearer"
assert caught.value.server_name == "blocked"
assert create.await_count == 2

View file

@ -2805,7 +2805,7 @@ async def test_initialize_request_tracks_active_session_after_response_header():
patch( # test-quality-ok: registry is empty in unit tests; key owns one server
"litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[MagicMock()],
return_value=[MCPServer(server_id="available", name="available", transport=MCPTransport.http)],
),
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
@ -2958,7 +2958,7 @@ async def test_initialize_request_records_client_name_in_gateway_sessions_report
patch( # test-quality-ok: registry is empty in unit tests; key owns one server
"litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[MagicMock()],
return_value=[MCPServer(server_id="available", name="available", transport=MCPTransport.http)],
),
patch( # test-quality-ok: session manager init is a module-level flag; the suite's only seam
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
@ -10744,7 +10744,7 @@ async def test_unified_preflight_challenges_only_when_all_authorized_servers_nee
@pytest.mark.asyncio
@pytest.mark.parametrize("failure", (HTTPException(status_code=503, detail="unavailable"), asyncio.CancelledError()))
@pytest.mark.parametrize("failure", (HTTPException(status_code=503, detail="unavailable"), RuntimeError("discovery failed"), asyncio.CancelledError()))
async def test_unified_preflight_does_not_misclassify_discovery_failure_as_oauth(
monkeypatch: pytest.MonkeyPatch, failure: BaseException
) -> None:
@ -10761,11 +10761,10 @@ async def test_unified_preflight_does_not_misclassify_discovery_failure_as_oauth
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:
with pytest.raises(type(failure)) as caught:
await request
if not isinstance(failure, asyncio.CancelledError):
assert caught.value is failure
discovery.assert_awaited_once_with(upstream)
@ -10788,3 +10787,108 @@ async def test_unified_preflight_preserves_delegated_oauth_challenge(monkeypatch
assert (caught.value.headers or {})["www-authenticate"] == (
'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/delegated"'
)
@pytest.mark.asyncio
@pytest.mark.parametrize("unified", (False, True))
@pytest.mark.parametrize(
"auth_type,bridge,upstream_status,expected_challenge",
(
(MCPAuth.none, False, 401, "gateway"),
(MCPAuth.oauth_delegate, False, 401, "gateway"),
(MCPAuth.true_passthrough, True, 401, "gateway"),
(MCPAuth.true_passthrough, False, 401, "upstream"),
(MCPAuth.true_passthrough, False, 200, None),
),
)
async def test_preflight_preserves_client_forwarded_auth_challenges(
monkeypatch: pytest.MonkeyPatch,
unified: bool,
auth_type: MCPAuth,
bridge: bool,
upstream_status: int,
expected_challenge: str | None,
) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
upstream: Final = _client_forwarded_mode_server("forwarded", auth_type).model_copy(
update={
"dcr_bridge": bridge,
"oauth_passthrough": auth_type == MCPAuth.none,
"extra_headers": ["Authorization"] if auth_type == MCPAuth.none else None,
}
)
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)
probe: Final = AsyncMock(return_value=(upstream_status, "Bearer realm=upstream"))
monkeypatch.setattr(server_module, "_probe_upstream_auth", probe)
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 if unified else [upstream.name],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(user_id="reader"),
client_ip=None,
)
if expected_challenge is None:
await request
else:
with pytest.raises(HTTPException) as caught:
await request
assert caught.value.status_code == 401
expected: Final = (
'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/forwarded"'
if expected_challenge == "gateway"
else "Bearer realm=upstream"
)
assert (caught.value.headers or {})["www-authenticate"] == expected
assert probe.await_count == int(auth_type == MCPAuth.true_passthrough and not bridge)
@pytest.mark.asyncio
async def test_preflight_gateway_subject_uses_server_resource_metadata(monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
upstream: Final = _make_oauth2_server("managed")
auth: Final = UserAPIKeyAuth(user_id="reader")
auth.mcp_admitted_user_subject = True
manager: Final = mcp_operations.global_mcp_server_manager
monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream)
monkeypatch.setattr(manager, "has_user_oauth_token", AsyncMock(return_value=False))
with pytest.raises(HTTPException) as caught:
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "path": "/mcp/managed", "method": "POST", "headers": [(b"host", b"gateway")]},
mcp_servers=["managed"],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=auth,
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/managed"'
)
@pytest.mark.asyncio
async def test_preflight_does_not_request_oauth_for_excluded_server(monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
upstream: Final = _make_oauth2_server("excluded")
manager: Final = mcp_operations.global_mcp_server_manager
monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream)
discovery: Final = AsyncMock(return_value=upstream)
tokens: Final = AsyncMock(return_value=False)
monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", discovery)
monkeypatch.setattr(manager, "has_user_oauth_token", tokens)
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "path": "/mcp/excluded", "method": "POST", "headers": []},
mcp_servers=["excluded"],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(user_id="reader"),
client_ip=None,
allowed_server_ids={"other"},
)
discovery.assert_not_awaited()
tokens.assert_not_awaited()

View file

@ -1842,7 +1842,11 @@ class TestMCPServerManager:
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None, **kwargs
): # pragma: no cover - helper
captured["subject_token"] = subject_token
return AsyncMock()
return AsyncMock(
discovery_auth_fingerprint=AsyncMock(return_value="test-credential-hash"),
list_prompts=AsyncMock(return_value=[]),
list_resources=AsyncMock(return_value=[]),
)
manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client)
manager._fetch_tools_with_timeout = AsyncMock(return_value=[])
@ -2630,7 +2634,11 @@ class TestMCPServerManager:
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None, **kwargs
): # pragma: no cover - helper
captured["subject_token"] = subject_token
return AsyncMock()
return AsyncMock(
discovery_auth_fingerprint=AsyncMock(return_value="test-credential-hash"),
list_prompts=AsyncMock(return_value=[]),
list_resources=AsyncMock(return_value=[]),
)
manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client)
await call(manager)
@ -13143,6 +13151,7 @@ class TestLitellmAdmissionKeyIsNeverTheSubjectToken:
client: Final = AsyncMock()
client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False))
client.list_prompts = AsyncMock(return_value=[])
client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash")
client.read_resource = AsyncMock(return_value=ReadResourceResult(contents=[]))
manager._create_mcp_client = AsyncMock(return_value=client)
return manager
@ -13922,8 +13931,12 @@ async def test_discovery_cache_empty_results_and_failures(kind: str, outcome: st
"templates": manager.get_resource_templates_from_server,
}[kind]
with _mcp_upstream(upstream.respond):
assert await operation(_discovery_server(), None) == []
assert await operation(_discovery_server(), None) == []
for _ in range(2):
if outcome == "failure":
with pytest.raises(MCPServerListError, match="discovery"):
await operation(_discovery_server(), None)
else:
assert await operation(_discovery_server(), None) == []
assert upstream.initializes == (2 if outcome == "failure" else 1)
if outcome == "failure":
upstream.outcome = "supported"
@ -13943,7 +13956,8 @@ async def test_discovery_cache_retries_failed_pagination_before_caching_complete
"templates": manager.get_resource_templates_from_server,
}[kind]
with _mcp_upstream(upstream.respond):
assert await operation(_discovery_server(), None) == []
with pytest.raises(MCPServerListError, match="discovery"):
await operation(_discovery_server(), None)
assert upstream.initializes == 1
upstream.outcome = "paged"
recovered: Final = await operation(_discovery_server(), None)