Merge branch 'litellm_internal_staging' into litellm_fix_stateful_statless_mcp

Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>
This commit is contained in:
mateo-berri 2026-05-13 21:30:27 +00:00
commit 38af29f2c9
No known key found for this signature in database
45 changed files with 2283 additions and 90 deletions

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "delegate_auth_to_upstream" BOOLEAN NOT NULL DEFAULT false;

View file

@ -323,6 +323,7 @@ model LiteLLM_MCPServerTable {
registration_url String?
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)
is_byok Boolean @default(false)
byok_description String[] @default([])
byok_api_key_help_url String?

View file

@ -41,7 +41,7 @@ def _get_redis_kwargs():
"retry",
}
include_args = [
include_args = {
"url",
"redis_connect_func",
"gcp_service_account",
@ -50,9 +50,9 @@ def _get_redis_kwargs():
"azure_client_id",
"azure_tenant_id",
"azure_client_secret",
]
}
available_args = [x for x in arg_spec.args if x not in exclude_args] + include_args
available_args = {x for x in arg_spec.args if x not in exclude_args} | include_args
return available_args
@ -84,23 +84,23 @@ def _get_redis_cluster_kwargs(client=None):
# Only allow primitive arguments
exclude_args = {"self", "connection_pool", "retry", "host", "port", "startup_nodes"}
available_args = [x for x in arg_spec.args if x not in exclude_args]
available_args.append("password")
available_args.append("username")
available_args.append("ssl")
available_args.append("ssl_cert_reqs")
available_args.append("ssl_check_hostname")
available_args.append("ssl_ca_certs")
available_args.append(
"redis_connect_func"
) # Needed for sync clusters and IAM detection
available_args.append("gcp_service_account")
available_args.append("gcp_ssl_ca_certs")
available_args.append("azure_redis_ad_token")
available_args.append("azure_client_id")
available_args.append("azure_tenant_id")
available_args.append("azure_client_secret")
available_args.append("max_connections")
available_args = {x for x in arg_spec.args if x not in exclude_args}
available_args |= {
"password",
"username",
"ssl",
"ssl_cert_reqs",
"ssl_check_hostname",
"ssl_ca_certs",
"redis_connect_func", # Needed for sync clusters and IAM detection
"gcp_service_account",
"gcp_ssl_ca_certs",
"azure_redis_ad_token",
"azure_client_id",
"azure_tenant_id",
"azure_client_secret",
"max_connections",
}
return available_args
@ -479,10 +479,24 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster:
return redis.RedisCluster(startup_nodes=new_startup_nodes, **cluster_kwargs) # type: ignore
def _get_redis_sentinel_connection_kwargs(redis_kwargs: dict) -> dict:
connection_kwargs = {}
args = _get_redis_kwargs()
for arg in redis_kwargs:
if arg in args:
connection_kwargs[arg] = redis_kwargs[arg]
return connection_kwargs
def _init_redis_sentinel(redis_kwargs) -> redis.Redis:
sentinel_nodes = redis_kwargs.get("sentinel_nodes")
sentinel_password = redis_kwargs.get("sentinel_password")
service_name = redis_kwargs.get("service_name")
connection_kwargs = _get_redis_sentinel_connection_kwargs(redis_kwargs)
connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT)
sentinel_kwargs = dict(connection_kwargs)
sentinel_kwargs["password"] = sentinel_password
if not sentinel_nodes or not service_name:
raise ValueError(
@ -494,19 +508,22 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis:
# Set up the Sentinel client
sentinel = redis.Sentinel(
sentinel_nodes,
socket_timeout=REDIS_SOCKET_TIMEOUT,
password=sentinel_password,
sentinel_kwargs=sentinel_kwargs,
)
# Return the master instance for the given service
return sentinel.master_for(service_name)
return sentinel.master_for(service_name, **connection_kwargs)
def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis:
sentinel_nodes = redis_kwargs.get("sentinel_nodes")
sentinel_password = redis_kwargs.get("sentinel_password")
service_name = redis_kwargs.get("service_name")
connection_kwargs = _get_redis_sentinel_connection_kwargs(redis_kwargs)
connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT)
sentinel_kwargs = dict(connection_kwargs)
sentinel_kwargs["password"] = sentinel_password
if not sentinel_nodes or not service_name:
raise ValueError(
@ -518,13 +535,12 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis:
# Set up the Sentinel client
sentinel = async_redis.Sentinel(
sentinel_nodes,
socket_timeout=REDIS_SOCKET_TIMEOUT,
password=sentinel_password,
sentinel_kwargs=sentinel_kwargs,
)
# Return the master instance for the given service
return sentinel.master_for(service_name)
return sentinel.master_for(service_name, **connection_kwargs)
def get_redis_client(**env_overrides):

View file

@ -178,6 +178,18 @@ class BudgetManager:
return list(self.user_dict.keys())
def reset_cost(self, user):
"""
Reset the tracked spend for a user back to zero.
Clears both the aggregate ``current_cost`` and the per-model
``model_cost`` breakdown stored for the given user.
Args:
user: The user identifier whose cost should be reset.
Returns:
dict: ``{"user": <updated user record>}`` reflecting the reset state.
"""
self.user_dict[user]["current_cost"] = 0
self.user_dict[user]["model_cost"] = {}
return {"user": self.user_dict[user]}

View file

@ -485,11 +485,16 @@ class MaskedHTTPStatusError(httpx.HTTPStatusError):
if k.lower() not in ("content-encoding", "content-length")
}
try:
request_content = original_error.request.content
except httpx.RequestNotRead:
request_content = b""
masked_request = httpx.Request(
method=original_error.request.method,
url=masked_url,
headers=original_error.request.headers,
content=original_error.request.content,
content=request_content,
)
super().__init__(

View file

@ -507,10 +507,10 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator):
# PROCESS REASONING CONTENT
reasoning_content: Optional[str] = None
content: Optional[str] = None
if chunk["message"].get("thinking") is not None:
if chunk["message"].get("thinking"):
reasoning_content = chunk["message"].get("thinking")
self.started_reasoning_content = True
elif chunk["message"].get("content") is not None:
if chunk["message"].get("content"):
if (
self.started_reasoning_content
and not self.finished_reasoning_content

View file

@ -108,7 +108,7 @@ class OllamaModelInfo(BaseLLMModelInfo):
continue
nm = entry.get("name") or entry.get("model")
if isinstance(nm, str):
names.add(nm)
names.add(nm if nm.startswith("ollama/") else f"ollama/{nm}")
except Exception as e:
verbose_logger.warning(f"Error retrieving ollama tag endpoint: {e}")
# If tags endpoint fails, fall back to static list

View file

@ -15551,14 +15551,17 @@
"uses_embed_content": true
},
"vertex_ai/gemini-embedding-2-preview": {
"input_cost_per_token": 1.5e-07,
"input_cost_per_audio_per_second": 0.00016,
"input_cost_per_image": 0.00012,
"input_cost_per_token": 2e-07,
"input_cost_per_video_per_second": 0.00079,
"litellm_provider": "vertex_ai",
"max_input_tokens": 8192,
"max_tokens": 8192,
"mode": "embedding",
"output_cost_per_token": 0,
"output_vector_size": 3072,
"source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal",
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supports_multimodal": true,
"uses_embed_content": true
},
@ -15573,7 +15576,7 @@
"mode": "embedding",
"output_cost_per_token": 0,
"output_vector_size": 3072,
"source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal",
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supports_multimodal": true,
"uses_embed_content": true
},

View file

@ -1,3 +1,4 @@
import re
from typing import Dict, List, Optional, Set, Tuple, cast
from fastapi import HTTPException
@ -122,6 +123,24 @@ class MCPRequestHandler:
# cannot be smuggled via query string, hostname, or a deeper URL segment.
if request.url.path.startswith("/.well-known/"):
validated_user_api_key_auth = UserAPIKeyAuth()
elif (
not litellm_api_key
and MCPRequestHandler._target_servers_delegate_auth_to_upstream( # noqa: E501
path=request.url.path, mcp_servers=mcp_servers
)
):
# Operator opted this oauth2 server into upstream-delegated auth
# (PKCE passthrough): skip LiteLLM API-key/SSO entirely so the
# client authenticates directly with the upstream MCP server.
# Fires ONLY when neither x-litellm-api-key nor Authorization is
# present. If any LiteLLM key is supplied (primary or secondary
# header), we fall through so user_id is resolved, spend/rate
# limiting apply, and any stored OAuth token can be retrieved
# and forwarded upstream. Gated by
# _target_servers_delegate_auth_to_upstream, which only returns
# True when EVERY target is auth_type=oauth2 AND has the
# delegate_auth_to_upstream flag set — fails closed otherwise.
validated_user_api_key_auth = UserAPIKeyAuth()
elif has_explicit_litellm_key:
# Explicit x-litellm-api-key provided - always validate normally
validated_user_api_key_auth = await user_api_key_auth(
@ -181,23 +200,62 @@ class MCPRequestHandler:
@staticmethod
def _extract_target_server_names_from_path(path: str) -> List[str]:
"""
Extract the target MCP server name from the standard MCP transport
URL patterns: ``/mcp/{server_name}[/...]`` and
Extract the target MCP server name(s) from the standard MCP transport
URL patterns: ``/mcp/{server_name_or_csv}[/...]`` and
``/{server_name}/mcp[/...]``. Returns ``[]`` for any other path so
callers fail closed when the target cannot be resolved.
Mirrors the regex-based parser in ``server.py::_get_mcp_servers_in_path``
so the names used for auth gating match the names used for downstream
filtering. Without this alignment, an attacker could craft
``/mcp/<delegated_server>/<garbage>`` so that auth treats the request
as targeting the delegate server (bypassing LiteLLM auth) while
downstream filtering sees a different (non-existent) target and falls
back to the caller's full allowed-server set.
REST/admin endpoints, OAuth2 server endpoints
(``/{server_name}/authorize``, ``/token`` etc.), and ``.well-known``
discovery routes intentionally fall through — those flows do not need
OAuth2 token passthrough. Clients aggregating multiple servers should
use ``x-mcp-servers``, which takes precedence over path parsing.
use ``x-mcp-servers`` on a path that does not encode a target.
"""
# ``/{server_name}/mcp[/...]`` form — single server. The literal
# ``mcp`` must be the second segment (not the first, which would be
# the ``/mcp/...`` form handled below). This branch must stay in sync
# with ``server.py::_get_mcp_servers_in_path``, which also accepts the
# un-rewritten form (some entry points may skip the
# ``dynamic_mcp_route`` rewrite).
segments = [s for s in path.split("/") if s]
if len(segments) >= 2 and segments[0] == "mcp":
return [segments[1]]
if len(segments) >= 2 and segments[1] == "mcp":
if len(segments) >= 2 and segments[1] == "mcp" and segments[0] != "mcp":
return [segments[0]]
return []
# ``/mcp/...`` form — server name(s) may contain a slash (e.g.
# ``custom_solutions/user_123``) and may be a comma-separated list.
# Use the same parsing logic as ``_get_mcp_servers_in_path`` so the
# parsed names match downstream routing.
mcp_path_match = re.match(r"^/mcp/([^?#]+)(?:\?.*)?(?:#.*)?$", path)
if not mcp_path_match:
return []
servers_and_path = mcp_path_match.group(1)
if not servers_and_path:
return []
if "," in servers_and_path:
# Comma-separated servers, possibly followed by a trailing path.
path_match = re.search(r"/([^/,]+(?:/[^/,]+)*)$", servers_and_path)
if path_match:
servers_part = servers_and_path[: -(len(path_match.group(1)) + 1)]
else:
servers_part = servers_and_path
return [s.strip() for s in servers_part.split(",") if s.strip()]
# Single-server case — server name may contain at most one slash.
single_server_match = re.match(
r"^([^/]+(?:/[^/]+)?)(?:/.*)?$", servers_and_path
)
if single_server_match:
return [single_server_match.group(1)]
return [servers_and_path]
@staticmethod
def _target_servers_use_oauth2(path: str, mcp_servers: Optional[List[str]]) -> bool:
@ -217,13 +275,13 @@ class MCPRequestHandler:
)
from litellm.types.mcp import MCPAuth
# Use the x-mcp-servers header verbatim when present (including the
# explicitly-empty list, which means "no targets" → fail closed).
# Only fall back to path parsing when the header was absent entirely.
target_names = (
mcp_servers
if mcp_servers is not None
else MCPRequestHandler._extract_target_server_names_from_path(path)
# Resolve the same target list downstream routing will use. For
# ``/mcp/...`` routes, ``extract_mcp_auth_context`` overrides the
# ``x-mcp-servers`` header with path-derived names, so we must mirror
# that here — otherwise a caller could set the header to a permissive
# server while the path targets a stricter one (header/path TOCTOU).
target_names = MCPRequestHandler._resolve_target_server_names(
path=path, mcp_servers_header=mcp_servers
)
if not target_names:
return False
@ -234,6 +292,78 @@ class MCPRequestHandler:
return False
return True
@staticmethod
def _target_servers_delegate_auth_to_upstream(
path: str, mcp_servers: Optional[List[str]]
) -> bool:
"""
True only when EVERY MCP server the request targets is configured for
``auth_type == oauth2`` AND has ``delegate_auth_to_upstream=True``.
Fails closed when any target does not opt in or cannot be resolved.
Used by :meth:`process_mcp_request` to skip LiteLLM API-key/SSO auth
entirely (PKCE passthrough) so the client authenticates directly with
the upstream MCP server. Mixed-target requests (e.g. one delegated +
one non-delegated server) fall back to normal LiteLLM auth.
"""
# Inline imports avoid a circular dependency: mcp_server_manager imports
# from this module.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.types.mcp import MCPAuth
# See _target_servers_use_oauth2: must mirror the downstream
# header-vs-path override or an attacker could set
# ``x-mcp-servers`` to a delegate-enabled server while the URL path
# targets a non-delegate server, skipping LiteLLM auth for it.
target_names = MCPRequestHandler._resolve_target_server_names(
path=path, mcp_servers_header=mcp_servers
)
if not target_names:
return False
for name in target_names:
server = global_mcp_server_manager.get_mcp_server_by_name(name)
if server is None or server.auth_type != MCPAuth.oauth2:
return False
# `is True` is intentional: opt-in must be an explicit boolean
# True. A MagicMock attribute (in tests) or any other truthy
# non-bool must not silently enable the bypass.
if getattr(server, "delegate_auth_to_upstream", False) is not True:
return False
if not getattr(server, "available_on_public_internet", True):
return False
# Never delegate for M2M (client_credentials) servers: LiteLLM
# fetches the upstream token automatically using stored credentials,
# so allowing anonymous bypass would let any external caller invoke
# tools authenticated as LiteLLM's service account.
if server.has_client_credentials:
return False
return True
@staticmethod
def _resolve_target_server_names(
path: str, mcp_servers_header: Optional[List[str]]
) -> List[str]:
"""
Resolve the target MCP server names exactly as downstream routing
does (``server.py::extract_mcp_auth_context``).
For ``/mcp/...`` paths, downstream routing **overrides** any
``x-mcp-servers`` header value with the path-derived names. Mirror
that here so an attacker cannot use a permissive header value to
flip an auth gate while the path targets a stricter server
(header/path TOCTOU). For non-``/mcp/...`` paths (where the path
does not encode targets), fall back to the header.
"""
path_targets = MCPRequestHandler._extract_target_server_names_from_path(path)
if path_targets:
return path_targets
# Path did not resolve to /mcp/... targets — trust the header
# (including an explicitly empty list, which means "no targets").
return mcp_servers_header if mcp_servers_header is not None else []
@staticmethod
def _get_mcp_auth_header_from_headers(headers: Headers) -> Optional[str]:
"""

View file

@ -402,6 +402,9 @@ class MCPServerManager:
available_on_public_internet=bool(
server_config.get("available_on_public_internet", True)
),
delegate_auth_to_upstream=bool(
server_config.get("delegate_auth_to_upstream", False)
),
# AWS SigV4 fields
aws_access_key_id=server_config.get("aws_access_key_id", None),
aws_secret_access_key=server_config.get("aws_secret_access_key", None),
@ -796,6 +799,9 @@ class MCPServerManager:
available_on_public_internet=bool(
getattr(mcp_server, "available_on_public_internet", True)
),
delegate_auth_to_upstream=bool(
getattr(mcp_server, "delegate_auth_to_upstream", False)
),
created_at=getattr(mcp_server, "created_at", None),
updated_at=getattr(mcp_server, "updated_at", None),
tool_name_to_display_name=_deserialize_json_dict(
@ -967,6 +973,34 @@ class MCPServerManager:
if not in_toolset_scope:
combined_servers.update(allow_all_server_ids)
# For anonymous callers (no user_id, no role), also surface any
# servers the operator has opted into upstream-delegated auth.
# These servers handle their own auth at the upstream level, so
# LiteLLM granting access here does not bypass any security gate.
is_anonymous = not (
user_api_key_auth
and (
getattr(user_api_key_auth, "user_id", None)
or getattr(user_api_key_auth, "user_role", None)
or getattr(user_api_key_auth, "api_key", None)
)
)
if is_anonymous:
delegate_server_ids = [
server.server_id
for server in self.get_registry().values()
if getattr(server, "auth_type", None) == MCPAuth.oauth2
and getattr(server, "delegate_auth_to_upstream", False) is True
# M2M servers must not be exposed anonymously: an
# unauthenticated caller would get LiteLLM to proxy tool
# calls using its stored client_credentials.
and not server.has_client_credentials
# Internal-only servers must not be reachable from public
# internet callers who happen to carry an upstream token.
and getattr(server, "available_on_public_internet", True)
]
combined_servers.update(delegate_server_ids)
if len(combined_servers) == 0:
verbose_logger.debug(
"No allowed MCP Servers found for user api key auth."

View file

@ -25,6 +25,7 @@ from typing import (
cast,
)
import httpx
from fastapi import FastAPI, HTTPException
from pydantic import AnyUrl, ConfigDict
from starlette.requests import Request as StarletteRequest
@ -53,13 +54,17 @@ from litellm.proxy._experimental.mcp_server.utils import (
get_server_prefix,
iter_known_server_prefixes,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
get_chain_id_from_headers,
)
from litellm.types.mcp import MCPAuth
from litellm.types.mcp import MCPAuth, MCPSpecVersion
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall
from litellm.utils import Rules, client, function_setup
@ -1423,8 +1428,24 @@ if MCP_AVAILABLE:
raw_headers=raw_headers,
)
# If no OAuth2 token came from request headers, fall back to pre-fetched creds
if extra_headers is None and server.auth_type == MCPAuth.oauth2:
# Prefer server-stored per-user OAuth when configured, so a stale
# Authorization header from the MCP client cannot override Redis/DB
# (same issue as call_tool in mcp_server_manager: VS Code caches tokens).
if (
server.auth_type == MCPAuth.oauth2
and getattr(server, "needs_user_oauth_token", False)
and user_api_key_auth is not None
):
db_headers = await _get_user_oauth_extra_headers_from_db(
server,
user_api_key_auth,
prefetched_creds=_prefetched_oauth_creds,
)
if db_headers:
extra_headers = db_headers
# If still no OAuth2 token, fall back to pre-fetched creds (non-stale-client path)
elif extra_headers is None and server.auth_type == MCPAuth.oauth2:
extra_headers = await _get_user_oauth_extra_headers_from_db(
server,
user_api_key_auth,
@ -2627,6 +2648,10 @@ if MCP_AVAILABLE:
import re
mcp_servers_from_path: Optional[List[str]] = None
segments = [s for s in path.split("/") if s]
if len(segments) >= 2 and segments[1] == "mcp" and segments[0] != "mcp":
return [segments[0]]
# Match /mcp/<servers_and_maybe_path>
# Where servers can be comma-separated list of server names
# Server names can contain slashes (e.g., "custom_solutions/user_123")
@ -2971,6 +2996,157 @@ if MCP_AVAILABLE:
)
return user_api_key_auth.model_copy(update={"object_permission": updated_op})
def _get_forwarded_auth_from_scope(scope: Scope) -> Optional[str]:
"""Return the upstream-bound ``Authorization`` header value, or None.
Only returns the ``Authorization`` header when ``x-litellm-api-key`` is
also present. In that case ``Authorization`` is unambiguously the
upstream token the caller wants forwarded to the MCP server. When
``x-litellm-api-key`` is absent the ``Authorization`` header may itself
be the LiteLLM proxy API key (backward-compat path in
``MCPRequestHandler.process_mcp_request``), and forwarding it upstream
would leak the proxy key to a third-party MCP server.
"""
authorization = None
has_litellm_key_header = False
for key, value in scope.get("headers", []):
key_lower = key.lower()
if key_lower == b"authorization":
authorization = value.decode("latin-1")
elif key_lower == b"x-litellm-api-key":
has_litellm_key_header = True
if not has_litellm_key_header:
return None
return authorization
async def _probe_upstream_auth(
url: str,
auth_header: str,
timeout: float = 5.0,
) -> tuple:
"""JSON-RPC initialize-probe the upstream URL to check whether the token is accepted.
Uses POST so StreamableHTTP MCP servers run the same auth path as a
real client request. Returns (status_code, www_authenticate).
Fails-open with (200, None) on network errors so a transient hiccup
does not block valid requests.
Uses the public ``AsyncHTTPHandler.post()`` interface and catches
``httpx.HTTPStatusError`` separately so the 401/403 we want to surface
is not swallowed by the broad fail-open ``except Exception`` below.
"""
client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.MCP,
params={"timeout": timeout},
)
probe_payload = {
"jsonrpc": "2.0",
"id": "litellm-mcp-auth-probe",
"method": "initialize",
"params": {
"protocolVersion": MCPSpecVersion.jun_2025.value,
"capabilities": {},
"clientInfo": {
"name": "litellm-mcp-auth-probe",
"version": "1.0.0",
},
},
}
probe_headers = {
"Authorization": auth_header,
"Accept": "application/json, text/event-stream",
}
try:
resp = await client.post(
url=url,
headers=probe_headers,
json=probe_payload,
timeout=timeout,
)
return resp.status_code, resp.headers.get("www-authenticate")
except httpx.HTTPStatusError as exc:
# AsyncHTTPHandler.post() calls raise_for_status(); a 401/403 from
# upstream lands here. Return its status so the caller can map it
# to the appropriate response.
return exc.response.status_code, exc.response.headers.get(
"www-authenticate"
)
except Exception as exc:
verbose_logger.debug(
f"_probe_upstream_auth: probe to {url} failed ({exc}), allowing request through"
)
return 200, None
async def _check_passthrough_upstream_auth(
scope: Scope,
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_servers: Optional[List[str]],
client_ip: Optional[str],
) -> None:
"""Probe pass-through upstream servers in parallel before the MCP session starts.
Only servers the caller's key is already authorized to reach are probed —
the list is derived from _get_allowed_mcp_servers so that a user cannot
trigger an upstream probe against a server their key is not permitted for.
The MCP SDK commits HTTP 200 headers before invoking handlers, so a 401
can only be returned before that point. This function raises HTTPException(401)
with a WWW-Authenticate header if any upstream rejects the client token.
Fails-open: network errors are logged and the request is allowed through.
"""
forwarded_auth = _get_forwarded_auth_from_scope(scope)
if not forwarded_auth:
return
# Use the authorized server set, not the raw user-supplied names, so that
# a caller cannot force a probe to a server their key is not allowed to use.
allowed_servers = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
client_ip=client_ip,
)
passthrough_servers = [
srv
for srv in allowed_servers
if srv.extra_headers
and any(h.lower() == "authorization" for h in srv.extra_headers)
# Exclude M2M servers: _prepare_mcp_server_headers skips caller
# Authorization when has_client_credentials is set, so probing
# those with the caller's token would send the wrong credential.
and not srv.has_client_credentials
]
if not passthrough_servers:
return
probe_results = await asyncio.gather(
*[
_probe_upstream_auth(srv.url or "", forwarded_auth)
for srv in passthrough_servers
]
)
request = StarletteRequest(scope)
base_url = get_request_base_url(request)
for srv, (probe_status, _) in zip(passthrough_servers, probe_results):
if probe_status == 401:
# Token is missing or expired — direct the client to re-authorize.
authorization_uri = (
f"Bearer authorization_uri="
f"{base_url}/.well-known/oauth-authorization-server/{srv.name}"
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"WWW-Authenticate": authorization_uri},
)
if probe_status == 403:
# Token is valid but the caller lacks permission — do not hint
# at re-authorization (RFC 9110: a fresh token with the same
# scopes would just hit 403 again and loop indefinitely).
raise HTTPException(
status_code=403,
detail="Forbidden",
)
async def handle_streamable_http_mcp( # noqa: PLR0915
scope: Scope, receive: Receive, send: Send
) -> None:
@ -3044,6 +3220,13 @@ if MCP_AVAILABLE:
user_api_key_auth, active_toolset_id
)
# Pre-flight auth check for pass-through servers. Must run after
# toolset scoping so the probe list is derived from the fully-authorized
# server set, not the raw user-supplied names.
await _check_passthrough_upstream_auth(
scope, user_api_key_auth, mcp_servers, _client_ip
)
# Inject masked debug headers when client sends x-litellm-mcp-debug: true
_debug_headers = MCPDebug.maybe_build_debug_headers(
raw_headers=raw_headers,

View file

@ -1274,6 +1274,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None
allow_all_keys: bool = False
available_on_public_internet: bool = True
delegate_auth_to_upstream: bool = False
is_byok: bool = False
byok_description: List[str] = Field(default_factory=list)
byok_api_key_help_url: Optional[str] = None
@ -1356,6 +1357,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
registration_url: Optional[str] = None
allow_all_keys: bool = False
available_on_public_internet: bool = True
delegate_auth_to_upstream: bool = False
is_byok: bool = False
byok_description: List[str] = Field(default_factory=list)
byok_api_key_help_url: Optional[str] = None
@ -1427,6 +1429,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
registration_url: Optional[str] = None
allow_all_keys: bool = False
available_on_public_internet: bool = True
delegate_auth_to_upstream: bool = False
is_byok: bool = False
byok_description: List[str] = Field(default_factory=list)
byok_api_key_help_url: Optional[str] = None

View file

@ -36,6 +36,31 @@ ADMIN_ONLY_HEALTH_DISPLAY_PARAMS = ("api_base", "api_version")
MINIMAL_DISPLAY_PARAMS = ["model", "mode_error"]
# Modes whose health-check probe is a chat-style completion call and
# therefore accept `max_tokens`. Other modes (embedding, image_generation,
# audio_*, rerank, video_generation, ocr, search, moderation, ...) hit
# endpoints that reject unknown fields with 400 "Unknown parameter:
# 'max_tokens'". Allow-list so new modes are safe by default.
# Per-deployment override: `model_info.health_check_supports_max_tokens`.
_MAX_TOKEN_SUPPORT_MODES: frozenset = frozenset({"chat", "completion", "responses"})
def _should_inject_health_check_max_tokens(model_info: dict) -> bool:
"""
Whether the health-check probe should include `max_tokens`.
Order:
1. `model_info.health_check_supports_max_tokens` (operator override).
2. `_MAX_TOKEN_SUPPORT_MODES`. Missing `mode` is treated as `chat`
for backward compatibility.
"""
explicit = model_info.get("health_check_supports_max_tokens")
if explicit is not None:
return bool(explicit)
mode = model_info.get("mode") or "chat"
return mode in _MAX_TOKEN_SUPPORT_MODES
# Health-check modes that forward `reasoning_effort` to the provider (chat-style calls).
_HEALTH_CHECK_MODES_SUPPORTING_REASONING_EFFORT = frozenset(
(None, "chat", "completion")
@ -389,14 +414,22 @@ def _update_litellm_params_for_health_check(
Update the litellm params for health check.
- gets a short `messages` param for health check
- adds a bounded `max_tokens` when the deployment is a chat-style mode
(`chat`, `completion`, `responses`) or the operator explicitly opts in
via `model_info.health_check_supports_max_tokens`. Non-chat endpoints
(image, embedding, audio_*, rerank, video, ocr, search, moderation, ...)
reject unknown fields with 400 "Unknown parameter: 'max_tokens'".
- updates the `model` param with the `health_check_model` if it exists Doc: https://docs.litellm.ai/docs/proxy/health#wildcard-routes
- updates the `voice` param with the `health_check_voice` for `audio_speech` mode if it exists Doc: https://docs.litellm.ai/docs/proxy/health#text-to-speech-models
- for Bedrock models with region routing (bedrock/region/model), strips the litellm routing prefix but preserves the model ID
"""
litellm_params["messages"] = _get_random_llm_message()
_resolved_max_tokens = _resolve_health_check_max_tokens(model_info, litellm_params)
if _resolved_max_tokens is not None:
litellm_params["max_tokens"] = _resolved_max_tokens
if _should_inject_health_check_max_tokens(model_info):
_resolved_max_tokens = _resolve_health_check_max_tokens(
model_info, litellm_params
)
if _resolved_max_tokens is not None:
litellm_params["max_tokens"] = _resolved_max_tokens
# Per-model reasoning effort for health checks only (e.g. reasoning_effort=none).
if model_info.get("mode", None) in _HEALTH_CHECK_MODES_SUPPORTING_REASONING_EFFORT:

View file

@ -1548,15 +1548,21 @@ def _allow_public_health_readiness_details() -> bool:
return general_settings.get("allow_public_health_readiness_details") is True
async def _set_public_readiness_status(response: Response) -> None:
async def _resolve_public_readiness_db(response: Response) -> str:
"""
Return the db status string for the public probe and flip the response to
503 when a configured DB is unreachable. Mirrors the legacy values:
"Not connected" (no DB configured), "connected", "disconnected".
"""
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return
return "Not connected"
db_health_status = await _db_health_readiness_check()
if db_health_status["status"] != "connected":
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
return db_health_status["status"]
@router.get(
@ -1565,15 +1571,17 @@ async def _set_public_readiness_status(response: Response) -> None:
)
async def health_readiness(response: Response):
"""
Public readiness probe. Keep this low-detail for unauthenticated load
balancers by default. Admins can opt into the legacy detailed public
payload with general_settings.allow_public_health_readiness_details.
Public readiness probe. Returns a low-detail payload safe to expose to
unauthenticated load balancers — `status` plus `db` so orchestrators and
external probes can distinguish "healthy" from "DB unreachable" without a
credential. Admins can opt into the legacy detailed payload with
general_settings.allow_public_health_readiness_details.
"""
if _allow_public_health_readiness_details():
return await _get_health_readiness_details(response=response)
await _set_public_readiness_status(response=response)
return {"status": "healthy"}
db_status = await _resolve_public_readiness_db(response=response)
return {"status": "healthy", "db": db_status}
@router.get(

View file

@ -163,7 +163,7 @@ 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 MCPCredentials
from litellm.types.mcp import MCPAuth, MCPCredentials
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@dataclass
@ -1551,6 +1551,57 @@ if MCP_AVAILABLE:
except _jwt.InvalidTokenError:
pass
# For delegate_auth_to_upstream servers the entire PKCE handshake
# (both /authorize browser redirect and /token authorization_code
# exchange) must work without a LiteLLM session. /authorize is opened
# in a VS Code webview that may have no cookie; /token is a programmatic
# POST from VS Code. PKCE security (code_verifier) guarantees the
# authorization_code exchange cannot be replayed, so anonymous access
# is safe for that grant only.
#
# Importantly, NOT safe for refresh_token grants: ``mcp_token`` will
# forward the request to the upstream issuer with LiteLLM's stored
# ``client_secret`` attached, so any caller holding a refresh token
# issued to this client could mint fresh upstream access tokens through
# us. Require normal LiteLLM auth for those.
if not api_key:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
server_id = request.path_params.get("server_id", "")
if server_id:
_s = global_mcp_server_manager.get_mcp_server_by_id(server_id)
if not _s:
_s = global_mcp_server_manager.get_mcp_server_by_name(server_id)
if (
_s
and getattr(_s, "auth_type", None) == MCPAuth.oauth2
and getattr(_s, "delegate_auth_to_upstream", False) is True
and getattr(_s, "available_on_public_internet", True)
# M2M servers fetch tokens with stored credentials; never
# expose their /authorize or /token endpoints anonymously.
and not _s.has_client_credentials
):
# For /token, require PKCE authorization_code; refresh_token
# grants must NOT bypass auth (see comment above).
path_lower = (request.url.path or "").rstrip("/").lower()
if path_lower.endswith("/token"):
body_data = await _read_request_body(request=request)
grant_type = (body_data or {}).get("grant_type", "")
if grant_type != "authorization_code":
# Fall through to normal LiteLLM auth (will 401 if
# no key supplied).
pass
else:
return UserAPIKeyAuth()
else:
# /authorize and other PKCE-flow GETs are safe to
# bypass: PKCE binds the upstream issuer's ``code``
# to the original ``code_challenge`` so no anonymous
# token can be minted via the redirect alone.
return UserAPIKeyAuth()
request_data = await _read_request_body(request=request)
request_data = populate_request_with_path_params(
request_data=request_data, request=request

View file

@ -15026,8 +15026,17 @@ async def _stream_mcp_asgi_response(
# If the handler task dies (exception or cancellation) without sending the EOF
# sentinel, body_iter() would block forever on body_queue.get(). The callback
# below guarantees the queue gets unblocked regardless of how the task ends.
# When this happens before response headers, propagate the original exception
# instead of waiting for the header timeout.
def _ensure_eof(task: asyncio.Task) -> None:
if task.cancelled() or task.exception() is not None:
if task.cancelled():
body_queue.put_nowait(None)
return
task_exception = task.exception()
if task_exception is not None:
if not headers_ready.done():
headers_ready.set_exception(task_exception)
body_queue.put_nowait(None)
handler_task.add_done_callback(_ensure_eof)

View file

@ -323,6 +323,7 @@ model LiteLLM_MCPServerTable {
registration_url String?
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)
is_byok Boolean @default(false)
byok_description String[] @default([])
byok_api_key_help_url String?

View file

@ -1238,6 +1238,8 @@ class LiteLLMCompletionResponsesConfig:
file_dict["file_data"] = item["file_data"]
new_item: Dict[str, Any] = {"type": "file", "file": file_dict}
if "cache_control" in item:
new_item["cache_control"] = item["cache_control"]
return new_item
@staticmethod
@ -1282,26 +1284,28 @@ class LiteLLMCompletionResponsesConfig:
)
)
elif item.get("type") == "input_image":
content_list.append(
dict(
LiteLLMCompletionResponsesConfig._transform_input_image_item_to_image_item(
item
)
image_block = dict(
LiteLLMCompletionResponsesConfig._transform_input_image_item_to_image_item(
item
)
)
if "cache_control" in item:
image_block["cache_control"] = item["cache_control"]
content_list.append(image_block)
else:
# Skip text blocks with None text to avoid downstream errors
text_value = item.get("text")
if text_value is None:
continue
content_list.append(
{
"type": LiteLLMCompletionResponsesConfig._get_chat_completion_request_content_type(
item.get("type") or "text"
),
"text": text_value,
}
)
content_block: Dict[str, Any] = {
"type": LiteLLMCompletionResponsesConfig._get_chat_completion_request_content_type(
item.get("type") or "text"
),
"text": text_value,
}
if "cache_control" in item:
content_block["cache_control"] = item["cache_control"]
content_list.append(content_block)
return content_list
else:
raise ValueError(f"Invalid content type: {type(content)}")

View file

@ -90,6 +90,13 @@ def timeout(timeout_duration: float = 0.0, exception_to_raise=Timeout):
class _LoopWrapper(Thread):
"""Daemon thread that owns a dedicated asyncio event loop.
Used by the sync branch of :func:`timeout` to run a coroutine on a
background event loop so the calling thread can wait on it with a
timeout via :func:`asyncio.run_coroutine_threadsafe`.
"""
def __init__(self):
super().__init__(daemon=True)
self.loop = asyncio.new_event_loop()

View file

@ -68,6 +68,12 @@ class MCPServer(BaseModel):
access_groups: Optional[List[str]] = None
allow_all_keys: bool = False
available_on_public_internet: bool = True
# When True AND auth_type == oauth2, MCP requests targeting this server
# bypass LiteLLM API-key/SSO auth (and the pre-emptive 401) so the client
# completes PKCE directly with the upstream MCP server. Honored only for
# auth_type=oauth2; ignored for any other auth_type. See
# MCPRequestHandler._target_servers_delegate_auth_to_upstream.
delegate_auth_to_upstream: bool = False
is_byok: bool = False
byok_description: List[str] = []
byok_api_key_help_url: Optional[str] = None

View file

@ -15556,14 +15556,17 @@
"uses_embed_content": true
},
"vertex_ai/gemini-embedding-2-preview": {
"input_cost_per_token": 1.5e-07,
"input_cost_per_audio_per_second": 0.00016,
"input_cost_per_image": 0.00012,
"input_cost_per_token": 2e-07,
"input_cost_per_video_per_second": 0.00079,
"litellm_provider": "vertex_ai",
"max_input_tokens": 8192,
"max_tokens": 8192,
"mode": "embedding",
"output_cost_per_token": 0,
"output_vector_size": 3072,
"source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal",
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supports_multimodal": true,
"uses_embed_content": true
},
@ -15578,7 +15581,7 @@
"mode": "embedding",
"output_cost_per_token": 0,
"output_vector_size": 3072,
"source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal",
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supports_multimodal": true,
"uses_embed_content": true
},

View file

@ -323,6 +323,7 @@ model LiteLLM_MCPServerTable {
registration_url String?
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)
is_byok Boolean @default(false)
byok_description String[] @default([])
byok_api_key_help_url String?

View file

@ -104,6 +104,31 @@ class TestMaskedHTTPStatusError:
# The attached request must be the masked one, not the original.
assert "KEY_X" not in str(req.url)
def test_handles_streaming_request_content(self):
"""MaskedHTTPStatusError must not crash when request body is streamed."""
streaming_request = httpx.Request(
"POST",
"https://api.openai.com/v1/images/edits?key=SECRET_KEY",
stream=httpx.ByteStream(b"multipart-data"),
)
response = httpx.Response(
400,
request=streaming_request,
content=b'{"error": "bad request"}',
)
orig = httpx.HTTPStatusError(
message="400 Bad Request",
request=streaming_request,
response=response,
)
masked = MaskedHTTPStatusError(orig)
assert masked.status_code == 400
assert masked.response.status_code == 400
assert masked.response.request is not None
assert "SECRET_KEY" not in str(masked.request.url)
def test_strips_content_encoding_to_avoid_double_decode(self):
"""If the upstream response declared Content-Encoding (e.g. gzip),
the rebuilt Response must not carry that header over — otherwise httpx

View file

@ -694,6 +694,71 @@ class TestOllamaReasoningContentStreaming:
# reasoning_content is not set when there's no thinking in the chunk
assert getattr(result2.choices[0].delta, "reasoning_content", None) is None
def test_thinking_and_content_in_same_chunk(self):
"""
Test that a chunk containing both thinking and content preserves both fields.
"""
iterator = OllamaChatCompletionResponseIterator(
streaming_response=iter([]),
sync_stream=True,
)
chunk = {
"model": "deepseek-r1",
"message": {
"role": "assistant",
"thinking": "Let me reason first.",
"content": "Final answer.",
},
"done": False,
}
result = iterator.chunk_parser(chunk)
assert result.choices[0].delta.reasoning_content == "Let me reason first."
assert result.choices[0].delta.content == "Final answer."
def test_streaming_chunks_ignore_inactive_empty_reasoning_fields(self):
"""
Test that Ollama chunks with inactive empty fields stay in the active delta.
"""
iterator = OllamaChatCompletionResponseIterator(
streaming_response=iter([]),
sync_stream=True,
)
chunk = {
"model": "deepseek-r1",
"message": {
"role": "assistant",
"thinking": "Let me reason first.",
"content": "",
},
"done": False,
}
result = iterator.chunk_parser(chunk)
assert result.choices[0].delta.reasoning_content == "Let me reason first."
assert result.choices[0].delta.content is None
assert iterator.finished_reasoning_content is False
content_chunk = {
"model": "deepseek-r1",
"message": {
"role": "assistant",
"thinking": "",
"content": "Final answer.",
},
"done": False,
}
result = iterator.chunk_parser(content_chunk)
assert getattr(result.choices[0].delta, "reasoning_content", None) is None
assert result.choices[0].delta.content == "Final answer."
assert iterator.finished_reasoning_content is True
def test_think_tags_in_content(self):
"""
Test that <think> tags embedded in content are properly parsed.

View file

@ -73,7 +73,7 @@ class TestOllamaModelInfo:
info = OllamaModelInfo()
models = info.get_models()
# Only 'alpha' and 'zeta' should be returned, sorted alphabetically
assert models == ["alpha", "zeta"]
assert models == ["ollama/alpha", "ollama/zeta"]
# Ensure correct endpoint was called
assert calls and calls[0].endswith("/api/tags")
assert call_headers and call_headers[0] == {}
@ -122,7 +122,7 @@ class TestOllamaModelInfo:
monkeypatch.setattr(httpx, "get", mock_get)
info = OllamaModelInfo()
models = info.get_models()
assert models == ["m1", "m2"]
assert models == ["ollama/m1", "ollama/m2"]
def test_get_models_fallback_on_error(self, monkeypatch):
"""
@ -139,6 +139,32 @@ class TestOllamaModelInfo:
# Default static ollama_models is ['llama2'], so expect ['ollama/llama2']
assert models == ["ollama/llama2"]
def test_get_models_no_double_prefix(self, monkeypatch):
"""
Names that already carry the 'ollama/' prefix (or are returned by an
Ollama server that's been configured to emit them) should not be
prefixed a second time.
"""
sample = {
"models": [
{"name": "ollama/already-prefixed"},
{"name": "fresh"},
{"name": "hf.co/Qwen/Qwen3-14B:latest"},
]
}
def mock_get(url, headers):
return DummyResponse(sample, status_code=200)
monkeypatch.setattr(httpx, "get", mock_get)
info = OllamaModelInfo()
models = info.get_models()
assert models == [
"ollama/already-prefixed",
"ollama/fresh",
"ollama/hf.co/Qwen/Qwen3-14B:latest",
]
class TestOllamaGetModelInfo:
"""Tests for OllamaConfig.get_model_info() api_base threading and graceful fallback."""

View file

@ -1075,6 +1075,727 @@ class TestMCPOAuth2FallbackTargetGating:
await MCPRequestHandler.process_mcp_request(scope)
@pytest.mark.asyncio
class TestMCPDelegateAuthToUpstream:
"""
Tests for the ``delegate_auth_to_upstream`` per-server flag.
When set on an ``auth_type=oauth2`` MCP server, LiteLLM must skip its own
API-key/SSO check entirely so the client completes PKCE directly with the
upstream MCP server. The gate must fail closed for any non-oauth2 server,
any mixed-target request, and any request where the target cannot be
resolved.
"""
@staticmethod
def _make_server(auth_type, delegate_auth_to_upstream=False):
from litellm.types.mcp_server.mcp_server_manager import MCPServer
return MCPServer(
server_id="test-server-id",
name="test-server",
transport="http",
auth_type=auth_type,
delegate_auth_to_upstream=delegate_auth_to_upstream,
)
async def test_delegate_skips_litellm_auth_with_no_authorization(self):
"""
oauth2 + delegate_auth_to_upstream=True, no Authorization header at
all → anonymous UserAPIKeyAuth and ``user_api_key_auth`` is never
called.
"""
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/delegated_oauth_server",
"headers": [],
}
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
) as mock_auth,
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = (
TestMCPDelegateAuthToUpstream._make_server(
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
)
)
(auth_result, *_rest) = await MCPRequestHandler.process_mcp_request(scope)
assert isinstance(auth_result, UserAPIKeyAuth)
mock_auth.assert_not_called()
async def test_delegate_with_upstream_token_in_authorization_falls_back_to_anonymous(
self,
):
"""
oauth2 + delegate_auth_to_upstream=True with an upstream OAuth token in
``Authorization`` (not a LiteLLM key): LiteLLM auth is attempted first
(and fails), then the existing oauth2 fallback returns anonymous so the
bearer is forwarded upstream untouched. The delegate branch itself does
not fire when Authorization is present — that is what protects spend
tracking for callers using Authorization-style LiteLLM keys.
"""
from fastapi import HTTPException
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/delegated_oauth_server",
"headers": [(b"authorization", b"Bearer upstream-pkce-token")],
}
async def mock_user_api_key_auth_fails(api_key, request):
raise HTTPException(status_code=401, detail="Invalid API key")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = (
TestMCPDelegateAuthToUpstream._make_server(
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
)
)
(
auth_result,
_,
_,
_,
oauth2_headers,
_,
) = await MCPRequestHandler.process_mcp_request(scope)
assert isinstance(auth_result, UserAPIKeyAuth)
assert oauth2_headers.get("Authorization") == "Bearer upstream-pkce-token"
async def test_delegate_off_still_requires_litellm_auth(self):
"""
oauth2 server but delegate flag is OFF → existing behaviour: a missing
/ invalid LiteLLM key still 401s (no anonymous fast-path).
"""
from fastapi import HTTPException
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/non_delegated_oauth_server",
"headers": [],
}
async def mock_user_api_key_auth_fails(api_key, request):
raise HTTPException(status_code=401, detail="Invalid API key")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = (
TestMCPDelegateAuthToUpstream._make_server(
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=False,
)
)
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
async def test_delegate_ignored_for_non_oauth2_server(self):
"""
Defense in depth: even if an operator turns on delegate_auth_to_upstream
for a non-oauth2 server (api_key, bearer_token, etc.), the gate must
not fire — only oauth2 servers may delegate.
"""
from fastapi import HTTPException
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/api_key_server",
"headers": [],
}
async def mock_user_api_key_auth_fails(api_key, request):
raise HTTPException(status_code=401, detail="Invalid API key")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = (
TestMCPDelegateAuthToUpstream._make_server(
auth_type=MCPAuth.api_key,
delegate_auth_to_upstream=True,
)
)
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
async def test_delegate_mixed_targets_fail_closed(self):
"""
x-mcp-servers can list multiple targets. If ANY of them does not opt in
to delegate_auth_to_upstream, the bypass must NOT fire — otherwise an
attacker could mix one delegated server in to skip auth on the others.
"""
from fastapi import HTTPException
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
"method": "POST",
"path": "/mcp",
"headers": [
(b"x-mcp-servers", b"delegated_oauth,plain_oauth"),
],
}
async def mock_user_api_key_auth_fails(api_key, request):
raise HTTPException(status_code=401, detail="Invalid API key")
def mock_lookup(name, client_ip=None):
if name == "delegated_oauth":
return TestMCPDelegateAuthToUpstream._make_server(
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
)
return TestMCPDelegateAuthToUpstream._make_server(
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=False,
)
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.side_effect = mock_lookup
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
async def test_delegate_no_resolvable_target_fail_closed(self):
"""
If the target server cannot be resolved at all (e.g. admin/REST path
that isn't ``/mcp/{name}`` or ``/{name}/mcp``), we cannot prove the
gate's preconditions, so we must fail closed and run normal auth.
"""
from fastapi import HTTPException
scope = {
"type": "http",
"method": "GET",
"path": "/admin/whatever",
"headers": [],
}
async def mock_user_api_key_auth_fails(api_key, request):
raise HTTPException(status_code=401, detail="Invalid API key")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = None
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
async def test_explicit_litellm_key_takes_precedence_over_delegate(self):
"""
When ``x-litellm-api-key`` is present, normal auth runs even for a
delegate server, so ``user_id`` is resolved and any stored upstream
OAuth credentials can be looked up and forwarded. The bypass only
fires when no LiteLLM key is supplied.
"""
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/delegated_oauth_server",
"headers": [(b"x-litellm-api-key", b"Bearer sk-1234")],
}
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
new_callable=AsyncMock,
return_value=UserAPIKeyAuth(user_id="real-user"),
) as mock_auth,
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = (
TestMCPDelegateAuthToUpstream._make_server(
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
)
)
(auth_result, *_rest) = await MCPRequestHandler.process_mcp_request(scope)
assert isinstance(auth_result, UserAPIKeyAuth)
assert auth_result.user_id == "real-user"
mock_auth.assert_called_once()
async def test_litellm_key_via_authorization_header_not_bypassed(self):
"""
Regression: a LiteLLM key sent via the secondary ``Authorization`` header
(e.g. ``Authorization: Bearer sk-...``) must still trigger normal auth
and not be silently swallowed by the delegate bypass — otherwise spend
tracking and rate limiting are skipped for those callers.
"""
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/delegated_oauth_server",
"headers": [(b"authorization", b"Bearer sk-1234")],
}
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
new_callable=AsyncMock,
return_value=UserAPIKeyAuth(user_id="real-user"),
) as mock_auth,
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = (
TestMCPDelegateAuthToUpstream._make_server(
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
)
)
(auth_result, *_rest) = await MCPRequestHandler.process_mcp_request(scope)
assert isinstance(auth_result, UserAPIKeyAuth)
assert auth_result.user_id == "real-user"
mock_auth.assert_called_once()
async def test_delegate_ignored_for_client_credentials_server(self):
"""
oauth2 + delegate_auth_to_upstream=True but oauth2_flow=client_credentials
→ bypass must NOT fire; normal LiteLLM auth must be attempted.
M2M servers fetch the upstream token automatically using stored
credentials, so allowing anonymous bypass would let any external
caller invoke tools as LiteLLM's service account.
"""
from fastapi import HTTPException
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/m2m_server",
"headers": [],
}
m2m_server = MCPServer(
server_id="m2m-server-id",
name="m2m_server",
transport="http",
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
oauth2_flow="client_credentials",
)
async def mock_auth_raises(*_args, **_kwargs):
raise HTTPException(status_code=401, detail="No key provided")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_auth_raises,
) as mock_auth,
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = m2m_server
# No delegate bypass → normal auth is attempted → 401 raised
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
mock_auth.assert_called_once()
async def test_delegate_ignored_for_non_public_server(self):
"""
Internal-only delegate servers must not bypass LiteLLM auth for
anonymous public callers.
"""
from fastapi import HTTPException
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/internal_server",
"headers": [],
}
internal_server = MCPServer(
server_id="internal-server-id",
name="internal_server",
transport="http",
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
available_on_public_internet=False,
)
async def mock_auth_raises(*_args, **_kwargs):
raise HTTPException(status_code=401, detail="No key provided")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_auth_raises,
) as mock_auth,
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = internal_server
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
mock_auth.assert_called_once()
async def test_get_allowed_servers_excludes_client_credentials_delegate(self):
"""
get_allowed_mcp_servers must not surface M2M (client_credentials) delegate
servers to anonymous callers even if delegate_auth_to_upstream=True.
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager = MCPServerManager()
pkce_server = MCPServer(
server_id="pkce-server",
name="pkce_server",
transport="http",
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
available_on_public_internet=True,
)
m2m_server = MCPServer(
server_id="m2m-server",
name="m2m_server",
transport="http",
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
oauth2_flow="client_credentials",
available_on_public_internet=True,
)
manager.registry = {
pkce_server.server_id: pkce_server,
m2m_server.server_id: m2m_server,
}
with patch.object(
MCPRequestHandler,
"get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[],
):
result = await manager.get_allowed_mcp_servers(None)
assert "pkce-server" in result
assert "m2m-server" not in result
async def test_get_allowed_servers_excludes_non_public_delegate(self):
"""
Internal-only (available_on_public_internet=False) delegate servers
must not appear in the anonymous allow-list.
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager = MCPServerManager()
public_server = MCPServer(
server_id="public-server",
name="public_server",
transport="http",
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
available_on_public_internet=True,
)
internal_server = MCPServer(
server_id="internal-server",
name="internal_server",
transport="http",
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
available_on_public_internet=False,
)
manager.registry = {
public_server.server_id: public_server,
internal_server.server_id: internal_server,
}
with patch.object(
MCPRequestHandler,
"get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[],
):
result = await manager.get_allowed_mcp_servers(None)
assert "public-server" in result
assert "internal-server" not in result
def test_extract_target_server_names_matches_routing_parser(self):
"""
Regression: _extract_target_server_names_from_path must match the
downstream regex parser in server.py::_get_mcp_servers_in_path.
Previously, a request to ``/mcp/<delegated>/garbage`` was parsed as
targeting ``<delegated>`` by the auth gate (bypassing LiteLLM auth)
while the routing layer parsed it as ``<delegated>/garbage`` — when
that name did not resolve, the request fell back to the anonymous
allow-list which can include ``allow_all_keys`` servers that normally
require a LiteLLM key.
"""
from litellm.proxy._experimental.mcp_server.server import (
_get_mcp_servers_in_path,
)
cases = [
# Single server, single segment.
("/mcp/foo", ["foo"]),
# Server name with one embedded slash (two segments).
("/mcp/foo/bar", ["foo/bar"]),
# Server name with embedded slash + extra path → name stays at two segments.
("/mcp/foo/bar/tools", ["foo/bar"]),
# Comma-separated servers, no trailing path.
("/mcp/foo,bar", ["foo", "bar"]),
# Comma-separated servers with trailing path.
("/mcp/foo,bar/tools", ["foo", "bar"]),
# Alternative form ``/<server>/mcp`` is also parsed (both auth
# parser and routing parser handle it for defense-in-depth — some
# entry points may not be rewritten by ``dynamic_mcp_route``).
("/foo/mcp", ["foo"]),
("/foo/mcp/tools", ["foo"]),
# Non-MCP paths → empty (fail closed).
("/.well-known/oauth-authorization-server", []),
("/v1/keys", []),
("/", []),
]
for path_input, expected in cases:
assert (
MCPRequestHandler._extract_target_server_names_from_path(path_input)
== expected
), f"path={path_input!r} → expected {expected!r}"
assert (
_get_mcp_servers_in_path(path_input) or []
) == expected, f"path={path_input!r} → routing expected {expected!r}"
async def test_delegate_does_not_bypass_on_extra_path_segment(self):
"""
Regression: ``/mcp/<delegated>/<garbage>`` must NOT bypass auth.
The bypass key check is now performed against the same parsed target
as downstream routing — ``<delegated>/<garbage>`` — which will not
resolve to a delegate-enabled server, so normal LiteLLM auth runs.
"""
from fastapi import HTTPException
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/delegated_server/extra",
"headers": [],
}
delegate_server = TestMCPDelegateAuthToUpstream._make_server(
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
)
def lookup_by_name(name):
# Only the *exact* delegated name resolves. Anything else (e.g.
# ``delegated_server/extra``) returns None so the bypass fails.
if name == "delegated_server":
return delegate_server
return None
async def mock_auth_raises(*_args, **_kwargs):
raise HTTPException(status_code=401, detail="No key provided")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_auth_raises,
) as mock_auth,
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.side_effect = lookup_by_name
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
# Auth was attempted (not bypassed) because the parsed target
# name does not match any registered delegate server.
mock_auth.assert_called_once()
async def test_delegate_ignores_x_mcp_servers_header_for_mcp_paths(self):
"""
Regression (header/path TOCTOU): For ``/mcp/...`` routes, downstream
routing overrides ``x-mcp-servers`` with the path-derived names.
The auth bypass must do the same — otherwise an attacker could send
``x-mcp-servers: <delegated>`` while the URL path targets a
non-delegate server, flipping the auth gate on a server that should
require a LiteLLM key.
"""
from fastapi import HTTPException
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/non_delegate_server",
"headers": [(b"x-mcp-servers", b"delegated_server")],
}
delegate_server = MCPServer(
server_id="delegate-id",
name="delegated_server",
transport="http",
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
available_on_public_internet=True,
)
non_delegate = MCPServer(
server_id="non-delegate-id",
name="non_delegate_server",
transport="http",
auth_type=MCPAuth.api_key,
)
def lookup_by_name(name):
return {
"delegated_server": delegate_server,
"non_delegate_server": non_delegate,
}.get(name)
async def mock_auth_raises(*_args, **_kwargs):
raise HTTPException(status_code=401, detail="No key provided")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_auth_raises,
) as mock_auth,
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.side_effect = lookup_by_name
# Bypass MUST NOT fire — path-derived target is the non-delegate
# server. Normal auth runs and 401s.
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
mock_auth.assert_called_once()
async def test_resolve_target_server_names_prefers_path_over_header(self):
"""
``_resolve_target_server_names`` must:
- For ``/mcp/<name>`` paths, return the path-derived list and ignore
the header (mirrors downstream routing).
- For non-MCP paths, fall back to the header (including the explicit
empty-list case, which fails closed).
"""
# Path matches /mcp/... — header is ignored.
assert MCPRequestHandler._resolve_target_server_names(
path="/mcp/foo", mcp_servers_header=["evil"]
) == ["foo"]
assert MCPRequestHandler._resolve_target_server_names(
path="/mcp/foo,bar", mcp_servers_header=["evil"]
) == ["foo", "bar"]
assert MCPRequestHandler._resolve_target_server_names(
path="/foo/mcp", mcp_servers_header=["evil"]
) == ["foo"]
# Path does not match — header is trusted.
assert MCPRequestHandler._resolve_target_server_names(
path="/.well-known/oauth-authorization-server",
mcp_servers_header=["foo"],
) == ["foo"]
# Explicit empty list on a non-MCP path → empty (fail closed).
assert (
MCPRequestHandler._resolve_target_server_names(
path="/.well-known/oauth-authorization-server",
mcp_servers_header=[],
)
== []
)
# No header on a non-MCP path → empty.
assert (
MCPRequestHandler._resolve_target_server_names(
path="/.well-known/oauth-authorization-server",
mcp_servers_header=None,
)
== []
)
class TestMCPCustomHeaderName:
"""Test suite for custom MCP authentication header name functionality"""

View file

@ -4466,3 +4466,137 @@ async def test_call_tool_empty_extra_headers_returns_none():
), "P2 API consistency issue: expected None for empty extra_headers, got: " + str(
captured_extra_headers
)
# ---------------------------------------------------------------------------
# Pre-flight upstream auth check tests
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_probe_upstream_auth_returns_upstream_status():
"""_probe_upstream_auth forwards the status code from the upstream server."""
from litellm.proxy._experimental.mcp_server.server import _probe_upstream_auth
mock_response = MagicMock()
mock_response.status_code = 401
mock_response.headers = {"www-authenticate": 'Bearer realm="test"'}
mock_client = MagicMock()
mock_client.post = AsyncMock(return_value=mock_response)
with patch(
"litellm.proxy._experimental.mcp_server.server.get_async_httpx_client",
return_value=mock_client,
):
status, www_auth = await _probe_upstream_auth(
"http://upstream/mcp", "Bearer some-token"
)
assert status == 401
assert www_auth == 'Bearer realm="test"'
mock_client.post.assert_awaited_once()
_, kwargs = mock_client.post.call_args
assert kwargs["headers"]["Authorization"] == "Bearer some-token"
assert kwargs["json"]["method"] == "initialize"
@pytest.mark.asyncio
async def test_probe_upstream_auth_surfaces_httpx_status_error():
"""Probe extracts status + WWW-Authenticate from httpx.HTTPStatusError.
AsyncHTTPHandler.post() calls raise_for_status() internally, so when the
upstream returns 401/403 the call raises httpx.HTTPStatusError rather than
returning the response. The probe must catch that specifically (before the
fail-open `except Exception`) so the auth check is not silently defeated.
"""
import httpx
from litellm.proxy._experimental.mcp_server.server import _probe_upstream_auth
mock_response = MagicMock()
mock_response.status_code = 401
mock_response.headers = {"www-authenticate": 'Bearer realm="test"'}
request = httpx.Request("POST", "http://upstream/mcp")
error = httpx.HTTPStatusError(
message="401 Unauthorized", request=request, response=mock_response
)
mock_client = MagicMock()
mock_client.post = AsyncMock(side_effect=error)
with patch(
"litellm.proxy._experimental.mcp_server.server.get_async_httpx_client",
return_value=mock_client,
):
status, www_auth = await _probe_upstream_auth(
"http://upstream/mcp", "Bearer some-token"
)
assert status == 401
assert www_auth == 'Bearer realm="test"'
@pytest.mark.asyncio
async def test_probe_upstream_auth_fails_open_on_network_error():
"""_probe_upstream_auth returns (200, None) when the network call fails."""
from litellm.proxy._experimental.mcp_server.server import _probe_upstream_auth
mock_client = MagicMock()
mock_client.post = AsyncMock(side_effect=Exception("connection refused"))
with patch(
"litellm.proxy._experimental.mcp_server.server.get_async_httpx_client",
return_value=mock_client,
):
status, www_auth = await _probe_upstream_auth(
"http://upstream/mcp", "Bearer some-token"
)
assert status == 200
assert www_auth is None
def test_get_forwarded_auth_from_scope_extracts_header():
"""Returns Authorization value when x-litellm-api-key is also present."""
from litellm.proxy._experimental.mcp_server.server import (
_get_forwarded_auth_from_scope,
)
scope = {
"headers": [
(b"content-type", b"application/json"),
(b"x-litellm-api-key", b"sk-litellm-proxy-key"),
(b"authorization", b"Bearer my-token"),
]
}
assert _get_forwarded_auth_from_scope(scope) == "Bearer my-token"
def test_get_forwarded_auth_from_scope_returns_none_when_missing():
from litellm.proxy._experimental.mcp_server.server import (
_get_forwarded_auth_from_scope,
)
assert _get_forwarded_auth_from_scope({"headers": []}) is None
def test_get_forwarded_auth_from_scope_skips_when_no_litellm_key_header():
"""Skip when ``x-litellm-api-key`` is absent.
Without ``x-litellm-api-key``, the ``Authorization`` header may itself be
the LiteLLM proxy API key (backward-compat). Forwarding it upstream would
leak the proxy key, so the helper must return None and the probe must
not fire.
"""
from litellm.proxy._experimental.mcp_server.server import (
_get_forwarded_auth_from_scope,
)
scope = {
"headers": [
(b"content-type", b"application/json"),
(b"authorization", b"Bearer ambiguous-token"),
]
}
assert _get_forwarded_auth_from_scope(scope) is None

View file

@ -2456,6 +2456,51 @@ class TestMCPServerManager:
assert "test_server_1" in result
assert "test_server_2" in result
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_anonymous_delegate_requires_oauth2(self):
"""Anonymous delegated auth listing should only include oauth2 servers."""
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
manager = MCPServerManager()
oauth_delegate_server = MCPServer(
server_id="oauth-delegate",
name="oauth_delegate",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
)
api_key_delegate_server = MCPServer(
server_id="api-key-delegate",
name="api_key_delegate",
transport=MCPTransport.http,
auth_type=MCPAuth.api_key,
delegate_auth_to_upstream=True,
)
oauth_non_delegate_server = MCPServer(
server_id="oauth-non-delegate",
name="oauth_non_delegate",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=False,
)
manager.registry = {
oauth_delegate_server.server_id: oauth_delegate_server,
api_key_delegate_server.server_id: api_key_delegate_server,
oauth_non_delegate_server.server_id: oauth_non_delegate_server,
}
with patch.object(
MCPRequestHandler,
"get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[],
):
result = await manager.get_allowed_mcp_servers(None)
assert set(result) == {"oauth-delegate"}
def test_get_mcp_server_from_tool_name_uses_server_name_not_name(self):
"""
Test that _get_mcp_server_from_tool_name uses server.server_name instead of server.name

View file

@ -736,3 +736,82 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
assert mock_get_stored_token.await_count == 1
assert mock_handle_request.await_count == 1
@pytest.mark.asyncio
async def test_handle_streamable_http_mcp_emits_401_for_delegated_server_without_token():
"""
OAuth2 server with ``delegate_auth_to_upstream=True`` and no Authorization
header must still emit a pre-emptive 401 with WWW-Authenticate so the
client kicks off PKCE. The 401 points at LiteLLM's discovery shim, which
in turn delegates to the upstream OAuth issuer.
"""
from fastapi import HTTPException
try:
from litellm.proxy._experimental.mcp_server.server import (
handle_streamable_http_mcp,
session_manager,
)
except ImportError:
pytest.skip("MCP server not available")
scope = {
"type": "http",
"method": "POST",
"path": "/mcp",
"headers": [
(b"content-type", b"application/json"),
(b"host", b"litellm.example.com"),
],
}
receive = AsyncMock()
send = AsyncMock()
user_auth = MagicMock()
user_auth.user_id = None
delegated_server = MagicMock()
delegated_server.auth_type = MCPAuth.oauth2
delegated_server.delegate_auth_to_upstream = True
delegated_server.needs_user_oauth_token = True
with (
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
return_value=(
user_auth,
None,
["delegated_oauth_server"],
None,
None,
None,
),
),
patch(
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
),
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
),
patch(
"litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session",
new_callable=AsyncMock,
return_value=False,
),
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name",
return_value=delegated_server,
),
patch.object(
session_manager,
"handle_request",
new_callable=AsyncMock,
) as mock_handle_request,
):
with pytest.raises(HTTPException) as exc_info:
await handle_streamable_http_mcp(scope, receive, send)
assert exc_info.value.status_code == 401
assert "www-authenticate" in exc_info.value.headers
assert mock_handle_request.await_count == 0

View file

@ -844,9 +844,14 @@ def test_health_readiness(proxy_client):
duration_ms < 500
), f"Health check took {duration_ms:.2f}ms, expected < 500ms for readiness endpoint"
# Assert response contains only low-detail public probe fields
# Assert response contains only low-detail public probe fields. `db` is
# included so unauthenticated probes can distinguish "DB unreachable"
# from a fully-healthy worker; its value depends on whether the test env
# exposes DATABASE_URL.
response_data = response.json()
assert response_data == {"status": "healthy"}
assert set(response_data.keys()) == {"status", "db"}
assert response_data["status"] == "healthy"
assert response_data["db"] in {"connected", "disconnected", "Not connected"}
print(f"Response time: {duration_ms:.2f}ms")
@ -1750,7 +1755,7 @@ async def test_health_readiness_returns_503_when_db_disconnected():
result = await health_readiness(response=response)
assert response.status_code == 503
assert result == {"status": "healthy"}
assert result == {"status": "healthy", "db": "disconnected"}
@pytest.mark.asyncio
@ -1773,7 +1778,7 @@ async def test_health_readiness_returns_200_when_db_connected():
result = await health_readiness(response=response)
assert response.status_code == 200
assert result == {"status": "healthy"}
assert result == {"status": "healthy", "db": "connected"}
@pytest.mark.asyncio
@ -1792,7 +1797,7 @@ async def test_health_readiness_returns_200_when_no_db_configured():
result = await health_readiness(response=response)
assert response.status_code == 200
assert result == {"status": "healthy"}
assert result == {"status": "healthy", "db": "Not connected"}
def test_clean_endpoint_data_strips_credentials_keeps_routing_fields():

View file

@ -1645,6 +1645,108 @@ class TestTemporaryMCPSessionEndpoints:
_, call_kwargs = auth_builder_mock.call_args
assert call_kwargs["api_key"] == "Bearer sk-header-key"
@pytest.mark.asyncio
async def test_mcp_oauth_user_api_key_auth_requires_oauth2_for_delegate_bypass(
self,
):
"""Non-oauth2 servers must not get anonymous access from the delegate flag."""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_mcp_oauth_user_api_key_auth,
)
expected_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
mock_request = MagicMock()
mock_request.headers = {}
mock_request.cookies = {}
mock_request.path_params = {"server_id": "server-1"}
non_oauth_server = MagicMock()
non_oauth_server.auth_type = MCPAuth.api_key
non_oauth_server.delegate_auth_to_upstream = True
mock_manager = MagicMock()
mock_manager.get_mcp_server_by_id.return_value = non_oauth_server
mock_manager.get_mcp_server_by_name.return_value = None
fake_proxy_server = types.SimpleNamespace(master_key=None)
with (
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
mock_manager,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_api_key_auth_builder",
AsyncMock(return_value=expected_auth),
) as auth_builder_mock,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body",
AsyncMock(return_value={}),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.populate_request_with_path_params",
side_effect=lambda request_data, request: request_data,
),
):
result = await _mcp_oauth_user_api_key_auth(mock_request)
assert result is expected_auth
auth_builder_mock.assert_awaited_once()
_, call_kwargs = auth_builder_mock.call_args
assert call_kwargs["api_key"] == ""
@pytest.mark.asyncio
async def test_mcp_oauth_user_api_key_auth_requires_public_server_for_delegate_bypass(
self,
):
"""Internal-only delegate servers must still require LiteLLM auth."""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_mcp_oauth_user_api_key_auth,
)
expected_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
mock_request = MagicMock()
mock_request.headers = {}
mock_request.cookies = {}
mock_request.path_params = {"server_id": "server-1"}
internal_server = MagicMock()
internal_server.auth_type = MCPAuth.oauth2
internal_server.delegate_auth_to_upstream = True
internal_server.available_on_public_internet = False
internal_server.has_client_credentials = False
mock_manager = MagicMock()
mock_manager.get_mcp_server_by_id.return_value = internal_server
mock_manager.get_mcp_server_by_name.return_value = None
fake_proxy_server = types.SimpleNamespace(master_key=None)
with (
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
mock_manager,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_api_key_auth_builder",
AsyncMock(return_value=expected_auth),
) as auth_builder_mock,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body",
AsyncMock(return_value={}),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.populate_request_with_path_params",
side_effect=lambda request_data, request: request_data,
),
):
result = await _mcp_oauth_user_api_key_auth(mock_request)
assert result is expected_auth
auth_builder_mock.assert_awaited_once()
_, call_kwargs = auth_builder_mock.call_args
assert call_kwargs["api_key"] == ""
@pytest.mark.asyncio
async def test_mcp_authorize_proxies_to_discoverable_endpoint(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (

View file

@ -227,6 +227,132 @@ def test_wildcard_ignores_reasoning_split_model_info(monkeypatch):
assert _resolve_health_check_max_tokens(model_info, litellm_params) is None
# ---------------------------------------------------------------------------
# image_generation must not receive max_tokens.
#
# _update_litellm_params_for_health_check injected `max_tokens` for every
# deployment. For `mode: image_generation` that leaked into OpenAI
# `/v1/images/generations`, which strictly rejects unknown fields with
# `400 "Unknown parameter: 'max_tokens'"`, marking dall-e-* and
# gpt-image-1 as permanently unhealthy even though their actual image
# calls succeed. `messages` still gets injected (downstream
# `_filter_model_params` already strips it for non-chat handlers).
# ---------------------------------------------------------------------------
def test_image_generation_mode_skips_max_tokens():
"""image_generation must not receive max_tokens."""
model_info = {"mode": "image_generation"}
litellm_params = {"model": "openai/dall-e-3", "api_key": "sk-test"}
updated = _update_litellm_params_for_health_check(model_info, litellm_params)
assert "max_tokens" not in updated
# connection-level params must still pass through unchanged
assert updated["api_key"] == "sk-test"
def test_health_check_max_tokens_value_is_ignored_for_non_chat_modes():
"""A configured `health_check_max_tokens` *value* (the int that controls
how many tokens to inject) is still skipped when the mode is outside the
allow-list — the inject decision runs before value resolution, so the
value never reaches `_resolve_health_check_max_tokens`. Note this is
distinct from `health_check_supports_max_tokens` (the bool that toggles
injection on/off per deployment)."""
model_info = {"mode": "image_generation", "health_check_max_tokens": 50}
litellm_params = {"model": "openai/dall-e-3"}
updated = _update_litellm_params_for_health_check(model_info, litellm_params)
assert "max_tokens" not in updated
def test_chat_mode_still_injects_max_tokens():
"""Regression guard: the chat-style probe payload is unchanged."""
model_info = {"mode": "chat"}
litellm_params = {"model": "gpt-4"}
updated = _update_litellm_params_for_health_check(model_info, litellm_params)
assert updated["max_tokens"] == 5
def test_no_mode_still_injects_max_tokens():
"""Regression guard: model_info without `mode` keeps the legacy path."""
model_info: dict = {}
litellm_params = {"model": "gpt-4"}
updated = _update_litellm_params_for_health_check(model_info, litellm_params)
assert updated["max_tokens"] == 5
# ---------------------------------------------------------------------------
# Allow-list behavior: only chat-style modes (chat / completion / responses)
# receive max_tokens. Every other mode is skipped by default.
#
# Per-deployment override via `health_check_supports_max_tokens` lets the
# operator force injection on (e.g. a non-listed but max_tokens-capable
# endpoint where they want to bound probe token usage) or off (e.g. a
# chat-style provider with a strict schema that rejects unknown fields).
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("mode", ["chat", "completion", "responses"])
def test_chat_style_modes_inject_max_tokens(mode):
updated = _update_litellm_params_for_health_check(
{"mode": mode}, {"model": f"openai/dummy-{mode}"}
)
assert updated["max_tokens"] == 5
@pytest.mark.parametrize(
"mode",
[
"embedding",
"image_generation",
"image_edit",
"audio_speech",
"audio_transcription",
"rerank",
"video_generation",
"ocr",
"search",
"moderation",
],
)
def test_non_chat_modes_skip_max_tokens(mode):
updated = _update_litellm_params_for_health_check(
{"mode": mode}, {"model": f"openai/dummy-{mode}"}
)
assert "max_tokens" not in updated
def test_explicit_override_true_forces_injection_outside_allowlist():
"""Operator opts a non-listed deployment in to bound probe token usage."""
model_info = {
"mode": "image_generation",
"health_check_supports_max_tokens": True,
}
litellm_params = {"model": "openai/some-future-image-model"}
updated = _update_litellm_params_for_health_check(model_info, litellm_params)
assert updated["max_tokens"] == 5
def test_explicit_override_false_suppresses_injection_inside_allowlist():
"""Operator opts a chat-style deployment out (strict-schema provider)."""
model_info = {"mode": "chat", "health_check_supports_max_tokens": False}
litellm_params = {"model": "openai/strict-schema-chat"}
updated = _update_litellm_params_for_health_check(model_info, litellm_params)
assert "max_tokens" not in updated
def test_update_litellm_params_health_check_reasoning_effort():
"""model_info.health_check_reasoning_effort sets reasoning_effort for chat-style health checks."""
model_info = {"health_check_reasoning_effort": "low"}

View file

@ -0,0 +1,36 @@
import asyncio
import pytest
from fastapi import HTTPException
from litellm.proxy.proxy_server import _stream_mcp_asgi_response
@pytest.mark.asyncio
async def test_stream_mcp_asgi_response_propagates_pre_header_http_exception():
async def handle_fn(_scope, _receive, _send):
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"WWW-Authenticate": "Bearer authorization_uri=https://example.test/auth"
},
)
async def receive():
return {"type": "http.request", "body": b"", "more_body": False}
with pytest.raises(HTTPException) as exc_info:
await asyncio.wait_for(
_stream_mcp_asgi_response(
handle_fn,
{"type": "http", "method": "POST", "path": "/mcp", "headers": []},
receive,
),
timeout=1.0,
)
assert exc_info.value.status_code == 401
assert exc_info.value.headers == {
"WWW-Authenticate": "Bearer authorization_uri=https://example.test/auth"
}

View file

@ -2170,3 +2170,86 @@ class TestEnsureOutputItemContentPartAdded:
events = iterator._pending_response_events
assert len(events) == 2
class TestCacheControlPreservation:
def test_cache_control_preserved_in_content_transformation(self):
"""cache_control injected by AnthropicCacheControlHook must survive
the Responses API -> Chat Completion content transformation."""
content = [
{
"type": "text",
"text": "hello",
"cache_control": {"type": "ephemeral"},
}
]
result = LiteLLMCompletionResponsesConfig._transform_responses_api_content_to_chat_completion_content(
content
)
assert isinstance(result, list)
assert len(result) == 1
assert result[0]["cache_control"] == {"type": "ephemeral"}
def test_content_without_cache_control_unaffected(self):
"""Content blocks that don't have cache_control should be unaffected."""
content = [{"type": "text", "text": "hello"}]
result = LiteLLMCompletionResponsesConfig._transform_responses_api_content_to_chat_completion_content(
content
)
assert isinstance(result, list)
assert len(result) == 1
assert "cache_control" not in result[0]
def test_cache_control_preserved_in_input_item_transformation(self):
"""cache_control survives the full input-item -> messages transformation."""
input_item = {
"role": "user",
"content": [
{
"type": "text",
"text": "long context",
"cache_control": {"type": "ephemeral"},
}
],
}
messages = LiteLLMCompletionResponsesConfig._transform_responses_api_input_item_to_chat_completion_message(
input_item
)
assert len(messages) == 1
msg_content = (
messages[0].get("content")
if isinstance(messages[0], dict)
else getattr(messages[0], "content", None)
)
assert isinstance(msg_content, list)
assert msg_content[0]["cache_control"] == {"type": "ephemeral"}
def test_cache_control_preserved_for_input_file_block(self):
content = [
{
"type": "input_file",
"file_id": "file-abc123",
"cache_control": {"type": "ephemeral"},
}
]
result = LiteLLMCompletionResponsesConfig._transform_responses_api_content_to_chat_completion_content(
content
)
assert isinstance(result, list)
assert len(result) == 1
assert result[0]["cache_control"] == {"type": "ephemeral"}
def test_cache_control_preserved_for_input_image_block(self):
content = [
{
"type": "input_image",
"image_url": "https://example.com/img.png",
"cache_control": {"type": "ephemeral"},
}
]
result = LiteLLMCompletionResponsesConfig._transform_responses_api_content_to_chat_completion_content(
content
)
assert isinstance(result, list)
assert len(result) == 1
assert result[0]["cache_control"] == {"type": "ephemeral"}

View file

@ -426,6 +426,120 @@ def test_sync_client_prefers_cluster_over_url_via_env_var(
assert len(call_kwargs["startup_nodes"]) == 1
@patch("litellm._redis.redis.Sentinel")
def test_sync_sentinel_uses_sentinel_password_and_master_password(mock_sentinel_cls):
"""Sentinel auth must be passed to the sentinel, not the Redis master client."""
mock_sentinel = MagicMock()
mock_sentinel_cls.return_value = mock_sentinel
get_redis_client(
sentinel_nodes=[("sentinel-1", 26379)],
sentinel_password="sentinel-secret",
service_name="mymaster",
password="redis-secret",
username="redis-user",
ssl=True,
ssl_cert_reqs="required",
ssl_check_hostname=True,
ssl_ca_certs="/tmp/test-ca.pem",
max_connections=17,
socket_timeout=5,
)
mock_sentinel_cls.assert_called_once()
sentinel_call_kwargs = mock_sentinel_cls.call_args[1]
assert "password" not in sentinel_call_kwargs
assert "username" not in sentinel_call_kwargs
assert "ssl" not in sentinel_call_kwargs
assert "ssl_cert_reqs" not in sentinel_call_kwargs
assert "ssl_check_hostname" not in sentinel_call_kwargs
assert "ssl_ca_certs" not in sentinel_call_kwargs
assert "max_connections" not in sentinel_call_kwargs
assert "socket_timeout" not in sentinel_call_kwargs
assert sentinel_call_kwargs["sentinel_kwargs"] == {
"password": "sentinel-secret",
"username": "redis-user",
"ssl": True,
"ssl_cert_reqs": "required",
"ssl_check_hostname": True,
"ssl_ca_certs": "/tmp/test-ca.pem",
"max_connections": 17,
"socket_timeout": 5,
}
assert "service_name" not in sentinel_call_kwargs["sentinel_kwargs"]
assert "sentinel_nodes" not in sentinel_call_kwargs["sentinel_kwargs"]
assert "sentinel_password" not in sentinel_call_kwargs["sentinel_kwargs"]
mock_sentinel.master_for.assert_called_once_with(
"mymaster",
password="redis-secret",
username="redis-user",
ssl=True,
ssl_cert_reqs="required",
ssl_check_hostname=True,
ssl_ca_certs="/tmp/test-ca.pem",
max_connections=17,
socket_timeout=5,
)
@patch("litellm._redis.async_redis.Sentinel")
def test_async_sentinel_uses_sentinel_password_and_master_password(
mock_sentinel_cls,
):
"""Async sentinel auth must mirror the sync sentinel password routing."""
mock_sentinel = MagicMock()
mock_sentinel_cls.return_value = mock_sentinel
get_redis_async_client(
sentinel_nodes=[("sentinel-1", 26379)],
sentinel_password="sentinel-secret",
service_name="mymaster",
password="redis-secret",
username="redis-user",
ssl=True,
ssl_cert_reqs="required",
ssl_check_hostname=True,
ssl_ca_certs="/tmp/test-ca.pem",
max_connections=17,
socket_timeout=5,
)
mock_sentinel_cls.assert_called_once()
sentinel_call_kwargs = mock_sentinel_cls.call_args[1]
assert "password" not in sentinel_call_kwargs
assert "username" not in sentinel_call_kwargs
assert "ssl" not in sentinel_call_kwargs
assert "ssl_cert_reqs" not in sentinel_call_kwargs
assert "ssl_check_hostname" not in sentinel_call_kwargs
assert "ssl_ca_certs" not in sentinel_call_kwargs
assert "max_connections" not in sentinel_call_kwargs
assert "socket_timeout" not in sentinel_call_kwargs
assert sentinel_call_kwargs["sentinel_kwargs"] == {
"password": "sentinel-secret",
"username": "redis-user",
"ssl": True,
"ssl_cert_reqs": "required",
"ssl_check_hostname": True,
"ssl_ca_certs": "/tmp/test-ca.pem",
"max_connections": 17,
"socket_timeout": 5,
}
assert "service_name" not in sentinel_call_kwargs["sentinel_kwargs"]
assert "sentinel_nodes" not in sentinel_call_kwargs["sentinel_kwargs"]
assert "sentinel_password" not in sentinel_call_kwargs["sentinel_kwargs"]
mock_sentinel.master_for.assert_called_once_with(
"mymaster",
password="redis-secret",
username="redis-user",
ssl=True,
ssl_cert_reqs="required",
ssl_check_hostname=True,
ssl_ca_certs="/tmp/test-ca.pem",
max_connections=17,
socket_timeout=5,
)
@patch("litellm._redis.init_redis_cluster")
def test_sync_client_preserves_password_for_cluster_when_url_also_set(
mock_init_cluster, monkeypatch

View file

@ -3002,7 +3002,7 @@ def test_model_info_for_openrouter_kimi_k2_5():
def test_gemini_embedding_2_ga_in_cost_map():
"""GA gemini-embedding-2 entries align with preview multimodal unit pricing."""
"""GA and Vertex preview gemini-embedding-2 entries align with multimodal unit pricing."""
import json
from pathlib import Path
@ -3013,6 +3013,7 @@ def test_gemini_embedding_2_ga_in_cost_map():
for key, provider in (
("gemini/gemini-embedding-2", "gemini"),
("vertex_ai/gemini-embedding-2", "vertex_ai"),
("vertex_ai/gemini-embedding-2-preview", "vertex_ai"),
("gemini-embedding-2", "vertex_ai-embedding-models"),
):
info = model_cost.get(key)

View file

@ -1,7 +1,7 @@
import React, { useEffect } from "react";
import { Form, Select, Tooltip, Collapse, Input, Space, Button, Switch } from "antd";
import { InfoCircleOutlined, MinusCircleOutlined, PlusOutlined } from "@ant-design/icons";
import { MCPServer } from "./types";
import { MCPServer, AUTH_TYPE } from "./types";
const { Panel } = Collapse;
interface MCPPermissionManagementProps {
@ -23,6 +23,8 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
getAccessGroupOptions,
}) => {
const form = Form.useFormInstance();
const watchedAuthType = Form.useWatch("auth_type", form);
const isOAuth2 = watchedAuthType === AUTH_TYPE.OAUTH2;
// Set initial values when mcpServer changes
useEffect(() => {
@ -40,12 +42,25 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
if (typeof mcpServer.available_on_public_internet === "boolean") {
form.setFieldValue("available_on_public_internet", mcpServer.available_on_public_internet);
}
if (typeof mcpServer.delegate_auth_to_upstream === "boolean") {
form.setFieldValue("delegate_auth_to_upstream", mcpServer.delegate_auth_to_upstream);
}
} else {
form.setFieldValue("allow_all_keys", false);
form.setFieldValue("available_on_public_internet", true);
form.setFieldValue("delegate_auth_to_upstream", false);
}
}, [mcpServer, form]);
// delegate_auth_to_upstream is only honored server-side when auth_type=oauth2.
// Force it back to false whenever the user switches away from oauth2 so a
// stale toggle value doesn't get persisted with another auth type.
useEffect(() => {
if (!isOAuth2) {
form.setFieldValue("delegate_auth_to_upstream", false);
}
}, [isOAuth2, form]);
return (
<Collapse className="bg-gray-50 border border-gray-200 rounded-lg" expandIconPosition="end" ghost={false}>
<Panel
@ -105,6 +120,30 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
</Form.Item>
</div>
{isOAuth2 && (
<div className="flex items-start justify-between gap-4">
<div>
<span className="text-sm font-medium text-gray-700 flex items-center">
Delegate auth to upstream (PKCE passthrough)
<Tooltip title="When on, LiteLLM skips its own API key/SSO check for this server and lets the client complete PKCE directly with the upstream MCP server. Only honored when Auth Type is oauth2. No spend tracking or per-key rate limiting will run on this route.">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
</span>
<p className="text-sm text-gray-600 mt-1">
Bypass LiteLLM auth so clients authenticate directly with the upstream OAuth MCP server.
</p>
</div>
<Form.Item
name="delegate_auth_to_upstream"
valuePropName="checked"
initialValue={mcpServer?.delegate_auth_to_upstream ?? false}
className="mb-0"
>
<Switch />
</Form.Item>
</div>
)}
<Form.Item
label={
<span className="text-sm font-medium text-gray-700 flex items-center">

View file

@ -284,6 +284,7 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
credentials: credentialValues,
allow_all_keys: allowAllKeysRaw,
available_on_public_internet: availableOnPublicInternetRaw,
delegate_auth_to_upstream: delegateAuthToUpstreamRaw,
token_validation_json: rawTokenValidationJson,
...restValues
} = values;
@ -388,6 +389,7 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
tool_name_to_description: Object.keys(toolNameToDescription).length > 0 ? toolNameToDescription : null,
allow_all_keys: Boolean(allowAllKeysRaw),
available_on_public_internet: Boolean(availableOnPublicInternetRaw),
delegate_auth_to_upstream: Boolean(delegateAuthToUpstreamRaw),
static_headers: staticHeaders,
...(tokenValidation !== null && { token_validation: tokenValidation }),
};

View file

@ -177,6 +177,47 @@ describe("MCPServerEdit (stdio)", () => {
});
});
describe("MCPServerEdit (delegate auth)", () => {
beforeEach(() => {
vi.clearAllMocks();
});
it("should clear delegate auth flag when saving a non-oauth2 server", async () => {
vi.mocked(networking.updateMCPServer).mockResolvedValue({
...interactiveOAuthServer,
auth_type: "none",
delegate_auth_to_upstream: false,
});
render(
<MCPServerEdit
mcpServer={{
...interactiveOAuthServer,
auth_type: "none",
delegate_auth_to_upstream: true,
}}
accessToken="access-token"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
const saveButtons = screen.getAllByRole("button", { name: "Save Changes" });
await act(async () => {
fireEvent.click(saveButtons[0]);
});
await waitFor(() => {
expect(networking.updateMCPServer).toHaveBeenCalledTimes(1);
});
const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0];
expect(payload.auth_type).toBe("none");
expect(payload.delegate_auth_to_upstream).toBe(false);
});
});
describe("MCPServerEdit (interactive OAuth)", () => {
beforeEach(() => {
vi.clearAllMocks();

View file

@ -384,6 +384,7 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
args: rawArgs,
allow_all_keys: allowAllKeysRaw,
available_on_public_internet: availableOnPublicInternetRaw,
delegate_auth_to_upstream: delegateAuthToUpstreamRaw,
token_validation_json: rawTokenValidationJson,
...restValues
} = values;
@ -552,6 +553,15 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
static_headers: staticHeaders,
allow_all_keys: Boolean(allowAllKeysRaw ?? mcpServer.allow_all_keys),
available_on_public_internet: Boolean(availableOnPublicInternetRaw ?? mcpServer.available_on_public_internet),
// ``delegate_auth_to_upstream`` is only honored server-side for
// ``auth_type=oauth2``. The Form.Item is conditionally rendered so the
// value drops out of the form on auth_type change; force false for any
// non-oauth2 server to avoid persisting a stale ``true`` that would
// silently re-activate if auth_type is later switched back to oauth2.
delegate_auth_to_upstream:
restValues.auth_type === AUTH_TYPE.OAUTH2
? Boolean(delegateAuthToUpstreamRaw ?? mcpServer.delegate_auth_to_upstream)
: false,
// Include token_validation when it is set (non-null) or when clearing an existing value
...(tokenValidation !== null || mcpServer.token_validation
? { token_validation: tokenValidation }

View file

@ -272,6 +272,23 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
)}
</div>
</div>
{handleAuth(mcpServer.auth_type) === "oauth2" && (
<div className="py-3 grid grid-cols-3 gap-4">
<Text className="text-sm font-medium text-gray-500">Delegate Auth to Upstream</Text>
<div className="col-span-2">
{mcpServer.delegate_auth_to_upstream ? (
<span className="inline-flex items-center gap-1 px-2 py-0.5 bg-green-50 text-green-700 rounded-full border border-green-200 text-xs font-medium">
<span className="h-1.5 w-1.5 rounded-full bg-green-500"></span>
Enabled (PKCE passthrough)
</span>
) : (
<span className="inline-flex items-center gap-1 px-2 py-0.5 bg-gray-50 text-gray-600 rounded-full border border-gray-200 text-xs font-medium">
Disabled
</span>
)}
</div>
</div>
)}
<div className="py-3 grid grid-cols-3 gap-4">
<Text className="text-sm font-medium text-gray-500">Access Groups</Text>
<div className="col-span-2">

View file

@ -202,6 +202,7 @@ export interface MCPServer {
tool_name_to_description?: Record<string, string>;
allow_all_keys?: boolean;
available_on_public_internet?: boolean;
delegate_auth_to_upstream?: boolean;
/** Stdio-only fields (present when transport === 'stdio') */
command?: string | null;

View file

@ -8800,6 +8800,7 @@ interface ExchangeMcpOAuthTokenParams {
clientSecret?: string;
codeVerifier: string;
redirectUri: string;
accessToken?: string | null;
}
export const exchangeMcpOAuthToken = async ({
@ -8809,6 +8810,7 @@ export const exchangeMcpOAuthToken = async ({
clientSecret,
codeVerifier,
redirectUri,
accessToken,
}: ExchangeMcpOAuthTokenParams) => {
const base = getProxyBaseUrl();
const normalizedServerId = encodeURIComponent(serverId.trim());
@ -8826,11 +8828,16 @@ export const exchangeMcpOAuthToken = async ({
body.set("code_verifier", codeVerifier);
body.set("redirect_uri", redirectUri);
const headers: Record<string, string> = {
"Content-Type": "application/x-www-form-urlencoded",
};
if (accessToken) {
headers["Authorization"] = `Bearer ${accessToken}`;
}
const response = await fetch(url, {
method: "POST",
headers: {
"Content-Type": "application/x-www-form-urlencoded",
},
headers,
body: body.toString(),
});

View file

@ -279,6 +279,7 @@ export const useMcpOAuthFlow = ({
clientSecret: flowState.clientSecret,
codeVerifier: flowState.codeVerifier,
redirectUri: flowState.redirectUri,
accessToken,
});
onTokenReceived(token);

View file

@ -224,6 +224,7 @@ export const useUserMcpOAuthFlow = ({
clientSecret: flowState.clientSecret,
codeVerifier: flowState.codeVerifier,
redirectUri: flowState.redirectUri,
accessToken,
});
// Persist the token for this user via the backend.