mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
fix(mcp): support client_secret_basic for upstream OAuth token endpoints (#31635)
The MCP gateway authenticated to upstream OAuth token endpoints only with
client_secret_post (client_secret placed in the POST body). Providers that
require HTTP Basic client authentication (client_secret_basic, the OIDC
default) reject that with invalid_client, which surfaced as a 500 on the
/<server>/token exchange and broke both the initial authorization_code
exchange and refresh.
Add a per-server token_endpoint_auth_method ("client_secret_basic" |
"client_secret_post") and a single helper that builds the right headers and
body for the configured method, then route every upstream token-endpoint POST
through it: the inbound exchange and refresh in discoverable_endpoints, the v1
per-user refresh in db, the v2 authorization_code refresher, the M2M
client_credentials fetch, and the RFC 8693 token exchange. The default stays
client_secret_post so existing servers are unaffected; basic sends
Authorization: Basic base64(form-urlencode(client_id):form-urlencode(secret))
per RFC 6749 section 2.3.1 and omits the secret from the body.
client_secret_basic is a confidential-client method, so a server configured for
it with a missing client_id/secret raises rather than silently downgrading to a
body request (no-silent-fallback); the inbound endpoint maps that to a 400 and
the refresh paths to a failed-refresh / needs-reauth. A secretless client_id
under the default method stays valid for public clients authenticating with PKCE.
Resolves LIT-4091
(cherry picked from commit 7baf25526f)
This commit is contained in:
parent
0ade44f4da
commit
2578d9557b
17 changed files with 601 additions and 35 deletions
|
|
@ -0,0 +1,78 @@
|
|||
"""Client authentication for OAuth 2.0 token-endpoint requests (RFC 6749 section 2.3.1).
|
||||
|
||||
A confidential MCP upstream may require ``client_secret_basic`` (HTTP Basic, the OIDC
|
||||
default) or ``client_secret_post`` (credentials in the form body). Every token-endpoint
|
||||
POST in the MCP gateway builds its client authentication here so the two methods are
|
||||
applied identically across the inbound exchange, the refresh grants, the M2M
|
||||
client_credentials fetch, and RFC 8693 token exchange. The default is
|
||||
``client_secret_post`` so servers that never set ``token_endpoint_auth_method`` keep
|
||||
their current behavior.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
from dataclasses import dataclass
|
||||
from urllib.parse import quote_plus
|
||||
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPTokenEndpointAuthMethod
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TokenEndpointClientAuth:
|
||||
headers: dict[str, str]
|
||||
body: dict[str, str]
|
||||
|
||||
|
||||
class TokenEndpointAuthConfigError(ValueError):
|
||||
"""``client_secret_basic`` is configured but the client credentials needed for it are missing.
|
||||
|
||||
Subclasses ``ValueError`` so existing call sites that already guard missing credentials with
|
||||
``except ValueError`` / ``except Exception`` keep mapping it to their own failure contract.
|
||||
"""
|
||||
|
||||
|
||||
def normalize_token_endpoint_auth_method(
|
||||
value: object,
|
||||
) -> MCPTokenEndpointAuthMethod | None:
|
||||
"""Narrow an untyped (DB/JSON-sourced) value to the auth-method literal, else ``None``."""
|
||||
if value == "client_secret_basic":
|
||||
return "client_secret_basic"
|
||||
if value == "client_secret_post":
|
||||
return "client_secret_post"
|
||||
return None
|
||||
|
||||
|
||||
def build_token_endpoint_client_auth(
|
||||
*,
|
||||
auth_method: MCPTokenEndpointAuthMethod | None,
|
||||
client_id: str | None,
|
||||
client_secret: str | None,
|
||||
) -> TokenEndpointClientAuth:
|
||||
"""Return the headers and body fields that authenticate the client to the token endpoint.
|
||||
|
||||
``client_secret_basic`` is a confidential-client method, so it requires both ``client_id`` and
|
||||
``client_secret`` and raises ``TokenEndpointAuthConfigError`` when either is missing rather than
|
||||
silently degrading to a weaker request (RFC 6749 section 2.3.1; matches the "absent credential
|
||||
must surface, never fall sideways" rule). It sends an HTTP Basic ``Authorization`` header and
|
||||
keeps the credentials out of the body. Any other method (including ``None``, the default) is the
|
||||
``client_secret_post`` path: it places whichever of ``client_id`` / ``client_secret`` are present
|
||||
into the body, so a secretless client_id (a public client authenticating with PKCE) stays valid.
|
||||
"""
|
||||
if auth_method == "client_secret_basic":
|
||||
if not client_id or not client_secret:
|
||||
raise TokenEndpointAuthConfigError(
|
||||
"token_endpoint_auth_method=client_secret_basic requires both client_id and client_secret"
|
||||
)
|
||||
# RFC 6749 section 2.3.1: form-urlencode each value before joining with ':' so a
|
||||
# client_id/secret containing reserved characters (':', '+', '%', ...) is transmitted intact.
|
||||
userpass = f"{quote_plus(client_id)}:{quote_plus(client_secret)}"
|
||||
encoded = base64.b64encode(userpass.encode()).decode()
|
||||
return TokenEndpointClientAuth(headers={"Authorization": f"Basic {encoded}"}, body={})
|
||||
return TokenEndpointClientAuth(
|
||||
headers={},
|
||||
body={
|
||||
**({"client_id": client_id} if client_id else {}),
|
||||
**({"client_secret": client_secret} if client_secret else {}),
|
||||
},
|
||||
)
|
||||
|
|
@ -24,6 +24,9 @@ from litellm.constants import (
|
|||
MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
||||
build_token_endpoint_client_auth,
|
||||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -113,12 +116,16 @@ class TokenExchangeHandler:
|
|||
f"but missing client_id or client_secret"
|
||||
)
|
||||
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
auth_method=server.token_endpoint_auth_method,
|
||||
client_id=server.client_id,
|
||||
client_secret=server.client_secret,
|
||||
)
|
||||
data: Dict[str, str] = {
|
||||
"grant_type": TOKEN_EXCHANGE_GRANT_TYPE,
|
||||
"subject_token": subject_token,
|
||||
"subject_token_type": server.subject_token_type or DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
"client_id": server.client_id,
|
||||
"client_secret": server.client_secret,
|
||||
**client_auth.body,
|
||||
}
|
||||
if server.audience:
|
||||
data["audience"] = server.audience
|
||||
|
|
@ -133,8 +140,9 @@ class TokenExchangeHandler:
|
|||
)
|
||||
|
||||
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
|
||||
post_kwargs = {"data": data, **({"headers": client_auth.headers} if client_auth.headers else {})}
|
||||
try:
|
||||
response = await client.post(endpoint, data=data)
|
||||
response = await client.post(endpoint, **post_kwargs)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as exc:
|
||||
verbose_logger.debug(
|
||||
|
|
|
|||
|
|
@ -9,6 +9,10 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm._uuid import uuid
|
||||
from litellm.constants import MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
||||
build_token_endpoint_client_auth,
|
||||
normalize_token_endpoint_auth_method,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
|
|
@ -1030,20 +1034,21 @@ async def refresh_user_oauth_token(
|
|||
)
|
||||
return None
|
||||
|
||||
token_data: Dict[str, str] = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": refresh_token,
|
||||
}
|
||||
if client_id:
|
||||
token_data["client_id"] = client_id
|
||||
if client_secret:
|
||||
token_data["client_secret"] = client_secret
|
||||
|
||||
try:
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
auth_method=normalize_token_endpoint_auth_method(getattr(server, "token_endpoint_auth_method", None)),
|
||||
client_id=client_id,
|
||||
client_secret=client_secret,
|
||||
)
|
||||
token_data: Dict[str, str] = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": refresh_token,
|
||||
**client_auth.body,
|
||||
}
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
response = await async_client.post(
|
||||
token_url,
|
||||
headers={"Accept": "application/json"},
|
||||
headers={"Accept": "application/json", **client_auth.headers},
|
||||
data=token_data,
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
|
|
|||
|
|
@ -14,6 +14,10 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
||||
TokenEndpointAuthConfigError,
|
||||
build_token_endpoint_client_auth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
TOKEN_NO_CACHE_HEADERS,
|
||||
get_request_base_url,
|
||||
|
|
@ -398,6 +402,14 @@ async def exchange_token_with_server(
|
|||
|
||||
resolved_client_id = mcp_server.client_id if mcp_server.client_id else client_id
|
||||
resolved_client_secret = mcp_server.client_secret if mcp_server.client_secret else client_secret
|
||||
try:
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
auth_method=mcp_server.token_endpoint_auth_method,
|
||||
client_id=resolved_client_id,
|
||||
client_secret=resolved_client_secret,
|
||||
)
|
||||
except TokenEndpointAuthConfigError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
if grant_type == "refresh_token":
|
||||
if not refresh_token:
|
||||
|
|
@ -408,10 +420,8 @@ async def exchange_token_with_server(
|
|||
token_data: dict = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": refresh_token,
|
||||
"client_id": resolved_client_id,
|
||||
**client_auth.body,
|
||||
}
|
||||
if resolved_client_secret is not None:
|
||||
token_data["client_secret"] = resolved_client_secret
|
||||
if scope:
|
||||
token_data["scope"] = scope
|
||||
else:
|
||||
|
|
@ -423,19 +433,17 @@ async def exchange_token_with_server(
|
|||
proxy_base_url = get_request_base_url(request)
|
||||
token_data = {
|
||||
"grant_type": "authorization_code",
|
||||
"client_id": resolved_client_id,
|
||||
"code": code,
|
||||
"redirect_uri": f"{proxy_base_url}/callback",
|
||||
**client_auth.body,
|
||||
}
|
||||
if resolved_client_secret is not None:
|
||||
token_data["client_secret"] = resolved_client_secret
|
||||
if code_verifier:
|
||||
token_data["code_verifier"] = code_verifier
|
||||
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
response = await async_client.post(
|
||||
mcp_server.token_url,
|
||||
headers={"Accept": "application/json"},
|
||||
headers={"Accept": "application/json", **client_auth.headers},
|
||||
data=token_data,
|
||||
)
|
||||
if response is None:
|
||||
|
|
|
|||
|
|
@ -754,6 +754,7 @@ class MCPServerManager:
|
|||
authorization_url=resolved_authorization_url,
|
||||
token_url=resolved_token_url,
|
||||
registration_url=resolved_registration_url,
|
||||
token_endpoint_auth_method=server_config.get("token_endpoint_auth_method", None),
|
||||
# TODO: utility fn the default values
|
||||
transport=server_config.get("transport", MCPTransport.http),
|
||||
auth_type=auth_type,
|
||||
|
|
@ -1127,6 +1128,9 @@ class MCPServerManager:
|
|||
authorization_url=mcp_server.authorization_url or getattr(mcp_oauth_metadata, "authorization_url", None),
|
||||
token_url=mcp_server.token_url or getattr(mcp_oauth_metadata, "token_url", None),
|
||||
registration_url=mcp_server.registration_url or getattr(mcp_oauth_metadata, "registration_url", None),
|
||||
token_endpoint_auth_method=(
|
||||
credentials_dict.get("token_endpoint_auth_method") if credentials_dict else None
|
||||
),
|
||||
command=getattr(mcp_server, "command", None),
|
||||
args=getattr(mcp_server, "args", None) or [],
|
||||
env=env_dict,
|
||||
|
|
|
|||
|
|
@ -27,6 +27,9 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
|||
encrypt_value_helper,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.auth import token_exchange
|
||||
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
||||
build_token_endpoint_client_auth,
|
||||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -103,10 +106,14 @@ class MCPOAuth2TokenCache(InMemoryCache):
|
|||
f"token_url={bool(server.token_url)}"
|
||||
)
|
||||
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
auth_method=server.token_endpoint_auth_method,
|
||||
client_id=server.client_id,
|
||||
client_secret=server.client_secret,
|
||||
)
|
||||
data: Dict[str, str] = {
|
||||
"grant_type": "client_credentials",
|
||||
"client_id": server.client_id,
|
||||
"client_secret": server.client_secret,
|
||||
**client_auth.body,
|
||||
}
|
||||
if server.scopes:
|
||||
data["scope"] = " ".join(server.scopes)
|
||||
|
|
@ -116,8 +123,9 @@ class MCPOAuth2TokenCache(InMemoryCache):
|
|||
server.server_id,
|
||||
)
|
||||
|
||||
post_kwargs = {"data": data, **({"headers": client_auth.headers} if client_auth.headers else {})}
|
||||
try:
|
||||
response = await client.post(server.token_url, data=data)
|
||||
response = await client.post(server.token_url, **post_kwargs)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise ValueError(
|
||||
|
|
|
|||
|
|
@ -14,6 +14,11 @@ import time
|
|||
from collections.abc import Awaitable, Callable
|
||||
from typing import TYPE_CHECKING, Protocol
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
||||
TokenEndpointAuthConfigError,
|
||||
build_token_endpoint_client_auth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
OAuthToken,
|
||||
)
|
||||
|
|
@ -22,7 +27,7 @@ if TYPE_CHECKING:
|
|||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
ServerLookup = Callable[[str], "MCPServer | None"]
|
||||
TokenEndpointPost = Callable[[str, dict[str, str]], Awaitable["dict[str, object] | None"]]
|
||||
TokenEndpointPost = Callable[[str, dict[str, str], dict[str, str]], Awaitable["dict[str, object] | None"]]
|
||||
|
||||
|
||||
class CredentialPersist(Protocol):
|
||||
|
|
@ -86,13 +91,21 @@ class AuthorizationCodeRefresher:
|
|||
if server is None or not server.token_url:
|
||||
return None
|
||||
|
||||
try:
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
auth_method=server.token_endpoint_auth_method,
|
||||
client_id=server.client_id,
|
||||
client_secret=server.client_secret,
|
||||
)
|
||||
except TokenEndpointAuthConfigError as exc:
|
||||
verbose_logger.warning("MCP OAuth refresh misconfigured for server %s: %s", server_id, exc)
|
||||
return None
|
||||
form = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": token.refresh_token,
|
||||
**({"client_id": server.client_id} if server.client_id else {}),
|
||||
**({"client_secret": server.client_secret} if server.client_secret else {}),
|
||||
**client_auth.body,
|
||||
}
|
||||
body = await self._token_endpoint(server.token_url, form)
|
||||
body = await self._token_endpoint(server.token_url, form, client_auth.headers)
|
||||
if body is None:
|
||||
return None
|
||||
access_token = body.get("access_token")
|
||||
|
|
|
|||
|
|
@ -92,7 +92,7 @@ async def _persist_credential(
|
|||
)
|
||||
|
||||
|
||||
async def _post_token_endpoint(url: str, form: dict[str, str]) -> dict[str, object] | None:
|
||||
async def _post_token_endpoint(url: str, form: dict[str, str], headers: dict[str, str]) -> dict[str, object] | None:
|
||||
from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415
|
||||
get_async_httpx_client, # pyright: ignore
|
||||
)
|
||||
|
|
@ -101,11 +101,11 @@ async def _post_token_endpoint(url: str, form: dict[str, str]) -> dict[str, obje
|
|||
# litellm's httpx handler and httpx.Response are only partially typed; the IdP returns a JSON
|
||||
# object and the refresher validates each field, so the untyped boundary is contained here.
|
||||
provider = httpxSpecialProvider.Oauth2Check
|
||||
headers = {"Accept": "application/json"}
|
||||
request_headers = {"Accept": "application/json", **headers}
|
||||
# A failed refresh is a miss, not a 500 (matches v1), so any error becomes None.
|
||||
try:
|
||||
client = get_async_httpx_client(llm_provider=provider) # pyright: ignore
|
||||
response = await client.post(url, headers=headers, data=form) # pyright: ignore
|
||||
response = await client.post(url, headers=request_headers, data=form) # pyright: ignore
|
||||
response.raise_for_status() # pyright: ignore
|
||||
body: dict[str, object] = response.json() # pyright: ignore
|
||||
except Exception as exc: # noqa: BLE001
|
||||
|
|
|
|||
|
|
@ -73,6 +73,10 @@ class MCPPublicServer(BaseModel):
|
|||
mcp_info: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
# OAuth 2.0 token-endpoint client authentication method (RFC 6749 section 2.3.1).
|
||||
MCPTokenEndpointAuthMethod = Literal["client_secret_basic", "client_secret_post"]
|
||||
|
||||
|
||||
class MCPCredentials(TypedDict, total=False):
|
||||
auth_value: Optional[str]
|
||||
"""
|
||||
|
|
@ -132,6 +136,12 @@ class MCPCredentials(TypedDict, total=False):
|
|||
Default: urn:ietf:params:oauth:token-type:access_token
|
||||
"""
|
||||
|
||||
token_endpoint_auth_method: Optional[MCPTokenEndpointAuthMethod]
|
||||
"""
|
||||
How the gateway authenticates to the upstream token endpoint. "client_secret_basic"
|
||||
sends HTTP Basic; defaults to "client_secret_post" when unset.
|
||||
"""
|
||||
|
||||
|
||||
class MCPServerCostInfo(TypedDict, total=False):
|
||||
default_cost_per_query: Optional[float]
|
||||
|
|
|
|||
|
|
@ -3,7 +3,12 @@ from typing import Any, Dict, List, Literal, Optional
|
|||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
from litellm.types.mcp import MCPAuth, MCPAuthType, MCPTransportType
|
||||
from litellm.types.mcp import (
|
||||
MCPAuth,
|
||||
MCPAuthType,
|
||||
MCPTokenEndpointAuthMethod,
|
||||
MCPTransportType,
|
||||
)
|
||||
|
||||
# MCPInfo now allows arbitrary additional fields for custom metadata
|
||||
MCPInfo = Dict[str, Any]
|
||||
|
|
@ -48,6 +53,10 @@ class MCPServer(BaseModel):
|
|||
authorization_url: Optional[str] = None
|
||||
token_url: Optional[str] = None
|
||||
registration_url: Optional[str] = None
|
||||
# How the gateway authenticates to the upstream token endpoint. When
|
||||
# "client_secret_basic" the credentials go in an HTTP Basic Authorization
|
||||
# header (omitted from the body); None defaults to "client_secret_post".
|
||||
token_endpoint_auth_method: Optional[MCPTokenEndpointAuthMethod] = None
|
||||
# AWS SigV4 fields
|
||||
aws_access_key_id: Optional[str] = None
|
||||
aws_secret_access_key: Optional[str] = None
|
||||
|
|
|
|||
|
|
@ -0,0 +1,80 @@
|
|||
"""Tests for token-endpoint client authentication (client_secret_basic vs client_secret_post)."""
|
||||
|
||||
import base64
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
||||
TokenEndpointAuthConfigError,
|
||||
build_token_endpoint_client_auth,
|
||||
normalize_token_endpoint_auth_method,
|
||||
)
|
||||
|
||||
|
||||
def _expected_basic(client_id: str, client_secret: str) -> str:
|
||||
return "Basic " + base64.b64encode(f"{client_id}:{client_secret}".encode()).decode()
|
||||
|
||||
|
||||
def test_basic_puts_credentials_in_header_and_not_body():
|
||||
auth = build_token_endpoint_client_auth(auth_method="client_secret_basic", client_id="cid", client_secret="sec")
|
||||
assert auth.headers == {"Authorization": _expected_basic("cid", "sec")}
|
||||
assert "client_secret" not in auth.body
|
||||
assert auth.body == {}
|
||||
|
||||
|
||||
def test_basic_form_urlencodes_reserved_characters():
|
||||
"""RFC 6749 2.3.1: client_id and client_secret are form-urlencoded before the ':' join, so reserved
|
||||
characters survive base64 transport instead of corrupting the username/password split."""
|
||||
auth = build_token_endpoint_client_auth(
|
||||
auth_method="client_secret_basic", client_id="client:one", client_secret="sec+ret:two"
|
||||
)
|
||||
decoded = base64.b64decode(auth.headers["Authorization"].removeprefix("Basic ")).decode()
|
||||
assert decoded == "client%3Aone:sec%2Bret%3Atwo"
|
||||
|
||||
|
||||
def test_post_default_puts_credentials_in_body_and_no_auth_header():
|
||||
auth = build_token_endpoint_client_auth(auth_method="client_secret_post", client_id="cid", client_secret="sec")
|
||||
assert auth.headers == {}
|
||||
assert auth.body == {"client_id": "cid", "client_secret": "sec"}
|
||||
|
||||
|
||||
def test_none_method_defaults_to_post():
|
||||
auth = build_token_endpoint_client_auth(auth_method=None, client_id="cid", client_secret="sec")
|
||||
assert auth.headers == {}
|
||||
assert auth.body == {"client_id": "cid", "client_secret": "sec"}
|
||||
|
||||
|
||||
def test_explicit_basic_without_secret_raises():
|
||||
"""client_secret_basic is a confidential-client method; a missing secret is a misconfiguration
|
||||
that must surface, not silently downgrade to a body request (RFC 6749; the no-silent-fallback rule)."""
|
||||
with pytest.raises(TokenEndpointAuthConfigError):
|
||||
build_token_endpoint_client_auth(auth_method="client_secret_basic", client_id="cid", client_secret=None)
|
||||
|
||||
|
||||
def test_explicit_basic_without_client_id_raises():
|
||||
with pytest.raises(TokenEndpointAuthConfigError):
|
||||
build_token_endpoint_client_auth(auth_method="client_secret_basic", client_id=None, client_secret="sec")
|
||||
|
||||
|
||||
def test_default_method_without_secret_is_public_client_post():
|
||||
"""A secretless client_id under the default method is the legitimate public-client / PKCE case:
|
||||
client_id goes in the body, no secret, no error."""
|
||||
auth = build_token_endpoint_client_auth(auth_method=None, client_id="cid", client_secret=None)
|
||||
assert auth.headers == {}
|
||||
assert auth.body == {"client_id": "cid"}
|
||||
|
||||
|
||||
def test_explicit_post_without_secret_does_not_raise():
|
||||
"""Unlike basic, explicit client_secret_post degrades to a valid public-client request, so it
|
||||
does not error on a missing secret."""
|
||||
auth = build_token_endpoint_client_auth(auth_method="client_secret_post", client_id="cid", client_secret=None)
|
||||
assert auth.headers == {}
|
||||
assert auth.body == {"client_id": "cid"}
|
||||
|
||||
|
||||
def test_normalize_only_accepts_known_methods():
|
||||
assert normalize_token_endpoint_auth_method("client_secret_basic") == "client_secret_basic"
|
||||
assert normalize_token_endpoint_auth_method("client_secret_post") == "client_secret_post"
|
||||
assert normalize_token_endpoint_auth_method("private_key_jwt") is None
|
||||
assert normalize_token_endpoint_auth_method(None) is None
|
||||
assert normalize_token_endpoint_auth_method(123) is None
|
||||
|
|
@ -509,3 +509,31 @@ async def test_database_loading_token_exchange_scopes_from_credentials():
|
|||
assert server.token_exchange_endpoint == "https://idp.example.com/oauth2/token"
|
||||
assert server.audience == "api://db-mcp"
|
||||
assert server.scopes == ["db.read", "db.write"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_token_uses_client_secret_basic_when_configured():
|
||||
"""LIT-4091: token exchange with token_endpoint_auth_method=client_secret_basic sends the
|
||||
client credentials as HTTP Basic and omits client_secret from the body."""
|
||||
import base64
|
||||
|
||||
handler = TokenExchangeHandler()
|
||||
server = _obo_server(
|
||||
server_id="srv-obo-basic", token_endpoint_auth_method="client_secret_basic"
|
||||
)
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = _exchange_response("scoped-basic")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
|
||||
return_value=mock_client,
|
||||
):
|
||||
result = await handler.exchange_token("user-jwt-basic", server)
|
||||
|
||||
assert result == "scoped-basic"
|
||||
_, kwargs = mock_client.post.call_args
|
||||
expected = "Basic " + base64.b64encode(b"litellm-client-id:litellm-client-secret").decode()
|
||||
assert kwargs["headers"]["Authorization"] == expected
|
||||
assert "client_secret" not in kwargs["data"]
|
||||
assert "client_id" not in kwargs["data"]
|
||||
assert kwargs["data"]["grant_type"] == TOKEN_EXCHANGE_GRANT_TYPE
|
||||
|
|
|
|||
|
|
@ -16,10 +16,12 @@ class _Server:
|
|||
token_url="https://idp.example.com/token",
|
||||
client_id="cid",
|
||||
client_secret="sec",
|
||||
token_endpoint_auth_method=None,
|
||||
):
|
||||
self.token_url = token_url
|
||||
self.client_id = client_id
|
||||
self.client_secret = client_secret
|
||||
self.token_endpoint_auth_method = token_endpoint_auth_method
|
||||
|
||||
|
||||
def _lookup(server):
|
||||
|
|
@ -27,9 +29,9 @@ def _lookup(server):
|
|||
|
||||
|
||||
def _endpoint(body, sink=None):
|
||||
async def post(url, form):
|
||||
async def post(url, form, headers):
|
||||
if sink is not None:
|
||||
sink.append((url, form))
|
||||
sink.append((url, form, headers))
|
||||
return body
|
||||
|
||||
return post
|
||||
|
|
@ -81,8 +83,8 @@ async def test_refreshes_persists_and_returns_typed_token():
|
|||
assert token.expires_at == 1000.0 + 3600 # clock + expires_in -> epoch
|
||||
# the rotated triple is persisted for (user, server) with parsed scopes
|
||||
assert persisted == [("alice", "srv", "new-at", "new-rt", 3600, ("a", "b"))]
|
||||
# the grant carried the refresh_token + client credentials
|
||||
url, form = posted[0]
|
||||
# the grant carried the refresh_token + client credentials in the body (client_secret_post default)
|
||||
url, form, headers = posted[0]
|
||||
assert url == "https://idp.example.com/token"
|
||||
assert form == {
|
||||
"grant_type": "refresh_token",
|
||||
|
|
@ -90,6 +92,44 @@ async def test_refreshes_persists_and_returns_typed_token():
|
|||
"client_id": "cid",
|
||||
"client_secret": "sec",
|
||||
}
|
||||
assert "Authorization" not in headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_secret_basic_sends_authorization_header_not_body():
|
||||
"""A server with token_endpoint_auth_method=client_secret_basic authenticates via HTTP Basic;
|
||||
the secret must not also leak into the form body."""
|
||||
import base64
|
||||
|
||||
posted = []
|
||||
server = _Server(token_endpoint_auth_method="client_secret_basic")
|
||||
refresher = _refresher(
|
||||
server=server,
|
||||
body={"access_token": "new-at"},
|
||||
post_sink=posted,
|
||||
)
|
||||
token = await refresher.refresh(
|
||||
"alice", "srv", OAuthToken(access_token="old", refresh_token="old-rt")
|
||||
)
|
||||
|
||||
assert token is not None
|
||||
_url, form, headers = posted[0]
|
||||
expected = "Basic " + base64.b64encode(b"cid:sec").decode()
|
||||
assert headers["Authorization"] == expected
|
||||
assert "client_secret" not in form
|
||||
assert "client_id" not in form
|
||||
assert form == {"grant_type": "refresh_token", "refresh_token": "old-rt"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_secret_basic_without_secret_is_a_failed_refresh():
|
||||
"""A server set to client_secret_basic but missing its secret cannot authenticate; the refresh
|
||||
returns None (failed refresh -> needs reauth) and never posts a downgraded request to the IdP."""
|
||||
posted = []
|
||||
server = _Server(client_secret=None, token_endpoint_auth_method="client_secret_basic")
|
||||
refresher = _refresher(server=server, body={"access_token": "x"}, post_sink=posted)
|
||||
assert await refresher.refresh("a", "s", OAuthToken("old", refresh_token="rt")) is None
|
||||
assert posted == [] # never hit the IdP
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -643,3 +643,82 @@ async def test_rotate_user_env_vars_skips_undecryptable_rows():
|
|||
assert prisma.db.litellm_mcpuserenvvars.update.call_count == 1
|
||||
where = prisma.db.litellm_mcpuserenvvars.update.call_args.kwargs["where"]
|
||||
assert where["user_id_server_id"]["server_id"] == "srv-ok"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_user_oauth_token_uses_client_secret_basic(monkeypatch):
|
||||
"""LIT-4091: a per-user refresh against a server with token_endpoint_auth_method=client_secret_basic
|
||||
sends HTTP Basic and keeps the secret out of the body."""
|
||||
import litellm.proxy._experimental.mcp_server.db as db_mod
|
||||
|
||||
server = MagicMock()
|
||||
server.token_url = "https://idp.example.com/oauth2/token"
|
||||
server.server_id = "srv"
|
||||
server.client_id = "cid"
|
||||
server.client_secret = "sec"
|
||||
server.token_endpoint_auth_method = "client_secret_basic"
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"access_token": "new-at", "expires_in": 3600}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = mock_response
|
||||
|
||||
monkeypatch.setattr(db_mod, "get_async_httpx_client", lambda **kwargs: mock_client)
|
||||
monkeypatch.setattr(db_mod, "store_user_oauth_credential", AsyncMock())
|
||||
monkeypatch.setattr(
|
||||
db_mod, "get_user_oauth_credential", AsyncMock(return_value={"access_token": "new-at"})
|
||||
)
|
||||
|
||||
result = await db_mod.refresh_user_oauth_token(
|
||||
prisma_client=MagicMock(),
|
||||
user_id="alice",
|
||||
server=server,
|
||||
cred={"refresh_token": "rt"},
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
_, kwargs = mock_client.post.call_args
|
||||
assert kwargs["headers"]["Authorization"] == "Basic " + base64.b64encode(b"cid:sec").decode()
|
||||
assert "client_secret" not in kwargs["data"]
|
||||
assert "client_id" not in kwargs["data"]
|
||||
assert kwargs["data"]["grant_type"] == "refresh_token"
|
||||
assert kwargs["data"]["refresh_token"] == "rt"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_user_oauth_token_defaults_to_client_secret_post(monkeypatch):
|
||||
"""Backward compatibility: with no token_endpoint_auth_method the refresh keeps credentials in
|
||||
the body (client_secret_post) and sends no Authorization header."""
|
||||
import litellm.proxy._experimental.mcp_server.db as db_mod
|
||||
|
||||
server = MagicMock()
|
||||
server.token_url = "https://idp.example.com/oauth2/token"
|
||||
server.server_id = "srv"
|
||||
server.client_id = "cid"
|
||||
server.client_secret = "sec"
|
||||
server.token_endpoint_auth_method = None
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"access_token": "new-at", "expires_in": 3600}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = mock_response
|
||||
|
||||
monkeypatch.setattr(db_mod, "get_async_httpx_client", lambda **kwargs: mock_client)
|
||||
monkeypatch.setattr(db_mod, "store_user_oauth_credential", AsyncMock())
|
||||
monkeypatch.setattr(
|
||||
db_mod, "get_user_oauth_credential", AsyncMock(return_value={"access_token": "new-at"})
|
||||
)
|
||||
|
||||
await db_mod.refresh_user_oauth_token(
|
||||
prisma_client=MagicMock(),
|
||||
user_id="alice",
|
||||
server=server,
|
||||
cred={"refresh_token": "rt"},
|
||||
)
|
||||
|
||||
_, kwargs = mock_client.post.call_args
|
||||
assert "Authorization" not in kwargs["headers"]
|
||||
assert kwargs["data"]["client_id"] == "cid"
|
||||
assert kwargs["data"]["client_secret"] == "sec"
|
||||
|
|
|
|||
|
|
@ -2733,3 +2733,128 @@ async def test_token_exchange_passes_through_upstream_expires_in():
|
|||
{"access_token": "tok", "token_type": "Bearer", "expires_in": 43200}
|
||||
)
|
||||
assert body["expires_in"] == 43200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_endpoint_uses_client_secret_basic_when_configured():
|
||||
"""LIT-4091: a server with token_endpoint_auth_method=client_secret_basic must send the
|
||||
credentials as an HTTP Basic Authorization header and omit client_secret from the body;
|
||||
providers requiring Basic rejected body credentials with invalid_client."""
|
||||
import base64
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
token_endpoint,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
oauth2_server = MCPServer(
|
||||
server_id="basic_mcp",
|
||||
name="basic_mcp",
|
||||
server_name="basic_mcp",
|
||||
alias="basic_mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="basic-client",
|
||||
client_secret="basic-secret",
|
||||
authorization_url="https://idp.example.com/authorize",
|
||||
token_url="https://idp.example.com/oauth2/token",
|
||||
token_endpoint_auth_method="client_secret_basic",
|
||||
)
|
||||
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm-proxy.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"access_token": "at",
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3599,
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
mock_get_client.return_value = mock_async_client
|
||||
|
||||
await token_endpoint(
|
||||
request=mock_request,
|
||||
grant_type="authorization_code",
|
||||
code="auth-code",
|
||||
redirect_uri="http://localhost/callback",
|
||||
client_id="basic-client",
|
||||
mcp_server_name="basic_mcp",
|
||||
client_secret="basic-secret",
|
||||
code_verifier="verifier",
|
||||
)
|
||||
|
||||
call_args = mock_async_client.post.call_args
|
||||
expected = "Basic " + base64.b64encode(b"basic-client:basic-secret").decode()
|
||||
assert call_args[1]["headers"]["Authorization"] == expected
|
||||
assert "client_secret" not in call_args[1]["data"]
|
||||
assert "client_id" not in call_args[1]["data"]
|
||||
assert call_args[1]["data"]["grant_type"] == "authorization_code"
|
||||
assert call_args[1]["data"]["code"] == "auth-code"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_endpoint_client_secret_basic_without_secret_returns_400():
|
||||
"""A server configured client_secret_basic but missing its secret is a misconfiguration; the
|
||||
inbound /token endpoint surfaces it as a 400 rather than silently posting a downgraded request."""
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
token_endpoint,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
oauth2_server = MCPServer(
|
||||
server_id="basic_no_secret",
|
||||
name="basic_no_secret",
|
||||
server_name="basic_no_secret",
|
||||
alias="basic_no_secret",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="basic-client",
|
||||
client_secret=None,
|
||||
authorization_url="https://idp.example.com/authorize",
|
||||
token_url="https://idp.example.com/oauth2/token",
|
||||
token_endpoint_auth_method="client_secret_basic",
|
||||
)
|
||||
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm-proxy.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await token_endpoint(
|
||||
request=mock_request,
|
||||
grant_type="authorization_code",
|
||||
code="auth-code",
|
||||
redirect_uri="http://localhost/callback",
|
||||
client_id="basic-client",
|
||||
mcp_server_name="basic_no_secret",
|
||||
client_secret=None,
|
||||
code_verifier="verifier",
|
||||
)
|
||||
assert exc_info.value.status_code == 400
|
||||
|
|
|
|||
|
|
@ -323,6 +323,28 @@ class TestMCPServerManager:
|
|||
assert cost_info["tool_name_to_cost_per_query"]["geocode"] == 1e-3
|
||||
assert isinstance(cost_info["tool_name_to_cost_per_query"]["geocode"], float)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_servers_from_config_sets_token_endpoint_auth_method(self):
|
||||
"""token_endpoint_auth_method from config is carried onto the MCPServer (LIT-4091)."""
|
||||
manager = MCPServerManager()
|
||||
config = {
|
||||
"basic_provider": {
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
"token_endpoint_auth_method": "client_secret_basic",
|
||||
},
|
||||
"default_provider": {
|
||||
"url": "https://example.com/mcp2",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
}
|
||||
|
||||
await manager.load_servers_from_config(config)
|
||||
|
||||
by_name = {s.server_name: s for s in manager.config_mcp_servers.values()}
|
||||
assert by_name["basic_provider"].token_endpoint_auth_method == "client_secret_basic"
|
||||
assert by_name["default_provider"].token_endpoint_auth_method is None
|
||||
|
||||
def test_normalize_mcp_server_cost_info_preserves_float_values(self):
|
||||
mcp_info = {
|
||||
"server_name": "maps",
|
||||
|
|
@ -3292,6 +3314,31 @@ class TestMCPServerTimestamps:
|
|||
assert mcp_server.created_at == created
|
||||
assert mcp_server.updated_at == updated
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_mcp_server_from_table_reads_token_endpoint_auth_method(self):
|
||||
"""token_endpoint_auth_method stored in the credentials JSON is loaded onto the MCPServer (LIT-4091)."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
basic_record = LiteLLM_MCPServerTable(
|
||||
server_id="basic-db-1",
|
||||
server_name="basic_db",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
credentials={"token_endpoint_auth_method": "client_secret_basic"},
|
||||
)
|
||||
basic_server = await manager.build_mcp_server_from_table(basic_record, credentials_are_encrypted=False)
|
||||
assert basic_server.token_endpoint_auth_method == "client_secret_basic"
|
||||
|
||||
default_record = LiteLLM_MCPServerTable(
|
||||
server_id="default-db-1",
|
||||
server_name="default_db",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
credentials={},
|
||||
)
|
||||
default_server = await manager.build_mcp_server_from_table(default_record, credentials_are_encrypted=False)
|
||||
assert default_server.token_endpoint_auth_method is None
|
||||
|
||||
def test_build_mcp_server_table_preserves_timestamps(self):
|
||||
"""_build_mcp_server_table must use the MCPServer's stored timestamps, not datetime.now()."""
|
||||
manager = MCPServerManager()
|
||||
|
|
|
|||
|
|
@ -173,3 +173,27 @@ async def test_non_dict_response_raises_value_error():
|
|||
pytest.raises(ValueError, match="non-object JSON"),
|
||||
):
|
||||
await resolve_mcp_auth(server)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_credentials_uses_client_secret_basic_when_configured():
|
||||
"""LIT-4091: a client_credentials server with token_endpoint_auth_method=client_secret_basic
|
||||
authenticates via HTTP Basic and keeps the secret out of the form body."""
|
||||
import base64
|
||||
|
||||
server = _server(server_id="srv-basic", token_endpoint_auth_method="client_secret_basic")
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = _token_response("m2m-basic")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
|
||||
return_value=mock_client,
|
||||
):
|
||||
result = await resolve_mcp_auth(server)
|
||||
|
||||
assert result == "m2m-basic"
|
||||
_, kwargs = mock_client.post.call_args
|
||||
assert kwargs["headers"]["Authorization"] == "Basic " + base64.b64encode(b"cid:csec").decode()
|
||||
assert "client_secret" not in kwargs["data"]
|
||||
assert "client_id" not in kwargs["data"]
|
||||
assert kwargs["data"]["grant_type"] == "client_credentials"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue