mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(mcp): make token_exchange (OBO) production-ready - discovery threading + audit hardening + RFC 9728 challenge (#31622)
* feat(mcp): thread the caller token into tools/list discovery for token_exchange
A token_exchange (OBO) server's tools could not be discovered through the aggregator: the list path
never threaded the caller's token, so every tools/list hit the no-subject branch. v1 masked this with
its client_credentials fallback (discovery used a service token); v2 dropped that fallback, so listing
had no credential and the OBO server's tools never appeared - and an MCP client lists before it calls.
Thread the inbound subject_token into the list path the same way the call path does, gated on
auth_type oauth2_token_exchange so the caller's bearer never leaks into other modes:
_get_tools_from_server takes an oauth2_headers param, extracts the token via _extract_bearer_token, and
passes it to _create_mcp_client; server.py forwards oauth2_headers at the list call site.
authorization_code (resolves off identity plus stored token), the static/config modes, and the
background registry refresh are unaffected, and the list path's existing graceful degradation
(catch -> empty list) is preserved.
* fix(mcp): harden token_exchange OBO from the audit (strip, TTL/expires_in, subject_token_type)
- _should_strip_caller_authorization returns True for oauth2_token_exchange, so the inbound subject
token is never forwarded upstream raw - only the IdP-exchanged token is (matches authorization_code).
- _parse_expires_in accepts a JSON float / numeric-string expires_in, and _ttl_seconds caps the cache
TTL at the token's real remaining lifetime so a short-lived exchanged token is never served stale.
- to_server_spec normalizes a falsy subject_token_type to the default URN, parity with v1.
The subject/key disambiguation (never exchange the LiteLLM key; Authorization: Bearer <litellm-key>
support for /mcp) is intentionally a separate cross-cutting PR off staging, not part of this OBO work.
* fix(mcp): stop caller header bypassing OBO exchange; thread subject into prompts/resources
The per-server x-mcp-* override guard in _create_mcp_client only kept the v2 spec
for authorization_code, so a caller-supplied header silently disabled the RFC 8693
exchange on a token_exchange server and forwarded the raw bearer upstream. Extend
the guard to token_exchange so the exchange always runs and the caller cannot
substitute an arbitrary upstream credential.
prompts/list+get, resources/list+read, and resource-templates/list never threaded
the OBO subject token, so those operations failed closed (401 / empty) on a
token_exchange server. Thread the caller's bearer as the subject for those paths
too, gated on the token_exchange mode via a shared _obo_subject_token helper.
* fix(mcp): keep the OBO/authz_code resolver credential authoritative; centralize OpenAPI strip
A guardrail (e.g. MCPJWTSigner), static_headers, or any other injected Authorization could
shadow the resolver-owned credential for token_exchange / authorization_code servers, so the
upstream would receive e.g. the signer's JWT instead of the exchanged token and reject it. In
_create_mcp_client the resolver-owned credential now wins: a conflicting header is dropped and
the minted/stored token reaches upstream. No behavior change for none/passthrough/static modes,
where an injected Authorization still wins as before.
The OpenAPI/local _request_extra_headers forwarder gated its Authorization strip on
has_client_credentials only, so an OpenAPI-backed token_exchange server with
extra_headers:[Authorization] forwarded the raw subject token upstream and never exchanged. It
now uses the centralized _should_strip_caller_authorization so it matches the managed paths.
* feat(mcp): RFC 9728 challenge for token_exchange (OBO) unauthorized
OBO previously returned an opaque 401 (Bearer error="invalid_request") with no discovery
info, and any IdP exchange failure collapsed to a retryable 503. Now an OBO server behaves like
a standards-compliant OAuth resource server:
- A missing/rejected subject token returns the RFC 9728 / RFC 6750 challenge: 401 +
WWW-Authenticate: Bearer resource_metadata="...", error="invalid_token", so a spec-compliant
MCP client can discover the IdP, SSO, and retry with a fresh subject token.
- The protected-resource metadata for a token_exchange server advertises the JWT-auth issuer(s)
(JWT_ISSUER / litellm_jwtauth.issuers) as authorization_servers -- the IdP that issues and
validates the subject -- instead of the gateway.
- An IdP 4xx (subject rejected) is now a non-retryable 401 (the challenge) instead of a 503, so a
caller with a dead token re-authenticates rather than looping; 5xx/transport stays retryable 503.
* fix(mcp): emit the OBO RFC 9728 challenge preemptively so a no-subject client can discover the IdP
A token_exchange server's tools are not discoverable without a subject token (list is lenient ->
empty), and a tool-call-time 401 is wrapped into a JSON-RPC error so the WWW-Authenticate header is
lost. So a cold-start client never saw the challenge and could not start discovery. Add a
token_exchange branch to the preemptive-401: a no-subject connect to an OBO server now returns
401 + WWW-Authenticate: Bearer resource_metadata=..., error="invalid_token" at the transport level,
so a spec-compliant client discovers the IdP (the PRM advertises the JWT-auth issuer), SSOs, and
retries with a subject token. Verified live on the per-server endpoint; the with-subject connect
still proceeds (no challenge).
(Also formats two lines from earlier commits in this stack.)
* refactor(mcp): inject root_path into the OBO/OAuth challenge edge
The adapter's raise_user_oauth_challenge and raise_token_exchange_challenge
reached into os.getenv("SERVER_ROOT_PATH") via get_server_root_path(), a
hidden ambient read in a module that is meant to be a pure edge. That coupling
made the preemptive-challenge test order-dependent under xdist: a sibling test
sets SERVER_ROOT_PATH at import without cleanup, leaking the prefix into the
challenge URL and failing the exact-match assertion.
Resolve the root path at the imperative-shell call sites and pass it in
keyword-only, so both challenge builders become pure functions of their inputs.
Extract the shared resource_metadata path construction into a single
oauth_protected_resource_path helper, collapsing the duplicated prefix/name
logic the two functions carried.
Also reduce _create_mcp_client below the strict complexity ceiling by extracting
the v2 credential resolution into _resolve_v2_auth, and extract the OBO
protected-resource-metadata branch into _obo_protected_resource_response (which
shipped without coverage) so discovery can be unit-tested directly.
Tests are now hermetic: the adapter tests pass root_path as a real input rather
than monkeypatching the environment, the stale-session preemptive test asserts
structural invariants instead of the exact prefixed URL, and five new tests
cover the OBO PRM issuer branch end to end.
* feat(mcp): OBO cache-key tenant isolation, reactive 401 retry, v1-parity logs
From a pass over the OBO behavior contract. Three changes to the
token_exchange arm, none of which alters any other auth mode.
The exchanged-token cache key now folds in the caller's tenant alongside
the subject token and exchange config, so two tenants presenting the same
opaque token can never share a cache entry; cross-tenant isolation is
structural rather than incidental to subject-token uniqueness. tenant_id is
threaded from the resolver's Subject; it is keyword-only with an empty
default so the no-tenant case and the existing call sites are unchanged.
The tool-call path gains one reactive retry. When an upstream rejects the
injected token with a 401/403, the gateway invalidates the cached exchange,
re-mints once through the IdP by rebuilding the client, and retries the call
exactly once before surfacing the upstream error, so a token revoked or
rotated upstream mid-TTL self-heals without an infinite loop. It is gated
strictly to oauth2_token_exchange; passthrough, authorization_code,
client_credentials, api_key, and none keep their single-call behavior.
MCPClient.call_tool gains a raise_on_error flag (mirroring list_tools) so
the path can tell an upstream 401 apart from an ordinary tool error and
avoid re-running a non-idempotent tool on a non-auth failure.
The exchanger also emits the v1-parity log lines it had dropped (attempt
with server, endpoint and audience; success; cache hit), while never
logging the form, subject token, secret, or minted token.
* fix(mcp): fail closed with 412 when a token_exchange server has no endpoint
A true token_exchange (OBO) server must use only an explicitly configured
token endpoint; it must never guess an IdP or silently fall back to a weaker
source. Previously an OBO server with client credentials but no
token_exchange_endpoint/token_url deferred to v1, which no-op'd and let the
request connect to the upstream with no credential (an upstream 401 rather
than a clear gateway error).
Now such a server is owned by the v2 arm: _token_exchange_spec builds the spec
even when the endpoint is absent, and the exchanger fails closed with a
precondition_required error that maps to HTTP 412 before any upstream or IdP
call, with the caller's subject token never sent anywhere. A missing
client_id/secret still maps to misconfigured (500); a present-but-rejected
subject still maps to 401; an unreachable IdP still maps to 503. The no-subject
case keeps its existing 401 RFC 9728 challenge.
* feat(mcp): log a refused non-Bearer token_type in the OBO exchange
* fix(mcp): surface OBO/authorization_code list-time 401 as a challenge instead of masking it
* feat(mcp): classify RFC 6749 gateway-fault token-exchange errors as 500, not a caller 401
* test(mcp): absorb fixture uses 500 now that 401/403 are challenge-class at list time
* style(mcp): PEP 604 union in the OBO retry signature to keep the UP007 budget flat
This commit is contained in:
parent
f19bf2c984
commit
0e56fc39e2
16 changed files with 1496 additions and 572 deletions
|
|
@ -520,13 +520,28 @@ class MCPClient:
|
|||
# Return empty list instead of raising to allow graceful degradation
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def error_tool_result(exc: Exception) -> MCPCallToolResult:
|
||||
"""The error result ``call_tool`` returns when it swallows a failure (no re-execution)."""
|
||||
return MCPCallToolResult(
|
||||
content=[TextContent(type="text", text=f"{type(exc).__name__}: {str(exc)}")],
|
||||
isError=True,
|
||||
)
|
||||
|
||||
async def call_tool(
|
||||
self,
|
||||
call_tool_request_params: MCPCallToolRequestParams,
|
||||
host_progress_callback: Optional[Callable] = None,
|
||||
raise_on_error: bool = False,
|
||||
) -> MCPCallToolResult:
|
||||
"""
|
||||
Call an MCP Tool.
|
||||
|
||||
Args:
|
||||
raise_on_error: When True, re-raise the underlying exception instead of returning an
|
||||
``isError=True`` result. The token-exchange (OBO) tool-call path uses this to detect
|
||||
an upstream 401 so it can re-mint the exchanged token and retry once; every other
|
||||
caller keeps the default and gets graceful ``isError`` degradation.
|
||||
"""
|
||||
verbose_logger.info(f"MCP client calling tool '{call_tool_request_params.name}'")
|
||||
|
||||
|
|
@ -579,11 +594,10 @@ class MCPClient:
|
|||
"MCP client detected broken connection/stream - "
|
||||
"the MCP server may have crashed, disconnected, or timed out."
|
||||
)
|
||||
if raise_on_error:
|
||||
raise
|
||||
# Return a default error result instead of raising
|
||||
return MCPCallToolResult(
|
||||
content=[TextContent(type="text", text=f"{error_type}: {str(e)}")], # Empty content for error case
|
||||
isError=True,
|
||||
)
|
||||
return self.error_tool_result(e)
|
||||
|
||||
async def list_prompts(self) -> List[Prompt]:
|
||||
"""List available prompts from the server."""
|
||||
|
|
|
|||
|
|
@ -1288,7 +1288,14 @@ async def _build_oauth_protected_resource_response(
|
|||
detail=(f"Upstream oauth-protected-resource metadata unavailable for MCP server {mcp_server.name!r}"),
|
||||
)
|
||||
|
||||
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth-protected resource")
|
||||
obo_response = _obo_protected_resource_response(mcp_server, resource_url)
|
||||
if obo_response is not None:
|
||||
return obo_response
|
||||
|
||||
# An OBO server with no configured issuer falls through to the gateway default so discovery still
|
||||
# returns metadata; every other non-oauth2 named server 404s to avoid enumeration.
|
||||
if mcp_server is None or mcp_server.auth_type != MCPAuth.oauth2_token_exchange:
|
||||
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth-protected resource")
|
||||
|
||||
return {
|
||||
"authorization_servers": [
|
||||
|
|
@ -1299,6 +1306,51 @@ async def _build_oauth_protected_resource_response(
|
|||
}
|
||||
|
||||
|
||||
def _obo_protected_resource_response(mcp_server: Optional[MCPServer], resource_url: str) -> Optional[dict]:
|
||||
"""The OBO (token_exchange) PRM, or None when this server is not OBO / no issuer is configured.
|
||||
|
||||
The client SSOs with the IdP to obtain a subject token, which LiteLLM then exchanges, so discovery
|
||||
points at the JWT-auth issuer(s) LiteLLM trusts (the same IdP that issues and validates the
|
||||
subject), not the gateway. None falls the caller back to the gateway default so discovery still
|
||||
returns metadata; it just can't name the IdP.
|
||||
"""
|
||||
if mcp_server is None or mcp_server.auth_type != MCPAuth.oauth2_token_exchange:
|
||||
return None
|
||||
issuers = _jwt_auth_issuers()
|
||||
if not issuers:
|
||||
return None
|
||||
return {
|
||||
"authorization_servers": issuers,
|
||||
"resource": resource_url,
|
||||
"scopes_supported": (mcp_server.scopes if mcp_server.scopes else []),
|
||||
}
|
||||
|
||||
|
||||
def _jwt_auth_issuers() -> list:
|
||||
"""The OAuth issuer identifier(s) LiteLLM's JWT auth trusts, for the OBO PRM authorization_servers.
|
||||
|
||||
In token_exchange the IdP that issues the subject JWT is the same one LiteLLM validates it
|
||||
against, so OBO discovery points clients at the JWT-auth issuer to obtain a subject token.
|
||||
Sourced from ``JWT_ISSUER`` and any configured ``litellm_jwtauth.issuers``.
|
||||
"""
|
||||
import os # noqa: PLC0415
|
||||
|
||||
from litellm.proxy.proxy_server import general_settings # noqa: PLC0415
|
||||
|
||||
issuers: list = []
|
||||
env_issuer = os.getenv("JWT_ISSUER")
|
||||
if env_issuer:
|
||||
issuers.append(env_issuer)
|
||||
|
||||
jwtauth = general_settings.get("litellm_jwtauth") if isinstance(general_settings, dict) else None
|
||||
raw_issuers = jwtauth.get("issuers") if isinstance(jwtauth, dict) else getattr(jwtauth, "issuers", None)
|
||||
for cfg in raw_issuers or []:
|
||||
issuer = cfg.get("issuer") if isinstance(cfg, dict) else getattr(cfg, "issuer", None)
|
||||
if issuer and issuer not in issuers:
|
||||
issuers.append(issuer)
|
||||
return issuers
|
||||
|
||||
|
||||
# Standard MCP pattern: /.well-known/oauth-protected-resource/mcp/{server_name}
|
||||
# This is the pattern expected by standard MCP clients (mcp-inspector, VSCode Copilot)
|
||||
@router.get(
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from typing import Any, AsyncIterator, Callable, Literal, Optional, Union, cast
|
|||
from urllib.parse import urlparse
|
||||
|
||||
import anyio
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
from httpx import HTTPStatusError
|
||||
from mcp import ReadResourceResult, Resource
|
||||
|
|
@ -64,6 +65,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials import (
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_public,
|
||||
raise_token_exchange_challenge,
|
||||
raise_user_oauth_challenge,
|
||||
to_server_spec,
|
||||
to_subject,
|
||||
|
|
@ -76,6 +78,8 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
AuthorizationCodeConfig,
|
||||
ServerSpec,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
MCP_TOOL_PREFIX_SEPARATOR,
|
||||
|
|
@ -108,7 +112,7 @@ from litellm.proxy._types import (
|
|||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
|
||||
from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.proxy.utils import ProxyLogging, get_server_root_path
|
||||
from litellm.repositories.table_repositories import MCPServerRepository
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
from litellm.types.mcp import MCPAuth, MCPStdioConfig
|
||||
|
|
@ -210,6 +214,10 @@ def _should_strip_caller_authorization(
|
|||
``Authorization`` is the upstream OAuth token and must be
|
||||
forwarded, so we keep it.
|
||||
"""
|
||||
if mcp_server.auth_type == MCPAuth.oauth2_token_exchange:
|
||||
# OBO: the inbound Authorization is the subject token. It is exchanged at the IdP and only the
|
||||
# exchanged token is sent upstream, so the raw caller bearer must never be forwarded.
|
||||
return True
|
||||
if mcp_server.has_client_credentials:
|
||||
return True
|
||||
if mcp_server.auth_type == MCPAuth.oauth2 and to_server_spec(mcp_server) is not None:
|
||||
|
|
@ -1672,6 +1680,21 @@ class MCPServerManager:
|
|||
return auth_value
|
||||
return None
|
||||
|
||||
def _obo_subject_token(
|
||||
self,
|
||||
server: MCPServer,
|
||||
raw_headers: Optional[dict[str, str]],
|
||||
) -> Optional[str]:
|
||||
"""The caller's bearer as the token_exchange (OBO) subject token, for that mode only.
|
||||
|
||||
Prompts/resources discovery and reads on a token_exchange server must exchange the caller's
|
||||
token like the tools paths do, not connect with no credential. Other modes never read the
|
||||
inbound bearer, so return None to avoid forwarding it.
|
||||
"""
|
||||
if server.auth_type != MCPAuth.oauth2_token_exchange:
|
||||
return None
|
||||
return self._extract_bearer_token(None, raw_headers)
|
||||
|
||||
def _build_stdio_env(
|
||||
self,
|
||||
server: MCPServer,
|
||||
|
|
@ -1862,6 +1885,54 @@ class MCPServerManager:
|
|||
_write_user_env_vars_cache(user_id, server.server_id, values)
|
||||
return values
|
||||
|
||||
async def _resolve_v2_auth(
|
||||
self,
|
||||
*,
|
||||
server: MCPServer,
|
||||
spec: ServerSpec,
|
||||
provider: UpstreamCredentialProvider,
|
||||
subject_token: Optional[str],
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
extra_headers: Optional[dict[str, str]],
|
||||
) -> tuple[Optional[httpx.Auth], Optional[dict[str, str]]]:
|
||||
"""Resolve a v2-owned server's upstream credential into ``(resolved_auth, extra_headers)``.
|
||||
|
||||
On a missing/rejected per-user credential this raises the mode's discovery challenge
|
||||
(authorization_code's browser-OAuth 401, token_exchange's RFC 9728 challenge) or maps any
|
||||
other ``CredError`` onto its public HTTP status; it never returns an error as a value.
|
||||
"""
|
||||
match await provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec):
|
||||
case Ok(auth):
|
||||
# NoOpAuth has no header_name and so never conflicts.
|
||||
header_name = getattr(auth, "header_name", None)
|
||||
conflicts = bool(
|
||||
header_name and extra_headers and any(key.lower() == header_name.lower() for key in extra_headers)
|
||||
)
|
||||
if not conflicts:
|
||||
return auth, extra_headers
|
||||
if isinstance(spec.config, (TokenExchangeConfig, AuthorizationCodeConfig)):
|
||||
# The resolver owns the per-user credential here (token_exchange's exchanged
|
||||
# token, authorization_code's stored token). It is authoritative: a guardrail such
|
||||
# as MCPJWTSigner, static_headers, or any other injected Authorization must NOT
|
||||
# shadow it (otherwise the upstream gets e.g. the signer's JWT instead of the
|
||||
# exchanged token and rejects it). Drop the conflicting header so the resolved
|
||||
# token reaches upstream.
|
||||
return auth, _without_authorization(extra_headers)
|
||||
# Other modes: an Authorization already supplied via extra_headers (a forwarded caller
|
||||
# header or static_headers) is intentional and wins; v1 applies those last.
|
||||
return None, extra_headers
|
||||
case Error(err):
|
||||
if err.tag == "unauthorized" and isinstance(spec.config, AuthorizationCodeConfig):
|
||||
# authorization_code's missing per-user token -> the per-server browser-OAuth
|
||||
# challenge, built here where the full MCPServer is in hand.
|
||||
raise_user_oauth_challenge(server, root_path=get_server_root_path())
|
||||
if err.tag == "unauthorized" and isinstance(spec.config, TokenExchangeConfig):
|
||||
# token_exchange (OBO): a missing/rejected subject token -> the RFC 9728 challenge
|
||||
# pointing at the IdP the client must SSO with to obtain one, rather than an opaque
|
||||
# 401. No gateway-side browser flow.
|
||||
raise_token_exchange_challenge(server, root_path=get_server_root_path())
|
||||
raise_public(err)
|
||||
|
||||
async def _create_mcp_client(
|
||||
self,
|
||||
server: MCPServer,
|
||||
|
|
@ -1896,11 +1967,17 @@ class MCPServerManager:
|
|||
spec = None if transport == MCPTransport.stdio else to_server_spec(server)
|
||||
provider = cred_provider or self._cred_provider
|
||||
# A caller-supplied per-request override (mcp_auth_header / x-mcp-*) defers to the v1 path
|
||||
# so it wins - except for authorization_code, whose per-user token the v2 resolver owns. A
|
||||
# caller must not be able to substitute another user's stored credential, so we keep the v2
|
||||
# spec and ignore the override there; the REST tools preview supplies its not-yet-persisted
|
||||
# token through the resolver (cred_provider), never this path.
|
||||
if spec is not None and mcp_auth_header and not isinstance(spec.config, AuthorizationCodeConfig):
|
||||
# so it wins - except for the per-user modes the v2 resolver owns (authorization_code's
|
||||
# stored token and token_exchange's RFC 8693 minted token). A caller must not be able to
|
||||
# substitute another user's stored credential, nor silently disable the OBO exchange and
|
||||
# forward an arbitrary bearer upstream, so we keep the v2 spec and ignore the override for
|
||||
# both; the REST tools preview supplies its not-yet-persisted token through the resolver
|
||||
# (cred_provider), never this path.
|
||||
if (
|
||||
spec is not None
|
||||
and mcp_auth_header
|
||||
and not isinstance(spec.config, (AuthorizationCodeConfig, TokenExchangeConfig))
|
||||
):
|
||||
spec = None
|
||||
auth_value = (
|
||||
await resolve_mcp_auth(server, mcp_auth_header, subject_token=subject_token) if spec is None else None
|
||||
|
|
@ -1966,28 +2043,14 @@ class MCPServerManager:
|
|||
server_url = server.url or ""
|
||||
|
||||
if spec is not None:
|
||||
match await provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec):
|
||||
case Ok(auth):
|
||||
resolved_auth = auth
|
||||
# Do not override an Authorization already supplied via extra_headers
|
||||
# (a guardrail hook such as the JWT signer, static_headers, or a
|
||||
# forwarded caller header): v1 applies those last, so they win. NoOpAuth
|
||||
# has no header_name and so never skips.
|
||||
header_name = getattr(resolved_auth, "header_name", None)
|
||||
if (
|
||||
header_name
|
||||
and extra_headers
|
||||
and any(key.lower() == header_name.lower() for key in extra_headers)
|
||||
):
|
||||
resolved_auth = None
|
||||
case Error(err):
|
||||
if err.tag == "unauthorized" and isinstance(spec.config, AuthorizationCodeConfig):
|
||||
# authorization_code's missing per-user token -> the per-server
|
||||
# browser-OAuth challenge, built here where the full MCPServer is in
|
||||
# hand. token_exchange and other modes carry their own 401 (e.g. OBO
|
||||
# needs a caller token, not a browser flow), so they go via raise_public.
|
||||
raise_user_oauth_challenge(server)
|
||||
raise_public(err)
|
||||
resolved_auth, extra_headers = await self._resolve_v2_auth(
|
||||
server=server,
|
||||
spec=spec,
|
||||
provider=provider,
|
||||
subject_token=subject_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
extra_headers=extra_headers,
|
||||
)
|
||||
return MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
|
|
@ -2032,6 +2095,7 @@ class MCPServerManager:
|
|||
add_prefix: bool = True,
|
||||
raw_headers: Optional[dict[str, str]] = None,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
oauth2_headers: Optional[dict[str, str]] = None,
|
||||
) -> list[MCPTool]:
|
||||
"""
|
||||
Helper method to get tools from a single MCP server with prefixed names.
|
||||
|
|
@ -2105,11 +2169,21 @@ class MCPServerManager:
|
|||
|
||||
stdio_env = self._build_stdio_env(server, raw_headers)
|
||||
|
||||
# token_exchange (OBO) discovery needs the caller's token too: list it with the user's own
|
||||
# token (mirrors the call path), not v1's deleted client_credentials fallback. Other modes
|
||||
# never read the inbound bearer, so leave subject_token None to avoid forwarding it.
|
||||
subject_token = (
|
||||
self._extract_bearer_token(oauth2_headers, raw_headers)
|
||||
if server.auth_type == MCPAuth.oauth2_token_exchange
|
||||
else None
|
||||
)
|
||||
|
||||
client = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
|
|
@ -2149,12 +2223,16 @@ class MCPServerManager:
|
|||
# aggregator catches this explicitly to keep absorbing.
|
||||
raise
|
||||
except HTTPException as e:
|
||||
headers = e.headers or {}
|
||||
www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate")
|
||||
if e.status_code == 401 and www_authenticate is not None:
|
||||
# A v2 resolver auth challenge (token_exchange's RFC 9728 401, authorization_code's
|
||||
# browser-OAuth 401, or a 403) is raised at client-build time, inside this try. Route it
|
||||
# through the same MCPUpstreamAuthError channel as pass-through so single-server routes
|
||||
# surface the challenge (the client re-authenticates) while the aggregator keeps absorbing.
|
||||
# Non-auth HTTP errors stay absorbed so one misconfigured server can't blank the listing.
|
||||
if e.status_code in (401, 403):
|
||||
headers = e.headers or {}
|
||||
raise MCPUpstreamAuthError(
|
||||
status_code=401,
|
||||
www_authenticate=www_authenticate,
|
||||
status_code=e.status_code,
|
||||
www_authenticate=headers.get("WWW-Authenticate") or headers.get("www-authenticate"),
|
||||
server_name=server.name,
|
||||
) from e
|
||||
verbose_logger.warning(f"Failed to get tools from server {server.name}: {str(e)}")
|
||||
|
|
@ -2194,12 +2272,14 @@ class MCPServerManager:
|
|||
extra_headers.update(server.static_headers)
|
||||
|
||||
stdio_env = self._build_stdio_env(server, raw_headers)
|
||||
subject_token = self._obo_subject_token(server, raw_headers)
|
||||
|
||||
client = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
)
|
||||
|
||||
prompts = await client.list_prompts()
|
||||
|
|
@ -2234,12 +2314,14 @@ class MCPServerManager:
|
|||
extra_headers.update(server.static_headers)
|
||||
|
||||
stdio_env = self._build_stdio_env(server, raw_headers)
|
||||
subject_token = self._obo_subject_token(server, raw_headers)
|
||||
|
||||
client = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
)
|
||||
|
||||
resources = await client.list_resources()
|
||||
|
|
@ -2274,12 +2356,14 @@ class MCPServerManager:
|
|||
extra_headers.update(server.static_headers)
|
||||
|
||||
stdio_env = self._build_stdio_env(server, raw_headers)
|
||||
subject_token = self._obo_subject_token(server, raw_headers)
|
||||
|
||||
client = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
)
|
||||
|
||||
resource_templates = await client.list_resource_templates()
|
||||
|
|
@ -2313,12 +2397,14 @@ class MCPServerManager:
|
|||
extra_headers.update(server.static_headers)
|
||||
|
||||
stdio_env = self._build_stdio_env(server, raw_headers)
|
||||
subject_token = self._obo_subject_token(server, raw_headers)
|
||||
|
||||
client = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
)
|
||||
|
||||
return await client.read_resource(url)
|
||||
|
|
@ -2343,12 +2429,14 @@ class MCPServerManager:
|
|||
extra_headers.update(server.static_headers)
|
||||
|
||||
stdio_env = self._build_stdio_env(server, raw_headers)
|
||||
subject_token = self._obo_subject_token(server, raw_headers)
|
||||
|
||||
client = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
)
|
||||
|
||||
get_prompt_request_params = GetPromptRequestParams(
|
||||
|
|
@ -3242,6 +3330,46 @@ class MCPServerManager:
|
|||
async with semaphore:
|
||||
yield
|
||||
|
||||
async def _obo_call_tool_with_retry(
|
||||
self,
|
||||
*,
|
||||
client: MCPClient,
|
||||
call_tool_params: MCPCallToolRequestParams,
|
||||
host_progress_callback: Optional[Callable],
|
||||
mcp_server: MCPServer,
|
||||
server_auth_header: str | dict[str, str] | None,
|
||||
extra_headers: Optional[dict[str, str]],
|
||||
stdio_env: Optional[dict[str, str]],
|
||||
subject_token: Optional[str],
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
) -> CallToolResult:
|
||||
"""Call a token_exchange (OBO) tool; on an upstream 401/403 re-mint the token once and retry.
|
||||
|
||||
The exchanged token is baked into the client at build time, so the retry invalidates the
|
||||
cached exchange and rebuilds the client (which re-exchanges). One retry only: a non-auth
|
||||
failure or a second auth failure degrades to the normal ``isError`` result, and a re-exchange
|
||||
that now fails surfaces its own 401 challenge from ``_create_mcp_client``.
|
||||
"""
|
||||
try:
|
||||
return await client.call_tool(
|
||||
call_tool_params, host_progress_callback=host_progress_callback, raise_on_error=True
|
||||
)
|
||||
except Exception as exc:
|
||||
if _extract_upstream_auth_failure(exc) is None:
|
||||
return MCPClient.error_tool_result(exc)
|
||||
spec = to_server_spec(mcp_server)
|
||||
if spec is not None:
|
||||
await self._cred_provider.invalidate_credentials(to_subject(user_api_key_auth, subject_token), spec)
|
||||
retry_client = await self._create_mcp_client(
|
||||
server=mcp_server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
return await retry_client.call_tool(call_tool_params, host_progress_callback=host_progress_callback)
|
||||
|
||||
async def _call_regular_mcp_tool(
|
||||
self,
|
||||
mcp_server: MCPServer,
|
||||
|
|
@ -3400,11 +3528,30 @@ class MCPServerManager:
|
|||
arguments=arguments,
|
||||
)
|
||||
|
||||
async def _call_tool_via_client(client, params):
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
return await client.call_tool(params, host_progress_callback=host_progress_callback)
|
||||
if mcp_server.auth_type == MCPAuth.oauth2_token_exchange and subject_token:
|
||||
# OBO: the exchanged token may have been revoked/rotated upstream since it was cached, so
|
||||
# an upstream 401 gets one re-mint + retry. Gated to this mode; all others keep the plain
|
||||
# single call below.
|
||||
tool_call_coro = self._obo_call_tool_with_retry(
|
||||
client=client,
|
||||
call_tool_params=call_tool_params,
|
||||
host_progress_callback=host_progress_callback,
|
||||
mcp_server=mcp_server,
|
||||
server_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
else:
|
||||
|
||||
tasks.append(asyncio.create_task(_call_tool_via_client(client, call_tool_params)))
|
||||
async def _call_tool_via_client(client, params):
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
return await client.call_tool(params, host_progress_callback=host_progress_callback)
|
||||
|
||||
tool_call_coro = _call_tool_via_client(client, call_tool_params)
|
||||
|
||||
tasks.append(asyncio.create_task(tool_call_coro))
|
||||
|
||||
_timeout = mcp_server.timeout if mcp_server.timeout is not None else MCP_CLIENT_TIMEOUT
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -102,21 +102,23 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]:
|
|||
|
||||
|
||||
def _token_exchange_spec(server: MCPServer, resource: str) -> Optional[ServerSpec]:
|
||||
"""Build a token_exchange (RFC 8693 OBO) spec, or defer (None) if the exchange config is absent.
|
||||
"""Build a token_exchange (RFC 8693 OBO) spec, or defer (None) when it is not OBO-configured.
|
||||
|
||||
Mirrors v1's ``has_token_exchange_config`` precondition: an endpoint (``token_exchange_endpoint``
|
||||
or ``token_url``) plus ``client_id``/``client_secret`` must all be present, else there is nothing
|
||||
to exchange against and the server stays on v1 (parity-safe). ``audience`` is forwarded only when
|
||||
the operator set it; a missing one is omitted, not derived.
|
||||
An OBO server with ``client_id``/``client_secret`` is owned by the v2 arm even if the
|
||||
``token_exchange_endpoint``/``token_url`` is absent: a missing endpoint then fails closed (412) at
|
||||
the exchanger rather than silently deferring to v1 and connecting unauthenticated, since the
|
||||
gateway must not guess the IdP or fall back to a weaker source. Without client credentials there is
|
||||
nothing to own, so the server stays on v1 (parity-safe). ``audience`` is forwarded only when the
|
||||
operator set it; a missing one is omitted, not derived.
|
||||
"""
|
||||
endpoint = server.token_exchange_endpoint or server.token_url
|
||||
if not endpoint or not server.client_id or not server.client_secret:
|
||||
if not server.client_id or not server.client_secret:
|
||||
return None
|
||||
return ServerSpec(
|
||||
server_id=server.server_id,
|
||||
resource=resource,
|
||||
config=TokenExchangeConfig(
|
||||
subject_token_type=server.subject_token_type,
|
||||
subject_token_type=server.subject_token_type or "urn:ietf:params:oauth:token-type:access_token",
|
||||
token_exchange_endpoint=endpoint,
|
||||
audience=server.audience,
|
||||
client_id=server.client_id,
|
||||
|
|
@ -178,23 +180,52 @@ def raise_public(error: CredError) -> NoReturn:
|
|||
assert_never(error.tag)
|
||||
|
||||
|
||||
def raise_user_oauth_challenge(server: MCPServer) -> NoReturn:
|
||||
def oauth_protected_resource_path(root_path: str, server: MCPServer) -> str:
|
||||
"""The server's RFC 9728 Protected Resource Metadata path, the shared anchor of both challenges.
|
||||
|
||||
``root_path`` is the proxy's ``SERVER_ROOT_PATH``, resolved by the caller (the imperative shell)
|
||||
so this stays a pure function of its inputs; ``"/"`` and ``""`` both mean no prefix. The path is
|
||||
relative, so it resolves against the caller's own host (correct even behind a reverse proxy).
|
||||
"""
|
||||
prefix = "" if root_path == "/" else root_path
|
||||
name = server.alias or server.server_name or server.name or server.server_id
|
||||
return f"/.well-known/oauth-protected-resource{prefix}/mcp/{name}"
|
||||
|
||||
|
||||
def raise_user_oauth_challenge(server: MCPServer, *, root_path: str) -> NoReturn:
|
||||
"""Raise the 401 an ``authorization_code`` server returns at egress when the user has no token.
|
||||
|
||||
Points at the server's RFC 9728 Protected Resource Metadata (``resource_metadata``), which names
|
||||
the upstream authorization server the client must complete OAuth with. The URL is per-server and
|
||||
relative, so it resolves against the caller's own host (correct even behind a reverse proxy)
|
||||
without needing request context. The listing-phase 401 still emits the RFC 8414 ``authorization_uri``
|
||||
form pending the format unification; both target the same server, so the difference is cosmetic.
|
||||
Points at the server's RFC 9728 Protected Resource Metadata, which names the upstream
|
||||
authorization server the client must complete OAuth with. The listing-phase 401 still emits the
|
||||
RFC 8414 ``authorization_uri`` form pending the format unification; both target the same server,
|
||||
so the difference is cosmetic.
|
||||
"""
|
||||
from litellm.proxy.utils import get_server_root_path # noqa: PLC0415
|
||||
|
||||
root = get_server_root_path()
|
||||
prefix = "" if root == "/" else root
|
||||
name = server.alias or server.server_name or server.name or server.server_id
|
||||
resource_metadata = f"/.well-known/oauth-protected-resource{prefix}/mcp/{name}"
|
||||
resource_metadata = oauth_protected_resource_path(root_path, server)
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Unauthorized",
|
||||
headers={"WWW-Authenticate": f'Bearer resource_metadata="{resource_metadata}"'},
|
||||
)
|
||||
|
||||
|
||||
def raise_token_exchange_challenge(server: MCPServer, *, root_path: str) -> NoReturn:
|
||||
"""Raise the RFC 9728 / RFC 6750 challenge an OBO (``token_exchange``) server returns when the
|
||||
caller's subject token is missing or the IdP rejected it.
|
||||
|
||||
Points at the server's Protected Resource Metadata, whose ``authorization_servers`` names the IdP
|
||||
the client must SSO with to obtain a subject token; ``error="invalid_token"`` tells a
|
||||
spec-compliant MCP client to discover that AS and retry with a fresh bearer. Mirrors
|
||||
``raise_user_oauth_challenge`` but for the exchange flow: there is no gateway-side browser OAuth —
|
||||
the client re-authenticates directly with the IdP, and LiteLLM then exchanges the resulting token.
|
||||
"""
|
||||
resource_metadata = oauth_protected_resource_path(root_path, server)
|
||||
www_authenticate = (
|
||||
f'Bearer resource_metadata="{resource_metadata}", '
|
||||
'error="invalid_token", '
|
||||
'error_description="Missing or invalid subject token; authenticate with the IdP and retry"'
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Unauthorized",
|
||||
headers={"WWW-Authenticate": www_authenticate},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -63,10 +63,15 @@ class _NullTokenExchanger:
|
|||
"""Fail-closed default: with no exchanger wired, token_exchange cannot produce a credential."""
|
||||
|
||||
async def exchange(
|
||||
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig
|
||||
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig, *, tenant_id: str = ""
|
||||
) -> Result[OAuthToken, CredError]:
|
||||
return Error(CredError.of_misconfigured("token exchange collaborator not wired"))
|
||||
|
||||
async def invalidate(
|
||||
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig, *, tenant_id: str = ""
|
||||
) -> None:
|
||||
return None
|
||||
|
||||
|
||||
class UpstreamCredentialProvider:
|
||||
"""Produces the one `httpx.Auth` for a `(subject, upstream)` pair, per declared mode.
|
||||
|
|
@ -146,12 +151,26 @@ class UpstreamCredentialProvider:
|
|||
www_authenticate='Bearer error="invalid_request"',
|
||||
)
|
||||
)
|
||||
match await self._token_exchanger.exchange(inbound.get_secret_value(), server, config):
|
||||
match await self._token_exchanger.exchange(
|
||||
inbound.get_secret_value(), server, config, tenant_id=subject.tenant_id
|
||||
):
|
||||
case Ok(token):
|
||||
return Ok(StaticHeaderAuth(f"Bearer {token.access_token}", header_name="Authorization"))
|
||||
case Error(err):
|
||||
return Error(err)
|
||||
|
||||
async def invalidate_credentials(self, subject: Subject, server: ServerSpec) -> None:
|
||||
"""Drop any cached credential the resolver owns for this `(subject, server)`.
|
||||
|
||||
Used after an upstream rejects the injected credential, so the next resolve re-mints rather
|
||||
than serving the same rejected token until TTL. Only `token_exchange` holds a re-mintable
|
||||
cached credential here; other modes are a no-op.
|
||||
"""
|
||||
if isinstance(server.config, TokenExchangeConfig) and subject.inbound_token is not None:
|
||||
await self._token_exchanger.invalidate(
|
||||
subject.inbound_token.get_secret_value(), server, server.config, tenant_id=subject.tenant_id
|
||||
)
|
||||
|
||||
async def _authz_token(self, subject: Subject, server: ServerSpec) -> OAuthToken | None:
|
||||
"""The user's authorization_code token, or None when absent or the store is unreachable.
|
||||
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ call), so it needs no lazy wrapper.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import (
|
||||
MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL,
|
||||
|
|
@ -21,8 +23,33 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import (
|
||||
Rfc8693TokenExchanger,
|
||||
SubjectTokenRejected,
|
||||
TokenExchangeClientError,
|
||||
)
|
||||
|
||||
# RFC 6749 5.2 error codes that mean the gateway's own request/credentials are wrong (not the
|
||||
# caller's subject token), so they surface as a 500 the caller can't fix by re-authenticating.
|
||||
_GATEWAY_FAULT_OAUTH_ERRORS = frozenset(
|
||||
{"invalid_client", "unauthorized_client", "unsupported_grant_type", "invalid_target", "invalid_scope"}
|
||||
)
|
||||
|
||||
|
||||
def _oauth_error_code(response: httpx.Response) -> str | None:
|
||||
"""Read the RFC 6749 5.2 ``error`` code from a token-endpoint error body, or None if absent.
|
||||
|
||||
The ``error_description`` is deliberately not read: it can carry IdP internals and must never
|
||||
reach the caller. Only the standard machine code drives classification.
|
||||
"""
|
||||
try:
|
||||
body: object = response.json()
|
||||
except Exception: # noqa: BLE001
|
||||
return None
|
||||
if isinstance(body, dict):
|
||||
code = body.get("error")
|
||||
if isinstance(code, str):
|
||||
return code
|
||||
return None
|
||||
|
||||
|
||||
async def _post_exchange_endpoint(
|
||||
url: str, form: dict[str, str], client_auth_headers: dict[str, str]
|
||||
|
|
@ -34,13 +61,29 @@ async def _post_exchange_endpoint(
|
|||
|
||||
# litellm's httpx handler and httpx.Response are only partially typed; the IdP returns a JSON
|
||||
# object and the exchanger validates each field, so the untyped boundary is contained here.
|
||||
# A failed exchange is a miss, not a 500 (matches v1), so any error becomes None.
|
||||
# A 4xx is the IdP rejecting the subject (non-retryable -> 401 via SubjectTokenRejected); any
|
||||
# other failure is a miss (-> None -> upstream_unavailable -> 503), matching v1's fail-closed.
|
||||
headers = {"Accept": "application/json", **client_auth_headers}
|
||||
try:
|
||||
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) # pyright: ignore
|
||||
response = await client.post(url, headers=headers, data=form) # pyright: ignore
|
||||
response.raise_for_status() # pyright: ignore
|
||||
parsed: object = response.json() # pyright: ignore
|
||||
except httpx.HTTPStatusError as status_err:
|
||||
status_code = status_err.response.status_code
|
||||
if 400 <= status_code < 500:
|
||||
oauth_error = _oauth_error_code(status_err.response)
|
||||
if oauth_error in _GATEWAY_FAULT_OAUTH_ERRORS:
|
||||
verbose_logger.warning(
|
||||
"MCP token exchange rejected as %s (HTTP %d); check the gateway client credentials, "
|
||||
"audience, and scope for this server",
|
||||
oauth_error,
|
||||
status_code,
|
||||
)
|
||||
raise TokenExchangeClientError(oauth_error) from status_err
|
||||
raise SubjectTokenRejected(f"IdP rejected the subject token (HTTP {status_code})") from status_err
|
||||
verbose_logger.warning("MCP token exchange request failed: %s", status_err)
|
||||
return None
|
||||
except Exception as exc: # noqa: BLE001
|
||||
verbose_logger.warning("MCP token exchange request failed: %s", exc)
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ import time
|
|||
from collections.abc import Awaitable, Callable
|
||||
from typing import Protocol
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
InMemoryTokenCacheBackend,
|
||||
InProcessRefreshCoordinator,
|
||||
|
|
@ -70,26 +71,51 @@ _NON_ACCESS_ISSUED_TOKEN_TYPES = frozenset(
|
|||
ExchangeHttpPost = Callable[[str, "dict[str, str]", "dict[str, str]"], Awaitable["dict[str, object] | None"]]
|
||||
|
||||
|
||||
class SubjectTokenRejected(Exception):
|
||||
"""The IdP refused to exchange the subject token (an RFC 8693 4xx, e.g. ``invalid_grant``).
|
||||
|
||||
Distinct from a transport / IdP-availability failure, which the post adapter maps to ``None`` ->
|
||||
``upstream_unavailable`` -> 503 (retryable). A rejected subject is the caller's problem, not the
|
||||
gateway's, so the arm surfaces it as a non-retryable 401 (the OBO challenge) instead.
|
||||
"""
|
||||
|
||||
|
||||
class TokenExchangeClientError(Exception):
|
||||
"""The IdP rejected the exchange for a reason that is the gateway's fault, not the caller's.
|
||||
|
||||
RFC 6749 5.2 codes such as ``invalid_client`` (the gateway's own STS credentials are wrong),
|
||||
``unauthorized_client`` / ``unsupported_grant_type`` (the gateway is not permitted to exchange),
|
||||
``invalid_target`` / ``invalid_scope`` (the gateway's audience/scope config for this server is
|
||||
wrong). The caller cannot fix these by re-authenticating, so the arm surfaces them as a 500
|
||||
(``misconfigured``), not the 401 OBO challenge. The IdP ``error_description`` is never carried.
|
||||
"""
|
||||
|
||||
|
||||
class TokenExchanger(Protocol):
|
||||
"""Exchanges a caller token for an upstream-bound one, per the server's token_exchange config."""
|
||||
|
||||
async def exchange(
|
||||
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig
|
||||
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig, *, tenant_id: str = ""
|
||||
) -> Result[OAuthToken, CredError]: ...
|
||||
|
||||
async def invalidate(
|
||||
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig, *, tenant_id: str = ""
|
||||
) -> None: ...
|
||||
|
||||
def _cache_key(subject_token: str, config: TokenExchangeConfig) -> str:
|
||||
"""Bind the cache entry to the caller token AND the exchange config that minted it.
|
||||
|
||||
A rotated caller token, endpoint, audience, scope, client_id, secret, auth method, or
|
||||
subject_token_type all change the key, so a config change forces a fresh exchange instead of
|
||||
serving a token minted for the old config until TTL. Everything is hashed, so no secret is held
|
||||
in the key.
|
||||
def _cache_key(subject_token: str, tenant_id: str, config: TokenExchangeConfig) -> str:
|
||||
"""Bind the cache entry to the caller token, the tenant, AND the exchange config that minted it.
|
||||
|
||||
A rotated caller token, a different tenant, endpoint, audience, scope, client_id, secret, auth
|
||||
method, or subject_token_type all change the key, so two tenants behind the same opaque token
|
||||
never share an entry and a config change forces a fresh exchange instead of serving a token
|
||||
minted for the old config until TTL. Everything is hashed, so no secret is held in the key.
|
||||
"""
|
||||
secret = config.client_secret.get_secret_value() if config.client_secret else ""
|
||||
material = "\x00".join(
|
||||
(
|
||||
subject_token,
|
||||
tenant_id,
|
||||
config.token_exchange_endpoint or "",
|
||||
config.audience or "",
|
||||
config.subject_token_type,
|
||||
|
|
@ -105,11 +131,11 @@ def _cache_key(subject_token: str, config: TokenExchangeConfig) -> str:
|
|||
def _parse_expires_in(raw: object) -> int | None:
|
||||
if isinstance(raw, bool):
|
||||
return None
|
||||
if isinstance(raw, int):
|
||||
return raw
|
||||
if isinstance(raw, (int, float)):
|
||||
return int(raw)
|
||||
if isinstance(raw, str):
|
||||
try:
|
||||
return int(raw)
|
||||
return int(float(raw))
|
||||
except ValueError:
|
||||
return None
|
||||
return None
|
||||
|
|
@ -160,22 +186,25 @@ class Rfc8693TokenExchanger:
|
|||
self._expiry_buffer_seconds = expiry_buffer_seconds
|
||||
|
||||
async def exchange(
|
||||
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig
|
||||
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig, *, tenant_id: str = ""
|
||||
) -> Result[OAuthToken, CredError]:
|
||||
endpoint = config.token_exchange_endpoint
|
||||
client_id = config.client_id
|
||||
client_secret = config.client_secret
|
||||
if not endpoint or not client_id or client_secret is None:
|
||||
if not endpoint:
|
||||
# No endpoint configured and none discoverable: fail closed (412) rather than guess an IdP
|
||||
# or fall back to a weaker source. The caller's token is never sent anywhere.
|
||||
return Error(
|
||||
CredError.of_misconfigured(
|
||||
"token_exchange requires token_exchange_endpoint, client_id and client_secret"
|
||||
)
|
||||
CredError.of_precondition_required("token exchange endpoint is not configured for this server")
|
||||
)
|
||||
if not client_id or client_secret is None:
|
||||
return Error(CredError.of_misconfigured("token_exchange requires client_id and client_secret"))
|
||||
|
||||
cache_key = _cache_key(subject_token, config)
|
||||
cache_key = _cache_key(subject_token, tenant_id, config)
|
||||
server_id = server.server_id
|
||||
cached = await self._cache.get(cache_key, server_id)
|
||||
if cached is not None:
|
||||
verbose_logger.debug("MCP token exchange cache hit for server %s", server_id)
|
||||
return Ok(cached)
|
||||
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
|
|
@ -197,6 +226,9 @@ class Rfc8693TokenExchanger:
|
|||
fresh = await self._cache.get(cache_key, server_id)
|
||||
if fresh is not None:
|
||||
return fresh
|
||||
verbose_logger.debug(
|
||||
"Exchanging token for MCP server %s at %s (audience=%s)", server_id, endpoint, config.audience
|
||||
)
|
||||
body = await self._http_post(endpoint, form, client_auth.headers)
|
||||
if body is None:
|
||||
return None
|
||||
|
|
@ -204,16 +236,38 @@ class Rfc8693TokenExchanger:
|
|||
if token is None:
|
||||
return None
|
||||
await self._cache.set(cache_key, server_id, token, self._ttl_seconds(token))
|
||||
verbose_logger.info("Token exchange succeeded for MCP server %s", server_id)
|
||||
return token
|
||||
|
||||
async def reread() -> OAuthToken | None:
|
||||
return await self._cache.get(cache_key, server_id)
|
||||
|
||||
token = await self._coordinator.run(cache_key, server_id, refresh=run_exchange, reread=reread)
|
||||
try:
|
||||
token = await self._coordinator.run(cache_key, server_id, refresh=run_exchange, reread=reread)
|
||||
except SubjectTokenRejected as rejected:
|
||||
# The IdP rejected the subject token (4xx). This is non-retryable: the caller must
|
||||
# re-authenticate with the IdP, so it surfaces as a 401 (the OBO challenge), not a 503.
|
||||
return Error(CredError.of_unauthorized(str(rejected) or "subject token rejected by the IdP"))
|
||||
except TokenExchangeClientError:
|
||||
# RFC 6749 5.2 gateway-fault code (invalid_client / invalid_target / ...): the caller can't
|
||||
# fix it by re-authenticating, so surface a 500 rather than the OBO 401 challenge. The
|
||||
# specific code is logged at the edge; the user-facing summary stays generic.
|
||||
return Error(
|
||||
CredError.of_misconfigured(
|
||||
"token exchange configuration error: the gateway's credentials, audience, or scope "
|
||||
"for this server were not accepted by the IdP"
|
||||
)
|
||||
)
|
||||
if token is None:
|
||||
return Error(CredError.of_upstream_unavailable("token exchange did not return a usable access token"))
|
||||
return Ok(token)
|
||||
|
||||
async def invalidate(
|
||||
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig, *, tenant_id: str = ""
|
||||
) -> None:
|
||||
"""Drop the cached exchanged token so the next call re-exchanges (e.g. after an upstream 401)."""
|
||||
await self._cache.delete(_cache_key(subject_token, tenant_id, config), server.server_id)
|
||||
|
||||
def _token_from_body(self, body: dict[str, object]) -> OAuthToken | None:
|
||||
access_token = body.get("access_token")
|
||||
if not isinstance(access_token, str) or not access_token:
|
||||
|
|
@ -222,6 +276,9 @@ class Rfc8693TokenExchanger:
|
|||
# must fail closed rather than be minted as a bogus Bearer; an absent type defaults to Bearer.
|
||||
token_type = body.get("token_type")
|
||||
if isinstance(token_type, str) and token_type.strip().lower() != "bearer":
|
||||
verbose_logger.warning(
|
||||
"MCP token exchange returned unusable token_type %r; refusing to forward it as Bearer", token_type
|
||||
)
|
||||
return None
|
||||
# issued_token_type says what representation was minted; reject a clearly-non-access type
|
||||
# (refresh/id/saml) even if token_type claimed Bearer. access_token / jwt / absent / unknown pass.
|
||||
|
|
@ -235,5 +292,7 @@ class Rfc8693TokenExchanger:
|
|||
def _ttl_seconds(self, token: OAuthToken) -> float:
|
||||
if token.expires_at is None:
|
||||
return self._default_ttl_seconds
|
||||
lifetime = token.expires_at - self._clock()
|
||||
return max(lifetime - self._expiry_buffer_seconds, self._min_ttl_seconds)
|
||||
lifetime = max(0.0, token.expires_at - self._clock())
|
||||
# Floor at min_ttl, but never cache past the token's own expiry: a token whose remaining
|
||||
# lifetime is below the buffer (or even below min_ttl) must not be served stale upstream.
|
||||
return min(max(lifetime - self._expiry_buffer_seconds, self._min_ttl_seconds), lifetime)
|
||||
|
|
|
|||
|
|
@ -1838,6 +1838,7 @@ if MCP_AVAILABLE:
|
|||
add_prefix=True, # Always add server prefix
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
oauth2_headers=oauth2_headers,
|
||||
)
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
|
|
@ -2692,12 +2693,18 @@ if MCP_AVAILABLE:
|
|||
|
||||
# Forward named client headers to OpenAPI tool upstream requests.
|
||||
# MCPServer.extra_headers lists header names to copy from raw_headers.
|
||||
# OAuth2 M2M: never take Authorization from the caller (matches
|
||||
# _prepare_mcp_server_headers for managed MCP).
|
||||
# The strip decision is centralized in _should_strip_caller_authorization so this
|
||||
# OpenAPI/local path agrees with the managed paths: M2M and the resolver-owned modes
|
||||
# (token_exchange's raw subject token, authorization_code's stored token) must never
|
||||
# have the caller's Authorization forwarded verbatim upstream.
|
||||
forwarded_headers: Optional[Dict[str, str]] = None
|
||||
if mcp_server and mcp_server.extra_headers and raw_headers:
|
||||
normalized_raw = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)}
|
||||
skip_caller_authorization = bool(mcp_server.has_client_credentials)
|
||||
skip_caller_authorization = _should_strip_caller_authorization(
|
||||
mcp_server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
for header_name in mcp_server.extra_headers:
|
||||
if not isinstance(header_name, str):
|
||||
continue
|
||||
|
|
@ -3466,6 +3473,19 @@ if MCP_AVAILABLE:
|
|||
headers={"www-authenticate": authorization_uri},
|
||||
)
|
||||
|
||||
# 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
|
||||
raise_token_exchange_challenge,
|
||||
)
|
||||
from litellm.proxy.utils import get_server_root_path # noqa: PLC0415
|
||||
|
||||
raise_token_exchange_challenge(server, root_path=get_server_root_path())
|
||||
|
||||
# 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
|
||||
|
|
|
|||
|
|
@ -955,6 +955,7 @@ async def test_get_tools_from_mcp_servers():
|
|||
add_prefix=False,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
oauth2_headers=None,
|
||||
):
|
||||
if server.server_id == "server1_id":
|
||||
return [mock_tool_1]
|
||||
|
|
|
|||
|
|
@ -7,12 +7,12 @@ maps each CredError onto its HTTP status. These pin the parity-critical mapping
|
|||
|
||||
import base64
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
oauth_protected_resource_path,
|
||||
raise_public,
|
||||
raise_user_oauth_challenge,
|
||||
to_server_spec,
|
||||
|
|
@ -144,6 +144,22 @@ def test_token_exchange_falls_back_to_token_url_when_no_exchange_endpoint():
|
|||
assert spec.config.token_exchange_endpoint == "https://idp.example.com/token"
|
||||
|
||||
|
||||
def test_token_exchange_with_creds_but_no_endpoint_is_owned_for_fail_closed():
|
||||
# An OBO server with client credentials but no endpoint is still owned by v2 (spec, not None) so
|
||||
# it fails closed at the exchanger (412) rather than silently deferring to v1 and connecting
|
||||
# unauthenticated. The endpoint stays None for the exchanger to reject.
|
||||
spec = to_server_spec(
|
||||
_server(
|
||||
auth_type=MCPAuth.oauth2_token_exchange,
|
||||
client_id="cid",
|
||||
client_secret="csec",
|
||||
)
|
||||
)
|
||||
assert spec is not None
|
||||
assert isinstance(spec.config, TokenExchangeConfig)
|
||||
assert spec.config.token_exchange_endpoint is None
|
||||
|
||||
|
||||
def test_token_exchange_omits_audience_when_unset():
|
||||
spec = to_server_spec(
|
||||
_server(
|
||||
|
|
@ -158,6 +174,22 @@ def test_token_exchange_omits_audience_when_unset():
|
|||
assert spec.config.audience is None
|
||||
|
||||
|
||||
def test_token_exchange_empty_subject_token_type_normalizes_to_default():
|
||||
# Parity with v1: a falsy subject_token_type must not be sent verbatim to the IdP.
|
||||
spec = to_server_spec(
|
||||
_server(
|
||||
auth_type=MCPAuth.oauth2_token_exchange,
|
||||
token_exchange_endpoint="https://idp.example.com/token",
|
||||
client_id="cid",
|
||||
client_secret="csec",
|
||||
subject_token_type="",
|
||||
)
|
||||
)
|
||||
assert spec is not None
|
||||
assert isinstance(spec.config, TokenExchangeConfig)
|
||||
assert spec.config.subject_token_type == "urn:ietf:params:oauth:token-type:access_token"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"server",
|
||||
[
|
||||
|
|
@ -229,29 +261,17 @@ def test_raise_public_plain_unauthorized_has_no_challenge():
|
|||
assert exc.headers is None
|
||||
|
||||
|
||||
_ROOT_PATH = "litellm.proxy.utils.get_server_root_path"
|
||||
|
||||
|
||||
def test_raise_user_oauth_challenge_points_at_per_server_prm():
|
||||
with patch(_ROOT_PATH, return_value="/"), pytest.raises(HTTPException) as exc_info:
|
||||
raise_user_oauth_challenge(_server(alias="my-srv"))
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 401
|
||||
assert (
|
||||
exc.headers["WWW-Authenticate"] == 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/my-srv"'
|
||||
)
|
||||
|
||||
|
||||
def test_raise_user_oauth_challenge_includes_server_root_path():
|
||||
with (
|
||||
patch(_ROOT_PATH, return_value="/api/v1"),
|
||||
pytest.raises(HTTPException) as exc_info,
|
||||
):
|
||||
raise_user_oauth_challenge(_server(alias="my-srv"))
|
||||
assert (
|
||||
exc_info.value.headers["WWW-Authenticate"]
|
||||
== 'Bearer resource_metadata="/.well-known/oauth-protected-resource/api/v1/mcp/my-srv"'
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"root_path, expected_prefix",
|
||||
[
|
||||
("/", ""), # "/" means no prefix
|
||||
("", ""), # empty means no prefix
|
||||
("/api/v1", "/api/v1"), # a real root path is prepended verbatim
|
||||
],
|
||||
)
|
||||
def test_oauth_protected_resource_path_honors_root_path(root_path, expected_prefix):
|
||||
path = oauth_protected_resource_path(root_path, _server(alias="my-srv"))
|
||||
assert path == f"/.well-known/oauth-protected-resource{expected_prefix}/mcp/my-srv"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -262,7 +282,51 @@ def test_raise_user_oauth_challenge_includes_server_root_path():
|
|||
({}, "n"), # then the name field (server_id is the last fallback)
|
||||
],
|
||||
)
|
||||
def test_raise_user_oauth_challenge_name_fallback(kwargs, expected_name):
|
||||
with patch(_ROOT_PATH, return_value="/"), pytest.raises(HTTPException) as exc_info:
|
||||
raise_user_oauth_challenge(_server(**kwargs))
|
||||
assert f'/mcp/{expected_name}"' in exc_info.value.headers["WWW-Authenticate"]
|
||||
def test_oauth_protected_resource_path_name_fallback(kwargs, expected_name):
|
||||
assert oauth_protected_resource_path("/", _server(**kwargs)).endswith(f"/mcp/{expected_name}")
|
||||
|
||||
|
||||
def test_raise_user_oauth_challenge_points_at_per_server_prm():
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
raise_user_oauth_challenge(_server(alias="my-srv"), root_path="/")
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 401
|
||||
assert (
|
||||
exc.headers["WWW-Authenticate"] == 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/my-srv"'
|
||||
)
|
||||
|
||||
|
||||
def test_raise_user_oauth_challenge_includes_server_root_path():
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
raise_user_oauth_challenge(_server(alias="my-srv"), root_path="/api/v1")
|
||||
assert (
|
||||
exc_info.value.headers["WWW-Authenticate"]
|
||||
== 'Bearer resource_metadata="/.well-known/oauth-protected-resource/api/v1/mcp/my-srv"'
|
||||
)
|
||||
|
||||
|
||||
def test_raise_token_exchange_challenge_is_rfc9728_invalid_token():
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_token_exchange_challenge,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
raise_token_exchange_challenge(_server(alias="obo-srv"), root_path="/")
|
||||
exc = exc_info.value
|
||||
www = exc.headers["WWW-Authenticate"]
|
||||
assert exc.status_code == 401
|
||||
# RFC 9728 resource_metadata so the client can discover the IdP, plus RFC 6750 invalid_token.
|
||||
assert 'resource_metadata="/.well-known/oauth-protected-resource/mcp/obo-srv"' in www
|
||||
assert 'error="invalid_token"' in www
|
||||
assert "error_description=" in www
|
||||
|
||||
|
||||
def test_raise_token_exchange_challenge_includes_server_root_path():
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_token_exchange_challenge,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
raise_token_exchange_challenge(_server(alias="obo-srv"), root_path="/api/v1")
|
||||
www = exc_info.value.headers["WWW-Authenticate"]
|
||||
assert 'resource_metadata="/.well-known/oauth-protected-resource/api/v1/mcp/obo-srv"' in www
|
||||
|
|
|
|||
|
|
@ -172,12 +172,16 @@ async def test_has_user_token_false_for_a_non_per_user_mode():
|
|||
class _FakeExchanger:
|
||||
def __init__(self, result: Result[OAuthToken, CredError]) -> None:
|
||||
self._result = result
|
||||
self.calls: list[tuple[str, str]] = []
|
||||
self.calls: list[tuple[str, str, str]] = []
|
||||
self.invalidations: list[tuple[str, str, str]] = []
|
||||
|
||||
async def exchange(self, subject_token, server, config):
|
||||
self.calls.append((subject_token, server.server_id))
|
||||
async def exchange(self, subject_token, server, config, *, tenant_id=""):
|
||||
self.calls.append((subject_token, tenant_id, server.server_id))
|
||||
return self._result
|
||||
|
||||
async def invalidate(self, subject_token, server, config, *, tenant_id=""):
|
||||
self.invalidations.append((subject_token, tenant_id, server.server_id))
|
||||
|
||||
|
||||
_OBO = TokenExchangeConfig(
|
||||
token_exchange_endpoint="https://idp.example.com/token",
|
||||
|
|
@ -189,12 +193,29 @@ _OBO = TokenExchangeConfig(
|
|||
@pytest.mark.asyncio
|
||||
async def test_token_exchange_emits_the_exchanged_bearer():
|
||||
exchanger = _FakeExchanger(Ok(OAuthToken(access_token="exchanged-at")))
|
||||
subject = Subject(tenant_id="", subject_id="alice", inbound_token=SecretStr("caller-jwt"))
|
||||
subject = Subject(tenant_id="acme", subject_id="alice", inbound_token=SecretStr("caller-jwt"))
|
||||
result = await UpstreamCredentialProvider(token_exchanger=exchanger).resolve_credentials(subject, _spec(_OBO))
|
||||
assert isinstance(result, Ok)
|
||||
assert _emitted(result.ok)["Authorization"] == "Bearer exchanged-at"
|
||||
# The arm hands the unwrapped caller token to the exchanger, never the upstream.
|
||||
assert exchanger.calls == [("caller-jwt", "s")]
|
||||
# The arm hands the unwrapped caller token AND the tenant to the exchanger, never the upstream.
|
||||
assert exchanger.calls == [("caller-jwt", "acme", "s")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalidate_credentials_drops_the_exchanged_token_for_the_subject_and_tenant():
|
||||
exchanger = _FakeExchanger(Ok(OAuthToken(access_token="exchanged-at")))
|
||||
provider = UpstreamCredentialProvider(token_exchanger=exchanger)
|
||||
subject = Subject(tenant_id="acme", subject_id="alice", inbound_token=SecretStr("caller-jwt"))
|
||||
await provider.invalidate_credentials(subject, _spec(_OBO))
|
||||
assert exchanger.invalidations == [("caller-jwt", "acme", "s")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalidate_credentials_is_a_noop_without_a_caller_token():
|
||||
exchanger = _FakeExchanger(Ok(OAuthToken(access_token="never")))
|
||||
provider = UpstreamCredentialProvider(token_exchanger=exchanger)
|
||||
await provider.invalidate_credentials(Subject(tenant_id="acme", subject_id="alice"), _spec(_OBO))
|
||||
assert exchanger.invalidations == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -14,11 +14,32 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import (
|
||||
Rfc8693TokenExchanger,
|
||||
SubjectTokenRejected,
|
||||
TokenExchangeClientError,
|
||||
)
|
||||
|
||||
_HTTP_CLIENT = "litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
|
||||
|
||||
|
||||
def _client_raising_4xx(body: object):
|
||||
"""An httpx client whose POST returns a 4xx whose ``raise_for_status`` raises an HTTPStatusError
|
||||
carrying ``body`` as its JSON, so the RFC 6749 error-code classification can be driven."""
|
||||
import httpx
|
||||
|
||||
request = httpx.Request("POST", "https://idp/token")
|
||||
response = httpx.Response(400, json=body, request=request)
|
||||
|
||||
class _Resp:
|
||||
def raise_for_status(self) -> None:
|
||||
raise httpx.HTTPStatusError("bad request", request=request, response=response)
|
||||
|
||||
class _Client:
|
||||
async def post(self, url, headers, data):
|
||||
return _Resp()
|
||||
|
||||
return _Client()
|
||||
|
||||
|
||||
def test_build_token_exchanger_returns_an_exchanger():
|
||||
assert isinstance(build_token_exchanger(), Rfc8693TokenExchanger)
|
||||
|
||||
|
|
@ -53,6 +74,30 @@ async def test_post_parses_json_body_on_success():
|
|||
assert result == {"access_token": "x", "expires_in": 60}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"code", ["invalid_client", "unauthorized_client", "unsupported_grant_type", "invalid_target", "invalid_scope"]
|
||||
)
|
||||
async def test_post_maps_gateway_fault_4xx_to_client_error(code):
|
||||
# RFC 6749 5.2 gateway-fault codes must raise TokenExchangeClientError (-> 500), not the caller 401.
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_4xx({"error": code})):
|
||||
with pytest.raises(TokenExchangeClientError):
|
||||
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[{"error": "invalid_grant"}, {"error": "invalid_request"}, {}, {"error": 123}, "not-json-object"],
|
||||
ids=["invalid_grant", "invalid_request", "no_error", "non_str_error", "non_dict"],
|
||||
)
|
||||
async def test_post_maps_subject_fault_4xx_to_subject_rejected(body):
|
||||
# A subject-fault code (or an unparseable/absent error) is the caller's problem -> SubjectTokenRejected (401).
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)):
|
||||
with pytest.raises(SubjectTokenRejected):
|
||||
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("payload", [["a", "b"], "a-string", 42], ids=["list", "str", "int"])
|
||||
async def test_post_returns_none_on_non_object_json(payload):
|
||||
|
|
|
|||
|
|
@ -111,6 +111,47 @@ async def test_client_secret_post_keeps_creds_in_body_with_no_auth_header():
|
|||
assert "Authorization" not in post.headers[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_maps_idp_rejection_to_unauthorized():
|
||||
"""An IdP 4xx (surfaced as SubjectTokenRejected by the post adapter) is non-retryable: it maps
|
||||
to ``unauthorized`` (the 401 OBO challenge), not the retryable ``upstream_unavailable`` (503)."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import (
|
||||
SubjectTokenRejected,
|
||||
)
|
||||
|
||||
async def _rejecting_post(url, form, headers):
|
||||
raise SubjectTokenRejected("IdP rejected the token exchange (HTTP 400)")
|
||||
|
||||
result = await Rfc8693TokenExchanger(_rejecting_post, clock=_Clock()).exchange("bad-jwt", _SERVER, _CONFIG)
|
||||
assert isinstance(result, Error)
|
||||
assert result.error.tag == "unauthorized"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_maps_gateway_fault_to_misconfigured():
|
||||
"""A gateway-fault RFC 6749 code (invalid_client / invalid_target / ..., surfaced as
|
||||
TokenExchangeClientError) is the gateway's problem, not the caller's, so it maps to misconfigured
|
||||
(500) rather than the retryable 503 or the 401 OBO challenge the caller can't act on."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import (
|
||||
TokenExchangeClientError,
|
||||
)
|
||||
|
||||
async def _client_error_post(url, form, headers):
|
||||
raise TokenExchangeClientError("invalid_client")
|
||||
|
||||
result = await Rfc8693TokenExchanger(_client_error_post, clock=_Clock()).exchange("jwt", _SERVER, _CONFIG)
|
||||
assert isinstance(result, Error)
|
||||
assert result.error.tag == "misconfigured"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_maps_transport_failure_to_upstream_unavailable():
|
||||
"""A post returning None (5xx / network / timeout / malformed body) stays retryable: 503."""
|
||||
result = await Rfc8693TokenExchanger(_RecordingPost(None), clock=_Clock()).exchange("jwt", _SERVER, _CONFIG)
|
||||
assert isinstance(result, Error)
|
||||
assert result.error.tag == "upstream_unavailable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_caches_per_caller_token():
|
||||
post = _RecordingPost({"access_token": "x", "expires_in": 3600})
|
||||
|
|
@ -130,6 +171,43 @@ async def test_rotated_caller_token_re_exchanges():
|
|||
assert len(post.calls) == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_same_token_different_tenant_does_not_share_cache():
|
||||
# Two tenants presenting the same opaque token (e.g. a shared/service token) must not collide on
|
||||
# one cache entry: tenant_id is part of the key, so each tenant gets its own exchange.
|
||||
post = _RecordingPost({"access_token": "x", "expires_in": 3600})
|
||||
exchanger = Rfc8693TokenExchanger(post, clock=_Clock())
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG, tenant_id="acme")
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG, tenant_id="globex")
|
||||
assert len(post.calls) == 2
|
||||
# Same tenant + token still hits the cache.
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG, tenant_id="acme")
|
||||
assert len(post.calls) == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalidate_forces_re_exchange():
|
||||
post = _RecordingPost({"access_token": "x", "expires_in": 3600})
|
||||
exchanger = Rfc8693TokenExchanger(post, clock=_Clock())
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG, tenant_id="acme")
|
||||
await exchanger.invalidate("jwt", _SERVER, _CONFIG, tenant_id="acme")
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG, tenant_id="acme")
|
||||
assert len(post.calls) == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalidate_targets_only_the_matching_tenant():
|
||||
post = _RecordingPost({"access_token": "x", "expires_in": 3600})
|
||||
exchanger = Rfc8693TokenExchanger(post, clock=_Clock())
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG, tenant_id="acme")
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG, tenant_id="globex")
|
||||
await exchanger.invalidate("jwt", _SERVER, _CONFIG, tenant_id="acme")
|
||||
# globex's entry survives; only acme re-exchanges.
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG, tenant_id="globex")
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG, tenant_id="acme")
|
||||
assert len(post.calls) == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rotated_config_re_exchanges_before_ttl():
|
||||
# Same caller token + server, but the operator rotated the audience/scope: the cached token was
|
||||
|
|
@ -209,6 +287,19 @@ async def test_non_bearer_token_type_is_refused(token_type):
|
|||
assert result.error.tag == "upstream_unavailable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_bearer_token_type_is_logged():
|
||||
from unittest.mock import patch
|
||||
|
||||
post = _RecordingPost({"access_token": "x", "token_type": "N_A", "expires_in": 3600})
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger.verbose_logger"
|
||||
) as mock_logger:
|
||||
await Rfc8693TokenExchanger(post, clock=_Clock()).exchange("jwt", _SERVER, _CONFIG)
|
||||
assert mock_logger.warning.called
|
||||
assert "N_A" in repr(mock_logger.warning.call_args)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("token_type", ["Bearer", "bearer", "BEARER"])
|
||||
async def test_bearer_token_type_is_accepted_case_insensitively(token_type):
|
||||
|
|
@ -264,7 +355,6 @@ async def test_access_or_unknown_issued_token_type_is_accepted(issued_token_type
|
|||
@pytest.mark.parametrize(
|
||||
"config",
|
||||
[
|
||||
TokenExchangeConfig(client_id="c", client_secret=SecretStr("s")),
|
||||
TokenExchangeConfig(token_exchange_endpoint="https://idp/token", client_secret=SecretStr("s")),
|
||||
TokenExchangeConfig(token_exchange_endpoint="https://idp/token", client_id="c"),
|
||||
],
|
||||
|
|
@ -277,6 +367,18 @@ async def test_incomplete_config_is_misconfigured_without_hitting_idp(config):
|
|||
assert post.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_endpoint_is_precondition_required_without_hitting_idp():
|
||||
# No endpoint configured (and none discoverable): fail closed with a 412-mapped precondition
|
||||
# rather than guessing an IdP or falling back. The subject token is never POSTed anywhere.
|
||||
config = TokenExchangeConfig(client_id="c", client_secret=SecretStr("s"))
|
||||
post = _RecordingPost({"access_token": "x"})
|
||||
result = await Rfc8693TokenExchanger(post, clock=_Clock()).exchange("jwt", _spec(config), config)
|
||||
assert isinstance(result, Error)
|
||||
assert result.error.tag == "precondition_required"
|
||||
assert post.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cached_token_expires_after_its_ttl():
|
||||
clock = _Clock(1000.0)
|
||||
|
|
@ -307,17 +409,33 @@ async def test_audience_and_scope_omitted_when_unset():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_string_expires_in_is_honored():
|
||||
@pytest.mark.parametrize("expires_in", ["120", 120.0, "120.0"], ids=["str", "float", "str_float"])
|
||||
async def test_numeric_expires_in_is_honored(expires_in):
|
||||
# A JSON int/float/numeric-string expires_in must drive the TTL, not fall back to the default.
|
||||
clock = _Clock(1000.0)
|
||||
# "120" parsed as int -> ttl max(120-60, 10) = 60 -> cached until 1060.
|
||||
post = _RecordingPost({"access_token": "x", "expires_in": "120"})
|
||||
exchanger = Rfc8693TokenExchanger(post, clock=clock)
|
||||
post = _RecordingPost({"access_token": "x", "expires_in": expires_in})
|
||||
exchanger = Rfc8693TokenExchanger(post, clock=clock) # ttl = max(120-60, 10) = 60 -> until 1060
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG)
|
||||
clock.now = 1061.0
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG)
|
||||
assert len(post.calls) == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_short_lived_token_is_not_cached_past_its_expiry():
|
||||
# expires_in (5s) below the buffer/min floor must NOT be served stale: cache only until expiry.
|
||||
clock = _Clock(1000.0)
|
||||
post = _RecordingPost({"access_token": "x", "expires_in": 5})
|
||||
exchanger = Rfc8693TokenExchanger(post, clock=clock)
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG)
|
||||
clock.now = 1004.0 # within the 5s lifetime -> still cached
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG)
|
||||
assert len(post.calls) == 1
|
||||
clock.now = 1006.0 # past expiry -> re-exchange, not a stale bearer
|
||||
await exchanger.exchange("jwt", _SERVER, _CONFIG)
|
||||
assert len(post.calls) == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
|
|
|
|||
|
|
@ -96,9 +96,7 @@ async def test_authorize_endpoint_includes_response_type():
|
|||
mock_request.headers = {}
|
||||
|
||||
# Mock the encryption functions to avoid needing a signing key
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper"
|
||||
) as mock_encrypt:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper") as mock_encrypt:
|
||||
mock_encrypt.return_value = "mocked_encrypted_state"
|
||||
|
||||
# Call authorize endpoint
|
||||
|
|
@ -160,9 +158,7 @@ async def test_authorize_endpoint_preserves_existing_query_params():
|
|||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper"
|
||||
) as mock_encrypt:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper") as mock_encrypt:
|
||||
mock_encrypt.return_value = "mocked_encrypted_state"
|
||||
|
||||
response = await authorize(
|
||||
|
|
@ -176,9 +172,7 @@ async def test_authorize_endpoint_preserves_existing_query_params():
|
|||
location = response.headers["location"]
|
||||
|
||||
# Must NOT have double '?' — existing params must be merged correctly
|
||||
assert (
|
||||
location.count("?") == 1
|
||||
), f"Expected exactly one '?' in URL but got {location.count('?')}: {location}"
|
||||
assert location.count("?") == 1, f"Expected exactly one '?' in URL but got {location.count('?')}: {location}"
|
||||
assert "tenant=system" in location
|
||||
assert "client_id=test_client_id" in location
|
||||
assert "response_type=code" in location
|
||||
|
|
@ -228,9 +222,7 @@ async def test_authorize_endpoint_forwards_pkce_parameters():
|
|||
mock_request.headers = {}
|
||||
|
||||
# Mock the encryption function
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper"
|
||||
) as mock_encrypt:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper") as mock_encrypt:
|
||||
mock_encrypt.return_value = "mocked_encrypted_state_with_pkce"
|
||||
|
||||
# Call authorize endpoint with PKCE parameters
|
||||
|
|
@ -338,10 +330,7 @@ async def test_token_endpoint_forwards_code_verifier():
|
|||
# Check the data parameter includes code_verifier
|
||||
assert call_args[1]["data"]["code_verifier"] == "test_code_verifier_from_client"
|
||||
assert call_args[1]["data"]["code"] == "4/test_authorization_code"
|
||||
assert (
|
||||
call_args[1]["data"]["client_id"]
|
||||
== "669428968603-test.apps.googleusercontent.com"
|
||||
)
|
||||
assert call_args[1]["data"]["client_id"] == "669428968603-test.apps.googleusercontent.com"
|
||||
assert call_args[1]["data"]["client_secret"] == "GOCSPX-test_secret"
|
||||
assert call_args[1]["data"]["grant_type"] == "authorization_code"
|
||||
|
||||
|
|
@ -428,9 +417,7 @@ async def test_register_client_returns_existing_server_credentials():
|
|||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
|
||||
new=AsyncMock(return_value={}),
|
||||
):
|
||||
result = await register_client(
|
||||
request=mock_request, mcp_server_name=oauth2_server.server_name
|
||||
)
|
||||
result = await register_client(request=mock_request, mcp_server_name=oauth2_server.server_name)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
|
@ -505,9 +492,7 @@ async def test_register_client_remote_registration_success():
|
|||
return_value=mock_async_client,
|
||||
),
|
||||
):
|
||||
response = await register_client(
|
||||
request=mock_request, mcp_server_name=oauth2_server.server_name
|
||||
)
|
||||
response = await register_client(request=mock_request, mcp_server_name=oauth2_server.server_name)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
|
@ -524,14 +509,9 @@ async def test_register_client_remote_registration_success():
|
|||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
assert call_args.kwargs["json"]["redirect_uris"] == [
|
||||
"https://proxy.litellm.example/callback"
|
||||
]
|
||||
assert call_args.kwargs["json"]["redirect_uris"] == ["https://proxy.litellm.example/callback"]
|
||||
assert call_args.kwargs["json"]["grant_types"] == request_payload["grant_types"]
|
||||
assert (
|
||||
call_args.kwargs["json"]["token_endpoint_auth_method"]
|
||||
== request_payload["token_endpoint_auth_method"]
|
||||
)
|
||||
assert call_args.kwargs["json"]["token_endpoint_auth_method"] == request_payload["token_endpoint_auth_method"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1120,9 +1100,7 @@ async def test_authorize_endpoint_respects_x_forwarded_proto():
|
|||
mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy
|
||||
|
||||
# Mock the encryption functions
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper"
|
||||
) as mock_encrypt:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper") as mock_encrypt:
|
||||
mock_encrypt.return_value = "mocked_encrypted_state"
|
||||
|
||||
# Call authorize endpoint
|
||||
|
|
@ -1217,10 +1195,7 @@ async def test_token_endpoint_respects_x_forwarded_proto():
|
|||
|
||||
# Verify that the redirect_uri sent to the provider uses HTTPS
|
||||
call_args = mock_async_client.post.call_args
|
||||
assert (
|
||||
call_args[1]["data"]["redirect_uri"]
|
||||
== "https://litellm-proxy.example.com/callback"
|
||||
)
|
||||
assert call_args[1]["data"]["redirect_uri"] == "https://litellm-proxy.example.com/callback"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1272,9 +1247,7 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto():
|
|||
)
|
||||
|
||||
# Verify response uses HTTPS URLs
|
||||
assert response["authorization_servers"][0].startswith(
|
||||
"https://litellm.example.com/"
|
||||
)
|
||||
assert response["authorization_servers"][0].startswith("https://litellm.example.com/")
|
||||
assert response["scopes_supported"] == oauth2_server.scopes
|
||||
|
||||
|
||||
|
|
@ -1421,9 +1394,7 @@ async def test_authorize_endpoint_respects_x_forwarded_host():
|
|||
}
|
||||
|
||||
# Mock the encryption functions
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper"
|
||||
) as mock_encrypt:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper") as mock_encrypt:
|
||||
mock_encrypt.return_value = "mocked_encrypted_state"
|
||||
|
||||
# Call authorize endpoint
|
||||
|
|
@ -1440,8 +1411,7 @@ async def test_authorize_endpoint_respects_x_forwarded_host():
|
|||
|
||||
# The redirect_uri parameter should use the external URL
|
||||
assert (
|
||||
"redirect_uri=https%3A%2F%2Fproxy.example.com%2Fgithub%2Fmcp%2Fcallback"
|
||||
in location
|
||||
"redirect_uri=https%3A%2F%2Fproxy.example.com%2Fgithub%2Fmcp%2Fcallback" in location
|
||||
or "redirect_uri=https://proxy.example.com/github/mcp/callback" in location
|
||||
)
|
||||
|
||||
|
|
@ -1522,10 +1492,7 @@ async def test_token_endpoint_respects_x_forwarded_host():
|
|||
|
||||
# Verify that the redirect_uri sent to the provider uses the external URL
|
||||
call_args = mock_async_client.post.call_args
|
||||
assert (
|
||||
call_args[1]["data"]["redirect_uri"]
|
||||
== "https://proxy.example.com/github/mcp/callback"
|
||||
)
|
||||
assert call_args[1]["data"]["redirect_uri"] == "https://proxy.example.com/github/mcp/callback"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -1733,9 +1700,7 @@ def test_get_request_base_url_comprehensive(
|
|||
),
|
||||
],
|
||||
)
|
||||
def test_get_request_base_url_xff_trust_gate(
|
||||
general_settings, direct_ip, expect_xff_honoured
|
||||
):
|
||||
def test_get_request_base_url_xff_trust_gate(general_settings, direct_ip, expect_xff_honoured):
|
||||
"""Verify the X-Forwarded-* trust gate.
|
||||
|
||||
With XFF poisoning attempted, the helper must return either the literal
|
||||
|
|
@ -1813,12 +1778,10 @@ def test_xff_misconfig_warning_emitted_once(caplog):
|
|||
for _ in range(3):
|
||||
get_request_base_url(mock_request)
|
||||
|
||||
matching = [
|
||||
rec for rec in caplog.records if "mcp_trusted_proxy_ranges" in rec.getMessage()
|
||||
]
|
||||
assert (
|
||||
len(matching) == 1
|
||||
), f"expected exactly one warning, got {len(matching)}: {[r.getMessage() for r in matching]}"
|
||||
matching = [rec for rec in caplog.records if "mcp_trusted_proxy_ranges" in rec.getMessage()]
|
||||
assert len(matching) == 1, (
|
||||
f"expected exactly one warning, got {len(matching)}: {[r.getMessage() for r in matching]}"
|
||||
)
|
||||
|
||||
|
||||
def test_get_request_base_url_honors_proxy_base_url_env(monkeypatch):
|
||||
|
|
@ -1849,9 +1812,7 @@ def test_get_request_base_url_honors_proxy_base_url_env(monkeypatch):
|
|||
assert get_request_base_url(mock_request) == "https://litellm.example.com"
|
||||
|
||||
|
||||
def test_validate_trusted_redirect_uri_logs_diagnostic_on_rejection(
|
||||
caplog, monkeypatch
|
||||
):
|
||||
def test_validate_trusted_redirect_uri_logs_diagnostic_on_rejection(caplog, monkeypatch):
|
||||
try:
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
|
|
@ -1898,8 +1859,7 @@ def test_validate_trusted_redirect_uri_logs_diagnostic_on_rejection(
|
|||
|
||||
matching = [r for r in caplog.records if "rejecting redirect_uri" in r.getMessage()]
|
||||
assert len(matching) == 1, (
|
||||
"expected exactly one diagnostic warning, got "
|
||||
f"{[r.getMessage() for r in caplog.records]}"
|
||||
f"expected exactly one diagnostic warning, got {[r.getMessage() for r in caplog.records]}"
|
||||
)
|
||||
msg = matching[0].getMessage()
|
||||
assert "https://litellm.example.com/ui/mcp/oauth/callback" in msg
|
||||
|
|
@ -1918,9 +1878,7 @@ def test_validate_trusted_redirect_uri_logs_diagnostic_on_rejection(
|
|||
"not a url at all",
|
||||
],
|
||||
)
|
||||
def test_get_request_base_url_rejects_malformed_proxy_base_url(
|
||||
bad_value, monkeypatch, caplog
|
||||
):
|
||||
def test_get_request_base_url_rejects_malformed_proxy_base_url(bad_value, monkeypatch, caplog):
|
||||
try:
|
||||
from fastapi import Request
|
||||
|
||||
|
|
@ -1950,26 +1908,16 @@ def test_get_request_base_url_rejects_malformed_proxy_base_url(
|
|||
result = get_request_base_url(mock_request)
|
||||
|
||||
assert result == "http://litellm-internal:4000", (
|
||||
f"malformed PROXY_BASE_URL={bad_value!r} should be ignored, " f"got {result!r}"
|
||||
f"malformed PROXY_BASE_URL={bad_value!r} should be ignored, got {result!r}"
|
||||
)
|
||||
matching = [
|
||||
r
|
||||
for r in caplog.records
|
||||
if "PROXY_BASE_URL" in r.getMessage() and "ignored" in r.getMessage()
|
||||
]
|
||||
matching = [r for r in caplog.records if "PROXY_BASE_URL" in r.getMessage() and "ignored" in r.getMessage()]
|
||||
assert len(matching) == 1, (
|
||||
"expected one diagnostic for malformed PROXY_BASE_URL, got "
|
||||
f"{[r.getMessage() for r in caplog.records]}"
|
||||
)
|
||||
assert (
|
||||
repr(bad_value) in matching[0].getMessage()
|
||||
or bad_value in matching[0].getMessage()
|
||||
f"expected one diagnostic for malformed PROXY_BASE_URL, got {[r.getMessage() for r in caplog.records]}"
|
||||
)
|
||||
assert repr(bad_value) in matching[0].getMessage() or bad_value in matching[0].getMessage()
|
||||
|
||||
|
||||
def test_get_request_base_url_malformed_proxy_base_url_warning_is_one_shot(
|
||||
monkeypatch, caplog
|
||||
):
|
||||
def test_get_request_base_url_malformed_proxy_base_url_warning_is_one_shot(monkeypatch, caplog):
|
||||
try:
|
||||
from fastapi import Request
|
||||
|
||||
|
|
@ -1999,14 +1947,8 @@ def test_get_request_base_url_malformed_proxy_base_url_warning_is_one_shot(
|
|||
for _ in range(5):
|
||||
get_request_base_url(mock_request)
|
||||
|
||||
matching = [
|
||||
r
|
||||
for r in caplog.records
|
||||
if "PROXY_BASE_URL" in r.getMessage() and "ignored" in r.getMessage()
|
||||
]
|
||||
assert (
|
||||
len(matching) == 1
|
||||
), f"expected exactly one warning across 5 calls, got {len(matching)}"
|
||||
matching = [r for r in caplog.records if "PROXY_BASE_URL" in r.getMessage() and "ignored" in r.getMessage()]
|
||||
assert len(matching) == 1, f"expected exactly one warning across 5 calls, got {len(matching)}"
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
|
|
@ -2219,12 +2161,8 @@ async def test_authorize_root_fails_with_multiple_oauth2_servers():
|
|||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
server1 = _create_oauth2_server(
|
||||
server_id="server1", name="server1", server_name="server1", alias="server1"
|
||||
)
|
||||
server2 = _create_oauth2_server(
|
||||
server_id="server2", name="server2", server_name="server2", alias="server2"
|
||||
)
|
||||
server1 = _create_oauth2_server(server_id="server1", name="server1", server_name="server1", alias="server1")
|
||||
server2 = _create_oauth2_server(server_id="server2", name="server2", server_name="server2", alias="server2")
|
||||
global_mcp_server_manager.registry[server1.server_id] = server1
|
||||
global_mcp_server_manager.registry[server2.server_id] = server2
|
||||
|
||||
|
|
@ -2582,9 +2520,7 @@ async def test_oauth_callback_redirects_with_state():
|
|||
"client_redirect_uri": "http://localhost:3000/ui/mcp/oauth/callback",
|
||||
}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
|
||||
) as mock_decode:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash") as mock_decode:
|
||||
mock_decode.return_value = mock_state_data
|
||||
|
||||
# Call callback endpoint with code and state
|
||||
|
|
@ -2596,10 +2532,7 @@ async def test_oauth_callback_redirects_with_state():
|
|||
|
||||
# Should redirect to the client callback URL with code and original state
|
||||
assert response.status_code == 302
|
||||
assert (
|
||||
"http://localhost:3000/ui/mcp/oauth/callback"
|
||||
in response.headers["location"]
|
||||
)
|
||||
assert "http://localhost:3000/ui/mcp/oauth/callback" in response.headers["location"]
|
||||
assert "code=test_authorization_code_12345" in response.headers["location"]
|
||||
assert "state=test-uuid-state-123" in response.headers["location"]
|
||||
|
||||
|
|
@ -2619,17 +2552,13 @@ async def test_oauth_callback_preserves_client_redirect_uri_query():
|
|||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
|
||||
) as mock_decode:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash") as mock_decode:
|
||||
mock_decode.return_value = {
|
||||
"base_url": "http://localhost:3000/ui/mcp/oauth/callback",
|
||||
"original_state": "test-uuid-state-123",
|
||||
"code_challenge": "test_challenge",
|
||||
"code_challenge_method": "S256",
|
||||
"client_redirect_uri": (
|
||||
"http://localhost:3000/ui/mcp/oauth/callback?session=abc"
|
||||
),
|
||||
"client_redirect_uri": ("http://localhost:3000/ui/mcp/oauth/callback?session=abc"),
|
||||
}
|
||||
|
||||
response = await callback(
|
||||
|
|
@ -2657,9 +2586,7 @@ async def test_oauth_callback_handles_invalid_state():
|
|||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
# Mock state decoding to raise an exception
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
|
||||
) as mock_decode:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash") as mock_decode:
|
||||
mock_decode.side_effect = Exception("Failed to decrypt state")
|
||||
|
||||
# Call callback endpoint with invalid state
|
||||
|
|
@ -2682,9 +2609,7 @@ async def test_oauth_callback_accepts_same_origin_ui_redirect():
|
|||
callback,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
|
||||
) as mock_decode:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash") as mock_decode:
|
||||
mock_decode.return_value = {
|
||||
"base_url": "https://proxy.example.com/ui/mcp/oauth/callback",
|
||||
"original_state": "state-123",
|
||||
|
|
@ -2700,10 +2625,7 @@ async def test_oauth_callback_accepts_same_origin_ui_redirect():
|
|||
)
|
||||
|
||||
assert response.status_code == 302
|
||||
assert (
|
||||
"https://proxy.example.com/ui/mcp/oauth/callback"
|
||||
in response.headers["location"]
|
||||
)
|
||||
assert "https://proxy.example.com/ui/mcp/oauth/callback" in response.headers["location"]
|
||||
assert "code=auth-code-123" in response.headers["location"]
|
||||
assert "state=state-123" in response.headers["location"]
|
||||
|
||||
|
|
@ -2740,9 +2662,7 @@ async def test_oauth_authorize_includes_scopes_from_server_config():
|
|||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper"
|
||||
) as mock_encrypt:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper") as mock_encrypt:
|
||||
mock_encrypt.return_value = "encrypted_state"
|
||||
|
||||
# Call authorize without explicit scope parameter
|
||||
|
|
@ -2762,8 +2682,7 @@ async def test_oauth_authorize_includes_scopes_from_server_config():
|
|||
assert response.status_code in (307, 302)
|
||||
redirect_url = response.headers["location"]
|
||||
assert (
|
||||
"scope=api+read_user+ai_workflows" in redirect_url
|
||||
or "scope=api%20read_user%20ai_workflows" in redirect_url
|
||||
"scope=api+read_user+ai_workflows" in redirect_url or "scope=api%20read_user%20ai_workflows" in redirect_url
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -2798,9 +2717,7 @@ async def test_oauth_authorize_prefers_request_scope_over_server_config():
|
|||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper"
|
||||
) as mock_encrypt:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper") as mock_encrypt:
|
||||
mock_encrypt.return_value = "encrypted_state"
|
||||
|
||||
# Call authorize WITH explicit scope parameter
|
||||
|
|
@ -2820,8 +2737,7 @@ async def test_oauth_authorize_prefers_request_scope_over_server_config():
|
|||
assert response.status_code in (307, 302)
|
||||
redirect_url = response.headers["location"]
|
||||
assert (
|
||||
"scope=custom_scope1+custom_scope2" in redirect_url
|
||||
or "scope=custom_scope1%20custom_scope2" in redirect_url
|
||||
"scope=custom_scope1+custom_scope2" in redirect_url or "scope=custom_scope1%20custom_scope2" in redirect_url
|
||||
)
|
||||
assert "default_scope" not in redirect_url
|
||||
|
||||
|
|
@ -3078,9 +2994,7 @@ async def test_callback_revalidates_loopback_on_decoded_base_url():
|
|||
callback,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
|
||||
) as mock_decode:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash") as mock_decode:
|
||||
mock_decode.return_value = {
|
||||
"base_url": "https://attacker.example.com/cb",
|
||||
"original_state": "s",
|
||||
|
|
@ -3104,9 +3018,7 @@ async def test_callback_revalidates_loopback_on_decoded_client_redirect_uri():
|
|||
callback,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
|
||||
) as mock_decode:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash") as mock_decode:
|
||||
mock_decode.return_value = {
|
||||
"base_url": "http://localhost:3000/cb",
|
||||
"original_state": "s",
|
||||
|
|
@ -3130,9 +3042,7 @@ async def test_callback_rejects_state_missing_redirect_uri():
|
|||
callback,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
|
||||
) as mock_decode:
|
||||
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash") as mock_decode:
|
||||
mock_decode.return_value = {
|
||||
"original_state": "s",
|
||||
"code_challenge": None,
|
||||
|
|
@ -3261,9 +3171,7 @@ async def test_token_exchange_omits_expires_in_when_upstream_omits_it():
|
|||
rotation) returns no ``expires_in``. The exchange must mirror that and omit
|
||||
``expires_in`` rather than fabricate a 1-hour TTL, so the stored credential
|
||||
is treated as non-expiring instead of dying after an hour."""
|
||||
body = await _exchange_with_upstream_token_response(
|
||||
{"access_token": "tok", "token_type": "Bearer"}
|
||||
)
|
||||
body = await _exchange_with_upstream_token_response({"access_token": "tok", "token_type": "Bearer"})
|
||||
assert "expires_in" not in body
|
||||
|
||||
|
||||
|
|
@ -3277,6 +3185,120 @@ async def test_token_exchange_passes_through_upstream_expires_in():
|
|||
assert body["expires_in"] == 43200
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# OBO (token_exchange) Protected Resource Metadata: discovery must name the
|
||||
# JWT-auth issuer the client SSOs with, not the gateway.
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
_OBO_RESOURCE = "https://litellm.example.com/mcp/obo_mcp"
|
||||
_PATCH_ISSUERS = "litellm.proxy._experimental.mcp_server.discoverable_endpoints._jwt_auth_issuers"
|
||||
|
||||
|
||||
def _obo_server(scopes=None):
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
return MCPServer(
|
||||
server_id="obo_mcp",
|
||||
name="obo_mcp",
|
||||
server_name="obo_mcp",
|
||||
alias="obo_mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2_token_exchange,
|
||||
scopes=scopes,
|
||||
)
|
||||
|
||||
|
||||
def test_obo_protected_resource_response_names_jwt_issuers():
|
||||
"""An OBO server's PRM points authorization_servers at the configured JWT issuers (the IdP that
|
||||
mints and validates the subject token), with the gateway resource echoed back."""
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_obo_protected_resource_response,
|
||||
)
|
||||
|
||||
with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]):
|
||||
response = _obo_protected_resource_response(_obo_server(scopes=["read"]), _OBO_RESOURCE)
|
||||
assert response == {
|
||||
"authorization_servers": ["https://idp.example.com"],
|
||||
"resource": _OBO_RESOURCE,
|
||||
"scopes_supported": ["read"],
|
||||
}
|
||||
|
||||
|
||||
def test_obo_protected_resource_response_scopes_default_empty():
|
||||
"""A scopeless OBO server reports scopes_supported as [] rather than None."""
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_obo_protected_resource_response,
|
||||
)
|
||||
|
||||
with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]):
|
||||
response = _obo_protected_resource_response(_obo_server(scopes=None), _OBO_RESOURCE)
|
||||
assert response["scopes_supported"] == []
|
||||
|
||||
|
||||
def test_obo_protected_resource_response_falls_back_when_no_issuer():
|
||||
"""With no JWT issuer configured, the OBO branch returns None so the caller falls back to the
|
||||
gateway-default PRM (discovery still works, it just can't name the IdP)."""
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_obo_protected_resource_response,
|
||||
)
|
||||
|
||||
with patch(_PATCH_ISSUERS, return_value=[]):
|
||||
assert _obo_protected_resource_response(_obo_server(), _OBO_RESOURCE) is None
|
||||
|
||||
|
||||
def test_obo_protected_resource_response_ignores_non_obo_server():
|
||||
"""Non-OBO servers are not handled by this branch (returns None -> gateway default)."""
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_obo_protected_resource_response,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
oauth2_server = MCPServer(
|
||||
server_id="oauth2_mcp",
|
||||
name="oauth2_mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
)
|
||||
assert _obo_protected_resource_response(oauth2_server, _OBO_RESOURCE) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_oauth_protected_resource_response_obo_end_to_end():
|
||||
"""End to end through the response builder: an OBO server's PRM advertises the JWT issuer as
|
||||
authorization_servers, proving the extracted branch is wired into the public discovery path."""
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_build_oauth_protected_resource_response,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
global_mcp_server_manager.registry["obo_mcp"] = _obo_server(scopes=["read"])
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
try:
|
||||
with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]):
|
||||
response = await _build_oauth_protected_resource_response(
|
||||
request=mock_request,
|
||||
mcp_server_name="obo_mcp",
|
||||
use_standard_pattern=True,
|
||||
)
|
||||
assert response["authorization_servers"] == ["https://idp.example.com"]
|
||||
assert response["resource"] == "https://litellm.example.com/mcp/obo_mcp"
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
||||
def _token_request(headers):
|
||||
"""A real Starlette request with case-insensitive headers (matches production)."""
|
||||
from starlette.requests import Request
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -294,15 +294,9 @@ async def test_stale_mcp_session_id_is_stripped():
|
|||
|
||||
# Verify the mcp-session-id header was stripped
|
||||
header_names = [k for k, v in captured_scope.get("headers", [])]
|
||||
assert (
|
||||
b"mcp-session-id" not in header_names
|
||||
), "Stale mcp-session-id header should have been stripped from the scope"
|
||||
assert (
|
||||
stateless_handle_request.called
|
||||
), "Stale non-initialize requests should route stateless"
|
||||
assert (
|
||||
not stateful_handle_request.called
|
||||
), "Stale non-initialize requests should not route stateful"
|
||||
assert b"mcp-session-id" not in header_names, "Stale mcp-session-id header should have been stripped from the scope"
|
||||
assert stateless_handle_request.called, "Stale non-initialize requests should route stateless"
|
||||
assert not stateful_handle_request.called, "Stale non-initialize requests should not route stateful"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -366,9 +360,7 @@ async def test_delete_stale_mcp_session_returns_success():
|
|||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
|
||||
# Verify session manager was NOT called (request was handled early)
|
||||
assert (
|
||||
not mock_handle_request.called
|
||||
), "Session manager should not be called for DELETE on non-existent session"
|
||||
assert not mock_handle_request.called, "Session manager should not be called for DELETE on non-existent session"
|
||||
|
||||
# Verify a success response was sent
|
||||
assert send.called, "A response should have been sent"
|
||||
|
|
@ -523,9 +515,7 @@ async def test_valid_mcp_session_id_is_preserved():
|
|||
|
||||
# Verify the mcp-session-id header was preserved
|
||||
header_names = [k for k, v in captured_scope.get("headers", [])]
|
||||
assert (
|
||||
b"mcp-session-id" in header_names
|
||||
), "Valid mcp-session-id header should have been preserved"
|
||||
assert b"mcp-session-id" in header_names, "Valid mcp-session-id header should have been preserved"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -967,6 +957,94 @@ async def test_handle_streamable_http_mcp_delegated_server_without_token_returns
|
|||
challenge = exc_info.value.headers["www-authenticate"]
|
||||
assert "resource_metadata=" in challenge
|
||||
assert "authorization_uri=" not in challenge
|
||||
assert (
|
||||
"/.well-known/oauth-protected-resource/delegated_oauth_server/mcp" in challenge
|
||||
assert "/.well-known/oauth-protected-resource/delegated_oauth_server/mcp" in challenge
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_streamable_http_mcp_token_exchange_without_subject_returns_preemptive_resource_metadata_401():
|
||||
"""An ``oauth2_token_exchange`` (OBO) server with no caller subject token must fail fast at
|
||||
connect with a 401 carrying the RFC 9728 ``resource_metadata`` + RFC 6750 ``invalid_token``
|
||||
challenge, so the client discovers the IdP and retries with a subject token. A tool-call-time
|
||||
401 would be wrapped into a JSON-RPC error and the WWW-Authenticate lost, so this preemptive
|
||||
challenge is what drives the discovery flow."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
handle_streamable_http_mcp,
|
||||
session_manager_stateful,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp/obo_server",
|
||||
"_original_path": "/mcp/obo_server",
|
||||
"scheme": "https",
|
||||
"query_string": b"",
|
||||
"root_path": "",
|
||||
"server": ("litellm.example.com", 443),
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
(b"host", b"litellm.example.com"),
|
||||
],
|
||||
}
|
||||
receive = AsyncMock(
|
||||
return_value={
|
||||
"type": "http.request",
|
||||
"body": b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}',
|
||||
"more_body": False,
|
||||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth.user_id = None
|
||||
obo_server = MagicMock()
|
||||
obo_server.auth_type = MCPAuth.oauth2_token_exchange
|
||||
obo_server.alias = None
|
||||
obo_server.server_name = "obo_server"
|
||||
obo_server.name = "obo_server"
|
||||
obo_server.server_id = "obo-server"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(user_auth, None, ["obo_server"], None, None, None),
|
||||
),
|
||||
patch("litellm.proxy._experimental.mcp_server.server.set_auth_context"),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
|
||||
True,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session",
|
||||
new_callable=AsyncMock,
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=obo_server,
|
||||
),
|
||||
patch.object(
|
||||
session_manager_stateful,
|
||||
"handle_request",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handle_request,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
|
||||
assert mock_handle_request.await_count == 0
|
||||
assert exc_info.value.status_code == 401
|
||||
headers = {k.lower(): v for k, v in (exc_info.value.headers or {}).items()}
|
||||
challenge = headers["www-authenticate"]
|
||||
# Structural invariants only: the exact root-path prefix is exercised in the adapter's
|
||||
# oauth_protected_resource_path unit test, so this handler test stays hermetic w.r.t.
|
||||
# SERVER_ROOT_PATH (which other tests in the shard may have left set in the environment).
|
||||
assert "resource_metadata=" in challenge
|
||||
assert "/.well-known/oauth-protected-resource" in challenge
|
||||
assert challenge.split('resource_metadata="', 1)[1].split('"', 1)[0].endswith("/mcp/obo_server")
|
||||
assert 'error="invalid_token"' in challenge
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue