mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/model-table-divider-center-b75b6d
This commit is contained in:
commit
984051f74d
40 changed files with 1858 additions and 332 deletions
|
|
@ -2887,7 +2887,7 @@
|
|||
"input_cost_per_token": 5e-06,
|
||||
"output_cost_per_token": 2.5e-05,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 200000,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
|
|
@ -2916,7 +2916,7 @@
|
|||
"input_cost_per_token": 5e-06,
|
||||
"output_cost_per_token": 2.5e-05,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 200000,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
|
|
@ -3010,7 +3010,7 @@
|
|||
"input_cost_per_token": 5e-06,
|
||||
"output_cost_per_token": 2.5e-05,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 200000,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
|
|
|
|||
|
|
@ -9,10 +9,7 @@ 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._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
|
|
@ -1248,11 +1245,12 @@ def _decrypted_credential_field(creds: Dict[str, object], field: str) -> object:
|
|||
|
||||
def mcp_oauth_token_identity(server: object) -> tuple[object, ...]:
|
||||
"""The upstream-OAuth-token-determining fields of an MCP server: the resource/audience (url, or
|
||||
spec_path for OpenAPI servers), the OAuth mode/grant (auth_type, oauth2_flow), the
|
||||
authorization-server endpoints, and the OAuth client + scopes. Mirrors the dashboard's
|
||||
getOAuthAuthorizationIdentity. When any of these change on a server update, previously stored
|
||||
per-user tokens were minted for the old identity and are stale. Excludes transport and
|
||||
delegate_auth_to_upstream, which do not affect what token is minted (RFC 8707/8693).
|
||||
spec_path for OpenAPI servers, plus the RFC 8707 upstream_resource sent on the authorize and
|
||||
token legs), the OAuth mode/grant (auth_type, oauth2_flow), the authorization-server endpoints,
|
||||
and the OAuth client + scopes. Mirrors the dashboard's getOAuthAuthorizationIdentity. When any
|
||||
of these change on a server update, previously stored per-user tokens were minted for the old
|
||||
identity and are stale. Excludes transport and delegate_auth_to_upstream, which do not affect
|
||||
what token is minted (RFC 8693).
|
||||
|
||||
client_id/client_secret are compared decrypted: stored values are NaCl-encrypted with a fresh
|
||||
nonce on every write, so comparing ciphertext would flag every routine save as an identity
|
||||
|
|
@ -1278,6 +1276,7 @@ def mcp_oauth_token_identity(server: object) -> tuple[object, ...]:
|
|||
_decrypted_credential_field(creds_dict, "client_id"),
|
||||
_decrypted_credential_field(creds_dict, "client_secret"),
|
||||
creds_dict.get("scopes"),
|
||||
creds_dict.get("upstream_resource"),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1367,20 +1366,21 @@ async def refresh_user_oauth_token(
|
|||
return None
|
||||
|
||||
try:
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
auth_method=normalize_token_endpoint_auth_method(getattr(server, "token_endpoint_auth_method", None)),
|
||||
token_request = build_upstream_oauth2_token_request(
|
||||
server,
|
||||
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,
|
||||
**token_request.body,
|
||||
}
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
response = await async_client.post(
|
||||
token_url,
|
||||
headers={"Accept": "application/json", **client_auth.headers},
|
||||
headers={"Accept": "application/json", **token_request.headers},
|
||||
data=token_data,
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
|
|
|||
|
|
@ -21,7 +21,6 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
||||
TokenEndpointAuthConfigError,
|
||||
build_token_endpoint_client_auth,
|
||||
normalize_token_endpoint_auth_method,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPTokenEndpointAuthMethod
|
||||
|
|
@ -54,7 +53,9 @@ from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import (
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
TOKEN_NO_CACHE_HEADERS,
|
||||
build_upstream_oauth2_token_request,
|
||||
get_request_base_url,
|
||||
resolve_upstream_resource,
|
||||
validate_trusted_redirect_uri,
|
||||
well_known_root_suffix,
|
||||
)
|
||||
|
|
@ -726,6 +727,7 @@ def _redirect_to_upstream_authorize(
|
|||
to the upstream authorize endpoint verbatim, no relay state cookie is set, and the upstream
|
||||
enforces its own registered redirect binding for the client."""
|
||||
scope_value = scope or (" ".join(mcp_server.scopes) if mcp_server.scopes else None)
|
||||
upstream_resource = resolve_upstream_resource(mcp_server)
|
||||
passthrough_params = {
|
||||
"client_id": client_id,
|
||||
"redirect_uri": redirect_uri,
|
||||
|
|
@ -734,6 +736,7 @@ def _redirect_to_upstream_authorize(
|
|||
"code_challenge": code_challenge,
|
||||
"code_challenge_method": code_challenge_method,
|
||||
**({"scope": scope_value} if scope_value else {}),
|
||||
**({"resource": upstream_resource} if upstream_resource else {}),
|
||||
}
|
||||
parsed_auth_url = urlparse(mcp_server.authorization_url or "")
|
||||
merged_params = {**dict(parse_qsl(parsed_auth_url.query)), **passthrough_params}
|
||||
|
|
@ -842,6 +845,10 @@ async def authorize_with_server(
|
|||
if code_challenge_method:
|
||||
params["code_challenge_method"] = code_challenge_method
|
||||
|
||||
upstream_resource = resolve_upstream_resource(mcp_server)
|
||||
if upstream_resource:
|
||||
params["resource"] = upstream_resource
|
||||
|
||||
parsed_auth_url = urlparse(mcp_server.authorization_url)
|
||||
existing_params = dict(parse_qsl(parsed_auth_url.query))
|
||||
existing_params.update(params)
|
||||
|
|
@ -902,7 +909,8 @@ async def exchange_token_with_server(
|
|||
else (client_token_endpoint_auth_method or mcp_server.token_endpoint_auth_method)
|
||||
)
|
||||
try:
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
token_request = build_upstream_oauth2_token_request(
|
||||
mcp_server,
|
||||
auth_method=resolved_auth_method,
|
||||
client_id=resolved_client_id,
|
||||
client_secret=resolved_client_secret,
|
||||
|
|
@ -941,7 +949,7 @@ async def exchange_token_with_server(
|
|||
token_data: dict = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": upstream_refresh_token,
|
||||
**client_auth.body,
|
||||
**token_request.body,
|
||||
}
|
||||
refresh_request_scope = scope or bridge_upstream_scope
|
||||
if refresh_request_scope:
|
||||
|
|
@ -980,7 +988,7 @@ async def exchange_token_with_server(
|
|||
"grant_type": "authorization_code",
|
||||
"code": code,
|
||||
"redirect_uri": resolved_redirect_uri,
|
||||
**client_auth.body,
|
||||
**token_request.body,
|
||||
}
|
||||
if code_verifier:
|
||||
token_data["code_verifier"] = code_verifier
|
||||
|
|
@ -991,11 +999,12 @@ async def exchange_token_with_server(
|
|||
if not isinstance(prepared, _BridgeMintReady):
|
||||
return _bridge_mint_error_response(prepared)
|
||||
bridge_mint_ready = prepared
|
||||
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
try:
|
||||
response = await async_client.post(
|
||||
mcp_server.token_url,
|
||||
headers={"Accept": "application/json", **client_auth.headers},
|
||||
headers={"Accept": "application/json", **token_request.headers},
|
||||
data=token_data,
|
||||
)
|
||||
if response is not None:
|
||||
|
|
|
|||
|
|
@ -63,18 +63,21 @@ def _classify_oauth_error_code(
|
|||
) -> UpstreamOAuthFault:
|
||||
"""Blame assignment for a contract-conformant OAuth error code, shared by the token and DCR
|
||||
classifiers. Codes by which the upstream blames itself keep that blame; ``invalid_target`` is a
|
||||
gateway capability gap (RFC 8707 resource indicators, LIT-4339) no matter whose credentials were
|
||||
presented; credential-indicting codes follow the credential source; everything else, including
|
||||
codes we do not recognize, is the caller's to act on. The upstream's HTTP status is deliberately
|
||||
never consulted: status derives from this classification at render time, which is what keeps
|
||||
status and code from contradicting each other."""
|
||||
gateway configuration gap (the RFC 8707 resource indicator this server sends, or fails to send)
|
||||
no matter whose credentials were presented; credential-indicting codes follow the credential
|
||||
source; everything else, including codes we do not recognize, is the caller's to act on. The
|
||||
upstream's HTTP status is deliberately never consulted: status derives from this classification
|
||||
at render time, which is what keeps status and code from contradicting each other."""
|
||||
if code == "server_error" or code == "temporarily_unavailable":
|
||||
return UpstreamReportedFault(code=code)
|
||||
if code in GATEWAY_CAPABILITY_CODES:
|
||||
verbose_logger.warning(
|
||||
"MCP server %s: the upstream authorization server rejected the request with "
|
||||
"invalid_target; it may require RFC 8707 resource indicators, which the gateway "
|
||||
"does not send yet (tracked as LIT-4339)",
|
||||
"invalid_target, meaning it did not accept the RFC 8707 resource indicator for this "
|
||||
"request. Set upstream_resource on this server to the exact resource identifier the "
|
||||
"authorization server expects (or to 'auto' to send the server's own canonical url); "
|
||||
"if it is already set and the authorization server does not support resource "
|
||||
"indicators, unset it and express the target audience through scopes instead",
|
||||
log_context,
|
||||
)
|
||||
return GatewayRejected(code=code)
|
||||
|
|
|
|||
|
|
@ -16,8 +16,10 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HE
|
|||
def _gateway_rejected_description(code: str) -> str:
|
||||
if code == "invalid_target":
|
||||
return (
|
||||
"the upstream authorization server rejected the request (invalid_target); "
|
||||
"it may require RFC 8707 resource indicators, which the gateway does not send yet"
|
||||
"the upstream authorization server rejected the request (invalid_target); it did not "
|
||||
"accept this server's RFC 8707 resource indicator. Set upstream_resource on the MCP "
|
||||
"server to the resource identifier the authorization server expects, or unset it if "
|
||||
"that authorization server does not support resource indicators"
|
||||
)
|
||||
return (
|
||||
f"the upstream authorization server rejected the gateway's configured client credentials "
|
||||
|
|
|
|||
|
|
@ -25,9 +25,10 @@ gateway presented its own stored credentials, these are gateway-side faults the
|
|||
when the caller supplied the credentials, they are the caller's to fix."""
|
||||
|
||||
GATEWAY_CAPABILITY_CODES: frozenset[str] = frozenset({"invalid_target"})
|
||||
"""Codes that indict a gateway capability regardless of whose credentials were presented:
|
||||
``invalid_target`` means the upstream wants RFC 8707 resource indicators, which the gateway does not
|
||||
send yet (LIT-4339). Never the caller's fault."""
|
||||
"""Codes that indict gateway configuration regardless of whose credentials were presented:
|
||||
``invalid_target`` means the upstream did not accept the RFC 8707 resource indicator the server
|
||||
sent, or requires one it was not configured to send (``upstream_resource``). Never the caller's
|
||||
fault."""
|
||||
|
||||
UPSTREAM_FAULT_CODES: frozenset[str] = frozenset({"server_error", "temporarily_unavailable"})
|
||||
"""Codes by which the upstream blames itself. Relaying them as caller faults would invert blame, so
|
||||
|
|
|
|||
|
|
@ -71,6 +71,7 @@ from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
_redact_mcp_resource_url,
|
||||
canonicalize_url_identity,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
|
||||
Error,
|
||||
|
|
@ -261,19 +262,11 @@ def _endpoints_yield_to_issuer(
|
|||
|
||||
|
||||
def _normalized_authorize_endpoint(url: str) -> str:
|
||||
"""Compare authorize endpoints on scheme, host, and path only. The default port is elided and
|
||||
the host is lowercased so ``https://IDP.example.com:443/authorize/`` and
|
||||
``https://idp.example.com/authorize`` are the same identity; query and trailing slash are not."""
|
||||
parsed = urlparse(url)
|
||||
scheme = parsed.scheme.lower()
|
||||
host = (parsed.hostname or "").lower()
|
||||
default_port = {"https": 443, "http": 80}.get(scheme)
|
||||
try:
|
||||
port = parsed.port
|
||||
except ValueError:
|
||||
port = None
|
||||
authority = host if port is None or port == default_port else f"{host}:{port}"
|
||||
return f"{scheme}://{authority}{parsed.path.rstrip('/')}"
|
||||
"""Compare authorize endpoints / issuers on scheme, host, and path only, through the shared URL
|
||||
canonicalizer: the default port is elided and the host is lowercased so
|
||||
``https://IDP.example.com:443/authorize/`` and ``https://idp.example.com/authorize`` are the same
|
||||
identity, while query, fragment and a trailing slash are dropped."""
|
||||
return canonicalize_url_identity(url)
|
||||
|
||||
|
||||
def _issuer_matches(claimed_issuer: object, configured_issuer: str) -> bool:
|
||||
|
|
@ -1517,6 +1510,7 @@ class MCPServerManager:
|
|||
"subject_token_type",
|
||||
DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
),
|
||||
upstream_resource=server_config.get("upstream_resource", None),
|
||||
# ID-JAG fields
|
||||
id_jag_resource_token_endpoint=server_config.get("id_jag_resource_token_endpoint", None),
|
||||
id_jag_resource=server_config.get("id_jag_resource", None),
|
||||
|
|
@ -2016,6 +2010,7 @@ class MCPServerManager:
|
|||
subject_token_type=mcp_server.subject_token_type
|
||||
or (credentials_dict.get("subject_token_type") if credentials_dict else None)
|
||||
or DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
upstream_resource=(credentials_dict.get("upstream_resource") if credentials_dict else None),
|
||||
# ID-JAG fields — read from credentials JSON blob
|
||||
id_jag_resource_token_endpoint=(
|
||||
credentials_dict.get("id_jag_resource_token_endpoint") if credentials_dict else None
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ with ``client_id``, ``client_secret``, and ``token_url``.
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
from typing import TYPE_CHECKING, Dict, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
|
@ -26,8 +27,9 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
|||
decrypt_value_helper,
|
||||
encrypt_value_helper,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
||||
build_token_endpoint_client_auth,
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
build_upstream_oauth2_token_request,
|
||||
resolve_upstream_resource,
|
||||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
|
|
@ -37,10 +39,18 @@ if TYPE_CHECKING:
|
|||
|
||||
class MCPOAuth2TokenCache(InMemoryCache):
|
||||
"""
|
||||
In-memory cache for OAuth2 client_credentials tokens, keyed by server_id.
|
||||
In-memory cache for OAuth2 client_credentials tokens, keyed by the identity of the token
|
||||
request rather than by server_id alone.
|
||||
|
||||
A minted token is only reusable for the exact request that produced it. Keying on server_id
|
||||
alone served a token minted under the previous configuration whenever any of those inputs
|
||||
changed, so editing scopes, rotating the client secret, or setting ``upstream_resource``
|
||||
silently kept handing out a token carrying the old scopes or audience until it expired. The
|
||||
identity below covers every input ``_fetch_token`` puts on the wire, so a change to any of
|
||||
them misses the cache and mints afresh.
|
||||
|
||||
Inherits from ``InMemoryCache`` for TTL-based storage and eviction.
|
||||
Adds per-server ``asyncio.Lock`` to prevent duplicate concurrent fetches.
|
||||
Adds a per-identity ``asyncio.Lock`` to prevent duplicate concurrent fetches.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
|
|
@ -50,8 +60,25 @@ class MCPOAuth2TokenCache(InMemoryCache):
|
|||
)
|
||||
self._locks: Dict[str, asyncio.Lock] = {}
|
||||
|
||||
def _get_lock(self, server_id: str) -> asyncio.Lock:
|
||||
return self._locks.setdefault(server_id, asyncio.Lock())
|
||||
@staticmethod
|
||||
def _token_identity(server: "MCPServer") -> str:
|
||||
"""Cache key for the token this server's config would mint, prefixed by server_id so a
|
||||
single server's entries stay greppable and invalidatable. The secret is hashed with the
|
||||
rest of the identity rather than stored in a key."""
|
||||
material = "\x00".join(
|
||||
(
|
||||
server.token_url or "",
|
||||
server.client_id or "",
|
||||
server.client_secret or "",
|
||||
" ".join(server.scopes or ()),
|
||||
resolve_upstream_resource(server) or "",
|
||||
server.token_endpoint_auth_method or "",
|
||||
)
|
||||
)
|
||||
return f"{server.server_id}:{hashlib.sha256(material.encode()).hexdigest()}"
|
||||
|
||||
def _get_lock(self, identity: str) -> asyncio.Lock:
|
||||
return self._locks.setdefault(identity, asyncio.Lock())
|
||||
|
||||
@staticmethod
|
||||
def _has_client_credentials_config(server: "MCPServer") -> bool:
|
||||
|
|
@ -67,21 +94,21 @@ class MCPOAuth2TokenCache(InMemoryCache):
|
|||
if not self._has_client_credentials_config(server):
|
||||
return None
|
||||
|
||||
server_id = server.server_id
|
||||
identity = self._token_identity(server)
|
||||
|
||||
# Fast path — cached token is still valid
|
||||
cached = self.get_cache(server_id)
|
||||
cached = self.get_cache(identity)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
# Slow path — acquire per-server lock then double-check
|
||||
async with self._get_lock(server_id):
|
||||
cached = self.get_cache(server_id)
|
||||
# Slow path — acquire per-identity lock then double-check
|
||||
async with self._get_lock(identity):
|
||||
cached = self.get_cache(identity)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
token, ttl = await self._fetch_token(server)
|
||||
self.set_cache(server_id, token, ttl=ttl)
|
||||
self.set_cache(identity, token, ttl=ttl)
|
||||
return token
|
||||
|
||||
async def _fetch_token(self, server: "MCPServer") -> Tuple[str, int]:
|
||||
|
|
@ -100,14 +127,15 @@ class MCPOAuth2TokenCache(InMemoryCache):
|
|||
f"token_url={bool(server.token_url)}"
|
||||
)
|
||||
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
token_request = build_upstream_oauth2_token_request(
|
||||
server,
|
||||
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_auth.body,
|
||||
**token_request.body,
|
||||
}
|
||||
if server.scopes:
|
||||
data["scope"] = " ".join(server.scopes)
|
||||
|
|
@ -117,7 +145,7 @@ class MCPOAuth2TokenCache(InMemoryCache):
|
|||
server.server_id,
|
||||
)
|
||||
|
||||
post_kwargs = {"data": data, **({"headers": client_auth.headers} if client_auth.headers else {})}
|
||||
post_kwargs = {"data": data, **({"headers": token_request.headers} if token_request.headers else {})}
|
||||
try:
|
||||
response = await client.post(server.token_url, **post_kwargs)
|
||||
response.raise_for_status()
|
||||
|
|
@ -159,8 +187,14 @@ class MCPOAuth2TokenCache(InMemoryCache):
|
|||
return access_token, ttl
|
||||
|
||||
def invalidate(self, server_id: str) -> None:
|
||||
"""Remove a cached token (e.g. after a 401)."""
|
||||
self.delete_cache(server_id)
|
||||
"""Remove every cached token for a server (e.g. after a 401).
|
||||
|
||||
Entries are keyed by token identity, so one server can hold more than one entry across a
|
||||
config change; a 401 invalidates all of them rather than only the current configuration's.
|
||||
"""
|
||||
prefix = f"{server_id}:"
|
||||
for key in [k for k in self.cache_dict if isinstance(k, str) and k.startswith(prefix)]:
|
||||
self.delete_cache(key)
|
||||
|
||||
|
||||
mcp_oauth2_token_cache = MCPOAuth2TokenCache()
|
||||
|
|
|
|||
|
|
@ -3,14 +3,22 @@
|
|||
|
||||
import os
|
||||
from ipaddress import ip_address
|
||||
from typing import Any, Dict, List, NoReturn, Optional
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, NoReturn, Optional
|
||||
from urllib.parse import ParseResult, urlparse, urlsplit, urlunparse, urlunsplit
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
||||
TokenEndpointClientAuth,
|
||||
build_token_endpoint_client_auth,
|
||||
normalize_token_endpoint_auth_method,
|
||||
)
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
# RFC 6749 §5.1 / OAuth 2.1 draft-15 §4.1.3: token-endpoint responses
|
||||
# must not be cached — both success and error bodies may reveal secrets.
|
||||
TOKEN_NO_CACHE_HEADERS = {"Cache-Control": "no-store", "Pragma": "no-cache"}
|
||||
|
|
@ -21,6 +29,10 @@ TOKEN_NO_CACHE_HEADERS = {"Cache-Control": "no-store", "Pragma": "no-cache"}
|
|||
# explicit port, which would otherwise break a literal netloc compare).
|
||||
_DEFAULT_PORTS = {"http": 80, "https": 443}
|
||||
|
||||
# Sentinel ``upstream_resource`` value meaning "derive the RFC 8707 resource identifier from the
|
||||
# server's own url". RFC 8707 requires an absolute URI, so this can never be a real resource value.
|
||||
UPSTREAM_RESOURCE_AUTO = "auto"
|
||||
|
||||
# Env var for ops to allowlist additional redirect_uri origins beyond
|
||||
# same-origin + loopback — needed for first-party OAuth clients hosted
|
||||
# on sister domains (e.g. a web app on app.example.com registering as
|
||||
|
|
@ -574,3 +586,112 @@ def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None:
|
|||
if _trusted_redirect_uri_is_allowed(parsed, redirect_netloc, proxy_base):
|
||||
return
|
||||
_raise_trusted_redirect_uri_rejected(request, redirect_uri, parsed, redirect_netloc, proxy_base)
|
||||
|
||||
|
||||
def canonicalize_url_identity(url: str) -> str:
|
||||
"""Normalize a URL to a comparable identity: lowercase scheme and host, drop the scheme's default
|
||||
port, and strip userinfo, params, query, fragment and a trailing slash while keeping IPv6
|
||||
brackets. The one URL-canonicalization primitive shared by the RFC 8707 resource emitter and the
|
||||
RFC 8414 issuer/authorize-endpoint comparison, so the default-port and IPv6 rules cannot be
|
||||
present in one and missing in the other. The netloc (not ``parsed.hostname``) carries the
|
||||
authority so ``[::1]:8080`` survives with its brackets intact."""
|
||||
parsed = urlparse(url)
|
||||
scheme = parsed.scheme.lower()
|
||||
netloc = _strip_default_port(scheme, parsed.netloc.rpartition("@")[2])
|
||||
return urlunparse((scheme, netloc, parsed.path.rstrip("/"), "", "", ""))
|
||||
|
||||
|
||||
def _canonical_resource_uri(url: str) -> str | None:
|
||||
"""Canonicalize an upstream MCP server URL into an RFC 8707 resource identifier.
|
||||
|
||||
Keeps only the scheme, host, port and path, which is the shape the MCP authorization spec's
|
||||
"Canonical Server URI" section describes and every one of its examples takes; the reference
|
||||
implementation is ``mcp.shared.auth_utils.resource_url_from_server_url``, and this is the stricter
|
||||
variant. The scheme and host are lowercased, the scheme's default port is dropped so
|
||||
``https://host:443/mcp`` and ``https://host/mcp`` never present as two resources, and a trailing
|
||||
slash is dropped so ``https://host/mcp/`` and ``https://host/mcp`` do not either.
|
||||
|
||||
Userinfo, query and fragment are dropped rather than carried. A transport URL routinely holds
|
||||
credentials in exactly those components (``user:password@``, ``?api_key=``), while a resource
|
||||
indicator names the resource and nothing else; this value is published somewhere the transport
|
||||
URL never goes, into the authorization redirect the browser follows and into token request
|
||||
bodies, so carrying them would disclose them to the authorization server, its logs, and browser
|
||||
history. RFC 8707 forbids a fragment outright and says a resource SHOULD NOT carry a query. An
|
||||
upstream whose identifier genuinely needs more than this is served by setting
|
||||
``upstream_resource`` explicitly, which is passed through untouched.
|
||||
|
||||
Returns ``None`` when the URL is not absolute, which cannot yield a valid resource identifier.
|
||||
"""
|
||||
parsed = urlparse(url)
|
||||
if not parsed.scheme or not parsed.netloc:
|
||||
return None
|
||||
return canonicalize_url_identity(url)
|
||||
|
||||
|
||||
def resolve_upstream_resource(mcp_server: "MCPServer") -> str | None:
|
||||
"""Resolve the RFC 8707 ``resource`` value this server's upstream OAuth legs must carry.
|
||||
|
||||
The MCP authorization spec requires an MCP client to send ``resource`` on both the
|
||||
authorization request and every token request, naming the canonical URI of the MCP server the
|
||||
token is for. Authorization server temperaments are irreconcilable and undetectable, so this
|
||||
stays an explicit per-server opt-in: most SaaS providers ignore the parameter, some hard-reject
|
||||
it and express audience through scopes instead, and strict or MCP-native ones refuse to mint a
|
||||
correctly scoped token without it (``invalid_target``).
|
||||
|
||||
``None`` or blank omits the parameter, which is the default and preserves the behavior of every
|
||||
server working today. ``"auto"`` derives the canonical URI from the server's own URL; it is not
|
||||
an absolute URI, so RFC 8707 guarantees it can never collide with a real resource value. Any
|
||||
other value is sent verbatim, because the identifier has to match what the authorization server
|
||||
expects exactly and normalizing it could break that match.
|
||||
|
||||
Every upstream leg for a server resolves through this one function, so the authorize request
|
||||
and the token requests cannot disagree; a token request naming a resource the authorization
|
||||
request never asked for is itself an ``invalid_target`` under RFC 8707.
|
||||
"""
|
||||
configured = (mcp_server.upstream_resource or "").strip()
|
||||
if not configured:
|
||||
return None
|
||||
if configured.lower() != UPSTREAM_RESOURCE_AUTO:
|
||||
return configured
|
||||
if not mcp_server.url:
|
||||
verbose_logger.warning(
|
||||
"MCP server %s sets upstream_resource=auto but has no url to derive a resource "
|
||||
"identifier from; omitting the RFC 8707 resource parameter. Set upstream_resource to "
|
||||
"the exact resource identifier the authorization server expects instead.",
|
||||
mcp_server.server_id,
|
||||
)
|
||||
return None
|
||||
canonical = _canonical_resource_uri(mcp_server.url)
|
||||
if canonical is None:
|
||||
verbose_logger.warning(
|
||||
"MCP server %s sets upstream_resource=auto but its url is not an absolute URI, so no "
|
||||
"RFC 8707 resource identifier could be derived; omitting the resource parameter",
|
||||
mcp_server.server_id,
|
||||
)
|
||||
return canonical
|
||||
|
||||
|
||||
def build_upstream_oauth2_token_request(
|
||||
mcp_server: "MCPServer",
|
||||
*,
|
||||
auth_method: object,
|
||||
client_id: str | None,
|
||||
client_secret: str | None,
|
||||
) -> TokenEndpointClientAuth:
|
||||
"""Client auth plus the RFC 8707 ``resource`` for one upstream plain-OAuth2 token request.
|
||||
|
||||
Resolving both in one call is what stops a leg authenticating without naming the resource its
|
||||
sibling legs named; the RFC 8693 legs (OBO, id_jag) carry ``audience`` and stay on
|
||||
``build_token_endpoint_client_auth``. The client-auth inputs are passed in because a leg may
|
||||
authenticate as the caller's own client rather than the server's; ``resource`` always comes from
|
||||
the server, so no leg can choose or forget it.
|
||||
"""
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
auth_method=normalize_token_endpoint_auth_method(auth_method),
|
||||
client_id=client_id,
|
||||
client_secret=client_secret,
|
||||
)
|
||||
resource = resolve_upstream_resource(mcp_server)
|
||||
if not resource:
|
||||
return client_auth
|
||||
return TokenEndpointClientAuth(headers=client_auth.headers, body={**client_auth.body, "resource": resource})
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from fastapi import HTTPException
|
|||
from pydantic import SecretStr
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import resolve_upstream_resource
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
ApiKeyConfig,
|
||||
AuthorizationCodeConfig,
|
||||
|
|
@ -144,6 +145,7 @@ def _client_credentials_spec(server: MCPServer, resource: str) -> ServerSpec:
|
|||
token_url=server.token_url,
|
||||
scopes=tuple(server.scopes or ()),
|
||||
audience=server.audience,
|
||||
upstream_resource=resolve_upstream_resource(server),
|
||||
token_endpoint_auth_method=server.token_endpoint_auth_method,
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -17,8 +17,8 @@ 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.oauth_utils import build_upstream_oauth2_token_request
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
OAuthToken,
|
||||
)
|
||||
|
|
@ -92,7 +92,8 @@ class AuthorizationCodeRefresher:
|
|||
return None
|
||||
|
||||
try:
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
token_request = build_upstream_oauth2_token_request(
|
||||
server,
|
||||
auth_method=server.token_endpoint_auth_method,
|
||||
client_id=server.client_id,
|
||||
client_secret=server.client_secret,
|
||||
|
|
@ -103,9 +104,9 @@ class AuthorizationCodeRefresher:
|
|||
form = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": token.refresh_token,
|
||||
**client_auth.body,
|
||||
**token_request.body,
|
||||
}
|
||||
body = await self._token_endpoint(server.token_url, form, client_auth.headers)
|
||||
body = await self._token_endpoint(server.token_url, form, token_request.headers)
|
||||
if body is None:
|
||||
return None
|
||||
access_token = body.get("access_token")
|
||||
|
|
|
|||
|
|
@ -292,6 +292,7 @@ def _prepare_grant(config: ClientCredentialsConfig) -> Result[_PreparedGrant, Cr
|
|||
**client_auth.body,
|
||||
**({"scope": " ".join(config.scopes)} if config.scopes else {}),
|
||||
**({"audience": config.audience} if config.audience else {}),
|
||||
**({"resource": config.upstream_resource} if config.upstream_resource else {}),
|
||||
}
|
||||
return Ok(
|
||||
_PreparedGrant(
|
||||
|
|
@ -313,6 +314,7 @@ def _identity_key(config: ClientCredentialsConfig) -> str:
|
|||
config.token_endpoint_auth_method or "",
|
||||
" ".join(config.scopes),
|
||||
config.audience or "",
|
||||
config.upstream_resource or "",
|
||||
)
|
||||
)
|
||||
return hashlib.sha256(material.encode("utf-8")).hexdigest()
|
||||
|
|
|
|||
|
|
@ -199,6 +199,7 @@ class ClientCredentialsConfig(BaseModel):
|
|||
token_url: str | None = None
|
||||
scopes: tuple[str, ...] = ()
|
||||
audience: str | None = None
|
||||
upstream_resource: str | None = None
|
||||
token_endpoint_auth_method: Literal["client_secret_post", "client_secret_basic"] | None = None
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,8 +1,9 @@
|
|||
"""LLM-as-a-Judge guardrail: uses an LLM to score responses against weighted criteria."""
|
||||
|
||||
import json
|
||||
import re
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Literal, Optional, Union, cast
|
||||
|
||||
import litellm
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -13,6 +14,7 @@ from litellm.types.guardrails import GuardrailEventHooks, SupportedGuardrailInte
|
|||
from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm import Router
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.utils import StandardLoggingEvalInformation
|
||||
|
|
@ -30,6 +32,38 @@ Return ONLY valid JSON in this exact format:
|
|||
_VALID_ON_FAILURE = frozenset({"block", "log"})
|
||||
|
||||
|
||||
def _default_router_provider() -> "Router | None":
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
except ImportError:
|
||||
return None
|
||||
|
||||
return llm_router
|
||||
|
||||
|
||||
_JSON_FENCE_RE = re.compile(r"```(?:json)?\s*(.*?)\s*```", re.DOTALL | re.IGNORECASE)
|
||||
|
||||
|
||||
def _parse_judge_verdict(raw: str) -> Dict[str, Any]:
|
||||
"""Parse the judge's JSON verdict, tolerating markdown fences and surrounding prose."""
|
||||
text = raw.strip()
|
||||
fenced = _JSON_FENCE_RE.search(text)
|
||||
if fenced is not None:
|
||||
text = fenced.group(1).strip()
|
||||
parsed: object
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
start = text.find("{")
|
||||
end = text.rfind("}")
|
||||
if start == -1 or end <= start:
|
||||
raise
|
||||
parsed = json.loads(text[start : end + 1])
|
||||
if not isinstance(parsed, dict):
|
||||
raise ValueError("judge response is not a JSON object")
|
||||
return cast(Dict[str, Any], parsed) # cast-ok: narrowed to dict by the isinstance guard above
|
||||
|
||||
|
||||
def _extract_text_from_content(content: Any) -> str:
|
||||
"""Return plain text from a message content field (str or multimodal list)."""
|
||||
if isinstance(content, str):
|
||||
|
|
@ -94,6 +128,7 @@ class LLMAsAJudgeGuardrail(CustomGuardrail):
|
|||
on_failure: Literal["block", "log"] = "block",
|
||||
event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks]]] = None,
|
||||
default_on: bool = False,
|
||||
router_provider: "Callable[[], Router | None] | None" = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
_event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks]]] = None
|
||||
|
|
@ -114,6 +149,7 @@ class LLMAsAJudgeGuardrail(CustomGuardrail):
|
|||
self.criteria = criteria
|
||||
self.overall_threshold = overall_threshold
|
||||
self.on_failure = on_failure
|
||||
self._router_provider = router_provider or _default_router_provider
|
||||
|
||||
@classmethod
|
||||
def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]:
|
||||
|
|
@ -131,14 +167,27 @@ class LLMAsAJudgeGuardrail(CustomGuardrail):
|
|||
"content": _build_judge_prompt(self.criteria, messages, response_text),
|
||||
},
|
||||
]
|
||||
response = await litellm.acompletion(
|
||||
model=self.judge_model,
|
||||
messages=judge_messages,
|
||||
response_format={"type": "json_object"},
|
||||
temperature=0,
|
||||
)
|
||||
router = self._router_provider()
|
||||
if router is not None and (
|
||||
self.judge_model in router.model_group_alias or router.get_model_list(model_name=self.judge_model)
|
||||
):
|
||||
response = await router.acompletion(
|
||||
model=self.judge_model,
|
||||
messages=judge_messages,
|
||||
response_format={"type": "json_object"},
|
||||
temperature=0,
|
||||
num_retries=0,
|
||||
fallbacks=[],
|
||||
)
|
||||
else:
|
||||
response = await litellm.acompletion(
|
||||
model=self.judge_model,
|
||||
messages=judge_messages,
|
||||
response_format={"type": "json_object"},
|
||||
temperature=0,
|
||||
)
|
||||
raw = response.choices[0].message.content or "{}" # type: ignore[union-attr]
|
||||
return json.loads(raw)
|
||||
return _parse_judge_verdict(raw)
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -181,7 +181,11 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
|
||||
from litellm.proxy.management_helpers.utils import management_endpoint_wrapper
|
||||
from litellm.types.mcp import MCPAuth, MCPCredentials
|
||||
from litellm.types.mcp import (
|
||||
MCP_ADMIN_CONFIG_CREDENTIAL_KEYS,
|
||||
MCPAuth,
|
||||
MCPCredentials,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
@dataclass
|
||||
|
|
@ -476,7 +480,8 @@ if MCP_AVAILABLE:
|
|||
def _redact_mcp_credentials(
|
||||
mcp_server: LiteLLM_MCPServerTable,
|
||||
) -> LiteLLM_MCPServerTable:
|
||||
"""Return a copy of the MCP server object with credentials removed."""
|
||||
"""Return a copy with secret credentials removed, keeping only non-secret admin config so the
|
||||
admin form can show and clear it. Non-admin and virtual-key views strip the whole blob."""
|
||||
|
||||
try:
|
||||
redacted_server = mcp_server.model_copy(deep=True)
|
||||
|
|
@ -484,10 +489,35 @@ if MCP_AVAILABLE:
|
|||
redacted_server = mcp_server.copy(deep=True) # type: ignore[attr-defined]
|
||||
|
||||
if hasattr(redacted_server, "credentials"):
|
||||
setattr(redacted_server, "credentials", None)
|
||||
setattr(redacted_server, "credentials", _preserved_admin_config_credentials(redacted_server.credentials))
|
||||
|
||||
return redacted_server
|
||||
|
||||
def _preserved_admin_config_credentials(
|
||||
credentials: "MCPCredentials | str | None",
|
||||
) -> "dict[str, str] | None":
|
||||
"""Keep only the non-secret admin-config keys, which are stored unencrypted so they lift out
|
||||
as plaintext; every secret and minted-token key is dropped.
|
||||
|
||||
Total over every stored shape: a dict is read directly, a JSON-object string is parsed, and
|
||||
anything else (a malformed or non-object JSON string, a scalar, ``None``) falls back to full
|
||||
redaction rather than raising, because this runs on every admin list and get and one bad row
|
||||
must not fail them all."""
|
||||
parsed: object = credentials
|
||||
if isinstance(credentials, str):
|
||||
try:
|
||||
parsed = json.loads(credentials)
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
if not isinstance(parsed, dict):
|
||||
return None
|
||||
preserved = {
|
||||
key: value
|
||||
for key in MCP_ADMIN_CONFIG_CREDENTIAL_KEYS
|
||||
if isinstance((value := parsed.get(key)), str) and value
|
||||
}
|
||||
return preserved or None
|
||||
|
||||
def _redact_mcp_credentials_list(
|
||||
mcp_servers: Iterable[LiteLLM_MCPServerTable],
|
||||
) -> List[LiteLLM_MCPServerTable]:
|
||||
|
|
@ -529,6 +559,7 @@ if MCP_AVAILABLE:
|
|||
``[]``/``{}`` for required list/dict fields).
|
||||
"""
|
||||
sanitized = _redact_mcp_credentials(mcp_server)
|
||||
sanitized.credentials = None
|
||||
# URL is the highest-impact vector: many MCP integrations embed
|
||||
# the upstream API key directly in the path. spec_path can carry
|
||||
# similar tokens in the OpenAPI spec URL.
|
||||
|
|
@ -572,6 +603,7 @@ if MCP_AVAILABLE:
|
|||
"""
|
||||
|
||||
sanitized = _redact_mcp_credentials(mcp_server)
|
||||
sanitized.credentials = None
|
||||
|
||||
# Remove potentially sensitive config + identity fields.
|
||||
sanitized.url = None
|
||||
|
|
@ -615,36 +647,47 @@ if MCP_AVAILABLE:
|
|||
) -> List[LiteLLM_MCPServerTable]:
|
||||
return [_sanitize_mcp_server_for_virtual_key(server) for server in mcp_servers]
|
||||
|
||||
# (server attribute, credentials key) a session server inherits from the server it derives from.
|
||||
# Declared as a table rather than a chain of ifs, which is how upstream_resource was missed.
|
||||
_INHERITED_CREDENTIAL_FIELDS: tuple[tuple[str, str], ...] = (
|
||||
("authentication_token", "auth_value"),
|
||||
("client_id", "client_id"),
|
||||
("client_secret", "client_secret"),
|
||||
("scopes", "scopes"),
|
||||
("aws_access_key_id", "aws_access_key_id"),
|
||||
("aws_secret_access_key", "aws_secret_access_key"),
|
||||
("aws_session_token", "aws_session_token"),
|
||||
("aws_region_name", "aws_region_name"),
|
||||
("aws_service_name", "aws_service_name"),
|
||||
("upstream_resource", "upstream_resource"),
|
||||
)
|
||||
|
||||
def _has_non_admin_config_credentials(credentials: "MCPCredentials | None") -> bool:
|
||||
"""Did the caller supply an actual credential? Admin config rides in the same blob but is not
|
||||
one, so a form that round-trips it must not read as "credentials supplied"."""
|
||||
if not credentials:
|
||||
return False
|
||||
as_dict: dict[str, Any] = dict(credentials)
|
||||
return any(value for key, value in as_dict.items() if key not in MCP_ADMIN_CONFIG_CREDENTIAL_KEYS)
|
||||
|
||||
def _inherit_credentials_from_existing_server(
|
||||
payload: NewMCPServerRequest,
|
||||
) -> NewMCPServerRequest:
|
||||
if not payload.server_id or payload.credentials:
|
||||
if not payload.server_id or _has_non_admin_config_credentials(payload.credentials):
|
||||
return payload
|
||||
|
||||
existing_server = global_mcp_server_manager.get_mcp_server_by_id(payload.server_id)
|
||||
if existing_server is None:
|
||||
return payload
|
||||
|
||||
inherited_credentials: MCPCredentials = {}
|
||||
if existing_server.authentication_token:
|
||||
inherited_credentials["auth_value"] = existing_server.authentication_token
|
||||
if existing_server.client_id:
|
||||
inherited_credentials["client_id"] = existing_server.client_id
|
||||
if existing_server.client_secret:
|
||||
inherited_credentials["client_secret"] = existing_server.client_secret
|
||||
if existing_server.scopes:
|
||||
inherited_credentials["scopes"] = existing_server.scopes
|
||||
# AWS SigV4 fields
|
||||
if existing_server.aws_access_key_id:
|
||||
inherited_credentials["aws_access_key_id"] = existing_server.aws_access_key_id
|
||||
if existing_server.aws_secret_access_key:
|
||||
inherited_credentials["aws_secret_access_key"] = existing_server.aws_secret_access_key
|
||||
if existing_server.aws_session_token:
|
||||
inherited_credentials["aws_session_token"] = existing_server.aws_session_token
|
||||
if existing_server.aws_region_name:
|
||||
inherited_credentials["aws_region_name"] = existing_server.aws_region_name
|
||||
if existing_server.aws_service_name:
|
||||
inherited_credentials["aws_service_name"] = existing_server.aws_service_name
|
||||
inherited_credentials: dict[str, Any] = {
|
||||
credential_key: value
|
||||
for server_attr, credential_key in _INHERITED_CREDENTIAL_FIELDS
|
||||
if (value := getattr(existing_server, server_attr, None))
|
||||
}
|
||||
# The gate above guarantees anything still supplied is admin config, which the admin just
|
||||
# typed, so it wins over the stored value.
|
||||
inherited_credentials = {**inherited_credentials, **dict(payload.credentials or {})}
|
||||
|
||||
if not inherited_credentials:
|
||||
return payload
|
||||
|
|
|
|||
|
|
@ -171,6 +171,15 @@ class MCPCredentials(TypedDict, total=False):
|
|||
Optional RFC 8707 resource indicator sent on ID-JAG leg 1
|
||||
"""
|
||||
|
||||
upstream_resource: str | None
|
||||
"""
|
||||
Optional RFC 8707 resource indicator sent on the upstream oauth2 legs (authorize, both token
|
||||
grants, and the client_credentials fetch). Omitted when unset, which is the default; "auto"
|
||||
derives the canonical URI from the server's url; any other value is sent verbatim.
|
||||
Distinct from ``id_jag_resource``, which is the same parameter on the ID-JAG exchange, and from
|
||||
``audience``, which is the RFC 8693 token-exchange parameter.
|
||||
"""
|
||||
|
||||
client_private_key: Optional[str]
|
||||
"""
|
||||
PEM private key used to sign the private-key-JWT client_assertion (RFC 7523)
|
||||
|
|
@ -213,6 +222,11 @@ class MCPCredentials(TypedDict, total=False):
|
|||
"""
|
||||
|
||||
|
||||
MCP_ADMIN_CONFIG_CREDENTIAL_KEYS: tuple[str, ...] = ("upstream_resource",)
|
||||
"""Non-secret credential keys returned on read so the admin form can show and clear them. Mirrors
|
||||
``ADMIN_CONFIG_CREDENTIAL_KEYS`` in ``ui/litellm-dashboard/src/components/mcp_tools/types.tsx``."""
|
||||
|
||||
|
||||
class MCPServerCostInfo(TypedDict, total=False):
|
||||
default_cost_per_query: Optional[float]
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -75,6 +75,11 @@ class MCPServer(BaseModel):
|
|||
# "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
|
||||
# RFC 8707 resource indicator sent on this server's upstream oauth2 legs (authorize, both
|
||||
# token grants, and the client_credentials fetch). None omits it, which is the default and
|
||||
# today's behavior; "auto" derives the canonical URI from ``url``; any other value is sent
|
||||
# verbatim. Resolved by ``oauth_utils.resolve_upstream_resource``.
|
||||
upstream_resource: str | None = None
|
||||
# AWS SigV4 fields
|
||||
aws_access_key_id: Optional[str] = None
|
||||
aws_secret_access_key: Optional[str] = None
|
||||
|
|
|
|||
|
|
@ -2887,7 +2887,7 @@
|
|||
"input_cost_per_token": 5e-06,
|
||||
"output_cost_per_token": 2.5e-05,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 200000,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
|
|
@ -2916,7 +2916,7 @@
|
|||
"input_cost_per_token": 5e-06,
|
||||
"output_cost_per_token": 2.5e-05,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 200000,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
|
|
@ -3010,7 +3010,7 @@
|
|||
"input_cost_per_token": 5e-06,
|
||||
"output_cost_per_token": 2.5e-05,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 200000,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ count actual model entries, not reserved meta keys) and the extraction of the
|
|||
``fallback_generalizations`` block out of the raw map.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
|
||||
|
|
@ -25,6 +26,14 @@ from litellm.litellm_core_utils.get_model_cost_map import (
|
|||
)
|
||||
|
||||
|
||||
def _load_root_cost_map() -> dict:
|
||||
path = os.path.join(
|
||||
os.path.dirname(__file__), "../../../model_prices_and_context_window.json"
|
||||
)
|
||||
with open(path) as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def _make_models(n: int) -> dict:
|
||||
return {
|
||||
f"model-{i}": {"litellm_provider": "openai", "mode": "chat"} for i in range(n)
|
||||
|
|
@ -209,3 +218,33 @@ def test_shipped_backup_marks_claude_4_6_plus_adaptive_not_4_0():
|
|||
"claude-opus-4-5",
|
||||
]:
|
||||
assert "supports_adaptive_thinking" not in backup[non_adaptive], non_adaptive
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"cost_map",
|
||||
[_load_root_cost_map(), GetModelCostMap.load_local_model_cost_map()],
|
||||
ids=["root", "bundled_backup"],
|
||||
)
|
||||
def test_azure_ai_claude_1m_context_entries(cost_map: dict):
|
||||
"""Microsoft Foundry serves a 1M-token context window for Opus 4.6+ and Sonnet
|
||||
4.6+, so the ``azure_ai`` entries must not advertise the 200k cap that made
|
||||
context-aware clients compact prompts early (LIT-4406). Both the root map (used
|
||||
by default network loading) and the bundled fallback are checked so the two can
|
||||
never drift apart."""
|
||||
for model in [
|
||||
"azure_ai/claude-opus-4-6",
|
||||
"azure_ai/claude-opus-4-7",
|
||||
"azure_ai/claude-opus-4-8",
|
||||
"azure_ai/claude-opus-5",
|
||||
"azure_ai/claude-sonnet-5",
|
||||
"azure_ai/claude-sonnet-4-6",
|
||||
]:
|
||||
assert cost_map[model]["max_input_tokens"] == 1000000, model
|
||||
|
||||
for model in [
|
||||
"azure_ai/claude-opus-4-1",
|
||||
"azure_ai/claude-opus-4-5",
|
||||
"azure_ai/claude-sonnet-4-5",
|
||||
"azure_ai/claude-haiku-4-5",
|
||||
]:
|
||||
assert cost_map[model]["max_input_tokens"] == 200000, model
|
||||
|
|
|
|||
|
|
@ -163,6 +163,26 @@ def test_client_credentials_omits_audience_when_unset():
|
|||
assert spec is not None
|
||||
assert isinstance(spec.config, ClientCredentialsConfig)
|
||||
assert spec.config.audience is None
|
||||
assert spec.config.upstream_resource is None
|
||||
|
||||
|
||||
def test_client_credentials_resolves_upstream_resource_onto_the_config():
|
||||
"""The adapter is the one MCPServer -> config chokepoint, so it resolves the RFC 8707 send value
|
||||
(auto here derives the canonical server URI) and every M2M token request inherits it."""
|
||||
spec = to_server_spec(
|
||||
_server(
|
||||
auth_type=MCPAuth.oauth2,
|
||||
oauth2_flow="client_credentials",
|
||||
url="https://up.example.com/mcp",
|
||||
token_url="https://idp.example.com/token",
|
||||
client_id="cid",
|
||||
client_secret="csec",
|
||||
upstream_resource="auto",
|
||||
)
|
||||
)
|
||||
assert spec is not None
|
||||
assert isinstance(spec.config, ClientCredentialsConfig)
|
||||
assert spec.config.upstream_resource == "https://up.example.com/mcp"
|
||||
|
||||
|
||||
def test_client_credentials_with_incomplete_grant_fields_is_owned_for_fail_closed():
|
||||
|
|
|
|||
|
|
@ -17,11 +17,17 @@ class _Server:
|
|||
client_id="cid",
|
||||
client_secret="sec",
|
||||
token_endpoint_auth_method=None,
|
||||
upstream_resource=None,
|
||||
url=None,
|
||||
server_id="srv",
|
||||
):
|
||||
self.token_url = token_url
|
||||
self.client_id = client_id
|
||||
self.client_secret = client_secret
|
||||
self.token_endpoint_auth_method = token_endpoint_auth_method
|
||||
self.upstream_resource = upstream_resource
|
||||
self.url = url
|
||||
self.server_id = server_id
|
||||
|
||||
|
||||
def _lookup(server):
|
||||
|
|
@ -209,6 +215,40 @@ async def test_unrecorded_scope_is_carried_forward():
|
|||
assert persisted[0][5] == ("read", "write")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_sends_upstream_resource_when_set_explicitly():
|
||||
"""A silent refresh must carry the same RFC 8707 resource its authorize/initial-token legs sent,
|
||||
or a strict authorization server rejects the refresh with invalid_target."""
|
||||
posted = []
|
||||
server = _Server(upstream_resource="https://api.example.com/mcp")
|
||||
refresher = _refresher(server=server, body={"access_token": "new-at"}, post_sink=posted)
|
||||
token = await refresher.refresh("a", "s", OAuthToken("old", refresh_token="rt"))
|
||||
assert token is not None
|
||||
_url, form, _headers = posted[0]
|
||||
assert form["resource"] == "https://api.example.com/mcp"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_sends_upstream_resource_auto_derived_from_url():
|
||||
posted = []
|
||||
server = _Server(upstream_resource="auto", url="https://mcp.example.com/mcp")
|
||||
refresher = _refresher(server=server, body={"access_token": "new-at"}, post_sink=posted)
|
||||
token = await refresher.refresh("a", "s", OAuthToken("old", refresh_token="rt"))
|
||||
assert token is not None
|
||||
_url, form, _headers = posted[0]
|
||||
assert form["resource"] == "https://mcp.example.com/mcp"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_omits_resource_when_unset():
|
||||
posted = []
|
||||
refresher = _refresher(body={"access_token": "new-at"}, post_sink=posted)
|
||||
token = await refresher.refresh("a", "s", OAuthToken("old", refresh_token="rt"))
|
||||
assert token is not None
|
||||
_url, form, _headers = posted[0]
|
||||
assert "resource" not in form
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returned_scope_overrides_prior_when_present():
|
||||
persisted = []
|
||||
|
|
|
|||
|
|
@ -81,6 +81,31 @@ async def test_grant_omits_scope_and_audience_when_not_configured():
|
|||
_url, form, _headers = poster.calls[0]
|
||||
assert "scope" not in form
|
||||
assert "audience" not in form
|
||||
assert "resource" not in form
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_grant_sends_rfc8707_resource_indicator():
|
||||
"""HTTP/SSE M2M tool traffic resolves through this v2 arm, so the RFC 8707 resource must ride it
|
||||
too or a strict authorization server keeps answering invalid_target on the primary M2M path."""
|
||||
poster = _FakePoster([_success()])
|
||||
await ClientCredentialsTokenSource(poster).get("s", _config(upstream_resource="api://finance-audience"))
|
||||
_url, form, _headers = poster.calls[0]
|
||||
assert form["resource"] == "api://finance-audience"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_changing_only_the_resource_mints_a_fresh_token():
|
||||
"""The resource is part of the mint identity: retargeting a live M2M server must not keep serving
|
||||
the token minted for the previous audience."""
|
||||
poster = _FakePoster([_success(access_token="tok-a", expires_in=3600), _success(access_token="tok-b", expires_in=3600)])
|
||||
source = ClientCredentialsTokenSource(poster)
|
||||
first = await source.get("s", _config(upstream_resource="api://one"))
|
||||
second = await source.get("s", _config(upstream_resource="api://two"))
|
||||
assert isinstance(first, Ok) and isinstance(second, Ok)
|
||||
assert first.ok.access_token == "tok-a"
|
||||
assert second.ok.access_token == "tok-b"
|
||||
assert len(poster.calls) == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -95,6 +95,17 @@ def _identity_server(**overrides):
|
|||
{"credentials": {"client_id": "new", "client_secret": "csec", "scopes": ["a"]}},
|
||||
{"credentials": {"client_id": "cid", "client_secret": "rotated", "scopes": ["a"]}},
|
||||
{"credentials": {"client_id": "cid", "client_secret": "csec", "scopes": ["b"]}},
|
||||
# RFC 8707: upstream_resource is the audience the token is minted for, so changing it
|
||||
# alone strands every stored per-user token on the previous audience.
|
||||
{"credentials": {"client_id": "cid", "client_secret": "csec", "scopes": ["a"], "upstream_resource": "auto"}},
|
||||
{
|
||||
"credentials": {
|
||||
"client_id": "cid",
|
||||
"client_secret": "csec",
|
||||
"scopes": ["a"],
|
||||
"upstream_resource": "api://new-audience",
|
||||
}
|
||||
},
|
||||
],
|
||||
)
|
||||
def test_mcp_oauth_token_identity_changes_on_mint_relevant_fields(overrides):
|
||||
|
|
@ -827,6 +838,82 @@ async def test_resolve_returns_none_for_missing_credential(monkeypatch):
|
|||
refresh.assert_not_called()
|
||||
|
||||
|
||||
class _RefreshResponse:
|
||||
def __init__(self, body):
|
||||
self._body = body
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def json(self):
|
||||
return self._body
|
||||
|
||||
|
||||
def _refresh_server(**overrides):
|
||||
base = dict(
|
||||
token_url="https://idp.example.com/token",
|
||||
server_id="srv-1",
|
||||
client_id="cid",
|
||||
client_secret="csec",
|
||||
token_endpoint_auth_method=None,
|
||||
upstream_resource=None,
|
||||
url="https://up.example.com/mcp",
|
||||
)
|
||||
base.update(overrides)
|
||||
return SimpleNamespace(**base)
|
||||
|
||||
|
||||
async def _run_refresh(monkeypatch, server, response_body=None):
|
||||
import litellm.proxy._experimental.mcp_server.db as db_mod
|
||||
|
||||
captured: dict = {}
|
||||
|
||||
async def _post(url, headers=None, data=None):
|
||||
captured["url"] = url
|
||||
captured["headers"] = headers
|
||||
captured["data"] = data
|
||||
return _RefreshResponse(response_body or {"access_token": "at-new", "expires_in": 3600})
|
||||
|
||||
monkeypatch.setattr(db_mod, "get_async_httpx_client", lambda **_: SimpleNamespace(post=_post))
|
||||
monkeypatch.setattr(db_mod, "store_user_oauth_credential", AsyncMock())
|
||||
monkeypatch.setattr(db_mod, "get_user_oauth_credential", AsyncMock(return_value={"access_token": "at-new"}))
|
||||
|
||||
result = await db_mod.refresh_user_oauth_token(
|
||||
prisma_client=MagicMock(),
|
||||
user_id="alice",
|
||||
server=server,
|
||||
cred={"refresh_token": "rt-old", "scopes": ["a"]},
|
||||
)
|
||||
return result, captured
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_user_oauth_token_sends_upstream_resource_when_set(monkeypatch):
|
||||
"""The server-side silent refresh must carry the same RFC 8707 resource the authorize and initial
|
||||
token legs sent; a strict authorization server rejects a refresh whose resource is absent with
|
||||
invalid_target, forcing a needless re-auth."""
|
||||
result, captured = await _run_refresh(monkeypatch, _refresh_server(upstream_resource="api://audience"))
|
||||
assert result is not None
|
||||
assert captured["data"]["grant_type"] == "refresh_token"
|
||||
assert captured["data"]["resource"] == "api://audience"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_user_oauth_token_sends_auto_derived_resource(monkeypatch):
|
||||
result, captured = await _run_refresh(
|
||||
monkeypatch, _refresh_server(upstream_resource="auto", url="https://mcp.example.com/mcp")
|
||||
)
|
||||
assert result is not None
|
||||
assert captured["data"]["resource"] == "https://mcp.example.com/mcp"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_user_oauth_token_omits_resource_when_unset(monkeypatch):
|
||||
result, captured = await _run_refresh(monkeypatch, _refresh_server(upstream_resource=None))
|
||||
assert result is not None
|
||||
assert "resource" not in captured["data"]
|
||||
|
||||
|
||||
# ── per-user env-var rotation ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
|
|
@ -1067,3 +1154,28 @@ async def test_delete_mcp_server_cleans_oauth_client_store():
|
|||
await delete_mcp_server(prisma, "s1", invalidate_token_cache=AsyncMock())
|
||||
|
||||
prisma.db.litellm_mcpserveroauthclient.delete_many.assert_awaited_once_with(where={"server_id": "s1"})
|
||||
|
||||
|
||||
def test_mcp_oauth_token_identity_changes_when_only_upstream_resource_is_edited():
|
||||
"""A resource-only update must purge stored per-user tokens.
|
||||
|
||||
Changing ``upstream_resource`` changes the audience the next token is minted for, so every
|
||||
token already stored for this server was minted for the old (or unbounded) audience. Without
|
||||
this field in the identity, an administrator retargeting a server leaves authenticated users
|
||||
calling tools with the previous audience's token until it expires, which is the token-reuse
|
||||
RFC 8707 exists to stop.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.db import mcp_oauth_token_identity
|
||||
|
||||
creds = {"client_id": "cid", "client_secret": "csec", "scopes": ["a"]}
|
||||
unset = _identity_server(credentials=dict(creds))
|
||||
set_to_auto = _identity_server(credentials={**creds, "upstream_resource": "auto"})
|
||||
set_to_explicit = _identity_server(credentials={**creds, "upstream_resource": "api://audience-one"})
|
||||
retargeted = _identity_server(credentials={**creds, "upstream_resource": "api://audience-two"})
|
||||
|
||||
assert mcp_oauth_token_identity(unset) != mcp_oauth_token_identity(set_to_auto)
|
||||
assert mcp_oauth_token_identity(unset) != mcp_oauth_token_identity(set_to_explicit)
|
||||
assert mcp_oauth_token_identity(set_to_explicit) != mcp_oauth_token_identity(retargeted)
|
||||
assert mcp_oauth_token_identity(set_to_explicit) == mcp_oauth_token_identity(
|
||||
_identity_server(credentials={**creds, "upstream_resource": "api://audience-one"})
|
||||
)
|
||||
|
|
|
|||
|
|
@ -8822,3 +8822,303 @@ async def test_token_exchange_authenticates_with_the_sealed_clients_own_auth_met
|
|||
assert "Authorization" not in sent_headers
|
||||
assert sent_body["client_id"] == "minted-77"
|
||||
assert sent_body["client_secret"] == "mint-secret"
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# LIT-4339: RFC 8707 resource indicators on the upstream OAuth legs
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _resource_server(**overrides) -> "MCPServer":
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
defaults = dict(
|
||||
server_id="res-srv",
|
||||
name="res-srv",
|
||||
server_name="res-srv",
|
||||
alias="res-srv",
|
||||
url="https://mcp.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="gateway-client",
|
||||
client_secret="gateway-secret",
|
||||
authorization_url="https://idp.example.com/oauth/authorize",
|
||||
token_url="https://idp.example.com/oauth/token",
|
||||
)
|
||||
defaults.update(overrides)
|
||||
return MCPServer(**defaults)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"url, configured, expected",
|
||||
[
|
||||
("https://mcp.example.com/mcp", None, None),
|
||||
("https://mcp.example.com/mcp", "", None),
|
||||
("https://mcp.example.com/mcp", " ", None),
|
||||
("https://mcp.example.com/mcp", "auto", "https://mcp.example.com/mcp"),
|
||||
("https://mcp.example.com/mcp", "AUTO", "https://mcp.example.com/mcp"),
|
||||
("https://mcp.example.com/mcp/", "auto", "https://mcp.example.com/mcp"),
|
||||
("https://mcp.example.com/", "auto", "https://mcp.example.com"),
|
||||
("https://MCP.Example.COM/mcp", "auto", "https://mcp.example.com/mcp"),
|
||||
("HTTPS://mcp.example.com/mcp", "auto", "https://mcp.example.com/mcp"),
|
||||
("https://mcp.example.com/mcp#frag", "auto", "https://mcp.example.com/mcp"),
|
||||
("https://mcp.example.com:8443/server/mcp", "auto", "https://mcp.example.com:8443/server/mcp"),
|
||||
# The scheme's default port is dropped, so :443/:80 never present as a different resource than
|
||||
# the portless form against the strict authorization servers this feature targets.
|
||||
("https://mcp.example.com:443/mcp", "auto", "https://mcp.example.com/mcp"),
|
||||
("http://mcp.example.com:80/mcp", "auto", "http://mcp.example.com/mcp"),
|
||||
# IPv6 authority keeps its brackets (a bare ::1:8080 would be a malformed authority).
|
||||
("https://[::1]:8080/mcp", "auto", "https://[::1]:8080/mcp"),
|
||||
("https://[::1]:443/mcp", "auto", "https://[::1]/mcp"),
|
||||
("https://mcp.example.com/Server/MCP", "auto", "https://mcp.example.com/Server/MCP"),
|
||||
("https://User:PaSs@MCP.Example.com/mcp", "auto", "https://mcp.example.com/mcp"),
|
||||
("https://token@MCP.Example.com/mcp", "auto", "https://mcp.example.com/mcp"),
|
||||
("https://mcp.example.com/mcp?api_key=s3cr3t", "auto", "https://mcp.example.com/mcp"),
|
||||
("https://u:p@MCP.Example.com:8443/mcp/?token=abc#frag", "auto", "https://mcp.example.com:8443/mcp"),
|
||||
("mcp.example.com/mcp", "auto", None),
|
||||
(None, "auto", None),
|
||||
("https://mcp.example.com/mcp", "api://custom-audience", "api://custom-audience"),
|
||||
("https://mcp.example.com/mcp", " https://Other.example.com/RS/ ", "https://Other.example.com/RS/"),
|
||||
],
|
||||
)
|
||||
def test_resolve_upstream_resource_tristate_and_canonicalization(url, configured, expected):
|
||||
"""The knob is a tri-state: unset/blank omits the parameter, ``auto`` derives the MCP spec's
|
||||
canonical server URI from the server url, and anything else is sent verbatim.
|
||||
|
||||
Canonicalization follows the MCP authorization spec: lowercase scheme and host, drop the scheme's
|
||||
default port, drop the fragment (RFC 8707 forbids one), drop the query and userinfo (credential
|
||||
hygiene), and drop a trailing slash, while preserving a non-default port and the path case. An
|
||||
explicit value is never canonicalized, because it has to match what the authorization server
|
||||
expects byte for byte."""
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import resolve_upstream_resource
|
||||
|
||||
assert resolve_upstream_resource(_resource_server(url=url, upstream_resource=configured)) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"configured, url, expected_resource",
|
||||
[
|
||||
(None, "https://mcp.example.com/mcp", None),
|
||||
("auto", "https://MCP.Example.com/mcp/", "https://mcp.example.com/mcp"),
|
||||
("api://audience", "https://mcp.example.com/mcp", "api://audience"),
|
||||
],
|
||||
)
|
||||
def test_build_upstream_oauth2_token_request_bundles_resource_with_client_auth(configured, url, expected_resource):
|
||||
"""Every plain-OAuth2 token leg (authorization_code, refresh_token, client_credentials) builds its
|
||||
request body through this one helper, so the RFC 8707 resource is resolved in the same call as the
|
||||
mandatory client authentication and no leg can authenticate without also naming the resource its
|
||||
sibling legs named. A leg that reverted to hand-building its body would drop the resource and
|
||||
diverge from the authorize leg, which a strict authorization server rejects as invalid_target."""
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request
|
||||
|
||||
req = build_upstream_oauth2_token_request(
|
||||
_resource_server(url=url, upstream_resource=configured),
|
||||
auth_method=None,
|
||||
client_id="cid",
|
||||
client_secret="sec",
|
||||
)
|
||||
assert req.body.get("resource") == expected_resource
|
||||
assert req.body["client_id"] == "cid"
|
||||
assert req.body["client_secret"] == "sec"
|
||||
|
||||
|
||||
def test_build_upstream_oauth2_token_request_client_secret_basic_keeps_secret_out_of_body():
|
||||
"""client_secret_basic authenticates through the Authorization header, so the secret must never
|
||||
also appear in the body, while the RFC 8707 resource still rides in the body."""
|
||||
import base64
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request
|
||||
|
||||
req = build_upstream_oauth2_token_request(
|
||||
_resource_server(upstream_resource="api://audience"),
|
||||
auth_method="client_secret_basic",
|
||||
client_id="cid",
|
||||
client_secret="sec",
|
||||
)
|
||||
assert req.headers["Authorization"] == "Basic " + base64.b64encode(b"cid:sec").decode()
|
||||
assert "client_secret" not in req.body
|
||||
assert "client_id" not in req.body
|
||||
assert req.body["resource"] == "api://audience"
|
||||
|
||||
|
||||
async def _authorize_query(server) -> dict:
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
authorize_with_server,
|
||||
)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
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:
|
||||
mock_encrypt.return_value = "encrypted_state"
|
||||
response = await authorize_with_server(
|
||||
request=mock_request,
|
||||
mcp_server=server,
|
||||
client_id="caller-client",
|
||||
redirect_uri="http://localhost:3000/callback",
|
||||
state="client-state",
|
||||
code_challenge="challenge",
|
||||
code_challenge_method="S256",
|
||||
response_type="code",
|
||||
scope=None,
|
||||
)
|
||||
return parse_qs(urlparse(response.headers["location"]).query)
|
||||
|
||||
|
||||
async def _token_body(server, grant_type: str) -> dict:
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
exchange_token_with_server,
|
||||
)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
mock_response.json.return_value = {"access_token": "at", "token_type": "Bearer"}
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
|
||||
return_value=mock_async_client,
|
||||
):
|
||||
await exchange_token_with_server(
|
||||
request=mock_request,
|
||||
mcp_server=server,
|
||||
grant_type=grant_type,
|
||||
code="auth-code" if grant_type == "authorization_code" else None,
|
||||
redirect_uri="https://litellm.example.com/callback",
|
||||
client_id="caller-client",
|
||||
client_secret=None,
|
||||
code_verifier="verifier",
|
||||
refresh_token="upstream-refresh" if grant_type == "refresh_token" else None,
|
||||
)
|
||||
return mock_async_client.post.call_args.kwargs["data"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_resource_unset_sends_no_resource_on_any_leg():
|
||||
"""Default behavior is unchanged: with the knob unset the gateway sends no RFC 8707 resource
|
||||
on the authorize leg or on either token grant, so every server working today keeps working
|
||||
(notably the authorization servers that hard-reject the parameter)."""
|
||||
server = _resource_server()
|
||||
|
||||
assert "resource" not in await _authorize_query(server)
|
||||
assert "resource" not in await _token_body(server, "authorization_code")
|
||||
assert "resource" not in await _token_body(server, "refresh_token")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_resource_auto_sends_same_canonical_uri_on_every_leg():
|
||||
"""The cross-leg invariant. RFC 8707 requires the token request to name a resource the
|
||||
authorization request already asked for, so the authorize leg and both token grants must send
|
||||
an identical value; they all resolve through one helper to make that structural. Deleting the
|
||||
resolve call at any single leg fails this test."""
|
||||
server = _resource_server(url="https://MCP.Example.com/mcp/", upstream_resource="auto")
|
||||
canonical = "https://mcp.example.com/mcp"
|
||||
|
||||
assert (await _authorize_query(server))["resource"] == [canonical]
|
||||
assert (await _token_body(server, "authorization_code"))["resource"] == canonical
|
||||
assert (await _token_body(server, "refresh_token"))["resource"] == canonical
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_resource_explicit_value_is_sent_verbatim_on_every_leg():
|
||||
"""An explicit identifier is never canonicalized or derived from the url; authorization servers
|
||||
match the resource exactly, so an operator-supplied value goes out byte for byte."""
|
||||
server = _resource_server(upstream_resource="api://7c9f-audience/.default")
|
||||
|
||||
assert (await _authorize_query(server))["resource"] == ["api://7c9f-audience/.default"]
|
||||
assert (await _token_body(server, "authorization_code"))["resource"] == "api://7c9f-audience/.default"
|
||||
assert (await _token_body(server, "refresh_token"))["resource"] == "api://7c9f-audience/.default"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_resource_auto_never_leaks_credentials_from_the_server_url():
|
||||
"""A resource indicator names the resource, never the credentials used to reach it. Transport
|
||||
URLs routinely carry secrets in userinfo and in the query string, and this value is published
|
||||
into the authorization redirect the browser follows and into token request bodies, so neither
|
||||
component may survive into the derived resource."""
|
||||
server = _resource_server(
|
||||
url="https://svc-account:s3cr3t@MCP.Example.com/mcp?api_key=qu3ry-s3cr3t",
|
||||
upstream_resource="auto",
|
||||
)
|
||||
leaks = ("s3cr3t", "svc-account", "qu3ry-s3cr3t", "api_key")
|
||||
|
||||
query = await _authorize_query(server)
|
||||
assert query["resource"] == ["https://mcp.example.com/mcp"]
|
||||
assert not any(leak in query["resource"][0] for leak in leaks)
|
||||
|
||||
body = await _token_body(server, "authorization_code")
|
||||
assert not any(leak in body["resource"] for leak in leaks)
|
||||
|
||||
|
||||
def test_upstream_resource_auto_keeps_the_path_because_it_identifies_the_server():
|
||||
"""The path is load-bearing identity and must survive canonicalization, unlike userinfo and
|
||||
query which are transport concerns.
|
||||
|
||||
The MCP authorization spec requires the most specific URI and lists
|
||||
``https://mcp.example.com/server/mcp`` as canonical "when path component is necessary to
|
||||
identify individual MCP server". Two servers behind one host differ only by path, so dropping
|
||||
it would collide them onto one resource identifier and bind each token to the wrong audience,
|
||||
which is the exact confusion RFC 8707 exists to prevent. An operator whose path embeds a secret
|
||||
sets ``upstream_resource`` explicitly instead of using ``auto``."""
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import resolve_upstream_resource
|
||||
|
||||
first = resolve_upstream_resource(_resource_server(url="https://gw.example.com/team-a/mcp", upstream_resource="auto"))
|
||||
second = resolve_upstream_resource(
|
||||
_resource_server(url="https://gw.example.com/team-b/mcp", upstream_resource="auto")
|
||||
)
|
||||
|
||||
assert first == "https://gw.example.com/team-a/mcp"
|
||||
assert second == "https://gw.example.com/team-b/mcp"
|
||||
assert first != second
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_resource_auto_without_url_omits_the_parameter():
|
||||
"""A server with no url (OpenAPI spec or stdio) has nothing to derive a canonical URI from, so
|
||||
``auto`` omits the parameter rather than sending an empty or malformed resource."""
|
||||
server = _resource_server(url=None, upstream_resource="auto")
|
||||
|
||||
assert "resource" not in await _authorize_query(server)
|
||||
assert "resource" not in await _token_body(server, "authorization_code")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_resource_sent_on_dcr_bridge_relay_authorize():
|
||||
"""The DCR-bridge relay arm builds its own upstream authorize params, so it needs the resource
|
||||
too. Without it the relayed authorize would omit the resource while the gateway's token leg
|
||||
still sent one, which is itself an invalid_target."""
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_dcr_bridge_relays_client_registration,
|
||||
)
|
||||
|
||||
server = _resource_server(
|
||||
client_id=None,
|
||||
client_secret=None,
|
||||
auth_type=MCPAuth.oauth_delegate,
|
||||
dcr_bridge=True,
|
||||
registration_url="https://idp.example.com/register",
|
||||
upstream_resource="auto",
|
||||
)
|
||||
assert _dcr_bridge_relays_client_registration(server), "test must exercise the relay arm"
|
||||
|
||||
query = await _authorize_query(server)
|
||||
assert query["resource"] == ["https://mcp.example.com/mcp"]
|
||||
assert query["client_id"] == ["caller-client"]
|
||||
|
|
|
|||
|
|
@ -218,6 +218,27 @@ async def test_auth_type_switch_clears_stale_flow_scoped_fields():
|
|||
assert _credentials_cleared(data_dict["credentials"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_explicit_null_clears_upstream_resource_and_keeps_the_rest_of_the_blob():
|
||||
"""The knob's own guidance tells an operator to unset it when the authorization server rejects
|
||||
resource indicators, so the edit form sends an explicit null for it rather than omitting it. The
|
||||
credential merge must drop that key while every omitted key still means keep-existing."""
|
||||
mock_prisma = _mock_prisma()
|
||||
existing = MagicMock()
|
||||
existing.auth_type = "oauth2"
|
||||
existing.url = "https://up.example.com/mcp"
|
||||
existing.credentials = json.dumps({"client_secret": "csec", "upstream_resource": "api://audience"})
|
||||
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing)
|
||||
|
||||
data = UpdateMCPServerRequest(server_id="my-test-server", credentials={"upstream_resource": None})
|
||||
await update_mcp_server(mock_prisma, data, "test-user")
|
||||
data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
|
||||
|
||||
merged = json.loads(data_dict["credentials"])
|
||||
assert merged["upstream_resource"] is None
|
||||
assert merged["client_secret"] == "csec"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_url_change_clears_stale_discovered_oauth_fields():
|
||||
"""Re-pointing the server url at a potentially different upstream must clear the discovered or
|
||||
|
|
|
|||
|
|
@ -1219,6 +1219,53 @@ class TestMCPServerManager:
|
|||
assert spec is not None and isinstance(spec.config, TokenExchangeConfig)
|
||||
assert spec.config.profile == "entra_obo"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_resource_survives_db_credentials_round_trip(self):
|
||||
"""A server persisted through the management API carries upstream_resource in its
|
||||
credentials blob, mirroring id_jag_resource. Without reading it back on the DB build, a
|
||||
UI-created server silently drops the knob and keeps hitting invalid_target."""
|
||||
manager = MCPServerManager()
|
||||
row = LiteLLM_MCPServerTable(
|
||||
server_id="res-db-1",
|
||||
alias="res_db",
|
||||
description="rfc8707 from db",
|
||||
url="https://up.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
credentials={
|
||||
"client_id": "cid",
|
||||
"client_secret": "csec",
|
||||
"authorization_url": "https://idp.example.com/authorize",
|
||||
"token_url": "https://idp.example.com/token",
|
||||
"upstream_resource": "https://up.example.com/mcp",
|
||||
},
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
|
||||
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
|
||||
assert built.upstream_resource == "https://up.example.com/mcp"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_resource_loads_from_config(self):
|
||||
"""The config.yaml arm of the same field: mcp_servers entries must carry the knob onto the
|
||||
registry entry, since a config-declared server never round-trips through the DB."""
|
||||
manager = MCPServerManager()
|
||||
await manager.load_servers_from_config(
|
||||
{
|
||||
"strict_as": {
|
||||
"url": "https://strict.example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
"auth_type": MCPAuth.oauth2,
|
||||
"oauth2_flow": "authorization_code",
|
||||
"upstream_resource": "auto",
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
loaded = next(s for s in manager.get_registry().values() if s.name == "strict_as")
|
||||
assert loaded.upstream_resource == "auto"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate])
|
||||
async def test_build_from_table_discovers_upstream_oauth_for_client_forwarded_modes(self, auth_type):
|
||||
|
|
|
|||
|
|
@ -290,3 +290,105 @@ def test_default_ttl_paths_unchanged_without_storage_ttl():
|
|||
server = _server(oauth2_flow=None)
|
||||
assert _compute_per_user_token_ttl(server, expires_in=86400) == 86400 - MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS
|
||||
assert _compute_per_user_token_ttl(server, expires_in=None) == MCP_PER_USER_TOKEN_DEFAULT_TTL
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"configured, expected",
|
||||
[
|
||||
(None, None),
|
||||
("auto", "https://mcp.example.com/mcp"),
|
||||
("api://m2m-audience", "api://m2m-audience"),
|
||||
],
|
||||
)
|
||||
async def test_client_credentials_sends_rfc8707_resource(configured, expected):
|
||||
"""The client_credentials fetch carries the RFC 8707 resource indicator too, resolved through
|
||||
the same helper the interactive legs use, so the knob means one thing for every oauth2 flow on
|
||||
a server. Unset omits it, which is the default and preserves today's request body."""
|
||||
server = _server(server_id=f"srv-{configured}", upstream_resource=configured)
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = _token_response("m2m-tok")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
|
||||
return_value=mock_client,
|
||||
):
|
||||
await resolve_mcp_auth(server)
|
||||
|
||||
post_data = mock_client.post.call_args[1]["data"]
|
||||
assert post_data.get("resource") == expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"changed",
|
||||
[
|
||||
{"upstream_resource": "api://new-audience"},
|
||||
{"scopes": ["other.scope"]},
|
||||
{"client_secret": "rotated-secret"},
|
||||
{"token_url": "https://auth.example.com/other/token"},
|
||||
],
|
||||
)
|
||||
async def test_token_cache_mints_afresh_when_the_token_request_changes(changed):
|
||||
"""A minted token is only reusable for the exact request that produced it. Keying the cache on
|
||||
server_id alone kept serving a token carrying the previous scopes, secret, or audience until it
|
||||
expired, so setting upstream_resource on a live server appeared to do nothing. Each input that
|
||||
reaches the wire must miss the cache."""
|
||||
cache = MCPOAuth2TokenCache()
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.side_effect = [_token_response("tok-before"), _token_response("tok-after")]
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
|
||||
return_value=mock_client,
|
||||
):
|
||||
before = await cache.async_get_token(_server())
|
||||
after = await cache.async_get_token(_server(**changed))
|
||||
|
||||
assert before == "tok-before"
|
||||
assert after == "tok-after"
|
||||
assert mock_client.post.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_cache_still_reuses_a_token_when_nothing_changed():
|
||||
"""The flip side: an unchanged config must keep hitting the cache, so the identity key does not
|
||||
turn every call into a fresh mint."""
|
||||
cache = MCPOAuth2TokenCache()
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = _token_response("tok-reused")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
|
||||
return_value=mock_client,
|
||||
):
|
||||
first = await cache.async_get_token(_server(upstream_resource="auto"))
|
||||
second = await cache.async_get_token(_server(upstream_resource="auto"))
|
||||
|
||||
assert first == second == "tok-reused"
|
||||
assert mock_client.post.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalidate_clears_every_identity_for_a_server():
|
||||
"""A 401 invalidates the server, not one configuration of it, so entries left behind by an
|
||||
earlier config cannot be served after the eviction."""
|
||||
cache = MCPOAuth2TokenCache()
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.side_effect = [
|
||||
_token_response("tok-a"),
|
||||
_token_response("tok-b"),
|
||||
_token_response("tok-after-invalidate"),
|
||||
]
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
|
||||
return_value=mock_client,
|
||||
):
|
||||
await cache.async_get_token(_server())
|
||||
await cache.async_get_token(_server(upstream_resource="api://second"))
|
||||
cache.invalidate("srv-1")
|
||||
refetched = await cache.async_get_token(_server())
|
||||
|
||||
assert refetched == "tok-after-invalidate"
|
||||
assert mock_client.post.call_count == 3
|
||||
|
|
|
|||
|
|
@ -407,3 +407,65 @@ def test_repeated_db_sync_does_not_accumulate_runner_instances():
|
|||
finally:
|
||||
for cb_list, snapshot in zip(lists, snapshots):
|
||||
cb_list[:] = snapshot
|
||||
|
||||
|
||||
def _judge_guardrail(guardrail_id: str) -> Guardrail:
|
||||
return Guardrail(
|
||||
guardrail_id=guardrail_id,
|
||||
guardrail_name="quality-judge",
|
||||
litellm_params={
|
||||
"guardrail": "llm_as_a_judge",
|
||||
"mode": "post_call",
|
||||
"judge_model": "my-judge-alias",
|
||||
"overall_threshold": 80,
|
||||
"on_failure": "log",
|
||||
"criteria": [{"name": "helpfulness", "weight": 100, "description": "helpful?"}],
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_db_synced_judge_guardrail_uses_lazy_router_provider():
|
||||
"""A judge guardrail created/synced through a DB path must resolve the active
|
||||
Router lazily at call time (issue: UI-created guardrails failed open because the
|
||||
Router was captured at construction; a guardrail created before the Router
|
||||
existed captured None and never recovered). Asserting the default provider is
|
||||
wired guarantees the instance reads the live global rather than a stale value."""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge import (
|
||||
LLMAsAJudgeGuardrail,
|
||||
_default_router_provider,
|
||||
)
|
||||
|
||||
handler = InMemoryGuardrailHandler()
|
||||
|
||||
lists = _all_callback_lists()
|
||||
snapshots = [list(cb_list) for cb_list in lists]
|
||||
try:
|
||||
handler.sync_guardrail_from_db(_judge_guardrail("judge-db"))
|
||||
|
||||
instance = handler.guardrail_id_to_custom_guardrail["judge-db"]
|
||||
assert isinstance(instance, LLMAsAJudgeGuardrail)
|
||||
assert instance._router_provider is _default_router_provider
|
||||
finally:
|
||||
for cb_list, snapshot in zip(lists, snapshots):
|
||||
cb_list[:] = snapshot
|
||||
|
||||
|
||||
def test_reinitialized_judge_guardrail_uses_lazy_router_provider():
|
||||
from litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge import (
|
||||
LLMAsAJudgeGuardrail,
|
||||
_default_router_provider,
|
||||
)
|
||||
|
||||
handler = InMemoryGuardrailHandler()
|
||||
|
||||
lists = _all_callback_lists()
|
||||
snapshots = [list(cb_list) for cb_list in lists]
|
||||
try:
|
||||
handler.reinitialize_guardrail(_judge_guardrail("judge-reinit"), source="db")
|
||||
|
||||
instance = handler.guardrail_id_to_custom_guardrail["judge-reinit"]
|
||||
assert isinstance(instance, LLMAsAJudgeGuardrail)
|
||||
assert instance._router_provider is _default_router_provider
|
||||
finally:
|
||||
for cb_list, snapshot in zip(lists, snapshots):
|
||||
cb_list[:] = snapshot
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge import (
|
|||
LLMAsAJudgeGuardrail,
|
||||
_build_judge_prompt,
|
||||
_extract_text_from_content,
|
||||
_parse_judge_verdict,
|
||||
initialize_guardrail,
|
||||
)
|
||||
|
||||
|
|
@ -198,6 +199,86 @@ async def test_apply_guardrail_log_mode_does_not_block(mock_completion):
|
|||
assert request_data["metadata"]["eval_information"]["passed"] is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _parse_judge_verdict — tolerate fenced/prose-wrapped JSON
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_parse_judge_verdict_plain_json():
|
||||
assert _parse_judge_verdict('{"overall_score": 90}')["overall_score"] == 90
|
||||
|
||||
|
||||
def test_parse_judge_verdict_strips_json_fence_and_prose():
|
||||
raw = 'Here is my verdict:\n```json\n{"overall_score": 42}\n```\nHope that helps'
|
||||
assert _parse_judge_verdict(raw)["overall_score"] == 42
|
||||
|
||||
|
||||
def test_parse_judge_verdict_strips_bare_fence():
|
||||
raw = '```\n{"overall_score": 7}\n```'
|
||||
assert _parse_judge_verdict(raw)["overall_score"] == 7
|
||||
|
||||
|
||||
def test_parse_judge_verdict_extracts_json_from_surrounding_prose():
|
||||
raw = 'Sure, here it is: {"overall_score": 55} let me know'
|
||||
assert _parse_judge_verdict(raw)["overall_score"] == 55
|
||||
|
||||
|
||||
def test_parse_judge_verdict_reraises_when_no_json():
|
||||
with pytest.raises(json.JSONDecodeError):
|
||||
_parse_judge_verdict("no json here")
|
||||
|
||||
|
||||
def test_parse_judge_verdict_rejects_json_non_object():
|
||||
"""Valid JSON that is not an object (e.g. a bare list) raises ValueError."""
|
||||
with pytest.raises(ValueError):
|
||||
_parse_judge_verdict("[1, 2, 3]")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion")
|
||||
async def test_apply_guardrail_enforces_fenced_verdict(mock_completion):
|
||||
"""A failing verdict wrapped in a code fence blocks with a 422."""
|
||||
fenced = "```json\n" + json.dumps(_make_verdict_response(50.0)) + "\n```"
|
||||
mock_completion.return_value = MagicMock(choices=[MagicMock(message=MagicMock(content=fenced))])
|
||||
guardrail = _make_guardrail(overall_threshold=80.0, on_failure="block", router_provider=lambda: None)
|
||||
inputs = {"texts": ["bad response"]}
|
||||
request_data: dict = {"messages": [], "metadata": {}}
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.apply_guardrail(inputs, request_data, "response")
|
||||
assert exc_info.value.status_code == 422
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion")
|
||||
async def test_apply_guardrail_non_object_verdict_fails_open_with_status(mock_completion):
|
||||
"""A non-object verdict fails open and logs guardrail_failed_to_respond."""
|
||||
mock_completion.return_value = MagicMock(
|
||||
choices=[MagicMock(message=MagicMock(content='[{"overall_score": 50}]'))]
|
||||
)
|
||||
guardrail = _make_guardrail(overall_threshold=80.0, on_failure="block", router_provider=lambda: None)
|
||||
inputs = {"texts": ["response"]}
|
||||
request_data: dict = {"messages": [], "metadata": {}}
|
||||
result = await guardrail.apply_guardrail(inputs, request_data, "response")
|
||||
assert result is inputs
|
||||
logged = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
assert logged[0]["guardrail_status"] == "guardrail_failed_to_respond"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion")
|
||||
async def test_apply_guardrail_parses_fenced_json_verdict(mock_completion):
|
||||
"""Fencing-prone judge models wrap the verdict in a ```json fence; the guardrail
|
||||
must parse it and evaluate rather than failing open on json.loads."""
|
||||
fenced = "```json\n" + json.dumps(_make_verdict_response(90.0)) + "\n```"
|
||||
mock_completion.return_value = MagicMock(choices=[MagicMock(message=MagicMock(content=fenced))])
|
||||
guardrail = _make_guardrail(overall_threshold=80.0)
|
||||
inputs = {"texts": ["good response"]}
|
||||
request_data: dict = {"messages": [], "metadata": {}}
|
||||
result = await guardrail.apply_guardrail(inputs, request_data, "response")
|
||||
assert result is inputs
|
||||
assert request_data["metadata"]["eval_information"]["passed"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion")
|
||||
async def test_apply_guardrail_judge_error_fails_open(mock_completion):
|
||||
|
|
@ -209,6 +290,167 @@ async def test_apply_guardrail_judge_error_fails_open(mock_completion):
|
|||
assert result is inputs
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# judge_model credential/provider resolution — route through the proxy Router
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _judge_response_mock() -> MagicMock:
|
||||
return MagicMock(choices=[MagicMock(message=MagicMock(content=json.dumps(_make_verdict_response(90.0))))])
|
||||
|
||||
|
||||
def _real_router(model_list, **router_kwargs):
|
||||
"""Build a real Router so the router-membership decision is exercised for
|
||||
real (wildcards, model_group_alias, exact names), stubbing only the outbound
|
||||
completion so no network call is made."""
|
||||
from litellm import Router
|
||||
|
||||
router = Router(model_list=model_list, **router_kwargs)
|
||||
router.acompletion = AsyncMock(return_value=_judge_response_mock())
|
||||
return router
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_list, router_kwargs, judge_model",
|
||||
[
|
||||
(
|
||||
[{"model_name": "my-judge-alias", "litellm_params": {"model": "anthropic/claude-sonnet-4-6", "api_key": "sk-ant-test"}}],
|
||||
{},
|
||||
"my-judge-alias",
|
||||
),
|
||||
(
|
||||
[{"model_name": "anthropic/*", "litellm_params": {"model": "anthropic/*", "api_key": "sk-ant-test"}}],
|
||||
{},
|
||||
"anthropic/claude-sonnet-4-6",
|
||||
),
|
||||
(
|
||||
[{"model_name": "backing-group", "litellm_params": {"model": "anthropic/claude-sonnet-4-6", "api_key": "sk-ant-test"}}],
|
||||
{"model_group_alias": {"my-judge-alias": "backing-group"}},
|
||||
"my-judge-alias",
|
||||
),
|
||||
(
|
||||
[{"model_name": "backing-group", "litellm_params": {"model": "anthropic/claude-sonnet-4-6", "api_key": "sk-ant-test"}}],
|
||||
{"model_group_alias": {"my-judge-alias": {"model": "backing-group", "hidden": True}}},
|
||||
"my-judge-alias",
|
||||
),
|
||||
],
|
||||
ids=["plain-deployment", "wildcard-route", "model-group-alias", "hidden-model-group-alias"],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion", new_callable=AsyncMock)
|
||||
async def test_judge_routes_through_router_for_router_served_model(
|
||||
mock_sdk_completion, model_list, router_kwargs, judge_model
|
||||
):
|
||||
"""Any judge_model the Router can serve must resolve its credentials via the
|
||||
Router. Wildcard and alias shapes regress the naive `judge_model in
|
||||
get_model_names()` check, which reports patterns/aliases literally and so
|
||||
routes a servable model to the SDK, where deployment creds do not resolve."""
|
||||
router = _real_router(model_list, **router_kwargs)
|
||||
guardrail = _make_guardrail(judge_model=judge_model, router_provider=lambda: router)
|
||||
inputs = {"texts": ["good response"]}
|
||||
request_data: dict = {"messages": [{"role": "user", "content": "hi"}], "metadata": {}}
|
||||
|
||||
result = await guardrail.apply_guardrail(inputs, request_data, "response")
|
||||
|
||||
assert result is inputs
|
||||
router.acompletion.assert_awaited_once()
|
||||
call_kwargs = router.acompletion.await_args.kwargs
|
||||
assert call_kwargs["model"] == judge_model
|
||||
assert call_kwargs["num_retries"] == 0
|
||||
assert call_kwargs["fallbacks"] == []
|
||||
mock_sdk_completion.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion", new_callable=AsyncMock)
|
||||
async def test_judge_call_falls_back_to_sdk_when_model_not_in_router(mock_sdk_completion):
|
||||
"""A judge_model the Router cannot serve (e.g. a raw provider model resolved
|
||||
from the environment) must fall back to the SDK."""
|
||||
mock_sdk_completion.return_value = _judge_response_mock()
|
||||
router = _real_router(
|
||||
[{"model_name": "some-other-model", "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-test"}}]
|
||||
)
|
||||
guardrail = _make_guardrail(judge_model="gpt-4o-mini", router_provider=lambda: router)
|
||||
inputs = {"texts": ["good response"]}
|
||||
request_data: dict = {"messages": [], "metadata": {}}
|
||||
|
||||
result = await guardrail.apply_guardrail(inputs, request_data, "response")
|
||||
|
||||
assert result is inputs
|
||||
router.acompletion.assert_not_called()
|
||||
mock_sdk_completion.assert_awaited_once()
|
||||
assert mock_sdk_completion.await_args.kwargs["model"] == "gpt-4o-mini"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion", new_callable=AsyncMock)
|
||||
async def test_judge_call_uses_sdk_when_no_router(mock_sdk_completion):
|
||||
mock_sdk_completion.return_value = _judge_response_mock()
|
||||
guardrail = _make_guardrail(judge_model="gpt-4o-mini", router_provider=lambda: None)
|
||||
inputs = {"texts": ["good response"]}
|
||||
request_data: dict = {"messages": [], "metadata": {}}
|
||||
|
||||
result = await guardrail.apply_guardrail(inputs, request_data, "response")
|
||||
|
||||
assert result is inputs
|
||||
mock_sdk_completion.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion", new_callable=AsyncMock)
|
||||
async def test_judge_resolves_router_lazily_per_call(mock_sdk_completion):
|
||||
"""The Router is resolved at call time, not captured at construction. A
|
||||
guardrail built before the proxy Router exists (provider returns None) starts
|
||||
routing through the Router as soon as it is available, with no re-init. This
|
||||
regresses the config-less DB-backed startup order where the guardrail was
|
||||
created while the global Router was still None and then never recovered."""
|
||||
mock_sdk_completion.return_value = _judge_response_mock()
|
||||
holder: dict = {"router": None}
|
||||
guardrail = _make_guardrail(judge_model="my-judge-alias", router_provider=lambda: holder["router"])
|
||||
|
||||
await guardrail.apply_guardrail({"texts": ["r"]}, {"messages": [], "metadata": {}}, "response")
|
||||
mock_sdk_completion.assert_awaited_once()
|
||||
|
||||
holder["router"] = _real_router(
|
||||
[{"model_name": "my-judge-alias", "litellm_params": {"model": "anthropic/claude-sonnet-4-6", "api_key": "sk-ant-test"}}]
|
||||
)
|
||||
await guardrail.apply_guardrail({"texts": ["r"]}, {"messages": [], "metadata": {}}, "response")
|
||||
holder["router"].acompletion.assert_awaited_once()
|
||||
mock_sdk_completion.assert_awaited_once()
|
||||
|
||||
|
||||
def test_default_router_provider_returns_none_when_proxy_not_importable():
|
||||
"""If the proxy dependency set is not importable, the provider must return None
|
||||
so the judge falls back to the SDK rather than the ImportError being swallowed
|
||||
by the fail-open handler and the guardrail silently no-opping."""
|
||||
import sys
|
||||
|
||||
from litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge import _default_router_provider
|
||||
|
||||
with patch.dict(sys.modules, {"litellm.proxy.proxy_server": None}):
|
||||
assert _default_router_provider() is None
|
||||
|
||||
|
||||
def test_default_router_provider_reads_global_router():
|
||||
"""The default provider must read the live proxy global so the router is
|
||||
resolved lazily rather than captured."""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge import _default_router_provider
|
||||
|
||||
sentinel = object()
|
||||
with patch("litellm.proxy.proxy_server.llm_router", sentinel):
|
||||
assert _default_router_provider() is sentinel
|
||||
|
||||
|
||||
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.logging_callback_manager")
|
||||
def test_initialize_guardrail_uses_default_router_provider(mock_mgr):
|
||||
from litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge import _default_router_provider
|
||||
|
||||
lp = _make_litellm_params()
|
||||
g = _make_guardrail_dict(judge_model="my-judge-alias")
|
||||
instance = initialize_guardrail(lp, g)
|
||||
assert instance._router_provider is _default_router_provider
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion")
|
||||
async def test_apply_guardrail_clamps_score(mock_completion):
|
||||
|
|
|
|||
|
|
@ -761,6 +761,151 @@ class TestListMCPServers:
|
|||
assert mock_server.credentials == {"auth_value": "top-secret"}
|
||||
assert result.status == "healthy"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_single_mcp_server_preserves_upstream_resource_for_admin(self):
|
||||
"""upstream_resource is non-secret admin config, so the admin edit form must receive its real
|
||||
value to change or clear it; secrets sharing the blob are still dropped."""
|
||||
mock_server = generate_mock_mcp_server_db_record(server_id="server-ur", alias="UR")
|
||||
mock_server.credentials = {"client_secret": "top-secret", "upstream_resource": "api://audience"}
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_health_result = generate_mock_mcp_server_db_record(server_id="server-ur", alias="UR")
|
||||
mock_health_result.status = "healthy"
|
||||
mock_health_result.last_health_check = datetime.now()
|
||||
mock_health_result.health_check_error = None
|
||||
mock_user_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=mock_prisma_client,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
AsyncMock(return_value=mock_server),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager.health_check_server",
|
||||
AsyncMock(return_value=mock_health_result),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
fetch_mcp_server,
|
||||
)
|
||||
|
||||
result = await fetch_mcp_server(
|
||||
request=_make_mock_request(),
|
||||
server_id="server-ur",
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
assert result.credentials == {"upstream_resource": "api://audience"}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"stored_credentials, expected",
|
||||
[
|
||||
({"client_secret": "s", "upstream_resource": "api://audience"}, {"upstream_resource": "api://audience"}),
|
||||
('{"client_secret": "s", "upstream_resource": "api://audience"}', {"upstream_resource": "api://audience"}),
|
||||
("not-json{{", None),
|
||||
("null", None),
|
||||
("{}", None),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_single_mcp_server_redaction_is_total_over_malformed_credentials(
|
||||
self, stored_credentials, expected
|
||||
):
|
||||
"""Redaction runs on every admin list and get, so a row whose credentials blob is a corrupt or
|
||||
non-object JSON string must fall back to full redaction rather than raise and fail the whole
|
||||
request. A valid JSON-object string still has its admin config lifted out. Bare non-object JSON
|
||||
(a list or scalar) is not a reachable stored shape, since writes always persist a JSON object."""
|
||||
mock_server = generate_mock_mcp_server_db_record(server_id="server-mal", alias="MAL")
|
||||
mock_server.credentials = stored_credentials
|
||||
|
||||
mock_health_result = generate_mock_mcp_server_db_record(server_id="server-mal", alias="MAL")
|
||||
mock_health_result.status = "healthy"
|
||||
mock_health_result.last_health_check = datetime.now()
|
||||
mock_health_result.health_check_error = None
|
||||
mock_user_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
AsyncMock(return_value=mock_server),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager.health_check_server",
|
||||
AsyncMock(return_value=mock_health_result),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
fetch_mcp_server,
|
||||
)
|
||||
|
||||
result = await fetch_mcp_server(
|
||||
request=_make_mock_request(),
|
||||
server_id="server-mal",
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
assert result.credentials == expected
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_single_mcp_server_strips_upstream_resource_for_non_admin(self):
|
||||
"""A non-full-admin viewer gets the whole blob nulled, including the non-secret admin config,
|
||||
so admin-typed settings never leak to a discovery-only caller."""
|
||||
mock_server = generate_mock_mcp_server_db_record(server_id="server-ur2", alias="UR2")
|
||||
mock_server.credentials = {"client_secret": "top-secret", "upstream_resource": "api://audience"}
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_health_result = generate_mock_mcp_server_db_record(server_id="server-ur2", alias="UR2")
|
||||
mock_health_result.status = "healthy"
|
||||
mock_health_result.last_health_check = datetime.now()
|
||||
mock_health_result.health_check_error = None
|
||||
mock_user_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=mock_prisma_client,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
AsyncMock(return_value=mock_server),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager.health_check_server",
|
||||
AsyncMock(return_value=mock_health_result),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
fetch_mcp_server,
|
||||
)
|
||||
|
||||
result = await fetch_mcp_server(
|
||||
request=_make_mock_request(),
|
||||
server_id="server-ur2",
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
assert result.credentials is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_single_mcp_server_handles_missing_credentials_field(self):
|
||||
mock_server = generate_mock_mcp_server_db_record(server_id="server-2", alias="Server 2")
|
||||
|
|
@ -1428,6 +1573,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
existing_server.aws_session_token = None
|
||||
existing_server.aws_region_name = None
|
||||
existing_server.aws_service_name = None
|
||||
existing_server.upstream_resource = None
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id.return_value = existing_server
|
||||
|
|
@ -1450,6 +1596,68 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
}
|
||||
mock_manager.get_mcp_server_by_id.assert_called_once_with("server-123")
|
||||
|
||||
@staticmethod
|
||||
def _inherit_with(payload_credentials, **server_overrides):
|
||||
existing_server = MagicMock()
|
||||
existing_server.authentication_token = None
|
||||
existing_server.client_id = "client-123"
|
||||
existing_server.client_secret = "secret-xyz"
|
||||
existing_server.scopes = None
|
||||
existing_server.aws_access_key_id = None
|
||||
existing_server.aws_secret_access_key = None
|
||||
existing_server.aws_session_token = None
|
||||
existing_server.aws_region_name = None
|
||||
existing_server.aws_service_name = None
|
||||
existing_server.upstream_resource = None
|
||||
for key, value in server_overrides.items():
|
||||
setattr(existing_server, key, value)
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
server_id="server-123",
|
||||
alias="Temp Server",
|
||||
url="https://temp.example.com",
|
||||
transport=MCPTransport.http,
|
||||
credentials=payload_credentials,
|
||||
)
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id.return_value = existing_server
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_inherit_credentials_from_existing_server,
|
||||
)
|
||||
|
||||
return _inherit_credentials_from_existing_server(payload)
|
||||
|
||||
def test_admin_config_alone_does_not_suppress_credential_inheritance(self):
|
||||
"""The edit form round-trips upstream_resource, which is admin config rather than a credential.
|
||||
Treating the blob as "credentials supplied" left the Authorize session with no declared app on
|
||||
the exact path where this knob is configured."""
|
||||
updated = self._inherit_with({"upstream_resource": "api://audience"})
|
||||
|
||||
assert updated.credentials["client_id"] == "client-123"
|
||||
assert updated.credentials["client_secret"] == "secret-xyz"
|
||||
|
||||
def test_supplied_credential_still_wins_over_inheritance(self):
|
||||
"""A caller that supplies a real credential keeps it; inheritance must not overwrite it."""
|
||||
updated = self._inherit_with({"auth_value": "caller-token"})
|
||||
|
||||
assert updated.credentials == {"auth_value": "caller-token"}
|
||||
|
||||
def test_inheritance_carries_upstream_resource_to_the_session_server(self):
|
||||
"""Without this the temporary server omits the resource indicator and the Authorize leg it
|
||||
exists for fails as invalid_target."""
|
||||
updated = self._inherit_with(None, upstream_resource="api://stored")
|
||||
|
||||
assert updated.credentials["upstream_resource"] == "api://stored"
|
||||
|
||||
def test_supplied_upstream_resource_wins_over_the_stored_one(self):
|
||||
updated = self._inherit_with({"upstream_resource": "api://typed"}, upstream_resource="api://stored")
|
||||
|
||||
assert updated.credentials["upstream_resource"] == "api://typed"
|
||||
|
||||
def test_cache_temporary_mcp_server_stores_entry_with_ttl(self):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_cache_temporary_mcp_server,
|
||||
|
|
@ -1686,6 +1894,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
aws_session_token=None,
|
||||
aws_region_name=None,
|
||||
aws_service_name=None,
|
||||
upstream_resource=None,
|
||||
)
|
||||
built_server = generate_mock_mcp_server_config_record(server_id="temp-server")
|
||||
mock_manager = MagicMock()
|
||||
|
|
|
|||
|
|
@ -102,7 +102,7 @@ def test_opus_4_6_model_pricing_and_capabilities():
|
|||
"azure_ai/claude-opus-4-6": {
|
||||
"provider": "azure_ai",
|
||||
"has_long_context_pricing": False,
|
||||
"max_input_tokens": 200000,
|
||||
"max_input_tokens": 1000000,
|
||||
},
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -60,10 +60,9 @@ def test_opus_4_8_model_pricing_and_capabilities():
|
|||
"provider": "vertex_ai-anthropic_models",
|
||||
"max_input_tokens": 1000000,
|
||||
},
|
||||
# Microsoft Foundry / Azure caps Opus 4.8 at a 200k context window.
|
||||
"azure_ai/claude-opus-4-8": {
|
||||
"provider": "azure_ai",
|
||||
"max_input_tokens": 200000,
|
||||
"max_input_tokens": 1000000,
|
||||
},
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -982,9 +982,6 @@
|
|||
"max-lines": {
|
||||
"count": 1
|
||||
},
|
||||
"no-nested-ternary": {
|
||||
"count": 1
|
||||
},
|
||||
"no-restricted-imports": {
|
||||
"count": 2
|
||||
},
|
||||
|
|
|
|||
|
|
@ -29,6 +29,63 @@ describe("OAuthFormFields", () => {
|
|||
|
||||
// ── visibility by flow type ─────────────────────────────────────────────────
|
||||
|
||||
// The RFC 8707 resource indicator applies to both OAuth arms: the interactive authorize/token legs
|
||||
// and the M2M client_credentials fetch. It must render in each, or the arm missing it can only be
|
||||
// configured through the API.
|
||||
describe("resource indicator field", () => {
|
||||
it("renders in interactive mode", () => {
|
||||
render(
|
||||
<WithForm>
|
||||
<OAuthFormFields isM2M={false} />
|
||||
</WithForm>,
|
||||
);
|
||||
expect(screen.getByText("Resource Indicator (optional)")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("renders in M2M mode", () => {
|
||||
render(
|
||||
<WithForm>
|
||||
<OAuthFormFields isM2M={true} />
|
||||
</WithForm>,
|
||||
);
|
||||
expect(screen.getByText("Resource Indicator (optional)")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("keeps one placeholder when editing, since the stored value is returned and shown", () => {
|
||||
// Non-secret admin config is no longer redacted out of responses, so the field mounts with its
|
||||
// real value and an emptied field clears it. There is no keep-existing state left to signal.
|
||||
render(
|
||||
<WithForm>
|
||||
<OAuthFormFields isM2M={false} isEditing={true} />
|
||||
</WithForm>,
|
||||
);
|
||||
expect(screen.getByPlaceholderText("auto, or https://mcp.example.com/mcp")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("submits its value under credentials.upstream_resource", async () => {
|
||||
const onFinish = vi.fn();
|
||||
render(
|
||||
<WithForm onFinish={onFinish}>
|
||||
<OAuthFormFields isM2M={false} />
|
||||
</WithForm>,
|
||||
);
|
||||
const input = screen.getByPlaceholderText("auto, or https://mcp.example.com/mcp");
|
||||
await act(async () => {
|
||||
fireEvent.change(input, { target: { value: "api://finance-api/.default" } });
|
||||
});
|
||||
await act(async () => {
|
||||
fireEvent.click(screen.getByText("Submit"));
|
||||
});
|
||||
await waitFor(() => {
|
||||
expect(onFinish).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
credentials: expect.objectContaining({ upstream_resource: "api://finance-api/.default" }),
|
||||
}),
|
||||
);
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("interactive mode (isM2M=false)", () => {
|
||||
it("renders Token Validation Rules field", () => {
|
||||
render(
|
||||
|
|
|
|||
|
|
@ -23,6 +23,13 @@ interface OAuthFormFieldsProps {
|
|||
|
||||
const fieldClassName = "rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500";
|
||||
|
||||
const UPSTREAM_RESOURCE_TOOLTIP =
|
||||
"RFC 8707 resource indicator sent to the authorization server so it mints a token audienced for this MCP server. " +
|
||||
"Leave blank to send nothing, which is the default and what most providers expect. Use 'auto' to send this server's " +
|
||||
"own URL. Set an exact identifier when the authorization server expects a specific one. Some providers reject this " +
|
||||
"parameter and take the audience from scopes instead; if you see AADSTS901002, leave it blank. If you see " +
|
||||
"invalid_target, the authorization server needs it set.";
|
||||
|
||||
const FieldLabel: React.FC<{ label: string; tooltip: string }> = ({ label, tooltip }) => (
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
{label}
|
||||
|
|
@ -32,6 +39,15 @@ const FieldLabel: React.FC<{ label: string; tooltip: string }> = ({ label, toolt
|
|||
</span>
|
||||
);
|
||||
|
||||
const UpstreamResourceField: React.FC = () => (
|
||||
<Form.Item
|
||||
label={<FieldLabel label="Resource Indicator (optional)" tooltip={UPSTREAM_RESOURCE_TOOLTIP} />}
|
||||
name={["credentials", "upstream_resource"]}
|
||||
>
|
||||
<TextInput placeholder="auto, or https://mcp.example.com/mcp" className={fieldClassName} />
|
||||
</Form.Item>
|
||||
);
|
||||
|
||||
const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
|
||||
isM2M,
|
||||
isEditing = false,
|
||||
|
|
@ -40,6 +56,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
|
|||
docsUrl,
|
||||
}) => {
|
||||
const placeholderSuffix = isEditing ? " (leave blank to keep existing)" : "";
|
||||
const requiredWhenCreating = (message: string) => (isEditing ? [] : [{ required: true, message }]);
|
||||
|
||||
return (
|
||||
<>
|
||||
|
|
@ -53,7 +70,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
|
|||
name="oauth_flow_type"
|
||||
{...(initialFlowType ? { initialValue: initialFlowType } : {})}
|
||||
>
|
||||
<Select className="rounded-lg" size="large">
|
||||
<Select placeholder="Select OAuth flow" className="rounded-lg" size="large">
|
||||
<Select.Option value={OAUTH_FLOW.M2M}>
|
||||
<div>
|
||||
<span className="font-medium">Machine-to-Machine (M2M)</span>
|
||||
|
|
@ -74,7 +91,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
|
|||
<Form.Item
|
||||
label={<FieldLabel label="Client ID" tooltip="OAuth2 client ID for the client_credentials grant." />}
|
||||
name={["credentials", "client_id"]}
|
||||
rules={[{ required: true, message: "Client ID is required for M2M OAuth" }]}
|
||||
rules={requiredWhenCreating("Client ID is required for M2M OAuth")}
|
||||
>
|
||||
<TextInput
|
||||
type="password"
|
||||
|
|
@ -87,7 +104,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
|
|||
<FieldLabel label="Client Secret" tooltip="OAuth2 client secret for the client_credentials grant." />
|
||||
}
|
||||
name={["credentials", "client_secret"]}
|
||||
rules={[{ required: true, message: "Client Secret is required for M2M OAuth" }]}
|
||||
rules={requiredWhenCreating("Client Secret is required for M2M OAuth")}
|
||||
>
|
||||
<TextInput
|
||||
type="password"
|
||||
|
|
@ -98,7 +115,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
|
|||
<Form.Item
|
||||
label={<FieldLabel label="Token URL" tooltip="Token endpoint URL for the client_credentials grant." />}
|
||||
name="token_url"
|
||||
rules={[{ required: true, message: "Token URL is required for M2M OAuth" }]}
|
||||
rules={requiredWhenCreating("Token URL is required for M2M OAuth")}
|
||||
>
|
||||
<TextInput placeholder="https://auth.example.com/oauth/token" className={fieldClassName} />
|
||||
</Form.Item>
|
||||
|
|
@ -114,6 +131,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
|
|||
>
|
||||
<Select mode="tags" tokenSeparators={[","]} placeholder="Add scopes" className="rounded-lg" size="large" />
|
||||
</Form.Item>
|
||||
<UpstreamResourceField />
|
||||
</>
|
||||
) : (
|
||||
<>
|
||||
|
|
@ -167,6 +185,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
|
|||
>
|
||||
<Select mode="tags" tokenSeparators={[","]} placeholder="Add scopes" className="rounded-lg" size="large" />
|
||||
</Form.Item>
|
||||
<UpstreamResourceField />
|
||||
<Form.Item
|
||||
label={
|
||||
<FieldLabel
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ import {
|
|||
getOAuthAuthorizationIdentity,
|
||||
CLEARED_ON_INVALIDATION,
|
||||
isHeldOAuthTokenStale,
|
||||
preservedAdminCredentials,
|
||||
preservedDeclaredAppCredentials,
|
||||
withoutMintedTokenCredentials,
|
||||
} from "@/components/mcp_tools/types";
|
||||
|
|
@ -207,7 +208,7 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
// registered client (useMcpOAuthFlow keys reuse off credentials.client_id) instead of re-DCRing;
|
||||
// the client-forwarded modes carry only the declared app.
|
||||
credentials: isClientForwardedTokenMode(values.auth_type)
|
||||
? preservedDeclaredAppCredentials(values.credentials)
|
||||
? preservedAdminCredentials(values.credentials)
|
||||
: { ...((values.credentials as Record<string, unknown> | undefined) ?? {}), ...(dcrClientRef.current ?? {}) },
|
||||
issuer: values.issuer,
|
||||
authorization_url: values.authorization_url,
|
||||
|
|
@ -251,7 +252,7 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
|
||||
const current = (form.getFieldValue("credentials") as Record<string, unknown> | undefined) ?? {};
|
||||
const nextCredentials = {
|
||||
...(preservedDeclaredAppCredentials(current) ?? {}),
|
||||
...(preservedAdminCredentials(current) ?? {}),
|
||||
...(current.scopes !== undefined && { scopes: current.scopes }),
|
||||
access_token: token.access_token,
|
||||
...(token.refresh_token && { refresh_token: token.refresh_token }),
|
||||
|
|
@ -288,10 +289,10 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
// Capture the admin-typed app before resetFields destroys it, then re-apply it: the app is
|
||||
// upstream-scoped config, not minted material, so it survives every invalidation (the token is
|
||||
// what gets discarded). Token-shaped keys are excluded by the helper's key filter.
|
||||
const keptAppCredentials = preservedDeclaredAppCredentials(form.getFieldValue("credentials"));
|
||||
const keptAdminCredentials = preservedAdminCredentials(form.getFieldValue("credentials"));
|
||||
form.resetFields([...CLEARED_ON_INVALIDATION]);
|
||||
if (keptAppCredentials) {
|
||||
form.setFieldsValue({ credentials: keptAppCredentials });
|
||||
if (keptAdminCredentials) {
|
||||
form.setFieldsValue({ credentials: keptAdminCredentials });
|
||||
}
|
||||
// Re-apply the in-flight edit last; rc-field-form deep-merges nested objects, so a changed
|
||||
// credentials sub-field composes with the preserved sibling instead of replacing the object.
|
||||
|
|
@ -568,7 +569,7 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
// Client-forwarded rows persist ONLY the declared app; strip any token material that lingered in
|
||||
// the form (e.g. from a prior oauth2 authorize on the same session) so it can never reach the row.
|
||||
const submitCredentials = isClientForwardedTokenMode(restValues.auth_type)
|
||||
? preservedDeclaredAppCredentials(credentialsPayload)
|
||||
? preservedAdminCredentials(credentialsPayload)
|
||||
: credentialsPayload;
|
||||
|
||||
if (includeCredentials && submitCredentials && Object.keys(submitCredentials).length > 0) {
|
||||
|
|
|
|||
|
|
@ -1052,6 +1052,62 @@ describe("MCPServerEdit (interactive OAuth)", () => {
|
|||
});
|
||||
});
|
||||
|
||||
describe("MCPServerEdit (resource indicator)", () => {
|
||||
const RESOURCE_PLACEHOLDER = "auto, or https://mcp.example.com/mcp";
|
||||
const serverWithResource = {
|
||||
...interactiveOAuthServer,
|
||||
credentials: { upstream_resource: "api://finance-api/.default" },
|
||||
};
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
mockOauth.tokenResponse = null;
|
||||
vi.mocked(networking.updateMCPServer).mockResolvedValue({ ...serverWithResource });
|
||||
});
|
||||
|
||||
async function renderAndSave() {
|
||||
render(
|
||||
<MCPServerEdit
|
||||
mcpServer={serverWithResource}
|
||||
accessToken="access-token"
|
||||
onCancel={vi.fn()}
|
||||
onSuccess={vi.fn()}
|
||||
availableAccessGroups={[]}
|
||||
/>,
|
||||
);
|
||||
const input = await screen.findByPlaceholderText(RESOURCE_PLACEHOLDER);
|
||||
await waitFor(() => expect(input).toHaveValue("api://finance-api/.default"));
|
||||
return async () => {
|
||||
await act(async () => {
|
||||
fireEvent.click(screen.getAllByRole("button", { name: "Save Changes" })[0]);
|
||||
});
|
||||
await waitFor(() => {
|
||||
expect(networking.updateMCPServer).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0];
|
||||
return payload;
|
||||
};
|
||||
}
|
||||
|
||||
// Regression: the edit form hand-rolled its own OAuth fields and never mounted this one, while the
|
||||
// submit path re-added every missing admin-config key as an explicit null. Saving any unrelated
|
||||
// change therefore wiped a configured resource indicator.
|
||||
it("leaves an untouched resource indicator alone instead of clearing it", async () => {
|
||||
const save = await renderAndSave();
|
||||
const payload = await save();
|
||||
expect(payload.credentials?.upstream_resource).toBe("api://finance-api/.default");
|
||||
});
|
||||
|
||||
it("sends an explicit null when the admin empties the field, so the backend merge clears it", async () => {
|
||||
const save = await renderAndSave();
|
||||
await act(async () => {
|
||||
fireEvent.change(screen.getByPlaceholderText(RESOURCE_PLACEHOLDER), { target: { value: "" } });
|
||||
});
|
||||
const payload = await save();
|
||||
expect(payload.credentials?.upstream_resource).toBeNull();
|
||||
});
|
||||
});
|
||||
|
||||
describe("MCPServerEdit (tool list fetch)", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
|
|
|
|||
|
|
@ -8,7 +8,9 @@ import {
|
|||
getOAuthAuthorizationIdentity,
|
||||
CLEARED_ON_INVALIDATION,
|
||||
isHeldOAuthTokenStale,
|
||||
preservedAdminCredentials,
|
||||
preservedDeclaredAppCredentials,
|
||||
ADMIN_CONFIG_CREDENTIAL_KEYS,
|
||||
withoutMintedTokenCredentials,
|
||||
OAUTH_FLOW,
|
||||
MCP_OAUTH2_FLOW_M2M,
|
||||
|
|
@ -34,9 +36,9 @@ import PassthroughAuthorizeSection from "./PassthroughAuthorizeSection";
|
|||
import MCPToolConfiguration from "./mcp_tool_configuration";
|
||||
import StdioConfiguration from "./StdioConfiguration";
|
||||
import TokenExchangeFormFields from "./TokenExchangeFormFields";
|
||||
import OAuthFormFields from "./OAuthFormFields";
|
||||
import MCPLogoSelector from "./MCPLogoSelector";
|
||||
import EnvVarsSection from "./EnvVarsSection";
|
||||
import TokenEndpointAuthMethodField from "./TokenEndpointAuthMethodField";
|
||||
import {
|
||||
validateMCPServerUrl,
|
||||
validateMCPServerName,
|
||||
|
|
@ -194,7 +196,7 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
transport,
|
||||
auth_type: isClientForwardedTokenMode(values.auth_type) ? values.auth_type : AUTH_TYPE.OAUTH2,
|
||||
credentials: isClientForwardedTokenMode(values.auth_type)
|
||||
? preservedDeclaredAppCredentials(values.credentials)
|
||||
? preservedAdminCredentials(values.credentials)
|
||||
: values.credentials,
|
||||
mcp_access_groups: values.mcp_access_groups || mcpServer.mcp_access_groups,
|
||||
static_headers: staticHeaders,
|
||||
|
|
@ -225,7 +227,7 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
|
||||
const current = (form.getFieldValue("credentials") as Record<string, unknown> | undefined) ?? {};
|
||||
const nextCredentials = {
|
||||
...(preservedDeclaredAppCredentials(current) ?? {}),
|
||||
...(preservedAdminCredentials(current) ?? {}),
|
||||
...(current.scopes !== undefined && { scopes: current.scopes }),
|
||||
access_token: token.access_token,
|
||||
...(token.refresh_token && { refresh_token: token.refresh_token }),
|
||||
|
|
@ -451,10 +453,10 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
resetOAuthFlow();
|
||||
// The admin-typed app is upstream-scoped config, not minted material, so it survives every
|
||||
// invalidation; only the held token is discarded. Token-shaped keys are excluded by the filter.
|
||||
const keptAppCredentials = preservedDeclaredAppCredentials(form.getFieldValue("credentials"));
|
||||
const keptAdminCredentials = preservedAdminCredentials(form.getFieldValue("credentials"));
|
||||
form.resetFields([...CLEARED_ON_INVALIDATION]);
|
||||
if (keptAppCredentials) {
|
||||
form.setFieldsValue({ credentials: keptAppCredentials });
|
||||
if (keptAdminCredentials) {
|
||||
form.setFieldsValue({ credentials: keptAdminCredentials });
|
||||
}
|
||||
const preserved = Object.fromEntries(
|
||||
CLEARED_ON_INVALIDATION.filter((key) => key in changedValues).map((key) => [key, changedValues[key]]),
|
||||
|
|
@ -718,6 +720,9 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
credentialValues && typeof credentialValues === "object"
|
||||
? Object.entries(credentialValues).reduce((acc: Record<string, any>, [key, value]) => {
|
||||
if (value === undefined || value === null || value === "") {
|
||||
if (value === "" && (ADMIN_CONFIG_CREDENTIAL_KEYS as readonly string[]).includes(key)) {
|
||||
acc[key] = null;
|
||||
}
|
||||
return acc;
|
||||
}
|
||||
if (key === "scopes") {
|
||||
|
|
@ -928,7 +933,7 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
// Client-forwarded rows persist ONLY the declared app; strip any token material lingering in the
|
||||
// form (e.g. from a prior oauth2 authorize this session) so it can never reach the row.
|
||||
const submitCredentials = isClientForwardedTokenMode(restValues.auth_type)
|
||||
? preservedDeclaredAppCredentials(credentialsPayload)
|
||||
? preservedAdminCredentials(credentialsPayload)
|
||||
: credentialsPayload;
|
||||
|
||||
if (includeCredentials && submitCredentials && Object.keys(submitCredentials).length > 0) {
|
||||
|
|
@ -1228,22 +1233,6 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
|
||||
{!isStdioTransport && isOAuthAuthType && (
|
||||
<>
|
||||
<Form.Item
|
||||
label={
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
OAuth Flow Type
|
||||
<Tooltip title="Machine-to-Machine (M2M) authenticates with client credentials and no user interaction. Interactive (PKCE) authorizes each user in the browser and stores per-user tokens. Servers created before this field existed have no stored value; choose one to persist it.">
|
||||
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="oauth_flow_type"
|
||||
>
|
||||
<Select placeholder="Select OAuth flow">
|
||||
<Select.Option value={OAUTH_FLOW.M2M}>Machine-to-Machine (M2M)</Select.Option>
|
||||
<Select.Option value={OAUTH_FLOW.INTERACTIVE}>Interactive (PKCE)</Select.Option>
|
||||
</Select>
|
||||
</Form.Item>
|
||||
{!oauthFlowTypeValue && !isDelegateAuth && (
|
||||
<Alert
|
||||
type="warning"
|
||||
|
|
@ -1253,192 +1242,16 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
description="Choose Machine-to-Machine (M2M) or Interactive (PKCE) so LiteLLM authenticates it the way you intend, then save. Until it is set, LiteLLM falls back to interactive per-user auth and treats a machine-to-machine credential shape conservatively."
|
||||
/>
|
||||
)}
|
||||
<Form.Item
|
||||
label={
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
OAuth Client ID (optional)
|
||||
<Tooltip title="Provide only if your MCP server cannot handle dynamic client registration.">
|
||||
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name={["credentials", "client_id"]}
|
||||
>
|
||||
<Input.Password
|
||||
placeholder="Enter OAuth client ID (leave blank to keep existing)"
|
||||
className="rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"
|
||||
/>
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label={
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
OAuth Client Secret (optional)
|
||||
<Tooltip title="Provide only if your MCP server cannot handle dynamic client registration.">
|
||||
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name={["credentials", "client_secret"]}
|
||||
>
|
||||
<Input.Password
|
||||
placeholder="Enter OAuth client secret (leave blank to keep existing)"
|
||||
className="rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"
|
||||
/>
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label={
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
OAuth Scopes (optional)
|
||||
<Tooltip title="Add scopes to override the default scope list used for this MCP server.">
|
||||
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name={["credentials", "scopes"]}
|
||||
>
|
||||
<Select
|
||||
mode="tags"
|
||||
tokenSeparators={[","]}
|
||||
placeholder="Add scopes"
|
||||
className="rounded-lg"
|
||||
size="large"
|
||||
/>
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label={
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
Issuer (optional)
|
||||
<Tooltip title="OAuth 2.0 authorization server issuer (RFC 8414). Auto-discovered on first connect; set it explicitly to pin the trust anchor so token and scope discovery is fetched from and validated against this issuer (RFC 8414 §3.3) instead of anything the resource advertises.">
|
||||
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="issuer"
|
||||
>
|
||||
<Input
|
||||
placeholder="https://issuer.example.com"
|
||||
className="rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"
|
||||
/>
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label={
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
Authorization URL Override (optional)
|
||||
<Tooltip title="Optional override for the authorization endpoint.">
|
||||
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="authorization_url"
|
||||
>
|
||||
<Input
|
||||
placeholder="https://example.com/oauth/authorize"
|
||||
className="rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"
|
||||
/>
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label={
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
Token URL Override (optional)
|
||||
<Tooltip title="Optional override for the token endpoint.">
|
||||
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="token_url"
|
||||
>
|
||||
<Input
|
||||
placeholder="https://example.com/oauth/token"
|
||||
className="rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"
|
||||
/>
|
||||
</Form.Item>
|
||||
<TokenEndpointAuthMethodField isEditing />
|
||||
<Form.Item
|
||||
label={
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
Registration URL Override (optional)
|
||||
<Tooltip title="Optional override for the dynamic client registration endpoint.">
|
||||
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="registration_url"
|
||||
>
|
||||
<Input
|
||||
placeholder="https://example.com/oauth/register"
|
||||
className="rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"
|
||||
/>
|
||||
</Form.Item>
|
||||
{!isM2MFlow && (
|
||||
<>
|
||||
<Form.Item
|
||||
label={
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
Token Validation Rules (optional)
|
||||
<Tooltip title='JSON object of key-value rules checked against the OAuth token response before storing. Supports dot-notation for nested fields (e.g. {"organization": "my-org", "team.id": "123"}). Tokens that fail validation are rejected with HTTP 403.'>
|
||||
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="token_validation_json"
|
||||
rules={[
|
||||
{
|
||||
validator: (_: any, value: string) => {
|
||||
if (!value || value.trim() === "") return Promise.resolve();
|
||||
try {
|
||||
JSON.parse(value);
|
||||
return Promise.resolve();
|
||||
} catch {
|
||||
return Promise.reject(new Error("Must be valid JSON"));
|
||||
}
|
||||
},
|
||||
},
|
||||
]}
|
||||
>
|
||||
<Input.TextArea
|
||||
placeholder={'{\n "organization": "my-org",\n "team.id": "123"\n}'}
|
||||
rows={4}
|
||||
className="font-mono text-sm rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"
|
||||
/>
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label={
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
Token Storage TTL (seconds, optional)
|
||||
<Tooltip title="How long to cache each user's OAuth access token in Redis before evicting it (never longer than the token's own expires_in). Leave blank to derive the TTL from the token's expires_in, or fall back to the 12-hour default.">
|
||||
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="token_storage_ttl_seconds"
|
||||
>
|
||||
<InputNumber min={1} placeholder="e.g. 3600" style={{ width: "100%" }} className="rounded-lg" />
|
||||
</Form.Item>
|
||||
</>
|
||||
)}
|
||||
<div className="rounded-lg border border-dashed border-gray-300 p-4 space-y-2">
|
||||
<p className="text-sm text-gray-600">
|
||||
Use OAuth to fetch a fresh access token and temporarily save it in the session as the authentication
|
||||
value.
|
||||
</p>
|
||||
<Button
|
||||
variant="secondary"
|
||||
onClick={startOAuthFlow}
|
||||
disabled={oauthStatus === "authorizing" || oauthStatus === "exchanging"}
|
||||
>
|
||||
{oauthStatus === "authorizing"
|
||||
? "Waiting for authorization..."
|
||||
: oauthStatus === "exchanging"
|
||||
? "Exchanging authorization code..."
|
||||
: "Authorize & Fetch Token"}
|
||||
</Button>
|
||||
{oauthError && <p className="text-sm text-red-500">{oauthError}</p>}
|
||||
{oauthStatus === "success" && oauthTokenResponse?.access_token && (
|
||||
<p className="text-sm text-green-600">
|
||||
Token fetched. Expires in {oauthTokenResponse.expires_in ?? "?"} seconds.
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
<OAuthFormFields
|
||||
isM2M={isM2MFlow}
|
||||
isEditing
|
||||
oauthFlow={{
|
||||
startOAuthFlow,
|
||||
status: oauthStatus,
|
||||
error: oauthError,
|
||||
tokenResponse: oauthTokenResponse,
|
||||
}}
|
||||
/>
|
||||
</>
|
||||
)}
|
||||
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import {
|
|||
gatewayMintsClientFor,
|
||||
getOAuthAuthorizationIdentity,
|
||||
isHeldOAuthTokenStale,
|
||||
preservedAdminCredentials,
|
||||
oauth2FlowToFormValue,
|
||||
preservedDeclaredAppCredentials,
|
||||
withoutMintedTokenCredentials,
|
||||
|
|
@ -34,6 +35,30 @@ describe("getOAuthAuthorizationIdentity", () => {
|
|||
expect(getOAuthAuthorizationIdentity(edited)).not.toBe(getOAuthAuthorizationIdentity(authorized));
|
||||
});
|
||||
|
||||
// Regression: upstream_resource is the RFC 8707 audience the upstream token is minted for, so
|
||||
// editing it strands a held token on the previous audience. It must invalidate here for the same
|
||||
// reason it belongs in the backend's mcp_oauth_token_identity, which this function mirrors.
|
||||
it("changes when the upstream_resource credential changes", () => {
|
||||
const authorized = {
|
||||
auth_type: AUTH_TYPE.OAUTH2,
|
||||
url: "https://a.example.com/mcp",
|
||||
credentials: { client_id: "cid", upstream_resource: "api://audience-one" },
|
||||
};
|
||||
const retargeted = {
|
||||
auth_type: AUTH_TYPE.OAUTH2,
|
||||
url: "https://a.example.com/mcp",
|
||||
credentials: { client_id: "cid", upstream_resource: "api://audience-two" },
|
||||
};
|
||||
const unset = {
|
||||
auth_type: AUTH_TYPE.OAUTH2,
|
||||
url: "https://a.example.com/mcp",
|
||||
credentials: { client_id: "cid" },
|
||||
};
|
||||
expect(getOAuthAuthorizationIdentity(retargeted)).not.toBe(getOAuthAuthorizationIdentity(authorized));
|
||||
expect(getOAuthAuthorizationIdentity(unset)).not.toBe(getOAuthAuthorizationIdentity(authorized));
|
||||
expect(isHeldOAuthTokenStale(retargeted, getOAuthAuthorizationIdentity(authorized))).toBe(true);
|
||||
});
|
||||
|
||||
it("is stable across non-mint fields", () => {
|
||||
const authorized = { auth_type: AUTH_TYPE.OAUTH2, url: "https://a.example.com/mcp", server_name: "one" };
|
||||
const renamed = { auth_type: AUTH_TYPE.OAUTH2, url: "https://a.example.com/mcp", server_name: "two" };
|
||||
|
|
@ -288,3 +313,37 @@ describe("isUnsupportedOnGatewayConnect", () => {
|
|||
expect(isUnsupportedOnGatewayConnect(undefined)).toBe(false);
|
||||
});
|
||||
});
|
||||
|
||||
describe("preservedAdminCredentials vs preservedDeclaredAppCredentials", () => {
|
||||
// Regression: upstream_resource is admin-typed config living in `credentials`, and the invalidation
|
||||
// reset wipes that whole object. If it is not preserved, editing an unrelated field like the URL
|
||||
// silently discards the admin's resource indicator and the server goes back to sending none.
|
||||
it("preserves upstream_resource across an invalidation reset", () => {
|
||||
const credentials = { client_id: "cid", client_secret: "csec", upstream_resource: "api://audience" };
|
||||
expect(preservedAdminCredentials(credentials)).toEqual(credentials);
|
||||
});
|
||||
|
||||
it("preserves upstream_resource even when no OAuth app is declared", () => {
|
||||
// A dynamic-client-registration server has no client_id/client_secret but can still pin a resource.
|
||||
expect(preservedAdminCredentials({ upstream_resource: "auto" })).toEqual({ upstream_resource: "auto" });
|
||||
});
|
||||
|
||||
it("strips minted token material", () => {
|
||||
const credentials = { client_id: "cid", upstream_resource: "auto", access_token: "tok", refresh_token: "r" };
|
||||
expect(preservedAdminCredentials(credentials)).toEqual({ client_id: "cid", upstream_resource: "auto" });
|
||||
});
|
||||
|
||||
// The two helpers answer different questions and must not be collapsed: "has the admin declared an
|
||||
// OAuth app" gates the app-may-not-match-upstream warning, so a resource-only server must read as
|
||||
// having no declared app.
|
||||
it("does not report a declared app for a resource-only server", () => {
|
||||
expect(preservedDeclaredAppCredentials({ upstream_resource: "auto" })).toBeUndefined();
|
||||
expect(preservedAdminCredentials({ upstream_resource: "auto" })).toBeDefined();
|
||||
});
|
||||
|
||||
it("still reports a declared app when client keys are present", () => {
|
||||
expect(preservedDeclaredAppCredentials({ client_id: "cid", upstream_resource: "auto" })).toEqual({
|
||||
client_id: "cid",
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -103,6 +103,7 @@ export const getOAuthAuthorizationIdentity = (values: Record<string, unknown>):
|
|||
client_id: credentials.client_id ?? null,
|
||||
client_secret: credentials.client_secret ?? null,
|
||||
scopes: credentials.scopes ?? null,
|
||||
upstream_resource: credentials.upstream_resource ?? null,
|
||||
issuer: values.issuer ?? null,
|
||||
authorization_url: values.authorization_url ?? null,
|
||||
token_url: values.token_url ?? null,
|
||||
|
|
@ -129,23 +130,46 @@ export const CLEARED_ON_INVALIDATION = ["credentials"] as const;
|
|||
// token-shaped keys so a preserve can never carry minted material through. Shared by both forms.
|
||||
const DECLARED_APP_CREDENTIAL_KEYS = ["client_id", "client_secret"] as const;
|
||||
|
||||
// Admin-typed credential config that is NOT part of the declared OAuth app. It is preserved across an
|
||||
// invalidation for the same reason the client keys are (nothing programmatic writes it, so a reset
|
||||
// would destroy admin input), but it must stay OUT of the declared-app set: whether an app exists is
|
||||
// a distinct question that gates the "app may not match upstream" warning, and a server using dynamic
|
||||
// client registration can set a resource indicator while having no app at all.
|
||||
export const ADMIN_CONFIG_CREDENTIAL_KEYS = ["upstream_resource"] as const;
|
||||
|
||||
// Minted token material the oauth2 authorize path writes beside the app keys; stripped from restored
|
||||
// snapshots and from any credentials that transit to the temp-session preview so a stale token never
|
||||
// reaches the backend or a client-forwarded server row.
|
||||
export const MINTED_TOKEN_CREDENTIAL_KEYS = ["access_token", "refresh_token", "expires_in", "scope"] as const;
|
||||
|
||||
export const preservedDeclaredAppCredentials = (
|
||||
const pickStringCredentials = (
|
||||
credentials: Record<string, unknown> | null | undefined,
|
||||
keys: readonly string[],
|
||||
): Record<string, string> | undefined => {
|
||||
if (!credentials) return undefined;
|
||||
const kept = Object.fromEntries(
|
||||
DECLARED_APP_CREDENTIAL_KEYS.filter((key) => typeof credentials[key] === "string" && credentials[key] !== "").map(
|
||||
(key) => [key, credentials[key] as string],
|
||||
),
|
||||
keys
|
||||
.filter((key) => typeof credentials[key] === "string" && credentials[key] !== "")
|
||||
.map((key) => [key, credentials[key] as string]),
|
||||
);
|
||||
return Object.keys(kept).length > 0 ? kept : undefined;
|
||||
};
|
||||
|
||||
// Does the admin have a declared OAuth client app? Answers only that question; use
|
||||
// preservedAdminCredentials for anything deciding what survives a reset or reaches the backend, or a
|
||||
// server that only carries admin config would read as having an app it never declared.
|
||||
export const preservedDeclaredAppCredentials = (
|
||||
credentials: Record<string, unknown> | null | undefined,
|
||||
): Record<string, string> | undefined => pickStringCredentials(credentials, DECLARED_APP_CREDENTIAL_KEYS);
|
||||
|
||||
// Everything the admin typed into `credentials` and nothing minted: the declared app plus the config
|
||||
// keys. This is what must survive the invalidation reset and what a client-forwarded row may persist,
|
||||
// so dropping a key from here silently discards admin input on an unrelated edit.
|
||||
export const preservedAdminCredentials = (
|
||||
credentials: Record<string, unknown> | null | undefined,
|
||||
): Record<string, string> | undefined =>
|
||||
pickStringCredentials(credentials, [...DECLARED_APP_CREDENTIAL_KEYS, ...ADMIN_CONFIG_CREDENTIAL_KEYS]);
|
||||
|
||||
// Drop minted token keys, keeping everything else (the declared app plus any non-token config).
|
||||
export const withoutMintedTokenCredentials = (
|
||||
credentials: Record<string, unknown> | null | undefined,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue