diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260513120000_add_delegate_auth_to_upstream_to_mcp_servers/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260513120000_add_delegate_auth_to_upstream_to_mcp_servers/migration.sql new file mode 100644 index 00000000000..50a48743901 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260513120000_add_delegate_auth_to_upstream_to_mcp_servers/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "delegate_auth_to_upstream" BOOLEAN NOT NULL DEFAULT false; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 84ce99557e3..b53507abe6a 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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? diff --git a/litellm/_redis.py b/litellm/_redis.py index 1c11ea829ba..65284162663 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -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): diff --git a/litellm/budget_manager.py b/litellm/budget_manager.py index b25967579e0..bbebb6042cb 100644 --- a/litellm/budget_manager.py +++ b/litellm/budget_manager.py @@ -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": }`` reflecting the reset state. + """ self.user_dict[user]["current_cost"] = 0 self.user_dict[user]["model_cost"] = {} return {"user": self.user_dict[user]} diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index af18c666679..e11d8532dbf 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -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__( diff --git a/litellm/llms/ollama/chat/transformation.py b/litellm/llms/ollama/chat/transformation.py index 48534799c97..e36150a4954 100644 --- a/litellm/llms/ollama/chat/transformation.py +++ b/litellm/llms/ollama/chat/transformation.py @@ -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 diff --git a/litellm/llms/ollama/common_utils.py b/litellm/llms/ollama/common_utils.py index 8aedd9b3500..8ca8b7d383a 100644 --- a/litellm/llms/ollama/common_utils.py +++ b/litellm/llms/ollama/common_utils.py @@ -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 diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 8df25d2c9b5..61e0dd07968 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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 }, diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index a05af66118c..c87e8c414cd 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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//`` 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]: """ diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 4901bc76d2f..31ed0918f3c 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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." diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 41673c7283d..e2b7dcdbb82 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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/ # 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, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index f513381c868..9a9d27b9b84 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 585ba883947..4a28143e617 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -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: diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 3032c9d38cc..ff3df11c448 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -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( diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index f2e64fdde83..587b80d4726 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5851d2550bd..ecef96351ae 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 84ce99557e3..b53507abe6a 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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? diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 48b12a5fba9..e2ba8353591 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -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)}") diff --git a/litellm/timeout.py b/litellm/timeout.py index f9bf036cea2..0d03a3e45e8 100644 --- a/litellm/timeout.py +++ b/litellm/timeout.py @@ -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() diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 268d064eacc..776c7fa67a6 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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 diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 70b065aa918..e62f8686e9d 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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 }, diff --git a/schema.prisma b/schema.prisma index 84ce99557e3..b53507abe6a 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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? diff --git a/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py b/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py index 72b4da7b38b..0a3bf403bf8 100644 --- a/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py +++ b/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py @@ -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 diff --git a/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py b/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py index 05b96b88228..906c51d8064 100644 --- a/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py +++ b/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py @@ -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 tags embedded in content are properly parsed. diff --git a/tests/test_litellm/llms/ollama/test_ollama_model_info.py b/tests/test_litellm/llms/ollama/test_ollama_model_info.py index 95fc80b7fd6..448a26bafe1 100644 --- a/tests/test_litellm/llms/ollama/test_ollama_model_info.py +++ b/tests/test_litellm/llms/ollama/test_ollama_model_info.py @@ -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.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 696f9de54b2..3b7ff909105 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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//garbage`` was parsed as + targeting ```` by the auth gate (bypassing LiteLLM auth) + while the routing layer parsed it as ``/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 ``//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//`` must NOT bypass auth. + + The bypass key check is now performed against the same parsed target + as downstream routing — ``/`` — 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: `` 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/`` 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""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index f8fab3aff67..24d1962e60a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index b53420f0000..794864f658b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py index b7b9c984bf9..0505f98cb8e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -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 diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index bcd7fcb37b3..80a4804956c 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -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(): diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 30ad84e18b8..47e058786a5 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -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 ( diff --git a/tests/test_litellm/proxy/test_health_check_max_tokens.py b/tests/test_litellm/proxy/test_health_check_max_tokens.py index 72d77862b5d..5cb7cdacc60 100644 --- a/tests/test_litellm/proxy/test_health_check_max_tokens.py +++ b/tests/test_litellm/proxy/test_health_check_max_tokens.py @@ -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"} diff --git a/tests/test_litellm/proxy/test_mcp_asgi_response.py b/tests/test_litellm/proxy/test_mcp_asgi_response.py new file mode 100644 index 00000000000..d030f65af4b --- /dev/null +++ b/tests/test_litellm/proxy/test_mcp_asgi_response.py @@ -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" + } diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py index 5d8ff8022e5..503a610e016 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -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"} diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 2483469db23..282c9d72d51 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -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 diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index d07af922ea6..bc60375f906 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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) diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx b/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx index dcb7298c830..8f60e50d2b2 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx @@ -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 = ({ 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 = ({ 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 ( = ({ + {isOAuth2 && ( +
+
+ + Delegate auth to upstream (PKCE passthrough) + + + + +

+ Bypass LiteLLM auth so clients authenticate directly with the upstream OAuth MCP server. +

+
+ + + +
+ )} + diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx index 17bcd59c43e..f8b0141b25d 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx @@ -284,6 +284,7 @@ const CreateMCPServer: React.FC = ({ 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 = ({ 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 }), }; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx index 760504d5797..1f2864f6759 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx @@ -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( + , + ); + + 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(); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx index 997e76982d4..9278d41c3e3 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx @@ -384,6 +384,7 @@ const MCPServerEdit: React.FC = ({ 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 = ({ 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 } diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx index f9a3d57e952..1f8f7f68d33 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx @@ -272,6 +272,23 @@ export const MCPServerView: React.FC = ({ )} + {handleAuth(mcpServer.auth_type) === "oauth2" && ( +
+ Delegate Auth to Upstream +
+ {mcpServer.delegate_auth_to_upstream ? ( + + + Enabled (PKCE passthrough) + + ) : ( + + Disabled + + )} +
+
+ )}
Access Groups
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 814c5d74f46..7cfe08d9ee5 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -202,6 +202,7 @@ export interface MCPServer { tool_name_to_description?: Record; allow_all_keys?: boolean; available_on_public_internet?: boolean; + delegate_auth_to_upstream?: boolean; /** Stdio-only fields (present when transport === 'stdio') */ command?: string | null; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 7f28dbe6da9..756348f4937 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -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 = { + "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(), }); diff --git a/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx b/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx index 9157fedbe21..7edeade4cbd 100644 --- a/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx +++ b/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx @@ -279,6 +279,7 @@ export const useMcpOAuthFlow = ({ clientSecret: flowState.clientSecret, codeVerifier: flowState.codeVerifier, redirectUri: flowState.redirectUri, + accessToken, }); onTokenReceived(token); diff --git a/ui/litellm-dashboard/src/hooks/useUserMcpOAuthFlow.tsx b/ui/litellm-dashboard/src/hooks/useUserMcpOAuthFlow.tsx index aa7ce84de5e..cf0a81dcadf 100644 --- a/ui/litellm-dashboard/src/hooks/useUserMcpOAuthFlow.tsx +++ b/ui/litellm-dashboard/src/hooks/useUserMcpOAuthFlow.tsx @@ -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.