mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
commit
38af29f2c9
45 changed files with 2283 additions and 90 deletions
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "delegate_auth_to_upstream" BOOLEAN NOT NULL DEFAULT false;
|
||||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]}
|
||||
|
|
|
|||
|
|
@ -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__(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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)}")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
36
tests/test_litellm/proxy/test_mcp_asgi_response.py
Normal file
36
tests/test_litellm/proxy/test_mcp_asgi_response.py
Normal 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"
|
||||
}
|
||||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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 }),
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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 }
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
});
|
||||
|
||||
|
|
|
|||
|
|
@ -279,6 +279,7 @@ export const useMcpOAuthFlow = ({
|
|||
clientSecret: flowState.clientSecret,
|
||||
codeVerifier: flowState.codeVerifier,
|
||||
redirectUri: flowState.redirectUri,
|
||||
accessToken,
|
||||
});
|
||||
|
||||
onTokenReceived(token);
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue