Merge pull request #25699 from BerriAI/litellm_ishaan_april14

Litellm ishaan april14
This commit is contained in:
ishaan-berri 2026-04-15 19:01:06 -07:00 • committed by GitHub
commit 0b7335201b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
29 changed files with 1800 additions and 457 deletions

View file

@ -487,6 +487,7 @@ router_settings:
| AZURE_STORAGE_CLIENT_ID | The Application Client ID to use for Authentication to Azure Blob Storage logging
| AZURE_STORAGE_CLIENT_SECRET | The Application Client Secret to use for Authentication to Azure Blob Storage logging
| AZURE_VECTOR_STORE_COST_PER_GB_PER_DAY | Cost per GB per day for Azure Vector Store service
| BACKGROUND_HEALTH_CHECK_MAX_TOKENS | Optional global default for `max_tokens` on proxy background health checks when a model has no `health_check_max_tokens`. If unset, non-wildcard models default to 1. Applies to wildcard routes when set. Default is unset
| BATCH_STATUS_POLL_INTERVAL_SECONDS | Interval in seconds for polling batch status. Default is 3600 (1 hour)
| BATCH_STATUS_POLL_MAX_ATTEMPTS | Maximum number of attempts for polling batch status. Default is 24 (for 24 hours)
| BEDROCK_MAX_POLICY_SIZE | Maximum size for Bedrock policy. Default is 75

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "instructions" TEXT;

View file

@ -289,6 +289,7 @@ model LiteLLM_MCPServerTable {
server_name String?
alias String?
description String?
instructions String?
url String?
spec_path String?
transport String @default("sse")

View file

@ -1,6 +1,6 @@
import os
import sys
from typing import List, Literal
from typing import List, Literal, Optional
from litellm.litellm_core_utils.env_utils import get_env_int
@ -1343,6 +1343,22 @@ BATCH_STATUS_POLL_MAX_ATTEMPTS = int(
HEALTH_CHECK_TIMEOUT_SECONDS = int(
os.getenv("HEALTH_CHECK_TIMEOUT_SECONDS", 60)
) # 60 seconds
_background_health_check_max_tokens_env = os.getenv(
"BACKGROUND_HEALTH_CHECK_MAX_TOKENS"
)
try:
_raw_background_health_check_max_tokens = (
_background_health_check_max_tokens_env.strip()
if _background_health_check_max_tokens_env is not None
else ""
)
BACKGROUND_HEALTH_CHECK_MAX_TOKENS: Optional[int] = (
int(_raw_background_health_check_max_tokens)
if _raw_background_health_check_max_tokens
else None
)
except (ValueError, TypeError):
BACKGROUND_HEALTH_CHECK_MAX_TOKENS = None
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME = "litellm-internal-health-check"
LITTELM_CLI_SERVICE_ACCOUNT_NAME = "litellm-cli"
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME = "litellm_internal_jobs"

View file

@ -221,6 +221,7 @@ class MCPClient:
self.extra_headers: Optional[Dict[str, str]] = extra_headers
self.ssl_verify: Optional[VerifyTypes] = ssl_verify
self._aws_auth: Optional[httpx.Auth] = aws_auth
self._last_initialize_instructions: Optional[str] = None
# handle the basic auth value if provided
if auth_value:
self.update_auth_value(auth_value)
@ -296,7 +297,12 @@ class MCPClient:
session_ctx = ClientSession(read_stream, write_stream)
session = await session_ctx.__aenter__()
try:
await session.initialize()
init_result = await session.initialize()
self._last_initialize_instructions = None
if init_result is not None:
ins = getattr(init_result, "instructions", None)
if isinstance(ins, str) and ins.strip():
self._last_initialize_instructions = ins.strip()
return await operation(session)
finally:
try:
@ -315,6 +321,7 @@ class MCPClient:
"""Open a session, run the provided coroutine, and clean up."""
http_client: Optional[httpx.AsyncClient] = None
try:
self._last_initialize_instructions = None
transport_ctx, http_client = self._create_transport_context()
return await self._execute_session_operation(transport_ctx, operation)
except Exception:

View file

@ -11,6 +11,10 @@ from openai.types.completion_create_params import (
CompletionCreateParamsStreaming as TextCompletionCreateParamsStreaming,
)
from openai.types.embedding_create_params import EmbeddingCreateParams
from openai.types.responses.response_create_params import (
ResponseCreateParamsNonStreaming,
ResponseCreateParamsStreaming,
)
from litellm._logging import verbose_logger
from litellm.types.rerank import RerankRequest
@ -65,6 +69,9 @@ class ModelParamHelper:
ModelParamHelper._get_litellm_supported_transcription_kwargs()
)
rerank_kwargs = ModelParamHelper._get_litellm_supported_rerank_kwargs()
responses_api_kwargs = (
ModelParamHelper._get_litellm_supported_responses_api_kwargs()
)
exclude_kwargs = ModelParamHelper._get_exclude_kwargs()
combined_kwargs = chat_completion_kwargs.union(
@ -72,6 +79,7 @@ class ModelParamHelper:
embedding_kwargs,
transcription_kwargs,
rerank_kwargs,
responses_api_kwargs,
)
combined_kwargs = combined_kwargs.difference(exclude_kwargs)
return combined_kwargs
@ -93,9 +101,9 @@ class ModelParamHelper:
streaming_params: Set[str] = set(
getattr(CompletionCreateParamsStreaming, "__annotations__", {}).keys()
)
litellm_provider_specific_params: Set[
str
] = ModelParamHelper.get_litellm_provider_specific_params_for_chat_params()
litellm_provider_specific_params: Set[str] = (
ModelParamHelper.get_litellm_provider_specific_params_for_chat_params()
)
all_chat_completion_kwargs: Set[str] = non_streaming_params.union(
streaming_params
).union(litellm_provider_specific_params)
@ -167,6 +175,21 @@ class ModelParamHelper:
verbose_logger.debug("Error getting transcription kwargs %s", str(e))
return set()
@staticmethod
def _get_litellm_supported_responses_api_kwargs() -> Set[str]:
"""
Get the litellm supported responses API kwargs
This follows the OpenAI API Spec
"""
non_streaming_params: Set[str] = set(
getattr(ResponseCreateParamsNonStreaming, "__annotations__", {}).keys()
)
streaming_params: Set[str] = set(
getattr(ResponseCreateParamsStreaming, "__annotations__", {}).keys()
)
return non_streaming_params.union(streaming_params)
@staticmethod
def _get_exclude_kwargs() -> Set[str]:
"""

View file

@ -2,6 +2,7 @@
This file contains common utils for anthropic calls.
"""
import copy
from typing import Any, Dict, List, Optional, Union
import httpx
@ -736,6 +737,69 @@ def strip_advisor_blocks_from_messages(
return messages
def is_anthropic_invalid_thinking_signature_error(error_text: str) -> bool:
"""
Detect Anthropic 400 when encrypted thinking signatures in history do not match
the current deployment (e.g. user rotated API key or switched model endpoint).
Example API message:
messages.N.content.M: Invalid `signature` in `thinking` block
"""
if not error_text:
return False
lower = error_text.lower()
return (
"invalid" in lower
and "signature" in lower
and "thinking" in lower
and "block" in lower
)
def strip_thinking_blocks_from_anthropic_messages(messages: List[Any]) -> List[Any]:
"""
Return a new message list with thinking / redacted_thinking content blocks removed
from each message. Used to recover from invalid thinking signatures on retry.
Messages whose content is a list and becomes empty after stripping are omitted,
since Anthropic rejects empty content arrays.
"""
out: List[Any] = []
for m in messages:
if not isinstance(m, dict):
out.append(m)
continue
mm = copy.deepcopy(m)
content = mm.get("content")
if isinstance(content, list):
filtered = [
b
for b in content
if not (
isinstance(b, dict)
and b.get("type") in ("thinking", "redacted_thinking")
)
]
if not filtered:
continue
mm["content"] = filtered
out.append(mm)
return out
def strip_thinking_blocks_from_anthropic_messages_request_dict(
data: Dict[str, Any],
) -> None:
"""
Mutate an Anthropic Messages-style request dict: strip thinking blocks from
``messages`` and remove the top-level ``thinking`` extended-thinking param.
"""
msgs = data.get("messages")
if isinstance(msgs, list):
data["messages"] = strip_thinking_blocks_from_anthropic_messages(msgs)
data.pop("thinking", None)
def process_anthropic_headers(headers: Union[httpx.Headers, dict]) -> dict:
openai_headers = {}
if "anthropic-ratelimit-requests-limit" in headers:

View file

@ -1,7 +1,9 @@
from typing import TYPE_CHECKING, List, Optional, Tuple
import httpx
from httpx import Response
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.secret_managers.main import get_secret_str
@ -11,6 +13,8 @@ from litellm.types.router import GenericLiteLLMParams
if TYPE_CHECKING:
from httpx import URL
from litellm.types.utils import CostResponseTypes
class AzurePassthroughConfig(BasePassthroughConfig):
def is_streaming_request(self, endpoint: str, request_data: dict) -> bool:
@ -83,3 +87,36 @@ class AzurePassthroughConfig(BasePassthroughConfig):
self, api_key: Optional[str] = None, api_base: Optional[str] = None
) -> List[str]:
return super().get_models(api_key, api_base)
def logging_non_streaming_response(
self,
model: str,
custom_llm_provider: str,
httpx_response: Response,
request_data: dict,
logging_obj: Logging,
endpoint: str,
) -> Optional["CostResponseTypes"]:
from litellm import encoding
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
from litellm.types.utils import ModelResponse
if "chat/completions" not in endpoint:
return None
openai_chat_config = OpenAIGPTConfig()
litellm_model_response: ModelResponse = openai_chat_config.transform_response(
model=model,
messages=[{"role": "user", "content": "no-message-pass-through-endpoint"}],
raw_response=httpx_response,
model_response=ModelResponse(),
logging_obj=logging_obj,
optional_params={},
litellm_params={},
api_key="",
request_data=request_data,
encoding=encoding,
)
return litellm_model_response

View file

@ -120,3 +120,46 @@ class BaseAnthropicMessagesConfig(ABC):
return BaseLLMException(
message=error_message, status_code=status_code, headers=headers
)
@property
def max_retry_on_anthropic_messages_http_error(self) -> int:
"""
Max HTTP attempts for /v1/messages when the handler may mutate the body and
retry (e.g. strip invalid encrypted thinking signatures after a deployment or
credential change).
"""
return 2
def should_retry_anthropic_messages_on_http_error(
self, e: httpx.HTTPStatusError, litellm_params: dict
) -> bool:
"""
When True, async_anthropic_messages_handler will transform the request body
and issue one more attempt (bounded by max_retry_on_anthropic_messages_http_error).
"""
from litellm.llms.anthropic.common_utils import (
is_anthropic_invalid_thinking_signature_error,
)
return (
e.response.status_code == 400
and is_anthropic_invalid_thinking_signature_error(e.response.text)
)
def transform_anthropic_messages_request_on_http_error(
self, e: httpx.HTTPStatusError, request_data: dict
) -> dict:
"""
Mutates request_data in place when retrying after a recoverable HTTP error.
"""
from litellm.llms.anthropic.common_utils import (
is_anthropic_invalid_thinking_signature_error,
strip_thinking_blocks_from_anthropic_messages_request_dict,
)
if (
e.response.status_code == 400
and is_anthropic_invalid_thinking_signature_error(e.response.text)
):
strip_thinking_blocks_from_anthropic_messages_request_dict(request_data)
return request_data

View file

@ -1816,6 +1816,73 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
)
async def _async_post_anthropic_messages_with_http_error_retry(
self,
async_httpx_client: AsyncHTTPHandler,
request_url: str,
headers: dict,
signed_json_body: Optional[bytes],
request_body: dict,
stream: bool,
logging_obj: LiteLLMLoggingObj,
provider_config: BaseAnthropicMessagesConfig,
litellm_params: GenericLiteLLMParams,
api_key: Optional[str],
model: str,
) -> httpx.Response:
max_attempts = max(
provider_config.max_retry_on_anthropic_messages_http_error, 1
)
litellm_params_dict = dict(litellm_params)
optional_params_dict = dict(litellm_params)
for attempt_idx in range(max_attempts):
try:
response = await async_httpx_client.post(
url=request_url,
headers=headers,
data=signed_json_body or json.dumps(request_body),
stream=stream or False,
logging_obj=logging_obj,
)
response.raise_for_status()
return response
except httpx.HTTPStatusError as e:
hit_max_attempt = attempt_idx + 1 == max_attempts
should_retry = (
provider_config.should_retry_anthropic_messages_on_http_error(
e=e, litellm_params=litellm_params_dict
)
)
if should_retry and not hit_max_attempt:
verbose_logger.debug(
"Anthropic /v1/messages: invalid thinking signature; "
"stripping thinking blocks and retrying (attempt %s/%s).",
attempt_idx + 2,
max_attempts,
)
provider_config.transform_anthropic_messages_request_on_http_error(
e=e, request_data=request_body
)
headers, signed_json_body = provider_config.sign_request(
headers=headers,
optional_params=optional_params_dict,
request_data=request_body,
api_base=request_url,
api_key=api_key,
stream=stream,
fake_stream=False,
model=model,
)
logging_obj.model_call_details.update(request_body)
continue
raise self._handle_error(e=e, provider_config=provider_config)
except Exception as e:
raise self._handle_error(e=e, provider_config=provider_config)
raise RuntimeError(
"unreachable: anthropic messages HTTP retry loop exited without return"
)
async def async_anthropic_messages_handler(
self,
model: str,
@ -1955,19 +2022,19 @@ class BaseLLMHTTPHandler:
},
)
try:
response = await async_httpx_client.post(
url=request_url,
headers=headers,
data=signed_json_body or json.dumps(request_body),
stream=stream or False,
logging_obj=logging_obj,
)
response.raise_for_status()
except Exception as e:
raise self._handle_error(
e=e, provider_config=anthropic_messages_provider_config
)
response = await self._async_post_anthropic_messages_with_http_error_retry(
async_httpx_client=async_httpx_client,
request_url=request_url,
headers=headers,
signed_json_body=signed_json_body,
request_body=request_body,
stream=stream or False,
logging_obj=logging_obj,
provider_config=anthropic_messages_provider_config,
litellm_params=litellm_params,
api_key=api_key,
model=model,
)
# used for logging + cost tracking
logging_obj.model_call_details["httpx_response"] = response
@ -4496,9 +4563,9 @@ class BaseLLMHTTPHandler:
# Second: Execute agentic loop
# Add custom_llm_provider to kwargs so the agentic loop can reconstruct the full model name
kwargs_with_provider = kwargs.copy() if kwargs else {}
kwargs_with_provider[
"custom_llm_provider"
] = custom_llm_provider
kwargs_with_provider["custom_llm_provider"] = (
custom_llm_provider
)
agentic_response = await callback.async_run_agentic_loop(
tools=tool_calls,
model=model,
@ -4614,9 +4681,9 @@ class BaseLLMHTTPHandler:
# Second: Execute agentic loop
# Add custom_llm_provider to kwargs so the agentic loop can reconstruct the full model name
kwargs_with_provider = kwargs.copy() if kwargs else {}
kwargs_with_provider[
"custom_llm_provider"
] = custom_llm_provider
kwargs_with_provider["custom_llm_provider"] = (
custom_llm_provider
)
agentic_response = (
await callback.async_run_chat_completion_agentic_loop(
tools=tool_calls,
@ -5110,7 +5177,10 @@ class BaseLLMHTTPHandler:
_is_async: bool = False,
fake_stream: bool = False,
litellm_metadata: Optional[Dict[str, Any]] = None,
) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]:
) -> Union[
ImageResponse,
Coroutine[Any, Any, ImageResponse],
]:
"""
Handles image edit requests.
@ -5322,7 +5392,10 @@ class BaseLLMHTTPHandler:
fake_stream: bool = False,
litellm_metadata: Optional[Dict[str, Any]] = None,
api_key: Optional[str] = None,
) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]:
) -> Union[
ImageResponse,
Coroutine[Any, Any, ImageResponse],
]:
"""
Handles image generation requests.
When _is_async=True, returns a coroutine instead of making the call directly.
@ -5562,7 +5635,10 @@ class BaseLLMHTTPHandler:
fake_stream: bool = False,
litellm_metadata: Optional[Dict[str, Any]] = None,
api_key: Optional[str] = None,
) -> Union[VideoObject, Coroutine[Any, Any, VideoObject],]:
) -> Union[
VideoObject,
Coroutine[Any, Any, VideoObject],
]:
"""
Handles video generation requests.
When _is_async=True, returns a coroutine instead of making the call directly.

View file

@ -14,3 +14,8 @@ from typing import Optional
_mcp_active_toolset_id: ContextVar[Optional[str]] = ContextVar(
"_mcp_active_toolset_id", default=None
)
# Per-request merged InitializeResult.instructions; set in MCP HTTP/SSE handlers.
_mcp_gateway_initialize_instructions: ContextVar[Optional[str]] = ContextVar(
"_mcp_gateway_initialize_instructions", default=None
)

View file

@ -184,6 +184,16 @@ class MCPServerManager:
"gmail_send_email": "zapier_mcp_server",
}
"""
self._upstream_initialize_instructions_by_server_id: Dict[str, str] = {}
def _remember_upstream_initialize_instructions(
self, server: MCPServer, client: MCPClient
) -> None:
raw = getattr(client, "_last_initialize_instructions", None)
if raw and str(raw).strip():
self._upstream_initialize_instructions_by_server_id[server.server_id] = str(
raw
).strip()
def get_registry(self) -> Dict[str, MCPServer]:
"""
@ -204,6 +214,7 @@ class MCPServerManager:
mcp_aliases: Optional dictionary mapping aliases to server names from litellm_settings
"""
verbose_logger.debug("Loading MCP Servers from config-----")
self._upstream_initialize_instructions_by_server_id.clear()
# Track which aliases have been used to ensure only first occurrence is used
used_aliases = set()
@ -351,6 +362,7 @@ class MCPServerManager:
aws_service_name=server_config.get("aws_service_name", None),
aws_role_name=server_config.get("aws_role_name", None),
aws_session_name=server_config.get("aws_session_name", None),
instructions=server_config.get("instructions", None),
)
self.config_mcp_servers[server_id] = new_server
@ -693,6 +705,7 @@ class MCPServerManager:
aws_service_name=aws_creds.get("aws_service_name"),
aws_role_name=aws_creds.get("aws_role_name"),
aws_session_name=aws_creds.get("aws_session_name"),
instructions=mcp_server.instructions,
)
return new_server
@ -1247,6 +1260,7 @@ class MCPServerManager:
return tools
else:
tools = await self._fetch_tools_with_timeout(client, server.name)
self._remember_upstream_initialize_instructions(server, client)
prefixed_or_original_tools = self._create_prefixed_tools(
tools, server, add_prefix=add_prefix
@ -2383,6 +2397,7 @@ class MCPServerManager:
# If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task)
result_index = 1 if proxy_logging_obj else 0
result = mcp_responses[result_index]
self._remember_upstream_initialize_instructions(mcp_server, client)
return cast(CallToolResult, result)
@ -2627,6 +2642,7 @@ class MCPServerManager:
)
verbose_logger.debug("Loading MCP servers from database into registry...")
self._upstream_initialize_instructions_by_server_id.clear()
# perform authz check to filter the mcp servers user has access to
prisma_client = get_prisma_client_or_throw(
@ -2910,6 +2926,7 @@ class MCPServerManager:
await asyncio.wait_for(
client.run_with_session(_noop), timeout=MCP_HEALTH_CHECK_TIMEOUT
)
self._remember_upstream_initialize_instructions(server, client)
status = "healthy"
except asyncio.TimeoutError:
health_check_error = (
@ -2951,6 +2968,7 @@ class MCPServerManager:
token_url=server.token_url,
registration_url=server.registration_url,
allow_all_keys=server.allow_all_keys,
instructions=server.instructions,
)
async def get_all_mcp_servers_with_health_and_teams(
@ -3046,6 +3064,7 @@ class MCPServerManager:
is_byok=server.is_byok,
byok_description=server.byok_description,
byok_api_key_help_url=server.byok_api_key_help_url,
instructions=server.instructions,
)
async def get_all_mcp_servers_unfiltered(self) -> List[LiteLLM_MCPServerTable]:

View file

@ -933,6 +933,7 @@ if MCP_AVAILABLE:
authorization_url=request.authorization_url,
registration_url=request.registration_url,
oauth2_flow=_oauth2_flow,
instructions=request.instructions,
)
stdio_env = global_mcp_server_manager._build_stdio_env(

View file

@ -7,6 +7,7 @@ LiteLLM MCP Server Routes
import asyncio
import contextlib
import time
import types
import traceback
import uuid
from datetime import datetime
@ -37,7 +38,10 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
get_request_base_url,
)
from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_active_toolset_id
from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_active_toolset_id,
_mcp_gateway_initialize_instructions,
)
from litellm.proxy._experimental.mcp_server.mcp_debug import MCPDebug
from litellm.proxy._experimental.mcp_server.utils import (
LITELLM_MCP_SERVER_DESCRIPTION,
@ -122,6 +126,8 @@ _INITIALIZATION_LOCK = asyncio.Lock()
if MCP_AVAILABLE:
from mcp.server import Server
from mcp.server.lowlevel.server import NotificationOptions
from mcp.server.models import InitializationOptions
# Import auth context variables and middleware
from mcp.server.auth.middleware.auth_context import (
@ -200,6 +206,21 @@ if MCP_AVAILABLE:
)
return normalized
def _gateway_create_initialization_options(
self,
notification_options: Optional[NotificationOptions] = None,
experimental_capabilities: Optional[Dict[str, Dict[str, Any]]] = None,
) -> InitializationOptions:
opts = Server.create_initialization_options(
self,
notification_options=notification_options,
experimental_capabilities=experimental_capabilities or {},
)
merged = _mcp_gateway_initialize_instructions.get()
if merged is not None:
return opts.model_copy(update={"instructions": merged})
return opts
########################################################
############ Initialize the MCP Server #################
########################################################
@ -207,6 +228,9 @@ if MCP_AVAILABLE:
name=LITELLM_MCP_SERVER_NAME,
version=LITELLM_MCP_SERVER_VERSION,
)
server.create_initialization_options = types.MethodType( # type: ignore[method-assign]
_gateway_create_initialization_options, server
)
sse: SseServerTransport = SseServerTransport("/mcp/sse/messages")
# Create session managers
@ -1021,7 +1045,9 @@ if MCP_AVAILABLE:
except (ValueError, TypeError):
pass
ttl = _compute_per_user_token_ttl(server, raw_expires)
await mcp_per_user_token_cache.set(user_id, server_id, access_token, ttl)
await mcp_per_user_token_cache.set(
user_id, server_id, access_token, ttl
)
return {"Authorization": f"Bearer {access_token}"}
except Exception as e:
@ -1103,6 +1129,57 @@ if MCP_AVAILABLE:
return server_auth_header, extra_headers
def _merge_gateway_initialize_instructions(
allowed_mcp_servers: List[MCPServer],
) -> Optional[str]:
"""YAML/DB override, else in-memory upstream text from list_tools / health_check / call_tool."""
if not allowed_mcp_servers:
return None
texts: List[Tuple[str, str]] = []
for server in allowed_mcp_servers:
label = (
server.alias
or server.server_name
or server.name
or server.server_id
or "mcp"
)
if server.instructions and server.instructions.strip():
texts.append((label, server.instructions.strip()))
continue
if server.spec_path:
continue
cached = global_mcp_server_manager._upstream_initialize_instructions_by_server_id.get(
server.server_id
)
if cached and cached.strip():
texts.append((label, cached.strip()))
if not texts:
return None
if len(texts) == 1:
return texts[0][1]
return "\n\n---\n\n".join(f"[{lbl}]\n{txt}" for lbl, txt in texts)
@contextlib.asynccontextmanager
async def _gateway_initialize_instructions_request_scope(
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_servers: Optional[List[str]],
client_ip: Optional[str],
) -> AsyncIterator[None]:
allowed = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
client_ip=client_ip,
)
merged = _merge_gateway_initialize_instructions(allowed_mcp_servers=allowed)
tok = _mcp_gateway_initialize_instructions.set(merged)
try:
yield
finally:
_mcp_gateway_initialize_instructions.reset(tok)
async def _get_tools_from_mcp_servers( # noqa: PLR0915
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_auth_header: Optional[str],
@ -2670,7 +2747,12 @@ if MCP_AVAILABLE:
# Request was fully handled (e.g., DELETE on non-existent session)
return
await session_manager.handle_request(scope, receive, send)
async with _gateway_initialize_instructions_request_scope(
user_api_key_auth,
mcp_servers,
_client_ip,
):
await session_manager.handle_request(scope, receive, send)
except HTTPException:
# Re-raise HTTP exceptions to preserve status codes and details
raise
@ -2729,7 +2811,12 @@ if MCP_AVAILABLE:
await initialize_session_managers()
await asyncio.sleep(0.1)
await sse_session_manager.handle_request(scope, receive, send)
async with _gateway_initialize_instructions_request_scope(
user_api_key_auth,
mcp_servers,
_sse_client_ip,
):
await sse_session_manager.handle_request(scope, receive, send)
except Exception as e:
verbose_logger.exception(f"Error handling MCP request: {e}")
# Instead of re-raising, try to send a graceful error response

View file

@ -904,9 +904,9 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
allowed_cache_controls: Optional[list] = []
config: Optional[dict] = {}
permissions: Optional[dict] = {}
model_max_budget: Optional[
dict
] = {} # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {}
model_max_budget: Optional[dict] = (
{}
) # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {}
model_config = ConfigDict(protected_namespaces=())
model_rpm_limit: Optional[dict] = None
@ -1048,9 +1048,9 @@ class RegenerateKeyRequest(GenerateKeyRequest):
spend: Optional[float] = None
metadata: Optional[dict] = None
new_master_key: Optional[str] = None
grace_period: Optional[
str
] = None # Duration to keep old key valid (e.g. "24h", "2d"); None = immediate revoke
grace_period: Optional[str] = (
None # Duration to keep old key valid (e.g. "24h", "2d"); None = immediate revoke
)
class ResetSpendRequest(LiteLLMPydanticObjectBase):
@ -1137,6 +1137,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
tool_name_to_description: Optional[Dict[str, str]] = None
extra_headers: Optional[List[str]] = None
static_headers: Optional[Dict[str, str]] = None
instructions: Optional[str] = None
# Stdio-specific fields
command: Optional[str] = None
args: List[str] = Field(default_factory=list)
@ -1219,6 +1220,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
tool_name_to_description: Optional[Dict[str, str]] = None
extra_headers: Optional[List[str]] = None
static_headers: Optional[Dict[str, str]] = None
instructions: Optional[str] = None
# Stdio-specific fields
command: Optional[str] = None
args: List[str] = Field(default_factory=list)
@ -1270,6 +1272,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
transport: MCPTransportType
auth_type: Optional[MCPAuthType] = None
credentials: Optional[MCPCredentials] = None
instructions: Optional[str] = None
created_at: Optional[datetime] = None
created_by: Optional[str] = None
updated_at: Optional[datetime] = None
@ -1574,12 +1577,12 @@ class NewCustomerRequest(BudgetNewRequest):
blocked: bool = False # allow/disallow requests for this end-user
budget_id: Optional[str] = None # give either a budget_id or max_budget
spend: Optional[float] = None
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
default_model: Optional[
str
] = None # if no equivalent model in allowed region - default all requests to this model
allowed_model_region: Optional[AllowedModelRegion] = (
None # require all user requests to use models in this specific region
)
default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
@model_validator(mode="before")
@ -1602,12 +1605,12 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase):
blocked: bool = False # allow/disallow requests for this end-user
max_budget: Optional[float] = None
budget_id: Optional[str] = None # give either a budget_id or max_budget
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
default_model: Optional[
str
] = None # if no equivalent model in allowed region - default all requests to this model
allowed_model_region: Optional[AllowedModelRegion] = (
None # require all user requests to use models in this specific region
)
default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
@ -1697,15 +1700,15 @@ class NewTeamRequest(TeamBase):
] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm
model_tpm_limit: Optional[Dict[str, int]] = None
team_member_budget: Optional[
float
] = None # allow user to set a budget for all team members
team_member_rpm_limit: Optional[
int
] = None # allow user to set RPM limit for all team members
team_member_tpm_limit: Optional[
int
] = None # allow user to set TPM limit for all team members
team_member_budget: Optional[float] = (
None # allow user to set a budget for all team members
)
team_member_rpm_limit: Optional[int] = (
None # allow user to set RPM limit for all team members
)
team_member_tpm_limit: Optional[int] = (
None # allow user to set TPM limit for all team members
)
team_member_key_duration: Optional[str] = None # e.g. "1d", "1w", "1m"
team_member_budget_duration: Optional[str] = None # e.g. "30d", "1mo"
allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None
@ -1802,9 +1805,9 @@ class BlockKeyRequest(LiteLLMPydanticObjectBase):
class AddTeamCallback(LiteLLMPydanticObjectBase):
callback_name: str
callback_type: Optional[
Literal["success", "failure", "success_and_failure"]
] = "success_and_failure"
callback_type: Optional[Literal["success", "failure", "success_and_failure"]] = (
"success_and_failure"
)
callback_vars: Dict[str, str]
@model_validator(mode="before")
@ -2146,9 +2149,9 @@ class ConfigList(LiteLLMPydanticObjectBase):
stored_in_db: Optional[bool]
field_default_value: Any
premium_field: bool = False
nested_fields: Optional[
List[FieldDetail]
] = None # For nested dictionary or Pydantic fields
nested_fields: Optional[List[FieldDetail]] = (
None # For nested dictionary or Pydantic fields
)
class UserHeaderMapping(LiteLLMPydanticObjectBase):
@ -2507,9 +2510,9 @@ class UserAPIKeyAuth(
user_max_budget: Optional[float] = None
request_route: Optional[str] = None
user: Optional[Any] = None # Expanded user object when expand=user is used
created_by_user: Optional[
Any
] = None # Expanded created_by user when expand=user is used
created_by_user: Optional[Any] = (
None # Expanded created_by user when expand=user is used
)
end_user_object_permission: Optional[LiteLLM_ObjectPermissionTable] = None
# Decoded upstream IdP claims (groups, roles, etc.) propagated by JWT auth machinery
# and forwarded into outbound tokens by guardrails such as MCPJWTSigner.
@ -2648,9 +2651,9 @@ class LiteLLM_OrganizationMembershipTable(LiteLLMPydanticObjectBase):
budget_id: Optional[str] = None
created_at: datetime
updated_at: datetime
user: Optional[
Any
] = None # You might want to replace 'Any' with a more specific type if available
user: Optional[Any] = (
None # You might want to replace 'Any' with a more specific type if available
)
litellm_budget_table: Optional[LiteLLM_BudgetTable] = None
user_email: Optional[str] = None
@ -3805,9 +3808,9 @@ class TeamModelDeleteRequest(BaseModel):
# Organization Member Requests
class OrganizationMemberAddRequest(OrgMemberAddRequest):
organization_id: str
max_budget_in_organization: Optional[
float
] = None # Users max budget within the organization
max_budget_in_organization: Optional[float] = (
None # Users max budget within the organization
)
class OrganizationMemberDeleteRequest(MemberDeleteRequest):
@ -4062,9 +4065,9 @@ class ProviderBudgetResponse(LiteLLMPydanticObjectBase):
Maps provider names to their budget configs.
"""
providers: Dict[
str, ProviderBudgetResponseObject
] = {} # Dictionary mapping provider names to their budget configurations
providers: Dict[str, ProviderBudgetResponseObject] = (
{}
) # Dictionary mapping provider names to their budget configurations
class ProxyStateVariables(TypedDict):
@ -4226,9 +4229,9 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
enforce_rbac: bool = False
roles_jwt_field: Optional[str] = None # v2 on role mappings
role_mappings: Optional[List[RoleMapping]] = None
object_id_jwt_field: Optional[
str
] = None # can be either user / team, inferred from the role mapping
object_id_jwt_field: Optional[str] = (
None # can be either user / team, inferred from the role mapping
)
scope_mappings: Optional[List[ScopeMapping]] = None
enforce_scope_based_access: bool = False
enforce_team_based_model_access: bool = False

View file

@ -11,7 +11,11 @@ from typing import List, Optional
import litellm
logger = logging.getLogger(__name__)
from litellm.constants import DEFAULT_HEALTH_CHECK_PROMPT, HEALTH_CHECK_TIMEOUT_SECONDS
from litellm.constants import (
BACKGROUND_HEALTH_CHECK_MAX_TOKENS,
DEFAULT_HEALTH_CHECK_PROMPT,
HEALTH_CHECK_TIMEOUT_SECONDS,
)
ILLEGAL_DISPLAY_PARAMS = [
"messages",
@ -242,7 +246,9 @@ async def _perform_health_check(
cleaned["model_id"] = _model_id
if isinstance(is_healthy, Exception):
exceptions_by_model_id[_model_id] = is_healthy
cleaned["exception_status"] = getattr(is_healthy, "status_code", 500)
cleaned["exception_status"] = getattr(
is_healthy, "status_code", 500
)
unhealthy_endpoints.append(cleaned)
return healthy_endpoints, unhealthy_endpoints, exceptions_by_model_id
@ -301,6 +307,8 @@ def _update_litellm_params_for_health_check(
_health_check_max_tokens = model_info.get("health_check_max_tokens", None)
if _health_check_max_tokens is not None:
litellm_params["max_tokens"] = _health_check_max_tokens
elif BACKGROUND_HEALTH_CHECK_MAX_TOKENS is not None:
litellm_params["max_tokens"] = BACKGROUND_HEALTH_CHECK_MAX_TOKENS
elif "*" not in (
model_info.get("health_check_model") or litellm_params.get("model") or ""
):

View file

@ -289,6 +289,7 @@ model LiteLLM_MCPServerTable {
server_name String?
alias String?
description String?
instructions String?
url String?
spec_path String?
transport String @default("sse")

View file

@ -27,20 +27,21 @@ class MCPServer(BaseModel):
spec_path: Optional[str] = None
auth_type: Optional[MCPAuthType] = None
authentication_token: Optional[str] = None
instructions: Optional[str] = None
mcp_info: Optional[MCPInfo] = None
extra_headers: Optional[
List[str]
] = None # allow admin to specify which headers to forward from client to the MCP server
extra_headers: Optional[List[str]] = (
None # allow admin to specify which headers to forward from client to the MCP server
)
allowed_tools: Optional[List[str]] = None
disallowed_tools: Optional[List[str]] = None
tool_name_to_display_name: Optional[Dict[str, str]] = None
tool_name_to_description: Optional[Dict[str, str]] = None
allowed_params: Optional[
Dict[str, List[str]]
] = None # map of tool names to allowed parameter lists
static_headers: Optional[
Dict[str, str]
] = None # static headers to forward to the MCP server
allowed_params: Optional[Dict[str, List[str]]] = (
None # map of tool names to allowed parameter lists
)
static_headers: Optional[Dict[str, str]] = (
None # static headers to forward to the MCP server
)
# OAuth-specific fields
client_id: Optional[str] = None
client_secret: Optional[str] = None

View file

@ -289,6 +289,7 @@ model LiteLLM_MCPServerTable {
server_name String?
alias String?
description String?
instructions String?
url String?
spec_path String?
transport String @default("sse")

View file

@ -131,6 +131,55 @@ def test_get_cache_key_text_completion():
assert cache_key_2 == cache_key_3
def test_get_cache_key_responses_api():
"""
Regression test: two /v1/responses calls that differ only in
`instructions` (or any Responses-API-only param) must produce
different cache keys. Mirrors the chat / embedding / text-completion
cache-key tests above.
"""
cache = Cache()
base_kwargs = {
"model": "openai/gpt-4.1",
"input": [{"role": "user", "content": "what is the weather"}],
"temperature": 0.3,
}
kwargs_a = {
**base_kwargs,
"instructions": "summarize the weather on 10th May",
}
kwargs_b = {
**base_kwargs,
"instructions": "summarize the weather on 7th May",
}
key_a = cache.get_cache_key(**kwargs_a)
key_b = cache.get_cache_key(**kwargs_b)
assert isinstance(key_a, str) and len(key_a) > 0
assert key_a != key_b, "instructions must be part of the Responses API cache key"
# Sanity: identical payloads must still collide (cache hits still work)
key_a_again = cache.get_cache_key(**kwargs_a)
assert key_a == key_a_again
# Spot-check a handful of other Responses-only params individually.
for param, value_x, value_y in [
("previous_response_id", "resp_aaa", "resp_bbb"),
("reasoning", {"effort": "low"}, {"effort": "high"}),
("include", ["reasoning.encrypted_content"], []),
("max_output_tokens", 100, 500),
("background", True, False),
]:
kx = {**base_kwargs, param: value_x}
ky = {**base_kwargs, param: value_y}
assert cache.get_cache_key(**kx) != cache.get_cache_key(
**ky
), f"Responses-API param `{param}` is not part of the cache key"
def test_get_hashed_cache_key():
cache = Cache()
cache_key = "model:gpt-3.5-turbo,messages:Hello world"

View file

@ -417,17 +417,22 @@ async def test_streamable_http_mcp_handler_mock():
# Mock extract_mcp_auth_context to bypass auth checks in the handler
mock_auth_context = (None, None, None, {}, {}, {})
with patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
), patch(
"litellm.proxy._experimental.mcp_server.server.session_manager",
mock_session_manager,
), patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
AsyncMock(return_value=mock_auth_context),
), patch(
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
with (
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
),
patch(
"litellm.proxy._experimental.mcp_server.server.session_manager",
mock_session_manager,
),
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
AsyncMock(return_value=mock_auth_context),
),
patch(
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
),
):
from litellm.proxy._experimental.mcp_server.server import (
handle_streamable_http_mcp,
@ -471,17 +476,22 @@ async def test_sse_mcp_handler_mock():
[],
)
with patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
), patch(
"litellm.proxy._experimental.mcp_server.server.sse_session_manager",
mock_sse_session_manager,
), patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new=AsyncMock(return_value=mock_auth_result),
), patch(
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
with (
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
),
patch(
"litellm.proxy._experimental.mcp_server.server.sse_session_manager",
mock_sse_session_manager,
),
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new=AsyncMock(return_value=mock_auth_result),
),
patch(
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
),
):
from litellm.proxy._experimental.mcp_server.server import handle_sse_mcp
@ -833,7 +843,9 @@ async def test_get_tools_from_mcp_servers():
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["server1_id", "server2_id"]
)
mock_manager.get_mcp_server_by_id = lambda server_id: mock_server_1 if server_id == "server1_id" else mock_server_2
mock_manager.get_mcp_server_by_id = lambda server_id: (
mock_server_1 if server_id == "server1_id" else mock_server_2
)
mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1])
# Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test)
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(
@ -859,7 +871,10 @@ async def test_get_tools_from_mcp_servers():
mock_manager_2.get_allowed_mcp_servers = AsyncMock(
return_value=["server1_id", "server2_id"]
)
mock_manager_2.get_mcp_server_by_id = lambda server_id: mock_server_1 if server_id == "server1_id" else mock_server_2
mock_manager_2.get_mcp_server_by_id = lambda server_id: (
mock_server_1 if server_id == "server1_id" else mock_server_2
)
async def mock_get_tools_side_effect(
server,
mcp_auth_header=None,
@ -900,7 +915,11 @@ async def test_get_tools_from_mcp_servers():
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["server1_id", "server2_id", "server3_id"]
)
mock_manager.get_mcp_server_by_id = lambda server_id: mock_server_1 if server_id == "server1_id" else (mock_server_2 if server_id == "server2_id" else mock_server_3)
mock_manager.get_mcp_server_by_id = lambda server_id: (
mock_server_1
if server_id == "server1_id"
else (mock_server_2 if server_id == "server2_id" else mock_server_3)
)
mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1])
# Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test)
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(
@ -1050,15 +1069,15 @@ async def test_mcp_server_manager_access_groups_from_config():
# Should find config_server for group-a, both for group-b, other_server for group-c
import asyncio
server_ids_a = await MCPRequestHandler._get_mcp_servers_from_access_groups([
"group-a"
])
server_ids_b = await MCPRequestHandler._get_mcp_servers_from_access_groups([
"group-b"
])
server_ids_c = await MCPRequestHandler._get_mcp_servers_from_access_groups([
"group-c"
])
server_ids_a = await MCPRequestHandler._get_mcp_servers_from_access_groups(
["group-a"]
)
server_ids_b = await MCPRequestHandler._get_mcp_servers_from_access_groups(
["group-b"]
)
server_ids_c = await MCPRequestHandler._get_mcp_servers_from_access_groups(
["group-c"]
)
assert any(config_server.server_id == sid for sid in server_ids_a)
assert set(server_ids_b) == set(
[
@ -1474,6 +1493,7 @@ async def test_add_update_server_with_alias():
mock_mcp_server.byok_api_key_help_url = None
mock_mcp_server.created_at = None
mock_mcp_server.updated_at = None
mock_mcp_server.instructions = None
# Add server to manager
await test_manager.add_server(mock_mcp_server)
@ -1530,6 +1550,7 @@ async def test_add_update_server_without_alias():
mock_mcp_server.byok_api_key_help_url = None
mock_mcp_server.created_at = None
mock_mcp_server.updated_at = None
mock_mcp_server.instructions = None
# Add server to manager
await test_manager.add_server(mock_mcp_server)
@ -1587,7 +1608,7 @@ async def test_add_update_server_fallback_to_server_id():
mock_mcp_server.byok_api_key_help_url = None
mock_mcp_server.created_at = None
mock_mcp_server.updated_at = None
mock_mcp_server.instructions = None
# Add server to manager
await test_manager.add_server(mock_mcp_server)
@ -2151,8 +2172,12 @@ async def test_list_tool_rest_api_all_servers_with_auth():
for call_args in mock_get_tools.call_args_list
}
assert server_auth_map.get(mock_zapier_server) == "Bearer zapier_token"
assert server_auth_map.get(mock_slack_server) == "Bearer slack_token"
assert (
server_auth_map.get(mock_zapier_server) == "Bearer zapier_token"
)
assert (
server_auth_map.get(mock_slack_server) == "Bearer slack_token"
)
@pytest.mark.asyncio
@ -2690,26 +2715,33 @@ async def test_call_mcp_tool_uses_manager_permission_lookup():
expected_response = [TextContent(type="text", text="ok")]
with patch.object(
global_mcp_server_manager,
"get_allowed_mcp_servers",
new_callable=AsyncMock,
) as mock_get_allowed, patch.object(
global_mcp_server_manager,
"get_mcp_server_by_id",
return_value=mock_server,
), patch.object(
global_mcp_server_manager,
"_get_mcp_server_from_tool_name",
return_value=mock_server,
) as mock_get_server, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry"
) as mock_tool_registry, patch(
"litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool",
new_callable=AsyncMock,
) as mock_handle_managed, patch(
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed",
return_value=True,
with (
patch.object(
global_mcp_server_manager,
"get_allowed_mcp_servers",
new_callable=AsyncMock,
) as mock_get_allowed,
patch.object(
global_mcp_server_manager,
"get_mcp_server_by_id",
return_value=mock_server,
),
patch.object(
global_mcp_server_manager,
"_get_mcp_server_from_tool_name",
return_value=mock_server,
) as mock_get_server,
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry"
) as mock_tool_registry,
patch(
"litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool",
new_callable=AsyncMock,
) as mock_handle_managed,
patch(
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed",
return_value=True,
),
):
mock_get_allowed.return_value = [mock_server.server_id]
mock_tool_registry.get_tool.return_value = None
@ -2759,27 +2791,34 @@ async def test_call_mcp_tool_resolves_unprefixed_tool_name_and_checks_permission
expected_response = [TextContent(type="text", text="ok")]
with patch.object(
global_mcp_server_manager,
"get_allowed_mcp_servers",
new_callable=AsyncMock,
) as mock_get_allowed, patch.object(
global_mcp_server_manager,
"get_mcp_server_by_id",
return_value=mock_server,
), patch.object(
global_mcp_server_manager,
"_get_mcp_server_from_tool_name",
return_value=mock_server,
) as mock_get_server, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry"
) as mock_tool_registry, patch(
"litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool",
new_callable=AsyncMock,
) as mock_handle_managed, patch(
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed",
return_value=True,
) as mock_is_allowed:
with (
patch.object(
global_mcp_server_manager,
"get_allowed_mcp_servers",
new_callable=AsyncMock,
) as mock_get_allowed,
patch.object(
global_mcp_server_manager,
"get_mcp_server_by_id",
return_value=mock_server,
),
patch.object(
global_mcp_server_manager,
"_get_mcp_server_from_tool_name",
return_value=mock_server,
) as mock_get_server,
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry"
) as mock_tool_registry,
patch(
"litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool",
new_callable=AsyncMock,
) as mock_handle_managed,
patch(
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed",
return_value=True,
) as mock_is_allowed,
):
mock_get_allowed.return_value = [mock_server.server_id]
mock_tool_registry.get_tool.return_value = None
mock_handle_managed.return_value = expected_response

View file

@ -40,6 +40,7 @@ class TestMCPClient:
with pytest.raises(
ValueError, match="stdio_config is required for stdio transport"
):
async def _noop(session):
return None
@ -251,11 +252,11 @@ class TestMCPClient:
server_url="http://example.com/sse",
transport_type="sse",
auth_type=MCPAuth.token,
auth_value="my-secret-token"
auth_value="my-secret-token",
)
headers = client._get_auth_headers()
assert "Authorization" in headers
assert headers["Authorization"] == "token my-secret-token"
@ -266,27 +267,27 @@ class TestMCPClient:
server_url="http://example.com/sse",
transport_type="sse",
auth_type=MCPAuth.bearer_token,
auth_value="bearer-token"
auth_value="bearer-token",
)
headers = client._get_auth_headers()
assert headers["Authorization"] == "Bearer bearer-token"
# Test API key
client = MCPClient(
server_url="http://example.com/sse",
transport_type="sse",
auth_type=MCPAuth.api_key,
auth_value="api-key"
auth_value="api-key",
)
headers = client._get_auth_headers()
assert headers["X-API-Key"] == "api-key"
# Test basic auth (gets base64 encoded)
client = MCPClient(
server_url="http://example.com/sse",
transport_type="sse",
auth_type=MCPAuth.basic,
auth_value="user:pass"
auth_value="user:pass",
)
headers = client._get_auth_headers()
assert headers["Authorization"].startswith("Basic ")
@ -298,11 +299,11 @@ class TestMCPClient:
transport_type="sse",
auth_type=MCPAuth.token,
auth_value="my-token",
extra_headers={"X-Custom-Header": "custom-value"}
extra_headers={"X-Custom-Header": "custom-value"},
)
headers = client._get_auth_headers()
assert headers["Authorization"] == "token my-token"
assert headers["X-Custom-Header"] == "custom-value"
@ -312,5 +313,80 @@ class TestMCPClient:
assert MCPAuth.token.value == "token"
# ---------------------------------------------------------------------------
# _last_initialize_instructions capture
# ---------------------------------------------------------------------------
class TestMCPClientInstructionsCapture:
"""Tests for _last_initialize_instructions capture during session init."""
def test_initial_value_is_none(self):
"""Fresh client has no cached instructions."""
client = MCPClient(
server_url="http://example.com/mcp",
transport_type="http",
)
assert client._last_initialize_instructions is None
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_captures_instructions_from_initialize(self, mock_session_cls):
"""Instructions from upstream initialize() are captured and stripped."""
client = MCPClient(
server_url="http://example.com/mcp",
transport_type="http",
)
mock_session = AsyncMock()
init_result = MagicMock()
init_result.instructions = " upstream says hello "
mock_session.initialize = AsyncMock(return_value=init_result)
session_ctx = MagicMock()
session_ctx.__aenter__ = AsyncMock(return_value=mock_session)
session_ctx.__aexit__ = AsyncMock(return_value=False)
mock_session_cls.return_value = session_ctx
transport_ctx = MagicMock()
transport_ctx.__aenter__ = AsyncMock(return_value=(MagicMock(), MagicMock()))
transport_ctx.__aexit__ = AsyncMock(return_value=False)
async def _op(session):
return "done"
await client._execute_session_operation(transport_ctx, _op)
assert client._last_initialize_instructions == "upstream says hello"
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_none_instructions_stays_none(self, mock_session_cls):
"""When upstream returns no instructions the field stays None."""
client = MCPClient(
server_url="http://example.com/mcp",
transport_type="http",
)
mock_session = AsyncMock()
init_result = MagicMock()
init_result.instructions = None
mock_session.initialize = AsyncMock(return_value=init_result)
session_ctx = MagicMock()
session_ctx.__aenter__ = AsyncMock(return_value=mock_session)
session_ctx.__aexit__ = AsyncMock(return_value=False)
mock_session_cls.return_value = session_ctx
transport_ctx = MagicMock()
transport_ctx.__aenter__ = AsyncMock(return_value=(MagicMock(), MagicMock()))
transport_ctx.__aexit__ = AsyncMock(return_value=False)
async def _op(session):
return "done"
await client._execute_session_operation(transport_ctx, _op)
assert client._last_initialize_instructions is None
if __name__ == "__main__":
pytest.main([__file__])

View file

@ -440,14 +440,18 @@ class TestProxyOAuthHeaderForwarding:
(b"content-type", b"application/json"),
]
)
# Should preserve OAuth even with flag=False
cleaned_without_flag = clean_headers(raw_headers, forward_llm_provider_auth_headers=False)
cleaned_without_flag = clean_headers(
raw_headers, forward_llm_provider_auth_headers=False
)
assert "authorization" in cleaned_without_flag
assert cleaned_without_flag["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
# Should also preserve OAuth with flag=True
cleaned_with_flag = clean_headers(raw_headers, forward_llm_provider_auth_headers=True)
cleaned_with_flag = clean_headers(
raw_headers, forward_llm_provider_auth_headers=True
)
assert "authorization" in cleaned_with_flag
assert cleaned_with_flag["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
@ -867,8 +871,6 @@ class TestValidateEnvironmentAuthToken:
assert "authorization" not in headers
class TestGetAuthToken:
"""Tests for AnthropicModelInfo.get_auth_token() static method."""
@ -1092,7 +1094,10 @@ class TestPassthroughAuthToken:
config = AnthropicMessagesConfig()
with mock_patch.dict(
"os.environ",
{"ANTHROPIC_API_KEY": FAKE_REGULAR_KEY, "ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN},
{
"ANTHROPIC_API_KEY": FAKE_REGULAR_KEY,
"ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN,
},
clear=True,
):
updated_headers, _ = config.validate_anthropic_messages_environment(
@ -1131,3 +1136,147 @@ class TestPassthroughAuthToken:
)
assert url == "https://custom.example.com/v1/messages"
class TestAnthropicThinkingSignatureSelfHeal:
"""Helpers for retrying after invalid encrypted thinking signatures."""
def test_is_anthropic_invalid_thinking_signature_error_positive(self):
from litellm.llms.anthropic.common_utils import (
is_anthropic_invalid_thinking_signature_error,
)
raw = (
'{"type":"error","error":{"type":"invalid_request_error",'
'"message":"messages.3.content.3: Invalid `signature` in `thinking` block"},'
'"request_id":"req_011Ca2EtQDxp7x6RGUY2jVn9"}'
)
assert is_anthropic_invalid_thinking_signature_error(raw) is True
def test_is_anthropic_invalid_thinking_signature_error_negative(self):
from litellm.llms.anthropic.common_utils import (
is_anthropic_invalid_thinking_signature_error,
)
assert is_anthropic_invalid_thinking_signature_error("") is False
assert (
is_anthropic_invalid_thinking_signature_error("rate limit exceeded")
is False
)
def test_strip_thinking_blocks_from_anthropic_messages(self):
from litellm.llms.anthropic.common_utils import (
strip_thinking_blocks_from_anthropic_messages,
)
messages = [
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": [
{"type": "thinking", "thinking": "plan", "signature": "sig"},
{"type": "text", "text": "hello"},
],
},
]
out = strip_thinking_blocks_from_anthropic_messages(messages)
assert len(out) == 2
assert out[0] == messages[0]
assert len(out[1]["content"]) == 1
assert out[1]["content"][0]["type"] == "text"
assert messages[1]["content"][0]["type"] == "thinking"
def test_strip_thinking_blocks_drops_message_when_only_thinking_blocks(self):
from litellm.llms.anthropic.common_utils import (
strip_thinking_blocks_from_anthropic_messages,
)
messages = [
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": [
{"type": "thinking", "thinking": "plan", "signature": "sig"},
],
},
]
out = strip_thinking_blocks_from_anthropic_messages(messages)
assert len(out) == 1
assert out[0]["role"] == "user"
def test_strip_thinking_blocks_from_anthropic_messages_request_dict(self):
from litellm.llms.anthropic.common_utils import (
strip_thinking_blocks_from_anthropic_messages_request_dict,
)
data = {
"model": "claude-sonnet-4-20250514",
"messages": [
{
"role": "assistant",
"content": [
{
"type": "thinking",
"thinking": "x",
"signature": "y",
},
],
}
],
"thinking": {"type": "enabled", "budget_tokens": 1024},
}
strip_thinking_blocks_from_anthropic_messages_request_dict(data)
assert "thinking" not in data
assert data["messages"] == []
def test_anthropic_messages_config_http_retry_helpers(self):
import httpx
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
config = AnthropicMessagesConfig()
assert config.max_retry_on_anthropic_messages_http_error == 2
req = httpx.Request("POST", "https://api.anthropic.com/v1/messages")
err_text = (
'{"type":"error","error":{"type":"invalid_request_error",'
'"message":"messages.3.content.3: Invalid `signature` in `thinking` block"},'
'"request_id":"req_011Ca2EtQDxp7x6RGUY2jVn9"}'
)
resp = httpx.Response(400, request=req, text=err_text)
err = httpx.HTTPStatusError("bad", request=req, response=resp)
assert config.should_retry_anthropic_messages_on_http_error(err, {}) is True
resp_bad = httpx.Response(400, request=req, text="rate limit exceeded")
err_bad = httpx.HTTPStatusError("bad", request=req, response=resp_bad)
assert (
config.should_retry_anthropic_messages_on_http_error(err_bad, {}) is False
)
resp_500 = httpx.Response(500, request=req, text=err_text)
err_500 = httpx.HTTPStatusError("bad", request=req, response=resp_500)
assert (
config.should_retry_anthropic_messages_on_http_error(err_500, {}) is False
)
data = {
"model": "claude-sonnet-4-20250514",
"messages": [
{
"role": "assistant",
"content": [
{
"type": "thinking",
"thinking": "x",
"signature": "y",
},
],
}
],
"thinking": {"type": "enabled", "budget_tokens": 1024},
}
config.transform_anthropic_messages_request_on_http_error(err, data)
assert "thinking" not in data
assert data["messages"] == []

View file

@ -0,0 +1,97 @@
import json
import os
import sys
from unittest.mock import MagicMock
import httpx
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.azure.passthrough.transformation import AzurePassthroughConfig
from litellm.types.utils import ModelResponse
def _azure_chat_completion_body():
return {
"id": "chatcmpl-abc123",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4.1-mini-2025-04-14",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "Hello! How can I assist you today?",
},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 10,
"completion_tokens": 8,
"total_tokens": 18,
},
}
def _make_httpx_response(body: dict) -> httpx.Response:
return httpx.Response(
status_code=200,
headers={"content-type": "application/json"},
content=json.dumps(body).encode("utf-8"),
request=httpx.Request(
"POST",
"https://example.openai.azure.com/openai/deployments/gpt-4.1-mini/chat/completions",
),
)
def test_azure_passthrough_logging_non_streaming_response_chat_completions():
"""
Returns a populated ModelResponse (with usage + content) for a chat/completions
endpoint. This is what _success_handler_helper_fn needs to build
standard_logging_object — without it, Datadog/cost-tracking/router-success all
raise on every Azure passthrough request.
"""
config = AzurePassthroughConfig()
logging_obj = MagicMock()
result = config.logging_non_streaming_response(
model="gpt-4.1-mini",
custom_llm_provider="azure",
httpx_response=_make_httpx_response(_azure_chat_completion_body()),
request_data={
"model": "gpt-4.1-mini",
"messages": [{"role": "user", "content": "hi"}],
},
logging_obj=logging_obj,
endpoint="openai/deployments/gpt-4.1-mini/chat/completions",
)
assert isinstance(result, ModelResponse)
assert result.choices[0].message.content == "Hello! How can I assist you today?"
assert result.usage.prompt_tokens == 10
assert result.usage.completion_tokens == 8
assert result.usage.total_tokens == 18
def test_azure_passthrough_logging_non_streaming_response_unknown_endpoint_returns_none():
"""
Endpoints other than chat/completions (responses, messages, images) fall
through to None — matches base-class behavior and Bedrock's "unknown
endpoint" handling. Not a regression; just scoping.
"""
config = AzurePassthroughConfig()
logging_obj = MagicMock()
result = config.logging_non_streaming_response(
model="gpt-4.1-mini",
custom_llm_provider="azure",
httpx_response=_make_httpx_response(_azure_chat_completion_body()),
request_data={},
logging_obj=logging_obj,
endpoint="openai/responses",
)
assert result is None

View file

@ -5,7 +5,12 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from mcp import ReadResourceResult, Resource
from mcp.types import BlobResourceContents, Prompt, ResourceTemplate, TextResourceContents
from mcp.types import (
BlobResourceContents,
Prompt,
ResourceTemplate,
TextResourceContents,
)
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
@ -157,15 +162,19 @@ async def test_get_prompts_from_mcp_servers_success():
server_b.auth_type = None
server_b.extra_headers = None
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server_a, server_b]),
) as mock_allowed, patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=(None, None),
) as mock_headers, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager:
with (
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server_a, server_b]),
) as mock_allowed,
patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=(None, None),
) as mock_headers,
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager,
):
mock_manager.get_prompts_from_server = AsyncMock(
side_effect=[
[Prompt(name="hello", description="hi")],
@ -213,15 +222,19 @@ async def test_get_resources_from_mcp_servers_success():
server_b.auth_type = None
server_b.extra_headers = None
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server_a, server_b]),
) as mock_allowed, patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=(None, None),
) as mock_headers, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager:
with (
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server_a, server_b]),
) as mock_allowed,
patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=(None, None),
) as mock_headers,
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager,
):
mock_manager.get_resources_from_server = AsyncMock(
side_effect=[
[
@ -274,15 +287,19 @@ async def test_get_resource_templates_from_mcp_servers_success():
server.auth_type = None
server.extra_headers = None
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server]),
) as mock_allowed, patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=(None, None),
) as mock_headers, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager:
with (
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server]),
) as mock_allowed,
patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=(None, None),
) as mock_headers,
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager,
):
mock_manager.get_resource_templates_from_server = AsyncMock(
return_value=[
ResourceTemplate(
@ -320,15 +337,19 @@ async def test_mcp_get_prompt_success():
prompt_result = MagicMock(name="prompt_result")
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server]),
) as mock_allowed, patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=({"Authorization": "token"}, {"X-Test": "1"}),
) as mock_headers, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager:
with (
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server]),
) as mock_allowed,
patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=({"Authorization": "token"}, {"X-Test": "1"}),
) as mock_headers,
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager,
):
mock_manager.get_prompt_from_server = AsyncMock(return_value=prompt_result)
result = await mcp_get_prompt(
@ -378,15 +399,19 @@ async def test_mcp_read_resource_success():
]
)
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server]),
) as mock_allowed, patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=({"Authorization": "token"}, {"X-Test": "1"}),
) as mock_headers, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager:
with (
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[server]),
) as mock_allowed,
patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=({"Authorization": "token"}, {"X-Test": "1"}),
) as mock_headers,
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager,
):
mock_manager.read_resource_from_server = AsyncMock(return_value=read_result)
result = await mcp_read_resource(
@ -591,7 +616,10 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
working_server if server_id == "working_server" else failing_server
)
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
server_ids,
0,
)
async def mock_get_tools_from_server(
server,
@ -693,7 +721,10 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing():
failing_server1 if server_id == "failing_server1" else failing_server2
)
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
server_ids,
0,
)
async def mock_get_tools_from_server(
server,
@ -830,12 +861,14 @@ async def test_concurrent_initialize_session_managers():
mcp_server._sse_session_manager_cm = None
# Mock the session managers to avoid actual MCP initialization
with patch(
"litellm.proxy._experimental.mcp_server.server.session_manager"
) as mock_session_manager, patch(
"litellm.proxy._experimental.mcp_server.server.sse_session_manager"
) as mock_sse_session_manager, patch(
"litellm.proxy._experimental.mcp_server.server.verbose_logger"
with (
patch(
"litellm.proxy._experimental.mcp_server.server.session_manager"
) as mock_session_manager,
patch(
"litellm.proxy._experimental.mcp_server.server.sse_session_manager"
) as mock_sse_session_manager,
patch("litellm.proxy._experimental.mcp_server.server.verbose_logger"),
):
# Mock the run() method to return a mock context manager
mock_cm = AsyncMock()
@ -961,15 +994,19 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name():
return_value=[specific_server.server_id, other_server.server_id]
)
with patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers",
mock_get_allowed,
), patch(
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler._get_mcp_servers_from_access_groups",
mock_db_lookup,
), patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server",
mock_get_tools_spy,
with (
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers",
mock_get_allowed,
),
patch(
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler._get_mcp_servers_from_access_groups",
mock_db_lookup,
),
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server",
mock_get_tools_spy,
),
):
mcp_servers_from_path = _get_mcp_servers_in_path(test_path)
@ -1062,17 +1099,21 @@ async def test_oauth2_headers_passed_to_mcp_client():
async def mock_fetch_tools_with_timeout(client, server_name):
return [] # Return empty list of tools
with patch.object(
global_mcp_server_manager,
"_create_mcp_client",
side_effect=mock_create_mcp_client,
) as mock_create_client, patch.object(
global_mcp_server_manager,
"_fetch_tools_with_timeout",
side_effect=mock_fetch_tools_with_timeout,
), patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[oauth2_server]),
with (
patch.object(
global_mcp_server_manager,
"_create_mcp_client",
side_effect=mock_create_mcp_client,
) as mock_create_client,
patch.object(
global_mcp_server_manager,
"_fetch_tools_with_timeout",
side_effect=mock_fetch_tools_with_timeout,
),
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
AsyncMock(return_value=[oauth2_server]),
),
):
# Call _get_tools_from_mcp_servers which should eventually call _create_mcp_client
await _get_tools_from_mcp_servers(
@ -1138,7 +1179,10 @@ async def test_list_tools_single_server_unprefixed_names():
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
mock_manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
server_ids,
0,
)
async def mock_get_tools_from_server(
server,
@ -1216,7 +1260,10 @@ async def test_list_tools_multiple_servers_prefixed_names():
server1 if server_id == "server1" else server2
)
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
server_ids,
0,
)
async def mock_get_tools_from_server(
server,
@ -1270,12 +1317,15 @@ async def test_mcp_manager_allows_public_servers_without_permissions():
)
manager.registry = {public_server.server_id: public_server}
with patch(
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view",
return_value=False,
), patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
AsyncMock(return_value=[]),
with (
patch(
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view",
return_value=False,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
AsyncMock(return_value=[]),
),
):
allowed = await manager.get_allowed_mcp_servers(UserAPIKeyAuth())
@ -1302,12 +1352,15 @@ async def test_mcp_manager_returns_public_when_permission_lookup_fails():
)
manager.registry = {public_server.server_id: public_server}
with patch(
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view",
return_value=False,
), patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
AsyncMock(side_effect=Exception("boom")),
with (
patch(
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view",
return_value=False,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
AsyncMock(side_effect=Exception("boom")),
),
):
allowed = await manager.get_allowed_mcp_servers(UserAPIKeyAuth())
@ -1342,12 +1395,15 @@ async def test_mcp_manager_merges_public_and_restricted_servers():
scoped_server.server_id: scoped_server,
}
with patch(
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view",
return_value=False,
), patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
AsyncMock(return_value=["restricted"]),
with (
patch(
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view",
return_value=False,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
AsyncMock(return_value=["restricted"]),
),
):
allowed = await manager.get_allowed_mcp_servers(UserAPIKeyAuth())
@ -1399,12 +1455,15 @@ async def test_call_mcp_tool_user_unauthorized_access():
return another_server_obj
return None
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers",
AsyncMock(return_value=["allowed_server", "another_server"]),
), patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_id",
side_effect=mock_get_server_by_id,
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers",
AsyncMock(return_value=["allowed_server", "another_server"]),
),
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_id",
side_effect=mock_get_server_by_id,
),
):
# Try to call a tool from "restricted_server" - should raise HTTPException with 403 status
with pytest.raises(HTTPException) as exc_info:
@ -1467,7 +1526,10 @@ async def test_list_tools_filters_by_key_team_permissions():
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
mock_manager.get_mcp_server_by_id = lambda server_id: server
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
server_ids,
0,
)
async def mock_get_tools_from_server(
server,
@ -1573,7 +1635,10 @@ async def test_list_tools_with_team_tool_permissions_inheritance():
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
mock_manager.get_mcp_server_by_id = lambda server_id: server
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
server_ids,
0,
)
async def mock_get_tools_from_server(
server,
@ -1665,7 +1730,10 @@ async def test_list_tools_with_no_tool_permissions_shows_all():
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
mock_manager.get_mcp_server_by_id = lambda server_id: server
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
server_ids,
0,
)
async def mock_get_tools_from_server(
server,
@ -1760,7 +1828,10 @@ async def test_list_tools_strips_prefix_when_matching_permissions():
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["gitmcp_server"])
mock_manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
server_ids,
0,
)
async def mock_get_tools_from_server(
server,
@ -2002,12 +2073,15 @@ class TestMCPServerManagerReload:
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(
return_value=[db_row]
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma,
), patch.object(
manager, "build_mcp_server_from_table", AsyncMock()
) as mock_build:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma,
),
patch.object(
manager, "build_mcp_server_from_table", AsyncMock()
) as mock_build,
):
await manager.reload_servers_from_database()
mock_build.assert_not_awaited()
@ -2045,14 +2119,17 @@ class TestMCPServerManagerReload:
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(
return_value=[db_row]
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma,
), patch.object(
manager,
"build_mcp_server_from_table",
AsyncMock(return_value=rebuilt_server),
) as mock_build:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma,
),
patch.object(
manager,
"build_mcp_server_from_table",
AsyncMock(return_value=rebuilt_server),
) as mock_build,
):
await manager.reload_servers_from_database()
mock_build.assert_awaited_once_with(db_row)
@ -2090,26 +2167,32 @@ async def test_call_mcp_tool_logs_failure_via_post_call_failure_hook():
user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
with patch.object(
global_mcp_server_manager,
"get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[mock_server.server_id],
), patch.object(
global_mcp_server_manager,
"get_mcp_server_by_id",
return_value=mock_server,
), patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names",
new_callable=AsyncMock,
return_value=[mock_server],
), patch(
"litellm.proxy._experimental.mcp_server.server.execute_mcp_tool",
new_callable=AsyncMock,
side_effect=Exception("boom"),
), patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
proxy_logging_mock,
with (
patch.object(
global_mcp_server_manager,
"get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[mock_server.server_id],
),
patch.object(
global_mcp_server_manager,
"get_mcp_server_by_id",
return_value=mock_server,
),
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names",
new_callable=AsyncMock,
return_value=[mock_server],
),
patch(
"litellm.proxy._experimental.mcp_server.server.execute_mcp_tool",
new_callable=AsyncMock,
side_effect=Exception("boom"),
),
patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
proxy_logging_mock,
),
):
with pytest.raises(Exception):
await call_mcp_tool(
@ -2157,23 +2240,30 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}}
dummy_logging_obj.async_success_handler = AsyncMock()
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
new=AsyncMock(return_value=[server_a]),
), patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=(None, None),
), patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager, patch(
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools",
side_effect=lambda tools, _server: tools,
), patch(
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions",
new=AsyncMock(side_effect=lambda tools, **_: tools),
), patch(
"litellm.proxy._experimental.mcp_server.server.function_setup",
return_value=(dummy_logging_obj, None),
with (
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
new=AsyncMock(return_value=[server_a]),
),
patch(
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
return_value=(None, None),
),
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager,
patch(
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools",
side_effect=lambda tools, _server: tools,
),
patch(
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions",
new=AsyncMock(side_effect=lambda tools, **_: tools),
),
patch(
"litellm.proxy._experimental.mcp_server.server.function_setup",
return_value=(dummy_logging_obj, None),
),
):
mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1])
@ -2188,7 +2278,9 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
assert tools == [tool_1]
dummy_logging_obj.async_success_handler.assert_awaited_once()
assert dummy_logging_obj.async_success_handler.await_args.kwargs["result"] == [tool_1]
assert dummy_logging_obj.async_success_handler.await_args.kwargs["result"] == [
tool_1
]
spend_meta = dummy_logging_obj.model_call_details["metadata"]["spend_logs_metadata"]
assert spend_meta["tool_count_total"] == 1
@ -2381,26 +2473,34 @@ async def test_get_tools_from_mcp_servers_injects_stored_oauth2_token():
oauth2_server.extra_headers = None
# Simulate the DB returning a valid credential for this user+server
prefetched_creds = {SERVER_ID: {"access_token": STORED_TOKEN, "server_id": SERVER_ID}}
prefetched_creds = {
SERVER_ID: {"access_token": STORED_TOKEN, "server_id": SERVER_ID}
}
tool_1 = MagicMock()
tool_1.name = "atlassian_test-search"
with patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
new=AsyncMock(return_value=[oauth2_server]),
), patch(
# Patch the bulk prefetch so no real DB connection is needed
"litellm.proxy._experimental.mcp_server.server._prefetch_oauth_creds_for_user",
new=AsyncMock(return_value=prefetched_creds),
) as mock_prefetch, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager, patch(
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools",
side_effect=lambda tools, _server: tools,
), patch(
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions",
new=AsyncMock(side_effect=lambda tools, **_: tools),
with (
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
new=AsyncMock(return_value=[oauth2_server]),
),
patch(
# Patch the bulk prefetch so no real DB connection is needed
"litellm.proxy._experimental.mcp_server.server._prefetch_oauth_creds_for_user",
new=AsyncMock(return_value=prefetched_creds),
) as mock_prefetch,
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
) as mock_manager,
patch(
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools",
side_effect=lambda tools, _server: tools,
),
patch(
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions",
new=AsyncMock(side_effect=lambda tools, **_: tools),
),
):
mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1])
@ -2421,3 +2521,201 @@ async def test_get_tools_from_mcp_servers_injects_stored_oauth2_token():
assert call_kwargs["extra_headers"] == {"Authorization": f"Bearer {STORED_TOKEN}"}
assert tools == [tool_1]
# ---------------------------------------------------------------------------
# _merge_gateway_initialize_instructions + ContextVar / InitializationOptions
# ---------------------------------------------------------------------------
def _make_instruction_server(
server_id="s1",
name="s1",
*,
alias=None,
server_name=None,
instructions=None,
spec_path=None,
url="https://example.com",
):
return MCPServer(
server_id=server_id,
name=name,
alias=alias,
server_name=server_name,
url=url,
transport=MCPTransport.http,
instructions=instructions,
spec_path=spec_path,
)
class TestMergeGatewayInitializeInstructions:
"""Tests for _merge_gateway_initialize_instructions."""
def _merge(self, servers):
try:
from litellm.proxy._experimental.mcp_server.server import (
_merge_gateway_initialize_instructions,
)
except ImportError:
pytest.skip("MCP server not available")
return _merge_gateway_initialize_instructions(servers)
def test_empty_server_list_returns_none(self):
"""No servers yields no instructions."""
assert self._merge([]) is None
def test_single_server_yaml_instructions(self):
"""A single server with YAML instructions returns them verbatim."""
s = _make_instruction_server(instructions="Use add() for sums.")
assert self._merge([s]) == "Use add() for sums."
def test_yaml_instructions_strips_whitespace(self):
"""Leading/trailing whitespace is stripped."""
s = _make_instruction_server(instructions=" padded \n")
assert self._merge([s]) == "padded"
def test_yaml_override_beats_upstream_cache(self):
"""YAML/DB instructions take precedence over upstream cache."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
global_mcp_server_manager._upstream_initialize_instructions_by_server_id[
"s1"
] = "upstream"
try:
s = _make_instruction_server(instructions="yaml wins")
assert self._merge([s]) == "yaml wins"
finally:
global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop(
"s1", None
)
def test_upstream_cache_used_when_no_yaml(self):
"""Upstream cached instructions are used when no YAML override is set."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
global_mcp_server_manager._upstream_initialize_instructions_by_server_id[
"s1"
] = "from upstream"
try:
s = _make_instruction_server(instructions=None)
assert self._merge([s]) == "from upstream"
finally:
global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop(
"s1", None
)
def test_spec_path_servers_skipped(self):
"""OpenAPI (spec_path) servers do not contribute instructions."""
s = _make_instruction_server(spec_path="/openapi.json", url=None)
assert self._merge([s]) is None
def test_no_instructions_no_cache_returns_none(self):
"""Server with no instructions and no cache yields None."""
s = _make_instruction_server()
assert self._merge([s]) is None
def test_multiple_servers_merged_with_labels(self):
"""Multiple servers get label-prefixed and separator-joined."""
s1 = _make_instruction_server(
server_id="a", name="a", alias="Alpha", instructions="instr A"
)
s2 = _make_instruction_server(
server_id="b", name="b", alias="Beta", instructions="instr B"
)
result = self._merge([s1, s2])
assert result is not None
assert "[Alpha]" in result and "[Beta]" in result
assert "instr A" in result and "instr B" in result
assert "---" in result
def test_single_server_no_label_wrapping(self):
"""A single server's instructions are not wrapped with a label."""
s = _make_instruction_server(alias="MyServer", instructions="single")
result = self._merge([s])
assert result == "single"
assert "[MyServer]" not in result
def test_mixed_yaml_cache_specpath(self):
"""YAML, upstream-cache, and spec_path servers are handled correctly together."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
global_mcp_server_manager._upstream_initialize_instructions_by_server_id[
"c"
] = "cached C"
try:
s_yaml = _make_instruction_server(
server_id="a", name="a", alias="A", instructions="yaml A"
)
s_spec = _make_instruction_server(
server_id="b", name="b", alias="B", spec_path="/spec.json", url=None
)
s_cached = _make_instruction_server(server_id="c", name="c", alias="C")
result = self._merge([s_yaml, s_spec, s_cached])
assert "yaml A" in result
assert "cached C" in result
assert "[B]" not in result
finally:
global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop(
"c", None
)
class TestGatewayCreateInitializationOptions:
"""Tests for the patched server.create_initialization_options via ContextVar."""
def test_no_contextvar_returns_default_options(self):
"""When ContextVar is None, instructions are absent."""
try:
from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_gateway_initialize_instructions,
)
from litellm.proxy._experimental.mcp_server.server import server
except ImportError:
pytest.skip("MCP server not available")
tok = _mcp_gateway_initialize_instructions.set(None)
try:
opts = server.create_initialization_options()
assert getattr(opts, "instructions", None) is None
finally:
_mcp_gateway_initialize_instructions.reset(tok)
def test_contextvar_set_injects_instructions(self):
"""When ContextVar has a value, it appears in InitializationOptions."""
try:
from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_gateway_initialize_instructions,
)
from litellm.proxy._experimental.mcp_server.server import server
except ImportError:
pytest.skip("MCP server not available")
tok = _mcp_gateway_initialize_instructions.set("hello from merge")
try:
opts = server.create_initialization_options()
assert opts.instructions == "hello from merge"
finally:
_mcp_gateway_initialize_instructions.reset(tok)
def test_contextvar_reset_removes_instructions(self):
"""After resetting the ContextVar, instructions disappear."""
try:
from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_gateway_initialize_instructions,
)
from litellm.proxy._experimental.mcp_server.server import server
except ImportError:
pytest.skip("MCP server not available")
tok = _mcp_gateway_initialize_instructions.set("temporary")
_mcp_gateway_initialize_instructions.reset(tok)
opts = server.create_initialization_options()
assert getattr(opts, "instructions", None) is None

View file

@ -43,10 +43,10 @@ def _reload_mcp_manager_module():
# After reload, server.py still holds a stale reference to the old
# global_mcp_server_manager. Update it so tests that exercise server.py
# functions (e.g. _get_tools_from_mcp_servers) use the fresh instance.
server_module = sys.modules.get(
"litellm.proxy._experimental.mcp_server.server"
)
if server_module is not None and hasattr(server_module, "global_mcp_server_manager"):
server_module = sys.modules.get("litellm.proxy._experimental.mcp_server.server")
if server_module is not None and hasattr(
server_module, "global_mcp_server_manager"
):
server_module.global_mcp_server_manager = reloaded.global_mcp_server_manager
return reloaded
@ -223,9 +223,7 @@ class TestMCPServerManager:
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
await manager.load_servers_from_config(config)
assert any(
"invalid alias 'bad/name'" in message for message in caplog.messages
)
assert any("invalid alias 'bad/name'" in message for message in caplog.messages)
@pytest.mark.asyncio
async def test_load_servers_from_config_accepts_valid_alias(self, caplog):
@ -492,7 +490,12 @@ class TestMCPServerManager:
mock_client = AsyncMock()
mock_client.list_prompts = AsyncMock(return_value=[mock_prompt])
with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client):
with patch.object(
manager,
"_create_mcp_client",
new_callable=AsyncMock,
return_value=mock_client,
):
prompts = await manager.get_prompts_from_server(server, add_prefix=True)
mock_client.list_prompts.assert_awaited_once()
@ -520,7 +523,12 @@ class TestMCPServerManager:
mock_client = AsyncMock()
mock_client.get_prompt = AsyncMock(return_value=mock_result)
with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client):
with patch.object(
manager,
"_create_mcp_client",
new_callable=AsyncMock,
return_value=mock_client,
):
result = await manager.get_prompt_from_server(
server=server,
prompt_name="hello",
@ -551,13 +559,23 @@ class TestMCPServerManager:
mock_client = AsyncMock()
mock_resources = [Resource(name="file", uri="https://example.com/file")]
mock_client.list_resources = AsyncMock(return_value=mock_resources)
prefixed_resources = [Resource(name="alias-server-file", uri="https://example.com/file")]
prefixed_resources = [
Resource(name="alias-server-file", uri="https://example.com/file")
]
with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client) as mock_create_client, patch.object(
manager,
"_create_prefixed_resources",
return_value=prefixed_resources,
) as mock_prefix:
with (
patch.object(
manager,
"_create_mcp_client",
new_callable=AsyncMock,
return_value=mock_client,
) as mock_create_client,
patch.object(
manager,
"_create_prefixed_resources",
return_value=prefixed_resources,
) as mock_prefix,
):
result = await manager.get_resources_from_server(
server=server,
mcp_auth_header="auth",
@ -602,11 +620,19 @@ class TestMCPServerManager:
)
]
with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client) as mock_create_client, patch.object(
manager,
"_create_prefixed_resource_templates",
return_value=prefixed_templates,
) as mock_prefix:
with (
patch.object(
manager,
"_create_mcp_client",
new_callable=AsyncMock,
return_value=mock_client,
) as mock_create_client,
patch.object(
manager,
"_create_prefixed_resource_templates",
return_value=prefixed_templates,
) as mock_prefix,
):
result = await manager.get_resource_templates_from_server(
server=server,
mcp_auth_header="auth",
@ -650,7 +676,12 @@ class TestMCPServerManager:
)
mock_client.read_resource = AsyncMock(return_value=read_result)
with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client) as mock_create_client:
with patch.object(
manager,
"_create_mcp_client",
new_callable=AsyncMock,
return_value=mock_client,
) as mock_create_client:
result = await manager.read_resource_from_server(
server=server,
url="https://example.com/resource",
@ -661,7 +692,9 @@ class TestMCPServerManager:
mock_create_client.assert_called_once()
called_kwargs = mock_create_client.call_args.kwargs
assert called_kwargs["extra_headers"] == {"X-Test": "1", "X-Static": "1"}
mock_client.read_resource.assert_awaited_once_with("https://example.com/resource")
mock_client.read_resource.assert_awaited_once_with(
"https://example.com/resource"
)
assert result is read_result
@pytest.mark.asyncio
@ -724,22 +757,27 @@ class TestMCPServerManager:
registration_url=None,
)
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
), patch.object(
manager,
"_fetch_oauth_metadata_from_resource",
AsyncMock(return_value=([], None)),
), patch.object(
manager,
"_attempt_well_known_discovery",
AsyncMock(return_value=([], None)),
), patch.object(
manager,
"_fetch_authorization_server_metadata",
AsyncMock(return_value=mock_metadata),
) as mock_fetch_auth:
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
),
patch.object(
manager,
"_fetch_oauth_metadata_from_resource",
AsyncMock(return_value=([], None)),
),
patch.object(
manager,
"_attempt_well_known_discovery",
AsyncMock(return_value=([], None)),
),
patch.object(
manager,
"_fetch_authorization_server_metadata",
AsyncMock(return_value=mock_metadata),
) as mock_fetch_auth,
):
result = await manager._descovery_metadata(server_url)
mock_fetch_auth.assert_awaited_once_with(["https://example.com"])
@ -779,9 +817,8 @@ class TestMCPServerManager:
assert server.scopes == ["config"] # config overrides discovery
assert server.authorization_url == "https://config.example.com/auth"
assert server.token_url == "https://discovered.example.com/token"
assert (
server.registration_url == "https://discovered.example.com/register"
)
assert server.registration_url == "https://discovered.example.com/register"
@pytest.mark.asyncio
async def test_config_oauth_initialize_tool_name_to_mcp_server_name_mapping(self):
manager = MCPServerManager()
@ -801,7 +838,7 @@ class TestMCPServerManager:
# Initialize the tool mapping
await manager._initialize_tool_name_to_mcp_server_name_mapping()
assert manager.tool_name_to_mcp_server_name_mapping == {}
@pytest.mark.asyncio
async def test_list_tools_handles_missing_server_alias(self):
"""Test that list_tools handles servers without alias gracefully"""
@ -1017,7 +1054,9 @@ class TestMCPServerManager:
# Capture the extra_headers passed to _create_mcp_client
captured_extra_headers = None
async def capture_create_mcp_client(server, mcp_auth_header, extra_headers, stdio_env):
async def capture_create_mcp_client(
server, mcp_auth_header, extra_headers, stdio_env
):
nonlocal captured_extra_headers
captured_extra_headers = extra_headers
return mock_client
@ -1314,15 +1353,19 @@ class TestMCPServerManager:
return tool_func
with patch(
"litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.create_tool_function",
side_effect=fake_create_tool_function,
), patch(
"litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.build_input_schema",
return_value={"type": "object", "properties": {}, "required": []},
), patch(
"litellm.proxy._experimental.mcp_server.tool_registry.global_mcp_tool_registry.register_tool",
return_value=None,
with (
patch(
"litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.create_tool_function",
side_effect=fake_create_tool_function,
),
patch(
"litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.build_input_schema",
return_value={"type": "object", "properties": {}, "required": []},
),
patch(
"litellm.proxy._experimental.mcp_server.tool_registry.global_mcp_tool_registry.register_tool",
return_value=None,
),
):
await manager._register_openapi_tools(
spec_path=str(spec_path),
@ -2161,7 +2204,9 @@ class TestMCPServerManager:
# Register the server and map a tool to it
manager.registry = {"test-server": server}
manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server"
manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server"
manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = (
"test-server"
)
# Create mock client that tracks call_tool usage
mock_client = AsyncMock()
@ -2252,11 +2297,16 @@ class TestMCPServerManager:
# Verify MCPRequestHandler.get_allowed_mcp_servers was called with user_api_key_auth
mock_get_allowed.assert_called_once()
call_args = mock_get_allowed.call_args
assert call_args[0][0] is user_api_key_auth # First positional arg should be user_api_key_auth
assert (
call_args[0][0] is user_api_key_auth
) # First positional arg should be user_api_key_auth
assert call_args[0][0].user_id == "user-123"
assert call_args[0][0].object_permission_id == "perm_123"
assert call_args[0][0].object_permission is not None
assert call_args[0][0].object_permission.mcp_servers == ["test_server_1", "test_server_2"]
assert call_args[0][0].object_permission.mcp_servers == [
"test_server_1",
"test_server_2",
]
# Verify result contains the expected servers
assert "test_server_1" in result
@ -2483,5 +2533,82 @@ class TestHasClientCredentialsOAuth2Flow:
assert server.needs_user_oauth_token is False
# ---------------------------------------------------------------------------
# Upstream initialize-instructions cache
# ---------------------------------------------------------------------------
class TestMCPServerManagerUpstreamInstructionsCache:
"""Tests for the upstream initialize-instructions cache."""
def test_get_returns_none_when_empty(self):
"""Empty cache returns None for any key."""
manager = MCPServerManager()
assert (
manager._upstream_initialize_instructions_by_server_id.get("nonexistent")
is None
)
def test_remember_stores_stripped_value(self):
"""_remember_upstream_initialize_instructions stores a stripped string."""
manager = MCPServerManager()
fake_server = MagicMock(server_id="srv")
fake_client = MagicMock(_last_initialize_instructions=" hello \n")
manager._remember_upstream_initialize_instructions(fake_server, fake_client)
assert (
manager._upstream_initialize_instructions_by_server_id.get("srv") == "hello"
)
def test_remember_ignores_empty_string(self):
"""Whitespace-only instructions are not stored."""
manager = MCPServerManager()
fake_server = MagicMock(server_id="srv")
fake_client = MagicMock(_last_initialize_instructions=" ")
manager._remember_upstream_initialize_instructions(fake_server, fake_client)
assert manager._upstream_initialize_instructions_by_server_id.get("srv") is None
def test_remember_ignores_none(self):
"""None instructions are not stored."""
manager = MCPServerManager()
fake_server = MagicMock(server_id="srv")
fake_client = MagicMock(_last_initialize_instructions=None)
manager._remember_upstream_initialize_instructions(fake_server, fake_client)
assert manager._upstream_initialize_instructions_by_server_id.get("srv") is None
@pytest.mark.asyncio
async def test_load_servers_from_config_clears_cache(self):
"""Reloading config clears any previously cached upstream instructions."""
manager = MCPServerManager()
manager._upstream_initialize_instructions_by_server_id["old"] = "stale"
await manager.load_servers_from_config(
mcp_servers_config={
"fresh_srv": {
"url": "https://example.com",
"instructions": "from yaml",
}
}
)
assert manager._upstream_initialize_instructions_by_server_id.get("old") is None
@pytest.mark.asyncio
async def test_load_servers_reads_instructions_from_config(self):
"""instructions field from YAML config is persisted on the MCPServer."""
manager = MCPServerManager()
await manager.load_servers_from_config(
mcp_servers_config={
"srv_a": {
"url": "https://a.example.com",
"instructions": "A instructions",
},
"srv_b": {
"url": "https://b.example.com",
},
}
)
by_name = {s.server_name: s for s in manager.config_mcp_servers.values()}
assert "srv_a" in by_name and by_name["srv_a"].instructions == "A instructions"
assert "srv_b" in by_name and by_name["srv_b"].instructions is None
if __name__ == "__main__":
pytest.main([__file__])

View file

@ -399,7 +399,9 @@ class TestMCPServerManagerSigV4:
server = next(iter(manager.config_mcp_servers.values()))
assert server.auth_type == MCPAuth.aws_sigv4
assert server.aws_access_key_id == "AKIAIOSFODNN7EXAMPLE"
assert server.aws_secret_access_key == "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
assert (
server.aws_secret_access_key == "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
)
assert server.aws_region_name == "us-east-1"
assert server.aws_service_name == "bedrock-agentcore"
@ -529,7 +531,9 @@ class TestMCPServerManagerSigV4:
"aws_session_name": "my-session",
}
result = manager._extract_aws_credentials(creds, credentials_are_encrypted=False)
result = manager._extract_aws_credentials(
creds, credentials_are_encrypted=False
)
assert result["aws_role_name"] == "arn:aws:iam::123456789012:role/TestRole"
assert result["aws_session_name"] == "my-session"
@ -615,12 +619,15 @@ class TestCredentialMergeOnUpdate:
credentials={"aws_region_name": "eu-west-1"},
)
with patch(
"litellm.proxy._experimental.mcp_server.db._get_salt_key",
return_value=None,
), patch(
"litellm.proxy._experimental.mcp_server.db.encrypt_value_helper",
side_effect=lambda value, new_encryption_key: value,
with (
patch(
"litellm.proxy._experimental.mcp_server.db._get_salt_key",
return_value=None,
),
patch(
"litellm.proxy._experimental.mcp_server.db.encrypt_value_helper",
side_effect=lambda value, new_encryption_key: value,
),
):
await update_mcp_server(mock_prisma, data, "test-user")
@ -685,12 +692,15 @@ class TestCredentialMergeOnUpdate:
credentials={"aws_region_name": "us-east-1"},
)
with patch(
"litellm.proxy._experimental.mcp_server.db._get_salt_key",
return_value=None,
), patch(
"litellm.proxy._experimental.mcp_server.db.encrypt_value_helper",
side_effect=lambda value, new_encryption_key: value,
with (
patch(
"litellm.proxy._experimental.mcp_server.db._get_salt_key",
return_value=None,
),
patch(
"litellm.proxy._experimental.mcp_server.db.encrypt_value_helper",
side_effect=lambda value, new_encryption_key: value,
),
):
await update_mcp_server(mock_prisma, data, "test-user")
@ -728,12 +738,15 @@ class TestCredentialMergeOnUpdate:
credentials={"auth_value": "my-key"},
)
with patch(
"litellm.proxy._experimental.mcp_server.db._get_salt_key",
return_value=None,
), patch(
"litellm.proxy._experimental.mcp_server.db.encrypt_value_helper",
side_effect=lambda value, new_encryption_key: f"enc:{value}",
with (
patch(
"litellm.proxy._experimental.mcp_server.db._get_salt_key",
return_value=None,
),
patch(
"litellm.proxy._experimental.mcp_server.db.encrypt_value_helper",
side_effect=lambda value, new_encryption_key: f"enc:{value}",
),
):
await update_mcp_server(mock_prisma, data, "test-user")
@ -772,12 +785,15 @@ class TestCredentialMergeOnUpdate:
credentials={"scopes": ["read", "write"]},
)
with patch(
"litellm.proxy._experimental.mcp_server.db._get_salt_key",
return_value=None,
), patch(
"litellm.proxy._experimental.mcp_server.db.encrypt_value_helper",
side_effect=lambda value, new_encryption_key: value,
with (
patch(
"litellm.proxy._experimental.mcp_server.db._get_salt_key",
return_value=None,
),
patch(
"litellm.proxy._experimental.mcp_server.db.encrypt_value_helper",
side_effect=lambda value, new_encryption_key: value,
),
):
await update_mcp_server(mock_prisma, data, "test-user")
@ -803,7 +819,9 @@ class TestSigV4BuildFromTable:
table_record.server_name = "sigv4_server"
table_record.alias = None
table_record.description = None
table_record.url = "https://bedrock-agentcore.us-east-1.amazonaws.com/invocations"
table_record.url = (
"https://bedrock-agentcore.us-east-1.amazonaws.com/invocations"
)
table_record.spec_path = None
table_record.transport = "http"
table_record.auth_type = "aws_sigv4"
@ -838,6 +856,7 @@ class TestSigV4BuildFromTable:
table_record.tool_name_to_description = None
table_record.byok_api_key_help_url = None
table_record.oauth2_flow = None
table_record.instructions = None
manager = MCPServerManager()
@ -895,6 +914,7 @@ class TestSigV4BuildFromTable:
table_record.tool_name_to_description = None
table_record.byok_api_key_help_url = None
table_record.oauth2_flow = None
table_record.instructions = None
manager = MCPServerManager()
@ -934,7 +954,9 @@ class TestDecryptCredentials:
with patch(
"litellm.proxy._experimental.mcp_server.db.decrypt_value_helper",
side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace("enc:", ""),
side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace(
"enc:", ""
),
):
result = decrypt_credentials(credentials=creds)
@ -956,7 +978,9 @@ class TestDecryptCredentials:
with patch(
"litellm.proxy._experimental.mcp_server.db.decrypt_value_helper",
side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace("enc:", ""),
side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace(
"enc:", ""
),
):
result = decrypt_credentials(credentials=creds)
@ -988,15 +1012,21 @@ class TestRotateCredentials:
)
mock_prisma.db.litellm_mcpservertable.update = AsyncMock()
with patch(
"litellm.proxy._experimental.mcp_server.db._get_salt_key",
return_value="old-key",
), patch(
"litellm.proxy._experimental.mcp_server.db.decrypt_value_helper",
side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace("enc_old:", ""),
), patch(
"litellm.proxy._experimental.mcp_server.db.encrypt_value_helper",
side_effect=lambda value, new_encryption_key: f"enc_new:{value}",
with (
patch(
"litellm.proxy._experimental.mcp_server.db._get_salt_key",
return_value="old-key",
),
patch(
"litellm.proxy._experimental.mcp_server.db.decrypt_value_helper",
side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace(
"enc_old:", ""
),
),
patch(
"litellm.proxy._experimental.mcp_server.db.encrypt_value_helper",
side_effect=lambda value, new_encryption_key: f"enc_new:{value}",
),
):
await rotate_mcp_server_credentials_master_key(
mock_prisma, "admin", "new-key"

View file

@ -1,7 +1,10 @@
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy.health_check import _update_litellm_params_for_health_check
from litellm.litellm_core_utils.health_check_helpers import HealthCheckHelpers
from unittest.mock import AsyncMock, patch, MagicMock
from litellm.proxy import health_check as hc_module
from litellm.proxy.health_check import _update_litellm_params_for_health_check
@pytest.mark.asyncio
@ -50,10 +53,13 @@ async def test_ahealth_check_wildcard_models_respects_max_tokens():
Test that ahealth_check_wildcard_models respects max_tokens if passed,
otherwise defaults to 10.
"""
with patch(
"litellm.litellm_core_utils.llm_request_utils.pick_cheapest_chat_models_from_llm_provider",
return_value=["gpt-4o-mini"],
), patch("litellm.acompletion", new_callable=AsyncMock):
with (
patch(
"litellm.litellm_core_utils.llm_request_utils.pick_cheapest_chat_models_from_llm_provider",
return_value=["gpt-4o-mini"],
),
patch("litellm.acompletion", new_callable=AsyncMock),
):
# Test Case 1: No max_tokens passed, should default to 10
model_params = {}
await HealthCheckHelpers.ahealth_check_wildcard_models(
@ -73,3 +79,50 @@ async def test_ahealth_check_wildcard_models_respects_max_tokens():
litellm_logging_obj=MagicMock(),
)
assert model_params["max_tokens"] == 3
@pytest.mark.asyncio
async def test_background_health_check_max_tokens_env_var(monkeypatch):
"""
Test that BACKGROUND_HEALTH_CHECK_MAX_TOKENS env var is used as global default
for explicit (non-wildcard) models.
"""
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", 10)
model_info = {}
litellm_params = {"model": "azure/gpt-4"}
updated_params = _update_litellm_params_for_health_check(model_info, litellm_params)
assert updated_params["max_tokens"] == 10
@pytest.mark.asyncio
async def test_per_model_overrides_global_env_var(monkeypatch):
"""
Test that per-model health_check_max_tokens takes priority over
BACKGROUND_HEALTH_CHECK_MAX_TOKENS env var.
"""
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", 10)
model_info = {"health_check_max_tokens": 5}
litellm_params = {"model": "azure/gpt-4"}
updated_params = _update_litellm_params_for_health_check(model_info, litellm_params)
assert updated_params["max_tokens"] == 5
@pytest.mark.asyncio
async def test_global_env_var_applies_to_wildcard_models(monkeypatch):
"""
Test that BACKGROUND_HEALTH_CHECK_MAX_TOKENS env var also applies to wildcard models.
"""
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", 15)
model_info = {}
litellm_params = {"model": "openai/*"}
updated_params = _update_litellm_params_for_health_check(model_info, litellm_params)
assert updated_params["max_tokens"] == 15

View file

@ -31,3 +31,32 @@ def test_get_standard_logging_model_parameters_excludes_prompt_content():
assert "prompt" not in result
assert "input" not in result
assert result == {"temperature": 0.5}
def test_get_all_llm_api_params_includes_responses_api():
"""
Regression guard for the Responses API cache-key bug:
Responses-API-only kwargs must be present in the cache-key allow-list,
otherwise Cache.get_cache_key() silently drops them and two requests
that differ only in (e.g.) `instructions` collide on the same key.
"""
all_params = ModelParamHelper._get_all_llm_api_params()
responses_only_params = {
"instructions",
"previous_response_id",
"reasoning",
"include",
"store",
"background",
"max_output_tokens",
"max_tool_calls",
"prompt_cache_key",
"prompt_cache_retention",
"context_management",
"conversation",
"safety_identifier",
}
missing = responses_only_params - all_params
assert (
missing == set()
), f"Responses-API kwargs missing from cache-key allow-list: {sorted(missing)}"