mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): simplify OAuth checks and preserve mixed-failure challenges
This commit is contained in:
parent
8c7a4372ee
commit
0891aa4b35
7 changed files with 357 additions and 259 deletions
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue