From 6621c40622bf6e37879529c111d0a1cc9e2f2cd2 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 30 Mar 2026 10:02:32 +0530 Subject: [PATCH 01/34] feat(health-check): add BACKGROUND_HEALTH_CHECK_MAX_TOKENS env var --- litellm/constants.py | 10 +++- litellm/proxy/health_check.py | 8 ++- .../proxy/test_health_check_max_tokens.py | 54 ++++++++++++++++++- 3 files changed, 68 insertions(+), 4 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 28c6c0cc0e3..a10a288d16c 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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 @@ -1307,6 +1307,14 @@ 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" +) +BACKGROUND_HEALTH_CHECK_MAX_TOKENS: Optional[int] = ( + int(_background_health_check_max_tokens_env) + if _background_health_check_max_tokens_env + else 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" diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 5d1bcf31f84..5377ff6c320 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -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", @@ -301,6 +305,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 "" ): diff --git a/tests/test_litellm/proxy/test_health_check_max_tokens.py b/tests/test_litellm/proxy/test_health_check_max_tokens.py index e26f7fb9f20..bb125f764b8 100644 --- a/tests/test_litellm/proxy/test_health_check_max_tokens.py +++ b/tests/test_litellm/proxy/test_health_check_max_tokens.py @@ -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 @@ -73,3 +76,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 From d8598abf150f0e3716718f0d6060dae858d016fc Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 8 Apr 2026 21:33:27 +0530 Subject: [PATCH 02/34] Fix code qa --- docs/my-website/docs/proxy/config_settings.md | 1 + 1 file changed, 1 insertion(+) diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index 3b090b3a44a..e2eedda1817 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -486,6 +486,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 From b08e05845158b89f11b9200a4c07d5a2ca9d6181 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 8 Apr 2026 21:35:41 +0530 Subject: [PATCH 03/34] Fix greptile reviews --- litellm/constants.py | 18 +++++++++++++----- 1 file changed, 13 insertions(+), 5 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index a10a288d16c..a35bad7b98c 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1310,11 +1310,19 @@ HEALTH_CHECK_TIMEOUT_SECONDS = int( _background_health_check_max_tokens_env = os.getenv( "BACKGROUND_HEALTH_CHECK_MAX_TOKENS" ) -BACKGROUND_HEALTH_CHECK_MAX_TOKENS: Optional[int] = ( - int(_background_health_check_max_tokens_env) - if _background_health_check_max_tokens_env - else None -) +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" From 85497740ff4b88d6a5d8b38fd7b3f29afa156c96 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Tue, 14 Apr 2026 05:40:42 +0200 Subject: [PATCH 04/34] fix(caching): add Responses API params to cache key allow-list --- .../litellm_core_utils/model_param_helper.py | 23 +++++++++ tests/local_testing/test_unit_test_caching.py | 51 +++++++++++++++++++ tests/test_litellm/test_model_param_helper.py | 29 +++++++++++ 3 files changed, 103 insertions(+) diff --git a/litellm/litellm_core_utils/model_param_helper.py b/litellm/litellm_core_utils/model_param_helper.py index 66b174feac4..35e744be3a6 100644 --- a/litellm/litellm_core_utils/model_param_helper.py +++ b/litellm/litellm_core_utils/model_param_helper.py @@ -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 @@ -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]: """ diff --git a/tests/local_testing/test_unit_test_caching.py b/tests/local_testing/test_unit_test_caching.py index fa5cf802546..e4ee65a2aa2 100644 --- a/tests/local_testing/test_unit_test_caching.py +++ b/tests/local_testing/test_unit_test_caching.py @@ -131,6 +131,57 @@ 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" diff --git a/tests/test_litellm/test_model_param_helper.py b/tests/test_litellm/test_model_param_helper.py index c6e4b864a22..2012abec547 100644 --- a/tests/test_litellm/test_model_param_helper.py +++ b/tests/test_litellm/test_model_param_helper.py @@ -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)}" + ) From c7f7708d27593a4c5ecf44f9a00c4ec2ffd6a0b1 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 14 Apr 2026 09:39:39 +0530 Subject: [PATCH 05/34] feat(anthropic): retry /v1/messages after invalid thinking signature Strip thinking blocks from the request body and retry once when Anthropic returns an invalid thinking signature error (e.g. after credential or deployment change). Applies to all BaseAnthropicMessagesConfig providers (direct Anthropic, Bedrock, Vertex, Azure AI). Made-with: Cursor --- litellm/llms/anthropic/common_utils.py | 58 +++++++++ .../anthropic_messages/transformation.py | 37 ++++++ litellm/llms/custom_httpx/llm_http_handler.py | 70 +++++++++-- .../anthropic/test_anthropic_common_utils.py | 118 ++++++++++++++++++ 4 files changed, 272 insertions(+), 11 deletions(-) diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index a0da14bcc2b..4f7f3814e74 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -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,63 @@ 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. + """ + 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): + mm["content"] = [ + b + for b in content + if not ( + isinstance(b, dict) + and b.get("type") in ("thinking", "redacted_thinking") + ) + ] + 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: diff --git a/litellm/llms/base_llm/anthropic_messages/transformation.py b/litellm/llms/base_llm/anthropic_messages/transformation.py index fdad1633e8f..40063c0fb9a 100644 --- a/litellm/llms/base_llm/anthropic_messages/transformation.py +++ b/litellm/llms/base_llm/anthropic_messages/transformation.py @@ -120,3 +120,40 @@ 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 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 is_anthropic_invalid_thinking_signature_error(e.response.text): + strip_thinking_blocks_from_anthropic_messages_request_dict(request_data) + return request_data diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 7a8820a8785..06f030fc2ce 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1955,18 +1955,66 @@ 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: + max_anthropic_messages_http_attempts = max( + anthropic_messages_provider_config.max_retry_on_anthropic_messages_http_error, + 1, + ) + response: Optional[httpx.Response] = None + litellm_params_dict = dict(litellm_params) + for attempt_idx in range(max_anthropic_messages_http_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() + except httpx.HTTPStatusError as e: + hit_max_attempt = ( + attempt_idx + 1 == max_anthropic_messages_http_attempts + ) + should_retry = anthropic_messages_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( + "Retrying on HTTPStatusError (attempt %s/%s).", + attempt_idx + 2, + max_anthropic_messages_http_attempts, + ) + + request_body = anthropic_messages_provider_config.transform_anthropic_messages_request_on_http_error( + e=e, request_data=request_body + ) + headers, signed_json_body = ( + anthropic_messages_provider_config.sign_request( + headers=headers, + optional_params=dict(litellm_params), + 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=anthropic_messages_provider_config + ) + except Exception as e: + raise self._handle_error( + e=e, provider_config=anthropic_messages_provider_config + ) + break + + if response is None: raise self._handle_error( - e=e, provider_config=anthropic_messages_provider_config + e=ValueError("No response from Anthropic /v1/messages"), + provider_config=anthropic_messages_provider_config, ) # used for logging + cost tracking diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index 22470b93540..ba344403912 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -1131,3 +1131,121 @@ 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_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"][0]["content"] == [] + + 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 + + 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]["content"] == [] From 0f453cc59d928b85ef7223c8bd92c63be9c0b9b8 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 14 Apr 2026 10:00:27 +0530 Subject: [PATCH 06/34] Fic code qa --- litellm/llms/custom_httpx/llm_http_handler.py | 139 ++++++++++-------- 1 file changed, 79 insertions(+), 60 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 06f030fc2ce..9155fbb4aca 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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) + response: Optional[httpx.Response] = None + 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() + 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) + break + + if response is None: + raise self._handle_error( + e=ValueError("No response from Anthropic /v1/messages"), + provider_config=provider_config, + ) + return response + async def async_anthropic_messages_handler( self, model: str, @@ -1955,67 +2022,19 @@ class BaseLLMHTTPHandler: }, ) - max_anthropic_messages_http_attempts = max( - anthropic_messages_provider_config.max_retry_on_anthropic_messages_http_error, - 1, + 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, ) - response: Optional[httpx.Response] = None - litellm_params_dict = dict(litellm_params) - for attempt_idx in range(max_anthropic_messages_http_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() - except httpx.HTTPStatusError as e: - hit_max_attempt = ( - attempt_idx + 1 == max_anthropic_messages_http_attempts - ) - should_retry = anthropic_messages_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( - "Retrying on HTTPStatusError (attempt %s/%s).", - attempt_idx + 2, - max_anthropic_messages_http_attempts, - ) - - request_body = anthropic_messages_provider_config.transform_anthropic_messages_request_on_http_error( - e=e, request_data=request_body - ) - headers, signed_json_body = ( - anthropic_messages_provider_config.sign_request( - headers=headers, - optional_params=dict(litellm_params), - 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=anthropic_messages_provider_config - ) - except Exception as e: - raise self._handle_error( - e=e, provider_config=anthropic_messages_provider_config - ) - break - - if response is None: - raise self._handle_error( - e=ValueError("No response from Anthropic /v1/messages"), - provider_config=anthropic_messages_provider_config, - ) # used for logging + cost tracking logging_obj.model_call_details["httpx_response"] = response From 5670f6c7d49bb2255d9a9460336b74e4ea2157be Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 14 Apr 2026 10:03:10 +0530 Subject: [PATCH 07/34] fix(anthropic): tighten thinking-signature retry (Greptile) - Omit messages whose list content is empty after stripping thinking blocks - Retry only on HTTP 400 plus invalid-signature body match - Return response inline from retry loop; drop unreachable None guard - Tests: thinking-only turn dropped, non-400 no retry Made-with: Cursor --- litellm/llms/anthropic/common_utils.py | 8 +++++- .../anthropic_messages/transformation.py | 8 ++++-- litellm/llms/custom_httpx/llm_http_handler.py | 12 +++------ .../anthropic/test_anthropic_common_utils.py | 26 +++++++++++++++++-- 4 files changed, 41 insertions(+), 13 deletions(-) diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 4f7f3814e74..3be5a0c816f 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -760,6 +760,9 @@ def strip_thinking_blocks_from_anthropic_messages(messages: List[Any]) -> List[A """ 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: @@ -769,7 +772,7 @@ def strip_thinking_blocks_from_anthropic_messages(messages: List[Any]) -> List[A mm = copy.deepcopy(m) content = mm.get("content") if isinstance(content, list): - mm["content"] = [ + filtered = [ b for b in content if not ( @@ -777,6 +780,9 @@ def strip_thinking_blocks_from_anthropic_messages(messages: List[Any]) -> List[A and b.get("type") in ("thinking", "redacted_thinking") ) ] + if not filtered: + continue + mm["content"] = filtered out.append(mm) return out diff --git a/litellm/llms/base_llm/anthropic_messages/transformation.py b/litellm/llms/base_llm/anthropic_messages/transformation.py index 40063c0fb9a..733faca8532 100644 --- a/litellm/llms/base_llm/anthropic_messages/transformation.py +++ b/litellm/llms/base_llm/anthropic_messages/transformation.py @@ -141,7 +141,9 @@ class BaseAnthropicMessagesConfig(ABC): is_anthropic_invalid_thinking_signature_error, ) - return is_anthropic_invalid_thinking_signature_error(e.response.text) + 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 @@ -154,6 +156,8 @@ class BaseAnthropicMessagesConfig(ABC): strip_thinking_blocks_from_anthropic_messages_request_dict, ) - if is_anthropic_invalid_thinking_signature_error(e.response.text): + 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 diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 9155fbb4aca..60efa45df62 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1833,7 +1833,6 @@ class BaseLLMHTTPHandler: 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) - response: Optional[httpx.Response] = None for attempt_idx in range(max_attempts): try: response = await async_httpx_client.post( @@ -1844,6 +1843,7 @@ class BaseLLMHTTPHandler: 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( @@ -1874,14 +1874,10 @@ class BaseLLMHTTPHandler: raise self._handle_error(e=e, provider_config=provider_config) except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) - break - if response is None: - raise self._handle_error( - e=ValueError("No response from Anthropic /v1/messages"), - provider_config=provider_config, - ) - return response + raise RuntimeError( + "unreachable: anthropic messages HTTP retry loop exited without return" + ) async def async_anthropic_messages_handler( self, diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index ba344403912..d48d7716a8e 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -1181,6 +1181,24 @@ class TestAnthropicThinkingSignatureSelfHeal: 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, @@ -1204,7 +1222,7 @@ class TestAnthropicThinkingSignatureSelfHeal: } strip_thinking_blocks_from_anthropic_messages_request_dict(data) assert "thinking" not in data - assert data["messages"][0]["content"] == [] + assert data["messages"] == [] def test_anthropic_messages_config_http_retry_helpers(self): import httpx @@ -1230,6 +1248,10 @@ class TestAnthropicThinkingSignatureSelfHeal: 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": [ @@ -1248,4 +1270,4 @@ class TestAnthropicThinkingSignatureSelfHeal: } config.transform_anthropic_messages_request_on_http_error(err, data) assert "thinking" not in data - assert data["messages"][0]["content"] == [] + assert data["messages"] == [] From 63281e8330109004281eac284919e9d56002ba78 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Tue, 14 Apr 2026 08:10:39 +0200 Subject: [PATCH 08/34] fix(azure/passthrough): populate standard_logging_object via logging hook --- .../llms/azure/passthrough/transformation.py | 37 +++++++ .../test_azure_passthrough_transformation.py | 97 +++++++++++++++++++ 2 files changed, 134 insertions(+) create mode 100644 tests/test_litellm/llms/azure/passthrough/test_azure_passthrough_transformation.py diff --git a/litellm/llms/azure/passthrough/transformation.py b/litellm/llms/azure/passthrough/transformation.py index 4e9de4b314f..9b1d95e5314 100644 --- a/litellm/llms/azure/passthrough/transformation.py +++ b/litellm/llms/azure/passthrough/transformation.py @@ -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 diff --git a/tests/test_litellm/llms/azure/passthrough/test_azure_passthrough_transformation.py b/tests/test_litellm/llms/azure/passthrough/test_azure_passthrough_transformation.py new file mode 100644 index 00000000000..529a7453d74 --- /dev/null +++ b/tests/test_litellm/llms/azure/passthrough/test_azure_passthrough_transformation.py @@ -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 From 96ed00e1840669accc54c1fbeb19639ce32023a2 Mon Sep 17 00:00:00 2001 From: Milan Date: Tue, 14 Apr 2026 14:19:31 +0300 Subject: [PATCH 09/34] feat(mcp): gateway InitializeResult.instructions from upstream or YAML - Add optional instructions on MCPServer (config/DB/types) and Prisma migration. - MCPClient: fetch_upstream_initialize_instructions() for one-shot initialize. - Gateway merges per-request instructions: YAML/API overrides; otherwise fetch upstream initialize instructions (skip spec_path/OpenAPI-only servers). - Pass auth headers into instruction merge; ContextVar for gateway Server. - REST: wire instructions on connection-test MCPServer payloads. Made-with: Cursor --- .../migration.sql | 2 + .../litellm_proxy_extras/schema.prisma | 1 + litellm/experimental_mcp_client/client.py | 46 +++++ .../_experimental/mcp_server/mcp_context.py | 5 + .../mcp_server/mcp_server_manager.py | 4 + .../mcp_server/rest_endpoints.py | 1 + .../proxy/_experimental/mcp_server/server.py | 178 ++++++++++++++++-- litellm/proxy/_types.py | 4 + litellm/proxy/schema.prisma | 1 + .../types/mcp_server/mcp_server_manager.py | 2 + schema.prisma | 1 + 11 files changed, 229 insertions(+), 16 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260414140000_add_mcp_server_instructions/migration.sql diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260414140000_add_mcp_server_instructions/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260414140000_add_mcp_server_instructions/migration.sql new file mode 100644 index 00000000000..531024c519f --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260414140000_add_mcp_server_instructions/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "instructions" TEXT; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index fce95465b55..a728d912715 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -289,6 +289,7 @@ model LiteLLM_MCPServerTable { server_name String? alias String? description String? + instructions String? // MCP InitializeResult.instructions (optional) url String? spec_path String? transport String @default("sse") diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 1423617cac0..fe56a418e7b 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -329,6 +329,52 @@ class MCPClient: except BaseException as e: verbose_logger.debug(f"Error during http_client cleanup: {e}") + async def fetch_upstream_initialize_instructions(self) -> Optional[str]: + """Open a transport, run ``initialize`` once, return upstream ``instructions``.""" + http_client: Optional[httpx.AsyncClient] = None + try: + transport_ctx, http_client = self._create_transport_context() + transport = await transport_ctx.__aenter__() + try: + read_stream, write_stream = transport[0], transport[1] + session_ctx = ClientSession(read_stream, write_stream) + session = await session_ctx.__aenter__() + try: + init = await session.initialize() + return init.instructions + finally: + try: + await session_ctx.__aexit__(None, None, None) + except BaseException as e: + verbose_logger.debug( + "Error during session context exit (instructions fetch): %s", + e, + ) + finally: + try: + await transport_ctx.__aexit__(None, None, None) + except BaseException as e: + verbose_logger.debug( + "Error during transport context exit (instructions fetch): %s", + e, + ) + except Exception as e: + verbose_logger.debug( + "fetch_upstream_initialize_instructions failed for %s: %s", + self.server_url or "stdio", + e, + ) + return None + finally: + if http_client is not None: + try: + await http_client.aclose() + except BaseException as e: + verbose_logger.debug( + "Error during http_client cleanup (instructions fetch): %s", + e, + ) + def update_auth_value(self, mcp_auth_value: Union[str, Dict[str, str]]): """ Set the authentication header for the MCP client. diff --git a/litellm/proxy/_experimental/mcp_server/mcp_context.py b/litellm/proxy/_experimental/mcp_server/mcp_context.py index 12830db1d6a..a60138dd340 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_context.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_context.py @@ -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 +) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 8d3831e75fb..dd7f092cb53 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -351,6 +351,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 +694,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 @@ -2946,6 +2948,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( @@ -3041,6 +3044,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]: diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 32560a2211d..8131c040136 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 99578d006e1..1402b385a8e 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -6,6 +6,7 @@ LiteLLM MCP Server Routes import asyncio import contextlib +import contextvars import time import traceback import uuid @@ -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,10 +206,27 @@ if MCP_AVAILABLE: ) return normalized + class _LitellmMcpGatewayServer(Server): + """Gateway server that injects per-request ``InitializeResult.instructions``.""" + + def create_initialization_options( # type: ignore[override] + self, + notification_options: Optional[NotificationOptions] = None, + experimental_capabilities: Optional[Dict[str, Dict[str, Any]]] = None, + ) -> InitializationOptions: + opts = super().create_initialization_options( + 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 ################# ######################################################## - server: Server = Server( + server: Server = _LitellmMcpGatewayServer( name=LITELLM_MCP_SERVER_NAME, version=LITELLM_MCP_SERVER_VERSION, ) @@ -814,10 +837,7 @@ if MCP_AVAILABLE: return tools def _get_client_ip_from_context() -> Optional[str]: - """ - Extract client_ip from auth context. - Returns None if context not set (caller should handle this as "no IP filtering"). - """ + """Return ``client_ip`` from MCP auth context (set by HTTP/SSE handlers), or None.""" try: auth_user = auth_context_var.get() if auth_user and isinstance(auth_user, MCPAuthenticatedUser): @@ -836,19 +856,15 @@ if MCP_AVAILABLE: Args: user_api_key_auth: The authenticated user's API key info. mcp_servers: Optional list of server names to filter to. - client_ip: Client IP for IP-based access control. If None, falls back to - auth context. Pass explicitly from request handlers for safety. - Note: If client_ip is None and auth context is not set, IP filtering is skipped. - This is intentional for internal callers but may indicate a bug if called - from a request handler without proper context setup. + client_ip: Client IP for IP-based access control. MCP HTTP/SSE handlers set auth context (including ``client_ip``) before MCP work; when this is + ``None``, ``client_ip`` is taken from that context. Callers may still + pass ``client_ip`` explicitly when already computed. """ - # Use explicit client_ip if provided, otherwise try auth context if client_ip is None: client_ip = _get_client_ip_from_context() if client_ip is None: verbose_logger.debug( - "MCP _get_allowed_mcp_servers called without client_ip and no auth context. " - "IP filtering will be skipped. This is expected for internal calls." + "MCP _get_allowed_mcp_servers: client IP unknown; skipping public-internet IP filter." ) allowed_mcp_server_ids = ( @@ -1103,6 +1119,112 @@ if MCP_AVAILABLE: return server_auth_header, extra_headers + async def _merge_gateway_initialize_instructions( + allowed_mcp_servers: List[MCPServer], + user_api_key_auth: Optional[UserAPIKeyAuth], + mcp_auth_header: Optional[str], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + oauth2_headers: Optional[Dict[str, str]], + raw_headers: Optional[Dict[str, str]], + ) -> Optional[str]: + """Merge ``instructions`` for gateway ``initialize``: YAML/API overrides upstream.""" + if not allowed_mcp_servers: + return None + + _has_oauth2_server = any( + getattr(s, "auth_type", None) == MCPAuth.oauth2 + for s in allowed_mcp_servers + ) + _prefetched_oauth_creds = ( + await _prefetch_oauth_creds_for_user(user_api_key_auth) + if _has_oauth2_server + else {} + ) + + async def _one(server: MCPServer) -> Optional[Tuple[str, str]]: + label = ( + server.alias + or server.server_name + or server.name + or server.server_id + or "mcp" + ) + if server.instructions and server.instructions.strip(): + return (label, server.instructions.strip()) + if server.spec_path: + return None + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=mcp_server_auth_headers, + mcp_auth_header=mcp_auth_header, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + ) + if extra_headers is None and server.auth_type == MCPAuth.oauth2: + extra_headers = await _get_user_oauth_extra_headers_from_db( + server, + user_api_key_auth, + prefetched_creds=_prefetched_oauth_creds, + ) + try: + if server.static_headers: + if extra_headers is None: + extra_headers = {} + extra_headers.update(server.static_headers) + stdio_env = global_mcp_server_manager._build_stdio_env( + server, raw_headers + ) + client = await global_mcp_server_manager._create_mcp_client( + server=server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + stdio_env=stdio_env, + ) + text = await client.fetch_upstream_initialize_instructions() + if text and text.strip(): + return (label, text.strip()) + except Exception as e: + verbose_logger.debug( + "MCP gateway: upstream instructions fetch failed for %s: %s", + server.name, + e, + ) + return None + + pairs = await asyncio.gather(*(_one(s) for s in allowed_mcp_servers)) + texts = [p for p in pairs if p is not None] + 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) + + async def _set_mcp_gateway_initialize_instructions_token( + user_api_key_auth: Optional[UserAPIKeyAuth], + mcp_servers: Optional[List[str]], + client_ip: Optional[str], + mcp_auth_header: Optional[str], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + oauth2_headers: Optional[Dict[str, str]], + raw_headers: Optional[Dict[str, str]], + ) -> contextvars.Token[Optional[str]]: + """Resolve merged gateway ``instructions``; return ContextVar token to reset.""" + allowed = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + merged = await _merge_gateway_initialize_instructions( + allowed_mcp_servers=allowed, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + ) + return _mcp_gateway_initialize_instructions.set(merged) + async def _get_tools_from_mcp_servers( # noqa: PLR0915 user_api_key_auth: Optional[UserAPIKeyAuth], mcp_auth_header: Optional[str], @@ -2670,7 +2792,19 @@ if MCP_AVAILABLE: # Request was fully handled (e.g., DELETE on non-existent session) return - await session_manager.handle_request(scope, receive, send) + _instr_tok = await _set_mcp_gateway_initialize_instructions_token( + user_api_key_auth, + mcp_servers, + _client_ip, + mcp_auth_header, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + ) + try: + await session_manager.handle_request(scope, receive, send) + finally: + _mcp_gateway_initialize_instructions.reset(_instr_tok) except HTTPException: # Re-raise HTTP exceptions to preserve status codes and details raise @@ -2729,7 +2863,19 @@ if MCP_AVAILABLE: await initialize_session_managers() await asyncio.sleep(0.1) - await sse_session_manager.handle_request(scope, receive, send) + _sse_instr_tok = await _set_mcp_gateway_initialize_instructions_token( + user_api_key_auth, + mcp_servers, + _sse_client_ip, + mcp_auth_header, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + ) + try: + await sse_session_manager.handle_request(scope, receive, send) + finally: + _mcp_gateway_initialize_instructions.reset(_sse_instr_tok) except Exception as e: verbose_logger.exception(f"Error handling MCP request: {e}") # Instead of re-raising, try to send a graceful error response diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0bbee56d5e0..6ba8d0b68a3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1137,6 +1137,8 @@ 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 + # Shown to MCP clients in InitializeResult.instructions (optional) + instructions: Optional[str] = None # Stdio-specific fields command: Optional[str] = None args: List[str] = Field(default_factory=list) @@ -1219,6 +1221,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 +1273,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 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index fce95465b55..a728d912715 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -289,6 +289,7 @@ model LiteLLM_MCPServerTable { server_name String? alias String? description String? + instructions String? // MCP InitializeResult.instructions (optional) url String? spec_path String? transport String @default("sse") diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index a7d0968c0ef..805494b1854 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -27,6 +27,8 @@ class MCPServer(BaseModel): spec_path: Optional[str] = None auth_type: Optional[MCPAuthType] = None authentication_token: Optional[str] = None + # Optional text returned on MCP `initialize` (InitializeResult.instructions) + instructions: Optional[str] = None mcp_info: Optional[MCPInfo] = None extra_headers: Optional[ List[str] diff --git a/schema.prisma b/schema.prisma index fce95465b55..a728d912715 100644 --- a/schema.prisma +++ b/schema.prisma @@ -289,6 +289,7 @@ model LiteLLM_MCPServerTable { server_name String? alias String? description String? + instructions String? // MCP InitializeResult.instructions (optional) url String? spec_path String? transport String @default("sse") From 8c2ebee4decd5a9bfe3b8ad5fb46fc8f5ca1e5fb Mon Sep 17 00:00:00 2001 From: Milan Date: Tue, 14 Apr 2026 14:36:06 +0300 Subject: [PATCH 10/34] refactor(mcp): reuse existing sessions for initialize instructions Remove the gateway-specific initialize fetch path and reuse instructions captured during existing MCP calls (list_tools/health_check/call_tool), while keeping YAML/DB instructions as immediate overrides. Made-with: Cursor --- .../litellm_proxy_extras/schema.prisma | 2 +- litellm/experimental_mcp_client/client.py | 55 +------ .../mcp_server/mcp_server_manager.py | 18 +++ .../proxy/_experimental/mcp_server/server.py | 143 +++++------------- litellm/proxy/_types.py | 1 - litellm/proxy/schema.prisma | 2 +- .../types/mcp_server/mcp_server_manager.py | 1 - schema.prisma | 2 +- 8 files changed, 71 insertions(+), 153 deletions(-) diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index a728d912715..9965c003b0a 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -289,7 +289,7 @@ model LiteLLM_MCPServerTable { server_name String? alias String? description String? - instructions String? // MCP InitializeResult.instructions (optional) + instructions String? url String? spec_path String? transport String @default("sse") diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index fe56a418e7b..e703a3956b9 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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: @@ -329,52 +336,6 @@ class MCPClient: except BaseException as e: verbose_logger.debug(f"Error during http_client cleanup: {e}") - async def fetch_upstream_initialize_instructions(self) -> Optional[str]: - """Open a transport, run ``initialize`` once, return upstream ``instructions``.""" - http_client: Optional[httpx.AsyncClient] = None - try: - transport_ctx, http_client = self._create_transport_context() - transport = await transport_ctx.__aenter__() - try: - read_stream, write_stream = transport[0], transport[1] - session_ctx = ClientSession(read_stream, write_stream) - session = await session_ctx.__aenter__() - try: - init = await session.initialize() - return init.instructions - finally: - try: - await session_ctx.__aexit__(None, None, None) - except BaseException as e: - verbose_logger.debug( - "Error during session context exit (instructions fetch): %s", - e, - ) - finally: - try: - await transport_ctx.__aexit__(None, None, None) - except BaseException as e: - verbose_logger.debug( - "Error during transport context exit (instructions fetch): %s", - e, - ) - except Exception as e: - verbose_logger.debug( - "fetch_upstream_initialize_instructions failed for %s: %s", - self.server_url or "stdio", - e, - ) - return None - finally: - if http_client is not None: - try: - await http_client.aclose() - except BaseException as e: - verbose_logger.debug( - "Error during http_client cleanup (instructions fetch): %s", - e, - ) - def update_auth_value(self, mcp_auth_value: Union[str, Dict[str, str]]): """ Set the authentication header for the MCP client. diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index dd7f092cb53..750c9204c55 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -184,6 +184,19 @@ class MCPServerManager: "gmail_send_email": "zapier_mcp_server", } """ + self._upstream_initialize_instructions_by_server_id: Dict[str, str] = {} + + def get_upstream_initialize_instructions(self, server_id: str) -> Optional[str]: + return self._upstream_initialize_instructions_by_server_id.get(server_id) + + 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 +217,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() @@ -1249,6 +1263,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 @@ -2385,6 +2400,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) @@ -2624,6 +2640,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( @@ -2907,6 +2924,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 = ( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 1402b385a8e..864b17afe33 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -8,6 +8,7 @@ import asyncio import contextlib import contextvars import time +import types import traceback import uuid from datetime import datetime @@ -206,30 +207,31 @@ if MCP_AVAILABLE: ) return normalized - class _LitellmMcpGatewayServer(Server): - """Gateway server that injects per-request ``InitializeResult.instructions``.""" - - def create_initialization_options( # type: ignore[override] + 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: Optional[NotificationOptions] = None, - experimental_capabilities: Optional[Dict[str, Dict[str, Any]]] = None, - ) -> InitializationOptions: - opts = super().create_initialization_options( - 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 + 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 ################# ######################################################## - server: Server = _LitellmMcpGatewayServer( + server: Server = Server( 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 @@ -837,7 +839,10 @@ if MCP_AVAILABLE: return tools def _get_client_ip_from_context() -> Optional[str]: - """Return ``client_ip`` from MCP auth context (set by HTTP/SSE handlers), or None.""" + """ + Extract client_ip from auth context. + Returns None if context not set (caller should handle this as "no IP filtering"). + """ try: auth_user = auth_context_var.get() if auth_user and isinstance(auth_user, MCPAuthenticatedUser): @@ -856,15 +861,19 @@ if MCP_AVAILABLE: Args: user_api_key_auth: The authenticated user's API key info. mcp_servers: Optional list of server names to filter to. - client_ip: Client IP for IP-based access control. MCP HTTP/SSE handlers set auth context (including ``client_ip``) before MCP work; when this is - ``None``, ``client_ip`` is taken from that context. Callers may still - pass ``client_ip`` explicitly when already computed. + client_ip: Client IP for IP-based access control. If None, falls back to + auth context. Pass explicitly from request handlers for safety. + Note: If client_ip is None and auth context is not set, IP filtering is skipped. + This is intentional for internal callers but may indicate a bug if called + from a request handler without proper context setup. """ + # Use explicit client_ip if provided, otherwise try auth context if client_ip is None: client_ip = _get_client_ip_from_context() if client_ip is None: verbose_logger.debug( - "MCP _get_allowed_mcp_servers: client IP unknown; skipping public-internet IP filter." + "MCP _get_allowed_mcp_servers called without client_ip and no auth context. " + "IP filtering will be skipped. This is expected for internal calls." ) allowed_mcp_server_ids = ( @@ -1119,29 +1128,15 @@ if MCP_AVAILABLE: return server_auth_header, extra_headers - async def _merge_gateway_initialize_instructions( + def _merge_gateway_initialize_instructions( allowed_mcp_servers: List[MCPServer], - user_api_key_auth: Optional[UserAPIKeyAuth], - mcp_auth_header: Optional[str], - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], - oauth2_headers: Optional[Dict[str, str]], - raw_headers: Optional[Dict[str, str]], ) -> Optional[str]: - """Merge ``instructions`` for gateway ``initialize``: YAML/API overrides upstream.""" + """YAML/DB override, else in-memory upstream text from list_tools / health_check / call_tool.""" if not allowed_mcp_servers: return None - _has_oauth2_server = any( - getattr(s, "auth_type", None) == MCPAuth.oauth2 - for s in allowed_mcp_servers - ) - _prefetched_oauth_creds = ( - await _prefetch_oauth_creds_for_user(user_api_key_auth) - if _has_oauth2_server - else {} - ) - - async def _one(server: MCPServer) -> Optional[Tuple[str, str]]: + texts: List[Tuple[str, str]] = [] + for server in allowed_mcp_servers: label = ( server.alias or server.server_name @@ -1150,50 +1145,16 @@ if MCP_AVAILABLE: or "mcp" ) if server.instructions and server.instructions.strip(): - return (label, server.instructions.strip()) + texts.append((label, server.instructions.strip())) + continue if server.spec_path: - return None - - server_auth_header, extra_headers = _prepare_mcp_server_headers( - server=server, - mcp_server_auth_headers=mcp_server_auth_headers, - mcp_auth_header=mcp_auth_header, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, + continue + cached = global_mcp_server_manager.get_upstream_initialize_instructions( + server.server_id ) - if extra_headers is None and server.auth_type == MCPAuth.oauth2: - extra_headers = await _get_user_oauth_extra_headers_from_db( - server, - user_api_key_auth, - prefetched_creds=_prefetched_oauth_creds, - ) - try: - if server.static_headers: - if extra_headers is None: - extra_headers = {} - extra_headers.update(server.static_headers) - stdio_env = global_mcp_server_manager._build_stdio_env( - server, raw_headers - ) - client = await global_mcp_server_manager._create_mcp_client( - server=server, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - stdio_env=stdio_env, - ) - text = await client.fetch_upstream_initialize_instructions() - if text and text.strip(): - return (label, text.strip()) - except Exception as e: - verbose_logger.debug( - "MCP gateway: upstream instructions fetch failed for %s: %s", - server.name, - e, - ) - return None + if cached and cached.strip(): + texts.append((label, cached.strip())) - pairs = await asyncio.gather(*(_one(s) for s in allowed_mcp_servers)) - texts = [p for p in pairs if p is not None] if not texts: return None if len(texts) == 1: @@ -1204,25 +1165,13 @@ if MCP_AVAILABLE: user_api_key_auth: Optional[UserAPIKeyAuth], mcp_servers: Optional[List[str]], client_ip: Optional[str], - mcp_auth_header: Optional[str], - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], - oauth2_headers: Optional[Dict[str, str]], - raw_headers: Optional[Dict[str, str]], ) -> contextvars.Token[Optional[str]]: - """Resolve merged gateway ``instructions``; return ContextVar token to reset.""" allowed = await _get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip, ) - merged = await _merge_gateway_initialize_instructions( - allowed_mcp_servers=allowed, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - ) + merged = _merge_gateway_initialize_instructions(allowed_mcp_servers=allowed) return _mcp_gateway_initialize_instructions.set(merged) async def _get_tools_from_mcp_servers( # noqa: PLR0915 @@ -2796,10 +2745,6 @@ if MCP_AVAILABLE: user_api_key_auth, mcp_servers, _client_ip, - mcp_auth_header, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, ) try: await session_manager.handle_request(scope, receive, send) @@ -2867,10 +2812,6 @@ if MCP_AVAILABLE: user_api_key_auth, mcp_servers, _sse_client_ip, - mcp_auth_header, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, ) try: await sse_session_manager.handle_request(scope, receive, send) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 6ba8d0b68a3..f13dc5efd84 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1137,7 +1137,6 @@ 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 - # Shown to MCP clients in InitializeResult.instructions (optional) instructions: Optional[str] = None # Stdio-specific fields command: Optional[str] = None diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index a728d912715..9965c003b0a 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -289,7 +289,7 @@ model LiteLLM_MCPServerTable { server_name String? alias String? description String? - instructions String? // MCP InitializeResult.instructions (optional) + instructions String? url String? spec_path String? transport String @default("sse") diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 805494b1854..81fb424c153 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -27,7 +27,6 @@ class MCPServer(BaseModel): spec_path: Optional[str] = None auth_type: Optional[MCPAuthType] = None authentication_token: Optional[str] = None - # Optional text returned on MCP `initialize` (InitializeResult.instructions) instructions: Optional[str] = None mcp_info: Optional[MCPInfo] = None extra_headers: Optional[ diff --git a/schema.prisma b/schema.prisma index a728d912715..9965c003b0a 100644 --- a/schema.prisma +++ b/schema.prisma @@ -289,7 +289,7 @@ model LiteLLM_MCPServerTable { server_name String? alias String? description String? - instructions String? // MCP InitializeResult.instructions (optional) + instructions String? url String? spec_path String? transport String @default("sse") From 7e656f4329becd0164ccedd56bf102f452fdcd72 Mon Sep 17 00:00:00 2001 From: Milan Date: Tue, 14 Apr 2026 15:54:27 +0300 Subject: [PATCH 11/34] test: add unit tests for MCP initialize instructions feature MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Extend existing test modules with coverage for the instructions merge logic, upstream cache, ContextVar-based injection, and client-side capture — following each file's established patterns. Made-with: Cursor --- .../proxy/_experimental/mcp_server/server.py | 26 ++- .../test_mcp_client.py | 75 ++++++++ .../mcp_server/test_mcp_server.py | 175 ++++++++++++++++++ .../mcp_server/test_mcp_server_manager.py | 72 +++++++ 4 files changed, 334 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 864b17afe33..b7dbdeed5fc 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -6,7 +6,6 @@ LiteLLM MCP Server Routes import asyncio import contextlib -import contextvars import time import types import traceback @@ -1161,18 +1160,23 @@ if MCP_AVAILABLE: return texts[0][1] return "\n\n---\n\n".join(f"[{lbl}]\n{txt}" for lbl, txt in texts) - async def _set_mcp_gateway_initialize_instructions_token( + @contextlib.asynccontextmanager + async def _gateway_initialize_instructions_request_scope( user_api_key_auth: Optional[UserAPIKeyAuth], mcp_servers: Optional[List[str]], client_ip: Optional[str], - ) -> contextvars.Token[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) - return _mcp_gateway_initialize_instructions.set(merged) + 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], @@ -2741,15 +2745,12 @@ if MCP_AVAILABLE: # Request was fully handled (e.g., DELETE on non-existent session) return - _instr_tok = await _set_mcp_gateway_initialize_instructions_token( + async with _gateway_initialize_instructions_request_scope( user_api_key_auth, mcp_servers, _client_ip, - ) - try: + ): await session_manager.handle_request(scope, receive, send) - finally: - _mcp_gateway_initialize_instructions.reset(_instr_tok) except HTTPException: # Re-raise HTTP exceptions to preserve status codes and details raise @@ -2808,15 +2809,12 @@ if MCP_AVAILABLE: await initialize_session_managers() await asyncio.sleep(0.1) - _sse_instr_tok = await _set_mcp_gateway_initialize_instructions_token( + async with _gateway_initialize_instructions_request_scope( user_api_key_auth, mcp_servers, _sse_client_ip, - ) - try: + ): await sse_session_manager.handle_request(scope, receive, send) - finally: - _mcp_gateway_initialize_instructions.reset(_sse_instr_tok) except Exception as e: verbose_logger.exception(f"Error handling MCP request: {e}") # Instead of re-raising, try to send a graceful error response diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index 13a09f54e68..46d483c248f 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -312,5 +312,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__]) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 384d428888f..acba06afee4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -2421,3 +2421,178 @@ 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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 656a9c616e8..503ef71173f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -2483,5 +2483,77 @@ 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.get_upstream_initialize_instructions("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.get_upstream_initialize_instructions("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.get_upstream_initialize_instructions("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.get_upstream_initialize_instructions("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.get_upstream_initialize_instructions("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__]) From e7c630ed1998174e00e660a104fe96c9829395cc Mon Sep 17 00:00:00 2001 From: Milan Date: Tue, 14 Apr 2026 15:59:00 +0300 Subject: [PATCH 12/34] refactor: inline get_upstream_initialize_instructions Remove the trivial one-line wrapper and access the dict directly. Made-with: Cursor --- .../_experimental/mcp_server/mcp_server_manager.py | 3 --- litellm/proxy/_experimental/mcp_server/server.py | 2 +- .../mcp_server/test_mcp_server_manager.py | 10 +++++----- 3 files changed, 6 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 750c9204c55..68b858868ad 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -186,9 +186,6 @@ class MCPServerManager: """ self._upstream_initialize_instructions_by_server_id: Dict[str, str] = {} - def get_upstream_initialize_instructions(self, server_id: str) -> Optional[str]: - return self._upstream_initialize_instructions_by_server_id.get(server_id) - def _remember_upstream_initialize_instructions( self, server: MCPServer, client: MCPClient ) -> None: diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index b7dbdeed5fc..adeacc06f8a 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1148,7 +1148,7 @@ if MCP_AVAILABLE: continue if server.spec_path: continue - cached = global_mcp_server_manager.get_upstream_initialize_instructions( + cached = global_mcp_server_manager._upstream_initialize_instructions_by_server_id.get( server.server_id ) if cached and cached.strip(): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 503ef71173f..aa95836a927 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -2494,7 +2494,7 @@ class TestMCPServerManagerUpstreamInstructionsCache: def test_get_returns_none_when_empty(self): """Empty cache returns None for any key.""" manager = MCPServerManager() - assert manager.get_upstream_initialize_instructions("nonexistent") is None + 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.""" @@ -2502,7 +2502,7 @@ class TestMCPServerManagerUpstreamInstructionsCache: 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.get_upstream_initialize_instructions("srv") == "hello" + assert manager._upstream_initialize_instructions_by_server_id.get("srv") == "hello" def test_remember_ignores_empty_string(self): """Whitespace-only instructions are not stored.""" @@ -2510,7 +2510,7 @@ class TestMCPServerManagerUpstreamInstructionsCache: fake_server = MagicMock(server_id="srv") fake_client = MagicMock(_last_initialize_instructions=" ") manager._remember_upstream_initialize_instructions(fake_server, fake_client) - assert manager.get_upstream_initialize_instructions("srv") is None + assert manager._upstream_initialize_instructions_by_server_id.get("srv") is None def test_remember_ignores_none(self): """None instructions are not stored.""" @@ -2518,7 +2518,7 @@ class TestMCPServerManagerUpstreamInstructionsCache: fake_server = MagicMock(server_id="srv") fake_client = MagicMock(_last_initialize_instructions=None) manager._remember_upstream_initialize_instructions(fake_server, fake_client) - assert manager.get_upstream_initialize_instructions("srv") is None + 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): @@ -2533,7 +2533,7 @@ class TestMCPServerManagerUpstreamInstructionsCache: } } ) - assert manager.get_upstream_initialize_instructions("old") is None + 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): From e20c1148111fe5eaa5c0e73aafbf88da4896e427 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Tue, 14 Apr 2026 11:05:22 -0700 Subject: [PATCH 13/34] fix(mcp): set instructions=None in SigV4BuildFromTable test mocks New MCPServer.instructions field requires a str; MagicMock attributes not explicitly set return a MagicMock object, which fails Pydantic validation. --- .../proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py index 7c142e3a771..eb9f4dde55f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py @@ -838,6 +838,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 +896,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() From 8c505634bd9771b15803ea2fb7fb30edf8cc3aac Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Tue, 14 Apr 2026 11:05:25 -0700 Subject: [PATCH 14/34] chore: sync uv.lock with pyproject.toml (v1.83.6 -> v1.83.7) --- uv.lock | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/uv.lock b/uv.lock index 04224dc5374..54000158017 100644 --- a/uv.lock +++ b/uv.lock @@ -11,7 +11,7 @@ resolution-markers = [ ] [options] -exclude-newer = "2026-04-08T16:01:27.663665Z" +exclude-newer = "2026-04-11T18:05:05.631902Z" exclude-newer-span = "P3D" [manifest] @@ -3602,7 +3602,7 @@ wheels = [ [[package]] name = "litellm" -version = "1.83.6" +version = "1.83.7" source = { editable = "." } dependencies = [ { name = "aiohttp" }, From 92a5ed4c3d620d64f90c161ad2b7a9661f870996 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Tue, 14 Apr 2026 12:15:54 -0700 Subject: [PATCH 15/34] fix(mcp): set instructions=None in test_add_update_server_fallback_to_server_id mock --- tests/mcp_tests/test_mcp_server.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index a4a28215e16..7c5dcc66a83 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -1587,7 +1587,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) From 2b5eb794fca2727aff280c18799aec7813f71121 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Tue, 14 Apr 2026 12:40:55 -0700 Subject: [PATCH 16/34] fix(mcp): set instructions=None in test_add_update_server_with_alias mock --- tests/mcp_tests/test_mcp_server.py | 192 +++++++++++++++++------------ 1 file changed, 115 insertions(+), 77 deletions(-) diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 7c5dcc66a83..d0a4e9d0c9f 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -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) @@ -2151,8 +2171,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 +2714,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 +2790,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 From 6126b47c8655f04a5f9cde119ac9f464e2eceb8c Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Tue, 14 Apr 2026 12:42:48 -0700 Subject: [PATCH 17/34] fix(mcp): set instructions=None in test_add_update_server_without_alias mock --- tests/mcp_tests/test_mcp_server.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index d0a4e9d0c9f..6af07585796 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -1550,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) From c3dbd782f4e1d7705ce66078862640e7c70976db Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:18:18 -0700 Subject: [PATCH 18/34] style: black format llm_http_handler.py --- litellm/llms/custom_httpx/llm_http_handler.py | 37 +++++++++++++------ 1 file changed, 25 insertions(+), 12 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 60efa45df62..ba76e13b6a8 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1830,7 +1830,9 @@ class BaseLLMHTTPHandler: api_key: Optional[str], model: str, ) -> httpx.Response: - max_attempts = max(provider_config.max_retry_on_anthropic_messages_http_error, 1) + 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): @@ -1846,8 +1848,10 @@ class BaseLLMHTTPHandler: 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 + 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( @@ -4559,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, @@ -4677,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, @@ -5173,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. @@ -5385,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. @@ -5625,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. From 563e05ebfa3b315aa5263e64798b6b8e2884b2d2 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:18:24 -0700 Subject: [PATCH 19/34] style: black format _types.py --- litellm/proxy/_types.py | 96 ++++++++++++++++++++--------------------- 1 file changed, 48 insertions(+), 48 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index f13dc5efd84..7fff640c497 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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): @@ -1577,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") @@ -1605,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 @@ -1700,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 @@ -1805,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") @@ -2149,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): @@ -2510,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. @@ -2651,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 @@ -3808,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): @@ -4065,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): @@ -4229,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 From d8dbb46dcf11aa87a950719b0ffa11ee7f978681 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:18:29 -0700 Subject: [PATCH 20/34] style: black format mcp_server_manager.py --- .../proxy/_experimental/mcp_server/mcp_server_manager.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 68b858868ad..bca7aa45207 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -191,9 +191,9 @@ class MCPServerManager: ) -> 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() - ) + self._upstream_initialize_instructions_by_server_id[server.server_id] = str( + raw + ).strip() def get_registry(self) -> Dict[str, MCPServer]: """ From 65061b1e3cc337db60e99bdb72a0a1c639e7e2cb Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:18:33 -0700 Subject: [PATCH 21/34] style: black format mcp server.py --- litellm/proxy/_experimental/mcp_server/server.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index adeacc06f8a..cefe2ef763b 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1045,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: From 0acd05207b814632b5d2aa6213e408b16167a5c5 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:18:36 -0700 Subject: [PATCH 22/34] style: black format health_check.py --- litellm/proxy/health_check.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 5377ff6c320..1518ed66ab1 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -246,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 From e5adafc7689f08d81eddeaca17d5dae8ce9e5511 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:18:41 -0700 Subject: [PATCH 23/34] style: black format anthropic_messages transformation.py --- .../llms/base_llm/anthropic_messages/transformation.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/litellm/llms/base_llm/anthropic_messages/transformation.py b/litellm/llms/base_llm/anthropic_messages/transformation.py index 733faca8532..49aa563781f 100644 --- a/litellm/llms/base_llm/anthropic_messages/transformation.py +++ b/litellm/llms/base_llm/anthropic_messages/transformation.py @@ -141,8 +141,9 @@ class BaseAnthropicMessagesConfig(ABC): is_anthropic_invalid_thinking_signature_error, ) - return e.response.status_code == 400 and is_anthropic_invalid_thinking_signature_error( - e.response.text + return ( + e.response.status_code == 400 + and is_anthropic_invalid_thinking_signature_error(e.response.text) ) def transform_anthropic_messages_request_on_http_error( @@ -156,8 +157,9 @@ class BaseAnthropicMessagesConfig(ABC): strip_thinking_blocks_from_anthropic_messages_request_dict, ) - if e.response.status_code == 400 and is_anthropic_invalid_thinking_signature_error( - e.response.text + 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 From 107003a7137de0d0c20a8a747e638cc1da3d6936 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:18:45 -0700 Subject: [PATCH 24/34] style: black format model_param_helper.py --- litellm/litellm_core_utils/model_param_helper.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/litellm/litellm_core_utils/model_param_helper.py b/litellm/litellm_core_utils/model_param_helper.py index 35e744be3a6..b4fa5cb60aa 100644 --- a/litellm/litellm_core_utils/model_param_helper.py +++ b/litellm/litellm_core_utils/model_param_helper.py @@ -101,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) From 13952b0b1b40626ef0d3b076e37133eed6e6536f Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:18:48 -0700 Subject: [PATCH 25/34] style: black format types/mcp_server/mcp_server_manager.py --- litellm/types/mcp_server/mcp_server_manager.py | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 81fb424c153..ace8c8a4188 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -29,19 +29,19 @@ class MCPServer(BaseModel): 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 From 3847a59d79d36b336515d53b43a0ccfb7f276d0d Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:18:52 -0700 Subject: [PATCH 26/34] style: black format test_model_param_helper.py --- tests/test_litellm/test_model_param_helper.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/test_litellm/test_model_param_helper.py b/tests/test_litellm/test_model_param_helper.py index 2012abec547..a62779aeab1 100644 --- a/tests/test_litellm/test_model_param_helper.py +++ b/tests/test_litellm/test_model_param_helper.py @@ -57,6 +57,6 @@ def test_get_all_llm_api_params_includes_responses_api(): "safety_identifier", } missing = responses_only_params - all_params - assert missing == set(), ( - f"Responses-API kwargs missing from cache-key allow-list: {sorted(missing)}" - ) + assert ( + missing == set() + ), f"Responses-API kwargs missing from cache-key allow-list: {sorted(missing)}" From f2a1dbe7c9854b5527474d6a2d1168b4127e098d Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:18:56 -0700 Subject: [PATCH 27/34] style: black format test_health_check_max_tokens.py --- .../proxy/test_health_check_max_tokens.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/tests/test_litellm/proxy/test_health_check_max_tokens.py b/tests/test_litellm/proxy/test_health_check_max_tokens.py index bb125f764b8..4d417b40b59 100644 --- a/tests/test_litellm/proxy/test_health_check_max_tokens.py +++ b/tests/test_litellm/proxy/test_health_check_max_tokens.py @@ -53,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( From 93a90a53bed14d56844b1a608eef3823798e0588 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:19:01 -0700 Subject: [PATCH 28/34] style: black format test_mcp_client.py --- .../test_mcp_client.py | 23 ++++++++++--------- 1 file changed, 12 insertions(+), 11 deletions(-) diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index 46d483c248f..dee689708c3 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -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" From c8a0fe193f1c5e53c9b7cf725a013878990c132e Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:19:04 -0700 Subject: [PATCH 29/34] style: black format test_unit_test_caching.py --- tests/local_testing/test_unit_test_caching.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/tests/local_testing/test_unit_test_caching.py b/tests/local_testing/test_unit_test_caching.py index e4ee65a2aa2..e25b75e658f 100644 --- a/tests/local_testing/test_unit_test_caching.py +++ b/tests/local_testing/test_unit_test_caching.py @@ -159,9 +159,7 @@ def test_get_cache_key_responses_api(): 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" - ) + 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) @@ -177,9 +175,9 @@ def test_get_cache_key_responses_api(): ]: 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" - ) + 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(): From 9a154a3be773709c4579c457c63eb9d8374700aa Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:19:08 -0700 Subject: [PATCH 30/34] style: black format test_mcp_sigv4_auth.py --- .../mcp_server/test_mcp_sigv4_auth.py | 104 +++++++++++------- 1 file changed, 66 insertions(+), 38 deletions(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py index eb9f4dde55f..32b988ddb22 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py @@ -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" @@ -936,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) @@ -958,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) @@ -990,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" From f7689465496523426153f2a65ffa6614a5b30b6c Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:19:12 -0700 Subject: [PATCH 31/34] style: black format test_anthropic_common_utils.py --- .../anthropic/test_anthropic_common_utils.py | 27 ++++++++++++------- 1 file changed, 18 insertions(+), 9 deletions(-) diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index d48d7716a8e..2f57ce5d180 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -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( @@ -1246,11 +1251,15 @@ class TestAnthropicThinkingSignatureSelfHeal: 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 + 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 + assert ( + config.should_retry_anthropic_messages_on_http_error(err_500, {}) is False + ) data = { "model": "claude-sonnet-4-20250514", From fcd71e0026bfd9885f96f6d560872d4c32063142 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:19:17 -0700 Subject: [PATCH 32/34] style: black format test_mcp_server_manager.py --- .../mcp_server/test_mcp_server_manager.py | 169 ++++++++++++------ 1 file changed, 112 insertions(+), 57 deletions(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index aa95836a927..ac5349e7105 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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 @@ -2494,7 +2544,10 @@ class TestMCPServerManagerUpstreamInstructionsCache: 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 + 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.""" @@ -2502,7 +2555,9 @@ class TestMCPServerManagerUpstreamInstructionsCache: 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" + assert ( + manager._upstream_initialize_instructions_by_server_id.get("srv") == "hello" + ) def test_remember_ignores_empty_string(self): """Whitespace-only instructions are not stored.""" From 537e72c74258da43a4996c7b6e28914b5b4d6ea9 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:19:21 -0700 Subject: [PATCH 33/34] style: black format test_mcp_server.py --- .../mcp_server/test_mcp_server.py | 487 +++++++++++------- 1 file changed, 305 insertions(+), 182 deletions(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index acba06afee4..9df6408b0d7 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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]) @@ -2481,24 +2581,34 @@ class TestMergeGatewayInitializeInstructions: 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" + + 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) + 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" + + 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) + 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.""" @@ -2512,8 +2622,12 @@ class TestMergeGatewayInitializeInstructions: 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") + 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 @@ -2532,17 +2646,26 @@ class TestMergeGatewayInitializeInstructions: 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" + + 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_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) + global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop( + "c", None + ) class TestGatewayCreateInitializationOptions: From 26136708bb1b16eace219542e6ece5595871b981 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:48:13 -0700 Subject: [PATCH 34/34] chore: trigger CI re-evaluation