mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge pull request #25699 from BerriAI/litellm_ishaan_april14
Litellm ishaan april14
This commit is contained in:
commit
0b7335201b
29 changed files with 1800 additions and 457 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "instructions" TEXT;
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ""
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
|
|
@ -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"] == []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)}"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue