diff --git a/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py index 0059498e60a..a71e0d10ec9 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py +++ b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py @@ -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.""" diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 375cf9fc200..9ff9dddacd3 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index ea4799d9712..17c64c0f5d3 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py index fc02fed0945..33c4c2a6c19 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py @@ -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 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 a7c7e1853a0..72af43d73f7 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 @@ -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 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 0ba8b3f2a40..4117e5c9ea2 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 @@ -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() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 9f713c4175b..165324850be 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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)