diff --git a/Makefile b/Makefile
index 8b657dcb465..e9b2fb9d8f1 100644
--- a/Makefile
+++ b/Makefile
@@ -176,6 +176,8 @@ lint-ruff-FULL-dev: install-dev
if [ -n "$$files" ]; then echo "$$files" | xargs $(UV_RUN) ruff check; \
else echo "No changed .py files to check."; fi
+lint-basedpyright lint-basedpyright-budget-update: export NODE_OPTIONS := --max-old-space-size=12288
+
lint-basedpyright: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging
diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json
index db3c2502e94..513115b7088 100644
--- a/basedpyright-code-budget.json
+++ b/basedpyright-code-budget.json
@@ -1,12 +1,12 @@
{
"reportAny": {
- "limit": 33216
+ "limit": 33129
},
"reportArgumentType": {
- "limit": 2648
+ "limit": 2645
},
"reportAssignmentType": {
- "limit": 330
+ "limit": 328
},
"reportAttributeAccessIssue": {
"limit": 516
@@ -18,13 +18,13 @@
"limit": 59
},
"reportDeprecated": {
- "limit": 326
+ "limit": 324
},
"reportDuplicateImport": {
"limit": 42
},
"reportExplicitAny": {
- "limit": 10228
+ "limit": 10227
},
"reportFunctionMemberAccess": {
"limit": 11
@@ -54,10 +54,10 @@
"limit": 0
},
"reportMissingParameterType": {
- "limit": 5893
+ "limit": 5845
},
"reportMissingTypeArgument": {
- "limit": 15886
+ "limit": 15846
},
"reportMissingTypeStubs": {
"limit": 41
@@ -72,7 +72,7 @@
"limit": 0
},
"reportOptionalMemberAccess": {
- "limit": 1085
+ "limit": 1073
},
"reportOptionalOperand": {
"limit": 0
@@ -84,13 +84,13 @@
"limit": 77
},
"reportPrivateUsage": {
- "limit": 2438
+ "limit": 2437
},
"reportRedeclaration": {
"limit": 12
},
"reportReturnType": {
- "limit": 225
+ "limit": 217
},
"reportTypedDictNotRequiredAccess": {
"limit": 27
@@ -99,19 +99,19 @@
"limit": 0
},
"reportUnknownArgumentType": {
- "limit": 45567
+ "limit": 45498
},
"reportUnknownLambdaType": {
- "limit": 113
+ "limit": 109
},
"reportUnknownMemberType": {
- "limit": 40525
+ "limit": 40458
},
"reportUnknownParameterType": {
- "limit": 20384
+ "limit": 20302
},
"reportUnknownVariableType": {
- "limit": 32099
+ "limit": 32026
},
"reportUnnecessaryCast": {
"limit": 177
@@ -123,7 +123,7 @@
"limit": 7
},
"reportUnnecessaryIsInstance": {
- "limit": 1206
+ "limit": 1205
},
"reportUntypedBaseClass": {
"limit": 165
@@ -138,7 +138,7 @@
"limit": 206
},
"reportUnusedImport": {
- "limit": 1005
+ "limit": 1001
},
"reportUnusedVariable": {
"limit": 1297
diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py
index 4e3d94e2ab3..d8ce48f05de 100644
--- a/litellm/litellm_core_utils/prompt_templates/factory.py
+++ b/litellm/litellm_core_utils/prompt_templates/factory.py
@@ -1267,7 +1267,7 @@ def _get_dummy_thought_signature() -> str:
def convert_to_gemini_tool_call_invoke(
message: ChatCompletionAssistantMessage,
model: Optional[str] = None,
- custom_llm_provider: Optional[str] = None,
+ forward_function_call_id: bool = False,
) -> List[VertexPartType]:
"""
OpenAI tool invokes:
@@ -1317,16 +1317,12 @@ def convert_to_gemini_tool_call_invoke(
VertexGeminiConfig,
)
- forward_tool_call_id = bool(
- model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider)
- )
-
if tool_calls is not None:
for idx, tool in enumerate(tool_calls):
if "function" in tool:
gemini_function_call: Optional[VertexFunctionCall] = _gemini_tool_call_invoke_helper(
function_call_params=tool["function"],
- tool_call_id=(tool.get("id") if forward_tool_call_id else None),
+ tool_call_id=(tool.get("id") if forward_function_call_id else None),
)
if gemini_function_call is not None:
part_dict: VertexPartType = {"function_call": gemini_function_call}
@@ -1378,8 +1374,7 @@ def convert_to_gemini_tool_call_invoke(
def convert_to_gemini_tool_call_result(
message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage],
last_message_with_tool_calls: Optional[dict],
- model: Optional[str] = None,
- custom_llm_provider: Optional[str] = None,
+ forward_function_call_id: bool = False,
) -> Union[VertexPartType, List[VertexPartType]]:
"""
OpenAI message with a tool result looks like:
@@ -1501,14 +1496,8 @@ def convert_to_gemini_tool_call_result(
name = tool.get("function", {}).get("name", "")
# Echo the OpenAI tool_call_id on functionResponse (strip thought-signature suffix).
- # Only Google AI Studio Gemini 3+ accepts `id` on function_response parts.
- # Vertex AI and older Gemini models reject the field with HTTP 400.
- from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
- VertexGeminiConfig,
- )
-
gemini_call_id: Optional[str] = None
- if model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider):
+ if forward_function_call_id:
raw_tool_call_id = message.get("tool_call_id")
if raw_tool_call_id and isinstance(raw_tool_call_id, str):
stripped_id = raw_tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0]
diff --git a/litellm/llms/bedrock/chat/__init__.py b/litellm/llms/bedrock/chat/__init__.py
index c1323b9192a..37dcb270743 100644
--- a/litellm/llms/bedrock/chat/__init__.py
+++ b/litellm/llms/bedrock/chat/__init__.py
@@ -5,7 +5,6 @@ from .invoke_handler import (
AmazonAnthropicClaudeStreamDecoder,
AmazonDeepSeekR1StreamDecoder,
AWSEventStreamDecoder,
- BedrockLLM,
)
diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py
index 4c256be1ab8..c28627d5aec 100644
--- a/litellm/llms/bedrock/chat/invoke_handler.py
+++ b/litellm/llms/bedrock/chat/invoke_handler.py
@@ -1,19 +1,10 @@
-"""
-TODO: DELETE FILE. Bedrock LLM is no longer used. Goto `litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py`
-"""
-
-import copy
-import time
import types
-from functools import partial
from typing import (
AsyncIterator,
- Callable,
Iterator,
Optional,
Tuple,
cast,
- get_args,
)
import httpx # type: ignore
@@ -25,16 +16,6 @@ from litellm.caching.caching import InMemoryCache
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.litellm_core_utils.core_helpers import map_finish_reason
from litellm.litellm_core_utils.litellm_logging import Logging
-from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
-from litellm.litellm_core_utils.prompt_templates.factory import (
- cohere_message_pt,
- construct_tool_use_system_prompt,
- contains_tag,
- custom_prompt,
- extract_between_tags,
- parse_xml_params,
- prompt_factory,
-)
from litellm.llms.anthropic.chat.handler import (
ModelResponseIterator as AnthropicModelResponseIterator,
)
@@ -64,12 +45,9 @@ from litellm.types.utils import (
StreamingChoices,
Usage,
)
-from litellm.utils import CustomStreamWrapper, get_secret
-from ..base_aws_llm import BaseAWSLLM
from ..common_utils import (
BedrockError,
- ModelResponseIterator,
build_bedrock_stream_error,
get_bedrock_response_stream_shape,
get_bedrock_tool_name,
@@ -77,9 +55,6 @@ from ..common_utils import (
bedrock_tool_name_mappings: InMemoryCache = InMemoryCache(max_size_in_memory=50, default_ttl=600)
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
-from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import (
- AmazonBedrockOpenAIConfig,
-)
converse_config = AmazonConverseConfig()
@@ -351,932 +326,6 @@ def make_sync_call(
raise BedrockError(status_code=500, message=str(e))
-class BedrockLLM(BaseAWSLLM):
- """
- Example call
-
- ```
- curl --location --request POST 'https://bedrock-runtime.{aws_region_name}.amazonaws.com/model/{bedrock_model_name}/invoke' \
- --header 'Content-Type: application/json' \
- --header 'Accept: application/json' \
- --user "$AWS_ACCESS_KEY_ID":"$AWS_SECRET_ACCESS_KEY" \
- --aws-sigv4 "aws:amz:us-east-1:bedrock" \
- --data-raw '{
- "prompt": "Hi",
- "temperature": 0,
- "p": 0.9,
- "max_tokens": 4096
- }'
- ```
- """
-
- def __init__(self) -> None:
- super().__init__()
-
- @staticmethod
- def is_claude_messages_api_model(model: str) -> bool:
- """
- Check if the model uses the Claude Messages API (Claude 3+).
-
- Handles:
- - Regional prefixes: eu.anthropic.claude-*, us.anthropic.claude-*
- - Claude 3 models: claude-3-haiku, claude-3-sonnet, claude-3-opus, claude-3-5-*, claude-3-7-*
- - Claude 4 models: claude-opus-4, claude-sonnet-4, claude-haiku-4
- """
- # Normalize model string to lowercase for matching
- model_lower = model.lower()
-
- # Claude 3+ indicators (all use Messages API)
- messages_api_indicators = [
- "claude-3", # Claude 3.x models
- "claude-opus-4", # Claude Opus 4
- "claude-sonnet-4", # Claude Sonnet 4
- "claude-haiku-4", # Claude Haiku 4
- ]
-
- return any(indicator in model_lower for indicator in messages_api_indicators)
-
- def convert_messages_to_prompt(self, model, messages, provider, custom_prompt_dict) -> Tuple[str, Optional[list]]:
- # handle anthropic prompts and amazon titan prompts
- prompt = ""
- chat_history: Optional[list] = None
- ## CUSTOM PROMPT
- if model in custom_prompt_dict:
- # check if the model has a registered custom prompt
- model_prompt_details = custom_prompt_dict[model]
- prompt = custom_prompt(
- role_dict=model_prompt_details["roles"],
- initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""),
- final_prompt_value=model_prompt_details.get("final_prompt_value", ""),
- messages=messages,
- )
- return prompt, None
- ## ELSE
- if provider == "anthropic" or provider == "amazon":
- prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock")
- elif provider == "mistral":
- prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock")
- elif provider == "meta" or provider == "llama":
- prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock")
- elif provider == "openai":
- # OpenAI uses messages directly, no prompt conversion needed
- # Return empty prompt as it won't be used
- prompt = ""
- elif provider == "cohere":
- prompt, chat_history = cohere_message_pt(messages=messages)
- else:
- prompt = ""
- for message in messages:
- if "role" in message:
- if message["role"] == "user":
- prompt += f"{message['content']}"
- else:
- prompt += f"{message['content']}"
- else:
- prompt += f"{message['content']}"
- return prompt, chat_history # type: ignore
-
- def process_response(
- self,
- model: str,
- response: httpx.Response,
- model_response: ModelResponse,
- stream: Optional[bool],
- logging_obj: Logging,
- optional_params: dict,
- api_key: str,
- data: Union[dict, str],
- messages: List,
- print_verbose,
- encoding,
- ) -> Union[ModelResponse, CustomStreamWrapper]:
- provider = self.get_bedrock_invoke_provider(model)
- ## LOGGING
- logging_obj.post_call(
- input=messages,
- api_key=api_key,
- original_response=response.text,
- additional_args={"complete_input_dict": data},
- )
- print_verbose(f"raw model_response: {response.text}")
-
- ## RESPONSE OBJECT
- try:
- completion_response = response.json()
- except Exception:
- raise BedrockError(message=response.text, status_code=422)
-
- outputText: Optional[str] = None
- try:
- if provider == "cohere":
- if "text" in completion_response:
- outputText = completion_response["text"] # type: ignore
- elif "generations" in completion_response:
- outputText = completion_response["generations"][0]["text"]
- model_response.choices[0].finish_reason = map_finish_reason(
- completion_response["generations"][0]["finish_reason"]
- )
- elif provider == "anthropic":
- if self.is_claude_messages_api_model(model):
- json_schemas: dict = {}
- _is_function_call = False
- ## Handle Tool Calling
- if "tools" in optional_params:
- _is_function_call = True
- for tool in optional_params["tools"]:
- json_schemas[tool["function"]["name"]] = tool["function"].get("parameters", None)
- outputText = completion_response.get("content")[0].get("text", None)
- if outputText is not None and contains_tag("invoke", outputText): # OUTPUT PARSE FUNCTION CALL
- function_name = extract_between_tags("tool_name", outputText)[0]
- function_arguments_str = extract_between_tags("invoke", outputText)[0].strip()
- function_arguments_str = f"{function_arguments_str}"
- function_arguments = parse_xml_params(
- function_arguments_str,
- json_schema=json_schemas.get(
- function_name, None
- ), # check if we have a json schema for this function name)
- )
- _message = litellm.Message(
- tool_calls=[
- {
- "id": f"call_{uuid.uuid4()}",
- "type": "function",
- "function": {
- "name": function_name,
- "arguments": json.dumps(function_arguments),
- },
- }
- ],
- content=None,
- )
- model_response.choices[0].message = _message # type: ignore
- model_response._hidden_params["original_response"] = (
- outputText # allow user to access raw anthropic tool calling response
- )
- if _is_function_call is True and stream is not None and stream is True:
- print_verbose("INSIDE BEDROCK STREAMING TOOL CALLING CONDITION BLOCK")
- # return an iterator
- streaming_model_response = ModelResponseStream()
- streaming_model_response.choices[0].finish_reason = getattr(
- model_response.choices[0], "finish_reason", "stop"
- )
- # streaming_model_response.choices = [litellm.utils.StreamingChoices()]
- streaming_choice = litellm.utils.StreamingChoices()
- streaming_choice.index = model_response.choices[0].index
- _tool_calls = []
- print_verbose(f"type of model_response.choices[0]: {type(model_response.choices[0])}")
- print_verbose(f"type of streaming_choice: {type(streaming_choice)}")
- if isinstance(model_response.choices[0], litellm.Choices):
- if getattr(
- model_response.choices[0].message, "tool_calls", None
- ) is not None and isinstance(model_response.choices[0].message.tool_calls, list):
- for tool_call in model_response.choices[0].message.tool_calls:
- _tool_call = {**tool_call.dict(), "index": 0}
- _tool_calls.append(_tool_call)
- delta_obj = Delta(
- content=getattr(model_response.choices[0].message, "content", None),
- role=model_response.choices[0].message.role,
- tool_calls=_tool_calls,
- )
- streaming_choice.delta = delta_obj
- streaming_model_response.choices = [streaming_choice]
- completion_stream = ModelResponseIterator(model_response=streaming_model_response)
- print_verbose(
- "Returns anthropic CustomStreamWrapper with 'cached_response' streaming object"
- )
- return litellm.CustomStreamWrapper(
- completion_stream=completion_stream,
- model=model,
- custom_llm_provider="cached_response",
- logging_obj=logging_obj,
- )
-
- model_response.choices[0].finish_reason = map_finish_reason(
- completion_response.get("stop_reason", "")
- )
- _usage = litellm.Usage(
- prompt_tokens=completion_response["usage"]["input_tokens"],
- completion_tokens=completion_response["usage"]["output_tokens"],
- total_tokens=completion_response["usage"]["input_tokens"]
- + completion_response["usage"]["output_tokens"],
- )
- setattr(model_response, "usage", _usage)
- else:
- outputText = completion_response["completion"]
-
- model_response.choices[0].finish_reason = completion_response["stop_reason"]
- elif provider == "ai21":
- outputText = completion_response.get("completions")[0].get("data").get("text")
- elif provider == "meta" or provider == "llama":
- outputText = completion_response["generation"]
- elif provider == "openai":
- # OpenAI imported models use OpenAI Chat Completions format
- if "choices" in completion_response and len(completion_response["choices"]) > 0:
- choice = completion_response["choices"][0]
- if "message" in choice:
- outputText = choice["message"].get("content")
- elif "text" in choice: # fallback for completion format
- outputText = choice["text"]
-
- # Set finish reason
- if "finish_reason" in choice:
- model_response.choices[0].finish_reason = map_finish_reason(choice["finish_reason"])
-
- # Set usage if available
- if "usage" in completion_response:
- usage = completion_response["usage"]
- _usage = litellm.Usage(
- prompt_tokens=usage.get("prompt_tokens", 0),
- completion_tokens=usage.get("completion_tokens", 0),
- total_tokens=usage.get("total_tokens", 0),
- )
- setattr(model_response, "usage", _usage)
- elif provider == "mistral":
- outputText = completion_response["outputs"][0]["text"]
- model_response.choices[0].finish_reason = completion_response["outputs"][0]["stop_reason"]
- else: # amazon titan
- outputText = completion_response.get("results")[0].get("outputText")
- except Exception as e:
- raise BedrockError(
- message="Error processing={}, Received error={}".format(response.text, str(e)),
- status_code=422,
- )
-
- try:
- if (
- outputText is not None
- and len(outputText) > 0
- and hasattr(model_response.choices[0], "message")
- and getattr(model_response.choices[0].message, "tool_calls", None) # type: ignore
- is None
- ):
- model_response.choices[0].message.content = outputText # type: ignore
- elif (
- hasattr(model_response.choices[0], "message")
- and getattr(model_response.choices[0].message, "tool_calls", None) # type: ignore
- is not None
- ):
- pass
- else:
- raise Exception()
- except Exception as e:
- raise BedrockError(
- message="Error parsing received text={}.\nError-{}".format(outputText, str(e)),
- status_code=response.status_code,
- )
-
- if stream and provider == "ai21":
- streaming_model_response = ModelResponseStream()
- streaming_model_response.choices[0].finish_reason = model_response.choices[ # type: ignore
- 0
- ].finish_reason
- # streaming_model_response.choices = [litellm.utils.StreamingChoices()]
- streaming_choice = litellm.utils.StreamingChoices()
- streaming_choice.index = model_response.choices[0].index
- delta_obj = litellm.utils.Delta(
- content=getattr(model_response.choices[0].message, "content", None), # type: ignore
- role=model_response.choices[0].message.role, # type: ignore
- )
- streaming_choice.delta = delta_obj
- streaming_model_response.choices = [streaming_choice]
- mri = ModelResponseIterator(model_response=streaming_model_response)
- return CustomStreamWrapper(
- completion_stream=mri,
- model=model,
- custom_llm_provider="cached_response",
- logging_obj=logging_obj,
- )
-
- ## CALCULATING USAGE - bedrock returns usage in the headers
- # Skip if usage was already set (e.g., from JSON response for OpenAI provider)
- if not hasattr(model_response, "usage") or getattr(model_response, "usage", None) is None:
- bedrock_input_tokens = response.headers.get("x-amzn-bedrock-input-token-count", None)
- bedrock_output_tokens = response.headers.get("x-amzn-bedrock-output-token-count", None)
-
- prompt_tokens = int(bedrock_input_tokens or litellm.token_counter(messages=messages))
-
- completion_tokens = int(
- bedrock_output_tokens
- or litellm.token_counter(
- text=model_response.choices[0].message.content, # type: ignore
- count_response_tokens=True,
- )
- )
-
- model_response.created = int(time.time())
- model_response.model = model
- usage = Usage(
- prompt_tokens=prompt_tokens,
- completion_tokens=completion_tokens,
- total_tokens=prompt_tokens + completion_tokens,
- )
- setattr(model_response, "usage", usage)
- else:
- # Ensure created and model are set even if usage was already set
- model_response.created = int(time.time())
- model_response.model = model
-
- return model_response
-
- def completion(
- self,
- model: str,
- messages: list,
- api_base: Optional[str],
- custom_prompt_dict: dict,
- model_response: ModelResponse,
- print_verbose: Callable,
- encoding,
- logging_obj: Logging,
- optional_params: dict,
- acompletion: bool,
- timeout: Optional[Union[float, httpx.Timeout]],
- litellm_params=None,
- logger_fn=None,
- extra_headers: Optional[dict] = None,
- client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None,
- ) -> Union[ModelResponse, CustomStreamWrapper]:
- try:
- from botocore.credentials import Credentials
- except ImportError:
- raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
-
- ## SETUP ##
- stream = optional_params.pop("stream", None)
- stream_chunk_size = optional_params.pop("stream_chunk_size", None)
-
- provider = self.get_bedrock_invoke_provider(model)
- modelId = self.get_bedrock_model_id(
- model=model,
- provider=provider,
- optional_params=optional_params,
- )
-
- ## CREDENTIALS ##
- # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them
- aws_secret_access_key = optional_params.pop("aws_secret_access_key", None)
- aws_access_key_id = optional_params.pop("aws_access_key_id", None)
- aws_session_token = optional_params.pop("aws_session_token", None)
- aws_region_name = optional_params.pop("aws_region_name", None)
- aws_role_name = optional_params.pop("aws_role_name", None)
- aws_session_name = optional_params.pop("aws_session_name", None)
- aws_profile_name = optional_params.pop("aws_profile_name", None)
- aws_bedrock_runtime_endpoint = optional_params.pop(
- "aws_bedrock_runtime_endpoint", None
- ) # https://bedrock-runtime.{region_name}.amazonaws.com
- aws_web_identity_token = optional_params.pop("aws_web_identity_token", None)
- aws_sts_endpoint = optional_params.pop("aws_sts_endpoint", None)
- ssl_verify = optional_params.pop("ssl_verify", None)
-
- ### SET REGION NAME ###
- if aws_region_name is None:
- # check env #
- litellm_aws_region_name = get_secret("AWS_REGION_NAME", None)
-
- if litellm_aws_region_name is not None and isinstance(litellm_aws_region_name, str):
- aws_region_name = litellm_aws_region_name
-
- standard_aws_region_name = get_secret("AWS_REGION", None)
- if standard_aws_region_name is not None and isinstance(standard_aws_region_name, str):
- aws_region_name = standard_aws_region_name
-
- if aws_region_name is None:
- aws_region_name = "us-west-2"
-
- credentials: Credentials = self.get_credentials(
- aws_access_key_id=aws_access_key_id,
- aws_secret_access_key=aws_secret_access_key,
- aws_session_token=aws_session_token,
- aws_region_name=aws_region_name,
- aws_session_name=aws_session_name,
- aws_profile_name=aws_profile_name,
- aws_role_name=aws_role_name,
- aws_web_identity_token=aws_web_identity_token,
- aws_sts_endpoint=aws_sts_endpoint,
- ssl_verify=ssl_verify,
- )
-
- ### SET RUNTIME ENDPOINT ###
- endpoint_url, proxy_endpoint_url = self.get_runtime_endpoint(
- api_base=api_base,
- aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
- aws_region_name=aws_region_name,
- )
-
- if (stream is not None and stream is True) and provider != "ai21":
- endpoint_url = f"{endpoint_url}/model/{modelId}/invoke-with-response-stream"
- proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke-with-response-stream"
- else:
- endpoint_url = f"{endpoint_url}/model/{modelId}/invoke"
- proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke"
-
- if acompletion and provider == "anthropic" and self.is_claude_messages_api_model(model):
- if isinstance(client, HTTPHandler):
- client = None
- return self._async_anthropic_messages_completion(
- model=model,
- messages=messages,
- endpoint_url=endpoint_url,
- proxy_endpoint_url=proxy_endpoint_url,
- credentials=credentials,
- aws_region_name=aws_region_name,
- model_response=model_response,
- print_verbose=print_verbose,
- encoding=encoding,
- logging_obj=logging_obj,
- optional_params=optional_params,
- stream=stream,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- extra_headers=extra_headers,
- timeout=timeout,
- client=client,
- stream_chunk_size=stream_chunk_size,
- ) # type: ignore[return-value]
-
- prompt, chat_history = self.convert_messages_to_prompt(model, messages, provider, custom_prompt_dict)
- inference_params = copy.deepcopy(optional_params)
- json_schemas: dict = {}
- if provider == "cohere":
- if model.startswith("cohere.command-r"):
- ## LOAD CONFIG
- config = litellm.AmazonCohereChatConfig().get_config()
- for k, v in config.items():
- if (
- k not in inference_params
- ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
- inference_params[k] = v
- _data = {"message": prompt, **inference_params}
- if chat_history is not None:
- _data["chat_history"] = chat_history
- data = json.dumps(_data)
- else:
- ## LOAD CONFIG
- config = litellm.AmazonCohereConfig.get_config()
- for k, v in config.items():
- if (
- k not in inference_params
- ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
- inference_params[k] = v
- if stream is True:
- inference_params["stream"] = True # cohere requires stream = True in inference params
- data = json.dumps({"prompt": prompt, **inference_params})
- elif provider == "anthropic":
- if self.is_claude_messages_api_model(model):
- # Separate system prompt from rest of message
- system_prompt_idx: list[int] = []
- system_messages: list[str] = []
- for idx, message in enumerate(messages):
- if message["role"] == "system":
- system_messages.append(message["content"])
- system_prompt_idx.append(idx)
- if len(system_prompt_idx) > 0:
- inference_params["system"] = "\n".join(system_messages)
- messages = [i for j, i in enumerate(messages) if j not in system_prompt_idx]
- # Format rest of message according to anthropic guidelines
- messages = prompt_factory(model=model, messages=messages, custom_llm_provider="anthropic_xml") # type: ignore
- ## LOAD CONFIG
- config = litellm.AmazonAnthropicClaudeConfig.get_config()
- for k, v in config.items():
- if (
- k not in inference_params
- ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
- inference_params[k] = v
- ## Handle Tool Calling
- if "tools" in inference_params:
- _is_function_call = True
- for tool in inference_params["tools"]:
- json_schemas[tool["function"]["name"]] = tool["function"].get("parameters", None)
- tool_calling_system_prompt = construct_tool_use_system_prompt(tools=inference_params["tools"])
- inference_params["system"] = (
- inference_params.get("system", "\n") + tool_calling_system_prompt
- ) # add the anthropic tool calling prompt to the system prompt
- inference_params.pop("tools")
- data = json.dumps({"messages": messages, **inference_params})
- else:
- ## LOAD CONFIG
- config = litellm.AmazonAnthropicConfig.get_config()
- for k, v in config.items():
- if (
- k not in inference_params
- ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
- inference_params[k] = v
- data = json.dumps({"prompt": prompt, **inference_params})
- elif provider == "ai21":
- ## LOAD CONFIG
- config = litellm.AmazonAI21Config.get_config()
- for k, v in config.items():
- if (
- k not in inference_params
- ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
- inference_params[k] = v
-
- data = json.dumps({"prompt": prompt, **inference_params})
- elif provider == "mistral":
- ## LOAD CONFIG
- config = litellm.AmazonMistralConfig.get_config()
- for k, v in config.items():
- if (
- k not in inference_params
- ): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in
- inference_params[k] = v
-
- data = json.dumps({"prompt": prompt, **inference_params})
- elif provider == "amazon": # amazon titan
- ## LOAD CONFIG
- config = litellm.AmazonTitanConfig.get_config()
- for k, v in config.items():
- if (
- k not in inference_params
- ): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in
- inference_params[k] = v
-
- data = json.dumps(
- {
- "inputText": prompt,
- "textGenerationConfig": inference_params,
- }
- )
- elif provider == "meta" or provider == "llama":
- ## LOAD CONFIG
- config = litellm.AmazonLlamaConfig.get_config()
- for k, v in config.items():
- if (
- k not in inference_params
- ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
- inference_params[k] = v
- data = json.dumps({"prompt": prompt, **inference_params})
- elif provider == "openai":
- ## OpenAI imported models use OpenAI Chat Completions format (messages-based)
- # Use AmazonBedrockOpenAIConfig for proper OpenAI transformation
- openai_config = AmazonBedrockOpenAIConfig()
- supported_params = openai_config.get_supported_openai_params(model=model)
-
- # Filter to only supported OpenAI params
- filtered_params = {k: v for k, v in inference_params.items() if k in supported_params}
-
- # OpenAI uses messages format, not prompt
- data = json.dumps({"messages": messages, **filtered_params})
- else:
- ## LOGGING
- logging_obj.pre_call(
- input=messages,
- api_key="",
- additional_args={
- "complete_input_dict": inference_params,
- },
- )
- raise BedrockError(
- status_code=404,
- message="Bedrock Invoke HTTPX: Unknown provider={}, model={}. Try calling via converse route - `bedrock/converse/`.".format(
- provider, model
- ),
- )
-
- ## COMPLETION CALL
-
- headers = {"Content-Type": "application/json"}
- if extra_headers is not None:
- headers = {"Content-Type": "application/json", **extra_headers}
- prepped = self.get_request_headers(
- credentials=credentials,
- aws_region_name=aws_region_name,
- extra_headers=extra_headers,
- endpoint_url=endpoint_url,
- data=data,
- headers=headers,
- )
-
- ## LOGGING
- logging_obj.pre_call(
- input=messages,
- api_key="",
- additional_args={
- "complete_input_dict": data,
- "api_base": proxy_endpoint_url,
- "headers": prepped.headers,
- },
- )
-
- ### ROUTING (ASYNC, STREAMING, SYNC)
- if acompletion:
- if isinstance(client, HTTPHandler):
- client = None
- if stream is True and provider != "ai21":
- return self.async_streaming(
- model=model,
- messages=messages,
- data=data,
- api_base=proxy_endpoint_url,
- model_response=model_response,
- print_verbose=print_verbose,
- encoding=encoding,
- logging_obj=logging_obj,
- optional_params=optional_params,
- stream=True,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- headers=prepped.headers,
- timeout=timeout,
- client=client,
- stream_chunk_size=stream_chunk_size,
- ) # type: ignore
- ### ASYNC COMPLETION
- return self.async_completion(
- model=model,
- messages=messages,
- data=data,
- api_base=proxy_endpoint_url,
- model_response=model_response,
- print_verbose=print_verbose,
- encoding=encoding,
- logging_obj=logging_obj,
- optional_params=optional_params,
- stream=stream, # type: ignore
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- headers=prepped.headers,
- timeout=timeout,
- client=client,
- ) # type: ignore
-
- if client is None or isinstance(client, AsyncHTTPHandler):
- _params = {}
- if timeout is not None:
- if isinstance(timeout, float) or isinstance(timeout, int):
- timeout = httpx.Timeout(timeout)
- _params["timeout"] = timeout
- self.client = _get_httpx_client(_params) # type: ignore
- else:
- self.client = client
- if (stream is not None and stream is True) and provider != "ai21":
- response = self.client.post(
- url=proxy_endpoint_url,
- headers=prepped.headers, # type: ignore
- data=data,
- stream=stream,
- logging_obj=logging_obj,
- )
-
- if response.status_code != 200:
- raise BedrockError(status_code=response.status_code, message=str(response.read()))
-
- decoder = AWSEventStreamDecoder(model=model)
-
- completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
- streaming_response = CustomStreamWrapper(
- completion_stream=completion_stream,
- model=model,
- custom_llm_provider="bedrock",
- logging_obj=logging_obj,
- )
-
- ## LOGGING
- logging_obj.post_call(
- input=messages,
- api_key="",
- original_response=streaming_response,
- additional_args={"complete_input_dict": data},
- )
- return streaming_response
-
- try:
- response = self.client.post(
- url=proxy_endpoint_url,
- headers=dict(prepped.headers),
- data=data,
- logging_obj=logging_obj,
- )
- response.raise_for_status()
- except httpx.HTTPStatusError as err:
- error_code = err.response.status_code
- raise BedrockError(status_code=error_code, message=err.response.text)
- except httpx.TimeoutException:
- raise BedrockError(status_code=408, message="Timeout error occurred.")
-
- return self.process_response(
- model=model,
- response=response,
- model_response=model_response,
- stream=stream,
- logging_obj=logging_obj,
- optional_params=optional_params,
- api_key="",
- data=data,
- messages=messages,
- print_verbose=print_verbose,
- encoding=encoding,
- )
-
- async def _async_anthropic_messages_completion(
- self,
- model: str,
- messages: list,
- endpoint_url: str,
- proxy_endpoint_url: str,
- credentials,
- aws_region_name: str,
- model_response: ModelResponse,
- print_verbose: Callable,
- encoding,
- logging_obj: Logging,
- optional_params: dict,
- stream,
- litellm_params=None,
- logger_fn=None,
- extra_headers: Optional[dict] = None,
- timeout: Optional[Union[float, httpx.Timeout]] = None,
- client: Optional[AsyncHTTPHandler] = None,
- stream_chunk_size: Optional[int] = None,
- ) -> Union[ModelResponse, CustomStreamWrapper]:
- transformed_request = await litellm.AmazonAnthropicClaudeConfig().async_transform_request(
- model=model,
- messages=messages,
- optional_params=optional_params,
- litellm_params=litellm_params or {},
- headers=extra_headers or {},
- )
- data = json.dumps(transformed_request)
-
- headers = {"Content-Type": "application/json"}
- if extra_headers is not None:
- headers = {"Content-Type": "application/json", **extra_headers}
- prepped = self.get_request_headers(
- credentials=credentials,
- aws_region_name=aws_region_name,
- extra_headers=extra_headers,
- endpoint_url=endpoint_url,
- data=data,
- headers=headers,
- )
-
- logging_obj.pre_call(
- input=messages,
- api_key="",
- additional_args={
- "complete_input_dict": data,
- "api_base": proxy_endpoint_url,
- "headers": prepped.headers,
- },
- )
-
- if stream is True:
- return await self.async_streaming(
- model=model,
- messages=messages,
- data=data,
- api_base=proxy_endpoint_url,
- model_response=model_response,
- print_verbose=print_verbose,
- encoding=encoding,
- logging_obj=logging_obj,
- optional_params=optional_params,
- stream=True,
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- headers=prepped.headers,
- timeout=timeout,
- client=client,
- stream_chunk_size=stream_chunk_size,
- )
- return await self.async_completion(
- model=model,
- messages=messages,
- data=data,
- api_base=proxy_endpoint_url,
- model_response=model_response,
- print_verbose=print_verbose,
- encoding=encoding,
- logging_obj=logging_obj,
- optional_params=optional_params,
- stream=stream, # type: ignore
- litellm_params=litellm_params,
- logger_fn=logger_fn,
- headers=prepped.headers,
- timeout=timeout,
- client=client,
- )
-
- async def async_completion(
- self,
- model: str,
- messages: list,
- api_base: str,
- model_response: ModelResponse,
- print_verbose: Callable,
- data: str,
- timeout: Optional[Union[float, httpx.Timeout]],
- encoding,
- logging_obj: Logging,
- stream,
- optional_params: dict,
- litellm_params=None,
- logger_fn=None,
- headers={},
- client: Optional[AsyncHTTPHandler] = None,
- ) -> Union[ModelResponse, CustomStreamWrapper]:
- if client is None:
- _params = {}
- if timeout is not None:
- if isinstance(timeout, float) or isinstance(timeout, int):
- timeout = httpx.Timeout(timeout)
- _params["timeout"] = timeout
- client = get_async_httpx_client(params=_params, llm_provider=litellm.LlmProviders.BEDROCK) # type: ignore
- else:
- client = client # type: ignore
-
- try:
- response = await client.post(
- api_base,
- headers=headers,
- data=data,
- timeout=timeout,
- logging_obj=logging_obj,
- )
- response.raise_for_status()
- except httpx.HTTPStatusError as err:
- error_code = err.response.status_code
- raise BedrockError(status_code=error_code, message=err.response.text)
- except httpx.TimeoutException:
- raise BedrockError(status_code=408, message="Timeout error occurred.")
-
- return self.process_response(
- model=model,
- response=response,
- model_response=model_response,
- stream=stream if isinstance(stream, bool) else False,
- logging_obj=logging_obj,
- api_key="",
- data=data,
- messages=messages,
- print_verbose=print_verbose,
- optional_params=optional_params,
- encoding=encoding,
- )
-
- @track_llm_api_timing() # for streaming, we need to instrument the function calling the wrapper
- async def async_streaming(
- self,
- model: str,
- messages: list,
- api_base: str,
- model_response: ModelResponse,
- print_verbose: Callable,
- data: str,
- timeout: Optional[Union[float, httpx.Timeout]],
- encoding,
- logging_obj: Logging,
- stream,
- optional_params: dict,
- litellm_params=None,
- logger_fn=None,
- headers={},
- client: Optional[AsyncHTTPHandler] = None,
- stream_chunk_size: Optional[int] = None,
- ) -> CustomStreamWrapper:
- # The call is not made here; instead, we prepare the necessary objects for the stream.
-
- streaming_response = CustomStreamWrapper(
- completion_stream=None,
- make_call=partial(
- make_call,
- client=client,
- api_base=api_base,
- headers=headers,
- data=data, # type: ignore
- model=model,
- messages=messages,
- logging_obj=logging_obj,
- fake_stream=True if "ai21" in api_base else False,
- stream_chunk_size=stream_chunk_size,
- ),
- model=model,
- custom_llm_provider="bedrock",
- logging_obj=logging_obj,
- )
- return streaming_response
-
- @staticmethod
- def _get_provider_from_model_path(
- model_path: str,
- ) -> Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL]:
- """
- Helper function to get the provider from a model path with format: provider/model-name
-
- Args:
- model_path (str): The model path (e.g., 'llama/arn:aws:bedrock:us-east-1:086734376398:imported-model/r4c4kewx2s0n' or 'anthropic/model-name')
-
- Returns:
- Optional[str]: The provider name, or None if no valid provider found
- """
- parts = model_path.split("/")
- if len(parts) >= 1:
- provider = parts[0]
- if provider in get_args(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL):
- return cast(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL, provider)
- return None
-
-
class AWSEventStreamDecoder:
def __init__(self, model: str, json_mode: Optional[bool] = False) -> None:
from botocore.parsers import EventStreamJSONParser
diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py
index 5114677ffc0..93998f0610e 100644
--- a/litellm/llms/bedrock/common_utils.py
+++ b/litellm/llms/bedrock/common_utils.py
@@ -1109,8 +1109,10 @@ def get_bedrock_chat_config(model: str):
Returns:
The appropriate Bedrock config class instance
"""
+ from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
+
bedrock_route = BedrockModelInfo.get_bedrock_route(model)
- bedrock_invoke_provider = litellm.BedrockLLM.get_bedrock_invoke_provider(model=model)
+ bedrock_invoke_provider = BaseAWSLLM.get_bedrock_invoke_provider(model=model)
base_model = BedrockModelInfo.get_base_model(model)
# Handle explicit routes first
diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py
index f0ce3323ef6..36c78974aca 100644
--- a/litellm/llms/vertex_ai/context_caching/transformation.py
+++ b/litellm/llms/vertex_ai/context_caching/transformation.py
@@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works
"""
import re
-from typing import List, Optional, Tuple, Literal
+from typing import List, Optional, Sequence, Tuple, Literal
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.vertex_ai import CachedContentRequestBody
@@ -152,6 +152,20 @@ def separate_cached_messages(
return cached_messages, non_cached_messages
+def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageValues]) -> bool:
+ """
+ The cachedContents API rejects contents ending on a model turn, which is how it
+ classifies both assistant messages and tool results, with HTTP 400
+ "Requests ending with a model turn are not supported". System messages are
+ extracted into system_instruction before contents are built, so the terminal
+ turn is the last non-system message.
+ """
+ non_system_messages = tuple(message for message in cached_messages if message.get("role") != "system")
+ if not non_system_messages:
+ return bool(cached_messages)
+ return non_system_messages[-1].get("role") not in ("assistant", "tool", "function")
+
+
def transform_openai_messages_to_gemini_context_caching(
model: str,
messages: List[AllMessageValues],
diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py
index 0bf3715f798..f8774e33ca4 100644
--- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py
+++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py
@@ -22,6 +22,7 @@ from litellm.types.llms.vertex_ai import (
from ..common_utils import VertexAIError, get_vertex_base_url
from ..vertex_llm_base import VertexBase
from .transformation import (
+ cached_messages_end_on_supported_turn,
separate_cached_messages,
transform_openai_messages_to_gemini_context_caching,
)
@@ -308,6 +309,14 @@ class ContextCachingEndpoints(VertexBase):
if len(cached_messages) == 0:
return messages, optional_params, None
+ if not cached_messages_end_on_supported_turn(cached_messages):
+ verbose_logger.debug(
+ "Vertex AI context caching: cached message block ends on a model turn once "
+ "system messages are extracted, which the cachedContents API rejects. "
+ "Skipping context caching."
+ )
+ return messages, optional_params, None
+
# Gemini requires a minimum of 1024 tokens for context caching.
# Skip caching if the cached content is too small to avoid API errors.
if not is_prompt_caching_valid_prompt(
@@ -459,6 +468,14 @@ class ContextCachingEndpoints(VertexBase):
if len(cached_messages) == 0:
return messages, optional_params, None
+ if not cached_messages_end_on_supported_turn(cached_messages):
+ verbose_logger.debug(
+ "Vertex AI context caching: cached message block ends on a model turn once "
+ "system messages are extracted, which the cachedContents API rejects. "
+ "Skipping context caching."
+ )
+ return messages, optional_params, None
+
# Gemini requires a minimum of 1024 tokens for context caching.
# Skip caching if the cached content is too small to avoid API errors.
if not is_prompt_caching_valid_prompt(
diff --git a/litellm/llms/vertex_ai/files/handler.py b/litellm/llms/vertex_ai/files/handler.py
index 3bc09139f8f..4d2a1e18eb5 100644
--- a/litellm/llms/vertex_ai/files/handler.py
+++ b/litellm/llms/vertex_ai/files/handler.py
@@ -1,7 +1,9 @@
import asyncio
+import json
+import os
import time
from urllib.parse import unquote
-from typing import Any, Coroutine, Optional, Tuple, Union
+from typing import Any, Coroutine, Mapping, Optional, Tuple, Union
import httpx
@@ -10,6 +12,7 @@ from litellm.integrations.gcs_bucket.gcs_bucket_base import (
GCSBucketBase,
GCSLoggingConfig,
)
+from litellm.types.utils import StandardCallbackDynamicParams
from litellm.litellm_core_utils.cloud_storage_security import (
VERTEX_AI_MANAGED_GCS_PREFIX,
should_allow_legacy_cloud_file_ids,
@@ -39,6 +42,35 @@ class VertexAIFilesHandler(GCSBucketBase):
llm_provider=LlmProviders.VERTEX_AI,
)
+ def _resolve_read_gcs_config(
+ self,
+ litellm_params: Mapping[str, object] | None,
+ vertex_credentials: VERTEX_CREDENTIALS_TYPES | None,
+ ) -> tuple[str | None, str | None]:
+ """
+ Resolve the GCS bucket and service-account credentials for the read/content path.
+
+ Sources them from the deployment's ``litellm_params`` (``gcs_bucket_name`` /
+ ``bucket_name`` and ``vertex_credentials``), mirroring the write path in
+ ``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the global
+ ``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT`` env vars. This lets Vertex batch
+ run entirely at the model-group level, so output written to a per-model bucket is
+ readable without setting the global env vars.
+ """
+ params: Mapping[str, object] = litellm_params or {}
+ bucket_candidate = params.get("gcs_bucket_name") or params.get("bucket_name")
+ configured_bucket_name = bucket_candidate if isinstance(bucket_candidate, str) else os.getenv("GCS_BUCKET_NAME")
+
+ credentials = params.get("vertex_credentials") or vertex_credentials
+ if isinstance(credentials, dict):
+ path_service_account: str | None = json.dumps(credentials)
+ elif isinstance(credentials, str):
+ path_service_account = credentials
+ else:
+ path_service_account = os.getenv("GCS_PATH_SERVICE_ACCOUNT")
+
+ return configured_bucket_name, path_service_account
+
def _extract_bucket_and_object_from_file_id(
self,
file_id: str,
@@ -91,7 +123,17 @@ class VertexAIFilesHandler(GCSBucketBase):
if not file_id:
raise ValueError("file_id is required in file_content_request")
- gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(kwargs={})
+ configured_bucket_name, path_service_account = self._resolve_read_gcs_config(
+ litellm_params=litellm_params,
+ vertex_credentials=vertex_credentials,
+ )
+ dynamic_params = StandardCallbackDynamicParams(
+ gcs_bucket_name=configured_bucket_name,
+ gcs_path_service_account=path_service_account,
+ )
+ gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(
+ kwargs={"standard_callback_dynamic_params": dynamic_params}
+ )
bucket_name, object_path = self._extract_bucket_and_object_from_file_id(
file_id=file_id,
configured_bucket_name=gcs_logging_config["bucket_name"],
diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py
index 0db1118a7b4..cbca57c5e62 100644
--- a/litellm/llms/vertex_ai/gemini/transformation.py
+++ b/litellm/llms/vertex_ai/gemini/transformation.py
@@ -661,6 +661,10 @@ def _gemini_convert_messages_with_history(
vertex_project = litellm_params.get("vertex_project") or litellm_params.get("vertex_ai_project")
vertex_credentials = litellm_params.get("vertex_credentials") or litellm_params.get("vertex_ai_credentials")
+ from .vertex_and_google_ai_studio_gemini import VertexGeminiConfig
+
+ forward_function_call_id = VertexGeminiConfig._forward_gemini_function_call_id(model or "")
+
try:
while msg_i < len(messages):
user_content: List[PartType] = []
@@ -910,7 +914,7 @@ def _gemini_convert_messages_with_history(
gemini_tool_call_parts = convert_to_gemini_tool_call_invoke(
assistant_msg,
model=model,
- custom_llm_provider=custom_llm_provider,
+ forward_function_call_id=forward_function_call_id,
)
## check if gemini_tool_call already exists in assistant_content
for gemini_tool_call_part in gemini_tool_call_parts:
@@ -973,8 +977,7 @@ def _gemini_convert_messages_with_history(
_part = convert_to_gemini_tool_call_result(
messages[msg_i], # type: ignore
last_message_with_tool_calls, # type: ignore
- model=model,
- custom_llm_provider=custom_llm_provider,
+ forward_function_call_id=forward_function_call_id,
)
msg_i += 1
# Handle both single part and list of parts (for Computer Use with images)
diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py
index 8dd0dc19b81..126f82436e8 100644
--- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py
+++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py
@@ -289,15 +289,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return False
@staticmethod
- def _forward_gemini_function_call_id(model: str, custom_llm_provider: Optional[str] = None) -> bool:
+ def _forward_gemini_function_call_id(model: str) -> bool:
"""
Whether to include `id` on function_call / function_response parts.
- Gemini 3+ on Google AI Studio accepts (and returns) `id` for strict
- tool-call matching. Vertex AI rejects the field with HTTP 400.
+ Gemini 3+ accepts (and returns) `id` for strict tool-call matching, on Vertex AI and
+ Google AI Studio alike. Older Gemini models reject the field with HTTP 400.
"""
- if custom_llm_provider != "gemini":
- return False
return VertexGeminiConfig._is_gemini_3_or_newer(model)
def _supports_penalty_parameters(self, model: str) -> bool:
diff --git a/litellm/main.py b/litellm/main.py
index dc3ec469a1b..acdec7385da 100644
--- a/litellm/main.py
+++ b/litellm/main.py
@@ -207,7 +207,7 @@ from .llms.azure.chat.o_series_handler import AzureOpenAIO1ChatCompletion
from .llms.azure.completion.handler import AzureTextCompletion
from .llms.azure_ai.anthropic.handler import AzureAnthropicChatCompletion
from .llms.azure_ai.embed import AzureAIEmbedding
-from .llms.bedrock.chat import BedrockConverseLLM, BedrockLLM
+from .llms.bedrock.chat import BedrockConverseLLM
from .llms.bedrock.embed.embedding import BedrockEmbedding
from .llms.bedrock.image_edit.handler import BedrockImageEdit
from .llms.bedrock.image_generation.image_handler import BedrockImageGeneration
diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
index 82b820d8cd9..3e0775ac09e 100644
--- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
+++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
@@ -200,6 +200,13 @@ _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: tuple[MCPAuth, ...] = (
)
+# OAuth discovery retry cooldown for servers whose endpoints stay unresolved. The base is one
+# reload cadence so a transient upstream failure recovers immediately; the cap bounds the request
+# amplification and log volume of a permanently broken configuration.
+_OAUTH_DISCOVERY_RETRY_BASE_SECONDS = 30.0
+_OAUTH_DISCOVERY_RETRY_MAX_SECONDS = 900.0
+
+
def _blank_to_none(value: str | None) -> str | None:
"""Collapse an absent, empty, or whitespace-only string to ``None``.
@@ -247,6 +254,7 @@ def _endpoints_yield_to_issuer(
authorization_url: str | None,
token_url: str | None,
registration_url: str | None,
+ server_ref: str,
) -> tuple[str | None, str | None, str | None]:
"""The single rule that makes an admin-configured ``issuer`` the sole authoritative endpoint
source (RFC 8414 §3.3): when it is set for a discovery auth type, the stored/manual
@@ -256,9 +264,29 @@ def _endpoints_yield_to_issuer(
i.e. all ``None`` when issuer-anchored, else the inputs unchanged. Called at every resolution site
so the invariant holds in one place instead of being re-derived per merge.
"""
- if issuer is not None and is_discovery_auth_type:
- return None, None, None
- return authorization_url, token_url, registration_url
+ if issuer is None or not is_discovery_auth_type:
+ return authorization_url, token_url, registration_url
+ discarded = sorted(
+ label
+ for label, value in (
+ ("authorization_url", authorization_url),
+ ("token_url", token_url),
+ ("registration_url", registration_url),
+ )
+ if value
+ )
+ if discarded:
+ verbose_logger.warning(
+ "MCP server %s has a pinned Issuer, so its stored %s %s not used: an anchored issuer is the "
+ "sole endpoint source (RFC 8414 section 3.3) and a failed issuer fetch fails closed rather "
+ "than falling back to them. To use manually configured endpoints instead, clear the Issuer "
+ "field and re-enter the endpoint urls (clearing the Issuer also clears endpoints that may "
+ "have been resolved under it), or clear the Issuer alone to re-discover from the server url.",
+ server_ref,
+ ", ".join(discarded),
+ "is" if len(discarded) == 1 else "are",
+ )
+ return None, None, None
def _normalized_authorize_endpoint(url: str) -> str:
@@ -280,6 +308,68 @@ def _issuer_matches(claimed_issuer: object, configured_issuer: str) -> bool:
return _normalized_authorize_endpoint(claimed_issuer) == _normalized_authorize_endpoint(configured_issuer)
+def _flow_endpoints_missing(
+ auth_type: MCPAuthType | None,
+ oauth2_flow: str | None,
+ authorization_url: str | None,
+ token_url: str | None,
+ token_exchange_endpoint: str | None = None,
+) -> bool:
+ """Whether a built server is missing an endpoint its flow needs to run at all.
+
+ Used by the reload fast-path exemption: discovery runs at build time only, and the fast path
+ reuses an unchanged row's registry entry verbatim, so a server whose discovery came back empty
+ (transient upstream failure, rate limiting) would stay broken until some unrelated config write
+ bumps ``updated_at``, serving its 400 the whole time. Rebuilding just these entries retries
+ discovery on the normal reload cadence. It costs no extra fetch for servers that resolved, and
+ none for those with no discovery source, since the build skips discovery for both.
+ """
+ if auth_type == MCPAuth.oauth2_token_exchange:
+ # A configured exchange endpoint replaces discovery entirely; only a server that must
+ # discover its token endpoint and still has none is unresolved.
+ return token_exchange_endpoint is None and token_url is None
+ if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
+ return False
+ if oauth2_flow == "client_credentials":
+ return token_url is None
+ return authorization_url is None or token_url is None
+
+
+def _oauth_endpoints_unresolved(server: MCPServer) -> bool:
+ """``_flow_endpoints_missing`` over a built registry entry, for the reload fast-path check.
+
+ The flow comes from ``effective_oauth2_flow``, the one column-first, shape-fallback judge every
+ flow decision uses, not from the raw column: a legacy row the startup backfill deliberately left
+ unstamped (the ambiguous M2M shape) serves M2M at request time, and reading the bare column here
+ would classify it as interactive-missing-endpoints and re-run discovery on every reload.
+ """
+ if (
+ server.auth_type == MCPAuth.oauth2_token_exchange
+ and server.token_exchange_profile == "entra_obo"
+ and not server.scopes
+ ):
+ # entra_obo fails closed at exchange time without a scope (token_exchanger.py), and scopes
+ # can come from resource discovery, so a server that resolved its endpoints but no scopes is
+ # still unresolved for its flow.
+ return True
+ if server.is_dcr_bridge and not server.client_id and server.registration_url is None:
+ # A DCR bridge with no admin-configured client can only register callers through the
+ # upstream's registration endpoint, so a build that resolved the authorize and token
+ # endpoints but not registration_endpoint (partial metadata) is still unresolved for its
+ # flow and must keep retrying; without this it silently degrades to the short-circuit arm
+ # until an unrelated config write. Scopes are deliberately NOT part of completeness: they
+ # are a request hint the authorization server bounds at consent (RFC 6749 section 3.3),
+ # and a server without them is fully functional.
+ return True
+ return _flow_endpoints_missing(
+ server.auth_type,
+ MCPServerManager.effective_oauth2_flow(server),
+ server.authorization_url,
+ server.token_url,
+ server.token_exchange_endpoint,
+ )
+
+
def _endpoints_corroborate_authorization_url(
source_authorization_url: str | None,
trusted_authorization_url: str | None,
@@ -311,11 +401,10 @@ def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serv
during re-discovery downgrades a working server (``authorization_url`` set) to a broken one
(``None``, /authorize 400s) with no configuration change. Mirrors the ``short_prefix``
carry-forward. Skipped when the server's ``url`` or ``auth_type`` changed, since the previous
- endpoints may then belong to a different upstream. ``registration_url`` IS carried even though
- ``_persist_discovered_oauth_endpoints`` refuses to write it to the row: carrying only restores
- the same in-memory value the previous build already ran with, while persisting it would flip
- ``_dcr_bridge_relays_client_registration`` (which keys off the stored column) for dcr_bridge
- servers that never had one configured.
+ endpoints may then belong to a different upstream. Discovery results live only on the in-memory
+ registry entry; the gateway never writes them to the row, whose OAuth columns carry admin intent
+ alone, so this carry is the sole last-known-good mechanism and restores exactly the values the
+ previous build already ran with.
Carry-forward is a non-manual endpoint source, so the same trust rule as discovery applies: the
previous ``token_url``/``registration_url``/``scopes`` are carried only when the previous
@@ -1182,6 +1271,40 @@ class MCPServerManager:
# empty result, or failure). Used to throttle re-probes for servers that do
# not return instructions, and to apply a short cooldown after failures.
self._upstream_initialize_instructions_probed_at: dict[str, float] = {}
+ # Per-server (consecutive failures, monotonic timestamp) for OAuth discovery retries, so a
+ # server whose endpoints never resolve backs off instead of re-running the full
+ # RFC 9728 -> 8414 chain, and re-logging its warning, on every reload forever.
+ self._oauth_discovery_retry_state: dict[
+ str, tuple[int, float]
+ ] = {} # mutable-ok: retry cooldown cache, keyed per server and pruned on success
+
+ def _oauth_discovery_retry_due(self, server_id: str) -> bool:
+ """Whether an unresolved server is due for another discovery attempt.
+
+ The reload fast-path exemption is what retries a failed discovery, so without a cooldown a
+ permanently unresolvable server re-runs the whole RFC 9728 -> RFC 8414 -> origin-fallback
+ chain and re-emits its unresolved-endpoints warning on every reload, per server, forever.
+ Delay doubles per consecutive failure from ``_OAUTH_DISCOVERY_RETRY_BASE_SECONDS`` up to
+ ``_OAUTH_DISCOVERY_RETRY_MAX_SECONDS``, so a transient outage still recovers on the next
+ reload while a broken configuration settles to one attempt per cap.
+ """
+ state = self._oauth_discovery_retry_state.get(server_id)
+ if state is None:
+ return True
+ failures, attempted_at = state
+ delay = min(
+ _OAUTH_DISCOVERY_RETRY_BASE_SECONDS * (2 ** max(failures - 1, 0)),
+ _OAUTH_DISCOVERY_RETRY_MAX_SECONDS,
+ )
+ return (time.monotonic() - attempted_at) >= delay
+
+ def _record_oauth_discovery_outcome(self, server: MCPServer) -> None:
+ """Advance or clear a server's retry cooldown after a rebuild resolved it or did not."""
+ if not _oauth_endpoints_unresolved(server):
+ self._oauth_discovery_retry_state.pop(server.server_id, None)
+ return
+ failures, _ = self._oauth_discovery_retry_state.get(server.server_id, (0, 0.0))
+ self._oauth_discovery_retry_state[server.server_id] = (failures + 1, time.monotonic())
def _remember_upstream_initialize_instructions(self, server: MCPServer, client: MCPClient) -> None:
raw = getattr(client, "_last_initialize_instructions", None)
@@ -1357,6 +1480,7 @@ class MCPServerManager:
manual_authorization_url,
manual_token_url,
manual_registration_url,
+ server_name or server_id,
)
should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and (
is_discovery_auth_type or obo_needs_discovery
@@ -1834,7 +1958,6 @@ class MCPServerManager:
*,
credentials_are_encrypted: bool = True,
env_vars_are_encrypted: Optional[bool] = None,
- persist_discovered_endpoints: bool = True,
) -> MCPServer:
_mcp_info: MCPInfo = mcp_server.mcp_info or {}
env_dict = _deserialize_json_dict(getattr(mcp_server, "env", None))
@@ -1925,7 +2048,12 @@ class MCPServerManager:
or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url),
)
manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer(
- manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url
+ manual_issuer,
+ is_discovery_auth_type,
+ manual_authorization_url,
+ manual_token_url,
+ manual_registration_url,
+ mcp_server.alias or mcp_server.server_name or mcp_server.server_id,
)
gated_oauth_metadata = await self._resolve_table_oauth_metadata(
mcp_server=mcp_server,
@@ -2033,143 +2161,8 @@ class MCPServerManager:
max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None),
)
_warn_internal_delegate_pkce_if_applicable(new_server, source="database")
- if persist_discovered_endpoints:
- await self._persist_discovered_obo_token_url(
- server_id=mcp_server.server_id,
- auth_type=auth_type,
- existing_token_url=manual_token_url,
- discovered_token_url=new_server.token_url,
- )
- await self._persist_discovered_oauth_endpoints(
- server_id=mcp_server.server_id,
- auth_type=auth_type,
- existing_issuer=manual_issuer,
- existing_authorization_url=manual_authorization_url,
- existing_token_url=manual_token_url,
- existing_scopes=scopes,
- metadata=gated_oauth_metadata,
- is_issuer_anchored=use_issuer_anchor,
- )
return new_server
- async def _persist_discovered_obo_token_url(
- self,
- *,
- server_id: str,
- auth_type: Optional[MCPAuthType],
- existing_token_url: Optional[str],
- discovered_token_url: Optional[str],
- ) -> None:
- """Write a freshly discovered OBO token endpoint back onto the DB row.
-
- ``build_mcp_server_from_table`` resolves ``token_url`` via RFC 9728 -> RFC 8414 for an
- ``oauth2_token_exchange`` server that has none configured, but that resolved value otherwise
- lives only on the returned in-memory object; the row keeps ``token_url=None`` so every rebuild
- re-runs discovery, and a transient upstream outage during a rebuild leaves the server with no
- endpoint until discovery next succeeds. Persisting it makes ``_obo_needs_endpoint_discovery``
- return False on the next build. Fires at most once per server (skipped once the row has a
- value), and is best-effort: a write failure just means discovery runs again next time.
- """
- if auth_type != MCPAuth.oauth2_token_exchange:
- return
- if existing_token_url or not discovered_token_url:
- return
- from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415
-
- if prisma_client is None:
- return
- try:
- await MCPServerRepository(prisma_client).table.update(
- where={"server_id": server_id},
- data={"token_url": discovered_token_url},
- )
- verbose_logger.debug("Persisted discovered OBO token_url for MCP server %s", server_id)
- except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build
- verbose_logger.warning("Failed to persist discovered OBO token_url for MCP server %s: %s", server_id, exc)
-
- async def _persist_discovered_oauth_endpoints(
- self,
- *,
- server_id: str,
- auth_type: MCPAuthType | None,
- existing_issuer: str | None,
- existing_authorization_url: str | None,
- existing_token_url: str | None,
- existing_scopes: list[str] | None,
- metadata: MCPOAuthMetadata | None,
- is_issuer_anchored: bool = False,
- ) -> None:
- """Write freshly discovered OAuth endpoints back onto the DB row.
-
- Same rationale as ``_persist_discovered_obo_token_url`` but for the interactive oauth2
- family: discovered ``authorization_url``/``token_url``/``scopes`` otherwise live only on
- the in-memory registry entry, which is rebuilt on every client connect (the DCR reuse path
- calls ``update_server``) and on every post-write DB reload, so one failed re-discovery
- serves the 400 "authorization url is not configured" from /authorize until a later rebuild succeeds.
- Only fills row fields that are currently empty, never persists origin-fallback guesses
- (RFC 9728/8414-advertised metadata only), and deliberately skips ``registration_url``
- because ``_dcr_bridge_relays_client_registration`` keys off that column. Best-effort: a
- failed write re-discovers on the next build. Scopes go through ``update_mcp_server`` so
- they merge into the credentials blob without touching the stored client credentials.
-
- For an issuer-anchored server (``is_issuer_anchored``) the endpoints are re-derived from the
- §3.3-validated issuer document on every build, so they are NOT persisted into the endpoint
- columns: persisting them would make the next build see populated endpoints and treat them as
- authoritative stored values, defeating the "endpoints come solely from the issuer" invariant.
- Only the resource-driven scopes are persisted for such servers.
- """
- if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
- return
- if metadata is None or metadata.from_origin_fallback:
- return
- issuer_update = (
- {"issuer": metadata.discovered_issuer} if metadata.discovered_issuer and not existing_issuer else {}
- )
- authorization_url_update = (
- {"authorization_url": metadata.authorization_url}
- if metadata.authorization_url and not existing_authorization_url and not is_issuer_anchored
- else {}
- )
- token_url_update = (
- {"token_url": metadata.token_url}
- if metadata.token_url and not existing_token_url and not is_issuer_anchored
- else {}
- )
- scopes_update = {"credentials": {"scopes": metadata.scopes}} if metadata.scopes and not existing_scopes else {}
- updates: dict[str, object] = {
- **issuer_update,
- **authorization_url_update,
- **token_url_update,
- **scopes_update,
- }
- if not updates:
- return
- from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # db.py imports this module at load
- update_mcp_server,
- )
- from litellm.proxy._types import UpdateMCPServerRequest # noqa: PLC0415 # heavy module; import at call time
- from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime value, set after startup
-
- if prisma_client is None:
- return
- try:
- await update_mcp_server(
- prisma_client=prisma_client,
- data=UpdateMCPServerRequest.model_validate({"server_id": server_id, **updates}),
- touched_by="mcp_oauth_discovery",
- )
- verbose_logger.info(
- "Persisted discovered OAuth endpoints for MCP server %s: %s",
- server_id,
- sorted(updates),
- )
- except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build
- verbose_logger.warning(
- "Failed to persist discovered OAuth endpoints for MCP server %s: %s",
- server_id,
- exc,
- )
-
async def _maybe_register_openapi_tools(self, server: MCPServer, *, initialize_mapping: bool = True):
"""Register OpenAPI tools if the server has a spec_path configured."""
if server.spec_path:
@@ -5347,6 +5340,10 @@ class MCPServerManager:
and existing_server.updated_at is not None
and server.updated_at is not None
and existing_server.updated_at == server.updated_at
+ and not (
+ _oauth_endpoints_unresolved(existing_server)
+ and self._oauth_discovery_retry_due(server.server_id)
+ )
):
# Re-use existing server instance to avoid re-running build_mcp_server_from_table()
# which can perform network discovery for OAuth2 servers.
@@ -5364,6 +5361,7 @@ class MCPServerManager:
# already-decrypted records add_server/update_server are handed.
# Decrypt them while building the registry entry.
new_server = await self.build_mcp_server_from_table(server, env_vars_are_encrypted=True)
+ self._record_oauth_discovery_outcome(new_server)
# Carry the cached short_prefix from the previous registry entry
# (if any) so the prefix is stable across reloads.
if existing_server is not None and existing_server.short_prefix:
diff --git a/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py b/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py
new file mode 100644
index 00000000000..874fcc64772
--- /dev/null
+++ b/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py
@@ -0,0 +1,148 @@
+"""One-time heal for MCP server rows whose ``issuer`` a released version wrote by itself.
+
+Until the write was removed, OAuth discovery stamped the issuer it discovered onto the ``issuer``
+column trust-on-first-use. That column means "the admin pinned this trust anchor", so the next
+registry build read the gateway's own output back as admin intent: the server turned issuer-anchored
+(RFC 8414 section 3.3), its stored authorization/token/registration URLs stopped applying, and a
+failed issuer-document fetch left it with no authorize endpoint (GH #34985).
+
+Deleting the write fixes every row created afterwards but cannot fix a row already stamped, which
+still reads as pinned. This heals those rows by clearing the stamp so their configured endpoints
+apply again.
+
+The signal is a heuristic, and deliberately a narrow one. ``updated_by`` records only the most recent
+writer, and no audit trail says which field that writer touched, so "discovery wrote this issuer" is
+not directly knowable. Two independent clauses bound it, and each rules out a different way of
+destroying a pin an admin meant.
+
+Configured endpoints must be present. A deliberately pinned row very often has none, both because the
+Issuer field is documented as overriding them and because ``update_mcp_server`` clears them when an
+issuer changes, so "issuer set, endpoints empty" is the canonical shape of a real pin and must never
+be cleared on this evidence. Skipping those rows costs little: with nothing configured to restore, the
+anchored and resource-rooted paths resolve from the same upstream document, and the row still gets the
+unresolved-endpoint retry and the anchored-discard warning.
+
+The configured endpoints must also share the issuer's origin. A stamped issuer is by construction the
+one self-attested by the authorization-server document discovery reached from this very server, so
+endpoints typed alongside it address that same authority. An admin who pinned an issuer and typed
+endpoints for a different authority is expressing an intent that clearing the issuer would discard, so
+that row is warned about and never healed.
+
+What survives both clauses is a row whose configured endpoints and stamped issuer share an origin,
+which is exactly the GH #34985 shape. An admin who pinned that same origin by hand lands here too, and
+for them the clear is close to a no-op: their typed endpoints keep serving and still anchor the
+RFC 9700 corroboration gate, with only the stricter section 3.3 anchoring lost. Every heal logs the
+cleared value so it can be restored, and the clear is recorded under this module's actor so the heal
+runs at most once per row.
+"""
+
+from typing import Protocol
+from urllib.parse import urlparse
+
+from litellm._logging import verbose_proxy_logger
+from litellm.proxy._experimental.mcp_server.oauth_utils import canonicalize_url_identity
+from litellm.proxy.utils import PrismaClient
+
+# The actor the removed discovery write-back stamped rows with.
+_DISCOVERY_ACTOR = "mcp_oauth_discovery"
+
+# The actor recorded on a healed row, which also makes the heal idempotent: once a row is cleared it
+# no longer matches ``updated_by == _DISCOVERY_ACTOR`` and is never reconsidered.
+_BACKFILL_ACTOR = "mcp_oauth_issuer_stamp_backfill"
+
+_AUTH_TYPES_WITH_ISSUER_ANCHORING = ("oauth2", "true_passthrough", "oauth_delegate")
+
+
+def _origin(url: str) -> str | None:
+ """The scheme-and-authority identity of ``url``, or ``None`` when it has none.
+
+ Built on the shared URL canonicalizer so the lowercase-host and default-port rules match the
+ RFC 8414 issuer comparison the resolution path uses, instead of being re-derived here.
+ """
+ parsed = urlparse(canonicalize_url_identity(url))
+ if not parsed.scheme or not parsed.netloc:
+ return None
+ return f"{parsed.scheme}://{parsed.netloc}"
+
+
+class _MCPServerRow(Protocol):
+ """The MCP server row fields this heal reads, so the untyped DB record is narrowed once here."""
+
+ server_id: str
+ alias: str | None
+ server_name: str | None
+ auth_type: str | None
+ issuer: str | None
+ authorization_url: str | None
+ token_url: str | None
+ registration_url: str | None
+ updated_by: str | None
+
+
+def _is_stamped_issuer_row(row: _MCPServerRow) -> bool:
+ """Whether this row carries the full signature of a gateway-written issuer stamp.
+
+ The whole rule lives here, including the writer check the query also filters on, so the decision
+ to clear an admin-visible field is auditable in one place rather than split between a predicate
+ and a query.
+ """
+ if getattr(row, "updated_by", None) != _DISCOVERY_ACTOR:
+ return False
+ if not (getattr(row, "issuer", None) or "").strip():
+ return False
+ if getattr(row, "auth_type", None) not in _AUTH_TYPES_WITH_ISSUER_ANCHORING:
+ return False
+ configured = tuple(
+ value.strip()
+ for value in (row.authorization_url, row.token_url, row.registration_url)
+ if value and value.strip()
+ )
+ if not configured:
+ return False
+ issuer_origin = _origin(row.issuer or "")
+ return issuer_origin is not None and all(_origin(endpoint) == issuer_origin for endpoint in configured)
+
+
+async def backfill_discovery_stamped_issuers(prisma_client: PrismaClient) -> int:
+ """Clear gateway-written issuer stamps, returning the number of rows healed."""
+ candidate_rows: list[_MCPServerRow] = await prisma_client.db.litellm_mcpservertable.find_many(
+ where={
+ "updated_by": _DISCOVERY_ACTOR,
+ "auth_type": {"in": list(_AUTH_TYPES_WITH_ISSUER_ANCHORING)},
+ },
+ )
+ stamped = tuple(row for row in candidate_rows if _is_stamped_issuer_row(row))
+ if not stamped:
+ return 0
+
+ healed = 0
+ for row in stamped:
+ try:
+ await prisma_client.db.litellm_mcpservertable.update(
+ where={"server_id": row.server_id},
+ data={"issuer": None, "updated_by": _BACKFILL_ACTOR},
+ )
+ except Exception as exc: # noqa: BLE001 - per-row best effort; the next boot retries
+ verbose_proxy_logger.warning(
+ "MCP issuer stamp backfill: could not heal server_id=%s: %s", row.server_id, exc
+ )
+ continue
+ healed += 1
+ verbose_proxy_logger.warning(
+ "MCP issuer stamp backfill: cleared issuer %r on server_id=%s (alias=%s). OAuth discovery "
+ "had written that value onto the Issuer column, which made the server issuer-anchored and "
+ "fail-closed, and its configured Authorization/Token/Registration URLs were being ignored "
+ "as a result; those now apply again. If you pinned this issuer deliberately, set it again "
+ "via the dashboard or PUT /v1/mcp/server to restore RFC 8414 section 3.3 anchoring.",
+ row.issuer,
+ row.server_id,
+ row.alias or row.server_name,
+ )
+
+ if healed:
+ verbose_proxy_logger.warning(
+ "MCP issuer stamp backfill: healed %d server(s) whose Issuer had been written by OAuth "
+ "discovery rather than by an admin",
+ healed,
+ )
+ return healed
diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py
index ecb37e67c14..644253ceac7 100644
--- a/litellm/proxy/auth/auth_utils.py
+++ b/litellm/proxy/auth/auth_utils.py
@@ -1432,7 +1432,7 @@ def _extract_models_from_managed_resource_id(
)
_append_model_candidates(
candidates=candidates,
- value=get_model_id_from_unified_batch_id(unified_file_id),
+ value=_resolve_model_id_with_router(get_model_id_from_unified_batch_id(unified_file_id), llm_router),
)
except Exception as e:
verbose_proxy_logger.debug("Unable to extract model from managed file/batch ID: %s", str(e))
@@ -1442,7 +1442,10 @@ def _extract_models_from_managed_resource_id(
parsed_id = parse_unified_id(resource_id)
if parsed_id:
- _append_model_candidates(candidates=candidates, value=parsed_id.get("model_id"))
+ _append_model_candidates(
+ candidates=candidates,
+ value=_resolve_model_id_with_router(parsed_id.get("model_id"), llm_router),
+ )
_append_model_candidates(candidates=candidates, value=parsed_id.get("target_model_names"))
except Exception as e:
verbose_proxy_logger.debug("Unable to extract model from unified managed resource ID: %s", str(e))
diff --git a/litellm/proxy/client/cli/README.md b/litellm/proxy/client/cli/README.md
index 2ad8a08b8c3..de9d38963c1 100644
--- a/litellm/proxy/client/cli/README.md
+++ b/litellm/proxy/client/cli/README.md
@@ -10,11 +10,32 @@ uv tool install 'litellm[proxy]'
## Configuration
-The CLI can be configured using environment variables or command-line options:
+The CLI can be configured using environment variables, command-line options, or a persistent config file:
- `LITELLM_PROXY_URL`: Base URL of the LiteLLM proxy server (default: http://localhost:4000)
- `LITELLM_PROXY_API_KEY`: API key for authentication
+To stop exporting `LITELLM_PROXY_URL` in every shell session, store the proxy URL once in `~/.litellm/config.json`:
+
+```bash
+lite config set base_url https://your-proxy.example.com
+```
+
+Manage the stored config with:
+
+```bash
+lite config get base_url # print the stored value
+lite config get # print all stored config
+lite config unset base_url # remove the stored value
+```
+
+The base URL is resolved in this order of precedence:
+
+1. `--base-url` command-line option
+2. `LITELLM_PROXY_URL` environment variable
+3. `base_url` from `~/.litellm/config.json`
+4. `http://localhost:4000`
+
## Global Options
- `--version`, `-v`: Print the LiteLLM Proxy client and server version and exit.
@@ -581,6 +602,8 @@ The CLI respects the following environment variables:
- `LITELLM_PROXY_URL`: Base URL of the proxy server
- `LITELLM_PROXY_API_KEY`: API key for authentication
+`LITELLM_PROXY_URL` takes precedence over a `base_url` stored via `lite config set`, and the `--base-url` option overrides both. See the Configuration section for the full precedence order.
+
## Examples
1. List all models in table format:
diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py
index 61495403407..970d801dc6d 100644
--- a/litellm/proxy/client/cli/commands/auth.py
+++ b/litellm/proxy/client/cli/commands/auth.py
@@ -15,6 +15,8 @@ from rich.table import Table
from litellm.constants import CLI_JWT_EXPIRATION_HOURS
from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh
+from .private_json import write_private_json
+
# Token storage utilities
def get_token_file_path() -> str:
@@ -27,11 +29,7 @@ def get_token_file_path() -> str:
def save_token(token_data: Dict[str, Any]) -> None:
"""Save token data to file"""
- token_file = get_token_file_path()
- with open(token_file, "w") as f:
- json.dump(token_data, f, indent=2)
- # Set file permissions to be readable only by owner
- os.chmod(token_file, 0o600)
+ write_private_json(get_token_file_path(), token_data)
def load_token() -> Optional[Dict[str, Any]]:
diff --git a/litellm/proxy/client/cli/commands/config.py b/litellm/proxy/client/cli/commands/config.py
new file mode 100644
index 00000000000..851a6c11529
--- /dev/null
+++ b/litellm/proxy/client/cli/commands/config.py
@@ -0,0 +1,108 @@
+import json
+import os
+import sys
+from collections.abc import Mapping
+from pathlib import Path
+from urllib.parse import urlparse
+
+import click
+from pydantic import TypeAdapter
+
+from .private_json import write_private_json
+
+ALLOWED_CONFIG_KEYS: tuple[str, ...] = ("base_url",)
+
+_config_adapter: TypeAdapter[Mapping[str, str]] = TypeAdapter(Mapping[str, str])
+
+
+def get_config_file_path() -> str:
+ """Get the path to the persistent CLI config file"""
+ home_dir = Path.home()
+ config_dir = home_dir / ".litellm"
+ return str(config_dir / "config.json")
+
+
+def load_config() -> Mapping[str, str]:
+ """Load CLI config from file; returns {} if missing or unreadable"""
+ try:
+ config_file = get_config_file_path()
+ except RuntimeError:
+ return {}
+ if not os.path.exists(config_file):
+ return {}
+ try:
+ with open(config_file, "r") as f:
+ return _config_adapter.validate_python(json.load(f))
+ except (OSError, ValueError) as e:
+ click.echo(f"Warning: ignoring invalid config file {config_file}: {e}", err=True)
+ return {}
+
+
+def save_config(config: Mapping[str, str]) -> None:
+ """Save CLI config to file"""
+ write_private_json(get_config_file_path(), config)
+
+
+def get_config_value(key: str) -> str | None:
+ """Get a single value from the persistent CLI config"""
+ return load_config().get(key)
+
+
+@click.group(name="config")
+def config_commands() -> None:
+ """Manage persistent CLI configuration (~/.litellm/config.json)"""
+
+
+@config_commands.command(name="set")
+@click.argument("key")
+@click.argument("value")
+def set_config(key: str, value: str) -> None:
+ """Set a config KEY to VALUE (e.g. `lite config set base_url https://your-proxy.example.com`)"""
+ if key not in ALLOWED_CONFIG_KEYS:
+ raise click.UsageError(f"Unknown config key '{key}'. Allowed keys: {', '.join(ALLOWED_CONFIG_KEYS)}")
+
+ if key == "base_url":
+ parsed = urlparse(value)
+ if parsed.scheme not in ("http", "https") or not parsed.netloc:
+ raise click.UsageError("base_url must be a full http:// or https:// URL including a host")
+ if "?" in value or "#" in value:
+ raise click.UsageError("base_url must not include a query string or fragment")
+
+ normalized_value = value.rstrip("/")
+ save_config({**load_config(), key: normalized_value})
+ click.echo(f"Set {key} = {normalized_value} in {get_config_file_path()}")
+
+
+@config_commands.command(name="get")
+@click.argument("key", required=False)
+def get_config(key: str | None) -> None:
+ """Print the value of KEY, or all stored config when KEY is omitted"""
+ config = load_config()
+
+ if key is not None:
+ value = config.get(key)
+ if value is None:
+ click.echo(f"{key} is not set", err=True)
+ sys.exit(1)
+ click.echo(value)
+ return
+
+ if not config:
+ click.echo("(no config set)")
+ return
+
+ for entry_key, entry_value in config.items():
+ click.echo(f"{entry_key} = {entry_value}")
+
+
+@config_commands.command(name="unset")
+@click.argument("key")
+def unset_config(key: str) -> None:
+ """Remove KEY from the config file"""
+ config = load_config()
+ if key not in config:
+ click.echo(f"{key} was not set")
+ return
+
+ save_config({k: v for k, v in config.items() if k != key})
+ click.echo(f"Removed {key} from {get_config_file_path()}")
diff --git a/litellm/proxy/client/cli/commands/private_json.py b/litellm/proxy/client/cli/commands/private_json.py
new file mode 100644
index 00000000000..70aac0c6de0
--- /dev/null
+++ b/litellm/proxy/client/cli/commands/private_json.py
@@ -0,0 +1,20 @@
+import json
+import os
+import tempfile
+from collections.abc import Mapping
+from pathlib import Path
+
+
+def write_private_json(path: str, data: Mapping[str, object]) -> None:
+ """Atomically write JSON to path with owner-only permissions (0600)"""
+ parent = Path(path).parent
+ parent.mkdir(parents=True, exist_ok=True)
+ fd, tmp_path = tempfile.mkstemp(dir=str(parent), prefix=".tmp-", suffix=".json")
+ try:
+ with os.fdopen(fd, "w") as f:
+ json.dump(data, f, indent=2)
+ f.flush()
+ os.fsync(f.fileno())
+ os.replace(tmp_path, path)
+ finally:
+ Path(tmp_path).unlink(missing_ok=True)
diff --git a/litellm/proxy/client/cli/main.py b/litellm/proxy/client/cli/main.py
index e641956b2c5..24e5cdf747b 100644
--- a/litellm/proxy/client/cli/main.py
+++ b/litellm/proxy/client/cli/main.py
@@ -11,6 +11,7 @@ from .commands.agents import agent_commands
from .commands.auth import auth_group, get_stored_api_key, login, logout, whoami
from .commands.autoroute.commands import autoroute_group
from .commands.chat import chat
+from .commands.config import config_commands, get_config_value
from .commands.credentials import credentials
from .commands.encryption import encryption
from .commands.http import http
@@ -45,27 +46,16 @@ def print_version(base_url: str, api_key: Optional[str]):
@click.option(
"--version",
"-v",
+ "show_version",
is_flag=True,
- is_eager=True,
- expose_value=False,
help="Show the LiteLLM Proxy CLI and server version and exit.",
- callback=lambda ctx, param, value: (
- (
- print_version(
- ctx.params.get("base_url") or "http://localhost:4000",
- ctx.params.get("api_key"),
- )
- or ctx.exit()
- )
- if value and not ctx.resilient_parsing
- else None
- ),
)
@click.option(
"--base-url",
envvar="LITELLM_PROXY_URL",
show_envvar=True,
- default="http://localhost:4000",
+ default=None,
+ show_default="base_url from `lite config`, else http://localhost:4000",
help="Base URL of the LiteLLM proxy server",
)
@click.option(
@@ -75,13 +65,16 @@ def print_version(base_url: str, api_key: Optional[str]):
help="API key for authentication",
)
@click.pass_context
-def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None:
+def cli(ctx: click.Context, show_version: bool, base_url: str | None, api_key: Optional[str]) -> None:
"""LiteLLM Proxy CLI - Manage your LiteLLM proxy server"""
ctx.ensure_object(dict)
+ stored_base_url = get_config_value("base_url")
+ base_url_provided = base_url is not None
+
# Normalize once here so every downstream command (login, agents, http, ...) can safely
# do f"{base_url}/some/path" without producing a double slash.
- base_url = base_url.rstrip("/")
+ base_url = ((stored_base_url or "http://localhost:4000") if base_url is None else base_url).rstrip("/")
# If no API key provided via flag or environment variable, try to load from saved token.
# Pass base_url so we only use the stored key when it was issued for this server.
@@ -94,8 +87,13 @@ def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None:
# apiKeyHelper is invoked bare (no flags) -- commands that must work
# unattended (print-token) need to tell "user didn't say" apart from
# "user said localhost:4000 on purpose" so they can fall back to
- # whatever server the stored token was actually issued for.
- ctx.obj["base_url_explicit"] = ctx.get_parameter_source("base_url") != click.core.ParameterSource.DEFAULT
+ # whatever server the stored token was actually issued for. A base_url
+ # saved via `lite config set` counts as the user saying it.
+ ctx.obj["base_url_explicit"] = base_url_provided or bool(stored_base_url)
+
+ if show_version:
+ print_version(base_url, api_key)
+ ctx.exit()
# If no subcommand was invoked, start interactive mode
if ctx.invoked_subcommand is None:
@@ -141,6 +139,7 @@ cli.add_command(down)
cli.add_command(model_groups)
# Add the autoroute command group (QA auto-routing against your real proxy)
cli.add_command(autoroute_group, name="autoroute")
+cli.add_command(config_commands)
if __name__ == "__main__":
diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
index 282184d6495..1205d23ce02 100644
--- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py
+++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
@@ -1526,7 +1526,6 @@ if MCP_AVAILABLE:
temporary_server = await global_mcp_server_manager.build_mcp_server_from_table(
temp_record,
credentials_are_encrypted=False,
- persist_discovered_endpoints=False,
)
_cache_temporary_mcp_server(
temporary_server,
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index 70484eb1e4e..18a927e7a44 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -6758,6 +6758,9 @@ class ProxyConfig:
from litellm.proxy._experimental.mcp_server.oauth2_flow_backfill import (
backfill_null_oauth2_flows,
)
+ from litellm.proxy._experimental.mcp_server.oauth_issuer_stamp_backfill import (
+ backfill_discovery_stamped_issuers,
+ )
try:
if prisma_client is not None:
@@ -6767,6 +6770,16 @@ class ProxyConfig:
"litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db backfill - {}".format(str(e))
)
+ try:
+ if prisma_client is not None:
+ await backfill_discovery_stamped_issuers(prisma_client)
+ except Exception as e: # noqa: BLE001
+ verbose_proxy_logger.exception(
+ "litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db issuer stamp backfill - {}".format(
+ str(e)
+ )
+ )
+
try:
await global_mcp_server_manager.reload_servers_from_database()
except Exception as e:
@@ -6778,6 +6791,31 @@ class ProxyConfig:
if self._should_load_db_object(object_type="mcp"):
await self._init_mcp_servers_in_db()
+ async def reload_mcp_servers_from_db(self) -> None:
+ """Registry refresh only, for the periodic job in store_model_in_db-off deployments.
+
+ Deliberately narrower than ``init_mcp_servers_from_db``: the oauth2_flow backfill is a write
+ path that only needs to run once at startup, so the cadence here is purely the read-side
+ reload whose fast-path exemption retries failed OAuth discovery. Gated the same way, so an
+ admin who excluded mcp from supported_db_objects opts out of this too.
+ """
+ if not self._should_load_db_object(object_type="mcp"):
+ return
+ from litellm.proxy._experimental.mcp_server.utils import is_mcp_available
+
+ if not is_mcp_available():
+ return
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
+ global_mcp_server_manager,
+ )
+
+ try:
+ await global_mcp_server_manager.reload_servers_from_database()
+ except Exception as e: # noqa: BLE001 # scheduled job: a reload failure must not kill the recurring retry
+ verbose_proxy_logger.exception(
+ "litellm.proxy.proxy_server.py::ProxyConfig:reload_mcp_servers_from_db - {}".format(str(e))
+ )
+
async def _init_agents_in_db(self, prisma_client: PrismaClient):
from litellm.proxy.agent_endpoints.agent_registry import (
global_agent_registry as AGENT_REGISTRY,
@@ -8099,6 +8137,22 @@ class ProxyStartupEvent:
if store_model_in_db is not True:
await proxy_config.init_mcp_servers_from_db()
+ if prisma_client is not None:
+ # DB-backed MCP servers are live objects in every mode, so the registry refresh that
+ # store_model_in_db=True deployments get via the add_deployment job must run here
+ # too; without it, a server whose OAuth discovery failed at startup is rebuilt only
+ # by a management write, since the reload fast path is the retry's only driver.
+ mcp_reload_interval_seconds = proxy_config_reload_interval_seconds
+ if not isinstance(mcp_reload_interval_seconds, int) or mcp_reload_interval_seconds <= 0:
+ mcp_reload_interval_seconds = 30
+ scheduler.add_job(
+ proxy_config.reload_mcp_servers_from_db,
+ "interval",
+ seconds=mcp_reload_interval_seconds,
+ id="reload_mcp_servers_job",
+ replace_existing=True,
+ misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
+ )
await cls._initialize_slack_alerting_jobs(
scheduler=scheduler,
diff --git a/litellm/router.py b/litellm/router.py
index 69535d7c74a..0234a8f3424 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -9560,7 +9560,12 @@ class Router:
return None
# Strategy 1: Check if model_id directly matches a model_name or deployment ID
- if model_id in self.model_names or self.has_model_id(model_id):
+ if model_id in self.model_names:
+ return model_id
+ if self.has_model_id(model_id):
+ deployment = self.get_deployment(model_id=model_id)
+ if deployment is not None and deployment.model_name:
+ return deployment.model_name
return model_id
# Strategy 2: Search through router's model_list to find by litellm_params.model
diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py
index e5268b5107b..1da8ee68c6e 100644
--- a/litellm/router_strategy/complexity_router/complexity_router.py
+++ b/litellm/router_strategy/complexity_router/complexity_router.py
@@ -18,6 +18,7 @@ from __future__ import annotations
import asyncio
import random
import re
+from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Literal, Union, cast
from pydantic import BaseModel
@@ -25,6 +26,7 @@ from pydantic import BaseModel
from litellm._logging import verbose_router_logger
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
+from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.types.utils import ModelResponse
from .config import (
@@ -112,6 +114,16 @@ def _classifier_call_metadata(metadata: dict[str, Any] | None) -> dict[str, Any]
}
+def _effective_turn_off_message_logging(request_kwargs: Mapping[str, Any] | None) -> bool | None:
+ from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
+ initialize_standard_callback_dynamic_params,
+ )
+
+ return initialize_standard_callback_dynamic_params(dict(request_kwargs) if request_kwargs else {}).get(
+ "turn_off_message_logging"
+ )
+
+
class DimensionScore:
"""Represents a score for a single dimension with optional signal."""
@@ -427,7 +439,17 @@ class ComplexityRouter(CustomLogger):
# attributed to the calling key/team instead of being dropped. Excludes the
# parent request's budget reservation, which the routed completion (not this
# internal classifier call) is responsible for reconciling.
- metadata = _classifier_call_metadata((request_kwargs or {}).get("litellm_metadata"))
+ request_metadata = (request_kwargs or {}).get("litellm_metadata") or (request_kwargs or {}).get("metadata")
+ metadata = _classifier_call_metadata(request_metadata)
+ turn_off_message_logging = _effective_turn_off_message_logging(request_kwargs)
+
+ proxy_server_request = {
+ "body": {
+ "model": llm_config.model,
+ "messages": [{"role": "user", "content": classification_prompt}],
+ "response_format": type_to_response_format_param(TierClassification),
+ }
+ }
response: ModelResponse = await self.litellm_router_instance.acompletion(
model=llm_config.model,
@@ -435,6 +457,8 @@ class ComplexityRouter(CustomLogger):
response_format=TierClassification,
timeout=llm_config.timeout_ms / 1000,
metadata=metadata,
+ proxy_server_request=proxy_server_request,
+ turn_off_message_logging=turn_off_message_logging,
)
content = response.choices[0].message.content
if not content:
@@ -821,8 +845,16 @@ class ComplexityRouter(CustomLogger):
# key/team budget. Key/team attribution fields are preserved for spend logging.
metadata = _classifier_call_metadata(request_kwargs.get("metadata"))
litellm_metadata = _classifier_call_metadata(request_kwargs.get("litellm_metadata"))
+ turn_off_message_logging = _effective_turn_off_message_logging(request_kwargs)
+ proxy_server_request = {"body": {"model": self.config.embedding_model, "input": [user_message]}}
query_vector = (
- await encoder.aencode_queries([user_message], metadata=metadata, litellm_metadata=litellm_metadata)
+ await encoder.aencode_queries(
+ [user_message],
+ metadata=metadata,
+ litellm_metadata=litellm_metadata,
+ proxy_server_request=proxy_server_request,
+ turn_off_message_logging=turn_off_message_logging,
+ )
)[0]
route_choice = await routelayer.acall(vector=query_vector)
diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py
index fb3ddeebf52..da1ff7eda67 100644
--- a/litellm/types/llms/vertex_ai.py
+++ b/litellm/types/llms/vertex_ai.py
@@ -16,7 +16,7 @@ GeminiEmbeddingInput = Union[EmbeddingInput, List[List[str]]]
class FunctionResponse(TypedDict, total=False):
# `id` correlates this response with the originating `functionCall` part.
- # Supported on Google AI Studio Gemini 3.5+; Vertex AI rejects this field.
+ # Supported on Gemini 3+; older Gemini models reject this field.
id: str
name: Required[str]
response: Optional[dict]
@@ -24,8 +24,8 @@ class FunctionResponse(TypedDict, total=False):
class FunctionCall(TypedDict, total=False):
- # `id` correlates the corresponding `functionResponse` on Google AI Studio
- # Gemini 3.5+. Vertex AI and older Gemini models omit/reject this field.
+ # `id` correlates the corresponding `functionResponse` on Gemini 3+.
+ # Older Gemini models omit/reject this field.
id: str
name: Required[str]
args: Optional[dict]
@@ -58,8 +58,8 @@ class PartType(TypedDict, total=False):
class HttpxFunctionCall(TypedDict, total=False):
- # `id` correlates the corresponding `functionResponse` on Google AI Studio
- # Gemini 3.5+. Vertex AI and older Gemini models omit/reject this field.
+ # `id` correlates the corresponding `functionResponse` on Gemini 3+.
+ # Older Gemini models omit/reject this field.
id: str
name: Required[str]
args: dict
diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json
index e5dd022143c..eec8489c192 100644
--- a/ruff-strict-budget.json
+++ b/ruff-strict-budget.json
@@ -1,6 +1,6 @@
{
"ANN001": {
- "limit": 3142
+ "limit": 3094
},
"ANN002": {
"limit": 69
@@ -24,7 +24,7 @@
"limit": 130
},
"ANN401": {
- "limit": 2013
+ "limit": 2012
},
"ASYNC230": {
"limit": 14
@@ -33,7 +33,7 @@
"limit": 4
},
"B006": {
- "limit": 190
+ "limit": 186
},
"B008": {
"limit": 505
@@ -42,7 +42,7 @@
"limit": 84
},
"B010": {
- "limit": 197
+ "limit": 191
},
"B018": {
"limit": 5
@@ -60,7 +60,7 @@
"limit": 4
},
"BLE001": {
- "limit": 2900
+ "limit": 2895
},
"C401": {
"limit": 11
@@ -81,7 +81,7 @@
"limit": 4
},
"C901": {
- "limit": 316
+ "limit": 312
},
"D419": {
"limit": 9
@@ -123,7 +123,7 @@
"limit": 52
},
"I001": {
- "limit": 267
+ "limit": 269
},
"LOG015": {
"limit": 8
@@ -180,7 +180,7 @@
"limit": 34
},
"PLR1714": {
- "limit": 265
+ "limit": 257
},
"PLR1730": {
"limit": 10
@@ -189,7 +189,7 @@
"limit": 4
},
"PLW0127": {
- "limit": 44
+ "limit": 42
},
"PLW0133": {
"limit": 4
@@ -222,7 +222,7 @@
"limit": 38
},
"RET504": {
- "limit": 717
+ "limit": 714
},
"RUF010": {
"limit": 874
@@ -237,7 +237,7 @@
"limit": 41
},
"RUF022": {
- "limit": 83
+ "limit": 84
},
"RUF023": {
"limit": 5
@@ -261,7 +261,7 @@
"limit": 24
},
"SIM101": {
- "limit": 63
+ "limit": 59
},
"SIM102": {
"limit": 324
@@ -273,7 +273,7 @@
"limit": 6
},
"SIM114": {
- "limit": 113
+ "limit": 109
},
"SIM115": {
"limit": 5
@@ -288,7 +288,7 @@
"limit": 4
},
"SIM210": {
- "limit": 12
+ "limit": 10
},
"SIM211": {
"limit": 4
@@ -306,10 +306,10 @@
"limit": 9
},
"TID251": {
- "limit": 2650
+ "limit": 2651
},
"TRY002": {
- "limit": 548
+ "limit": 546
},
"TRY004": {
"limit": 98
@@ -324,7 +324,7 @@
"limit": 883
},
"UP006": {
- "limit": 12147
+ "limit": 12145
},
"UP007": {
"limit": 2526
@@ -348,13 +348,13 @@
"limit": 5
},
"UP032": {
- "limit": 629
+ "limit": 625
},
"UP034": {
"limit": 4
},
"UP035": {
- "limit": 2230
+ "limit": 2232
},
"UP036": {
"limit": 4
@@ -363,6 +363,6 @@
"limit": 103
},
"UP045": {
- "limit": 17816
+ "limit": 17777
}
}
diff --git a/tests/litellm_utils_tests/test_secret_manager.py b/tests/litellm_utils_tests/test_secret_manager.py
index 0a2419d0bea..0f95fd75c53 100644
--- a/tests/litellm_utils_tests/test_secret_manager.py
+++ b/tests/litellm_utils_tests/test_secret_manager.py
@@ -19,7 +19,8 @@ sys.path.insert(
import pytest
import litellm
from litellm.llms.azure.azure import get_azure_ad_token_from_oidc
-from litellm.llms.bedrock.chat import BedrockConverseLLM, BedrockLLM
+from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
+from litellm.llms.bedrock.chat import BedrockConverseLLM
from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2
from litellm.secret_managers.main import (
get_secret,
@@ -160,7 +161,7 @@ def test_oidc_circle_v1_with_amazon():
aws_role_name = "arn:aws:iam::335785316107:role/litellm-github-unit-tests-circleci-v1-assume-only"
aws_web_identity_token = "oidc/circleci/"
- bllm = BedrockLLM()
+ bllm = BaseAWSLLM()
creds = bllm.get_credentials(
aws_region_name="ca-west-1",
aws_web_identity_token=aws_web_identity_token,
diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py
index f4c307e9c8a..8ab2feaf896 100644
--- a/tests/llm_translation/test_bedrock_completion.py
+++ b/tests/llm_translation/test_bedrock_completion.py
@@ -33,7 +33,7 @@ from litellm import (
completion_cost,
embedding,
)
-from litellm.llms.bedrock.chat import BedrockLLM
+from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_tools_pt
from base_llm_unit_tests import BaseLLMChatTest, BaseAnthropicChatTest
@@ -225,7 +225,7 @@ def bedrock_session_token_creds():
aws_region_name = os.environ["AWS_REGION_NAME"]
aws_session_token = os.environ.get("AWS_SESSION_TOKEN")
- bllm = BedrockLLM()
+ bllm = BaseAWSLLM()
if aws_session_token is not None:
# For local testing
creds = bllm.get_credentials(
@@ -3573,40 +3573,11 @@ def test_bedrock_openai_model_id_extraction():
print(f"✓ Model ID extracted and encoded: {model_id}")
-def test_bedrock_openai_convert_messages_to_prompt():
- """
- Test that convert_messages_to_prompt returns empty string for OpenAI models.
- """
- from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
-
- bedrock_llm = BedrockLLM()
- messages = [
- {"role": "system", "content": "You are helpful"},
- {"role": "user", "content": "Hello"},
- ]
-
- prompt, chat_history = bedrock_llm.convert_messages_to_prompt(
- model="test-model", messages=messages, provider="openai", custom_prompt_dict={}
+def test_bedrock_openai_response_parsing():
+ from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import (
+ AmazonBedrockOpenAIConfig,
)
- # OpenAI models use messages directly, no prompt conversion
- assert prompt == ""
- assert chat_history is None
- print("✓ convert_messages_to_prompt returns empty for OpenAI")
-
-
-def test_bedrock_openai_response_parsing():
- """
- Test that OpenAI responses are correctly parsed.
- """
- from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
- from litellm import ModelResponse
- from unittest.mock import Mock
- import json
-
- bedrock_llm = BedrockLLM()
-
- # Mock OpenAI-style response
openai_response = {
"choices": [
{
@@ -3627,34 +3598,24 @@ def test_bedrock_openai_response_parsing():
mock_response.status_code = 200
mock_response.headers = {}
- model_response = ModelResponse()
- mock_logging = Mock()
-
- result = bedrock_llm.process_response(
+ result = AmazonBedrockOpenAIConfig().transform_response(
model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test",
- response=mock_response,
- model_response=model_response,
- stream=False,
- logging_obj=mock_logging,
- optional_params={},
- api_key="",
- data={},
+ raw_response=mock_response,
+ model_response=ModelResponse(),
+ logging_obj=Mock(),
+ request_data={},
messages=[{"role": "user", "content": "What is the capital of France?"}],
- print_verbose=lambda x: None,
+ optional_params={},
+ litellm_params={},
encoding=None,
)
- # Verify response content
assert result.choices[0].message.content == "The capital of France is Paris."
assert result.choices[0].finish_reason == "stop"
-
- # Verify usage
assert result.usage.prompt_tokens == 10
assert result.usage.completion_tokens == 8
assert result.usage.total_tokens == 18
- print("✓ OpenAI response parsing works correctly")
-
def test_bedrock_openai_request_transformation():
"""
@@ -3846,43 +3807,20 @@ def test_bedrock_openai_multiple_message_types():
def test_bedrock_openai_error_handling():
- """
- Test that errors from OpenAI models are properly handled.
- """
- from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
- from litellm import ModelResponse
+ from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import (
+ AmazonBedrockOpenAIConfig,
+ )
from litellm.llms.bedrock.common_utils import BedrockError
- from unittest.mock import Mock
- import json
- bedrock_llm = BedrockLLM()
+ error = AmazonBedrockOpenAIConfig().get_error_class(
+ error_message="ValidationException: bad request",
+ status_code=422,
+ headers={},
+ )
- # Mock error response
- mock_response = Mock()
- mock_response.json.side_effect = Exception("Invalid JSON")
- mock_response.text = "Invalid response"
- mock_response.status_code = 422
-
- model_response = ModelResponse()
- mock_logging = Mock()
-
- with pytest.raises(BedrockError) as exc_info:
- bedrock_llm.process_response(
- model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test",
- response=mock_response,
- model_response=model_response,
- stream=False,
- logging_obj=mock_logging,
- optional_params={},
- api_key="",
- data={},
- messages=[],
- print_verbose=lambda x: None,
- encoding=None,
- )
-
- assert exc_info.value.status_code == 422
- print("✓ Error handling works correctly")
+ assert isinstance(error, BedrockError)
+ assert error.status_code == 422
+ assert "ValidationException: bad request" in str(error)
# ============================================================================
diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py
index a969d21a681..bcc70fae67c 100644
--- a/tests/router_unit_tests/test_router_helper_utils.py
+++ b/tests/router_unit_tests/test_router_helper_utils.py
@@ -2659,6 +2659,23 @@ def test_resolve_model_name_from_model_id():
result = router.resolve_model_name_from_model_id("gpt-5-mini")
assert result == "gpt-5-mini"
+ # Test case 10: model_id is a deployment ID (hash) that differs from the
+ # public model_name. Regression for #32580: managed batch/file IDs embed the
+ # deployment model_id, and it must resolve back to the public model_name so
+ # team model-access checks compare against the model group, not the hash.
+ model_list = [
+ {
+ "model_name": "bedrock-batch-model",
+ "litellm_params": {
+ "model": "bedrock/global.anthropic.claude-haiku-4-5-20251001-v1:0",
+ },
+ "model_info": {"id": "8d0eaa7e6c6f54a425dfd0062cb6b0dc"},
+ },
+ ]
+ router = Router(model_list=model_list)
+ result = router.resolve_model_name_from_model_id("8d0eaa7e6c6f54a425dfd0062cb6b0dc")
+ assert result == "bedrock-batch-model"
+
def test_get_valid_args():
"""Test get_valid_args static method returns valid Router.__init__ arguments"""
diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py
index 3169b9b08e0..2580197d6d2 100644
--- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py
+++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py
@@ -5,6 +5,8 @@ Regression test for afile_retrieve called without credentials in
async_post_call_success_hook when processing completed batch responses.
"""
+import json
+
import pytest
from typing import Optional
from unittest.mock import AsyncMock, MagicMock, patch
@@ -385,3 +387,59 @@ async def test_afile_content_error_reports_unified_id_not_provider_uri():
message = str(exc_info.value)
assert unified_file_id in message
assert s3_uri not in message
+
+
+def _make_real_managed_files_instance():
+ """Create a _PROXY_LiteLLMManagedFiles with a real store_unified_file_id but
+ an AsyncMock prisma client, so the DB write path itself can be asserted."""
+ from litellm_enterprise.proxy.hooks.managed_files import (
+ _PROXY_LiteLLMManagedFiles,
+ )
+
+ mock_cache = MagicMock()
+ mock_cache.async_set_cache = AsyncMock()
+
+ mock_prisma = MagicMock()
+ mock_prisma.db.litellm_managedfiletable.upsert = AsyncMock()
+ mock_prisma.db.litellm_managedfiletable.create = AsyncMock(
+ side_effect=AssertionError(
+ "store_unified_file_id must upsert, not create, on the retrieve path"
+ )
+ )
+
+ return (
+ _PROXY_LiteLLMManagedFiles(
+ internal_usage_cache=mock_cache,
+ prisma_client=mock_prisma,
+ ),
+ mock_prisma,
+ )
+
+
+@pytest.mark.asyncio
+async def test_store_unified_file_id_is_idempotent_via_upsert():
+ """Regression test for the managed-batch retrieve 500 (UniqueViolationError on
+ unified_file_id): re-registering an already-stored output file id must upsert on
+ unified_file_id, never do an unconditional create that raises on conflict."""
+ managed_files, mock_prisma = _make_real_managed_files_instance()
+ file_id = "litellm_proxy_unified_output_id_abc"
+ model_mappings = {"model-deploy-xyz": "file-output-abc"}
+
+ for _ in range(2):
+ await managed_files.store_unified_file_id(
+ file_id=file_id,
+ file_object=_make_file_object(),
+ litellm_parent_otel_span=None,
+ model_mappings=model_mappings,
+ user_api_key_dict=_make_user_api_key_dict(),
+ )
+
+ mock_prisma.db.litellm_managedfiletable.create.assert_not_awaited()
+ upsert_mock = mock_prisma.db.litellm_managedfiletable.upsert
+ assert upsert_mock.await_count == 2
+ for upsert_call in upsert_mock.await_args_list:
+ assert upsert_call.kwargs["where"] == {"unified_file_id": file_id}
+ upsert_data = upsert_call.kwargs["data"]
+ assert upsert_data["create"]["unified_file_id"] == file_id
+ assert json.loads(upsert_data["create"]["model_mappings"]) == model_mappings
+ assert json.loads(upsert_data["update"]["model_mappings"]) == model_mappings
diff --git a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py
index 61987d25d9c..ee50b9db015 100644
--- a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py
+++ b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py
@@ -8,14 +8,11 @@ sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
-import litellm
from litellm.llms.bedrock.chat.invoke_handler import (
AWSEventStreamDecoder,
- BedrockLLM,
make_call,
make_sync_call,
)
-from litellm.llms.custom_httpx.http_handler import HTTPHandler
def test_transform_thinking_blocks_with_redacted_content():
@@ -296,33 +293,3 @@ def test_make_sync_call_honors_explicit_stream_chunk_size():
response.iter_bytes.assert_called_once_with(chunk_size=2048)
-
-def test_legacy_bedrock_llm_streaming_does_not_rechunk_by_default():
- mock_response = MagicMock()
- mock_response.status_code = 200
- mock_response.iter_bytes = MagicMock(return_value=iter([]))
- client = HTTPHandler()
- client.post = MagicMock(return_value=mock_response)
-
- BedrockLLM().completion(
- model="cohere.command-text-v14",
- messages=[{"role": "user", "content": "hi"}],
- api_base=None,
- custom_prompt_dict={},
- model_response=litellm.ModelResponse(),
- print_verbose=lambda *args, **kwargs: None,
- encoding=litellm.encoding,
- logging_obj=MagicMock(),
- optional_params={
- "stream": True,
- "aws_access_key_id": "fake",
- "aws_secret_access_key": "fake",
- "aws_region_name": "us-east-1",
- },
- acompletion=False,
- timeout=None,
- litellm_params={},
- client=client,
- )
-
- mock_response.iter_bytes.assert_called_once_with(chunk_size=None)
diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py
index cf75964ddb7..ad890d0c7ea 100644
--- a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py
+++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py
@@ -1452,6 +1452,157 @@ class TestContextCachingEndpoints:
# Restart the patcher so teardown_method can stop it cleanly
self._token_check_patcher.start()
+ def _model_turn_final_messages(self, final_cached_role):
+ tool_call = {
+ "id": "call_abc123",
+ "type": "function",
+ "function": {"name": "get_weather", "arguments": '{"location": "Boston"}'},
+ }
+ cached_tail = {
+ "assistant": [],
+ "tool": [
+ {
+ "role": "tool",
+ "tool_call_id": "call_abc123",
+ "content": "72F and sunny",
+ "cache_control": {"type": "ephemeral"},
+ }
+ ],
+ "system": [
+ {
+ "role": "system",
+ "content": "Tool results are authoritative.",
+ "cache_control": {"type": "ephemeral"},
+ }
+ ],
+ }[final_cached_role]
+ return [
+ {
+ "role": "user",
+ "content": [
+ {
+ "type": "text",
+ "text": "Use the weather tool for every answer.",
+ "cache_control": {"type": "ephemeral"},
+ }
+ ],
+ },
+ {
+ "role": "assistant",
+ "content": "",
+ "tool_calls": [tool_call],
+ "cache_control": {"type": "ephemeral"},
+ },
+ *cached_tail,
+ {"role": "user", "content": "What is the weather in Boston?"},
+ ]
+
+ @pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"])
+ def test_check_and_create_cache_skips_when_cached_block_ends_on_model_turn(
+ self, final_cached_role
+ ):
+ """The cachedContents API rejects contents ending on an assistant or tool turn
+ with HTTP 400 "Requests ending with a model turn are not supported", so the
+ request must proceed uncached instead of failing.
+ """
+ all_messages = self._model_turn_final_messages(final_cached_role)
+ optional_params = self.sample_optional_params.copy()
+
+ result = self.context_caching.check_and_create_cache(
+ messages=all_messages,
+ optional_params=optional_params,
+ api_key="test_key",
+ api_base=None,
+ model="gemini-3.6-flash",
+ client=self.mock_client,
+ timeout=30.0,
+ logging_obj=self.mock_logging,
+ cached_content=None,
+ custom_llm_provider="vertex_ai",
+ vertex_project="test_project",
+ vertex_location="us-central1",
+ vertex_auth_header="test_token",
+ )
+
+ messages, returned_params, returned_cache = result
+ assert messages == all_messages
+ assert returned_cache is None
+ assert "tools" in returned_params
+ self.mock_client.get.assert_not_called()
+ self.mock_client.post.assert_not_called()
+
+ @pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"])
+ @pytest.mark.asyncio
+ async def test_async_check_and_create_cache_skips_when_cached_block_ends_on_model_turn(
+ self, final_cached_role
+ ):
+ """Async variant: an unsupported terminal turn skips caching instead of failing."""
+ all_messages = self._model_turn_final_messages(final_cached_role)
+ optional_params = self.sample_optional_params.copy()
+
+ result = await self.context_caching.async_check_and_create_cache(
+ messages=all_messages,
+ optional_params=optional_params,
+ api_key="test_key",
+ api_base=None,
+ model="gemini-3.6-flash",
+ client=self.mock_async_client,
+ timeout=30.0,
+ logging_obj=self.mock_logging,
+ cached_content=None,
+ custom_llm_provider="vertex_ai",
+ vertex_project="test_project",
+ vertex_location="us-central1",
+ vertex_auth_header="test_token",
+ )
+
+ messages, returned_params, returned_cache = result
+ assert messages == all_messages
+ assert returned_cache is None
+ assert "tools" in returned_params
+ self.mock_async_client.get.assert_not_called()
+ self.mock_async_client.post.assert_not_called()
+
+
+def test_cached_messages_end_on_supported_turn():
+ from litellm.llms.vertex_ai.context_caching.transformation import (
+ cached_messages_end_on_supported_turn,
+ )
+
+ assert (
+ cached_messages_end_on_supported_turn(
+ [{"role": "assistant", "content": "hi"}, {"role": "user", "content": "hello"}]
+ )
+ is True
+ )
+ assert cached_messages_end_on_supported_turn([{"role": "system", "content": "be brief"}]) is True
+ assert cached_messages_end_on_supported_turn([{"role": "assistant", "content": "hi"}]) is False
+ assert (
+ cached_messages_end_on_supported_turn(
+ [
+ {"role": "user", "content": "hello"},
+ {"role": "assistant", "content": "hi"},
+ {"role": "system", "content": "be brief"},
+ ]
+ )
+ is False
+ )
+ assert (
+ cached_messages_end_on_supported_turn(
+ [{"role": "system", "content": "be brief"}, {"role": "user", "content": "hello"}]
+ )
+ is True
+ )
+ assert (
+ cached_messages_end_on_supported_turn([{"role": "tool", "tool_call_id": "x", "content": "y"}])
+ is False
+ )
+ assert (
+ cached_messages_end_on_supported_turn([{"role": "function", "name": "f", "content": "y"}])
+ is False
+ )
+ assert cached_messages_end_on_supported_turn([]) is False
+
class TestCheckCachePagination:
"""Test pagination logic in check_cache and async_check_cache methods."""
diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py
index 453a0c14bf9..5e854bbad70 100644
--- a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py
+++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py
@@ -31,10 +31,7 @@ class TestVertexAIFilesHandler:
def test_extract_bucket_and_object_from_file_id_standard_path(self):
"""Test extraction of bucket and object from URL-encoded file_id with standard path"""
# Sample file_id with nested folder structure
- file_id = (
- "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
- "%2Ftest-folder%2Fsub-folder%2Ftest-file.txt"
- )
+ file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Ftest-folder%2Fsub-folder%2Ftest-file.txt"
bucket_name, object_path = self.handler._extract_bucket_and_object_from_file_id(
file_id=file_id,
@@ -105,21 +102,14 @@ class TestVertexAIFilesHandler:
async def test_afile_content_success(self):
"""Test successful async file content retrieval"""
# Setup test data
- file_id = (
- "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
- "%2Fuploads%2Fabc-test-file.txt"
- )
+ file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-test-file.txt"
expected_content = b"test file content"
- file_content_request = FileContentRequest(
- file_id=file_id, extra_headers=None, extra_body=None
- )
+ file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
# Mock the download_gcs_object method
with (
- patch.object(
- self.handler, "download_gcs_object", new_callable=AsyncMock
- ) as mock_download,
+ patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
patch.object(
self.handler,
"get_gcs_logging_config",
@@ -148,15 +138,9 @@ class TestVertexAIFilesHandler:
# Verify the download was called with correct parameters
mock_download.assert_called_once()
call_args = mock_download.call_args
- assert (
- call_args.kwargs["object_name"]
- == "litellm-vertex-files/uploads/abc-test-file.txt"
- )
+ assert call_args.kwargs["object_name"] == "litellm-vertex-files/uploads/abc-test-file.txt"
assert "standard_callback_dynamic_params" in call_args.kwargs
- assert (
- call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"]
- == "test-bucket"
- )
+ assert call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"] == "test-bucket"
@pytest.mark.asyncio
async def test_afile_content_missing_file_id(self):
@@ -164,9 +148,7 @@ class TestVertexAIFilesHandler:
file_content_request = FileContentRequest(extra_headers=None, extra_body=None)
# Should raise ValueError for missing file_id
- with pytest.raises(
- ValueError, match="file_id is required in file_content_request"
- ):
+ with pytest.raises(ValueError, match="file_id is required in file_content_request"):
await self.handler.afile_content(
file_content_request=file_content_request,
vertex_credentials=None,
@@ -179,20 +161,13 @@ class TestVertexAIFilesHandler:
@pytest.mark.asyncio
async def test_afile_content_download_failure(self):
"""Test async file content retrieval when download fails"""
- file_id = (
- "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files"
- "%2Fuploads%2Fabc-test-file.txt"
- )
+ file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-test-file.txt"
- file_content_request = FileContentRequest(
- file_id=file_id, extra_headers=None, extra_body=None
- )
+ file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
# Mock download to return None (failure)
with (
- patch.object(
- self.handler, "download_gcs_object", new_callable=AsyncMock
- ) as mock_download,
+ patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
patch.object(
self.handler,
"get_gcs_logging_config",
@@ -216,14 +191,130 @@ class TestVertexAIFilesHandler:
max_retries=3,
)
+ def test_resolve_read_gcs_config_prefers_per_model_bucket(self, monkeypatch):
+ monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket")
+ monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json")
+
+ bucket, service_account = self.handler._resolve_read_gcs_config(
+ litellm_params={
+ "gcs_bucket_name": "my-model-bucket",
+ "vertex_credentials": "/model/sa.json",
+ },
+ vertex_credentials=None,
+ )
+
+ assert bucket == "my-model-bucket"
+ assert service_account == "/model/sa.json"
+
+ def test_resolve_read_gcs_config_falls_back_to_env(self, monkeypatch):
+ monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket")
+ monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json")
+
+ bucket, service_account = self.handler._resolve_read_gcs_config(litellm_params={}, vertex_credentials=None)
+
+ assert bucket == "env-default-bucket"
+ assert service_account == "/env/sa.json"
+
+ def test_resolve_read_gcs_config_serializes_dict_credentials(self, monkeypatch):
+ monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
+
+ _, service_account = self.handler._resolve_read_gcs_config(
+ litellm_params={"gcs_bucket_name": "my-model-bucket"},
+ vertex_credentials={"type": "service_account", "project_id": "p"},
+ )
+
+ assert service_account == '{"type": "service_account", "project_id": "p"}'
+
+ @pytest.mark.asyncio
+ async def test_afile_content_honors_per_model_bucket_over_env(self, monkeypatch):
+ """
+ Regression for #32640: a batch output written to a per-model gcs_bucket_name must be
+ readable even when the global GCS_BUCKET_NAME points at a different bucket. Before the
+ fix the read path resolved the bucket from env only and raised
+ "file_id bucket does not match the configured storage bucket".
+ """
+ monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket")
+ monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
+
+ file_id = "gs%3A%2F%2Fmy-model-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-batch-output.jsonl"
+ file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
+
+ with (
+ patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
+ patch.object(
+ self.handler,
+ "get_or_create_vertex_instance",
+ new_callable=AsyncMock,
+ return_value=object(),
+ ),
+ ):
+ mock_download.return_value = b"batch output"
+
+ result = await self.handler.afile_content(
+ file_content_request=file_content_request,
+ vertex_credentials="/model/sa.json",
+ vertex_project="test-project",
+ vertex_location="us-central1",
+ timeout=60.0,
+ max_retries=0,
+ litellm_params={
+ "gcs_bucket_name": "my-model-bucket",
+ "vertex_credentials": "/model/sa.json",
+ },
+ )
+
+ assert isinstance(result, HttpxBinaryResponseContent)
+ assert result.response.content == b"batch output"
+
+ dynamic_params = mock_download.call_args.kwargs["standard_callback_dynamic_params"]
+ assert dynamic_params["gcs_bucket_name"] == "my-model-bucket"
+ assert dynamic_params["gcs_path_service_account"] == "/model/sa.json"
+ assert mock_download.call_args.kwargs["object_name"] == "litellm-vertex-files/uploads/abc-batch-output.jsonl"
+
+ @pytest.mark.asyncio
+ async def test_afile_content_reads_without_global_env_bucket(self, monkeypatch):
+ """
+ Regression for #32640: with no global GCS_BUCKET_NAME set, a model-group-level
+ deployment (per-model gcs_bucket_name) must still be readable. Before the fix the read
+ path raised "GCS_BUCKET_NAME is not set in the environment".
+ """
+ monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
+ monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
+
+ file_id = "gs%3A%2F%2Fmy-model-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-batch-output.jsonl"
+ file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
+
+ with (
+ patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download,
+ patch.object(
+ self.handler,
+ "get_or_create_vertex_instance",
+ new_callable=AsyncMock,
+ return_value=object(),
+ ),
+ ):
+ mock_download.return_value = b"batch output"
+
+ result = await self.handler.afile_content(
+ file_content_request=file_content_request,
+ vertex_credentials="/model/sa.json",
+ vertex_project="test-project",
+ vertex_location="us-central1",
+ timeout=60.0,
+ max_retries=0,
+ litellm_params={"gcs_bucket_name": "my-model-bucket"},
+ )
+
+ assert isinstance(result, HttpxBinaryResponseContent)
+ dynamic_params = mock_download.call_args.kwargs["standard_callback_dynamic_params"]
+ assert dynamic_params["gcs_bucket_name"] == "my-model-bucket"
+
def test_file_content_sync_success(self):
"""Test successful sync file content retrieval"""
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
expected_content = b"test file content"
- file_content_request = FileContentRequest(
- file_id=file_id, extra_headers=None, extra_body=None
- )
+ file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
# Create expected response
mock_response = httpx.Response(
@@ -261,25 +352,17 @@ class TestVertexAIFilesHandler:
file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt"
expected_content = b"test file content"
- file_content_request = FileContentRequest(
- file_id=file_id, extra_headers=None, extra_body=None
- )
+ file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None)
# Mock the afile_content method
- with patch.object(
- self.handler, "afile_content", new_callable=AsyncMock
- ) as mock_afile_content:
+ with patch.object(self.handler, "afile_content", new_callable=AsyncMock) as mock_afile_content:
mock_response = httpx.Response(
status_code=200,
content=expected_content,
headers={"content-type": "application/octet-stream"},
- request=httpx.Request(
- method="GET", url="gs://test-bucket/test-file.txt"
- ),
- )
- mock_afile_content.return_value = HttpxBinaryResponseContent(
- response=mock_response
+ request=httpx.Request(method="GET", url="gs://test-bucket/test-file.txt"),
)
+ mock_afile_content.return_value = HttpxBinaryResponseContent(response=mock_response)
# Call the method with _is_async=True
result = self.handler.file_content(
diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py
index 5871644ca5f..51cc2857252 100644
--- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py
+++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py
@@ -2276,82 +2276,8 @@ def test_is_gemini_3_or_newer():
assert VertexGeminiConfig._is_gemini_3_or_newer("") == False
-def test_forward_gemini_function_call_id_vertex_vs_google_ai_studio():
- """Vertex AI rejects `id` on function_call/function_response; Google AI Studio accepts it on Gemini 3.5+."""
- from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
- VertexGeminiConfig,
- )
-
- model = "gemini-3.5-flash"
- assert (
- VertexGeminiConfig._forward_gemini_function_call_id(model, "vertex_ai") is False
- )
- assert (
- VertexGeminiConfig._forward_gemini_function_call_id(model, "vertex_ai_beta")
- is False
- )
- assert VertexGeminiConfig._forward_gemini_function_call_id(model, "gemini") is True
- assert VertexGeminiConfig._forward_gemini_function_call_id(model, None) is False
- assert (
- VertexGeminiConfig._forward_gemini_function_call_id(
- "gemini-2.5-flash", "gemini"
- )
- is False
- )
-
-
-def test_vertex_ai_gemini_35_tool_calls_omit_function_call_id():
- """Regression: Vertex must not send OpenAI tool_call id inside Gemini function_call parts."""
- from litellm.llms.vertex_ai.gemini.transformation import (
- _gemini_convert_messages_with_history,
- )
-
- messages = [
- {"role": "user", "content": "Explore this directory"},
- {
- "role": "assistant",
- "content": "",
- "tool_calls": [
- {
- "id": "call_50e7e0fe0989464a89f188eda443",
- "type": "function",
- "function": {
- "name": "read",
- "arguments": '{"filePath": "/tmp"}',
- },
- }
- ],
- },
- {
- "role": "tool",
- "tool_call_id": "call_50e7e0fe0989464a89f188eda443",
- "content": "ok",
- },
- ]
-
- contents = _gemini_convert_messages_with_history(
- messages=messages,
- model="gemini-3.5-flash",
- custom_llm_provider="vertex_ai",
- )
-
- for content in contents:
- for part in content.get("parts", []):
- fc = part.get("function_call")
- if fc is not None:
- assert "id" not in fc, f"Vertex payload must not include id: {fc}"
- fr = part.get("function_response")
- if fr is not None:
- assert "id" not in fr, f"Vertex payload must not include id: {fr}"
-
-
-def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id():
- from litellm.llms.vertex_ai.gemini.transformation import (
- _gemini_convert_messages_with_history,
- )
-
- tool_call_id = "call_50e7e0fe0989464a89f188eda443"
- messages = [
+def _tool_call_messages(tool_call_id: str):
+ return [
{"role": "user", "content": "hi"},
{
"role": "assistant",
@@ -2374,12 +2300,8 @@ def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id():
},
]
- contents = _gemini_convert_messages_with_history(
- messages=messages,
- model="gemini-3.5-flash",
- custom_llm_provider="gemini",
- )
+def _collect_function_call_ids(contents):
function_call_ids = []
function_response_ids = []
for content in contents:
@@ -2390,9 +2312,120 @@ def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id():
fr = part.get("function_response")
if fr is not None:
function_response_ids.append(fr.get("id"))
+ return function_call_ids, function_response_ids
- assert function_call_ids == [tool_call_id]
- assert function_response_ids == [tool_call_id]
+
+def test_forward_gemini_function_call_id_is_gated_on_model_version_only():
+ """Gemini 3+ takes `id` on Vertex AI and Google AI Studio alike; older models reject it."""
+ from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
+ VertexGeminiConfig,
+ )
+
+ assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-3.5-flash") is True
+ assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-3-pro") is True
+ assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-2.5-flash") is False
+ assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-2.0-flash") is False
+
+
+@pytest.mark.parametrize("custom_llm_provider", ["vertex_ai", "vertex_ai_beta", "gemini"])
+def test_gemini_35_tool_calls_include_function_call_id(custom_llm_provider):
+ """Vertex AI accepts `id` on Gemini 3+, so it must be sent there and not just on AI Studio.
+
+ Both parts are asserted together: Vertex pairs a result to its call by id, so emitting one
+ side without the other would break strict tool-call matching.
+ """
+ from litellm.llms.vertex_ai.gemini.transformation import (
+ _gemini_convert_messages_with_history,
+ )
+
+ tool_call_id = "call_50e7e0fe0989464a89f188eda443"
+ contents = _gemini_convert_messages_with_history(
+ messages=_tool_call_messages(tool_call_id),
+ model="gemini-3.5-flash",
+ custom_llm_provider=custom_llm_provider,
+ )
+
+ assert _collect_function_call_ids(contents) == ([tool_call_id], [tool_call_id])
+
+
+@pytest.mark.parametrize("custom_llm_provider", ["vertex_ai", "gemini"])
+def test_gemini_25_tool_calls_omit_function_call_id(custom_llm_provider):
+ """Regression: models older than Gemini 3 reject `id`, so the key must be absent entirely."""
+ from litellm.llms.vertex_ai.gemini.transformation import (
+ _gemini_convert_messages_with_history,
+ )
+
+ contents = _gemini_convert_messages_with_history(
+ messages=_tool_call_messages("call_50e7e0fe0989464a89f188eda443"),
+ model="gemini-2.5-flash",
+ custom_llm_provider=custom_llm_provider,
+ )
+
+ for content in contents:
+ for part in content.get("parts", []):
+ fc = part.get("function_call")
+ if fc is not None:
+ assert "id" not in fc, f"gemini-2.5 payload must not include id: {fc}"
+ fr = part.get("function_response")
+ if fr is not None:
+ assert "id" not in fr, f"gemini-2.5 payload must not include id: {fr}"
+
+
+def test_vertex_ai_forwarded_function_call_id_strips_thought_signature_suffix():
+ """The thought signature rides along on the OpenAI id but must not reach Vertex.
+
+ Vertex now sees this code path for the first time, so the suffix has to be stripped here too.
+ """
+ from litellm.llms.vertex_ai.gemini.transformation import (
+ _gemini_convert_messages_with_history,
+ )
+ from litellm.litellm_core_utils.prompt_templates.factory import (
+ THOUGHT_SIGNATURE_SEPARATOR,
+ )
+
+ bare_id = "call_50e7e0fe0989464a89f188eda443"
+ contents = _gemini_convert_messages_with_history(
+ messages=_tool_call_messages(f"{bare_id}{THOUGHT_SIGNATURE_SEPARATOR}sig123"),
+ model="gemini-3.5-flash",
+ custom_llm_provider="vertex_ai",
+ )
+
+ _, function_response_ids = _collect_function_call_ids(contents)
+ assert function_response_ids == [bare_id]
+
+
+@pytest.mark.parametrize("model", ["gemini-3.5-flash", "gemini-2.5-flash"])
+def test_tool_response_without_matching_tool_call_is_rejected(model):
+ """An unpairable tool result must raise, not ship a functionResponse with no matching call."""
+ from litellm.llms.vertex_ai.gemini.transformation import (
+ _gemini_convert_messages_with_history,
+ )
+
+ messages = [
+ {"role": "user", "content": "hi"},
+ {
+ "role": "assistant",
+ "content": "",
+ "tool_calls": [
+ {
+ "id": "call_50e7e0fe0989464a89f188eda443",
+ "type": "function",
+ "function": {
+ "name": "read",
+ "arguments": '{"filePath": "/tmp"}',
+ },
+ }
+ ],
+ },
+ {"role": "tool", "content": "ok"},
+ ]
+
+ with pytest.raises(Exception, match="Missing corresponding tool call"):
+ _gemini_convert_messages_with_history(
+ messages=messages,
+ model=model,
+ custom_llm_provider="vertex_ai",
+ )
def test_reasoning_effort_maps_to_thinking_level_gemini_3():
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py
index c063915e2e8..f6bd79c5d2d 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py
@@ -240,10 +240,10 @@ async def test_explicit_null_clears_upstream_resource_and_keeps_the_rest_of_the_
@pytest.mark.asyncio
-async def test_url_change_clears_stale_discovered_oauth_fields():
- """Re-pointing the server url at a potentially different upstream must clear the discovered or
- trust-on-first-use OAuth issuer and endpoints, so the new upstream re-discovers instead of
- anchoring on the previous upstream's issuer (RFC 8414 §3.3 against a stale anchor)."""
+async def test_url_change_clears_stale_oauth_fields():
+ """Re-pointing the server url at a potentially different upstream must clear the OAuth issuer and
+ endpoints, so the new upstream re-discovers instead of anchoring on the previous upstream's issuer
+ (RFC 8414 §3.3 against a stale anchor)."""
mock_prisma = _mock_prisma()
existing = MagicMock()
existing.auth_type = "oauth2"
@@ -350,11 +350,13 @@ async def test_repointing_pinned_issuer_clears_stale_endpoints_keeps_new_issuer(
@pytest.mark.asyncio
-async def test_establishing_issuer_first_time_preserves_discovered_fields():
- """Establishing an issuer for the first time (None -> X), which is exactly what the trust-on-first-use
- discovery write-back does, must NOT clear the endpoints or oauth2_flow it discovered in the same
- write. Only an issuer that was already pinned and is now changed or cleared invalidates its
- endpoints, so the discovery persist cannot wipe the fields it just resolved."""
+async def test_establishing_issuer_first_time_preserves_endpoints_set_in_the_same_write():
+ """Establishing an issuer for the first time (None -> X) must NOT clear endpoints or oauth2_flow
+ submitted in the same write. Only an issuer that was already pinned and is now changed or cleared
+ invalidates its endpoints, so an admin configuring an issuer and its endpoints together keeps
+ both. The write-back this once guarded (trust-on-first-use discovery stamping the issuer it had
+ just resolved) no longer exists; the db.py rule it relies on still governs admin writes, which is
+ what this now covers."""
mock_prisma = _mock_prisma()
existing = MagicMock()
existing.auth_type = "oauth2"
@@ -370,7 +372,7 @@ async def test_establishing_issuer_first_time_preserves_discovered_fields():
token_url="https://discovered-idp.example.com/token",
oauth2_flow="authorization_code",
)
- await update_mcp_server(mock_prisma, data, "mcp_oauth_discovery")
+ await update_mcp_server(mock_prisma, data, "some-admin@example.com")
data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
assert data_dict["issuer"] == "https://discovered-idp.example.com"
@@ -380,9 +382,9 @@ async def test_establishing_issuer_first_time_preserves_discovered_fields():
@pytest.mark.asyncio
-async def test_unchanged_url_does_not_clear_discovered_oauth_fields():
- """A partial update that resends the same url (or omits it) must not clear the discovered OAuth
- fields, so a routine save does not force needless re-discovery."""
+async def test_unchanged_url_does_not_clear_oauth_fields():
+ """A partial update that resends the same url (or omits it) must not clear the OAuth fields, so a
+ routine save does not force needless re-discovery."""
mock_prisma = _mock_prisma()
existing = MagicMock()
existing.auth_type = "oauth2"
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 5f7f2267fc7..8a8dea0ba28 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
@@ -2,6 +2,7 @@ import importlib
import asyncio
import json
import logging
+import time
import os
import sys
from datetime import datetime
@@ -35,6 +36,8 @@ from mcp.types import Tool as MCPTool
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
_deserialize_json_dict,
+ _flow_endpoints_missing,
+ _oauth_endpoints_unresolved,
_deserialize_json_list,
_normalize_mcp_server_cost_info,
_should_strip_caller_authorization,
@@ -1594,21 +1597,15 @@ class TestMCPServerManager:
token_url="https://idp.example.com/token",
scopes=["read"],
)
- with (
- patch.object(
- manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=issuer_resolved)
- ) as anchored,
- patch.object(manager, "_persist_discovered_oauth_endpoints", new=AsyncMock()) as mock_persist,
- ):
+ with patch.object(
+ manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=issuer_resolved)
+ ) as anchored:
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
anchored.assert_awaited_once_with("https://idp.example.com", "https://up.example.com/mcp")
assert built.authorization_url == "https://idp.example.com/authorize"
assert built.token_url == "https://idp.example.com/token"
assert built.token_url != "https://attacker.example.com/steal"
- # The issuer-anchored endpoints are never persisted into the endpoint columns, so a later
- # build cannot treat them as authoritative stored values.
- assert mock_persist.await_args.kwargs["is_issuer_anchored"] is True
@pytest.mark.asyncio
@pytest.mark.parametrize(
@@ -1624,8 +1621,8 @@ class TestMCPServerManager:
and PKCE verifier to the attacker (config-time RFC 9700 mix-up). The resource-driven scopes
are kept, because scope selection is resource-driven (MCP Scope Selection Strategy) and scope
inflation is bounded by the authorization server at consent (RFC 6749 §3.3), not by dropping
- scopes on an endpoint mismatch. Both the in-memory merge and the persisted metadata drop only
- the uncorroborated endpoints."""
+ scopes on an endpoint mismatch. The gateway persists nothing, so the in-memory merge is the
+ entire behavior."""
manager = MCPServerManager()
row = LiteLLM_MCPServerTable(
server_id="manual-auth-url-3",
@@ -1645,20 +1642,13 @@ class TestMCPServerManager:
registration_url="https://attacker.example.com/register",
scopes=["read", "admin"],
)
- with (
- patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)),
- patch.object(manager, "_persist_discovered_oauth_endpoints", new=AsyncMock()) as mock_persist,
- ):
+ with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)):
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
assert built.authorization_url == "https://idp.example.com/authorize"
assert built.token_url is None
assert built.registration_url is None
assert built.scopes == ["read", "admin"]
- persisted_metadata = mock_persist.await_args.kwargs["metadata"]
- assert persisted_metadata.token_url is None
- assert persisted_metadata.registration_url is None
- assert persisted_metadata.scopes == ["read", "admin"]
@pytest.mark.asyncio
async def test_build_from_table_skips_discovery_when_all_upstream_oauth_fields_present(self):
@@ -5586,388 +5576,300 @@ class TestMCPServerTimestamps:
assert server.token_exchange_endpoint == "https://idp.example.com/token"
@pytest.mark.asyncio
- async def test_build_mcp_server_from_table_persists_discovered_obo_token_url(self):
- """A DB-backed OBO server with no configured endpoint discovers token_url and must write it
- back to the row, so the next rebuild skips discovery instead of re-running it every time."""
+ async def test_discovery_never_writes_the_database(self):
+ """The #34985 regression, stated as the design invariant that fixes it: the gateway never
+ writes discovery results to the row. The OAuth columns and credentials.scopes carry admin
+ intent alone, so nothing the gateway learns can read back as an admin pin on a later build
+ (which is what anchored stamped servers fail-closed and 400ed /authorize). Discovery output
+ lives on the in-memory registry entry only, for oauth2 and OBO alike."""
manager = MCPServerManager()
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False):
- assert server_url == "https://example.com/mcp"
- assert allow_origin_fallback is False # OBO never guesses the origin
return MCPOAuthMetadata(
- scopes=None,
- authorization_url=None,
- token_url="https://discovered.example.com/token",
- registration_url=None,
- )
-
- manager._descovery_metadata = fake_discovery # type: ignore[attr-defined]
-
- record = LiteLLM_MCPServerTable(
- server_id="obo-persist-1",
- server_name="obo_persist",
- url="https://example.com/mcp",
- transport=MCPTransport.http,
- auth_type=MCPAuth.oauth2_token_exchange,
- credentials={"client_id": "cid", "client_secret": "csec", "audience": "aud"},
- )
-
- update_mock = AsyncMock()
- repo_instance = MagicMock()
- repo_instance.table.update = update_mock
- with (
- patch(
- "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository",
- return_value=repo_instance,
- ),
- patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
- ):
- server = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False)
-
- assert server.token_url == "https://discovered.example.com/token"
- update_mock.assert_awaited_once()
- assert update_mock.call_args.kwargs["where"] == {"server_id": "obo-persist-1"}
- assert update_mock.call_args.kwargs["data"] == {"token_url": "https://discovered.example.com/token"}
-
- @pytest.mark.asyncio
- async def test_persist_discovered_obo_token_url_skips_when_not_needed(self):
- """The write-back fires only for an OBO server that discovered a new endpoint: a row that
- already has token_url, a non-OBO auth_type, or a discovery that found nothing all no-op."""
- manager = MCPServerManager()
- update_mock = AsyncMock()
- repo_instance = MagicMock()
- repo_instance.table.update = update_mock
-
- with (
- patch(
- "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository",
- return_value=repo_instance,
- ),
- patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
- ):
- # already populated -> no write
- await manager._persist_discovered_obo_token_url(
- server_id="s",
- auth_type=MCPAuth.oauth2_token_exchange,
- existing_token_url="https://already.example.com/token",
- discovered_token_url="https://new.example.com/token",
- )
- # not an OBO server -> no write
- await manager._persist_discovered_obo_token_url(
- server_id="s",
- auth_type=MCPAuth.oauth2,
- existing_token_url=None,
- discovered_token_url="https://new.example.com/token",
- )
- # discovery found nothing -> no write
- await manager._persist_discovered_obo_token_url(
- server_id="s",
- auth_type=MCPAuth.oauth2_token_exchange,
- existing_token_url=None,
- discovered_token_url=None,
- )
-
- update_mock.assert_not_awaited()
-
- @pytest.mark.asyncio
- async def test_persist_discovered_obo_token_url_is_best_effort(self):
- """A write-back failure must not propagate; discovery just re-runs on the next build."""
- manager = MCPServerManager()
- update_mock = AsyncMock(side_effect=Exception("db unavailable"))
- repo_instance = MagicMock()
- repo_instance.table.update = update_mock
-
- with (
- patch(
- "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository",
- return_value=repo_instance,
- ),
- patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
- ):
- await manager._persist_discovered_obo_token_url(
- server_id="s",
- auth_type=MCPAuth.oauth2_token_exchange,
- existing_token_url=None,
- discovered_token_url="https://new.example.com/token",
- )
-
- update_mock.assert_awaited_once()
-
- @pytest.mark.asyncio
- async def test_build_mcp_server_from_table_persists_discovered_oauth_endpoints(self):
- """A DB-backed oauth2 server with no configured endpoints discovers them and must write
- authorization_url, token_url, and scopes back to the row; otherwise the resolved values
- live only in memory and one failed re-discovery serves the 400 "authorization url is not configured"
- from /authorize. registration_url must never be persisted because
- _dcr_bridge_relays_client_registration keys off that column."""
- manager = MCPServerManager()
-
- async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False):
- assert allow_origin_fallback is True
- return MCPOAuthMetadata(
- scopes=["mcp.read", "mcp.write"],
+ scopes=["mcp.read"],
authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/token",
registration_url="https://idp.example.com/register",
- )
-
- manager._descovery_metadata = fake_discovery # type: ignore[attr-defined]
-
- record = LiteLLM_MCPServerTable(
- server_id="oauth-persist-1",
- server_name="oauth_persist",
- url="https://example.com/mcp",
- transport=MCPTransport.http,
- auth_type=MCPAuth.oauth2,
- oauth2_flow="authorization_code",
- credentials={"client_id": "cid", "client_secret": "csec"},
- )
-
- update_mcp_server_mock = AsyncMock()
- with (
- patch(
- "litellm.proxy._experimental.mcp_server.db.update_mcp_server",
- new=update_mcp_server_mock,
- ),
- patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
- ):
- server = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False)
-
- assert server.authorization_url == "https://idp.example.com/authorize"
- update_mcp_server_mock.assert_awaited_once()
- persisted = update_mcp_server_mock.call_args.kwargs["data"]
- assert persisted.server_id == "oauth-persist-1"
- assert persisted.authorization_url == "https://idp.example.com/authorize"
- assert persisted.token_url == "https://idp.example.com/token"
- assert persisted.credentials == {"scopes": ["mcp.read", "mcp.write"]}
- assert "registration_url" not in persisted.fields_set()
- assert update_mcp_server_mock.call_args.kwargs["touched_by"] == "mcp_oauth_discovery"
-
- @pytest.mark.asyncio
- async def test_persist_discovered_oauth_endpoints_guards(self):
- """The write-back must no-op for non-discovery auth types, empty discovery, origin-fallback
- guesses (never harden an inferred authorization server into configuration), and rows whose
- fields are all already populated."""
- manager = MCPServerManager()
- advertised = MCPOAuthMetadata(
- scopes=["s1"],
- authorization_url="https://idp.example.com/authorize",
- token_url="https://idp.example.com/token",
- )
-
- update_mcp_server_mock = AsyncMock()
- with (
- patch(
- "litellm.proxy._experimental.mcp_server.db.update_mcp_server",
- new=update_mcp_server_mock,
- ),
- patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
- ):
- await manager._persist_discovered_oauth_endpoints(
- server_id="s",
- auth_type=MCPAuth.api_key,
- existing_issuer=None,
- existing_authorization_url=None,
- existing_token_url=None,
- existing_scopes=None,
- metadata=advertised,
- )
- await manager._persist_discovered_oauth_endpoints(
- server_id="s",
- auth_type=MCPAuth.oauth2,
- existing_issuer=None,
- existing_authorization_url=None,
- existing_token_url=None,
- existing_scopes=None,
- metadata=None,
- )
- await manager._persist_discovered_oauth_endpoints(
- server_id="s",
- auth_type=MCPAuth.oauth2,
- existing_issuer=None,
- existing_authorization_url=None,
- existing_token_url=None,
- existing_scopes=None,
- metadata=advertised.model_copy(update={"from_origin_fallback": True}),
- )
- await manager._persist_discovered_oauth_endpoints(
- server_id="s",
- auth_type=MCPAuth.oauth2,
- existing_issuer=None,
- existing_authorization_url="https://configured.example.com/authorize",
- existing_token_url="https://configured.example.com/token",
- existing_scopes=["configured"],
- metadata=advertised,
- )
-
- update_mcp_server_mock.assert_not_awaited()
-
- @pytest.mark.asyncio
- async def test_persist_discovered_oauth_endpoints_only_fills_empty_fields(self):
- """A row that already has token_url keeps it; only the missing authorization_url and
- scopes are written, so admin-typed values always win over discovery."""
- manager = MCPServerManager()
-
- update_mcp_server_mock = AsyncMock()
- with (
- patch(
- "litellm.proxy._experimental.mcp_server.db.update_mcp_server",
- new=update_mcp_server_mock,
- ),
- patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
- ):
- await manager._persist_discovered_oauth_endpoints(
- server_id="s",
- auth_type=MCPAuth.oauth2,
- existing_issuer=None,
- existing_authorization_url=None,
- existing_token_url="https://configured.example.com/token",
- existing_scopes=None,
- metadata=MCPOAuthMetadata(
- scopes=["s1"],
- authorization_url="https://idp.example.com/authorize",
- token_url="https://idp.example.com/token",
- ),
- )
-
- update_mcp_server_mock.assert_awaited_once()
- persisted = update_mcp_server_mock.call_args.kwargs["data"]
- assert persisted.authorization_url == "https://idp.example.com/authorize"
- assert persisted.credentials == {"scopes": ["s1"]}
- assert "token_url" not in persisted.fields_set()
-
- @pytest.mark.asyncio
- async def test_persist_discovered_oauth_endpoints_writes_discovered_issuer_trust_on_first_use(self):
- """A server with no configured issuer records the discovered issuer trust-on-first-use, so the
- next rebuild anchors discovery on it (RFC 8414 §3.3) instead of re-trusting the resource. When
- an issuer is already set (admin-typed or a prior discovery), it is never overwritten."""
- manager = MCPServerManager()
- metadata = MCPOAuthMetadata(
- authorization_url="https://idp.example.com/authorize",
- token_url="https://idp.example.com/token",
- discovered_issuer="https://idp.example.com",
- )
-
- update_mcp_server_mock = AsyncMock()
- with (
- patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock),
- patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
- ):
- await manager._persist_discovered_oauth_endpoints(
- server_id="s",
- auth_type=MCPAuth.oauth2,
- existing_issuer=None,
- existing_authorization_url=None,
- existing_token_url=None,
- existing_scopes=None,
- metadata=metadata,
- )
- await manager._persist_discovered_oauth_endpoints(
- server_id="s",
- auth_type=MCPAuth.oauth2,
- existing_issuer="https://admin-configured.example.com",
- existing_authorization_url="https://admin-configured.example.com/authorize",
- existing_token_url="https://admin-configured.example.com/token",
- existing_scopes=["cfg"],
- metadata=metadata,
- )
-
- assert update_mcp_server_mock.await_count == 1
- persisted = update_mcp_server_mock.call_args.kwargs["data"]
- assert persisted.issuer == "https://idp.example.com"
-
- @pytest.mark.asyncio
- async def test_persist_discovered_oauth_endpoints_does_not_persist_endpoints_for_issuer_anchored(self):
- """For an issuer-anchored server the endpoints are re-derived from the §3.3-validated issuer
- document every build, so they must NOT be written into the endpoint columns: persisting them
- would make the next build see populated endpoints and treat them as authoritative stored
- values, defeating the issuer-only invariant. Only the resource-driven scopes are persisted."""
- manager = MCPServerManager()
- metadata = MCPOAuthMetadata(
- authorization_url="https://idp.example.com/authorize",
- token_url="https://idp.example.com/token",
- scopes=["read"],
- )
-
- update_mcp_server_mock = AsyncMock()
- with (
- patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock),
- patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
- ):
- await manager._persist_discovered_oauth_endpoints(
- server_id="s",
- auth_type=MCPAuth.oauth2,
- existing_issuer="https://idp.example.com",
- existing_authorization_url=None,
- existing_token_url=None,
- existing_scopes=None,
- metadata=metadata,
- is_issuer_anchored=True,
- )
-
- update_mcp_server_mock.assert_awaited_once()
- persisted = update_mcp_server_mock.call_args.kwargs["data"]
- assert "authorization_url" not in persisted.fields_set()
- assert "token_url" not in persisted.fields_set()
- assert persisted.credentials == {"scopes": ["read"]}
-
- @pytest.mark.asyncio
- async def test_build_mcp_server_from_table_skips_persistence_for_temporary_servers(self):
- """The session endpoint builds temporary servers whose server_id has no DB row; with
- persist_discovered_endpoints=False neither the oauth2 nor the OBO write-back may fire."""
- manager = MCPServerManager()
-
- async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False):
- return MCPOAuthMetadata(
- scopes=["s1"],
- authorization_url="https://idp.example.com/authorize",
- token_url="https://idp.example.com/token",
+ discovered_issuer="https://idp.example.com",
)
manager._descovery_metadata = fake_discovery # type: ignore[attr-defined]
update_mcp_server_mock = AsyncMock()
- obo_update_mock = AsyncMock()
repo_instance = MagicMock()
- repo_instance.table.update = obo_update_mock
+ repo_instance.table.update = AsyncMock()
with (
- patch(
- "litellm.proxy._experimental.mcp_server.db.update_mcp_server",
- new=update_mcp_server_mock,
- ),
+ patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository",
return_value=repo_instance,
),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
):
- oauth2_record = LiteLLM_MCPServerTable(
- server_id="temp-oauth-1",
- server_name="temp_oauth",
- url="https://example.com/mcp",
+ for auth_type, flow in ((MCPAuth.oauth2, "authorization_code"), (MCPAuth.oauth2_token_exchange, None)):
+ record = LiteLLM_MCPServerTable(
+ server_id=f"no-write-{auth_type}",
+ server_name=f"no_write_{auth_type}",
+ url="https://example.com/mcp",
+ transport=MCPTransport.http,
+ auth_type=auth_type,
+ oauth2_flow=flow,
+ credentials={"client_id": "cid", "client_secret": "csec", "audience": "aud"},
+ )
+ built = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False)
+ assert built.token_url == "https://idp.example.com/token"
+
+ update_mcp_server_mock.assert_not_awaited()
+ repo_instance.table.update.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_declared_endpoints_survive_a_failed_discovery(self):
+ """The reporter's configuration: explicit authorization_url/token_url/registration_url,
+ issuer left empty. With the gateway never stamping the issuer column, the server never turns
+ anchored, so the declared endpoints resolve on every build, including one whose discovery
+ fails entirely; /authorize keeps redirecting instead of serving the 400."""
+ manager = MCPServerManager()
+ record = LiteLLM_MCPServerTable(
+ server_id="declared-1",
+ alias="declared",
+ url="https://up.example.com/mcp",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ oauth2_flow="authorization_code",
+ authorization_url="https://idp.example.com/authorize",
+ token_url="https://idp.example.com/token",
+ registration_url="https://idp.example.com/register",
+ created_at=datetime.now(),
+ updated_at=datetime.now(),
+ )
+
+ with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)):
+ built = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False)
+
+ assert built.issuer_is_anchored is False
+ assert built.authorization_url == "https://idp.example.com/authorize"
+ assert built.token_url == "https://idp.example.com/token"
+ assert built.registration_url == "https://idp.example.com/register"
+
+ def test_flow_endpoints_missing_arms(self):
+ """The reload fast-path exemption's completeness rule. Interactive needs authorize+token,
+ client_credentials and OBO need token only, an OBO server with a configured exchange
+ endpoint never discovers and must not be sent into a rebuild loop, and non-OAuth auth types
+ are never unresolved."""
+ assert _flow_endpoints_missing(MCPAuth.oauth2, "authorization_code", "https://idp/auth", None) is True
+ assert _flow_endpoints_missing(MCPAuth.oauth2, "authorization_code", None, "https://idp/token") is True
+ assert (
+ _flow_endpoints_missing(MCPAuth.oauth2, "authorization_code", "https://idp/auth", "https://idp/token")
+ is False
+ )
+ assert _flow_endpoints_missing(MCPAuth.oauth2, "client_credentials", None, "https://idp/token") is False
+ assert _flow_endpoints_missing(MCPAuth.oauth2, "client_credentials", None, None) is True
+ assert _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, None) is True
+ assert _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, "https://idp/token") is False
+ assert (
+ _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, None, "https://idp/exchange") is False
+ )
+ assert _flow_endpoints_missing(MCPAuth.api_key, None, None, None) is False
+
+ def test_unresolved_check_uses_the_flow_judge_not_the_raw_column(self):
+ """A legacy row the startup backfill deliberately left unstamped (token_url plus client
+ credentials, no authorization_url: the ambiguous M2M shape) serves client_credentials at
+ request time via effective_oauth2_flow. The reload check must reach the same verdict, or the
+ row is classified as interactive-missing-endpoints and re-runs discovery on every reload
+ forever. A null-flow row without the M2M shape stays interactive and genuinely unresolved."""
+ m2m_shaped = MCPServer(
+ server_id="null-flow-m2m",
+ name="null_flow_m2m",
+ server_name="null_flow_m2m",
+ url="https://up.example.com/mcp",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ oauth2_flow=None,
+ token_url="https://idp.example.com/token",
+ client_id="cid",
+ client_secret="csec",
+ )
+ assert _oauth_endpoints_unresolved(m2m_shaped) is False
+
+ interactive_unresolved = m2m_shaped.model_copy(update={"client_id": None, "client_secret": None})
+ assert _oauth_endpoints_unresolved(interactive_unresolved) is True
+
+ def test_dcr_bridge_relay_arm_needs_its_registration_endpoint(self):
+ """A dcr_bridge server with no admin-configured client can only register callers through the
+ upstream registration endpoint, so a partial discovery that resolved authorize and token but
+ not registration_endpoint leaves it silently degraded to the short-circuit arm. That counts as
+ unresolved so it keeps retrying. A bridge with a configured client_id uses the short-circuit
+ arm by design and is unaffected."""
+ relay_arm = MCPServer(
+ server_id="bridge-partial",
+ name="bridge_partial",
+ url="https://up.example.com/mcp",
+ transport=MCPTransport.http,
+ # dcr_bridge is only valid on the client-forwarded modes (see MCPServer.is_dcr_bridge)
+ auth_type=MCPAuth.oauth_delegate,
+ dcr_bridge=True,
+ authorization_url="https://idp.example.com/authorize",
+ token_url="https://idp.example.com/token",
+ registration_url=None,
+ )
+ assert _oauth_endpoints_unresolved(relay_arm) is True
+ assert _oauth_endpoints_unresolved(relay_arm.model_copy(update={"registration_url": "https://idp/reg"})) is False
+ assert _oauth_endpoints_unresolved(relay_arm.model_copy(update={"client_id": "admin-client"})) is False
+
+ def test_entra_obo_without_scopes_is_unresolved(self):
+ """entra_obo token exchange fails closed without a scope, and scopes can come from resource
+ discovery, so an entra_obo server that resolved its token endpoint but no scopes is still
+ unresolved for its flow. The default rfc8693 profile has no such requirement."""
+ entra = MCPServer(
+ server_id="entra-noscope",
+ name="entra_noscope",
+ url="https://up.example.com/mcp",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2_token_exchange,
+ token_exchange_profile="entra_obo",
+ token_url="https://idp.example.com/token",
+ scopes=None,
+ )
+ assert _oauth_endpoints_unresolved(entra) is True
+ assert _oauth_endpoints_unresolved(entra.model_copy(update={"scopes": ["api://app/.default"]})) is False
+ assert _oauth_endpoints_unresolved(entra.model_copy(update={"token_exchange_profile": "rfc8693"})) is False
+
+ def test_oauth_discovery_retry_backs_off_per_server(self):
+ """Without a cooldown the fast-path exemption re-runs the full discovery chain, and re-emits
+ the unresolved warning, on every reload forever for a server that can never resolve. Delay
+ doubles per consecutive failure up to the cap, a success clears the state so the next failure
+ starts from the base delay again, and the cooldown is per server."""
+ manager = MCPServerManager()
+
+ def unresolved(server_id):
+ return MCPServer(
+ server_id=server_id,
+ name=server_id,
+ url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
- credentials={"client_id": "cid", "client_secret": "csec"},
- )
- obo_record = LiteLLM_MCPServerTable(
- server_id="temp-obo-1",
- server_name="temp_obo",
- url="https://example.com/mcp",
- transport=MCPTransport.http,
- auth_type=MCPAuth.oauth2_token_exchange,
- credentials={"client_id": "cid", "client_secret": "csec"},
- )
- built_oauth2 = await manager.build_mcp_server_from_table(
- oauth2_record, credentials_are_encrypted=False, persist_discovered_endpoints=False
- )
- await manager.build_mcp_server_from_table(
- obo_record, credentials_are_encrypted=False, persist_discovered_endpoints=False
)
- assert built_oauth2.authorization_url == "https://idp.example.com/authorize"
- update_mcp_server_mock.assert_not_awaited()
- obo_update_mock.assert_not_awaited()
+ assert manager._oauth_discovery_retry_due("a") is True
+
+ manager._record_oauth_discovery_outcome(unresolved("a"))
+ assert manager._oauth_discovery_retry_due("a") is False
+ assert manager._oauth_discovery_retry_due("b") is True, "cooldown must be per server"
+
+ failures_before, _ = manager._oauth_discovery_retry_state["a"]
+ manager._record_oauth_discovery_outcome(unresolved("a"))
+ failures_after, _ = manager._oauth_discovery_retry_state["a"]
+ assert failures_after == failures_before + 1
+
+ # An elapsed cooldown lets the retry through, and the delay grows with the failure count
+ manager._oauth_discovery_retry_state["a"] = (1, time.monotonic() - 31.0)
+ assert manager._oauth_discovery_retry_due("a") is True
+ manager._oauth_discovery_retry_state["a"] = (5, time.monotonic() - 31.0)
+ assert manager._oauth_discovery_retry_due("a") is False
+
+ resolved = unresolved("a").model_copy(
+ update={
+ "authorization_url": "https://idp.example.com/authorize",
+ "token_url": "https://idp.example.com/token",
+ }
+ )
+ manager._record_oauth_discovery_outcome(resolved)
+ assert "a" not in manager._oauth_discovery_retry_state
+ assert manager._oauth_discovery_retry_due("a") is True
+
+ @pytest.mark.asyncio
+ async def test_reload_fast_path_retries_unresolved_oauth_servers(self):
+ """A server whose discovery failed must not be pinned broken by the updated_at fast path:
+ the next reload rebuilds it, retrying discovery on the normal cadence instead of waiting for
+ an unrelated config write. A resolved server with an unchanged row still takes the fast path,
+ so the exemption costs nothing in the steady state."""
+ manager = MCPServerManager()
+ stamp = datetime.now()
+ row = LiteLLM_MCPServerTable(
+ server_id="retry-1",
+ server_name="retry_server",
+ url="https://up.example.com/mcp",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ oauth2_flow="authorization_code",
+ created_at=stamp,
+ updated_at=stamp,
+ )
+
+ def entry(authorization_url, token_url):
+ return MCPServer(
+ server_id="retry-1",
+ name="retry_server",
+ server_name="retry_server",
+ url="https://up.example.com/mcp",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ oauth2_flow="authorization_code",
+ authorization_url=authorization_url,
+ token_url=token_url,
+ updated_at=stamp,
+ )
+
+ raw_row = MagicMock()
+ raw_row.model_dump.return_value = row.model_dump()
+ repo_instance = MagicMock()
+ repo_instance.table.find_many = AsyncMock(return_value=[raw_row])
+
+ async def run_reload(previous_entry):
+ manager.registry = {"retry-1": previous_entry}
+ build_mock = AsyncMock(return_value=previous_entry)
+ with (
+ patch(
+ "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository",
+ return_value=repo_instance,
+ ),
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
+ return_value=MagicMock(),
+ ),
+ patch.object(manager, "build_mcp_server_from_table", new=build_mock),
+ ):
+ await manager.reload_servers_from_database()
+ return build_mock
+
+ unresolved_build = await run_reload(entry(None, None))
+ unresolved_build.assert_awaited_once()
+
+ resolved_build = await run_reload(entry("https://idp.example.com/authorize", "https://idp.example.com/token"))
+ resolved_build.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_anchored_issuer_discarding_stored_endpoints_warns(self, caplog):
+ """An anchored server ignoring stored endpoint columns must say so: that state is exactly
+ what a row stamped by an earlier release looks like after upgrade, and the warning names the
+ remedy (clear the Issuer field) instead of leaving the 400 undiagnosable."""
+ manager = MCPServerManager()
+ record = LiteLLM_MCPServerTable(
+ server_id="stamped-1",
+ alias="stamped_row",
+ url="https://up.example.com/mcp",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ oauth2_flow="authorization_code",
+ issuer="https://idp.example.com",
+ authorization_url="https://idp.example.com/authorize",
+ token_url="https://idp.example.com/token",
+ created_at=datetime.now(),
+ updated_at=datetime.now(),
+ )
+
+ with (
+ patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=None)),
+ caplog.at_level(logging.WARNING, logger="LiteLLM"),
+ ):
+ built = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False)
+
+ assert built.issuer_is_anchored is True
+ assert built.authorization_url is None
+ assert "stamped_row" in caplog.text
+ assert "authorization_url, token_url" in caplog.text
+ assert "clear the Issuer" in caplog.text
@pytest.mark.asyncio
async def test_update_server_carries_forward_last_known_good_oauth_endpoints(self):
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py
new file mode 100644
index 00000000000..b6c946b95fa
--- /dev/null
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py
@@ -0,0 +1,129 @@
+"""Tests for the one-time heal of issuer values a released version's discovery write-back stamped."""
+
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, MagicMock
+
+import pytest
+
+from litellm.proxy._experimental.mcp_server.oauth_issuer_stamp_backfill import (
+ backfill_discovery_stamped_issuers,
+)
+
+
+def _row(**overrides):
+ fields = {
+ "server_id": "srv-1",
+ "alias": "srv_one",
+ "server_name": "srv_one",
+ "auth_type": "oauth2",
+ "issuer": "https://idp.example.com",
+ "authorization_url": "https://idp.example.com/authorize",
+ "token_url": "https://idp.example.com/token",
+ "registration_url": None,
+ "updated_by": "mcp_oauth_discovery",
+ }
+ fields.update(overrides)
+ return SimpleNamespace(**fields)
+
+
+def _prisma(rows):
+ prisma_client = MagicMock()
+ prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=rows)
+ prisma_client.db.litellm_mcpservertable.update = AsyncMock()
+ return prisma_client
+
+
+@pytest.mark.asyncio
+async def test_clears_the_stamp_and_records_its_own_actor():
+ """The GH #34985 row: discovery wrote the issuer, so the server reads as issuer-anchored and its
+ configured endpoints are ignored. Clearing the stamp makes them apply again. The heal records its
+ own actor, which is also what makes it idempotent: the row no longer matches the discovery-actor
+ filter, so it is never reconsidered on a later boot."""
+ prisma_client = _prisma([_row()])
+
+ assert await backfill_discovery_stamped_issuers(prisma_client) == 1
+
+ call = prisma_client.db.litellm_mcpservertable.update.call_args
+ assert call.kwargs["where"] == {"server_id": "srv-1"}
+ assert call.kwargs["data"]["issuer"] is None
+ assert call.kwargs["data"]["updated_by"] == "mcp_oauth_issuer_stamp_backfill"
+
+ where = prisma_client.db.litellm_mcpservertable.find_many.call_args.kwargs["where"]
+ assert where["updated_by"] == "mcp_oauth_discovery"
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ "overrides, reason",
+ [
+ ({"updated_by": "some-admin@example.com"}, "an admin was the last writer, so the pin is theirs"),
+ ({"issuer": None}, "nothing to heal"),
+ ({"issuer": " "}, "blank issuer is not a pin"),
+ (
+ {"authorization_url": None, "token_url": None, "registration_url": None},
+ "issuer set with no configured endpoints is the canonical shape of a deliberate pin, and "
+ "there is nothing configured for anchoring to discard anyway",
+ ),
+ (
+ {"authorization_url": "https://other-idp.example.com/authorize", "token_url": None},
+ "endpoints addressing a different authority than the issuer are an intent a clear would "
+ "discard, so the row is warned about rather than healed",
+ ),
+ (
+ {"issuer": "https://pinned.example.com"},
+ "same shape from the other side: a pinned issuer whose origin differs from the configured "
+ "endpoints cannot have been derived from them by discovery",
+ ),
+ ],
+)
+async def test_leaves_rows_alone_that_do_not_carry_the_defect_signature(overrides, reason):
+ """updated_by records only the most recent writer and no audit trail says which field it touched,
+ so the heal is deliberately narrow: it fires only on the full signature of the defect. Every
+ exclusion here protects a row whose issuer may be a deliberate admin pin."""
+ prisma_client = _prisma([_row(**overrides)])
+
+ assert await backfill_discovery_stamped_issuers(prisma_client) == 0, reason
+ prisma_client.db.litellm_mcpservertable.update.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_heals_across_url_forms_that_denote_the_same_origin():
+ """Origin comparison runs through the shared canonicalizer, so a default port or host casing
+ difference between the stamped issuer and the endpoints an admin typed does not make a #34985 row
+ look like a deliberate pin at a different authority."""
+ prisma_client = _prisma(
+ [
+ _row(
+ issuer="https://IDP.example.com:443",
+ authorization_url="https://idp.example.com/authorize",
+ token_url="https://idp.example.com/token",
+ )
+ ]
+ )
+
+ assert await backfill_discovery_stamped_issuers(prisma_client) == 1
+
+
+@pytest.mark.asyncio
+async def test_query_is_scoped_to_auth_types_where_an_issuer_anchors():
+ """Only the discovery auth types read an issuer as a trust anchor; clearing it elsewhere would be
+ an unrelated mutation."""
+ prisma_client = _prisma([])
+
+ await backfill_discovery_stamped_issuers(prisma_client)
+
+ where = prisma_client.db.litellm_mcpservertable.find_many.call_args.kwargs["where"]
+ assert set(where["auth_type"]["in"]) == {"oauth2", "true_passthrough", "oauth_delegate"}
+
+
+@pytest.mark.asyncio
+async def test_a_failed_row_does_not_abort_the_rest():
+ """Per-row best effort: one write failure must not leave later rows unhealed, and the next boot
+ retries the failed one since its updated_by is unchanged."""
+ prisma_client = _prisma([_row(server_id="bad"), _row(server_id="good")])
+ prisma_client.db.litellm_mcpservertable.update = AsyncMock(
+ side_effect=[Exception("write failed"), MagicMock()]
+ )
+
+ assert await backfill_discovery_stamped_issuers(prisma_client) == 1
+ assert prisma_client.db.litellm_mcpservertable.update.await_count == 2
diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py
index 9f24c662581..1610d76efb7 100644
--- a/tests/test_litellm/proxy/auth/test_auth_utils.py
+++ b/tests/test_litellm/proxy/auth/test_auth_utils.py
@@ -569,6 +569,103 @@ def test_get_model_from_request_resolves_video_id_model_with_router():
)
+_BATCH_DEPLOYMENT_ID = "8d0eaa7e6c6f54a425dfd0062cb6b0dc"
+
+
+def _managed_batch_router():
+ from litellm.router import Router
+
+ return Router(
+ model_list=[
+ {
+ "model_name": "bedrock-batch-model",
+ "litellm_params": {
+ "model": "bedrock/global.anthropic.claude-haiku-4-5-20251001-v1:0",
+ },
+ "model_info": {"id": _BATCH_DEPLOYMENT_ID},
+ },
+ {
+ "model_name": "some-other-model",
+ "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"},
+ "model_info": {"id": "a-different-deployment-id"},
+ },
+ ]
+ )
+
+
+def _encode_managed_id(decoded: str) -> str:
+ return base64.urlsafe_b64encode(decoded.encode()).decode().rstrip("=")
+
+
+_MANAGED_BATCH_ID = _encode_managed_id(
+ f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123"
+)
+_MANAGED_BATCH_OUTPUT_FILE_ID = _encode_managed_id(
+ f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123;"
+ "llm_output_file_id:provider-file-456"
+)
+
+
+@pytest.mark.parametrize(
+ "route, request_data",
+ [
+ ("/v1/batches/{batch_id}", {"batch_id": _MANAGED_BATCH_ID}),
+ ("/v1/batches/{batch_id}/cancel", {"batch_id": _MANAGED_BATCH_ID}),
+ ("/v1/files/{file_id}", {"file_id": _MANAGED_BATCH_OUTPUT_FILE_ID}),
+ ("/v1/files/{file_id}/content", {"file_id": _MANAGED_BATCH_OUTPUT_FILE_ID}),
+ ],
+)
+def test_get_model_from_request_resolves_batch_id_deployment_to_model_name(route, request_data):
+ """Regression for #32580: managed batch retrieve/cancel and managed batch output
+ file reads encode the deployment model_id into the resource id. The auth layer must
+ resolve that id back to the public model group name so model-access checks compare
+ against the model group, not the raw deployment id."""
+ assert (
+ get_model_from_request(
+ request_data=request_data,
+ route=route,
+ llm_router=_managed_batch_router(),
+ )
+ == "bedrock-batch-model"
+ )
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ "route, request_data",
+ [
+ ("/v1/batches/{batch_id}", {"batch_id": _MANAGED_BATCH_ID}),
+ ("/v1/batches/{batch_id}/cancel", {"batch_id": _MANAGED_BATCH_ID}),
+ ("/v1/files/{file_id}/content", {"file_id": _MANAGED_BATCH_OUTPUT_FILE_ID}),
+ ],
+)
+async def test_managed_batch_routes_pass_team_model_access_check(route, request_data):
+ """End-to-end regression for #32580: a team scoped to the batch model group got
+ ``team_model_access_denied`` on retrieve/cancel because the deployment id, not the
+ model group, was authorized. Fails pre-fix with the deployment id in the message."""
+ from litellm.proxy._types import LiteLLM_TeamTable
+ from litellm.proxy.auth.auth_checks import can_team_access_model
+
+ llm_router = _managed_batch_router()
+ model = get_model_from_request(request_data=request_data, route=route, llm_router=llm_router)
+
+ assert (
+ await can_team_access_model(
+ model=model,
+ team_object=LiteLLM_TeamTable(team_id="team-batch", models=["bedrock-batch-model"]),
+ llm_router=llm_router,
+ )
+ is True
+ )
+
+ with pytest.raises(Exception, match="team not allowed to access model"):
+ await can_team_access_model(
+ model=model,
+ team_object=LiteLLM_TeamTable(team_id="team-other", models=["some-other-model"]),
+ llm_router=llm_router,
+ )
+
+
def test_get_model_from_request_resolves_character_id_model_with_router():
from litellm.types.videos.utils import encode_character_id_with_provider
diff --git a/tests/test_litellm/proxy/client/cli/test_auth_commands.py b/tests/test_litellm/proxy/client/cli/test_auth_commands.py
index 2fbc9c5c82f..f0aa49ff123 100644
--- a/tests/test_litellm/proxy/client/cli/test_auth_commands.py
+++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py
@@ -1,5 +1,6 @@
import json
import os
+import stat
import sys
import time
from pathlib import Path
@@ -12,6 +13,7 @@ import pytest
from click.testing import CliRunner
from litellm.constants import CLI_JWT_EXPIRATION_HOURS
+from litellm.proxy.client.cli import cli
from litellm.proxy.client.cli.commands.auth import (
clear_token,
get_stored_api_key,
@@ -201,31 +203,22 @@ class TestTokenUtilities:
mock_mkdir.assert_called_once_with(exist_ok=True)
- def test_save_token(self):
+ def test_save_token(self, tmp_path):
"""Test saving token data to file"""
token_data = {
"key": "test-key",
"user_id": "test-user",
"timestamp": 1234567890,
}
+ token_file = tmp_path / "token.json"
- with (
- patch("builtins.open", mock_open()) as mock_file,
- patch("litellm.proxy.client.cli.commands.auth.get_token_file_path") as mock_path,
- patch("os.chmod") as mock_chmod,
- ):
- mock_path.return_value = "/test/path/token.json"
+ with patch("litellm.proxy.client.cli.commands.auth.get_token_file_path") as mock_path:
+ mock_path.return_value = str(token_file)
save_token(token_data)
- mock_file.assert_called_once_with("/test/path/token.json", "w")
- mock_file().write.assert_called()
- mock_chmod.assert_called_once_with("/test/path/token.json", 0o600)
-
- # Verify JSON content was written correctly
- written_content = "".join(call[0][0] for call in mock_file().write.call_args_list)
- parsed_content = json.loads(written_content)
- assert parsed_content == token_data
+ assert json.loads(token_file.read_text()) == token_data
+ assert stat.S_IMODE(token_file.stat().st_mode) == 0o600
def test_load_token_success(self):
"""Test loading token data from file successfully"""
@@ -808,7 +801,8 @@ class TestPrintTokenCommand:
since there is no explicit target to check it against. `--base-url`/
`LITELLM_PROXY_URL` only enforces the match when a caller explicitly
passes it (tracked via ctx.obj["base_url_explicit"], set by the `cli`
- group from click's ParameterSource).
+ group from click's ParameterSource); a base_url saved via
+ `lite config set` counts as explicit too.
"""
def setup_method(self):
@@ -928,3 +922,110 @@ class TestPrintTokenCommand:
assert "sk-stale-key" not in result.output
assert "lite login" in result.output
mock_post.assert_not_called()
+
+
+def _write_home_json(home: Path, filename: str, payload: dict[str, object]) -> None:
+ litellm_dir = home / ".litellm"
+ litellm_dir.mkdir(exist_ok=True)
+ (litellm_dir / filename).write_text(json.dumps(payload))
+
+
+class TestPrintTokenWithConfigFile:
+ """A config-file base_url is a drop-in replacement for exporting
+ LITELLM_PROXY_URL, so print-token must treat it as an explicit server
+ choice: a token minted for a different proxy is never handed out."""
+
+ @pytest.fixture
+ def isolated_home(self, monkeypatch, tmp_path):
+ monkeypatch.setenv("HOME", str(tmp_path))
+ monkeypatch.setenv("USERPROFILE", str(tmp_path))
+ monkeypatch.delenv("LITELLM_PROXY_URL", raising=False)
+ monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False)
+ return tmp_path
+
+ def test_config_base_url_mismatch_fails_closed(self, isolated_home):
+ _write_home_json(
+ isolated_home,
+ "token.json",
+ {"base_url": "https://server-a.example.com", "key": "sk-issued-for-a", "timestamp": time.time()},
+ )
+ _write_home_json(isolated_home, "config.json", {"base_url": "https://server-b.example.com"})
+
+ result = CliRunner().invoke(cli, ["auth", "print-token"])
+
+ assert result.exit_code == 1
+ assert "sk-issued-for-a" not in result.output
+ assert "Not authenticated for this server" in result.output
+
+ def test_config_base_url_match_prints_token(self, isolated_home):
+ _write_home_json(
+ isolated_home,
+ "token.json",
+ {"base_url": "https://server-a.example.com", "key": "sk-issued-for-a", "timestamp": time.time()},
+ )
+ _write_home_json(isolated_home, "config.json", {"base_url": "https://server-a.example.com"})
+
+ result = CliRunner().invoke(cli, ["auth", "print-token"])
+
+ assert result.exit_code == 0
+ assert result.stdout.strip() == "sk-issued-for-a"
+
+ def test_empty_config_base_url_treated_as_unset(self, isolated_home):
+ """A hand-edited config.json with base_url "" must behave like no config at all:
+ base_url falls back to the default AND explicitness stays False."""
+ _write_home_json(
+ isolated_home,
+ "token.json",
+ {"base_url": "https://server-a.example.com", "key": "sk-issued-for-a", "timestamp": time.time()},
+ )
+ _write_home_json(isolated_home, "config.json", {"base_url": ""})
+
+ result = CliRunner().invoke(cli, ["auth", "print-token"])
+
+ assert result.exit_code == 0
+ assert result.stdout.strip() == "sk-issued-for-a"
+
+ def test_bare_invocation_without_config_file_unchanged(self, isolated_home):
+ """No config file means base_url_explicit stays False, so the stored
+ token's own server is trusted (pre-config behavior must not regress)."""
+ _write_home_json(
+ isolated_home,
+ "token.json",
+ {"base_url": "https://server-a.example.com", "key": "sk-issued-for-a", "timestamp": time.time()},
+ )
+
+ result = CliRunner().invoke(cli, ["auth", "print-token"])
+
+ assert result.exit_code == 0
+ assert result.stdout.strip() == "sk-issued-for-a"
+
+
+class TestSaveTokenPrivateWrite:
+ """token.json holds the real API key: it must never be world-readable at any
+ instant, and a failed write must not destroy the previously stored token."""
+
+ @pytest.fixture
+ def isolated_home(self, monkeypatch, tmp_path):
+ monkeypatch.setenv("HOME", str(tmp_path))
+ monkeypatch.setenv("USERPROFILE", str(tmp_path))
+ monkeypatch.delenv("LITELLM_PROXY_URL", raising=False)
+ monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False)
+ return tmp_path
+
+ def test_save_token_owner_only_permissions_and_no_temp_leftovers(self, isolated_home):
+ save_token({"key": "sk-secret", "user_id": "u-1", "timestamp": 1234567890})
+
+ token_file = isolated_home / ".litellm" / "token.json"
+ assert json.loads(token_file.read_text()) == {"key": "sk-secret", "user_id": "u-1", "timestamp": 1234567890}
+ assert stat.S_IMODE(token_file.stat().st_mode) == 0o600
+ assert list(token_file.parent.glob(".tmp-*")) == []
+
+ def test_save_token_failure_mid_write_preserves_existing_token(self, isolated_home):
+ _write_home_json(isolated_home, "token.json", {"key": "sk-original", "timestamp": 1234567890})
+ token_file = isolated_home / ".litellm" / "token.json"
+
+ with pytest.raises(TypeError):
+ save_token({"key": object()})
+
+ assert json.loads(token_file.read_text()) == {"key": "sk-original", "timestamp": 1234567890}
+ assert list(token_file.parent.glob(".tmp-*")) == []
diff --git a/tests/test_litellm/proxy/client/cli/test_config_commands.py b/tests/test_litellm/proxy/client/cli/test_config_commands.py
new file mode 100644
index 00000000000..698d6188768
--- /dev/null
+++ b/tests/test_litellm/proxy/client/cli/test_config_commands.py
@@ -0,0 +1,284 @@
+import json
+import os
+import stat
+import sys
+from pathlib import Path
+
+import pytest
+from click.testing import CliRunner
+
+sys.path.insert(0, os.path.abspath("../../.."))
+
+
+from litellm.proxy.client.cli import cli
+from litellm.proxy.client.cli.commands.config import (
+ get_config_file_path,
+ get_config_value,
+ load_config,
+ save_config,
+)
+from litellm.proxy.client.cli.commands.private_json import write_private_json
+
+
+@pytest.fixture
+def cli_runner():
+ return CliRunner()
+
+
+@pytest.fixture
+def isolated_home(monkeypatch, tmp_path):
+ """Point HOME at tmp_path so tests never touch the developer's real ~/.litellm."""
+ monkeypatch.setenv("HOME", str(tmp_path))
+ monkeypatch.setenv("USERPROFILE", str(tmp_path))
+ monkeypatch.delenv("LITELLM_PROXY_URL", raising=False)
+ monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False)
+ return tmp_path
+
+
+def _config_path(home: Path) -> Path:
+ return home / ".litellm" / "config.json"
+
+
+def _raise_home_unresolvable() -> str:
+ raise RuntimeError("Could not determine home directory.")
+
+
+class TestConfigSet:
+ @pytest.mark.parametrize(
+ "value",
+ ["https://your-proxy.example.com", "http://your-proxy.example.com:8080"],
+ )
+ def test_set_stores_value_with_owner_only_permissions(self, cli_runner, isolated_home, value):
+ result = cli_runner.invoke(cli, ["config", "set", "base_url", value])
+
+ assert result.exit_code == 0
+ config_file = _config_path(isolated_home)
+ assert json.loads(config_file.read_text()) == {"base_url": value}
+ assert stat.S_IMODE(config_file.stat().st_mode) == 0o600
+ assert str(config_file) in result.output
+
+ def test_set_strips_trailing_slash(self, cli_runner, isolated_home):
+ """Downstream commands join paths onto base_url; a stored trailing
+ slash would produce double slashes in every request URL."""
+ result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com/"])
+
+ assert result.exit_code == 0
+ assert json.loads(_config_path(isolated_home).read_text()) == {"base_url": "https://your-proxy.example.com"}
+
+ def test_set_unknown_key_rejected_and_names_allowed_keys(self, cli_runner, isolated_home):
+ result = cli_runner.invoke(cli, ["config", "set", "api_key", "sk-secret"])
+
+ assert result.exit_code != 0
+ assert "base_url" in result.output
+ assert not _config_path(isolated_home).exists()
+
+ @pytest.mark.parametrize("value", ["your-proxy.example.com", "ftp://your-proxy.example.com"])
+ def test_set_base_url_without_http_scheme_rejected(self, cli_runner, isolated_home, value):
+ result = cli_runner.invoke(cli, ["config", "set", "base_url", value])
+
+ assert result.exit_code != 0
+ assert "http" in result.output
+ assert not _config_path(isolated_home).exists()
+
+ @pytest.mark.parametrize("value", ["https://", "http://", "https:///some-path"])
+ def test_set_base_url_without_host_rejected(self, cli_runner, isolated_home, value):
+ """rstrip("/") would otherwise persist a bare "https:" that breaks every later request."""
+ result = cli_runner.invoke(cli, ["config", "set", "base_url", value])
+
+ assert result.exit_code != 0
+ assert not _config_path(isolated_home).exists()
+
+ @pytest.mark.parametrize(
+ "value",
+ [
+ "https://proxy.example.com?env=prod",
+ "https://proxy.example.com#prod",
+ "https://proxy.example.com/?",
+ "https://proxy.example.com/#",
+ ],
+ )
+ def test_set_base_url_with_query_or_fragment_rejected(self, cli_runner, isolated_home, value):
+ """Downstream commands join paths onto base_url; a stored query string or
+ fragment would silently corrupt every request URL built from it. Bare
+ trailing '?' / '#' parse as EMPTY query/fragment yet still break every
+ joined path, so rejection must key off the raw characters."""
+ result = cli_runner.invoke(cli, ["config", "set", "base_url", value])
+
+ assert result.exit_code != 0
+ assert "query" in result.output or "fragment" in result.output
+ assert not _config_path(isolated_home).exists()
+
+ def test_set_base_url_with_path_prefix_accepted(self, cli_runner, isolated_home):
+ """Proxies are commonly served under a path prefix; the query/fragment
+ rejection must not over-reach into legitimate paths."""
+ result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://proxy.example.com/litellm"])
+
+ assert result.exit_code == 0
+ assert json.loads(_config_path(isolated_home).read_text()) == {"base_url": "https://proxy.example.com/litellm"}
+
+ def test_set_leaves_no_temp_files_behind(self, cli_runner, isolated_home):
+ """The atomic write goes through a .tmp-* sibling; it must be renamed away,
+ never abandoned next to the config."""
+ result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com"])
+
+ assert result.exit_code == 0
+ config_file = _config_path(isolated_home)
+ assert stat.S_IMODE(config_file.stat().st_mode) == 0o600
+ assert list(config_file.parent.glob(".tmp-*")) == []
+
+
+class TestConfigGet:
+ def test_get_prints_only_the_value(self, cli_runner, isolated_home):
+ """stdout must be exactly the value so scripts can do URL=$(lite config get base_url)."""
+ set_result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com"])
+ assert set_result.exit_code == 0
+
+ result = cli_runner.invoke(cli, ["config", "get", "base_url"])
+
+ assert result.exit_code == 0
+ assert result.stdout.strip() == "https://your-proxy.example.com"
+
+ def test_get_unset_key_exits_one_with_stderr_message(self, cli_runner, isolated_home):
+ result = cli_runner.invoke(cli, ["config", "get", "base_url"])
+
+ assert result.exit_code == 1
+ assert result.stdout.strip() == ""
+ assert result.stderr != ""
+
+ def test_get_without_key_lists_entries(self, cli_runner, isolated_home):
+ set_result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com"])
+ assert set_result.exit_code == 0
+
+ result = cli_runner.invoke(cli, ["config", "get"])
+
+ assert result.exit_code == 0
+ assert "base_url = https://your-proxy.example.com" in result.output
+
+ def test_get_without_key_when_nothing_set(self, cli_runner, isolated_home):
+ result = cli_runner.invoke(cli, ["config", "get"])
+
+ assert result.exit_code == 0
+ assert "no config" in result.output.lower()
+
+
+class TestConfigUnset:
+ def test_unset_removes_key_from_file(self, cli_runner, isolated_home):
+ set_result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com"])
+ assert set_result.exit_code == 0
+
+ result = cli_runner.invoke(cli, ["config", "unset", "base_url"])
+
+ assert result.exit_code == 0
+ assert "base_url" not in load_config()
+ assert cli_runner.invoke(cli, ["config", "get", "base_url"]).exit_code == 1
+
+ def test_unset_missing_key_is_idempotent(self, cli_runner, isolated_home):
+ result = cli_runner.invoke(cli, ["config", "unset", "base_url"])
+
+ assert result.exit_code == 0
+ assert "not set" in result.output.lower()
+
+
+class TestConfigHelpers:
+ def test_get_config_file_path_under_home(self, isolated_home):
+ assert get_config_file_path() == str(isolated_home / ".litellm" / "config.json")
+
+ def test_load_config_missing_file_returns_empty(self, isolated_home):
+ assert load_config() == {}
+
+ def test_home_unresolvable_does_not_crash_cli(self, cli_runner, isolated_home, monkeypatch):
+ """Path.home() raises RuntimeError in HOME-less containers; invocations that
+ never needed the home dir (--api-key supplied) must keep working."""
+ monkeypatch.setattr(
+ "litellm.proxy.client.cli.commands.config.get_config_file_path",
+ _raise_home_unresolvable,
+ )
+
+ assert load_config() == {}
+
+ result = cli_runner.invoke(cli, ["--api-key", "sk-test", "config", "get"])
+ assert result.exit_code == 0
+ assert "(no config set)" in result.output
+
+ @pytest.mark.parametrize(
+ "content",
+ [
+ "{not json",
+ '{"base_url": 123}',
+ '["https://your-proxy.example.com"]',
+ '"https://your-proxy.example.com"',
+ ],
+ )
+ def test_load_config_invalid_content_returns_empty(self, isolated_home, content):
+ """A corrupt or wrongly-shaped config file must degrade to defaults, never crash the CLI."""
+ config_file = _config_path(isolated_home)
+ config_file.parent.mkdir(parents=True, exist_ok=True)
+ config_file.write_text(content)
+
+ assert load_config() == {}
+
+ def test_load_config_invalid_utf8_returns_empty(self, isolated_home):
+ """json.load raises UnicodeDecodeError (a ValueError but not a JSONDecodeError)
+ on undecodable bytes; before catching ValueError this crashed every CLI
+ invocation, including the `config set` needed to repair the file."""
+ config_file = _config_path(isolated_home)
+ config_file.parent.mkdir(parents=True, exist_ok=True)
+ config_file.write_bytes(b"\xff\xfe{}")
+
+ assert load_config() == {}
+
+ def test_save_config_round_trip_creates_dir_and_restricts_permissions(self, isolated_home):
+ save_config({"base_url": "https://your-proxy.example.com"})
+
+ assert load_config() == {"base_url": "https://your-proxy.example.com"}
+ assert stat.S_IMODE(_config_path(isolated_home).stat().st_mode) == 0o600
+
+ def test_get_config_value_unset_then_set(self, isolated_home):
+ assert get_config_value("base_url") is None
+
+ save_config({"base_url": "https://your-proxy.example.com"})
+
+ assert get_config_value("base_url") == "https://your-proxy.example.com"
+
+ def test_corrupt_config_file_warns_on_stderr_but_command_succeeds(self, cli_runner, isolated_home):
+ """Silently ignoring a broken config file leaves users debugging why their
+ stored base_url stopped applying; the CLI must keep working but say why."""
+ config_file = _config_path(isolated_home)
+ config_file.parent.mkdir(parents=True, exist_ok=True)
+ config_file.write_text("{not json")
+
+ result = cli_runner.invoke(cli, ["config", "get"])
+
+ assert result.exit_code == 0
+ assert "Warning: ignoring invalid config file" in result.stderr
+
+
+class TestWritePrivateJson:
+ def test_failed_write_preserves_previous_file_and_removes_temp(self, tmp_path):
+ """json.dump can fail partway through serializing; writing to a temp file
+ and renaming keeps the previous file intact through a crash mid-write."""
+ target = tmp_path / "config.json"
+ original = '{"base_url": "https://original.example.com"}'
+ target.write_text(original)
+
+ with pytest.raises(TypeError):
+ write_private_json(str(target), {"bad": object()})
+
+ assert target.read_text() == original
+ assert list(tmp_path.glob(".tmp-*")) == []
+
+ def test_interrupted_write_removes_temp_file(self, tmp_path, monkeypatch):
+ """Ctrl-C is BaseException, which `except Exception` misses; an interrupt
+ mid-write must not abandon a .tmp-* file next to the config forever."""
+
+ def _interrupt(*args: object, **kwargs: object) -> None:
+ raise KeyboardInterrupt()
+
+ monkeypatch.setattr("litellm.proxy.client.cli.commands.private_json.json.dump", _interrupt)
+ target = tmp_path / "config.json"
+
+ with pytest.raises(KeyboardInterrupt):
+ write_private_json(str(target), {"base_url": "https://your-proxy.example.com"})
+
+ assert not target.exists()
+ assert list(tmp_path.glob(".tmp-*")) == []
diff --git a/tests/test_litellm/proxy/client/cli/test_global_options.py b/tests/test_litellm/proxy/client/cli/test_global_options.py
index 8df763d35c2..9995cb1bca5 100644
--- a/tests/test_litellm/proxy/client/cli/test_global_options.py
+++ b/tests/test_litellm/proxy/client/cli/test_global_options.py
@@ -1,4 +1,5 @@
# stdlib imports
+import json
import os
import sys
from pathlib import Path
@@ -7,9 +8,7 @@ from unittest.mock import Mock, patch
import pytest
from click.testing import CliRunner
-sys.path.insert(
- 0, os.path.abspath("../../..")
-) # Adds the parent directory to the system path
+sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path
import litellm.proxy.client.cli
@@ -71,13 +70,9 @@ def test_base_url_trailing_slash_normalized(cli_runner):
) as mock_post,
patch("requests.get", side_effect=ValueError("stop after start request")),
):
- cli_runner.invoke(
- cli, ["--base-url", "https://gateway.litellm-sandbox.ai/", "login"]
- )
+ cli_runner.invoke(cli, ["--base-url", "https://gateway.litellm-sandbox.ai/", "login"])
- mock_post.assert_called_once_with(
- "https://gateway.litellm-sandbox.ai/sso/cli/start", timeout=10
- )
+ mock_post.assert_called_once_with("https://gateway.litellm-sandbox.ai/sso/cli/start", timeout=10)
def test_cli_version_command(cli_runner):
@@ -94,3 +89,152 @@ def test_cli_version_command(cli_runner):
assert f"LiteLLM Proxy CLI Version: {litellm_version}" in result.output
assert "LiteLLM Proxy Server URL: http://localhost:4000" in result.output
assert "LiteLLM Proxy Server Version: 1.2.3" in result.output
+
+
+@pytest.fixture
+def isolated_home(monkeypatch, tmp_path):
+ """Point HOME at tmp_path so tests never touch the developer's real ~/.litellm."""
+ monkeypatch.setenv("HOME", str(tmp_path))
+ monkeypatch.setenv("USERPROFILE", str(tmp_path))
+ monkeypatch.delenv("LITELLM_PROXY_URL", raising=False)
+ monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False)
+ return tmp_path
+
+
+def _write_config_file(home: Path, config: dict[str, str]) -> None:
+ config_dir = home / ".litellm"
+ config_dir.mkdir(exist_ok=True)
+ (config_dir / "config.json").write_text(json.dumps(config))
+
+
+def _invoke_version(cli_runner: CliRunner, *args: str):
+ with patch(
+ "litellm.proxy.client.health.HealthManagementClient.get_server_version",
+ return_value="1.2.3",
+ ):
+ return cli_runner.invoke(cli, [*args, "version"])
+
+
+def test_base_url_read_from_config_file(cli_runner, isolated_home):
+ """base_url precedence: flag > env > config file > default."""
+ _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"})
+
+ result = _invoke_version(cli_runner)
+
+ assert result.exit_code == 0
+ assert "LiteLLM Proxy Server URL: https://config-proxy.example.com" in result.output
+
+
+def test_env_var_beats_config_file_base_url(cli_runner, isolated_home, monkeypatch):
+ _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"})
+ monkeypatch.setenv("LITELLM_PROXY_URL", "http://env-proxy.example.com:5000")
+
+ result = _invoke_version(cli_runner)
+
+ assert result.exit_code == 0
+ assert "LiteLLM Proxy Server URL: http://env-proxy.example.com:5000" in result.output
+
+
+def test_base_url_flag_beats_env_var_and_config_file(cli_runner, isolated_home, monkeypatch):
+ _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"})
+ monkeypatch.setenv("LITELLM_PROXY_URL", "http://env-proxy.example.com:5000")
+
+ result = _invoke_version(cli_runner, "--base-url", "http://flag-proxy.example.com:9000")
+
+ assert result.exit_code == 0
+ assert "LiteLLM Proxy Server URL: http://flag-proxy.example.com:9000" in result.output
+
+
+def test_default_base_url_unchanged_without_config_file(cli_runner, isolated_home):
+ result = _invoke_version(cli_runner)
+
+ assert result.exit_code == 0
+ assert "LiteLLM Proxy Server URL: http://localhost:4000" in result.output
+
+
+def test_corrupt_config_file_falls_back_to_default(cli_runner, isolated_home):
+ """A corrupt config file must never crash the CLI. Exactly one warning proves
+ the config file is read once per invocation, not once per lookup."""
+ config_dir = isolated_home / ".litellm"
+ config_dir.mkdir(exist_ok=True)
+ (config_dir / "config.json").write_text("{not json")
+
+ result = _invoke_version(cli_runner)
+
+ assert result.exit_code == 0
+ assert "LiteLLM Proxy Server URL: http://localhost:4000" in result.output
+ assert result.stderr.count("Warning: ignoring invalid config file") == 1
+
+
+def test_empty_base_url_flag_is_not_treated_as_unset(cli_runner, isolated_home):
+ """`--base-url ""` explicitly provided an (empty) value; falling back to the
+ config file or localhost would silently redirect auth-sensitive commands."""
+ _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"})
+
+ result = _invoke_version(cli_runner, "--base-url", "")
+
+ assert result.exit_code == 0
+ assert "LiteLLM Proxy Server URL:" not in result.output
+
+
+def test_version_flag_reads_config_file_base_url(cli_runner, isolated_home):
+ """--version resolves through the same precedence chain as every other command."""
+ _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"})
+
+ with patch(
+ "litellm.proxy.client.health.HealthManagementClient.get_server_version",
+ return_value="1.2.3",
+ ):
+ result = cli_runner.invoke(cli, ["--version"])
+
+ assert result.exit_code == 0
+ assert "LiteLLM Proxy Server URL: https://config-proxy.example.com" in result.output
+
+
+def test_version_flag_prefers_env_var_over_config_file(cli_runner, isolated_home, monkeypatch):
+ _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"})
+ monkeypatch.setenv("LITELLM_PROXY_URL", "http://env-proxy.example.com:5000")
+
+ with patch(
+ "litellm.proxy.client.health.HealthManagementClient.get_server_version",
+ return_value="1.2.3",
+ ):
+ result = cli_runner.invoke(cli, ["--version"])
+
+ assert result.exit_code == 0
+ assert "LiteLLM Proxy Server URL: http://env-proxy.example.com:5000" in result.output
+
+
+def test_version_flag_prefers_explicit_base_url_over_config_file(cli_runner, isolated_home):
+ """An eager --version could not see the flag and silently queried the config
+ server instead of the one the user named."""
+ _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"})
+
+ with patch(
+ "litellm.proxy.client.health.HealthManagementClient.get_server_version",
+ return_value="1.2.3",
+ ):
+ result = cli_runner.invoke(cli, ["--base-url", "https://flag-proxy.example.com", "--version"])
+
+ assert result.exit_code == 0
+ assert "LiteLLM Proxy Server URL: https://flag-proxy.example.com" in result.output
+ assert "config-proxy.example.com" not in result.output
+
+
+def test_version_flag_never_sends_api_key_to_unnamed_server(cli_runner, isolated_home, monkeypatch):
+ """The version request carries a bearer token; it must reach only the server the
+ user named, never whichever host happens to sit in the config file."""
+ _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"})
+ monkeypatch.setenv("LITELLM_PROXY_API_KEY", "sk-intended-for-flag-proxy")
+
+ with patch("litellm.proxy.client.http_client.requests.request") as mock_request:
+ mock_request.return_value.json.return_value = {"litellm_version": "1.2.3"}
+ mock_request.return_value.raise_for_status.return_value = None
+ result = cli_runner.invoke(cli, ["--base-url", "https://flag-proxy.example.com", "--version"])
+
+ assert result.exit_code == 0
+ requested_urls = [call.kwargs["url"] for call in mock_request.call_args_list]
+ assert requested_urls
+ assert all(url.startswith("https://flag-proxy.example.com") for url in requested_urls)
+ sent_keys = [call.kwargs["headers"].get("Authorization") for call in mock_request.call_args_list]
+ assert sent_keys == ["Bearer sk-intended-for-flag-proxy"] * len(requested_urls)
diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py
index ef70687bd97..f31ca32f4c5 100644
--- a/tests/test_litellm/router_strategy/test_complexity_router.py
+++ b/tests/test_litellm/router_strategy/test_complexity_router.py
@@ -1417,6 +1417,102 @@ class TestLLMClassifier:
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
assert call_kwargs["metadata"] == request_metadata
+ @pytest.mark.asyncio
+ async def test_aclassify_forwards_metadata_key_used_by_chat_completions(
+ self, llm_complexity_router, mock_router_instance
+ ):
+ """/v1/chat/completions puts the request metadata under "metadata", not "litellm_metadata".
+
+ Only the routes in LITELLM_METADATA_ROUTES (/v1/messages, /v1/responses, ...) get a
+ "litellm_metadata" bucket; chat completions gets "metadata". Reading only
+ "litellm_metadata" leaves the classifier call unattributed on the most common route,
+ so _should_track_cost_callback drops it and no spend-log row is written at all,
+ which also makes the captured request body unreachable in the Logs UI.
+ """
+ mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}'))
+ request_metadata = {"user_api_key": "sk-abc", "user_api_key_team_id": "team-1"}
+ await llm_complexity_router.aclassify("hi", request_kwargs={"metadata": request_metadata})
+ call_kwargs = mock_router_instance.acompletion.call_args.kwargs
+ assert call_kwargs["metadata"] == request_metadata
+
+ @pytest.mark.asyncio
+ async def test_aclassify_captures_request_body_in_proxy_server_request(
+ self, llm_complexity_router, mock_router_instance
+ ):
+ """The classifier call must supply proxy_server_request so its request body is logged.
+
+ proxy_server_request["body"] is populated only by the proxy's HTTP ingress
+ middleware, which never runs for this internally-initiated router.acompletion
+ call. Without it _get_proxy_server_request_for_spend_logs_payload reads nothing
+ and stores "{}" for the request, so the classifier's spend-log row shows a
+ populated response but an empty request and the log cannot show which prompt
+ drove the tier decision. The captured body must carry the classification prompt
+ actually sent, so the classifier model, the classification prompt, and the user
+ text are all asserted here.
+ """
+ mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "COMPLEX"}'))
+ await llm_complexity_router.aclassify("explain quantum tunneling in depth")
+ call_kwargs = mock_router_instance.acompletion.call_args.kwargs
+ body = call_kwargs["proxy_server_request"]["body"]
+ assert body["model"] == "haiku-classifier"
+ assert body["messages"] == call_kwargs["messages"]
+ assert "explain quantum tunneling in depth" in body["messages"][0]["content"]
+ assert body["response_format"]["type"] == "json_schema"
+ assert body["response_format"]["json_schema"]["schema"]["properties"]["tier"]["enum"] == [
+ "SIMPLE",
+ "MEDIUM",
+ "COMPLEX",
+ "REASONING",
+ ]
+
+ @pytest.mark.asyncio
+ async def test_aclassify_propagates_top_level_turn_off_message_logging(
+ self, llm_complexity_router, mock_router_instance
+ ):
+ """A caller's top-level turn_off_message_logging must reach the classifier call.
+
+ Without this, a caller who opts a request out of message logging still has their
+ prompt captured in full by the classifier's proxy_server_request: the spend-log
+ redaction gate (should_redact_message_logging) reads turn_off_message_logging off
+ the classifier call's own kwargs, and this internal call is not the caller's
+ request, so it never inherits the opt-out unless it's forwarded explicitly.
+ """
+ mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}'))
+ await llm_complexity_router.aclassify("secret prompt", request_kwargs={"turn_off_message_logging": True})
+ call_kwargs = mock_router_instance.acompletion.call_args.kwargs
+ assert call_kwargs["turn_off_message_logging"] is True
+
+ @pytest.mark.asyncio
+ async def test_aclassify_propagates_metadata_slot_turn_off_message_logging(
+ self, llm_complexity_router, mock_router_instance
+ ):
+ """turn_off_message_logging set inside metadata/litellm_metadata must also propagate.
+
+ initialize_standard_callback_dynamic_params reads this flag from either the
+ top-level request kwargs or the metadata/litellm_metadata dicts (the same slots a
+ real HTTP request populates), so the classifier call must resolve it from there too.
+ """
+ mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}'))
+ await llm_complexity_router.aclassify(
+ "secret prompt", request_kwargs={"litellm_metadata": {"turn_off_message_logging": True}}
+ )
+ call_kwargs = mock_router_instance.acompletion.call_args.kwargs
+ assert call_kwargs["turn_off_message_logging"] is True
+
+ @pytest.mark.asyncio
+ async def test_aclassify_defaults_turn_off_message_logging_to_none(
+ self, llm_complexity_router, mock_router_instance
+ ):
+ """With no caller opt-out, the classifier call must not force redaction on or off.
+
+ Passing None (rather than omitting the kwarg or defaulting to False) preserves the
+ existing header- and global-setting fallbacks in should_redact_message_logging.
+ """
+ mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}'))
+ await llm_complexity_router.aclassify("hi")
+ call_kwargs = mock_router_instance.acompletion.call_args.kwargs
+ assert call_kwargs["turn_off_message_logging"] is None
+
@pytest.mark.asyncio
async def test_aclassify_strips_budget_reservation_from_classifier_metadata(
self, llm_complexity_router, mock_router_instance
@@ -2169,6 +2265,69 @@ class TestSemanticKeywordTierRules:
assert fake_router.async_embedding_kwargs[0]["metadata"] == caller_metadata
assert fake_router.async_embedding_kwargs[0]["litellm_metadata"] == caller_litellm_metadata
+ @pytest.mark.asyncio
+ async def test_semantic_embedding_call_captures_request_body_in_proxy_server_request(self, basic_config):
+ """The query embedding call must supply proxy_server_request so its request is logged.
+
+ Like the LLM classifier, this embedding is fired internally and never passes
+ through the proxy's HTTP ingress middleware, so proxy_server_request is unset and
+ the embedding's spend-log row stores "{}" for the request while its response is
+ captured. The captured body must carry the embedded input so the log shows what
+ was classified.
+ """
+ fake_router = FakeEmbeddingRouter()
+ config = {
+ **basic_config,
+ "keyword_tier_rules": [{"keywords": ["kubernetes deployment"], "tier": "REASONING"}],
+ "semantic_keyword_matching": True,
+ "embedding_model": "fake-embed",
+ "match_threshold": 0.5,
+ }
+ router = ComplexityRouter(
+ model_name="test-router",
+ litellm_router_instance=fake_router,
+ complexity_router_config=config,
+ )
+ await router.async_pre_routing_hook(
+ model="test-model",
+ request_kwargs={},
+ messages=[{"role": "user", "content": "roll out my k8s cluster"}],
+ )
+ assert fake_router.async_embedding_kwargs, "expected an embedding call for the prompt"
+ body = fake_router.async_embedding_kwargs[0]["proxy_server_request"]["body"]
+ assert body["model"] == "fake-embed"
+ assert body["input"] == ["roll out my k8s cluster"]
+
+ @pytest.mark.asyncio
+ async def test_semantic_embedding_call_propagates_turn_off_message_logging(self, basic_config):
+ """A caller's turn_off_message_logging must reach the query embedding call.
+
+ The embedding now captures the user's prompt in proxy_server_request, so a caller
+ who opts out of message logging must have that opt-out forwarded; otherwise the
+ embedding's spend-log row stores the prompt in the clear despite the parent request
+ being redacted, exposing it to anyone authorized to read the team's spend logs.
+ """
+ fake_router = FakeEmbeddingRouter()
+ config = {
+ **basic_config,
+ "keyword_tier_rules": [{"keywords": ["kubernetes deployment"], "tier": "REASONING"}],
+ "semantic_keyword_matching": True,
+ "embedding_model": "fake-embed",
+ "match_threshold": 0.5,
+ }
+ router = ComplexityRouter(
+ model_name="test-router",
+ litellm_router_instance=fake_router,
+ complexity_router_config=config,
+ )
+ await router.async_pre_routing_hook(
+ model="test-model",
+ request_kwargs={"turn_off_message_logging": True},
+ messages=[{"role": "user", "content": "roll out my k8s cluster"}],
+ )
+ assert fake_router.async_embedding_kwargs, "expected an embedding call for the prompt"
+ assert fake_router.async_embedding_kwargs[0]["turn_off_message_logging"] is True
+
@pytest.mark.asyncio
async def test_semantic_embedding_call_strips_budget_reservation(self, basic_config):
"""The embedding call must not carry the parent request's budget reservation.
diff --git a/tests/test_litellm/test_ssl_verify_unit.py b/tests/test_litellm/test_ssl_verify_unit.py
index 7cc15703a3b..c39362c01a2 100644
--- a/tests/test_litellm/test_ssl_verify_unit.py
+++ b/tests/test_litellm/test_ssl_verify_unit.py
@@ -17,7 +17,6 @@ sys.path.insert(0, str(Path(__file__).parent))
import litellm.proxy.guardrails.guardrail_hooks.aim.aim as _aim_module
import litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks as _cato_networks_module
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
-from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
from litellm.proxy.guardrails.guardrail_hooks.aim.aim import AimGuardrail
from litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks import CatoNetworksGuardrail
@@ -87,23 +86,6 @@ class TestBaseAWSLLMSSLVerify:
assert True # If we got here without error, parameter was accepted
-class TestBedrockLLMSSLVerify:
- """Test SSL verification parameter handling in BedrockLLM."""
-
- def test_bedrock_llm_accepts_ssl_verify_in_optional_params(self):
- """Test that BedrockLLM can receive ssl_verify in optional_params."""
- # This is a simple test to verify the parameter is accepted
- # The actual propagation is tested in integration tests
- bedrock_llm = BedrockLLM()
-
- # Verify the class exists and can be instantiated
- assert bedrock_llm is not None
-
- # Verify _get_ssl_verify method exists and works
- result = bedrock_llm._get_ssl_verify(ssl_verify="/path/to/cert.pem")
- assert result == "/path/to/cert.pem"
-
-
class TestAimGuardrailSSLVerify:
"""Test SSL verification parameter handling in AimGuardrail."""
diff --git a/type-discipline-budget.json b/type-discipline-budget.json
index a655768d22d..703c9ab64ff 100644
--- a/type-discipline-budget.json
+++ b/type-discipline-budget.json
@@ -1,9 +1,9 @@
{
"LIT001": {
- "limit": 23280
+ "limit": 23240
},
"LIT002": {
- "limit": 27473
+ "limit": 27434
},
"LIT003": {
"limit": 292
@@ -15,7 +15,7 @@
"limit": 0
},
"LIT006": {
- "limit": 1109
+ "limit": 1107
},
"LIT007": {
"limit": 0
@@ -24,6 +24,6 @@
"limit": 1004
},
"LIT009": {
- "limit": 2495
+ "limit": 2453
}
}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx
index 8dffc80a70e..5650bd1d7e4 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx
@@ -190,7 +190,7 @@ const OAuthFormFields: React.FC = ({
label={
}
name="issuer"