From ec004a54e9e5e3f3984e9f18a4f587532f6149e7 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 2 Apr 2026 09:56:45 +0530 Subject: [PATCH 01/75] feat(vertex_ai): propagate metadata labels to embedding, Imagen, rerank - Add pop_vertex_request_labels / vertex_request_labels_from_litellm_params in common_utils - Vertex embeddings: pass litellm_params, set predict body labels; Gemini uses shared helper - Imagen: top-level labels from metadata; rerank: userLabels for Discovery Engine Rank API - Thread litellm_params through rerank handler and all BaseRerankConfig implementations Made-with: Cursor --- .../llms/base_llm/rerank/transformation.py | 1 + litellm/llms/cohere/rerank/transformation.py | 1 + .../llms/cohere/rerank_v2/transformation.py | 1 + litellm/llms/custom_httpx/llm_http_handler.py | 2 + .../llms/deepinfra/rerank/transformation.py | 1 + .../fireworks_ai/rerank/transformation.py | 1 + .../llms/hosted_vllm/rerank/transformation.py | 1 + .../llms/huggingface/rerank/transformation.py | 1 + litellm/llms/jina_ai/rerank/transformation.py | 6 +- .../rerank/ranking_transformation.py | 2 + .../llms/nvidia_nim/rerank/transformation.py | 1 + litellm/llms/vertex_ai/common_utils.py | 40 +++ .../llms/vertex_ai/gemini/transformation.py | 13 +- .../vertex_imagen_transformation.py | 12 +- .../llms/vertex_ai/rerank/transformation.py | 10 +- .../vertex_embeddings/embedding_handler.py | 15 +- .../vertex_embeddings/transformation.py | 25 +- .../llms/vertex_ai/vertex_embeddings/types.py | 3 +- litellm/llms/voyage/rerank/transformation.py | 6 +- litellm/llms/watsonx/rerank/transformation.py | 1 + litellm/main.py | 1 + litellm/rerank_api/main.py | 29 ++- ...rtex_ai_image_generation_transformation.py | 37 ++- .../test_vertex_ai_rerank_transformation.py | 222 +++++++++-------- .../vertex_ai/test_vertex_ai_common_utils.py | 235 +++++++++++------- 25 files changed, 426 insertions(+), 241 deletions(-) diff --git a/litellm/llms/base_llm/rerank/transformation.py b/litellm/llms/base_llm/rerank/transformation.py index 7874201f7f0..166f876ba04 100644 --- a/litellm/llms/base_llm/rerank/transformation.py +++ b/litellm/llms/base_llm/rerank/transformation.py @@ -33,6 +33,7 @@ class BaseRerankConfig(ABC): model: str, optional_rerank_params: Dict, headers: dict, + litellm_params: Optional[dict] = None, ) -> dict: return {} diff --git a/litellm/llms/cohere/rerank/transformation.py b/litellm/llms/cohere/rerank/transformation.py index 531b94d1805..64ae8e8ffa7 100644 --- a/litellm/llms/cohere/rerank/transformation.py +++ b/litellm/llms/cohere/rerank/transformation.py @@ -111,6 +111,7 @@ class CohereRerankConfig(BaseRerankConfig): model: str, optional_rerank_params: Dict, headers: dict, + litellm_params: Optional[dict] = None, ) -> dict: if "query" not in optional_rerank_params: raise ValueError("query is required for Cohere rerank") diff --git a/litellm/llms/cohere/rerank_v2/transformation.py b/litellm/llms/cohere/rerank_v2/transformation.py index 60d22ff4be0..4c800d6455d 100644 --- a/litellm/llms/cohere/rerank_v2/transformation.py +++ b/litellm/llms/cohere/rerank_v2/transformation.py @@ -71,6 +71,7 @@ class CohereRerankV2Config(CohereRerankConfig): model: str, optional_rerank_params: Dict, headers: dict, + litellm_params: Optional[dict] = None, ) -> dict: if "query" not in optional_rerank_params: raise ValueError("query is required for Cohere rerank") diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 4c9abaad908..b8635085bcc 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -978,6 +978,7 @@ class BaseLLMHTTPHandler: api_key: Optional[str] = None, api_base: Optional[str] = None, client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + litellm_params: Optional[Dict[str, Any]] = None, ) -> RerankResponse: # get config from model, custom llm provider headers = provider_config.validate_environment( @@ -997,6 +998,7 @@ class BaseLLMHTTPHandler: model=model, optional_rerank_params=optional_rerank_params, headers=headers, + litellm_params=litellm_params, ) ## LOGGING diff --git a/litellm/llms/deepinfra/rerank/transformation.py b/litellm/llms/deepinfra/rerank/transformation.py index 71e300d258c..276735f4758 100644 --- a/litellm/llms/deepinfra/rerank/transformation.py +++ b/litellm/llms/deepinfra/rerank/transformation.py @@ -132,6 +132,7 @@ class DeepinfraRerankConfig(BaseRerankConfig): model: str, optional_rerank_params: Dict, headers: dict, + litellm_params: Optional[dict] = None, ) -> dict: # Convert OptionalRerankParams to dict as expected by parent class if optional_rerank_params is None: diff --git a/litellm/llms/fireworks_ai/rerank/transformation.py b/litellm/llms/fireworks_ai/rerank/transformation.py index eb92399a058..4a7b64b9b77 100644 --- a/litellm/llms/fireworks_ai/rerank/transformation.py +++ b/litellm/llms/fireworks_ai/rerank/transformation.py @@ -127,6 +127,7 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): model: str, optional_rerank_params: Dict, headers: dict, + litellm_params: Optional[dict] = None, ) -> dict: """ Transform request to Fireworks AI rerank format diff --git a/litellm/llms/hosted_vllm/rerank/transformation.py b/litellm/llms/hosted_vllm/rerank/transformation.py index 8066e53afc7..60b6dc7d23d 100644 --- a/litellm/llms/hosted_vllm/rerank/transformation.py +++ b/litellm/llms/hosted_vllm/rerank/transformation.py @@ -121,6 +121,7 @@ class HostedVLLMRerankConfig(BaseRerankConfig): model: str, optional_rerank_params: Dict, headers: dict, + litellm_params: Optional[dict] = None, ) -> dict: if "query" not in optional_rerank_params: raise ValueError("query is required for Hosted VLLM rerank") diff --git a/litellm/llms/huggingface/rerank/transformation.py b/litellm/llms/huggingface/rerank/transformation.py index 3f83b8e422d..2c847b617ef 100644 --- a/litellm/llms/huggingface/rerank/transformation.py +++ b/litellm/llms/huggingface/rerank/transformation.py @@ -146,6 +146,7 @@ class HuggingFaceRerankConfig(BaseRerankConfig): model: str, optional_rerank_params: Union[OptionalRerankParams, dict], headers: dict, + litellm_params: Optional[dict] = None, ) -> dict: if "query" not in optional_rerank_params: raise ValueError("query is required for HuggingFace rerank") diff --git a/litellm/llms/jina_ai/rerank/transformation.py b/litellm/llms/jina_ai/rerank/transformation.py index 48d876f8ea2..ad4416925a6 100644 --- a/litellm/llms/jina_ai/rerank/transformation.py +++ b/litellm/llms/jina_ai/rerank/transformation.py @@ -74,7 +74,11 @@ class JinaAIRerankConfig(BaseRerankConfig): return cleaned_base def transform_rerank_request( - self, model: str, optional_rerank_params: Dict, headers: Dict + self, + model: str, + optional_rerank_params: Dict, + headers: Dict, + litellm_params: Optional[dict] = None, ) -> Dict: return {"model": model, **optional_rerank_params} diff --git a/litellm/llms/nvidia_nim/rerank/ranking_transformation.py b/litellm/llms/nvidia_nim/rerank/ranking_transformation.py index 757d874bf31..b9a46b8ac2b 100644 --- a/litellm/llms/nvidia_nim/rerank/ranking_transformation.py +++ b/litellm/llms/nvidia_nim/rerank/ranking_transformation.py @@ -66,6 +66,7 @@ class NvidiaNimRankingConfig(NvidiaNimRerankConfig): model: str, optional_rerank_params: Dict, headers: dict, + litellm_params: Optional[dict] = None, ) -> dict: """ Transform request, using clean model name without 'ranking/' prefix. @@ -75,4 +76,5 @@ class NvidiaNimRankingConfig(NvidiaNimRerankConfig): model=clean_model, optional_rerank_params=optional_rerank_params, headers=headers, + litellm_params=litellm_params, ) diff --git a/litellm/llms/nvidia_nim/rerank/transformation.py b/litellm/llms/nvidia_nim/rerank/transformation.py index bd5abac60c8..fc317293acc 100644 --- a/litellm/llms/nvidia_nim/rerank/transformation.py +++ b/litellm/llms/nvidia_nim/rerank/transformation.py @@ -177,6 +177,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig): model: str, optional_rerank_params: Dict, headers: dict, + litellm_params: Optional[dict] = None, ) -> dict: """ Transform request to Nvidia NIM format. diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index 5895a91f3aa..a8cc5e7bdec 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -27,6 +27,46 @@ class VertexAIError(BaseLLMException): super().__init__(message=message, status_code=status_code, headers=headers) +def vertex_request_labels_from_litellm_params( + litellm_params: Optional[dict], +) -> Optional[Dict[str, str]]: + """ + Build Vertex/GCP billing labels from LiteLLM ``litellm_params["metadata"]``, + using ``requester_metadata`` string key-value pairs (same convention as Gemini). + """ + if not litellm_params or "metadata" not in litellm_params: + return None + metadata = litellm_params["metadata"] + if metadata is None or not isinstance(metadata, dict): + return None + if "requester_metadata" not in metadata: + return None + rm = metadata["requester_metadata"] + if not isinstance(rm, dict): + return None + labels = {k: v for k, v in rm.items() if isinstance(v, str)} + return labels if labels else None + + +def pop_vertex_request_labels( + optional_params: Optional[dict], + litellm_params: Optional[dict], +) -> Optional[Dict[str, str]]: + """ + Resolve labels from optional ``labels`` (Gemini-style) and/or + ``litellm_params["metadata"]["requester_metadata"]``. Pops ``labels`` from + optional_params when present. + """ + labels: Optional[Dict[str, str]] = None + if optional_params is not None and "labels" in optional_params: + raw = optional_params.pop("labels") + if isinstance(raw, dict): + labels = {k: v for k, v in raw.items() if isinstance(v, str)} + if labels is None: + labels = vertex_request_labels_from_litellm_params(litellm_params) + return labels if labels else None + + class VertexAIModelRoute(str, Enum): """Enum for Vertex AI model routing""" diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 7945c44d44c..723fa2b1616 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -23,6 +23,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( response_schema_prompt, ) from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.llms.vertex_ai.common_utils import pop_vertex_request_labels from litellm.types.files import ( get_file_mime_type_for_file_type, get_file_type_from_extension, @@ -711,16 +712,8 @@ def _transform_request_body( # noqa: PLR0915 optional_params.pop("output_config", None) config_fields = GenerationConfig.__annotations__.keys() - # If the LiteLLM client sends Gemini-supported parameter "labels", add it - # as "labels" field to the request sent to the Gemini backend. - labels: Optional[dict[str, str]] = optional_params.pop("labels", None) - # If the LiteLLM client sends OpenAI-supported parameter "metadata", add it - # as "labels" field to the request sent to the Gemini backend. - if labels is None and "metadata" in litellm_params: - metadata = litellm_params["metadata"] - if metadata is not None and "requester_metadata" in metadata: - rm = metadata["requester_metadata"] - labels = {k: v for k, v in rm.items() if isinstance(v, str)} + # labels: optional explicit param and/or metadata.requester_metadata (OpenAI metadata) + labels = pop_vertex_request_labels(optional_params, litellm_params) filtered_params = { k: v diff --git a/litellm/llms/vertex_ai/image_generation/vertex_imagen_transformation.py b/litellm/llms/vertex_ai/image_generation/vertex_imagen_transformation.py index 1c7696d55a2..05ebd685d91 100644 --- a/litellm/llms/vertex_ai/image_generation/vertex_imagen_transformation.py +++ b/litellm/llms/vertex_ai/image_generation/vertex_imagen_transformation.py @@ -7,7 +7,10 @@ import litellm from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, ) -from litellm.llms.vertex_ai.common_utils import get_vertex_base_url +from litellm.llms.vertex_ai.common_utils import ( + get_vertex_base_url, + pop_vertex_request_labels, +) from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import ( @@ -203,13 +206,16 @@ class VertexAIImagenImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): "sampleCount": 1, } - # Merge with optional params + labels = pop_vertex_request_labels(optional_params, litellm_params) + # Merge with optional params (after popping labels so they are not sent as Imagen parameters) parameters = {**default_params, **optional_params} - request_body = { + request_body: dict = { "instances": [{"prompt": prompt}], "parameters": parameters, } + if labels: + request_body["labels"] = labels return request_body diff --git a/litellm/llms/vertex_ai/rerank/transformation.py b/litellm/llms/vertex_ai/rerank/transformation.py index 53651839671..3b84972e946 100644 --- a/litellm/llms/vertex_ai/rerank/transformation.py +++ b/litellm/llms/vertex_ai/rerank/transformation.py @@ -11,12 +11,15 @@ import httpx import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig +from litellm.llms.vertex_ai.common_utils import ( + vertex_request_labels_from_litellm_params, +) from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.secret_managers.main import get_secret_str from litellm.types.rerank import ( + RerankBilledUnits, RerankResponse, RerankResponseMeta, - RerankBilledUnits, RerankResponseResult, ) @@ -109,6 +112,7 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase): model: str, optional_rerank_params: Dict, headers: dict, + litellm_params: Optional[dict] = None, ) -> dict: """ Transform the request from Cohere format to Vertex AI Discovery Engine format @@ -145,6 +149,10 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase): # When return_documents is False, we want to ignore record details (return only IDs) request_data["ignoreRecordDetailsInResponse"] = not return_documents + user_labels = vertex_request_labels_from_litellm_params(litellm_params) + if user_labels: + request_data["userLabels"] = user_labels + return request_data def transform_rerank_response( diff --git a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py index 5fffd983c24..729cc9c3ead 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py @@ -1,4 +1,4 @@ -from typing import Literal, Optional, Union +from typing import Dict, Literal, Optional, Union import httpx @@ -44,6 +44,7 @@ class VertexEmbedding(VertexBase): vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None, gemini_api_key: Optional[str] = None, extra_headers: Optional[dict] = None, + litellm_params: Optional[Dict] = None, ) -> EmbeddingResponse: if aembedding is True: return self.async_embedding( # type: ignore @@ -61,6 +62,7 @@ class VertexEmbedding(VertexBase): vertex_credentials=vertex_credentials, gemini_api_key=gemini_api_key, extra_headers=extra_headers, + litellm_params=litellm_params, ) should_use_v1beta1_features = self.is_using_v1beta1_features( @@ -91,7 +93,10 @@ class VertexEmbedding(VertexBase): ) headers = self.set_headers(auth_header=auth_header, extra_headers=extra_headers) vertex_request: VertexEmbeddingRequest = litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request( - input=input, optional_params=optional_params, model=model + input=input, + optional_params=optional_params, + model=model, + litellm_params=litellm_params, ) _client_params = {} @@ -154,6 +159,7 @@ class VertexEmbedding(VertexBase): gemini_api_key: Optional[str] = None, extra_headers: Optional[dict] = None, encoding=None, + litellm_params: Optional[Dict] = None, ) -> EmbeddingResponse: """ Async embedding implementation @@ -185,7 +191,10 @@ class VertexEmbedding(VertexBase): ) headers = self.set_headers(auth_header=auth_header, extra_headers=extra_headers) vertex_request: VertexEmbeddingRequest = litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request( - input=input, optional_params=optional_params, model=model + input=input, + optional_params=optional_params, + model=model, + litellm_params=litellm_params, ) _async_client_params = {} diff --git a/litellm/llms/vertex_ai/vertex_embeddings/transformation.py b/litellm/llms/vertex_ai/vertex_embeddings/transformation.py index 132f29987af..a3018c5bd6b 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/transformation.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/transformation.py @@ -3,6 +3,7 @@ from typing import List, Literal, Optional, Union from pydantic import BaseModel +from litellm.llms.vertex_ai.common_utils import pop_vertex_request_labels from litellm.types.utils import EmbeddingResponse, Usage from .types import * @@ -100,7 +101,11 @@ class VertexAITextEmbeddingConfig(BaseModel): return optional_params def transform_openai_request_to_vertex_embedding_request( - self, input: Union[list, str], optional_params: dict, model: str + self, + input: Union[list, str], + optional_params: dict, + model: str, + litellm_params: Optional[dict] = None, ) -> VertexEmbeddingRequest: """ Transforms an openai request to a vertex embedding request. @@ -108,14 +113,24 @@ class VertexAITextEmbeddingConfig(BaseModel): # Import here to avoid circular import issues with litellm.__init__ from litellm.llms.vertex_ai.vertex_embeddings.bge import VertexBGEConfig + labels = pop_vertex_request_labels(optional_params, litellm_params) + if model.isdigit(): - return self._transform_openai_request_to_fine_tuned_embedding_request( - input, optional_params, model + vertex_request = ( + self._transform_openai_request_to_fine_tuned_embedding_request( + input, optional_params, model + ) ) + if labels: + vertex_request["labels"] = labels + return vertex_request if VertexBGEConfig.is_bge_model(model): - return VertexBGEConfig.transform_request( + vertex_request = VertexBGEConfig.transform_request( input=input, optional_params=optional_params, model=model ) + if labels: + vertex_request["labels"] = labels + return vertex_request vertex_request: VertexEmbeddingRequest = VertexEmbeddingRequest() vertex_text_embedding_input_list: List[TextEmbeddingInput] = [] @@ -133,6 +148,8 @@ class VertexAITextEmbeddingConfig(BaseModel): vertex_request["instances"] = vertex_text_embedding_input_list vertex_request["parameters"] = EmbeddingParameters(**optional_params) + if labels: + vertex_request["labels"] = labels return vertex_request diff --git a/litellm/llms/vertex_ai/vertex_embeddings/types.py b/litellm/llms/vertex_ai/vertex_embeddings/types.py index 317b9c4fb81..bf73f4d193a 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/types.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/types.py @@ -3,7 +3,7 @@ Types for Vertex Embeddings Requests """ from enum import Enum -from typing import List, Optional, Union +from typing import Dict, List, Optional, Union from typing_extensions import TypedDict @@ -56,6 +56,7 @@ class VertexEmbeddingRequest(TypedDict, total=False): List[TextEmbeddingFineTunedInput], ] parameters: Optional[Union[EmbeddingParameters, TextEmbeddingFineTunedParameters]] + labels: Optional[Dict[str, str]] # Example usage: diff --git a/litellm/llms/voyage/rerank/transformation.py b/litellm/llms/voyage/rerank/transformation.py index 521dae980d5..d64450a1211 100644 --- a/litellm/llms/voyage/rerank/transformation.py +++ b/litellm/llms/voyage/rerank/transformation.py @@ -67,7 +67,11 @@ class VoyageRerankConfig(BaseRerankConfig): return api_base def transform_rerank_request( - self, model: str, optional_rerank_params: Dict, headers: Dict + self, + model: str, + optional_rerank_params: Dict, + headers: Dict, + litellm_params: Optional[dict] = None, ) -> Dict: return {"model": model, **optional_rerank_params} diff --git a/litellm/llms/watsonx/rerank/transformation.py b/litellm/llms/watsonx/rerank/transformation.py index 4f8e196f254..202760f68a6 100644 --- a/litellm/llms/watsonx/rerank/transformation.py +++ b/litellm/llms/watsonx/rerank/transformation.py @@ -143,6 +143,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): model: str, optional_rerank_params: Dict, headers: dict, + litellm_params: Optional[dict] = None, ) -> dict: """ Transform request to IBM watsonx.ai rerank format diff --git a/litellm/main.py b/litellm/main.py index eace9c630ba..35a965d50f5 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -5302,6 +5302,7 @@ def embedding( # noqa: PLR0915 api_key=api_key, api_base=api_base, client=client, + litellm_params=litellm_params_dict, ) elif custom_llm_provider == "oobabooga": response = oobabooga.embedding( diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index 9868634362f..0dd18b6a5bf 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -163,19 +163,21 @@ def rerank( # noqa: PLR0915 model_response = RerankResponse() + rerank_litellm_params = { + "litellm_call_id": litellm_call_id, + "proxy_server_request": proxy_server_request, + "model_info": model_info, + "preset_cache_key": None, + "stream_response": {}, + **optional_params.model_dump(exclude_unset=True), + } + litellm_logging_obj.update_from_kwargs( kwargs=kwargs, model=model, user=user, optional_params=dict(optional_rerank_params), - litellm_params={ - "litellm_call_id": litellm_call_id, - "proxy_server_request": proxy_server_request, - "model_info": model_info, - "preset_cache_key": None, - "stream_response": {}, - **optional_params.model_dump(exclude_unset=True), - }, + litellm_params=rerank_litellm_params, custom_llm_provider=_custom_llm_provider, ) @@ -214,6 +216,7 @@ def rerank( # noqa: PLR0915 headers=headers or litellm.headers or {}, client=client, model_response=model_response, + litellm_params=rerank_litellm_params, ) elif _custom_llm_provider == litellm.LlmProviders.AZURE_AI: api_base = ( @@ -235,6 +238,7 @@ def rerank( # noqa: PLR0915 headers=headers or litellm.headers or {}, client=client, model_response=model_response, + litellm_params=rerank_litellm_params, ) elif _custom_llm_provider == litellm.LlmProviders.INFINITY: # Implement Infinity rerank logic @@ -265,6 +269,7 @@ def rerank( # noqa: PLR0915 headers=headers or litellm.headers or {}, client=client, model_response=model_response, + litellm_params=rerank_litellm_params, ) elif _custom_llm_provider == litellm.LlmProviders.TOGETHER_AI: # Implement Together AI rerank logic @@ -318,6 +323,7 @@ def rerank( # noqa: PLR0915 headers=headers or litellm.headers or {}, client=client, model_response=model_response, + litellm_params=rerank_litellm_params, ) elif _custom_llm_provider == litellm.LlmProviders.NVIDIA_NIM: if dynamic_api_key is None: @@ -346,6 +352,7 @@ def rerank( # noqa: PLR0915 headers=headers or litellm.headers or {}, client=client, model_response=model_response, + litellm_params=rerank_litellm_params, ) elif _custom_llm_provider == litellm.LlmProviders.BEDROCK: api_base = ( @@ -409,6 +416,7 @@ def rerank( # noqa: PLR0915 headers=headers or litellm.headers or {}, client=client, model_response=model_response, + litellm_params=rerank_litellm_params, ) elif _custom_llm_provider == litellm.LlmProviders.DEEPINFRA: @@ -442,6 +450,7 @@ def rerank( # noqa: PLR0915 headers=headers or litellm.headers or {}, client=client, model_response=model_response, + litellm_params=rerank_litellm_params, ) elif _custom_llm_provider == litellm.LlmProviders.FIREWORKS_AI: api_key = ( @@ -472,6 +481,7 @@ def rerank( # noqa: PLR0915 headers=headers or litellm.headers or {}, client=client, model_response=model_response, + litellm_params=rerank_litellm_params, ) elif _custom_llm_provider == litellm.LlmProviders.VOYAGE: api_key = ( @@ -500,6 +510,7 @@ def rerank( # noqa: PLR0915 headers=headers or litellm.headers or {}, client=client, model_response=model_response, + litellm_params=rerank_litellm_params, ) elif _custom_llm_provider == litellm.LlmProviders.WATSONX: credentials = IBMWatsonXMixin.get_watsonx_credentials( @@ -527,6 +538,7 @@ def rerank( # noqa: PLR0915 headers=headers or litellm.headers or {}, client=client, model_response=model_response, + litellm_params=rerank_litellm_params, ) else: # Generic handler for all providers that use base_llm_http_handler @@ -559,6 +571,7 @@ def rerank( # noqa: PLR0915 headers=headers or litellm.headers or {}, client=client, model_response=model_response, + litellm_params=rerank_litellm_params, ) # Placeholder return diff --git a/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py b/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py index 350fd75d3d8..61b5e3c9249 100644 --- a/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py @@ -67,7 +67,9 @@ class TestVertexAIGeminiImageGenerationConfig: def test_get_supported_openai_params_includes_native_gemini_params(self): """Test that native Gemini imageConfig params are supported""" - supported = self.config.get_supported_openai_params("gemini-3-pro-image-preview") + supported = self.config.get_supported_openai_params( + "gemini-3-pro-image-preview" + ) assert "aspectRatio" in supported assert "aspect_ratio" in supported assert "imageSize" in supported @@ -188,11 +190,11 @@ class TestVertexAIGeminiImageGenerationConfig: { "modality": "IMAGE", "tokenCount": 39, - } + }, ], "candidatesTokenCount": 17, "totalTokenCount": 110, - } + }, } mock_response.headers = {} @@ -219,7 +221,6 @@ class TestVertexAIGeminiImageGenerationConfig: assert result.usage.output_tokens == 17 assert result.usage.total_tokens == 110 - def test_transform_image_generation_response_multiple_images(self): """Test response transformation with multiple images""" mock_response = MagicMock(spec=httpx.Response) @@ -305,7 +306,10 @@ class TestVertexAIGeminiImageGenerationConfig: assert len(result.data) == 1 assert result.data[0].b64_json == "base64_encoded_image_data" - assert result.data[0].provider_specific_fields["thought_signature"] == "test_signature_abc123" + assert ( + result.data[0].provider_specific_fields["thought_signature"] + == "test_signature_abc123" + ) class TestVertexAIImagenImageGenerationConfig: @@ -369,14 +373,26 @@ class TestVertexAIImagenImageGenerationConfig: assert request["parameters"]["sampleCount"] == 2 assert request["parameters"]["aspectRatio"] == "16:9" + def test_transform_image_generation_request_labels_from_metadata(self): + """Billing labels from litellm_params.metadata.requester_metadata on predict body.""" + request = self.config.transform_image_generation_request( + model="imagegeneration@006", + prompt="A cat", + optional_params={}, + litellm_params={ + "metadata": {"requester_metadata": {"team": "platform", "env": "prod"}} + }, + headers={}, + ) + assert request["labels"] == {"team": "platform", "env": "prod"} + assert "labels" not in request["parameters"] + def test_transform_image_generation_response(self): """Test response transformation""" mock_response = MagicMock(spec=httpx.Response) mock_response.status_code = 200 mock_response.json.return_value = { - "predictions": [ - {"bytesBase64Encoded": "base64_encoded_image_data"} - ] + "predictions": [{"bytesBase64Encoded": "base64_encoded_image_data"}] } mock_response.headers = {} @@ -453,9 +469,7 @@ class TestGetVertexAIImageGenerationConfig: config = get_vertex_ai_image_generation_config("imagen-4.0-generate-001") assert isinstance(config, VertexAIImagenImageGenerationConfig) - config = get_vertex_ai_image_generation_config( - "vertex_ai/imagegeneration@006" - ) + config = get_vertex_ai_image_generation_config("vertex_ai/imagegeneration@006") assert isinstance(config, VertexAIImagenImageGenerationConfig) def test_get_non_gemini_model_config(self): @@ -548,4 +562,3 @@ class TestVertexAIImageGenerationIntegration: assert "us-central1" in url assert "imagegeneration@006" in url assert "predict" in url - diff --git a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py index 5bf2cb97fa9..fb5b0c76eaa 100644 --- a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py @@ -63,13 +63,16 @@ class TestVertexAIRerankTransform: import litellm # Set vertex_project attribute if it doesn't exist - if not hasattr(litellm, 'vertex_project'): + if not hasattr(litellm, "vertex_project"): litellm.vertex_project = None original_project = litellm.vertex_project litellm.vertex_project = "litellm-project-456" # Reset mock call count mock_ensure_access_token.reset_mock() - mock_ensure_access_token.return_value = ("mock-token", "litellm-project-456") + mock_ensure_access_token.return_value = ( + "mock-token", + "litellm-project-456", + ) try: url = self.config.get_complete_url(api_base=None, model=self.model) expected_url = "https://discoveryengine.googleapis.com/v1/projects/litellm-project-456/locations/global/rankingConfigs/default_ranking_config:rank" @@ -82,15 +85,19 @@ class TestVertexAIRerankTransform: import litellm # Set vertex_project to None to ensure no project ID is available - if not hasattr(litellm, 'vertex_project'): + if not hasattr(litellm, "vertex_project"): litellm.vertex_project = None original_project = litellm.vertex_project litellm.vertex_project = None # Reset mock and set it to raise an error mock_ensure_access_token.reset_mock() - mock_ensure_access_token.side_effect = ValueError("Vertex AI project ID is required") + mock_ensure_access_token.side_effect = ValueError( + "Vertex AI project ID is required" + ) try: - with pytest.raises(ValueError, match="Vertex AI project ID is required"): + with pytest.raises( + ValueError, match="Vertex AI project ID is required" + ): self.config.get_complete_url(api_base=None, model=self.model) finally: litellm.vertex_project = original_project @@ -109,15 +116,13 @@ class TestVertexAIRerankTransform: self.config._ensure_access_token = mock_ensure_access_token headers = self.config.validate_environment( - headers={}, - model=self.model, - api_key=None + headers={}, model=self.model, api_key=None ) expected_headers = { "Authorization": "Bearer test-access-token", "Content-Type": "application/json", - "X-Goog-User-Project": "test-project-123" + "X-Goog-User-Project": "test-project-123", } assert headers == expected_headers @@ -127,24 +132,22 @@ class TestVertexAIRerankTransform: "query": "What is Google Gemini?", "documents": [ "Gemini is a cutting edge large language model created by Google.", - "The Gemini zodiac symbol often depicts two figures standing side-by-side." + "The Gemini zodiac symbol often depicts two figures standing side-by-side.", ], - "top_n": 2 + "top_n": 2, } - + request_data = self.config.transform_rerank_request( - model=self.model, - optional_rerank_params=optional_params, - headers={} + model=self.model, optional_rerank_params=optional_params, headers={} ) - + # Verify basic structure assert request_data["model"] == self.model assert request_data["query"] == "What is Google Gemini?" assert request_data["topN"] == 2 assert "records" in request_data assert len(request_data["records"]) == 2 - + # Verify record structure for i, record in enumerate(request_data["records"]): assert "id" in record @@ -158,20 +161,25 @@ class TestVertexAIRerankTransform: optional_params = { "query": "What is Google Gemini?", "documents": [ - {"text": "Gemini is a cutting edge large language model created by Google.", "title": "Custom Title 1"}, - {"text": "The Gemini zodiac symbol often depicts two figures standing side-by-side."} - ] + { + "text": "Gemini is a cutting edge large language model created by Google.", + "title": "Custom Title 1", + }, + { + "text": "The Gemini zodiac symbol often depicts two figures standing side-by-side." + }, + ], } - + request_data = self.config.transform_rerank_request( - model=self.model, - optional_rerank_params=optional_params, - headers={} + model=self.model, optional_rerank_params=optional_params, headers={} ) - + # Verify record structure with custom titles assert request_data["records"][0]["title"] == "Custom Title 1" - assert request_data["records"][1]["title"] == "The Gemini zodiac" # First 3 words + assert ( + request_data["records"][1]["title"] == "The Gemini zodiac" + ) # First 3 words def test_transform_rerank_request_return_documents_mapping(self): """Test return_documents to ignoreRecordDetailsInResponse mapping.""" @@ -179,43 +187,50 @@ class TestVertexAIRerankTransform: optional_params_true = { "query": "test query", "documents": ["doc1", "doc2"], - "return_documents": True + "return_documents": True, } - + request_data_true = self.config.transform_rerank_request( - model=self.model, - optional_rerank_params=optional_params_true, - headers={} + model=self.model, optional_rerank_params=optional_params_true, headers={} ) assert request_data_true["ignoreRecordDetailsInResponse"] == False - + # Test return_documents=False optional_params_false = { "query": "test query", "documents": ["doc1", "doc2"], - "return_documents": False + "return_documents": False, } - + request_data_false = self.config.transform_rerank_request( - model=self.model, - optional_rerank_params=optional_params_false, - headers={} + model=self.model, optional_rerank_params=optional_params_false, headers={} ) assert request_data_false["ignoreRecordDetailsInResponse"] == True - + # Test return_documents not specified (should default to True) - optional_params_default = { - "query": "test query", - "documents": ["doc1", "doc2"] - } - + optional_params_default = {"query": "test query", "documents": ["doc1", "doc2"]} + request_data_default = self.config.transform_rerank_request( - model=self.model, - optional_rerank_params=optional_params_default, - headers={} + model=self.model, optional_rerank_params=optional_params_default, headers={} ) assert request_data_default["ignoreRecordDetailsInResponse"] == False + def test_transform_rerank_request_user_labels_from_metadata(self): + """Discovery Engine Rank API uses userLabels (string map) for billing.""" + optional_params = { + "query": "q", + "documents": ["a", "b"], + } + request_data = self.config.transform_rerank_request( + model=self.model, + optional_rerank_params=optional_params, + headers={}, + litellm_params={ + "metadata": {"requester_metadata": {"app": "litellm", "tier": "1"}} + }, + ) + assert request_data["userLabels"] == {"app": "litellm", "tier": "1"} + def test_transform_rerank_request_missing_required_params(self): """Test that transform_rerank_request handles missing required parameters.""" # Test missing query @@ -223,15 +238,17 @@ class TestVertexAIRerankTransform: self.config.transform_rerank_request( model=self.model, optional_rerank_params={"documents": ["doc1"]}, - headers={} + headers={}, ) - + # Test missing documents - with pytest.raises(ValueError, match="documents is required for Vertex AI rerank"): + with pytest.raises( + ValueError, match="documents is required for Vertex AI rerank" + ): self.config.transform_rerank_request( model=self.model, optional_rerank_params={"query": "test query"}, - headers={} + headers={}, ) def test_transform_rerank_response_success(self): @@ -243,34 +260,34 @@ class TestVertexAIRerankTransform: "id": "1", "score": 0.98, "title": "The Science of a Blue Sky", - "content": "The sky appears blue due to a phenomenon called Rayleigh scattering." + "content": "The sky appears blue due to a phenomenon called Rayleigh scattering.", }, { "id": "0", "score": 0.64, "title": "The Color of the Sky: A Poem", - "content": "A canvas stretched across the day, Where sunlight learns to dance and play." - } + "content": "A canvas stretched across the day, Where sunlight learns to dance and play.", + }, ] } - + # Create mock httpx response mock_response = MagicMock(spec=httpx.Response) mock_response.json.return_value = response_data mock_response.text = json.dumps(response_data) - + # Create mock logging object mock_logging = MagicMock() - + model_response = RerankResponse() - + result = self.config.transform_rerank_response( model=self.model, raw_response=mock_response, model_response=model_response, logging_obj=mock_logging, ) - + # Verify response structure assert result.id == f"vertex_ai_rerank_{self.model}" assert len(result.results) == 2 @@ -278,34 +295,29 @@ class TestVertexAIRerankTransform: assert result.results[0]["relevance_score"] == 0.98 assert result.results[1]["index"] == 0 assert result.results[1]["relevance_score"] == 0.64 - + # Verify metadata assert result.meta["billed_units"]["search_units"] == 2 def test_transform_rerank_response_with_ignore_record_details(self): """Test response transformation when ignoreRecordDetailsInResponse=true.""" # Mock response with only IDs (when ignoreRecordDetailsInResponse=true) - response_data = { - "records": [ - {"id": "1"}, - {"id": "0"} - ] - } - + response_data = {"records": [{"id": "1"}, {"id": "0"}]} + mock_response = MagicMock(spec=httpx.Response) mock_response.json.return_value = response_data mock_response.text = json.dumps(response_data) - + mock_logging = MagicMock() model_response = RerankResponse() - + result = self.config.transform_rerank_response( model=self.model, raw_response=mock_response, model_response=model_response, logging_obj=mock_logging, ) - + # Verify response structure with default scores assert len(result.results) == 2 assert result.results[0]["index"] == 1 # 0-based index @@ -318,10 +330,10 @@ class TestVertexAIRerankTransform: mock_response = MagicMock(spec=httpx.Response) mock_response.json.side_effect = json.JSONDecodeError("Invalid JSON", "doc", 0) mock_response.text = "Invalid JSON response" - + mock_logging = MagicMock() model_response = RerankResponse() - + with pytest.raises(ValueError, match="Failed to parse response"): self.config.transform_rerank_response( model=self.model, @@ -345,14 +357,14 @@ class TestVertexAIRerankTransform: query="test query", documents=["doc1", "doc2"], top_n=2, - return_documents=True + return_documents=True, ) - + expected_params = { "query": "test query", "documents": ["doc1", "doc2"], "top_n": 2, - "return_documents": True + "return_documents": True, } assert params == expected_params @@ -363,34 +375,32 @@ class TestVertexAIRerankTransform: "documents": [ "This is a very long document with many words that should be truncated to only the first three words for the title", "Short doc", - "Another document with multiple words here" - ] + "Another document with multiple words here", + ], } - + request_data = self.config.transform_rerank_request( - model=self.model, - optional_rerank_params=optional_params, - headers={} + model=self.model, optional_rerank_params=optional_params, headers={} ) - + # Verify title generation assert request_data["records"][0]["title"] == "This is a" # First 3 words assert request_data["records"][1]["title"] == "Short doc" # Less than 3 words - assert request_data["records"][2]["title"] == "Another document with" # First 3 words + assert ( + request_data["records"][2]["title"] == "Another document with" + ) # First 3 words def test_record_id_generation(self): """Test that record IDs are generated correctly with 0-based indexing.""" optional_params = { "query": "test query", - "documents": ["doc1", "doc2", "doc3", "doc4"] + "documents": ["doc1", "doc2", "doc3", "doc4"], } - + request_data = self.config.transform_rerank_request( - model=self.model, - optional_rerank_params=optional_params, - headers={} + model=self.model, optional_rerank_params=optional_params, headers={} ) - + # Verify 0-based indexing for i, record in enumerate(request_data["records"]): assert record["id"] == str(i) @@ -402,9 +412,9 @@ class TestVertexAIRerankTransform: "documents": ["doc1", "doc2"], "vertex_credentials": "path/to/credentials.json", "vertex_project": "my-project-id", - "vertex_location": "us-central1" + "vertex_location": "us-central1", } - + params = self.config.map_cohere_rerank_params( non_default_params=non_default_params, model=self.model, @@ -412,14 +422,14 @@ class TestVertexAIRerankTransform: query="test query", documents=["doc1", "doc2"], top_n=2, - return_documents=True + return_documents=True, ) - + # Verify vertex-specific parameters are preserved assert params["vertex_credentials"] == "path/to/credentials.json" assert params["vertex_project"] == "my-project-id" assert params["vertex_location"] == "us-central1" - + # Verify standard params are still present assert params["query"] == "test query" assert params["documents"] == ["doc1", "doc2"] @@ -428,10 +438,8 @@ class TestVertexAIRerankTransform: def test_map_cohere_rerank_params_without_vertex_credentials(self): """Test that map_cohere_rerank_params works when vertex credentials are not provided.""" - non_default_params = { - "documents": ["doc1", "doc2"] - } - + non_default_params = {"documents": ["doc1", "doc2"]} + params = self.config.map_cohere_rerank_params( non_default_params=non_default_params, model=self.model, @@ -439,14 +447,14 @@ class TestVertexAIRerankTransform: query="test query", documents=["doc1", "doc2"], top_n=2, - return_documents=True + return_documents=True, ) - + # Verify no vertex-specific parameters are added when not provided assert "vertex_credentials" not in params assert "vertex_project" not in params assert "vertex_location" not in params - + # Verify standard params are still present assert params["query"] == "test query" assert params["documents"] == ["doc1", "doc2"] @@ -470,14 +478,11 @@ class TestVertexAIRerankTransform: "vertex_credentials": "path/to/credentials.json", "vertex_project": "custom-project-id", "query": "test query", - "documents": ["doc1"] + "documents": ["doc1"], } headers = self.config.validate_environment( - headers={}, - model=self.model, - api_key=None, - optional_params=optional_params + headers={}, model=self.model, api_key=None, optional_params=optional_params ) # Verify that _ensure_access_token was called with the credentials from optional_params @@ -490,7 +495,7 @@ class TestVertexAIRerankTransform: expected_headers = { "Authorization": "Bearer test-access-token", "Content-Type": "application/json", - "X-Goog-User-Project": "test-project-123" + "X-Goog-User-Project": "test-project-123", } assert headers == expected_headers @@ -527,7 +532,10 @@ class TestVertexAIRerankTransform: assert optional_params["vertex_project"] == "custom-project-id" # get_complete_url should still be able to access the vertex params - with patch('litellm.llms.vertex_ai.rerank.transformation.get_secret_str', return_value=None): + with patch( + "litellm.llms.vertex_ai.rerank.transformation.get_secret_str", + return_value=None, + ): url = self.config.get_complete_url( api_base=None, model=self.model, diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py index d483a81a349..8a5c397d6f4 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py @@ -15,7 +15,9 @@ from litellm.llms.vertex_ai.common_utils import ( convert_anyof_null_to_nullable, get_vertex_location_from_url, get_vertex_project_id_from_url, + pop_vertex_request_labels, set_schema_property_ordering, + vertex_request_labels_from_litellm_params, ) @@ -440,7 +442,9 @@ def test_vertex_ai_complex_response_schema(): optional_params = {} v.apply_response_schema_transformation( - value=non_default_params["response_format"], optional_params=optional_params, model="gemini-1.5-pro-preview-0409" + value=non_default_params["response_format"], + optional_params=optional_params, + model="gemini-1.5-pro-preview-0409", ) # Assertions for the transformed schema @@ -558,7 +562,6 @@ def test_get_vertex_url_global_region(stream, expected_endpoint_suffix): assert url == expected_url - @pytest.mark.parametrize( "model_cost_entry, vertex_region, expected_region", [ @@ -571,9 +574,17 @@ def test_get_vertex_url_global_region(stream, expected_endpoint_suffix): # Model with supported_regions=["us-west2"], no user region -> use "us-west2" ({"supported_regions": ["us-west2"]}, None, "us-west2"), # Model with supported_regions=["us-west2", "us-central1"], user passes supported region -> respect it - ({"supported_regions": ["us-west2", "us-central1"]}, "us-central1", "us-central1"), + ( + {"supported_regions": ["us-west2", "us-central1"]}, + "us-central1", + "us-central1", + ), # Model with supported_regions=["us-west2", "us-central1"], user passes unsupported region -> override - ({"supported_regions": ["us-west2", "us-central1"]}, "europe-west1", "us-west2"), + ( + {"supported_regions": ["us-west2", "us-central1"]}, + "europe-west1", + "us-west2", + ), # No model_cost entry, no user region -> default us-central1 ({}, None, "us-central1"), # No model_cost entry, user specifies region -> use specified region @@ -656,11 +667,12 @@ def test_vertex_filter_format_uri(): assert "uri" not in json.dumps(new_parameters) + def test_convert_schema_types_type_array_conversion(): """ Test _convert_schema_types function handles type arrays and case conversion. - - This test verifies the fix for the issue where type arrays like ["string", "number"] + + This test verifies the fix for the issue where type arrays like ["string", "number"] would raise an exception in Vertex AI schema validation. Relevant issue: https://github.com/BerriAI/litellm/issues/14091 @@ -673,12 +685,12 @@ def test_convert_schema_types_type_array_conversion(): "properties": { "studio": { "type": ["string", "number"], - "description": "The studio ID or name" + "description": "The studio ID or name", } }, "required": ["studio"], "additionalProperties": False, - "$schema": "http://json-schema.org/draft-07/schema#" + "$schema": "http://json-schema.org/draft-07/schema#", } # Expected output: Vertex AI compatible schema with anyOf and uppercase types @@ -686,16 +698,13 @@ def test_convert_schema_types_type_array_conversion(): "type": "object", "properties": { "studio": { - "anyOf": [ - {"type": "string"}, - {"type": "number"} - ], - "description": "The studio ID or name" + "anyOf": [{"type": "string"}, {"type": "number"}], + "description": "The studio ID or name", } }, "required": ["studio"], "additionalProperties": False, - "$schema": "http://json-schema.org/draft-07/schema#" + "$schema": "http://json-schema.org/draft-07/schema#", } # Apply the transformation @@ -718,15 +727,17 @@ def test_convert_schema_types_type_array_conversion(): assert anyof_types[1]["type"] == "number" # 4. Other properties preserved - assert input_schema["properties"]["studio"]["description"] == "The studio ID or name" + assert ( + input_schema["properties"]["studio"]["description"] == "The studio ID or name" + ) assert input_schema["required"] == ["studio"] def test_fix_enum_empty_strings(): """ Test _fix_enum_empty_strings function replaces empty strings with None in enum arrays. - - This test verifies the fix for the issue where Gemini rejects tool definitions + + This test verifies the fix for the issue where Gemini rejects tool definitions with empty strings in enum values, causing API failures. Relevant issue: Gemini does not accept empty strings in enum values @@ -740,23 +751,23 @@ def test_fix_enum_empty_strings(): "user_agent_type": { "enum": ["", "desktop", "mobile", "tablet"], "type": "string", - "description": "Device type for user agent" + "description": "Device type for user agent", } }, - "required": ["user_agent_type"] + "required": ["user_agent_type"], } # Expected output: Empty strings replaced with None expected_output = { - "type": "object", + "type": "object", "properties": { "user_agent_type": { "enum": [None, "desktop", "mobile", "tablet"], "type": "string", - "description": "Device type for user agent" + "description": "Device type for user agent", } }, - "required": ["user_agent_type"] + "required": ["user_agent_type"], } # Apply the transformation @@ -859,7 +870,7 @@ def test_construct_target_url_with_version_prefix(): def test_fix_enum_types(): """ Test _fix_enum_types function removes enum fields when type is not string. - + This test verifies the fix for the issue where Gemini rejects cached content with function parameter enums on non-string types, causing API failures. @@ -874,38 +885,41 @@ def test_fix_enum_types(): "truncateMode": { "enum": ["auto", "none", "start", "end"], "type": "string", # This should keep the enum - "description": "How to truncate content" + "description": "How to truncate content", }, "maxLength": { "enum": [100, 200, 500], # This should be removed "type": "integer", - "description": "Maximum length" + "description": "Maximum length", }, "enabled": { "enum": [True, False], # This should be removed "type": "boolean", - "description": "Whether feature is enabled" + "description": "Whether feature is enabled", }, "nested": { "type": "object", "properties": { "innerEnum": { "enum": ["a", "b", "c"], # This should be kept - "type": "string" + "type": "string", }, "innerNonStringEnum": { "enum": [1, 2, 3], # This should be removed - "type": "integer" - } - } + "type": "integer", + }, + }, }, "anyOfField": { "anyOf": [ - {"type": "string", "enum": ["option1", "option2"]}, # This should be kept - {"type": "integer", "enum": [1, 2, 3]} # This should be removed + { + "type": "string", + "enum": ["option1", "option2"], + }, # This should be kept + {"type": "integer", "enum": [1, 2, 3]}, # This should be removed ] - } - } + }, + }, } # Expected output: Non-string enums removed, string enums kept @@ -919,31 +933,32 @@ def test_fix_enum_types(): }, "maxLength": { # enum removed "type": "integer", - "description": "Maximum length" + "description": "Maximum length", }, "enabled": { # enum removed "type": "boolean", - "description": "Whether feature is enabled" + "description": "Whether feature is enabled", }, "nested": { "type": "object", "properties": { "innerEnum": { "enum": ["a", "b", "c"], # Kept - string type - "type": "string" + "type": "string", }, - "innerNonStringEnum": { # enum removed - "type": "integer" - } - } + "innerNonStringEnum": {"type": "integer"}, # enum removed + }, }, "anyOfField": { "anyOf": [ - {"type": "string", "enum": ["option1", "option2"]}, # Kept - has string type - {"type": "integer"} # enum removed + { + "type": "string", + "enum": ["option1", "option2"], + }, # Kept - has string type + {"type": "integer"}, # enum removed ] - } - } + }, + }, } # Apply the transformation @@ -955,15 +970,27 @@ def test_fix_enum_types(): # Verify specific transformations: # 1. String enums are preserved assert "enum" in input_schema["properties"]["truncateMode"] - assert input_schema["properties"]["truncateMode"]["enum"] == ["auto", "none", "start", "end"] - + assert input_schema["properties"]["truncateMode"]["enum"] == [ + "auto", + "none", + "start", + "end", + ] + assert "enum" in input_schema["properties"]["nested"]["properties"]["innerEnum"] - assert input_schema["properties"]["nested"]["properties"]["innerEnum"]["enum"] == ["a", "b", "c"] + assert input_schema["properties"]["nested"]["properties"]["innerEnum"]["enum"] == [ + "a", + "b", + "c", + ] # 2. Non-string enums are removed assert "enum" not in input_schema["properties"]["maxLength"] assert "enum" not in input_schema["properties"]["enabled"] - assert "enum" not in input_schema["properties"]["nested"]["properties"]["innerNonStringEnum"] + assert ( + "enum" + not in input_schema["properties"]["nested"]["properties"]["innerNonStringEnum"] + ) # 3. anyOf with string type keeps enum, non-string removes it assert "enum" in input_schema["properties"]["anyOfField"]["anyOf"][0] @@ -1003,8 +1030,6 @@ def test_get_token_url(): print("url=", url) - - should_use_v1beta1_features = vertex_llm.is_using_v1beta1_features( optional_params={"temperature": 0.1} ) @@ -1210,9 +1235,7 @@ def test_vertex_ai_minimax_uses_openai_handler(): VertexAIPartnerModels, ) - assert VertexAIPartnerModels.should_use_openai_handler( - "minimaxai/minimax-m2-maas" - ) + assert VertexAIPartnerModels.should_use_openai_handler("minimaxai/minimax-m2-maas") def test_vertex_ai_moonshot_uses_openai_handler(): @@ -1236,9 +1259,7 @@ def test_vertex_ai_zai_uses_openai_handler(): VertexAIPartnerModels, ) - assert VertexAIPartnerModels.should_use_openai_handler( - "zai-org/glm-4.7-maas" - ) + assert VertexAIPartnerModels.should_use_openai_handler("zai-org/glm-4.7-maas") def test_vertex_ai_zai_is_partner_model(): @@ -1255,14 +1276,14 @@ def test_vertex_ai_zai_is_partner_model(): def test_build_vertex_schema_empty_properties(): """ Test _build_vertex_schema handles empty properties objects correctly. - - This test verifies the fix for the issue where Gemini rejects schemas + + This test verifies the fix for the issue where Gemini rejects schemas with empty properties objects like {"properties": {}, "type": "object"}. - + Error from Gemini: "GenerateContentRequest.generation_config.response_schema - .properties[\"action\"].items.any_of[0].properties[\"go_back\"].properties: + .properties[\"action\"].items.any_of[0].properties[\"go_back\"].properties: should be non-empty for OBJECT type" - + The fix removes empty properties objects and their associated type/required fields. """ from litellm.llms.vertex_ai.common_utils import _build_vertex_schema @@ -1281,20 +1302,20 @@ def test_build_vertex_schema_empty_properties(): "type": "object", "additionalProperties": False, "description": "Go back", - "required": [] + "required": [], } }, "required": ["go_back"], "type": "object", - "additionalProperties": False + "additionalProperties": False, } ] }, - "type": "array" + "type": "array", } }, "type": "object", - "additionalProperties": False + "additionalProperties": False, } # Apply the transformation @@ -1302,24 +1323,36 @@ def test_build_vertex_schema_empty_properties(): # Verify the transformation removed empty properties # Navigate to the go_back schema - go_back_schema = result["properties"]["action"]["items"]["anyOf"][0]["properties"]["go_back"] - + go_back_schema = result["properties"]["action"]["items"]["anyOf"][0]["properties"][ + "go_back" + ] + # Verify empty properties was removed assert "properties" not in go_back_schema, "Empty properties should be removed" - + # Verify type is kept as object (Gemini requires type: object even without properties) - assert go_back_schema.get("type") == "object", "Type should be kept as object when properties is empty" - + assert ( + go_back_schema.get("type") == "object" + ), "Type should be kept as object when properties is empty" + # Verify required was also removed - assert "required" not in go_back_schema, "Required should be removed when properties is empty" - + assert ( + "required" not in go_back_schema + ), "Required should be removed when properties is empty" + # Verify description is preserved - assert go_back_schema.get("description") == "Go back", "Description should be preserved" - + assert ( + go_back_schema.get("description") == "Go back" + ), "Description should be preserved" + # Verify parent schema still has proper structure parent_schema = result["properties"]["action"]["items"]["anyOf"][0] - assert parent_schema["type"] == "object", "Parent schema should still have object type" - assert "go_back" in parent_schema["properties"], "go_back should still be in parent properties" + assert ( + parent_schema["type"] == "object" + ), "Parent schema should still have object type" + assert ( + "go_back" in parent_schema["properties"] + ), "go_back should still be in parent properties" def test_add_object_type_schema_with_no_properties_and_no_type(): @@ -1330,9 +1363,7 @@ def test_add_object_type_schema_with_no_properties_and_no_type(): from litellm.llms.vertex_ai.common_utils import add_object_type # Input: Schema with no properties and no type (the problematic case) - input_schema = { - "$schema": "https://json-schema.org/draft/2020-12/schema" - } + input_schema = {"$schema": "https://json-schema.org/draft/2020-12/schema"} # Apply the transformation add_object_type(input_schema) @@ -1351,10 +1382,7 @@ def test_add_object_type_does_not_override_existing_type(): from litellm.llms.vertex_ai.common_utils import add_object_type # Input: Schema with existing type - input_schema = { - "type": "string", - "description": "A string field" - } + input_schema = {"type": "string", "description": "A string field"} # Apply the transformation add_object_type(input_schema) @@ -1370,15 +1398,42 @@ def test_add_object_type_does_not_add_type_when_anyof_present(): from litellm.llms.vertex_ai.common_utils import add_object_type # Input: Schema with anyOf but no type - input_schema = { - "anyOf": [ - {"type": "string"}, - {"type": "null"} - ] - } + input_schema = {"anyOf": [{"type": "string"}, {"type": "null"}]} # Apply the transformation add_object_type(input_schema) # Verify type was not added (anyOf handles the type) assert "type" not in input_schema, "type should not be added when anyOf is present" + + +def test_vertex_request_labels_from_litellm_params_extracts_requester_metadata(): + assert vertex_request_labels_from_litellm_params(None) is None + assert vertex_request_labels_from_litellm_params({}) is None + assert vertex_request_labels_from_litellm_params({"metadata": None}) is None + lp = {"metadata": {"requester_metadata": {"team": "analytics", "count": 3}}} + assert vertex_request_labels_from_litellm_params(lp) == {"team": "analytics"} + + +def test_pop_vertex_request_labels_prefers_explicit_labels_then_metadata(): + optional = {"labels": {"env": "prod"}} + litellm_params = {"metadata": {"requester_metadata": {"team": "x"}}} + assert pop_vertex_request_labels(optional, litellm_params) == {"env": "prod"} + assert "labels" not in optional + + optional2: dict = {} + assert pop_vertex_request_labels(optional2, litellm_params) == {"team": "x"} + + +def test_vertex_text_embedding_request_includes_labels_from_metadata(): + import litellm + + req = litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request( + input="hi", + optional_params={}, + model="text-embedding-004", + litellm_params={ + "metadata": {"requester_metadata": {"project_id": "cost-center-1"}} + }, + ) + assert req.get("labels") == {"project_id": "cost-center-1"} From ebdc34e226b8399b93654f948d01070bebca99e4 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Sat, 11 Apr 2026 00:48:06 +0530 Subject: [PATCH 02/75] Fix mypy --- litellm/llms/vertex_ai/vertex_embeddings/transformation.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/llms/vertex_ai/vertex_embeddings/transformation.py b/litellm/llms/vertex_ai/vertex_embeddings/transformation.py index a3018c5bd6b..24396628dbd 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/transformation.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/transformation.py @@ -132,7 +132,7 @@ class VertexAITextEmbeddingConfig(BaseModel): vertex_request["labels"] = labels return vertex_request - vertex_request: VertexEmbeddingRequest = VertexEmbeddingRequest() + vertex_request = VertexEmbeddingRequest() vertex_text_embedding_input_list: List[TextEmbeddingInput] = [] task_type: Optional[TaskType] = optional_params.get("task_type") title = optional_params.get("title") From b19951a0d401a54c38dbb03a9ace03383f99dcf8 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 17 Apr 2026 18:32:33 +0530 Subject: [PATCH 03/75] Add litellm_metadata support --- litellm/llms/vertex_ai/common_utils.py | 35 +++++++++++-------- .../vertex_ai/test_vertex_ai_common_utils.py | 27 ++++++++++++++ 2 files changed, 48 insertions(+), 14 deletions(-) diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index a8cc5e7bdec..a89db0affdb 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -31,21 +31,28 @@ def vertex_request_labels_from_litellm_params( litellm_params: Optional[dict], ) -> Optional[Dict[str, str]]: """ - Build Vertex/GCP billing labels from LiteLLM ``litellm_params["metadata"]``, + Build Vertex/GCP billing labels from LiteLLM user metadata on ``litellm_params``: + ``metadata`` (``completion(..., metadata=...)``) or ``litellm_metadata``, using ``requester_metadata`` string key-value pairs (same convention as Gemini). + ``metadata`` is tried first when both are present. """ - if not litellm_params or "metadata" not in litellm_params: + if not litellm_params: return None - metadata = litellm_params["metadata"] - if metadata is None or not isinstance(metadata, dict): - return None - if "requester_metadata" not in metadata: - return None - rm = metadata["requester_metadata"] - if not isinstance(rm, dict): - return None - labels = {k: v for k, v in rm.items() if isinstance(v, str)} - return labels if labels else None + for key in ("metadata", "litellm_metadata"): + if key not in litellm_params: + continue + metadata = litellm_params[key] + if metadata is None or not isinstance(metadata, dict): + continue + if "requester_metadata" not in metadata: + continue + rm = metadata["requester_metadata"] + if not isinstance(rm, dict): + continue + labels = {k: v for k, v in rm.items() if isinstance(v, str)} + if labels: + return labels + return None def pop_vertex_request_labels( @@ -54,8 +61,8 @@ def pop_vertex_request_labels( ) -> Optional[Dict[str, str]]: """ Resolve labels from optional ``labels`` (Gemini-style) and/or - ``litellm_params["metadata"]["requester_metadata"]``. Pops ``labels`` from - optional_params when present. + ``litellm_params["metadata"]`` / ``litellm_params["litellm_metadata"]`` + (``requester_metadata``). Pops ``labels`` from optional_params when present. """ labels: Optional[Dict[str, str]] = None if optional_params is not None and "labels" in optional_params: diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py index 8a5c397d6f4..74aaa4d42b6 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py @@ -1415,6 +1415,23 @@ def test_vertex_request_labels_from_litellm_params_extracts_requester_metadata() assert vertex_request_labels_from_litellm_params(lp) == {"team": "analytics"} +def test_vertex_request_labels_from_litellm_params_accepts_litellm_metadata(): + lp = { + "litellm_metadata": { + "requester_metadata": {"team": "platform", "count": 3} + } + } + assert vertex_request_labels_from_litellm_params(lp) == {"team": "platform"} + + +def test_vertex_request_labels_prefers_metadata_over_litellm_metadata(): + lp = { + "metadata": {"requester_metadata": {"source": "metadata"}}, + "litellm_metadata": {"requester_metadata": {"source": "litellm_metadata"}}, + } + assert vertex_request_labels_from_litellm_params(lp) == {"source": "metadata"} + + def test_pop_vertex_request_labels_prefers_explicit_labels_then_metadata(): optional = {"labels": {"env": "prod"}} litellm_params = {"metadata": {"requester_metadata": {"team": "x"}}} @@ -1425,6 +1442,16 @@ def test_pop_vertex_request_labels_prefers_explicit_labels_then_metadata(): assert pop_vertex_request_labels(optional2, litellm_params) == {"team": "x"} +def test_pop_vertex_request_labels_uses_litellm_metadata_when_metadata_absent(): + optional: dict = {} + litellm_params = { + "litellm_metadata": {"requester_metadata": {"team": "from_litellm_meta"}} + } + assert pop_vertex_request_labels(optional, litellm_params) == { + "team": "from_litellm_meta" + } + + def test_vertex_text_embedding_request_includes_labels_from_metadata(): import litellm From 7020b0a64b17aacc4a8971e3822be94a2dc3610d Mon Sep 17 00:00:00 2001 From: shivam Date: Fri, 17 Apr 2026 13:32:05 -0700 Subject: [PATCH 04/75] fix(proxy): replay ASGI receive after metrics auth to avoid /metrics hang --- .../middleware/prometheus_auth_middleware.py | 27 ++++++++++++++-- .../test_prometheus_auth_middleware.py | 31 +++++++++++++++++++ 2 files changed, 56 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/middleware/prometheus_auth_middleware.py b/litellm/proxy/middleware/prometheus_auth_middleware.py index 5915e4aa07d..d0e38f98946 100644 --- a/litellm/proxy/middleware/prometheus_auth_middleware.py +++ b/litellm/proxy/middleware/prometheus_auth_middleware.py @@ -2,6 +2,7 @@ Prometheus Auth Middleware - Pure ASGI implementation """ import json +from typing import List from fastapi import Request from starlette.types import ASGIApp, Receive, Scope, Send @@ -39,8 +40,17 @@ class PrometheusAuthMiddleware: # Only run auth if configured to do so if litellm.require_auth_for_metrics_endpoint is True: - # Construct Request only when auth is actually needed - request = Request(scope, receive) + # user_api_key_auth reads the request body, which consumes ASGI `receive`. + # Buffer those messages and replay them for the inner app; otherwise a + # successful auth would forward an exhausted receive and /metrics hangs. + buffered_messages: List[dict] = [] + + async def receive_for_auth() -> dict: + message = await receive() + buffered_messages.append(message) + return message + + request = Request(scope, receive_for_auth) api_key = request.headers.get(_AUTHORIZATION_HEADER) or "" try: @@ -69,5 +79,18 @@ class PrometheusAuthMiddleware: ) return + replay_idx = 0 + + async def receive_replay() -> dict: + nonlocal replay_idx + if replay_idx < len(buffered_messages): + msg = buffered_messages[replay_idx] + replay_idx += 1 + return msg + return await receive() + + await self.app(scope, receive_replay, send) + return + # Pass through to the inner application await self.app(scope, receive, send) diff --git a/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py b/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py index 9fd244d9c3f..310ee11573b 100644 --- a/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py +++ b/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py @@ -26,6 +26,15 @@ async def fake_valid_auth(request, api_key): return +async def fake_valid_auth_reads_body(request, api_key, **kwargs): + """ + Like real user_api_key_auth, consumes the ASGI body stream. Regression test + for successful auth passing a drained receive to the inner app (hang). + """ + await request.body() + return + + async def fake_invalid_auth(request, api_key): print("running fake invalid auth", request, api_key) # Simulate invalid auth by raising an exception. @@ -62,6 +71,28 @@ def app_with_middleware(): return app +def test_valid_auth_metrics_after_body_consumed(app_with_middleware, monkeypatch): + """ + Auth that reads the request body must not cause /metrics to hang on success. + """ + litellm.require_auth_for_metrics_endpoint = True + monkeypatch.setattr( + "litellm.proxy.middleware.prometheus_auth_middleware.user_api_key_auth", + fake_valid_auth_reads_body, + ) + + client = TestClient(app_with_middleware) + headers = {SpecialHeaders.openai_authorization.value: "valid"} + + response = client.get("/metrics", headers=headers) + assert response.status_code == 200, response.text + assert response.json() == {"msg": "metrics OK"} + + response = client.get("/metrics/", headers=headers) + assert response.status_code == 200, response.text + assert response.json() == {"msg": "metrics OK"} + + def test_valid_auth_metrics(app_with_middleware, monkeypatch): """ Test that a request to /metrics (and /metrics/) with valid auth headers passes. From 733ccd6e0920252a99fdbe8a5bad892503026499 Mon Sep 17 00:00:00 2001 From: shivam Date: Fri, 17 Apr 2026 16:36:34 -0700 Subject: [PATCH 05/75] fixed linting --- litellm/proxy/middleware/prometheus_auth_middleware.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/middleware/prometheus_auth_middleware.py b/litellm/proxy/middleware/prometheus_auth_middleware.py index d0e38f98946..d5dec6c4bda 100644 --- a/litellm/proxy/middleware/prometheus_auth_middleware.py +++ b/litellm/proxy/middleware/prometheus_auth_middleware.py @@ -2,7 +2,7 @@ Prometheus Auth Middleware - Pure ASGI implementation """ import json -from typing import List +from typing import Any, List, MutableMapping from fastapi import Request from starlette.types import ASGIApp, Receive, Scope, Send @@ -43,9 +43,9 @@ class PrometheusAuthMiddleware: # user_api_key_auth reads the request body, which consumes ASGI `receive`. # Buffer those messages and replay them for the inner app; otherwise a # successful auth would forward an exhausted receive and /metrics hangs. - buffered_messages: List[dict] = [] + buffered_messages: List[MutableMapping[str, Any]] = [] - async def receive_for_auth() -> dict: + async def receive_for_auth() -> MutableMapping[str, Any]: message = await receive() buffered_messages.append(message) return message @@ -81,7 +81,7 @@ class PrometheusAuthMiddleware: replay_idx = 0 - async def receive_replay() -> dict: + async def receive_replay() -> MutableMapping[str, Any]: nonlocal replay_idx if replay_idx < len(buffered_messages): msg = buffered_messages[replay_idx] From 49ddb0e5bad5d2f3dbb2f5bb4014ece36c295eef Mon Sep 17 00:00:00 2001 From: shivam Date: Fri, 17 Apr 2026 16:43:01 -0700 Subject: [PATCH 06/75] style: apply Black formatting to Prometheus integration modules Reformat prometheus logger, types, and metrics auth middleware so lint CI (black --check) passes. Made-with: Cursor --- litellm/integrations/prometheus.py | 46 +++++++++++-------- .../middleware/prometheus_auth_middleware.py | 1 + litellm/types/integrations/prometheus.py | 12 ++--- 3 files changed, 33 insertions(+), 26 deletions(-) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index b3bf792e93b..d5fde2a4864 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -88,7 +88,9 @@ class PrometheusLogger(CustomLogger): _custom_buckets = litellm.prometheus_latency_buckets self.latency_buckets = ( - tuple(_custom_buckets) if _custom_buckets is not None else LATENCY_BUCKETS + tuple(_custom_buckets) + if _custom_buckets is not None + else LATENCY_BUCKETS ) # Create metric factory functions @@ -1097,9 +1099,11 @@ class PrometheusLogger(CustomLogger): ), client_ip=standard_logging_payload["metadata"].get("requester_ip_address"), user_agent=standard_logging_payload["metadata"].get("user_agent"), - stream=str(standard_logging_payload.get("stream")) - if litellm.prometheus_emit_stream_label - else None, + stream=( + str(standard_logging_payload.get("stream")) + if litellm.prometheus_emit_stream_label + else None + ), ) if ( @@ -1767,9 +1771,11 @@ class PrometheusLogger(CustomLogger): client_ip=_metadata.get("requester_ip_address"), user_agent=_metadata.get("user_agent"), model_id=model_id, - stream=str(request_data.get("stream")) - if litellm.prometheus_emit_stream_label - else None, + stream=( + str(request_data.get("stream")) + if litellm.prometheus_emit_stream_label + else None + ), ) _labels = prometheus_label_factory( supported_enum_labels=self.get_labels_for_metric( @@ -2093,9 +2099,9 @@ class PrometheusLogger(CustomLogger): ): try: verbose_logger.debug("setting remaining tokens requests metric") - standard_logging_payload: Optional[ - StandardLoggingPayload - ] = request_kwargs.get("standard_logging_object") + standard_logging_payload: Optional[StandardLoggingPayload] = ( + request_kwargs.get("standard_logging_object") + ) if standard_logging_payload is None: return @@ -2728,9 +2734,7 @@ class PrometheusLogger(CustomLogger): ) return - async def fetch_keys( - page_size: int, page: int - ) -> Tuple[ + async def fetch_keys(page_size: int, page: int) -> Tuple[ List[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]], Optional[int], ]: @@ -2921,9 +2925,11 @@ class PrometheusLogger(CustomLogger): org_alias=org.organization_alias or "", spend=org.spend or 0.0, max_budget=budget_table.max_budget if budget_table else None, - budget_reset_at=getattr(budget_table, "budget_reset_at", None) - if budget_table - else None, + budget_reset_at=( + getattr(budget_table, "budget_reset_at", None) + if budget_table + else None + ), ) async def _set_team_budget_metrics_after_api_request( @@ -3405,10 +3411,10 @@ class PrometheusLogger(CustomLogger): from litellm.constants import PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES from litellm.integrations.custom_logger import CustomLogger - prometheus_loggers: List[ - CustomLogger - ] = litellm.logging_callback_manager.get_custom_loggers_for_type( - callback_type=PrometheusLogger + prometheus_loggers: List[CustomLogger] = ( + litellm.logging_callback_manager.get_custom_loggers_for_type( + callback_type=PrometheusLogger + ) ) # we need to get the initialized prometheus logger instance(s) and call logger.initialize_remaining_budget_metrics() on them verbose_logger.debug("found %s prometheus loggers", len(prometheus_loggers)) diff --git a/litellm/proxy/middleware/prometheus_auth_middleware.py b/litellm/proxy/middleware/prometheus_auth_middleware.py index d5dec6c4bda..3b30fd3d63c 100644 --- a/litellm/proxy/middleware/prometheus_auth_middleware.py +++ b/litellm/proxy/middleware/prometheus_auth_middleware.py @@ -1,6 +1,7 @@ """ Prometheus Auth Middleware - Pure ASGI implementation """ + import json from typing import Any, List, MutableMapping diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 51a41f97e03..1b36ad5daa6 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -676,9 +676,9 @@ class PrometheusMetricLabels: litellm_managed_batch_created_total = _batch_user_labels - litellm_managed_file_size_bytes: List[ - str - ] = [] # labels: purpose, file_type, model, api_provider, user (custom) + litellm_managed_file_size_bytes: List[str] = ( + [] + ) # labels: purpose, file_type, model, api_provider, user (custom) litellm_managed_batch_duration_seconds = [ UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, @@ -687,9 +687,9 @@ class PrometheusMetricLabels: litellm_managed_file_created_total = _batch_user_labels - litellm_managed_file_deleted_total: List[ - str - ] = [] # only "result" label, added at metric creation + litellm_managed_file_deleted_total: List[str] = ( + [] + ) # only "result" label, added at metric creation litellm_check_batch_cost_jobs_polled: List[str] = [] From 1170bd55c97612fe996de924b6ad95c87bb229ab Mon Sep 17 00:00:00 2001 From: shivam Date: Fri, 17 Apr 2026 16:58:25 -0700 Subject: [PATCH 07/75] fix(prometheus): remove dead prometheus_label_factory call; apply Black async_post_call_failure_hook passed only supported_enum_labels to prometheus_label_factory; enum_values is required. The call was unused because _inc_labeled_counter builds labels internally. Reformat Prometheus-related modules and passthrough/copilot helpers for CI. Made-with: Cursor --- litellm/integrations/prometheus.py | 18 ++++++++---------- litellm/integrations/prometheus_helpers.py | 3 +-- litellm/llms/github_copilot/authenticator.py | 4 +--- litellm/passthrough/utils.py | 4 +++- litellm/types/integrations/prometheus.py | 2 +- 5 files changed, 14 insertions(+), 17 deletions(-) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index f4fa1e6c036..723b142dfad 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -51,6 +51,7 @@ if TYPE_CHECKING: else: AsyncIOScheduler = Any + class PrometheusLogger(CustomLogger): # Class variables or attributes @@ -991,9 +992,7 @@ class PrometheusLogger(CustomLogger): amount: float = 1.0, ) -> None: _labels = prometheus_label_factory( - supported_enum_labels=self.get_labels_for_metric( - metric_name=metric_name - ), + supported_enum_labels=self.get_labels_for_metric(metric_name=metric_name), enum_values=enum_values, label_context=label_context, ) @@ -1118,7 +1117,9 @@ class PrometheusLogger(CustomLogger): user_api_key = hash_token(user_api_key) - label_context = PrometheusLabelFactoryContext(enum_values) #amortized per request. + label_context = PrometheusLabelFactoryContext( + enum_values + ) # amortized per request. # increment total LLM requests and spend metric self._increment_top_level_request_and_spend_metrics( @@ -1791,11 +1792,6 @@ class PrometheusLogger(CustomLogger): else None ), ) - _labels = prometheus_label_factory( - supported_enum_labels=self.get_labels_for_metric( - metric_name="litellm_proxy_failed_requests_metric" - ), - ) _label_ctx = PrometheusLabelFactoryContext(enum_values) PrometheusLogger._inc_labeled_counter( self, @@ -3495,7 +3491,9 @@ def _prometheus_labels_from_context( } if UserAPIKeyLabelNames.END_USER.value in filtered_labels: - filtered_labels[UserAPIKeyLabelNames.END_USER.value] = ctx.get_resolved_end_user() + filtered_labels[UserAPIKeyLabelNames.END_USER.value] = ( + ctx.get_resolved_end_user() + ) for sk, val in ctx._custom_by_sanitized_key.items(): if sk in supported_enum_labels: diff --git a/litellm/integrations/prometheus_helpers.py b/litellm/integrations/prometheus_helpers.py index 34f4855863e..784ab524dd5 100644 --- a/litellm/integrations/prometheus_helpers.py +++ b/litellm/integrations/prometheus_helpers.py @@ -51,8 +51,7 @@ class PrometheusLabelFactoryContext: self.enum_values = enum_values enum_dict = enum_values.model_dump() self._sanitized_enum: Dict[str, Optional[str]] = { - k: _sanitize_prometheus_label_value(v) - for k, v in enum_dict.items() + k: _sanitize_prometheus_label_value(v) for k, v in enum_dict.items() } self._custom_by_sanitized_key: Dict[str, Optional[str]] = {} if enum_values.custom_metadata_labels is not None: diff --git a/litellm/llms/github_copilot/authenticator.py b/litellm/llms/github_copilot/authenticator.py index f4698861edc..9de2987b9f6 100644 --- a/litellm/llms/github_copilot/authenticator.py +++ b/litellm/llms/github_copilot/authenticator.py @@ -294,9 +294,7 @@ class Authenticator: access_token_url = os.getenv( "GITHUB_COPILOT_ACCESS_TOKEN_URL", DEFAULT_GITHUB_ACCESS_TOKEN_URL ) - client_id = os.getenv( - "GITHUB_COPILOT_CLIENT_ID", DEFAULT_GITHUB_CLIENT_ID - ) + client_id = os.getenv("GITHUB_COPILOT_CLIENT_ID", DEFAULT_GITHUB_CLIENT_ID) for attempt in range(max_attempts): try: diff --git a/litellm/passthrough/utils.py b/litellm/passthrough/utils.py index 5dde13f0078..d39a0dda152 100644 --- a/litellm/passthrough/utils.py +++ b/litellm/passthrough/utils.py @@ -79,7 +79,9 @@ class BasePassthroughUtils: for header_name, header_value in request_headers.items(): if header_name.lower().startswith(PASS_THROUGH_HEADER_PREFIX): # Strip the 'x-pass-' prefix and normalize to lowercase - actual_header_name = header_name[len(PASS_THROUGH_HEADER_PREFIX) :].lower() + actual_header_name = header_name[ + len(PASS_THROUGH_HEADER_PREFIX) : + ].lower() if actual_header_name in _PASS_THROUGH_PROTECTED_HEADERS or any( actual_header_name.startswith(p) for p in _PASS_THROUGH_PROTECTED_HEADER_PREFIXES diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 338c5a79ce6..43a287f29bc 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -784,7 +784,7 @@ class UserAPIKeyLabelValues: org_id: Optional[str] = None org_alias: Optional[str] = None - #Added for test compatibility. + # Added for test compatibility. def __init__(self, **kwargs: Any) -> None: """ Match former Pydantic behavior: unknown keys are ignored; ``api_key_hash`` maps to From 30885467ff0e8cb5a609744a521f5c55acd2a554 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Tue, 21 Apr 2026 19:57:02 +0000 Subject: [PATCH 08/75] Add debugger settings to debug single worker proxy_server per request. --- .gitignore | 1 - .vscode/launch.json | 24 ++++++++++++++++++++++++ 2 files changed, 24 insertions(+), 1 deletion(-) create mode 100644 .vscode/launch.json diff --git a/.gitignore b/.gitignore index 38bf9554b5b..f5591b7f2c6 100644 --- a/.gitignore +++ b/.gitignore @@ -90,7 +90,6 @@ test.py litellm_config.yaml !.github/observatory/litellm_config.yaml .cursor -.vscode/launch.json litellm/proxy/to_delete_loadtest_work/* update_model_cost_map.py tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py diff --git a/.vscode/launch.json b/.vscode/launch.json new file mode 100644 index 00000000000..e48948319b0 --- /dev/null +++ b/.vscode/launch.json @@ -0,0 +1,24 @@ +{ + "version": "0.2.0", + "configurations": [ + { + "name": "LiteLLM proxy (debug)", + "type": "debugpy", + "request": "launch", + "program": "${workspaceFolder}/litellm/proxy/proxy_cli.py", + "cwd": "${workspaceFolder}", + "args": [ + "--host", + "0.0.0.0", + "--config", + "dev_config.yaml", + "--port", + "4000", + "--num_workers", + "1" + ], + "console": "integratedTerminal", + "justMyCode": false + } + ] +} \ No newline at end of file From d58f657fa281b72b27d0c2596d7e710a673e42c0 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Tue, 21 Apr 2026 20:12:42 +0000 Subject: [PATCH 09/75] Replace assertions with simple if conditions. assertions raise an exception which are not great for performance (specifically repeated throws). Logic remains the same and the tests are still passing --- litellm/proxy/auth/user_api_key_auth.py | 18 ++++++++++-------- 1 file changed, 10 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index b6366e41d5d..205e6176750 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1115,9 +1115,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 is_master_key_valid = False ## VALIDATE MASTER KEY ## - try: - assert isinstance(master_key, str) - except Exception: + if not isinstance(master_key, str): raise HTTPException( status_code=500, detail={ @@ -1183,11 +1181,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 if len(api_key) > 8 else "****" ) - assert api_key.startswith( - "sk-" - ), "LiteLLM Virtual Key expected. Received={}, expected to start with 'sk-'.".format( - _masked_key - ) # prevent token hashes from being used + if not api_key.startswith("sk-"): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=( + "LiteLLM Virtual Key expected. Received={}, expected to start with 'sk-'.".format( + _masked_key + ) + ), + ) # prevent token hashes from being used else: verbose_logger.warning( "litellm.proxy.proxy_server.user_api_key_auth(): Warning - Key is not a string. Got type={}".format( From 984287daaaf462cff6b4f3b0a861b69f0bd34965 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Tue, 21 Apr 2026 22:53:54 +0000 Subject: [PATCH 10/75] Implement CacheCodec for DualCache serialization and deserialization. Add attach_redis_cache method to DualCache for lazy Redis integration. Update RedisCache to handle None keys and improve logging. Enhance user_api_key_auth caching logic and introduce tests for CacheCodec functionality. --- litellm/caching/dual_cache.py | 19 ++++ litellm/caching/redis_cache.py | 10 ++- litellm/proxy/auth/auth_checks.py | 34 +++++--- litellm/proxy/auth/user_api_key_auth.py | 8 +- .../common_utils/cache_pydantic_utils.py | 86 +++++++++++++++++++ litellm/proxy/proxy_server.py | 43 ++++++++-- tests/local_testing/test_dual_cache.py | 69 +++++++++++++++ .../proxy/common_utils/test_cache_codec.py | 85 ++++++++++++++++++ 8 files changed, 331 insertions(+), 23 deletions(-) create mode 100644 litellm/proxy/common_utils/cache_pydantic_utils.py create mode 100644 tests/test_litellm/proxy/common_utils/test_cache_codec.py diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index 34ae3638a5b..7436f95f5b4 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -92,6 +92,25 @@ class DualCache(BaseCache): if default_redis_ttl is not None: self.default_redis_ttl = default_redis_ttl + def attach_redis_cache( + self, + redis_cache: Optional[RedisCache] = None, + *, + default_redis_ttl: Optional[float] = None, + ) -> None: + """ + Attach a Redis backend if this DualCache does not already have one. + + No-op when ``redis_cache`` is None or when Redis was already set (constructor + or a prior attach). Use this for lazy wiring after a shared Redis client exists. + Does not backfill in-memory-only keys to Redis. + """ + if redis_cache is None or self.redis_cache is not None: + return + self.redis_cache = redis_cache + if default_redis_ttl is not None: + self.default_redis_ttl = default_redis_ttl + def set_cache(self, key, value, local_only: bool = False, **kwargs): # Update both Redis and in-memory cache try: diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 84a2887f527..f8abb9e95c8 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -551,6 +551,13 @@ class RedisCache(BaseCache): async def async_set_cache(self, key, value, **kwargs): from redis.asyncio import Redis + if key is None: + verbose_logger.debug( + "LiteLLM Redis Caching: async set() skipped — key is None, value=%r", + value, + ) + return None + start_time = time.time() try: _redis_client: Redis = self.init_async_client() # type: ignore @@ -569,8 +576,9 @@ class RedisCache(BaseCache): ) ) verbose_logger.error( - "LiteLLM Redis Caching: async set() - Got exception from REDIS %s, Writing value=%s", + "LiteLLM Redis Caching: async set() - Got exception from REDIS %s, key=%r, value=%r", str(e), + key, value, ) raise e diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 290b34be836..1577b5bda9c 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -12,7 +12,7 @@ Run checks for: import asyncio import re import time -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type, Union, cast from fastapi import HTTPException, Request, status from pydantic import BaseModel @@ -65,6 +65,7 @@ from litellm.proxy.guardrails.tool_name_extraction import ( TOOL_CAPABLE_CALL_TYPES, extract_request_tool_names, ) +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec from litellm.proxy.route_llm_request import route_request from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics from litellm.router import Router @@ -1435,13 +1436,23 @@ async def get_user_object( async def _cache_management_object( key: str, - value: BaseModel, + value: Union[BaseModel, Dict[str, Any]], user_api_key_cache: DualCache, proxy_logging_obj: Optional[ProxyLogging], + *, + model_type: Type[BaseModel], ): + """ + Persist management objects to DualCache (in-memory + optional Redis). + + Values must be JSON-serializable for the Redis path (``json.dumps``). Payloads are + encoded with ``CacheCodec.serialize(..., model_type=...)`` so writes match reads + via ``CacheCodec.deserialize(..., model_type)``. + """ + cache_payload = CacheCodec.serialize(value, model_type=model_type) await user_api_key_cache.async_set_cache( key=key, - value=value, + value=cache_payload, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) @@ -1462,6 +1473,7 @@ async def _cache_team_object( value=team_table, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + model_type=LiteLLM_TeamTableCachedObj, ) @@ -1481,6 +1493,7 @@ async def _cache_key_object( value=user_api_key_obj, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + model_type=UserAPIKeyAuth, ) @@ -2196,15 +2209,13 @@ async def get_key_object( # check if in cache key = hashed_token - cached_key_obj: Optional[UserAPIKeyAuth] = await user_api_key_cache.async_get_cache( - key=key - ) - + # Same flow as before: use cache only when we have a hit we can turn into UserAPIKeyAuth + # (dict from Redis / model_dump, or UserAPIKeyAuth from in-memory). Otherwise fall through to DB. + cached_key_obj = await user_api_key_cache.async_get_cache(key=key) if cached_key_obj is not None: - if isinstance(cached_key_obj, dict): - return UserAPIKeyAuth(**cached_key_obj) - elif isinstance(cached_key_obj, UserAPIKeyAuth): - return cached_key_obj + user_api_key_auth = CacheCodec.deserialize(cached_key_obj, UserAPIKeyAuth) + if user_api_key_auth is not None: + return user_api_key_auth if check_cache_only: raise Exception( @@ -3465,6 +3476,7 @@ async def get_project_object( value=project_obj, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + model_type=LiteLLM_ProjectTableCachedObj, ) return project_obj diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 205e6176750..a6c8fd02ee1 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1452,9 +1452,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 else: _team_obj = None - await user_api_key_cache.async_set_cache( - key=valid_token.team_id, value=_team_obj - ) # save team table in cache - used for tpm/rpm limiting - tpm_rpm_limiter.py + # Only cache when the key is a real team_id (non-team keys must not use key=None). + if valid_token.team_id is not None and _team_obj is not None: + await user_api_key_cache.async_set_cache( + key=valid_token.team_id, value=_team_obj + ) # save team table in cache - used for tpm/rpm limiting - tpm_rpm_limiter.py # Fetch project object if key belongs to a project _project_obj = None diff --git a/litellm/proxy/common_utils/cache_pydantic_utils.py b/litellm/proxy/common_utils/cache_pydantic_utils.py new file mode 100644 index 00000000000..af0c3eef18d --- /dev/null +++ b/litellm/proxy/common_utils/cache_pydantic_utils.py @@ -0,0 +1,86 @@ +""" +DualCache presents a single API for reads and writes, but the two backends behave +differently: the in-memory layer can store arbitrary Python objects (including live +``BaseModel`` instances), while Redis persists strings and therefore needs JSON-safe +payloads (``json.dumps`` on the Redis side). + +Call sites therefore see cache ``value`` / ``cached`` as effectively ``Any``: the same +key may deserialize to a model on one process (memory hit) or to a ``dict`` after a +Redis round-trip. ``CacheCodec`` centralizes encode/decode at that boundary: +``CacheCodec.serialize`` before ``set``, ``CacheCodec.deserialize`` after ``get`` +when you need a typed ``BaseModel``. + +``dataclasses`` are not supported: only ``dict`` and Pydantic ``BaseModel`` inputs +are encoded; pass a Pydantic model or convert with e.g. ``dataclasses.asdict`` first. +""" + +from __future__ import annotations + +from typing import Any, Optional, Type, TypeVar + +from pydantic import BaseModel, ValidationError + +from litellm._logging import verbose_proxy_logger + +T = TypeVar("T", bound=BaseModel) + + +class CacheCodec: + """ + Encode/decode Pydantic models for DualCache (memory vs Redis safe payloads). + + Dataclasses are not supported yet (only ``dict`` and ``BaseModel``). + + Use ``serialize`` with ``model_type`` when writing so the same schema is used + as on read (``deserialize``). Pass ``model_type`` whenever you know it + (validates ``dict`` payloads and normalizes ``BaseModel`` instances). + """ + + @staticmethod + def serialize(value: Any, model_type: Optional[Type[T]] = None) -> Any: + """ + Encode a value for DualCache / Redis (``json.dumps``-safe). + + If ``model_type`` is set, the payload is validated with that model, then + ``model_dump(mode="json", exclude_none=True)`` — symmetric with ``deserialize``. + + If ``model_type`` is omitted, any ``BaseModel`` is dumped as above; other + values (e.g. plain ``dict``) are returned unchanged. + """ + if model_type is not None: + if isinstance(value, (dict, BaseModel)): + return model_type.model_validate(value).model_dump( + mode="json", exclude_none=True + ) + return value + if isinstance(value, BaseModel): + return value.model_dump(mode="json", exclude_none=True) + return value + + @staticmethod + def deserialize(cached: Any, model_type: Type[T]) -> Optional[T]: + """ + Decode a cache entry to ``model_type``. + + - ``None`` → ``None`` + - Already an instance of ``model_type`` (including subclasses) → returned as-is + - ``dict`` → ``model_type.model_validate(...)``; on ``ValidationError``, + logs a warning and returns ``None`` (treat as cache miss; avoids serving + malformed or schema-drifted entries) + - Any other type → ``None`` (caller should treat as cache miss or log) + """ + if cached is None: + return None + if isinstance(cached, model_type): + return cached + if isinstance(cached, dict): + try: + return model_type.model_validate(cached) + except ValidationError as e: + verbose_proxy_logger.warning( + "CacheCodec.deserialize: validation failed for %s (%s)", + model_type.__name__, + e, + ) + return None + return None diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f3ff3f5e23f..72ec41ef97f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -93,6 +93,7 @@ from litellm.proxy._types import ( TransformRequestBody, UserAPIKeyAuth, ) +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec from litellm.proxy.common_utils.callback_utils import ( normalize_callback_names, process_callback, @@ -1943,14 +1944,22 @@ async def update_cache( # noqa: PLR0915 else: hashed_token = token verbose_proxy_logger.debug("_update_key_cache: hashed_token=%s", hashed_token) - existing_spend_obj: LiteLLM_VerificationTokenView = await user_api_key_cache.async_get_cache(key=hashed_token) # type: ignore + cached_key = await user_api_key_cache.async_get_cache(key=hashed_token) verbose_proxy_logger.debug( - f"_update_key_cache: existing_spend_obj={existing_spend_obj}" + f"_update_key_cache: existing_spend_obj={cached_key}" ) - if existing_spend_obj is None: + if cached_key is None: return - else: - existing_spend = existing_spend_obj.spend + existing_spend_obj = CacheCodec.deserialize(cached_key, UserAPIKeyAuth) + if existing_spend_obj is None: + verbose_proxy_logger.warning( + "_update_key_cache: unexpected cached key type %s for hashed_token=%s; skipping spend update", + type(cached_key).__name__, + hashed_token, + ) + return + + existing_spend = existing_spend_obj.spend or 0.0 # Calculate the new cost by adding the existing cost and response_cost new_spend = existing_spend + response_cost @@ -2008,9 +2017,14 @@ async def update_cache( # noqa: PLR0915 existing_team_member_spend + response_cost ) - # Update the cost column for the given token + # Update the cost column for the given token (dict for Redis pipeline json.dumps) existing_spend_obj.spend = new_spend - values_to_update_in_cache.append((hashed_token, existing_spend_obj)) + values_to_update_in_cache.append( + ( + hashed_token, + CacheCodec.serialize(existing_spend_obj, model_type=UserAPIKeyAuth), + ) + ) ### UPDATE USER SPEND ### async def _update_user_cache(): @@ -2851,9 +2865,22 @@ class ProxyConfig: ): ## INIT PROXY REDIS USAGE CLIENT ## redis_usage_cache = litellm.cache.cache - spend_counter_cache.redis_cache = redis_usage_cache + spend_counter_cache.attach_redis_cache( + redis_usage_cache, + default_redis_ttl=litellm.default_redis_ttl, + ) # Note: PKCE verifier storage uses redis_usage_cache directly (not # user_api_key_cache) to avoid routing all API-key lookups through Redis. + # Share the same Redis client for virtual-key lookups (same DualCache as + # model_max_budget_limiter). attach_redis_cache is a no-op if Redis is + # already set (e.g. config reload). + user_api_key_cache.attach_redis_cache( + redis_usage_cache, + default_redis_ttl=litellm.default_redis_ttl, + ) + verbose_proxy_logger.debug( + "Attached redis_usage_cache Redis client to user_api_key_cache" + ) def switch_on_llm_response_caching(self): """ diff --git a/tests/local_testing/test_dual_cache.py b/tests/local_testing/test_dual_cache.py index 5a1cdf86487..30f45033c58 100644 --- a/tests/local_testing/test_dual_cache.py +++ b/tests/local_testing/test_dual_cache.py @@ -245,3 +245,72 @@ async def test_dual_cache_delete(is_async): result = dual_cache.get_cache(test_key) assert result is None + + +def test_dual_cache_late_attach_redis_wires_writes_and_ttl_sync(): + """ + Typical lazy startup (sync): DualCache runs with in-memory only, then Redis + becomes available and is attached. New writes must reach Redis; keys written + before attach are not backfilled. Optional default_redis_ttl is applied on attach. + """ + in_memory = InMemoryCache() + dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=None) + + mock_redis = MagicMock() + mock_redis.set_cache = MagicMock() + mock_redis.async_set_cache = AsyncMock() + + key_before = f"before_attach_{uuid.uuid4()}" + val_before = {"phase": "memory_only"} + dual_cache.set_cache(key_before, val_before) + + assert in_memory.get_cache(key_before) == val_before + + dual_cache.attach_redis_cache(mock_redis, default_redis_ttl=99.0) + assert dual_cache.redis_cache is mock_redis + assert dual_cache.default_redis_ttl == 99.0 + + mock_redis.set_cache.assert_not_called() + + key_after = f"after_attach_{uuid.uuid4()}" + val_after = {"phase": "memory_and_redis"} + dual_cache.set_cache(key_after, val_after) + mock_redis.set_cache.assert_called_once() + assert mock_redis.set_cache.call_args[0][:2] == (key_after, val_after) + + assert in_memory.get_cache(key_after) == val_after + + +@pytest.mark.asyncio +async def test_dual_cache_late_attach_redis_wires_writes_and_ttl_async(): + """ + Typical lazy startup (async): DualCache runs with in-memory only, then Redis + becomes available and is attached. New writes must reach Redis; keys written + before attach are not backfilled. Optional default_redis_ttl is applied on attach. + """ + in_memory = InMemoryCache() + dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=None) + + mock_redis = MagicMock() + mock_redis.set_cache = MagicMock() + mock_redis.async_set_cache = AsyncMock() + + key_before = f"before_attach_{uuid.uuid4()}" + val_before = {"phase": "memory_only"} + await dual_cache.async_set_cache(key_before, val_before) + + assert in_memory.get_cache(key_before) == val_before + + dual_cache.attach_redis_cache(mock_redis, default_redis_ttl=99.0) + assert dual_cache.redis_cache is mock_redis + assert dual_cache.default_redis_ttl == 99.0 + + mock_redis.async_set_cache.assert_not_called() + + key_after = f"after_attach_{uuid.uuid4()}" + val_after = {"phase": "memory_and_redis"} + await dual_cache.async_set_cache(key_after, val_after) + mock_redis.async_set_cache.assert_called_once() + assert mock_redis.async_set_cache.call_args[0][:2] == (key_after, val_after) + + assert in_memory.get_cache(key_after) == val_after diff --git a/tests/test_litellm/proxy/common_utils/test_cache_codec.py b/tests/test_litellm/proxy/common_utils/test_cache_codec.py new file mode 100644 index 00000000000..e99f22a75d4 --- /dev/null +++ b/tests/test_litellm/proxy/common_utils/test_cache_codec.py @@ -0,0 +1,85 @@ +import logging +from typing import Optional + +import pytest +from pydantic import BaseModel, ValidationError + +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec + + +class _SampleModel(BaseModel): + name: str + count: Optional[int] = None + + +class _SampleSubModel(_SampleModel): + pass + + +class TestCacheCodecSerialize: + def test_without_model_type_base_model_dumped_json_safe(self): + m = _SampleModel(name="a", count=1) + out = CacheCodec.serialize(m) + assert out == {"name": "a", "count": 1} + + def test_without_model_type_dict_unchanged(self): + d = {"name": "x"} + assert CacheCodec.serialize(d) is d + + def test_without_model_type_primitive_unchanged(self): + assert CacheCodec.serialize(42) == 42 + + def test_with_model_type_dict_validated_and_dumped(self): + out = CacheCodec.serialize({"name": "b", "count": 2}, model_type=_SampleModel) + assert out == {"name": "b", "count": 2} + + def test_with_model_type_base_model_validated_and_dumped(self): + m = _SampleModel(name="c", count=None) + out = CacheCodec.serialize(m, model_type=_SampleModel) + assert out == {"name": "c"} + + def test_with_model_type_exclude_none_on_dump(self): + out = CacheCodec.serialize({"name": "d"}, model_type=_SampleModel) + assert out == {"name": "d"} + assert "count" not in out + + def test_with_model_type_non_dict_non_model_passthrough(self): + assert CacheCodec.serialize("raw", model_type=_SampleModel) == "raw" + + def test_with_model_type_invalid_dict_raises(self): + with pytest.raises(ValidationError): + CacheCodec.serialize({"count": 1}, model_type=_SampleModel) + + +class TestCacheCodecDeserialize: + def test_none_returns_none(self): + assert CacheCodec.deserialize(None, _SampleModel) is None + + def test_dict_validates_to_model(self): + m = CacheCodec.deserialize({"name": "e", "count": 3}, _SampleModel) + assert isinstance(m, _SampleModel) + assert m.name == "e" + assert m.count == 3 + + def test_instance_same_type_returned_as_is(self): + original = _SampleModel(name="f") + m = CacheCodec.deserialize(original, _SampleModel) + assert m is original + + def test_subclass_instance_accepted(self): + sub = _SampleSubModel(name="g") + m = CacheCodec.deserialize(sub, _SampleModel) + assert m is sub + + def test_wrong_type_returns_none(self): + assert CacheCodec.deserialize("not-a-dict", _SampleModel) is None + + def test_invalid_dict_returns_none_and_logs_warning(self, caplog): + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + out = CacheCodec.deserialize({"count": 1}, _SampleModel) + assert out is None + assert any( + "CacheCodec.deserialize" in r.message and "_SampleModel" in r.message + for r in caplog.records + if r.levelno >= logging.WARNING + ), f"Expected deserialize validation warning. Records: {[r.message for r in caplog.records]}" From 9cfc5fc20951ddf7d2aae0e409fe9c1b5d458cc4 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Tue, 21 Apr 2026 23:30:46 +0000 Subject: [PATCH 11/75] update tests --- tests/local_testing/test_dual_cache.py | 69 ------------------ tests/test_litellm/caching/test_dual_cache.py | 70 +++++++++++++++++++ 2 files changed, 70 insertions(+), 69 deletions(-) diff --git a/tests/local_testing/test_dual_cache.py b/tests/local_testing/test_dual_cache.py index 30f45033c58..5a1cdf86487 100644 --- a/tests/local_testing/test_dual_cache.py +++ b/tests/local_testing/test_dual_cache.py @@ -245,72 +245,3 @@ async def test_dual_cache_delete(is_async): result = dual_cache.get_cache(test_key) assert result is None - - -def test_dual_cache_late_attach_redis_wires_writes_and_ttl_sync(): - """ - Typical lazy startup (sync): DualCache runs with in-memory only, then Redis - becomes available and is attached. New writes must reach Redis; keys written - before attach are not backfilled. Optional default_redis_ttl is applied on attach. - """ - in_memory = InMemoryCache() - dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=None) - - mock_redis = MagicMock() - mock_redis.set_cache = MagicMock() - mock_redis.async_set_cache = AsyncMock() - - key_before = f"before_attach_{uuid.uuid4()}" - val_before = {"phase": "memory_only"} - dual_cache.set_cache(key_before, val_before) - - assert in_memory.get_cache(key_before) == val_before - - dual_cache.attach_redis_cache(mock_redis, default_redis_ttl=99.0) - assert dual_cache.redis_cache is mock_redis - assert dual_cache.default_redis_ttl == 99.0 - - mock_redis.set_cache.assert_not_called() - - key_after = f"after_attach_{uuid.uuid4()}" - val_after = {"phase": "memory_and_redis"} - dual_cache.set_cache(key_after, val_after) - mock_redis.set_cache.assert_called_once() - assert mock_redis.set_cache.call_args[0][:2] == (key_after, val_after) - - assert in_memory.get_cache(key_after) == val_after - - -@pytest.mark.asyncio -async def test_dual_cache_late_attach_redis_wires_writes_and_ttl_async(): - """ - Typical lazy startup (async): DualCache runs with in-memory only, then Redis - becomes available and is attached. New writes must reach Redis; keys written - before attach are not backfilled. Optional default_redis_ttl is applied on attach. - """ - in_memory = InMemoryCache() - dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=None) - - mock_redis = MagicMock() - mock_redis.set_cache = MagicMock() - mock_redis.async_set_cache = AsyncMock() - - key_before = f"before_attach_{uuid.uuid4()}" - val_before = {"phase": "memory_only"} - await dual_cache.async_set_cache(key_before, val_before) - - assert in_memory.get_cache(key_before) == val_before - - dual_cache.attach_redis_cache(mock_redis, default_redis_ttl=99.0) - assert dual_cache.redis_cache is mock_redis - assert dual_cache.default_redis_ttl == 99.0 - - mock_redis.async_set_cache.assert_not_called() - - key_after = f"after_attach_{uuid.uuid4()}" - val_after = {"phase": "memory_and_redis"} - await dual_cache.async_set_cache(key_after, val_after) - mock_redis.async_set_cache.assert_called_once() - assert mock_redis.async_set_cache.call_args[0][:2] == (key_after, val_after) - - assert in_memory.get_cache(key_after) == val_after diff --git a/tests/test_litellm/caching/test_dual_cache.py b/tests/test_litellm/caching/test_dual_cache.py index 8e502175761..64774726201 100644 --- a/tests/test_litellm/caching/test_dual_cache.py +++ b/tests/test_litellm/caching/test_dual_cache.py @@ -1,5 +1,6 @@ import asyncio import time +import uuid from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -260,3 +261,72 @@ async def test_async_increment_cache_returns_none_when_no_in_memory_cache_and_re f"Expected None when in_memory_cache is absent and Redis fails, got {result!r}. " "Returning the delta (1.0) would silently miscalculate rate-limit counters." ) + + +def test_dual_cache_late_attach_redis_wires_writes_and_ttl_sync(): + """ + Typical lazy startup (sync): DualCache runs with in-memory only, then Redis + becomes available and is attached. New writes must reach Redis; keys written + before attach are not backfilled. Optional default_redis_ttl is applied on attach. + """ + in_memory = InMemoryCache() + dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=None) + + mock_redis = MagicMock() + mock_redis.set_cache = MagicMock() + mock_redis.async_set_cache = AsyncMock() + + key_before = f"before_attach_{uuid.uuid4()}" + val_before = {"phase": "memory_only"} + dual_cache.set_cache(key_before, val_before) + + assert in_memory.get_cache(key_before) == val_before + + dual_cache.attach_redis_cache(mock_redis, default_redis_ttl=99.0) + assert dual_cache.redis_cache is mock_redis + assert dual_cache.default_redis_ttl == 99.0 + + mock_redis.set_cache.assert_not_called() + + key_after = f"after_attach_{uuid.uuid4()}" + val_after = {"phase": "memory_and_redis"} + dual_cache.set_cache(key_after, val_after) + mock_redis.set_cache.assert_called_once() + assert mock_redis.set_cache.call_args[0][:2] == (key_after, val_after) + + assert in_memory.get_cache(key_after) == val_after + + +@pytest.mark.asyncio +async def test_dual_cache_late_attach_redis_wires_writes_and_ttl_async(): + """ + Typical lazy startup (async): DualCache runs with in-memory only, then Redis + becomes available and is attached. New writes must reach Redis; keys written + before attach are not backfilled. Optional default_redis_ttl is applied on attach. + """ + in_memory = InMemoryCache() + dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=None) + + mock_redis = MagicMock() + mock_redis.set_cache = MagicMock() + mock_redis.async_set_cache = AsyncMock() + + key_before = f"before_attach_{uuid.uuid4()}" + val_before = {"phase": "memory_only"} + await dual_cache.async_set_cache(key_before, val_before) + + assert in_memory.get_cache(key_before) == val_before + + dual_cache.attach_redis_cache(mock_redis, default_redis_ttl=99.0) + assert dual_cache.redis_cache is mock_redis + assert dual_cache.default_redis_ttl == 99.0 + + mock_redis.async_set_cache.assert_not_called() + + key_after = f"after_attach_{uuid.uuid4()}" + val_after = {"phase": "memory_and_redis"} + await dual_cache.async_set_cache(key_after, val_after) + mock_redis.async_set_cache.assert_called_once() + assert mock_redis.async_set_cache.call_args[0][:2] == (key_after, val_after) + + assert in_memory.get_cache(key_after) == val_after From f4e976e22570d23d9aed896d85a8bbcbd6c9b2db Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 22 Apr 2026 09:45:21 +0530 Subject: [PATCH 12/75] fix(anthropic): handle response_format tool alongside user tools in non-streaming Non-streaming path required len(tool_calls)==1 to unwrap json_tool_call, so mixed user tools leaked the internal tool. Align with Bedrock converse handling: strip internal tools, merge structured JSON into content. Made-with: Cursor --- litellm/llms/anthropic/chat/transformation.py | 79 ++++++++++++++----- .../test_anthropic_chat_transformation.py | 34 ++++++++ 2 files changed, 93 insertions(+), 20 deletions(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index cd5bb731717..4a6743b74be 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1549,25 +1549,56 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) data["output_config"] = output_config + def _resolve_json_mode_non_streaming( + self, + json_mode: Optional[bool], + tool_calls: List[ChatCompletionToolCallChunk], + ) -> Tuple[ + Optional[LitellmMessage], + List[ChatCompletionToolCallChunk], + Optional[str], + ]: + """Strip internal response_format tool calls; merge payload into content when mixed with user tools.""" + if json_mode is not True or not tool_calls: + return None, tool_calls, None + + json_indices = [ + i + for i, t in enumerate(tool_calls) + if t.get("function", {}).get("name") == RESPONSE_FORMAT_TOOL_NAME + ] + if not json_indices: + return None, tool_calls, None + + if len(json_indices) == len(tool_calls): + json_tool = tool_calls[json_indices[0]] + if json_tool.get("function", {}).get("arguments") is None: + return None, tool_calls, None + _message = AnthropicConfig._convert_tool_response_to_message( + tool_calls=[json_tool] + ) + return _message, [], None + + first_json = tool_calls[json_indices[0]] + json_msg = AnthropicConfig._convert_tool_response_to_message([first_json]) + extra_content: Optional[str] = ( + json_msg.content if json_msg is not None else None + ) + filtered_tools = [ + t for i, t in enumerate(tool_calls) if i not in json_indices + ] + return None, filtered_tools, extra_content + def _transform_response_for_json_mode( self, json_mode: Optional[bool], tool_calls: List[ChatCompletionToolCallChunk], ) -> Optional[LitellmMessage]: - _message: Optional[LitellmMessage] = None - if json_mode is True and len(tool_calls) == 1: - # check if tool name is the default tool name - json_mode_content_str: Optional[str] = None - if ( - "name" in tool_calls[0]["function"] - and tool_calls[0]["function"]["name"] == RESPONSE_FORMAT_TOOL_NAME - ): - json_mode_content_str = tool_calls[0]["function"].get("arguments") - if json_mode_content_str is not None: - _message = AnthropicConfig._convert_tool_response_to_message( - tool_calls=tool_calls, - ) - return _message + replacement, _, _ = self._resolve_json_mode_non_streaming( + json_mode=json_mode, + tool_calls=tool_calls, + ) + return replacement def extract_response_content(self, completion_response: dict) -> Tuple[ str, @@ -1927,19 +1958,27 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): tool_calls, ) + json_mode_message, tool_calls_for_message, json_extra_content = ( + self._resolve_json_mode_non_streaming( + json_mode=json_mode, + tool_calls=tool_calls, + ) + ) + merged_text = text_content or "" + if json_extra_content: + merged_text = ( + merged_text + json_extra_content if merged_text else json_extra_content + ) + _message = litellm.Message( - tool_calls=tool_calls, - content=text_content or None, + tool_calls=tool_calls_for_message, + content=merged_text or None, provider_specific_fields=provider_specific_fields, thinking_blocks=thinking_blocks, reasoning_content=reasoning_content, ) _message.provider_specific_fields = provider_specific_fields - json_mode_message = self._transform_response_for_json_mode( - json_mode=json_mode, - tool_calls=tool_calls, - ) if json_mode_message is not None: completion_response["stop_reason"] = "stop" _message = json_mode_message diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index a7f5f92ab05..84a985b275f 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -8,6 +8,7 @@ sys.path.insert( ) # Adds the parent directory to the system path from unittest.mock import MagicMock, patch +from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.llms.anthropic.chat.transformation import AnthropicConfig from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( AnthropicMessagesConfig, @@ -38,6 +39,39 @@ def test_response_format_transformation_unit_test(): print(result) +def test_anthropic_json_mode_non_streaming_mixed_internal_and_user_tools(): + """Non-streaming + response_format: internal json tool must not require len(tool_calls)==1.""" + config = AnthropicConfig() + tool_calls = [ + { + "id": "toolu_json", + "type": "function", + "function": { + "name": RESPONSE_FORMAT_TOOL_NAME, + "arguments": '{"values": {"answer": 42}}', + }, + "index": 0, + }, + { + "id": "toolu_user", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "NY"}', + }, + "index": 1, + }, + ] + replacement, filtered, extra = config._resolve_json_mode_non_streaming( + json_mode=True, + tool_calls=tool_calls, + ) + assert replacement is None + assert len(filtered) == 1 + assert filtered[0]["function"]["name"] == "get_weather" + assert extra == '{"answer": 42}' + + def test_calculate_usage(): """ Do not include cache_creation_input_tokens in the prompt_tokens From 595f42d22a81321d46ffe335f2a70347e77d7bcf Mon Sep 17 00:00:00 2001 From: harish-berri Date: Thu, 23 Apr 2026 23:42:03 +0000 Subject: [PATCH 13/75] Refactor caching logic in auth_checks and user_api_key_auth to utilize CacheCodec for serialization and deserialization. Simplify cache retrieval by removing unnecessary type checks and streamline cache storage with consistent key formatting. --- litellm/proxy/auth/auth_checks.py | 25 +++++++++++-------------- litellm/proxy/auth/user_api_key_auth.py | 9 ++++++++- 2 files changed, 19 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 1577b5bda9c..760c0fb12d6 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1627,13 +1627,10 @@ async def _get_team_object_from_cache( if cached_team_obj is None: cached_team_obj = await user_api_key_cache.async_get_cache(key=key) - if cached_team_obj is not None: - if isinstance(cached_team_obj, dict): - return LiteLLM_TeamTableCachedObj(**cached_team_obj) - elif isinstance(cached_team_obj, LiteLLM_TeamTableCachedObj): - return cached_team_obj - - return None + if cached_team_obj is None: + return None + + return CacheCodec.deserialize(cached_team_obj, LiteLLM_TeamTableCachedObj) async def get_team_object( @@ -1709,9 +1706,12 @@ async def _cache_access_object( proxy_logging_obj: Optional[ProxyLogging] = None, ): key = "access_group_id:{}".format(access_group_id) + cache_payload = CacheCodec.serialize( + access_group_table, model_type=LiteLLM_AccessGroupTable + ) await user_api_key_cache.async_set_cache( key=key, - value=access_group_table, + value=cache_payload, ttl=DEFAULT_ACCESS_GROUP_CACHE_TTL, ) @@ -1758,13 +1758,10 @@ async def get_access_object( key = "access_group_id:{}".format(access_group_id) - # Always check cache first cached_access_obj = await user_api_key_cache.async_get_cache(key=key) - if cached_access_obj is not None: - if isinstance(cached_access_obj, dict): - return LiteLLM_AccessGroupTable(**cached_access_obj) - elif isinstance(cached_access_obj, LiteLLM_AccessGroupTable): - return cached_access_obj + deserialized = CacheCodec.deserialize(cached_access_obj, LiteLLM_AccessGroupTable) + if deserialized is not None: + return deserialized # Not in cache - fetch from DB try: diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index a6c8fd02ee1..855ce50ec0c 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -58,6 +58,7 @@ from litellm.proxy.auth.oauth2_check import Oauth2Handler from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, _safe_get_request_headers, @@ -1454,8 +1455,14 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 # Only cache when the key is a real team_id (non-team keys must not use key=None). if valid_token.team_id is not None and _team_obj is not None: + # Match get_team_object / spend counters: "team_id:{id}". Serialize for Redis + # (json.dumps) — same as _cache_team_object in auth_checks. + team_cache_key = f"team_id:{valid_token.team_id}" await user_api_key_cache.async_set_cache( - key=valid_token.team_id, value=_team_obj + key=team_cache_key, + value=CacheCodec.serialize( + _team_obj, model_type=LiteLLM_TeamTableCachedObj + ), ) # save team table in cache - used for tpm/rpm limiting - tpm_rpm_limiter.py # Fetch project object if key belongs to a project From d4a26ff364a85980cdef5cae86beb77c6628b167 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Fri, 24 Apr 2026 01:30:01 +0000 Subject: [PATCH 14/75] Enhance caching mechanism by integrating CacheCodec for serialization across various components. Introduce the enable_redis_auth_cache flag to control Redis integration for user_api_key_cache, improving performance in multi-worker deployments. Update documentation and tests to reflect these changes. --- .gitignore | 3 +- docs/my-website/docs/proxy/config_settings.md | 61 +++++++ litellm/proxy/_types.py | 13 ++ litellm/proxy/auth/auth_checks.py | 156 ++++++++++------- litellm/proxy/auth/handle_jwt.py | 9 +- litellm/proxy/auth/user_api_key_auth.py | 26 ++- .../common_utils/cache_pydantic_utils.py | 7 + .../access_group_endpoints.py | 19 +- litellm/proxy/management_endpoints/ui_sso.py | 7 +- litellm/proxy/proxy_server.py | 164 ++++++++++-------- .../proxy/common_utils/test_cache_codec.py | 36 ++++ .../proxy/test_redis_auth_cache_flag.py | 130 ++++++++++++++ 12 files changed, 465 insertions(+), 166 deletions(-) create mode 100644 tests/test_litellm/proxy/test_redis_auth_cache_flag.py diff --git a/.gitignore b/.gitignore index f5591b7f2c6..59812ed6ed4 100644 --- a/.gitignore +++ b/.gitignore @@ -99,4 +99,5 @@ STABILIZATION_TODO.md **/test-results **/playwright-report **/*.storageState.json -**/coverage \ No newline at end of file +**/coverage +test-config \ No newline at end of file diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index be819a04ca1..7187690375d 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -307,9 +307,70 @@ router_settings: | token_rate_limit_type | string | Rate limit counting method: "total", "output", or "input" tokens | | use_redis_transaction_buffer | boolean | If true, buffers database transactions in Redis before writing | | use_shared_health_check | boolean | If true, uses Redis-backed shared health check state across multiple proxy instances | +| enable_redis_auth_cache | boolean | **[Beta]** When `true`, attaches Redis to the virtual-key auth cache (`user_api_key_cache`) so all proxy workers/pods share the same cache instead of each pod resolving keys independently against the database. Requires `litellm_settings.cache: true` with a Redis backend. Significantly reduces database load in multi-worker deployments by eliminating per-pod cache misses on the `combined_view` query. Off by default for a safe phased rollout — enable once your Redis cluster is healthy. Will become the default in a future release. See [Redis Auth Cache](#redis-auth-cache-multi-worker-db-load-reduction). | | user_header_mappings | dict | Map custom request headers to user IDs using lookup rules | | user_header_name | string | HTTP header name to extract user identity from requests | +## Redis Auth Cache — Multi-Worker DB Load Reduction + +### Problem + +In multi-worker or multi-pod deployments each worker process keeps its own **in-memory** virtual-key cache (`user_api_key_cache`). When a request is routed to a worker that has not seen a key before, it runs a `combined_view` SQL query against the database (8+ JOINs). With many workers and a large number of unique keys this causes: + +- Redundant DB queries on every pod's cold-start +- Sustained CPU spikes on the database (especially visible in RDS Performance Insights) +- Increased p99 latency for requests that miss the local cache + +### Solution + +Setting `enable_redis_auth_cache: true` attaches Redis to `user_api_key_cache` so the resolved key object is stored in a **shared** Redis cache. A cache hit on any worker prevents the DB query entirely. + +```yaml +# config.yaml + +litellm_settings: + cache: true + cache_params: + type: redis + host: os.environ/REDIS_HOST + port: os.environ/REDIS_PORT + +general_settings: + master_key: sk-1234 + enable_redis_auth_cache: true # ← share the auth cache across workers +``` + +### Requirements + +| Requirement | Notes | +|---|---| +| `litellm_settings.cache: true` | Redis must be configured as the cache backend | +| Redis cluster is healthy | Auth lookups now depend on Redis availability — monitor it | +| LiteLLM ≥ version with CacheCodec support | All read/write paths now use `CacheCodec` for safe serialisation across Redis round-trips | + +### Rollout recommendation + +1. Deploy with `enable_redis_auth_cache: false` (default) and baseline DB CPU. +2. Enable on a canary pod first; confirm cache-hit rate via proxy debug logs (`LITELLM_LOG=DEBUG`). +3. Roll out to all pods; monitor DB CPU drop. +4. The flag will be removed and the behaviour made permanent in a future LiteLLM release. + +### Debugging + +Set `LITELLM_LOG=DEBUG` and look for: + +``` +enable_redis_auth_cache=True: attached Redis to user_api_key_cache — virtual-key lookups are now shared across all proxy workers. +``` + +If the flag is off you will see: + +``` +enable_redis_auth_cache is not set: user_api_key_cache remains in-memory only (per-worker). Set general_settings.enable_redis_auth_cache: true to share the auth cache across workers and reduce DB load. +``` + +--- + ### router_settings - Reference :::info diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index a592a81ed86..960261d5cf1 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2358,6 +2358,19 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="List of MCP server fields that must be filled in for a submission to pass standards checks (e.g. ['description', 'source_url', 'alias']).", ) + enable_redis_auth_cache: Optional[bool] = Field( + None, + description=( + "When True, attaches Redis to user_api_key_cache so virtual-key lookups " + "are shared across all proxy workers/pods instead of being resolved " + "per-process. Requires a Redis cache to be configured under litellm_settings. " + "Reduces DB load significantly in multi-worker deployments by eliminating " + "redundant combined_view SQL queries caused by per-pod cache misses. " + "Off by default for a safe phased rollout — set to True once your Redis " + "cluster is healthy and the CacheCodec serialisation has been validated in " + "your environment. Will be enabled by default in a future release." + ), + ) class ConfigYAML(LiteLLMPydanticObjectBase): diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 760c0fb12d6..02d41041e86 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -823,7 +823,9 @@ async def get_default_end_user_budget( # Check cache first cached_budget = await user_api_key_cache.async_get_cache(key=cache_key) if cached_budget is not None: - return LiteLLM_BudgetTable(**cached_budget) + deserialized = CacheCodec.deserialize(cached_budget, LiteLLM_BudgetTable) + if deserialized is not None: + return deserialized # Fetch from database try: @@ -837,14 +839,15 @@ async def get_default_end_user_budget( ) return None + _budget_obj = LiteLLM_BudgetTable(**budget_record.dict()) # Cache the budget for 60 seconds await user_api_key_cache.async_set_cache( key=cache_key, - value=budget_record.dict(), + value=CacheCodec.serialize(_budget_obj, model_type=LiteLLM_BudgetTable), ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) - return LiteLLM_BudgetTable(**budget_record.dict()) + return _budget_obj except Exception as e: verbose_proxy_logger.error(f"Error fetching default end user budget: {str(e)}") @@ -960,20 +963,20 @@ async def get_end_user_object( # Check cache first cached_user_obj = await user_api_key_cache.async_get_cache(key=_key) if cached_user_obj is not None: - return_obj = LiteLLM_EndUserTable(**cached_user_obj) + return_obj = CacheCodec.deserialize(cached_user_obj, LiteLLM_EndUserTable) + if return_obj is not None: + # Apply default budget if needed + return_obj = await _apply_default_budget_to_end_user( + end_user_obj=return_obj, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + ) - # Apply default budget if needed - return_obj = await _apply_default_budget_to_end_user( - end_user_obj=return_obj, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - ) + # Check budget limits + _check_end_user_budget(end_user_obj=return_obj, route=route) - # Check budget limits - _check_end_user_budget(end_user_obj=return_obj, route=route) - - return return_obj + return return_obj # Fetch from database try: @@ -996,9 +999,10 @@ async def get_end_user_object( parent_otel_span=parent_otel_span, ) - # Save to cache (always store as dict for consistency) + # Save to cache await user_api_key_cache.async_set_cache( - key="end_user_id:{}".format(end_user_id), value=_response.dict() + key="end_user_id:{}".format(end_user_id), + value=CacheCodec.serialize(_response, model_type=LiteLLM_EndUserTable), ) # Check budget limits @@ -1051,10 +1055,11 @@ async def get_tag_objects_batch( cache_key = f"tag:{tag_name}" cached_tag = await user_api_key_cache.async_get_cache(key=cache_key) if cached_tag is not None: - if isinstance(cached_tag, dict): - tag_objects[tag_name] = LiteLLM_TagTable(**cached_tag) + deserialized_tag = CacheCodec.deserialize(cached_tag, LiteLLM_TagTable) + if deserialized_tag is not None: + tag_objects[tag_name] = deserialized_tag else: - tag_objects[tag_name] = cached_tag + uncached_tags.append(tag_name) else: uncached_tags.append(tag_name) @@ -1070,11 +1075,12 @@ async def get_tag_objects_batch( for db_tag in db_tags: tag_name = db_tag.tag_name cache_key = f"tag:{tag_name}" - # Cache with default TTL (same as end_user objects) + _tag_obj = LiteLLM_TagTable(**db_tag.dict()) await user_api_key_cache.async_set_cache( - key=cache_key, value=db_tag.dict() + key=cache_key, + value=CacheCodec.serialize(_tag_obj, model_type=LiteLLM_TagTable), ) - tag_objects[tag_name] = LiteLLM_TagTable(**db_tag.dict()) + tag_objects[tag_name] = _tag_obj except Exception as e: verbose_proxy_logger.debug(f"Error batch fetching tags from database: {e}") @@ -1146,7 +1152,11 @@ async def get_team_membership( # check if in cache cached_membership_obj = await user_api_key_cache.async_get_cache(key=_key) if cached_membership_obj is not None: - return LiteLLM_TeamMembership(**cached_membership_obj) + deserialized_membership = CacheCodec.deserialize( + cached_membership_obj, LiteLLM_TeamMembership + ) + if deserialized_membership is not None: + return deserialized_membership # else, check db try: @@ -1158,10 +1168,11 @@ async def get_team_membership( if response is None: return None - # save the team membership object to cache (store as dict) - await user_api_key_cache.async_set_cache(key=_key, value=response.dict()) - _response = LiteLLM_TeamMembership(**response.dict()) + await user_api_key_cache.async_set_cache( + key=_key, + value=CacheCodec.serialize(_response, model_type=LiteLLM_TeamMembership), + ) return _response except Exception: @@ -1350,10 +1361,11 @@ async def get_user_object( if not check_db_only: cached_user_obj = await user_api_key_cache.async_get_cache(key=user_id) if cached_user_obj is not None: - if isinstance(cached_user_obj, dict): - return LiteLLM_UserTable(**cached_user_obj) - elif isinstance(cached_user_obj, LiteLLM_UserTable): - return cached_user_obj + deserialized_user = CacheCodec.deserialize( + cached_user_obj, LiteLLM_UserTable + ) + if deserialized_user is not None: + return deserialized_user # else, check db if prisma_client is None: raise Exception("No db connected") @@ -1415,7 +1427,7 @@ async def get_user_object( # save the user object to cache await user_api_key_cache.async_set_cache( key=user_id, - value=response_dict, + value=CacheCodec.serialize(_response, model_type=LiteLLM_UserTable), ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) @@ -1886,16 +1898,19 @@ async def get_team_object_by_alias( ) # Cache the result by both alias and team_id + _serialized_team = CacheCodec.serialize( + team_obj, model_type=LiteLLM_TeamTableCachedObj + ) await user_api_key_cache.async_set_cache( key=cache_key, - value=team_obj, + value=_serialized_team, ttl=DEFAULT_IN_MEMORY_TTL, ) # Also cache by team_id for consistency team_id_cache_key = "team_id:{}".format(team_obj.team_id) await user_api_key_cache.async_set_cache( key=team_id_cache_key, - value=team_obj, + value=_serialized_team, ttl=DEFAULT_IN_MEMORY_TTL, ) @@ -1946,10 +1961,11 @@ async def get_org_object_by_alias( cache_key = "org_alias:{}".format(org_alias) cached_org_obj = await user_api_key_cache.async_get_cache(key=cache_key) if cached_org_obj is not None: - if isinstance(cached_org_obj, dict): - return LiteLLM_OrganizationTable(**cached_org_obj) - elif isinstance(cached_org_obj, LiteLLM_OrganizationTable): - return cached_org_obj + deserialized_org = CacheCodec.deserialize( + cached_org_obj, LiteLLM_OrganizationTable + ) + if deserialized_org is not None: + return deserialized_org # Query database by organization_alias try: @@ -1976,16 +1992,19 @@ async def get_org_object_by_alias( org = orgs[0] org_obj = LiteLLM_OrganizationTable(**org.model_dump()) + _serialized_org = CacheCodec.serialize( + org_obj, model_type=LiteLLM_OrganizationTable + ) # Cache the result await user_api_key_cache.async_set_cache( key=cache_key, - value=org_obj.model_dump(), + value=_serialized_org, ttl=DEFAULT_IN_MEMORY_TTL, ) # Also cache by org_id for consistency await user_api_key_cache.async_set_cache( key="org_id:{}".format(org_obj.organization_id), - value=org_obj.model_dump(), + value=_serialized_org, ttl=DEFAULT_IN_MEMORY_TTL, ) @@ -2287,10 +2306,11 @@ async def get_object_permission( key = "object_permission_id:{}".format(object_permission_id) cached_obj_permission = await user_api_key_cache.async_get_cache(key=key) if cached_obj_permission is not None: - if isinstance(cached_obj_permission, dict): - return LiteLLM_ObjectPermissionTable(**cached_obj_permission) - elif isinstance(cached_obj_permission, LiteLLM_ObjectPermissionTable): - return cached_obj_permission + deserialized_perm = CacheCodec.deserialize( + cached_obj_permission, LiteLLM_ObjectPermissionTable + ) + if deserialized_perm is not None: + return deserialized_perm # else, check db try: @@ -2301,14 +2321,16 @@ async def get_object_permission( if response is None: return None - # save the object permission to cache + _perm_obj = LiteLLM_ObjectPermissionTable(**response.dict()) await user_api_key_cache.async_set_cache( key=key, - value=response.model_dump(), + value=CacheCodec.serialize( + _perm_obj, model_type=LiteLLM_ObjectPermissionTable + ), ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) - return LiteLLM_ObjectPermissionTable(**response.dict()) + return _perm_obj except Exception: return None @@ -2339,10 +2361,11 @@ async def get_managed_vector_store_rows_by_uuids( key = "managed_vector_store_id:{}".format(uuid) cached = await user_api_key_cache.async_get_cache(key=key) if cached is not None: - if isinstance(cached, dict): - result.append(LiteLLM_ManagedVectorStoresTable(**cached)) - elif isinstance(cached, LiteLLM_ManagedVectorStoresTable): - result.append(cached) + deserialized_vs = CacheCodec.deserialize( + cached, LiteLLM_ManagedVectorStoresTable + ) + if deserialized_vs is not None: + result.append(deserialized_vs) else: cache_misses.append(uuid) else: @@ -2370,7 +2393,9 @@ async def get_managed_vector_store_rows_by_uuids( key = "managed_vector_store_id:{}".format(cached_obj.vector_store_id) await user_api_key_cache.async_set_cache( key=key, - value=row_dict, + value=CacheCodec.serialize( + cached_obj, model_type=LiteLLM_ManagedVectorStoresTable + ), ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) result.append(cached_obj) @@ -2413,12 +2438,13 @@ async def get_org_object( cache_key = "org_id:{}:with_budget".format(org_id) # check if in cache - cached_org_obj = user_api_key_cache.async_get_cache(key=cache_key) + cached_org_obj = await user_api_key_cache.async_get_cache(key=cache_key) if cached_org_obj is not None: - if isinstance(cached_org_obj, dict): - return LiteLLM_OrganizationTable(**cached_org_obj) - elif isinstance(cached_org_obj, LiteLLM_OrganizationTable): - return cached_org_obj + deserialized_org = CacheCodec.deserialize( + cached_org_obj, LiteLLM_OrganizationTable + ) + if deserialized_org is not None: + return deserialized_org # else, check db try: query_kwargs: Dict[str, Any] = {"where": {"organization_id": org_id}} @@ -2432,16 +2458,15 @@ async def get_org_object( if response is None: raise Exception + _org_obj = LiteLLM_OrganizationTable(**response.model_dump()) # Cache the result await user_api_key_cache.async_set_cache( key=cache_key, - value=( - response.model_dump() if hasattr(response, "model_dump") else response - ), + value=CacheCodec.serialize(_org_obj, model_type=LiteLLM_OrganizationTable), ttl=DEFAULT_IN_MEMORY_TTL, ) - return response + return _org_obj except Exception: raise Exception( f"Organization doesn't exist in db. Organization={org_id}. Create organization via `/organization/new` call." @@ -3451,10 +3476,11 @@ async def get_project_object( cache_key = "project_id:{}".format(project_id) cached_obj = await user_api_key_cache.async_get_cache(key=cache_key) if cached_obj is not None: - if isinstance(cached_obj, dict): - return LiteLLM_ProjectTableCachedObj(**cached_obj) - elif isinstance(cached_obj, LiteLLM_ProjectTableCachedObj): - return cached_obj + deserialized_project = CacheCodec.deserialize( + cached_obj, LiteLLM_ProjectTableCachedObj + ) + if deserialized_project is not None: + return deserialized_project # Fetch from DB project_row = await prisma_client.db.litellm_projecttable.find_unique( diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 4a6856b6d14..07d944ae931 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -45,6 +45,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_checks import can_team_access_model +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec from litellm.proxy.utils import PrismaClient, ProxyLogging from .auth_checks import ( @@ -1376,7 +1377,9 @@ class JWTAuthManager: if user_api_key_cache is not None: await user_api_key_cache.async_set_cache( key=user_object.user_id, - value=user_object.model_dump(), + value=CacheCodec.serialize( + user_object, model_type=LiteLLM_UserTable + ), ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) @@ -1399,7 +1402,9 @@ class JWTAuthManager: if user_api_key_cache is not None: await user_api_key_cache.async_set_cache( key=user_object.user_id, - value=user_object.model_dump(), + value=CacheCodec.serialize( + user_object, model_type=LiteLLM_UserTable + ), ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) return None diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 855ce50ec0c..3fe8863e267 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1297,27 +1297,39 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 if prisma_client is not None: _cache_key = f"{valid_token.team_id}_{valid_token.user_id}" - team_member_info = await user_api_key_cache.async_get_cache( + _cached_member = await user_api_key_cache.async_get_cache( key=_cache_key ) + team_member_info: Optional[LiteLLM_TeamMembership] = ( + CacheCodec.deserialize(_cached_member, LiteLLM_TeamMembership) + if _cached_member is not None + else None + ) if team_member_info is None: # read from DB _user_id = valid_token.user_id _team_id = valid_token.team_id if _user_id is not None and _team_id is not None: - team_member_info = await prisma_client.db.litellm_teammembership.find_first( + _db_member = await prisma_client.db.litellm_teammembership.find_first( where={ "user_id": _user_id, "team_id": _team_id, }, # type: ignore include={"litellm_budget_table": True}, ) - await user_api_key_cache.async_set_cache( - key=_cache_key, - value=team_member_info, - ttl=5, - ) + if _db_member is not None: + team_member_info = LiteLLM_TeamMembership( + **_db_member.dict() + ) + await user_api_key_cache.async_set_cache( + key=_cache_key, + value=CacheCodec.serialize( + team_member_info, + model_type=LiteLLM_TeamMembership, + ), + ttl=5, + ) if ( team_member_info is not None diff --git a/litellm/proxy/common_utils/cache_pydantic_utils.py b/litellm/proxy/common_utils/cache_pydantic_utils.py index af0c3eef18d..80a8d6281a1 100644 --- a/litellm/proxy/common_utils/cache_pydantic_utils.py +++ b/litellm/proxy/common_utils/cache_pydantic_utils.py @@ -44,10 +44,17 @@ class CacheCodec: If ``model_type`` is set, the payload is validated with that model, then ``model_dump(mode="json", exclude_none=True)`` — symmetric with ``deserialize``. + If the value is already an instance of ``model_type`` (or a subclass), + ``model_validate`` is skipped to avoid an unnecessary Pydantic copy — the + value is dumped directly. + If ``model_type`` is omitted, any ``BaseModel`` is dumped as above; other values (e.g. plain ``dict``) are returned unchanged. """ if model_type is not None: + if isinstance(value, model_type): + # Already the right type: dump directly, skip re-validation copy. + return value.model_dump(mode="json", exclude_none=True) if isinstance(value, (dict, BaseModel)): return model_type.model_validate(value).model_dump( mode="json", exclude_none=True diff --git a/litellm/proxy/management_endpoints/access_group_endpoints.py b/litellm/proxy/management_endpoints/access_group_endpoints.py index caaec12f7a3..30fed096232 100644 --- a/litellm/proxy/management_endpoints/access_group_endpoints.py +++ b/litellm/proxy/management_endpoints/access_group_endpoints.py @@ -9,6 +9,7 @@ from litellm.proxy._types import ( LitellmUserRoles, UserAPIKeyAuth, ) +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec from litellm.proxy.auth.auth_checks import ( _cache_access_object, _cache_key_object, @@ -236,12 +237,11 @@ async def _patch_key_caches_add_access_group( ) -> None: """Patch cached key objects to include access_group_id.""" for token in key_tokens: - cached_key = await user_api_key_cache.async_get_cache(key=token) - if cached_key is None: + raw = await user_api_key_cache.async_get_cache(key=token) + if raw is None: continue - if isinstance(cached_key, dict): - cached_key = UserAPIKeyAuth(**cached_key) - if not isinstance(cached_key, UserAPIKeyAuth): + cached_key = CacheCodec.deserialize(raw, UserAPIKeyAuth) + if cached_key is None: continue if cached_key.access_group_ids is None: cached_key.access_group_ids = [access_group_id] @@ -267,12 +267,11 @@ async def _patch_key_caches_remove_access_group( ) -> None: """Patch cached key objects to remove access_group_id.""" for token in key_tokens: - cached_key = await user_api_key_cache.async_get_cache(key=token) - if cached_key is None: + raw = await user_api_key_cache.async_get_cache(key=token) + if raw is None: continue - if isinstance(cached_key, dict): - cached_key = UserAPIKeyAuth(**cached_key) - if isinstance(cached_key, UserAPIKeyAuth) and cached_key.access_group_ids: + cached_key = CacheCodec.deserialize(raw, UserAPIKeyAuth) + if cached_key is not None and cached_key.access_group_ids: cached_key.access_group_ids = [ ag for ag in cached_key.access_group_ids if ag != access_group_id ] diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 46e7963da7c..0362c3ff6b2 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -70,6 +70,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_user_object +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec from litellm.proxy.auth.auth_utils import _has_user_setup_sso from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -1233,11 +1234,7 @@ async def _sync_user_role_from_jwt_role_map( user_info.user_role = mapped_role.value await user_api_key_cache.async_set_cache( key=user_info.user_id, - value=( - user_info.model_dump() - if hasattr(user_info, "model_dump") - else dict(user_info) - ), + value=CacheCodec.serialize(user_info, model_type=LiteLLM_UserTable), ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 72ec41ef97f..f8c52a5289b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -78,8 +78,11 @@ from litellm.proxy._types import ( InvitationNew, InvitationUpdate, Litellm_EntityType, + LiteLLM_EndUserTable, LiteLLM_JWTAuth, + LiteLLM_TagTable, LiteLLM_TeamTable, + LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, LitellmUserRoles, PassThroughGenericEndpoint, @@ -2029,34 +2032,38 @@ async def update_cache( # noqa: PLR0915 ### UPDATE USER SPEND ### async def _update_user_cache(): ## UPDATE CACHE FOR USER ID + GLOBAL PROXY + if response_cost is None: + return user_ids = [user_id] try: for _id in user_ids: # Fetch the existing cost for the given user if _id is None: continue - existing_spend_obj = await user_api_key_cache.async_get_cache(key=_id) - if existing_spend_obj is None: + cached_user = await user_api_key_cache.async_get_cache(key=_id) + if cached_user is None: # do nothing if there is no cache value return + existing_spend_obj = CacheCodec.deserialize(cached_user, LiteLLM_UserTable) + if existing_spend_obj is None: + return verbose_proxy_logger.debug( f"_update_user_db: existing spend: {existing_spend_obj}; response_cost: {response_cost}" ) - if isinstance(existing_spend_obj, dict): - existing_spend = existing_spend_obj["spend"] - else: - existing_spend = existing_spend_obj.spend + existing_spend = existing_spend_obj.spend or 0.0 # Calculate the new cost by adding the existing cost and response_cost new_spend = existing_spend + response_cost - # Update the cost column for the given user - if isinstance(existing_spend_obj, dict): - existing_spend_obj["spend"] = new_spend - values_to_update_in_cache.append((_id, existing_spend_obj)) - else: - existing_spend_obj.spend = new_spend - values_to_update_in_cache.append((_id, existing_spend_obj.json())) + existing_spend_obj.spend = new_spend + values_to_update_in_cache.append( + ( + _id, + CacheCodec.serialize( + existing_spend_obj, model_type=LiteLLM_UserTable + ), + ) + ) ## UPDATE GLOBAL PROXY ## global_proxy_spend = await user_api_key_cache.async_get_cache( key="{}:spend".format(litellm_proxy_admin_name) @@ -2088,31 +2095,33 @@ async def update_cache( # noqa: PLR0915 _id = "end_user_id:{}".format(end_user_id) try: # Fetch the existing cost for the given user - existing_spend_obj = await user_api_key_cache.async_get_cache(key=_id) - if existing_spend_obj is None: + cached_end_user = await user_api_key_cache.async_get_cache(key=_id) + if cached_end_user is None: # if user does not exist in LiteLLM_UserTable, create a new user # do nothing if end-user not in api key cache return + existing_spend_obj = CacheCodec.deserialize( + cached_end_user, LiteLLM_EndUserTable + ) + if existing_spend_obj is None: + return verbose_proxy_logger.debug( f"_update_end_user_db: existing spend: {existing_spend_obj}; response_cost: {response_cost}" ) - if existing_spend_obj is None: - existing_spend = 0 - else: - if isinstance(existing_spend_obj, dict): - existing_spend = existing_spend_obj["spend"] - else: - existing_spend = existing_spend_obj.spend + + existing_spend = existing_spend_obj.spend or 0.0 # Calculate the new cost by adding the existing cost and response_cost new_spend = existing_spend + response_cost - # Update the cost column for the given user - if isinstance(existing_spend_obj, dict): - existing_spend_obj["spend"] = new_spend - values_to_update_in_cache.append((_id, existing_spend_obj)) - else: - existing_spend_obj.spend = new_spend - values_to_update_in_cache.append((_id, existing_spend_obj.json())) + existing_spend_obj.spend = new_spend + values_to_update_in_cache.append( + ( + _id, + CacheCodec.serialize( + existing_spend_obj, model_type=LiteLLM_EndUserTable + ), + ) + ) except Exception as e: verbose_proxy_logger.warning( "Spend tracking - failed to update end user spend in cache. " @@ -2131,36 +2140,32 @@ async def update_cache( # noqa: PLR0915 _id = "team_id:{}".format(team_id) try: - # Fetch the existing cost for the given user - existing_spend_obj: Optional[LiteLLM_TeamTable] = ( - await user_api_key_cache.async_get_cache(key=_id) + cached_team = await user_api_key_cache.async_get_cache(key=_id) + if cached_team is None: + # do nothing if team not in api key cache + return + existing_spend_obj: Optional[LiteLLM_TeamTableCachedObj] = ( + CacheCodec.deserialize(cached_team, LiteLLM_TeamTableCachedObj) ) if existing_spend_obj is None: - # do nothing if team not in api key cache return verbose_proxy_logger.debug( f"_update_team_db: existing spend: {existing_spend_obj}; response_cost: {response_cost}" ) - if existing_spend_obj is None: - existing_spend: Optional[float] = 0.0 - else: - if isinstance(existing_spend_obj, dict): - existing_spend = existing_spend_obj["spend"] - else: - existing_spend = existing_spend_obj.spend - if existing_spend is None: - existing_spend = 0.0 + existing_spend: float = existing_spend_obj.spend or 0.0 # Calculate the new cost by adding the existing cost and response_cost new_spend = existing_spend + response_cost - # Update the cost column for the given user - if isinstance(existing_spend_obj, dict): - existing_spend_obj["spend"] = new_spend - values_to_update_in_cache.append((_id, existing_spend_obj)) - else: - existing_spend_obj.spend = new_spend - values_to_update_in_cache.append((_id, existing_spend_obj)) + existing_spend_obj.spend = new_spend + values_to_update_in_cache.append( + ( + _id, + CacheCodec.serialize( + existing_spend_obj, model_type=LiteLLM_TeamTableCachedObj + ), + ) + ) except Exception as e: verbose_proxy_logger.warning( "Spend tracking - failed to update team spend in cache. " @@ -2187,32 +2192,32 @@ async def update_cache( # noqa: PLR0915 cache_key = f"tag:{tag_name}" # Fetch the existing tag object from cache - existing_tag_obj = await user_api_key_cache.async_get_cache( - key=cache_key - ) - if existing_tag_obj is None: + cached_tag = await user_api_key_cache.async_get_cache(key=cache_key) + if cached_tag is None: # do nothing if tag not in api key cache continue + existing_tag_obj = CacheCodec.deserialize(cached_tag, LiteLLM_TagTable) + if existing_tag_obj is None: + continue + verbose_proxy_logger.debug( f"_update_tag_cache: existing spend for tag={tag_name}: {existing_tag_obj}; response_cost: {response_cost}" ) - if isinstance(existing_tag_obj, dict): - existing_spend = existing_tag_obj.get("spend", 0) or 0 - else: - existing_spend = getattr(existing_tag_obj, "spend", 0) or 0 - + existing_spend = existing_tag_obj.spend or 0.0 # Calculate the new cost by adding the existing cost and response_cost new_spend = existing_spend + response_cost - # Update the spend column for the given tag - if isinstance(existing_tag_obj, dict): - existing_tag_obj["spend"] = new_spend - values_to_update_in_cache.append((cache_key, existing_tag_obj)) - else: - existing_tag_obj.spend = new_spend - values_to_update_in_cache.append((cache_key, existing_tag_obj)) + existing_tag_obj.spend = new_spend + values_to_update_in_cache.append( + ( + cache_key, + CacheCodec.serialize( + existing_tag_obj, model_type=LiteLLM_TagTable + ), + ) + ) except Exception as e: verbose_proxy_logger.warning( "Spend tracking - failed to update tag spend in cache. " @@ -2849,7 +2854,7 @@ class ProxyConfig: self, cache_params: dict, ): - global redis_usage_cache, llm_router + global redis_usage_cache, llm_router, general_settings from litellm import Cache if "default_in_memory_ttl" in cache_params: @@ -2871,16 +2876,23 @@ class ProxyConfig: ) # Note: PKCE verifier storage uses redis_usage_cache directly (not # user_api_key_cache) to avoid routing all API-key lookups through Redis. - # Share the same Redis client for virtual-key lookups (same DualCache as - # model_max_budget_limiter). attach_redis_cache is a no-op if Redis is - # already set (e.g. config reload). - user_api_key_cache.attach_redis_cache( - redis_usage_cache, - default_redis_ttl=litellm.default_redis_ttl, - ) - verbose_proxy_logger.debug( - "Attached redis_usage_cache Redis client to user_api_key_cache" - ) + if general_settings.get("enable_redis_auth_cache") is True: + user_api_key_cache.attach_redis_cache( + redis_usage_cache, + default_redis_ttl=litellm.default_redis_ttl, + ) + verbose_proxy_logger.info( + "enable_redis_auth_cache=True: attached Redis to " + "user_api_key_cache — virtual-key lookups are now " + "shared across all proxy workers." + ) + else: + verbose_proxy_logger.info( + "enable_redis_auth_cache is not set: user_api_key_cache " + "remains in-memory only (per-worker). Set " + "general_settings.enable_redis_auth_cache: true to share " + "the auth cache across workers and reduce DB load." + ) def switch_on_llm_response_caching(self): """ diff --git a/tests/test_litellm/proxy/common_utils/test_cache_codec.py b/tests/test_litellm/proxy/common_utils/test_cache_codec.py index e99f22a75d4..ef52eab4b59 100644 --- a/tests/test_litellm/proxy/common_utils/test_cache_codec.py +++ b/tests/test_litellm/proxy/common_utils/test_cache_codec.py @@ -1,5 +1,6 @@ import logging from typing import Optional +from unittest.mock import patch import pytest from pydantic import BaseModel, ValidationError @@ -50,6 +51,41 @@ class TestCacheCodecSerialize: with pytest.raises(ValidationError): CacheCodec.serialize({"count": 1}, model_type=_SampleModel) + def test_with_model_type_already_correct_instance_skips_revalidation(self): + """Fast-path: value is already model_type — model_validate must NOT be called.""" + m = _SampleModel(name="fast", count=7) + with patch.object(_SampleModel, "model_validate", wraps=_SampleModel.model_validate) as mock_validate: + out = CacheCodec.serialize(m, model_type=_SampleModel) + assert out == {"name": "fast", "count": 7} + mock_validate.assert_not_called() + + def test_with_model_type_subclass_instance_skips_revalidation(self): + """Subclass is isinstance of base → should also take the fast path.""" + sub = _SampleSubModel(name="sub", count=2) + with patch.object(_SampleModel, "model_validate", wraps=_SampleModel.model_validate) as mock_validate: + out = CacheCodec.serialize(sub, model_type=_SampleModel) + assert out == {"name": "sub", "count": 2} + mock_validate.assert_not_called() + + def test_with_model_type_dict_input_goes_through_model_validate(self): + """A dict value (not yet an instance) must still go through model_validate.""" + raw = {"name": "via-dict", "count": 5} + with patch.object( + _SampleModel, "model_validate", wraps=_SampleModel.model_validate + ) as mock_validate: + out = CacheCodec.serialize(raw, model_type=_SampleModel) + assert out == {"name": "via-dict", "count": 5} + mock_validate.assert_called_once() + + def test_with_model_type_incompatible_model_raises_validation_error(self): + """Passing an instance of a completely different model is a caller error and raises.""" + + class _IncompatibleModel(BaseModel): + name: str + + with pytest.raises(Exception): + CacheCodec.serialize(_IncompatibleModel(name="x"), model_type=_SampleModel) + class TestCacheCodecDeserialize: def test_none_returns_none(self): diff --git a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py new file mode 100644 index 00000000000..f824fd2f9c0 --- /dev/null +++ b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py @@ -0,0 +1,130 @@ +""" +Tests for the enable_redis_auth_cache general_settings flag. + +Verifies that _init_cache attaches Redis to user_api_key_cache only when +the flag is explicitly set to True, and leaves it in-memory-only otherwise. +""" + +from contextlib import contextmanager +from unittest.mock import MagicMock, patch + +import pytest + +import litellm +import litellm.proxy.proxy_server as ps +from litellm.caching.caching import RedisCache +from litellm.caching.dual_cache import DualCache + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +class _FakeRedisCache(RedisCache): + """ + Minimal RedisCache subclass that passes isinstance checks without + requiring a real Redis connection. __init__ is bypassed so no + network calls are made. + """ + + def __init__(self): # noqa: super().__init__ skipped intentionally + pass + + +@contextmanager +def _patched_init_cache(general_settings: dict, cache_params: dict): + """ + Context manager that: + 1. Replaces the module-level globals with fresh DualCache instances. + 2. Patches ``litellm.Cache`` (locally imported inside _init_cache) so + it returns a fake cache whose ``.cache`` attribute is a + _FakeRedisCache (passes the isinstance guard in _init_cache). + 3. Yields (user_api_key_cache, spend_counter_cache) after calling + _init_cache, then restores everything. + """ + fake_redis = _FakeRedisCache() + + mock_litellm_cache = MagicMock() + mock_litellm_cache.cache = fake_redis + + fresh_user_cache = DualCache() + fresh_spend_cache = DualCache() + + with ( + patch.object(ps, "general_settings", general_settings), + patch.object(ps, "user_api_key_cache", fresh_user_cache), + patch.object(ps, "spend_counter_cache", fresh_spend_cache), + patch.object(ps, "llm_router", None), + # Cache is locally imported inside _init_cache: patch it at source. + patch("litellm.Cache", return_value=mock_litellm_cache), + ): + litellm.cache = None + ps.ProxyConfig()._init_cache(cache_params) + yield fresh_user_cache, fresh_spend_cache + + +# --------------------------------------------------------------------------- +# Tests +# --------------------------------------------------------------------------- + + +class TestRedisAuthCacheFlag: + def test_flag_true_attaches_redis_to_user_api_key_cache(self): + """When enable_redis_auth_cache=True, user_api_key_cache.redis_cache must be set.""" + with _patched_init_cache( + general_settings={"enable_redis_auth_cache": True}, + cache_params={"type": "redis", "host": "localhost", "port": 6379}, + ) as (user_cache, _): + assert user_cache.redis_cache is not None, ( + "Redis should be attached to user_api_key_cache when " + "enable_redis_auth_cache=True" + ) + + def test_flag_false_leaves_user_api_key_cache_in_memory_only(self): + """When enable_redis_auth_cache=False, user_api_key_cache must stay in-memory.""" + with _patched_init_cache( + general_settings={"enable_redis_auth_cache": False}, + cache_params={"type": "redis", "host": "localhost", "port": 6379}, + ) as (user_cache, _): + assert user_cache.redis_cache is None, ( + "user_api_key_cache must remain in-memory-only when " + "enable_redis_auth_cache=False" + ) + + def test_flag_absent_leaves_user_api_key_cache_in_memory_only(self): + """When enable_redis_auth_cache is not set at all, default is in-memory-only.""" + with _patched_init_cache( + general_settings={}, + cache_params={"type": "redis", "host": "localhost", "port": 6379}, + ) as (user_cache, _): + assert user_cache.redis_cache is None, ( + "user_api_key_cache must remain in-memory-only when " + "enable_redis_auth_cache is absent from general_settings" + ) + + def test_spend_counter_cache_always_gets_redis_regardless_of_flag(self): + """spend_counter_cache must receive Redis regardless of the auth-cache flag.""" + for flag_value in (True, False, None): + gs = ( + {"enable_redis_auth_cache": flag_value} + if flag_value is not None + else {} + ) + with _patched_init_cache( + general_settings=gs, + cache_params={"type": "redis", "host": "localhost", "port": 6379}, + ) as (_, spend_cache): + assert spend_cache.redis_cache is not None, ( + f"spend_counter_cache must always get Redis " + f"(enable_redis_auth_cache={flag_value!r})" + ) + + def test_flag_false_spend_gets_redis_but_user_cache_does_not(self): + """Explicit False: spend cache wired, auth cache left in-memory.""" + with _patched_init_cache( + general_settings={"enable_redis_auth_cache": False}, + cache_params={"type": "redis", "host": "localhost", "port": 6379}, + ) as (user_cache, spend_cache): + assert spend_cache.redis_cache is not None + assert user_cache.redis_cache is None From 655e75276e156cb28d0df8c214523182076a28af Mon Sep 17 00:00:00 2001 From: harish-berri Date: Fri, 24 Apr 2026 01:44:08 +0000 Subject: [PATCH 15/75] Enhance team endpoint tests by integrating AsyncMock for cache methods, ensuring proper asynchronous behavior in test_update_team_guardrails_with_org_id. This improves test reliability and aligns with recent caching improvements. --- .../proxy/management_endpoints/test_team_endpoints.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index f1e077b5e4a..65015b9d6b6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -5151,6 +5151,9 @@ async def test_update_team_guardrails_with_org_id(): return_value=mock_updated_team ) mock_prisma.jsonify_team_object = MagicMock(side_effect=lambda db_data: db_data) + # async_get_cache must be an AsyncMock so `await` in get_org_object works + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() # Mock llm_router mock_router = MagicMock() From f503c061a59ed329de1967914e4e55124d352b8b Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 24 Apr 2026 09:11:55 +0530 Subject: [PATCH 16/75] Fix black formatting --- litellm/llms/anthropic/chat/transformation.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 4a6743b74be..9aaaa39fd7a 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1584,9 +1584,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): extra_content: Optional[str] = ( json_msg.content if json_msg is not None else None ) - filtered_tools = [ - t for i, t in enumerate(tool_calls) if i not in json_indices - ] + filtered_tools = [t for i, t in enumerate(tool_calls) if i not in json_indices] return None, filtered_tools, extra_content def _transform_response_for_json_mode( From f0ff848e3662808f7c600de69534c76800aa47a9 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 24 Apr 2026 14:36:08 +0530 Subject: [PATCH 17/75] fix(proxy): route azure container file requests by decoded deployment Use decoded managed container model_id to resolve deployment credentials for container file calls and add regressions to verify provider/model metadata decoding and api_base selection. Made-with: Cursor --- .../container_endpoints/handler_factory.py | 111 +++++++++++++---- .../test_azure_container_transformation.py | 117 ++++++++++++++++++ 2 files changed, 205 insertions(+), 23 deletions(-) diff --git a/litellm/proxy/container_endpoints/handler_factory.py b/litellm/proxy/container_endpoints/handler_factory.py index fae7f939aed..60c98ca0a46 100644 --- a/litellm/proxy/container_endpoints/handler_factory.py +++ b/litellm/proxy/container_endpoints/handler_factory.py @@ -22,6 +22,54 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( from litellm.responses.utils import ResponsesAPIRequestUtils +def _decode_container_routing_metadata( + container_id: str, + custom_llm_provider: str, +) -> tuple[str, str, Any]: + """ + Decode managed container IDs into original ID + routing metadata. + """ + decoded = ResponsesAPIRequestUtils._decode_container_id(container_id) + original_container_id = decoded.get("response_id", container_id) + decoded_provider = decoded.get("custom_llm_provider") + decoded_model_id = decoded.get("model_id") + + if decoded_provider and custom_llm_provider == "openai": + custom_llm_provider = decoded_provider + + return original_container_id, custom_llm_provider, decoded_model_id + + +def _get_deployment_credentials_for_model_id(model_id: Any) -> Dict[str, Any]: + """ + Fetch router deployment credentials for a decoded model_id when available. + """ + if not model_id: + return {} + + from litellm.proxy.proxy_server import llm_router + + if llm_router is None: + return {} + + deployment_creds = llm_router.get_deployment_credentials(model_id=model_id) + if deployment_creds: + return deployment_creds + + # Fallback: some ids may map via model group name instead of explicit model id. + deployment = llm_router.get_deployment_by_model_group_name( + model_group_name=str(model_id) + ) + if deployment and deployment.litellm_params: + return deployment.litellm_params.model_dump(exclude_none=True) + + # Final fallback: lookup helper that supports provider-aware resolution. + deployment_creds = llm_router.get_deployment_credentials_with_provider( + model_id=str(model_id) + ) + return deployment_creds or {} + + def _load_endpoints_config() -> Dict: """Load the endpoints configuration from JSON file.""" config_path = Path(__file__).parent.parent.parent / "containers" / "endpoints.json" @@ -176,14 +224,28 @@ async def _process_binary_request( # Build litellm_params - credentials are resolved by provider config from env litellm_params = GenericLiteLLMParams() - # Decode container ID and extract provider info - decoded = ResponsesAPIRequestUtils._decode_container_id(container_id) - original_container_id = decoded.get("response_id", container_id) + ( + original_container_id, + custom_llm_provider, + decoded_model_id, + ) = _decode_container_routing_metadata( + container_id=container_id, + custom_llm_provider=custom_llm_provider, + ) - # If container ID has encoded provider info and user didn't explicitly set provider, use it - decoded_provider = decoded.get("custom_llm_provider") - if decoded_provider and custom_llm_provider == "openai": - custom_llm_provider = decoded_provider + if decoded_model_id and not litellm_params.get("model_id"): + litellm_params["model_id"] = decoded_model_id + + # Resolve provider credentials for routed deployment, e.g. Azure multi-region api_base. + deployment_creds = _get_deployment_credentials_for_model_id( + model_id=decoded_model_id + ) + if deployment_creds.get("api_base"): + litellm_params["api_base"] = deployment_creds["api_base"] + if deployment_creds.get("api_key"): + litellm_params["api_key"] = deployment_creds["api_key"] + if deployment_creds.get("api_version"): + litellm_params["api_version"] = deployment_creds["api_version"] # Get the provider config container_provider_config = _get_container_provider_config(custom_llm_provider) @@ -284,17 +346,19 @@ async def _process_multipart_upload_request( or "openai" ) - # Decode container ID and extract provider info - decoded = ResponsesAPIRequestUtils._decode_container_id(container_id) - original_container_id = decoded.get("response_id", container_id) - - # If container ID has encoded provider info and user didn't explicitly set provider, use it - decoded_provider = decoded.get("custom_llm_provider") - if decoded_provider and custom_llm_provider == "openai": - custom_llm_provider = decoded_provider + ( + original_container_id, + custom_llm_provider, + decoded_model_id, + ) = _decode_container_routing_metadata( + container_id=container_id, + custom_llm_provider=custom_llm_provider, + ) data["container_id"] = original_container_id # Use decoded original ID data["custom_llm_provider"] = custom_llm_provider + if decoded_model_id: + data["model_id"] = decoded_model_id processor = ProxyBaseLLMRequestProcessing(data=data) try: @@ -361,18 +425,19 @@ async def _process_request( # Decode container_id if present in path_params if "container_id" in path_params: - decoded = ResponsesAPIRequestUtils._decode_container_id( - path_params["container_id"] + ( + original_container_id, + custom_llm_provider, + decoded_model_id, + ) = _decode_container_routing_metadata( + container_id=path_params["container_id"], + custom_llm_provider=custom_llm_provider, ) - original_container_id = decoded.get("response_id", path_params["container_id"]) - - # If container ID has encoded provider info and user didn't explicitly set provider, use it - decoded_provider = decoded.get("custom_llm_provider") - if decoded_provider and custom_llm_provider == "openai": - custom_llm_provider = decoded_provider # Update path_params with decoded original ID data["container_id"] = original_container_id + if decoded_model_id: + data["model_id"] = decoded_model_id data["custom_llm_provider"] = custom_llm_provider diff --git a/tests/test_litellm/containers/test_azure_container_transformation.py b/tests/test_litellm/containers/test_azure_container_transformation.py index a46046b318b..4f4006c1584 100644 --- a/tests/test_litellm/containers/test_azure_container_transformation.py +++ b/tests/test_litellm/containers/test_azure_container_transformation.py @@ -11,6 +11,7 @@ sys.path.insert(0, os.path.abspath("../../../")) import litellm from litellm.llms.azure.containers.transformation import AzureContainerConfig from litellm.llms.base_llm.containers.transformation import BaseContainerConfig +from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.containers.main import ( ContainerFileListResponse, ContainerListResponse, @@ -518,3 +519,119 @@ class TestAzureContainerKnownFailureRegressions: c2 = _get_container_provider_config("azure_text") assert type(c1) is type(c2) assert isinstance(c1, AzureContainerConfig) + + def test_regression_decode_container_routing_metadata_includes_model_id(self): + from litellm.proxy.container_endpoints.handler_factory import ( + _decode_container_routing_metadata, + ) + + encoded_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="azure", + model_id="model_abc123", + container_id="cntr_123", + ) + original_id, provider, model_id = _decode_container_routing_metadata( + container_id=encoded_id, + custom_llm_provider="openai", + ) + + assert original_id == "cntr_123" + assert provider == "azure" + assert model_id == "model_abc123" + + @pytest.mark.asyncio + async def test_regression_binary_file_request_uses_deployment_api_base( + self, monkeypatch + ): + from starlette.requests import Request + + from litellm.proxy.container_endpoints import handler_factory + from litellm.types.router import GenericLiteLLMParams + + encoded_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="azure", + model_id="model_abc123", + container_id="cntr_123", + ) + + class _FakeContainerProviderConfig: + def validate_environment(self, headers, api_key=None): + return headers + + def get_complete_url(self, api_base, litellm_params): + return f"{api_base}/openai/v1/containers?api-version=v1" + + def transform_container_file_content_request( + self, container_id, file_id, api_base, litellm_params, headers + ): + return ( + f"{api_base}/{container_id}/files/{file_id}/content", + {}, + ) + + def transform_container_file_content_response( + self, raw_response, logging_obj + ): + return raw_response + + monkeypatch.setattr( + handler_factory, + "_get_container_provider_config", + lambda *_: _FakeContainerProviderConfig(), + ) + monkeypatch.setattr( + handler_factory, + "_get_deployment_credentials_for_model_id", + lambda model_id: { + "api_base": "https://swedencentral.openai.azure.com", + "api_key": "test-key", + "api_version": "v1", + }, + ) + + async def _mock_async_content_handler( + self, + container_id, + file_id, + container_provider_config, + litellm_params: GenericLiteLLMParams, + logging_obj, + ): + assert container_id == "cntr_123" + assert file_id == "cfile_abc" + assert litellm_params.get("model_id") == "model_abc123" + assert ( + litellm_params.get("api_base") + == "https://swedencentral.openai.azure.com" + ) + assert litellm_params.get("api_key") == "test-key" + assert litellm_params.get("api_version") == "v1" + return b"csv-bytes" + + from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler + + monkeypatch.setattr( + BaseLLMHTTPHandler, + "async_container_file_content_handler", + _mock_async_content_handler, + ) + + request = Request( + { + "type": "http", + "method": "GET", + "path": "/v1/containers/id/files/id/content", + "headers": [], + "query_string": b"", + } + ) + + response = await handler_factory._process_binary_request( + request=request, + container_id=encoded_id, + file_id="cfile_abc", + user_api_key_dict=MagicMock(), + ) + + assert response.status_code == 200 + assert response.body == b"csv-bytes" From 05d95fc15a2711a24b7102ed9b9fc6913ceb163e Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 24 Apr 2026 14:39:46 +0530 Subject: [PATCH 18/75] chore(proxy): address greptile feedback on container routing helper Remove redundant model_id guard assignment and drop duplicate provider-aware fallback lookup that repeated earlier router checks. Made-with: Cursor --- litellm/proxy/container_endpoints/handler_factory.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/container_endpoints/handler_factory.py b/litellm/proxy/container_endpoints/handler_factory.py index 60c98ca0a46..8b1740f177a 100644 --- a/litellm/proxy/container_endpoints/handler_factory.py +++ b/litellm/proxy/container_endpoints/handler_factory.py @@ -63,11 +63,7 @@ def _get_deployment_credentials_for_model_id(model_id: Any) -> Dict[str, Any]: if deployment and deployment.litellm_params: return deployment.litellm_params.model_dump(exclude_none=True) - # Final fallback: lookup helper that supports provider-aware resolution. - deployment_creds = llm_router.get_deployment_credentials_with_provider( - model_id=str(model_id) - ) - return deployment_creds or {} + return {} def _load_endpoints_config() -> Dict: @@ -233,7 +229,7 @@ async def _process_binary_request( custom_llm_provider=custom_llm_provider, ) - if decoded_model_id and not litellm_params.get("model_id"): + if decoded_model_id: litellm_params["model_id"] = decoded_model_id # Resolve provider credentials for routed deployment, e.g. Azure multi-region api_base. From d9292e7bcfdf984fcf94a6ffdbe2919e06c964a7 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Fri, 24 Apr 2026 16:13:42 +0000 Subject: [PATCH 19/75] Update test for CacheCodec serialization to clarify validation error handling. --- .../proxy/common_utils/test_cache_codec.py | 13 +++++++++---- 1 file changed, 9 insertions(+), 4 deletions(-) diff --git a/tests/test_litellm/proxy/common_utils/test_cache_codec.py b/tests/test_litellm/proxy/common_utils/test_cache_codec.py index ef52eab4b59..044d4c2d1a7 100644 --- a/tests/test_litellm/proxy/common_utils/test_cache_codec.py +++ b/tests/test_litellm/proxy/common_utils/test_cache_codec.py @@ -78,13 +78,18 @@ class TestCacheCodecSerialize: mock_validate.assert_called_once() def test_with_model_type_incompatible_model_raises_validation_error(self): - """Passing an instance of a completely different model is a caller error and raises.""" + """Passing a BaseModel whose fields don't satisfy model_type's required fields raises. + + _IncompatibleModel only has `foo: int`, so when Pydantic v2 extracts its + data and validates it against _SampleModel (which requires `name: str`), + a ValidationError is raised. + """ class _IncompatibleModel(BaseModel): - name: str + foo: int # missing required 'name' field of _SampleModel - with pytest.raises(Exception): - CacheCodec.serialize(_IncompatibleModel(name="x"), model_type=_SampleModel) + with pytest.raises(ValidationError): + CacheCodec.serialize(_IncompatibleModel(foo=1), model_type=_SampleModel) class TestCacheCodecDeserialize: From ce087aa796f05bced22e297953f99baabc4fba08 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Fri, 24 Apr 2026 23:44:57 +0000 Subject: [PATCH 20/75] refactor(proxy): update documentation and code references for enable_redis_auth_cache - Removed the deprecated enable_redis_auth_cache field from ConfigGeneralSettings. - Updated references in proxy_server.py to use litellm_settings for enable_redis_auth_cache. - Clarified documentation to reflect the new configuration structure and its impact on Redis integration for user_api_key_cache. --- docs/my-website/docs/proxy/config_settings.md | 8 ++++---- litellm/proxy/_types.py | 13 ------------- litellm/proxy/proxy_server.py | 13 ++++++++++--- 3 files changed, 14 insertions(+), 20 deletions(-) diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index 33b6006ac2c..cc9c9de58f5 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -307,7 +307,6 @@ router_settings: | token_rate_limit_type | string | Rate limit counting method: "total", "output", or "input" tokens | | use_redis_transaction_buffer | boolean | If true, buffers database transactions in Redis before writing | | use_shared_health_check | boolean | If true, uses Redis-backed shared health check state across multiple proxy instances | -| enable_redis_auth_cache | boolean | **[Beta]** When `true`, attaches Redis to the virtual-key auth cache (`user_api_key_cache`) so all proxy workers/pods share the same cache instead of each pod resolving keys independently against the database. Requires `litellm_settings.cache: true` with a Redis backend. Significantly reduces database load in multi-worker deployments by eliminating per-pod cache misses on the `combined_view` query. Off by default for a safe phased rollout — enable once your Redis cluster is healthy. Will become the default in a future release. See [Redis Auth Cache](#redis-auth-cache-multi-worker-db-load-reduction). | | user_header_mappings | dict | Map custom request headers to user IDs using lookup rules | | user_header_name | string | HTTP header name to extract user identity from requests | @@ -323,21 +322,22 @@ In multi-worker or multi-pod deployments each worker process keeps its own **in- ### Solution -Setting `enable_redis_auth_cache: true` attaches Redis to `user_api_key_cache` so the resolved key object is stored in a **shared** Redis cache. A cache hit on any worker prevents the DB query entirely. +Setting `enable_redis_auth_cache: true` under `litellm_settings` attaches Redis to `user_api_key_cache` so the resolved key object is stored in a **shared** Redis cache. A cache hit on any worker prevents the DB query entirely. ```yaml # config.yaml litellm_settings: + enable_redis_auth_cache: true # ← share the auth cache across workers cache: true cache_params: type: redis host: os.environ/REDIS_HOST port: os.environ/REDIS_PORT + max_connections: 100 # optional — tune per-worker pool size (default: 50) general_settings: master_key: sk-1234 - enable_redis_auth_cache: true # ← share the auth cache across workers ``` ### Requirements @@ -366,7 +366,7 @@ enable_redis_auth_cache=True: attached Redis to user_api_key_cache — virtual-k If the flag is off you will see: ``` -enable_redis_auth_cache is not set: user_api_key_cache remains in-memory only (per-worker). Set general_settings.enable_redis_auth_cache: true to share the auth cache across workers and reduce DB load. +enable_redis_auth_cache is not set: user_api_key_cache remains in-memory only (per-worker). Set litellm_settings.enable_redis_auth_cache: true to share the auth cache across workers and reduce DB load. ``` --- diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index ee7e7ce9463..b8e0aa0d128 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2383,19 +2383,6 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="List of MCP server fields that must be filled in for a submission to pass standards checks (e.g. ['description', 'source_url', 'alias']).", ) - enable_redis_auth_cache: Optional[bool] = Field( - None, - description=( - "When True, attaches Redis to user_api_key_cache so virtual-key lookups " - "are shared across all proxy workers/pods instead of being resolved " - "per-process. Requires a Redis cache to be configured under litellm_settings. " - "Reduces DB load significantly in multi-worker deployments by eliminating " - "redundant combined_view SQL queries caused by per-pod cache misses. " - "Off by default for a safe phased rollout — set to True once your Redis " - "cluster is healthy and the CacheCodec serialisation has been validated in " - "your environment. Will be enabled by default in a future release." - ), - ) class ConfigYAML(LiteLLMPydanticObjectBase): diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 73bc46ba391..8d199d2c092 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2979,6 +2979,7 @@ class ProxyConfig: def _init_cache( self, cache_params: dict, + enable_redis_auth_cache: bool = False, ): global redis_usage_cache, llm_router, general_settings from litellm import Cache @@ -3002,7 +3003,7 @@ class ProxyConfig: ) # Note: PKCE verifier storage uses redis_usage_cache directly (not # user_api_key_cache) to avoid routing all API-key lookups through Redis. - if general_settings.get("enable_redis_auth_cache") is True: + if enable_redis_auth_cache is True: user_api_key_cache.attach_redis_cache( redis_usage_cache, default_redis_ttl=litellm.default_redis_ttl, @@ -3016,7 +3017,7 @@ class ProxyConfig: verbose_proxy_logger.info( "enable_redis_auth_cache is not set: user_api_key_cache " "remains in-memory only (per-worker). Set " - "general_settings.enable_redis_auth_cache: true to share " + "litellm_settings.enable_redis_auth_cache: true to share " "the auth cache across workers and reduce DB load." ) @@ -3337,7 +3338,13 @@ class ProxyConfig: cache_params[key] = get_secret(value) ## to pass a complete url, or set ssl=True, etc. just set it as `os.environ[REDIS_URL] = `, _redis.py checks for REDIS specific environment variables - self._init_cache(cache_params=cache_params) + self._init_cache( + cache_params=cache_params, + enable_redis_auth_cache=litellm_settings.get( + "enable_redis_auth_cache", False + ) + is True, + ) if litellm.cache is not None: verbose_proxy_logger.debug( f"{blue_color_code}Set Cache on LiteLLM Proxy{reset_color_code}" From 77847752fafa253aeac2e34b4b3ceacc3ba16c6b Mon Sep 17 00:00:00 2001 From: harish-berri Date: Sat, 25 Apr 2026 00:21:33 +0000 Subject: [PATCH 21/75] feat(auth): add Prometheus metrics for combined_view SQL queries - Introduced `litellm_auth_combined_view_queries_total` metric to track the number of combined_view SQL queries issued for virtual-key authentication. - Added `AuthMetrics` class to encapsulate metric incrementing logic. - Updated relevant code to increment the new metric on database lookups, aiding in the validation of Redis cache effectiveness. - Updated Prometheus metric definitions to include the new auth diagnostic metrics. --- litellm/integrations/prometheus.py | 14 +++++++ litellm/proxy/auth/auth_checks.py | 2 + litellm/proxy/auth/auth_metrics.py | 52 ++++++++++++++++++++++++ litellm/types/integrations/prometheus.py | 7 ++++ 4 files changed, 75 insertions(+) create mode 100644 litellm/proxy/auth/auth_metrics.py diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 723b142dfad..1a899cf0e83 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -265,6 +265,20 @@ class PrometheusLogger(CustomLogger): ######################################## # LiteLLM Virtual API KEY metrics ######################################## + + # Auth DB load diagnostic: count direct combined_view SQL queries. + # Each increment means a virtual-key cache miss that hit the DB. + # Useful for validating that enable_redis_auth_cache is working. + self.litellm_auth_combined_view_queries_total = self._counter_factory( + "litellm_auth_combined_view_queries_total", + "Number of times the combined_view SQL query was issued for virtual-key auth. " + "Each count is a cache miss that hit the database. Use to validate " + "enable_redis_auth_cache is reducing DB load.", + labelnames=self.get_labels_for_metric( + "litellm_auth_combined_view_queries_total" + ), + ) + # Remaining MODEL RPM limit for API Key self.litellm_remaining_api_key_requests_for_model = self._gauge_factory( "litellm_remaining_api_key_requests_for_model", diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 7b1d1c234d1..3cce911cd0c 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -60,6 +60,7 @@ from litellm.proxy._types import ( SpecialModelNames, UserAPIKeyAuth, ) +from litellm.proxy.auth.auth_metrics import AuthMetrics from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.guardrails.tool_name_extraction import ( @@ -2257,6 +2258,7 @@ async def _fetch_key_object_from_db_with_reconnect( Fetch key object from DB and retry once if a DB connection error can be healed. """ try: + AuthMetrics.inc_combined_view_query(hashed_token) return await prisma_client.get_data( token=hashed_token, table_name="combined_view", diff --git a/litellm/proxy/auth/auth_metrics.py b/litellm/proxy/auth/auth_metrics.py new file mode 100644 index 00000000000..0b463cc920f --- /dev/null +++ b/litellm/proxy/auth/auth_metrics.py @@ -0,0 +1,52 @@ +""" +Prometheus metric helpers for the auth layer. + +All metrics are thin wrappers around the shared ``PrometheusLogger`` instance so +that every counter follows the same registration path (``_counter_factory``, +label-filter config) as the rest of LiteLLM's metrics. + +Usage:: + + from litellm.proxy.auth.auth_metrics import AuthMetrics + + AuthMetrics.inc_combined_view_query(hashed_token="sk-xxx") +""" + +from litellm._logging import verbose_proxy_logger + + +class AuthMetrics: + """Static helpers for incrementing auth-layer Prometheus counters.""" + + @staticmethod + def _get_prom(): + """Return the active PrometheusLogger, or None if Prometheus is not configured.""" + try: + from litellm.router_utils.cooldown_callbacks import ( + _get_prometheus_logger_from_callbacks, + ) + + return _get_prometheus_logger_from_callbacks() + except Exception: + return None + + @staticmethod + def inc_combined_view_query(hashed_token: str) -> None: + """ + Increment ``litellm_auth_combined_view_queries_total``. + + Called once per virtual-key DB lookup (combined_view query). Each + increment represents a cache miss that hit the database — use this to + validate that ``enable_redis_auth_cache`` is reducing DB load. + """ + try: + prom = AuthMetrics._get_prom() + if prom is not None: + prom.litellm_auth_combined_view_queries_total.labels( + hashed_api_key=hashed_token + ).inc() + except Exception as e: + verbose_proxy_logger.debug( + "AuthMetrics.inc_combined_view_query: failed to increment counter: %s", + e, + ) diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 43a287f29bc..a4573535f72 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -228,6 +228,8 @@ DEFINED_PROMETHEUS_METRICS = Literal[ "litellm_guardrail_latency_seconds", "litellm_guardrail_errors_total", "litellm_guardrail_requests_total", + # Auth DB diagnostic metrics + "litellm_auth_combined_view_queries_total", # Cache metrics "litellm_cache_hits_metric", "litellm_cache_misses_metric", @@ -307,6 +309,11 @@ class PrometheusMetricLabels: litellm_guardrail_errors_total: List[str] = [] litellm_guardrail_requests_total: List[str] = [] + # Auth DB diagnostic - label by key so you can see which virtual key causes DB hits + litellm_auth_combined_view_queries_total = [ + UserAPIKeyLabelNames.API_KEY_HASH.value, + ] + litellm_proxy_total_requests_metric = [ UserAPIKeyLabelNames.END_USER.value, UserAPIKeyLabelNames.API_KEY_HASH.value, From f08a8534ab199e3adfb1cb190a3a6d994e987a0a Mon Sep 17 00:00:00 2001 From: harish-berri Date: Sat, 25 Apr 2026 00:44:59 +0000 Subject: [PATCH 22/75] fix failing test --- .../proxy/test_redis_auth_cache_flag.py | 27 ++++++++++--------- 1 file changed, 15 insertions(+), 12 deletions(-) diff --git a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py index f824fd2f9c0..261222be9a2 100644 --- a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py +++ b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py @@ -1,5 +1,5 @@ """ -Tests for the enable_redis_auth_cache general_settings flag. +Tests for the enable_redis_auth_cache litellm_settings flag. Verifies that _init_cache attaches Redis to user_api_key_cache only when the flag is explicitly set to True, and leaves it in-memory-only otherwise. @@ -33,14 +33,16 @@ class _FakeRedisCache(RedisCache): @contextmanager -def _patched_init_cache(general_settings: dict, cache_params: dict): +def _patched_init_cache(litellm_settings: dict, cache_params: dict): """ Context manager that: 1. Replaces the module-level globals with fresh DualCache instances. 2. Patches ``litellm.Cache`` (locally imported inside _init_cache) so it returns a fake cache whose ``.cache`` attribute is a _FakeRedisCache (passes the isinstance guard in _init_cache). - 3. Yields (user_api_key_cache, spend_counter_cache) after calling + 3. Extracts enable_redis_auth_cache from litellm_settings and passes it + as the second argument to _init_cache (matching production behaviour). + 4. Yields (user_api_key_cache, spend_counter_cache) after calling _init_cache, then restores everything. """ fake_redis = _FakeRedisCache() @@ -51,8 +53,9 @@ def _patched_init_cache(general_settings: dict, cache_params: dict): fresh_user_cache = DualCache() fresh_spend_cache = DualCache() + enable_redis_auth_cache = litellm_settings.get("enable_redis_auth_cache", False) + with ( - patch.object(ps, "general_settings", general_settings), patch.object(ps, "user_api_key_cache", fresh_user_cache), patch.object(ps, "spend_counter_cache", fresh_spend_cache), patch.object(ps, "llm_router", None), @@ -60,7 +63,7 @@ def _patched_init_cache(general_settings: dict, cache_params: dict): patch("litellm.Cache", return_value=mock_litellm_cache), ): litellm.cache = None - ps.ProxyConfig()._init_cache(cache_params) + ps.ProxyConfig()._init_cache(cache_params, enable_redis_auth_cache) yield fresh_user_cache, fresh_spend_cache @@ -73,7 +76,7 @@ class TestRedisAuthCacheFlag: def test_flag_true_attaches_redis_to_user_api_key_cache(self): """When enable_redis_auth_cache=True, user_api_key_cache.redis_cache must be set.""" with _patched_init_cache( - general_settings={"enable_redis_auth_cache": True}, + litellm_settings={"enable_redis_auth_cache": True}, cache_params={"type": "redis", "host": "localhost", "port": 6379}, ) as (user_cache, _): assert user_cache.redis_cache is not None, ( @@ -84,7 +87,7 @@ class TestRedisAuthCacheFlag: def test_flag_false_leaves_user_api_key_cache_in_memory_only(self): """When enable_redis_auth_cache=False, user_api_key_cache must stay in-memory.""" with _patched_init_cache( - general_settings={"enable_redis_auth_cache": False}, + litellm_settings={"enable_redis_auth_cache": False}, cache_params={"type": "redis", "host": "localhost", "port": 6379}, ) as (user_cache, _): assert user_cache.redis_cache is None, ( @@ -95,24 +98,24 @@ class TestRedisAuthCacheFlag: def test_flag_absent_leaves_user_api_key_cache_in_memory_only(self): """When enable_redis_auth_cache is not set at all, default is in-memory-only.""" with _patched_init_cache( - general_settings={}, + litellm_settings={}, cache_params={"type": "redis", "host": "localhost", "port": 6379}, ) as (user_cache, _): assert user_cache.redis_cache is None, ( "user_api_key_cache must remain in-memory-only when " - "enable_redis_auth_cache is absent from general_settings" + "enable_redis_auth_cache is absent from litellm_settings" ) def test_spend_counter_cache_always_gets_redis_regardless_of_flag(self): """spend_counter_cache must receive Redis regardless of the auth-cache flag.""" for flag_value in (True, False, None): - gs = ( + ls = ( {"enable_redis_auth_cache": flag_value} if flag_value is not None else {} ) with _patched_init_cache( - general_settings=gs, + litellm_settings=ls, cache_params={"type": "redis", "host": "localhost", "port": 6379}, ) as (_, spend_cache): assert spend_cache.redis_cache is not None, ( @@ -123,7 +126,7 @@ class TestRedisAuthCacheFlag: def test_flag_false_spend_gets_redis_but_user_cache_does_not(self): """Explicit False: spend cache wired, auth cache left in-memory.""" with _patched_init_cache( - general_settings={"enable_redis_auth_cache": False}, + litellm_settings={"enable_redis_auth_cache": False}, cache_params={"type": "redis", "host": "localhost", "port": 6379}, ) as (user_cache, spend_cache): assert spend_cache.redis_cache is not None From 33371284435c34e85a0fafc798fbdaa4d9b26a20 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 27 Apr 2026 09:25:42 +0530 Subject: [PATCH 23/75] Fix black --- .../vertex_embeddings/embedding_handler.py | 56 ++++++++----------- 1 file changed, 23 insertions(+), 33 deletions(-) diff --git a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py index 729cc9c3ead..0e7afd5da3f 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py @@ -65,9 +65,7 @@ class VertexEmbedding(VertexBase): litellm_params=litellm_params, ) - should_use_v1beta1_features = self.is_using_v1beta1_features( - optional_params=optional_params - ) + should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params) _auth_header, vertex_project = self._ensure_access_token( credentials=vertex_credentials, @@ -92,11 +90,13 @@ class VertexEmbedding(VertexBase): use_psc_endpoint_format=use_psc_endpoint_format, ) headers = self.set_headers(auth_header=auth_header, extra_headers=extra_headers) - vertex_request: VertexEmbeddingRequest = litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request( - input=input, - optional_params=optional_params, - model=model, - litellm_params=litellm_params, + vertex_request: VertexEmbeddingRequest = ( + litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request( + input=input, + optional_params=optional_params, + model=model, + litellm_params=litellm_params, + ) ) _client_params = {} @@ -128,14 +128,10 @@ class VertexEmbedding(VertexBase): _json_response = response.json() ## LOGGING POST-CALL - logging_obj.post_call( - input=input, api_key=None, original_response=_json_response - ) + logging_obj.post_call(input=input, api_key=None, original_response=_json_response) - model_response = ( - litellm.vertexAITextEmbeddingConfig.transform_vertex_response_to_openai( - response=_json_response, model=model, model_response=model_response - ) + model_response = litellm.vertexAITextEmbeddingConfig.transform_vertex_response_to_openai( + response=_json_response, model=model, model_response=model_response ) return model_response @@ -164,9 +160,7 @@ class VertexEmbedding(VertexBase): """ Async embedding implementation """ - should_use_v1beta1_features = self.is_using_v1beta1_features( - optional_params=optional_params - ) + should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params) _auth_header, vertex_project = await self._ensure_access_token_async( credentials=vertex_credentials, project_id=vertex_project, @@ -190,20 +184,20 @@ class VertexEmbedding(VertexBase): use_psc_endpoint_format=use_psc_endpoint_format, ) headers = self.set_headers(auth_header=auth_header, extra_headers=extra_headers) - vertex_request: VertexEmbeddingRequest = litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request( - input=input, - optional_params=optional_params, - model=model, - litellm_params=litellm_params, + vertex_request: VertexEmbeddingRequest = ( + litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request( + input=input, + optional_params=optional_params, + model=model, + litellm_params=litellm_params, + ) ) _async_client_params = {} if timeout: _async_client_params["timeout"] = timeout if client is None or not isinstance(client, AsyncHTTPHandler): - client = get_async_httpx_client( - params=_async_client_params, llm_provider=litellm.LlmProviders.VERTEX_AI - ) + client = get_async_httpx_client(params=_async_client_params, llm_provider=litellm.LlmProviders.VERTEX_AI) else: client = client # type: ignore ## LOGGING @@ -228,14 +222,10 @@ class VertexEmbedding(VertexBase): _json_response = response.json() ## LOGGING POST-CALL - logging_obj.post_call( - input=input, api_key=None, original_response=_json_response - ) + logging_obj.post_call(input=input, api_key=None, original_response=_json_response) - model_response = ( - litellm.vertexAITextEmbeddingConfig.transform_vertex_response_to_openai( - response=_json_response, model=model, model_response=model_response - ) + model_response = litellm.vertexAITextEmbeddingConfig.transform_vertex_response_to_openai( + response=_json_response, model=model, model_response=model_response ) return model_response From 920443c4c931b73386eec609910f45f2a50c6d71 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 27 Apr 2026 11:29:42 +0530 Subject: [PATCH 24/75] Fix all container endpoints --- .../container_endpoints/handler_factory.py | 179 +++++------------- litellm/router.py | 25 ++- .../test_router_endpoints.py | 45 ++++- .../test_azure_container_transformation.py | 157 +++++++-------- 4 files changed, 197 insertions(+), 209 deletions(-) diff --git a/litellm/proxy/container_endpoints/handler_factory.py b/litellm/proxy/container_endpoints/handler_factory.py index 8b1740f177a..465d1fa9d79 100644 --- a/litellm/proxy/container_endpoints/handler_factory.py +++ b/litellm/proxy/container_endpoints/handler_factory.py @@ -19,51 +19,6 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_headers, get_custom_llm_provider_from_request_query, ) -from litellm.responses.utils import ResponsesAPIRequestUtils - - -def _decode_container_routing_metadata( - container_id: str, - custom_llm_provider: str, -) -> tuple[str, str, Any]: - """ - Decode managed container IDs into original ID + routing metadata. - """ - decoded = ResponsesAPIRequestUtils._decode_container_id(container_id) - original_container_id = decoded.get("response_id", container_id) - decoded_provider = decoded.get("custom_llm_provider") - decoded_model_id = decoded.get("model_id") - - if decoded_provider and custom_llm_provider == "openai": - custom_llm_provider = decoded_provider - - return original_container_id, custom_llm_provider, decoded_model_id - - -def _get_deployment_credentials_for_model_id(model_id: Any) -> Dict[str, Any]: - """ - Fetch router deployment credentials for a decoded model_id when available. - """ - if not model_id: - return {} - - from litellm.proxy.proxy_server import llm_router - - if llm_router is None: - return {} - - deployment_creds = llm_router.get_deployment_credentials(model_id=model_id) - if deployment_creds: - return deployment_creds - - # Fallback: some ids may map via model group name instead of explicit model id. - deployment = llm_router.get_deployment_by_model_group_name( - model_group_name=str(model_id) - ) - if deployment and deployment.litellm_params: - return deployment.litellm_params.model_dump(exclude_none=True) - - return {} def _load_endpoints_config() -> Dict: @@ -108,10 +63,12 @@ def _create_handler_for_path_params( request: Request, container_id: str, file_id: str, + fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): return await _process_binary_request( request=request, + fastapi_response=fastapi_response, container_id=container_id, file_id=file_id, user_api_key_dict=user_api_key_dict, @@ -196,77 +153,61 @@ def _create_handler_for_path_params( async def _process_binary_request( request: Request, + fastapi_response: Response, container_id: str, file_id: str, user_api_key_dict: UserAPIKeyAuth, ): """ - Process binary content requests using the proper transformation pattern. + Process binary content requests through the standard proxy/router pipeline. - This uses the provider config transformations and llm_http_handler - to maintain consistency with the established pattern. + The router owns managed container ID decoding and deployment selection. This + handler only adapts the byte response to FastAPI. """ - from litellm.litellm_core_utils.litellm_logging import Logging - from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler - from litellm.types.router import GenericLiteLLMParams + from litellm.proxy.proxy_server import ( + general_settings, + llm_router, + proxy_config, + proxy_logging_obj, + select_data_generator, + user_api_base, + user_max_tokens, + user_model, + user_request_timeout, + user_temperature, + version, + ) - # Extract custom_llm_provider custom_llm_provider = ( get_custom_llm_provider_from_request_headers(request=request) or get_custom_llm_provider_from_request_query(request=request) or "openai" ) - - # Build litellm_params - credentials are resolved by provider config from env - litellm_params = GenericLiteLLMParams() - - ( - original_container_id, - custom_llm_provider, - decoded_model_id, - ) = _decode_container_routing_metadata( - container_id=container_id, - custom_llm_provider=custom_llm_provider, - ) - - if decoded_model_id: - litellm_params["model_id"] = decoded_model_id - - # Resolve provider credentials for routed deployment, e.g. Azure multi-region api_base. - deployment_creds = _get_deployment_credentials_for_model_id( - model_id=decoded_model_id - ) - if deployment_creds.get("api_base"): - litellm_params["api_base"] = deployment_creds["api_base"] - if deployment_creds.get("api_key"): - litellm_params["api_key"] = deployment_creds["api_key"] - if deployment_creds.get("api_version"): - litellm_params["api_version"] = deployment_creds["api_version"] - - # Get the provider config - container_provider_config = _get_container_provider_config(custom_llm_provider) - - # Create logging object - logging_obj = Logging( - model="container-file-content", - messages=[], - stream=False, - call_type="container_file_content", - start_time=None, - litellm_call_id="", - function_id="", - ) - - # Use the HTTP handler to make the request - handler = BaseLLMHTTPHandler() + data: Dict[str, Any] = { + "container_id": container_id, + "file_id": file_id, + "custom_llm_provider": custom_llm_provider, + } + processor = ProxyBaseLLMRequestProcessing(data=data) try: - content = await handler.async_container_file_content_handler( - container_id=original_container_id, # Use decoded original ID - file_id=file_id, - container_provider_config=container_provider_config, - litellm_params=litellm_params, - logging_obj=logging_obj, + content = await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="aretrieve_container_file_content", + proxy_logging_obj=proxy_logging_obj, + llm_router=llm_router, + general_settings=general_settings, + proxy_config=proxy_config, + select_data_generator=select_data_generator, + model=None, + user_model=user_model, + user_temperature=user_temperature, + user_request_timeout=user_request_timeout, + user_max_tokens=user_max_tokens, + user_api_base=user_api_base, + version=version, ) # Determine content type based on common file extensions in the file_id @@ -293,7 +234,12 @@ async def _process_binary_request( ) except Exception as e: - raise e + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + version=version, + ) async def _process_multipart_upload_request( @@ -342,19 +288,8 @@ async def _process_multipart_upload_request( or "openai" ) - ( - original_container_id, - custom_llm_provider, - decoded_model_id, - ) = _decode_container_routing_metadata( - container_id=container_id, - custom_llm_provider=custom_llm_provider, - ) - - data["container_id"] = original_container_id # Use decoded original ID + data["container_id"] = container_id data["custom_llm_provider"] = custom_llm_provider - if decoded_model_id: - data["model_id"] = decoded_model_id processor = ProxyBaseLLMRequestProcessing(data=data) try: @@ -419,22 +354,6 @@ async def _process_request( or "openai" ) - # Decode container_id if present in path_params - if "container_id" in path_params: - ( - original_container_id, - custom_llm_provider, - decoded_model_id, - ) = _decode_container_routing_metadata( - container_id=path_params["container_id"], - custom_llm_provider=custom_llm_provider, - ) - - # Update path_params with decoded original ID - data["container_id"] = original_container_id - if decoded_model_id: - data["model_id"] = decoded_model_id - data["custom_llm_provider"] = custom_llm_provider processor = ProxyBaseLLMRequestProcessing(data=data) diff --git a/litellm/router.py b/litellm/router.py index b275c264ebc..a11281bb151 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -5261,11 +5261,32 @@ class Router: """ Initialize the Containers API endpoints on the router. - Container operations don't need model-based routing, so we call the - original function directly with the custom_llm_provider. + LiteLLM-managed container IDs (``cntr_...``) encode ``model_id`` and provider + metadata. When present, decode the ID, replace ``container_id`` with the + upstream value, and route through ``_ageneric_api_call_with_fallbacks`` so + deployment credentials (e.g. regional ``api_base`` for Azure) match + :meth:`_init_responses_api_endpoints`. Otherwise call the handler directly. """ if custom_llm_provider and "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = custom_llm_provider + + from litellm.responses.utils import ResponsesAPIRequestUtils + + container_id = kwargs.get("container_id") + if isinstance(container_id, str): + decoded = ResponsesAPIRequestUtils._decode_container_id(container_id) + model_id = decoded.get("model_id") + if model_id: + kwargs["container_id"] = decoded.get("response_id", container_id) + kwargs["model"] = model_id + decoded_provider = decoded.get("custom_llm_provider") + if decoded_provider and kwargs.get("custom_llm_provider") == "openai": + kwargs["custom_llm_provider"] = decoded_provider + return await self._ageneric_api_call_with_fallbacks( + original_function=original_function, + **kwargs, + ) + return await original_function(**kwargs) async def _init_responses_api_endpoints( diff --git a/tests/router_unit_tests/test_router_endpoints.py b/tests/router_unit_tests/test_router_endpoints.py index b93502e8152..93bb84dfe35 100644 --- a/tests/router_unit_tests/test_router_endpoints.py +++ b/tests/router_unit_tests/test_router_endpoints.py @@ -1110,7 +1110,7 @@ def test_initialize_skills_endpoints(): async def test_init_containers_api_endpoints(): """ Test that _init_containers_api_endpoints calls the original function - directly without model-based routing. + directly when there is no managed container ID (no embedded model_id). """ router = Router(model_list=[]) @@ -1127,3 +1127,46 @@ async def test_init_containers_api_endpoints(): custom_llm_provider="openai", name="Test Container" ) assert result == mock_response + + +@pytest.mark.asyncio +async def test_init_containers_api_endpoints_managed_id_routes_via_generic_fallbacks(): + """ + Managed ``cntr_`` IDs embed ``model_id``; router should decode and use + ``_ageneric_api_call_with_fallbacks`` so deployment credentials apply. + """ + from litellm.responses.utils import ResponsesAPIRequestUtils + + router = Router( + model_list=[ + { + "model_name": "azure-router-model", + "litellm_params": { + "model": "azure/gpt-4", + "api_key": "fake-key", + "api_base": "https://westus.api.cognitive.microsoft.com", + }, + } + ] + ) + router._ageneric_api_call_with_fallbacks = AsyncMock() + + managed_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="azure", + model_id="azure-router-model", + container_id="cfile_upstream_abc", + ) + + await router._init_containers_api_endpoints( + original_function=AsyncMock(), + custom_llm_provider="openai", + container_id=managed_id, + file_id="cfile_xyz", + ) + + router._ageneric_api_call_with_fallbacks.assert_called_once() + call_kw = router._ageneric_api_call_with_fallbacks.call_args.kwargs + assert call_kw["model"] == "azure-router-model" + assert call_kw["container_id"] == "cfile_upstream_abc" + assert call_kw["file_id"] == "cfile_xyz" + assert call_kw["custom_llm_provider"] == "azure" diff --git a/tests/test_litellm/containers/test_azure_container_transformation.py b/tests/test_litellm/containers/test_azure_container_transformation.py index 4f4006c1584..94ab1ae6b3d 100644 --- a/tests/test_litellm/containers/test_azure_container_transformation.py +++ b/tests/test_litellm/containers/test_azure_container_transformation.py @@ -520,100 +520,39 @@ class TestAzureContainerKnownFailureRegressions: assert type(c1) is type(c2) assert isinstance(c1, AzureContainerConfig) - def test_regression_decode_container_routing_metadata_includes_model_id(self): - from litellm.proxy.container_endpoints.handler_factory import ( - _decode_container_routing_metadata, - ) - - encoded_id = ResponsesAPIRequestUtils._build_container_id( - custom_llm_provider="azure", - model_id="model_abc123", - container_id="cntr_123", - ) - original_id, provider, model_id = _decode_container_routing_metadata( - container_id=encoded_id, - custom_llm_provider="openai", - ) - - assert original_id == "cntr_123" - assert provider == "azure" - assert model_id == "model_abc123" - @pytest.mark.asyncio - async def test_regression_binary_file_request_uses_deployment_api_base( + async def test_proxy_process_request_preserves_managed_container_id( self, monkeypatch ): from starlette.requests import Request from litellm.proxy.container_endpoints import handler_factory - from litellm.types.router import GenericLiteLLMParams encoded_id = ResponsesAPIRequestUtils._build_container_id( custom_llm_provider="azure", model_id="model_abc123", container_id="cntr_123", ) + captured = {} - class _FakeContainerProviderConfig: - def validate_environment(self, headers, api_key=None): - return headers - - def get_complete_url(self, api_base, litellm_params): - return f"{api_base}/openai/v1/containers?api-version=v1" - - def transform_container_file_content_request( - self, container_id, file_id, api_base, litellm_params, headers - ): - return ( - f"{api_base}/{container_id}/files/{file_id}/content", - {}, - ) - - def transform_container_file_content_response( - self, raw_response, logging_obj - ): - return raw_response - - monkeypatch.setattr( - handler_factory, - "_get_container_provider_config", - lambda *_: _FakeContainerProviderConfig(), - ) - monkeypatch.setattr( - handler_factory, - "_get_deployment_credentials_for_model_id", - lambda model_id: { - "api_base": "https://swedencentral.openai.azure.com", - "api_key": "test-key", - "api_version": "v1", - }, - ) - - async def _mock_async_content_handler( + async def _mock_base_process_llm_request( self, - container_id, - file_id, - container_provider_config, - litellm_params: GenericLiteLLMParams, - logging_obj, + request, + fastapi_response, + user_api_key_dict, + route_type, + **kwargs, ): - assert container_id == "cntr_123" - assert file_id == "cfile_abc" - assert litellm_params.get("model_id") == "model_abc123" - assert ( - litellm_params.get("api_base") - == "https://swedencentral.openai.azure.com" - ) - assert litellm_params.get("api_key") == "test-key" - assert litellm_params.get("api_version") == "v1" - return b"csv-bytes" + captured["data"] = self.data + captured["route_type"] = route_type + return {"id": "cfile_abc"} - from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler + from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing monkeypatch.setattr( - BaseLLMHTTPHandler, - "async_container_file_content_handler", - _mock_async_content_handler, + ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + _mock_base_process_llm_request, ) request = Request( @@ -625,13 +564,79 @@ class TestAzureContainerKnownFailureRegressions: "query_string": b"", } ) + fastapi_response = MagicMock() + + await handler_factory._process_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=MagicMock(), + route_type="alist_container_files", + path_params={"container_id": encoded_id}, + ) + + assert captured["route_type"] == "alist_container_files" + assert captured["data"]["container_id"] == encoded_id + assert captured["data"]["custom_llm_provider"] == "openai" + assert "model_id" not in captured["data"] + assert "api_base" not in captured["data"] + + @pytest.mark.asyncio + async def test_regression_binary_file_request_routes_through_proxy_processor( + self, monkeypatch + ): + from starlette.requests import Request + + from litellm.proxy.container_endpoints import handler_factory + + encoded_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="azure", + model_id="model_abc123", + container_id="cntr_123", + ) + captured = {} + + async def _mock_base_process_llm_request( + self, + request, + fastapi_response, + user_api_key_dict, + route_type, + **kwargs, + ): + captured["data"] = self.data + captured["route_type"] = route_type + return b"csv-bytes" + + from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + + monkeypatch.setattr( + ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + _mock_base_process_llm_request, + ) + + request = Request( + { + "type": "http", + "method": "GET", + "path": "/v1/containers/id/files/id/content", + "headers": [], + "query_string": b"", + } + ) + fastapi_response = MagicMock() response = await handler_factory._process_binary_request( request=request, + fastapi_response=fastapi_response, container_id=encoded_id, file_id="cfile_abc", user_api_key_dict=MagicMock(), ) + assert captured["route_type"] == "aretrieve_container_file_content" + assert captured["data"]["container_id"] == encoded_id + assert captured["data"]["file_id"] == "cfile_abc" + assert captured["data"]["custom_llm_provider"] == "openai" assert response.status_code == 200 assert response.body == b"csv-bytes" From 5d2b7bf613f4788fdfdecdb292ee316dc8c0c9b1 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 27 Apr 2026 11:39:17 +0530 Subject: [PATCH 25/75] Fix mypy --- litellm/proxy/common_request_processing.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 97801baaf0c..dae63f87107 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -683,6 +683,11 @@ class ProxyBaseLLMRequestProcessing: "aingest", "aretrieve_container", "adelete_container", + "aupload_container_file", + "alist_container_files", + "aretrieve_container_file", + "adelete_container_file", + "aretrieve_container_file_content", "acreate_skill", "alist_skills", "aget_skill", From 153223da8487dd546a9ac952ed89028a423cd3e6 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 27 Apr 2026 11:48:58 +0530 Subject: [PATCH 26/75] Fix mypy --- litellm/proxy/common_request_processing.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index dae63f87107..0325263d6b8 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -945,6 +945,11 @@ class ProxyBaseLLMRequestProcessing: "aingest", "aretrieve_container", "adelete_container", + "aupload_container_file", + "alist_container_files", + "aretrieve_container_file", + "adelete_container_file", + "aretrieve_container_file_content", "acreate_skill", "alist_skills", "aget_skill", From 09b648aa5745f89857729fe242b2781ddfd4397c Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 27 Apr 2026 19:01:47 +0530 Subject: [PATCH 27/75] fix lint --- .../vertex_embeddings/embedding_handler.py | 32 +++++++++++++------ 1 file changed, 23 insertions(+), 9 deletions(-) diff --git a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py index 0e7afd5da3f..696341598e5 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py @@ -65,7 +65,9 @@ class VertexEmbedding(VertexBase): litellm_params=litellm_params, ) - should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params) + should_use_v1beta1_features = self.is_using_v1beta1_features( + optional_params=optional_params + ) _auth_header, vertex_project = self._ensure_access_token( credentials=vertex_credentials, @@ -128,10 +130,14 @@ class VertexEmbedding(VertexBase): _json_response = response.json() ## LOGGING POST-CALL - logging_obj.post_call(input=input, api_key=None, original_response=_json_response) + logging_obj.post_call( + input=input, api_key=None, original_response=_json_response + ) - model_response = litellm.vertexAITextEmbeddingConfig.transform_vertex_response_to_openai( - response=_json_response, model=model, model_response=model_response + model_response = ( + litellm.vertexAITextEmbeddingConfig.transform_vertex_response_to_openai( + response=_json_response, model=model, model_response=model_response + ) ) return model_response @@ -160,7 +166,9 @@ class VertexEmbedding(VertexBase): """ Async embedding implementation """ - should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params) + should_use_v1beta1_features = self.is_using_v1beta1_features( + optional_params=optional_params + ) _auth_header, vertex_project = await self._ensure_access_token_async( credentials=vertex_credentials, project_id=vertex_project, @@ -197,7 +205,9 @@ class VertexEmbedding(VertexBase): if timeout: _async_client_params["timeout"] = timeout if client is None or not isinstance(client, AsyncHTTPHandler): - client = get_async_httpx_client(params=_async_client_params, llm_provider=litellm.LlmProviders.VERTEX_AI) + client = get_async_httpx_client( + params=_async_client_params, llm_provider=litellm.LlmProviders.VERTEX_AI + ) else: client = client # type: ignore ## LOGGING @@ -222,10 +232,14 @@ class VertexEmbedding(VertexBase): _json_response = response.json() ## LOGGING POST-CALL - logging_obj.post_call(input=input, api_key=None, original_response=_json_response) + logging_obj.post_call( + input=input, api_key=None, original_response=_json_response + ) - model_response = litellm.vertexAITextEmbeddingConfig.transform_vertex_response_to_openai( - response=_json_response, model=model, model_response=model_response + model_response = ( + litellm.vertexAITextEmbeddingConfig.transform_vertex_response_to_openai( + response=_json_response, model=model, model_response=model_response + ) ) return model_response From 3c2c61e1e4c1d4f54ef2d796710505d56436faa8 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Tue, 28 Apr 2026 19:15:03 +0000 Subject: [PATCH 28/75] refactor(proxy): replace DualCache with UserApiKeyCache for user API key management - Updated instances of DualCache to UserApiKeyCache across multiple files to enhance cache handling for user API keys. - Adjusted cache retrieval and storage methods to ensure proper serialization and deserialization of cached objects. - Introduced a new UserApiKeyCache class to streamline caching logic and improve type safety. - Updated relevant tests to reflect changes in caching behavior and ensure compatibility with the new cache implementation. --- litellm/proxy/auth/auth_checks.py | 310 +++++++++--------- litellm/proxy/auth/handle_jwt.py | 29 +- litellm/proxy/auth/user_api_key_auth.py | 30 +- .../proxy/common_utils/user_api_key_cache.py | 156 +++++++++ .../access_group_endpoints.py | 17 +- litellm/proxy/management_endpoints/ui_sso.py | 12 +- litellm/proxy/proxy_server.py | 29 +- .../proxy/auth/test_auth_checks.py | 36 +- .../common_utils/test_user_api_key_cache.py | 113 +++++++ .../proxy/test_redis_auth_cache_flag.py | 14 +- 10 files changed, 503 insertions(+), 243 deletions(-) create mode 100644 litellm/proxy/common_utils/user_api_key_cache.py create mode 100644 tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 3cce911cd0c..8bd8eb794e0 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -68,6 +68,7 @@ from litellm.proxy.guardrails.tool_name_extraction import ( extract_request_tool_names, ) from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.route_llm_request import route_request from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics from litellm.router import Router @@ -854,7 +855,7 @@ def get_actual_routes(allowed_routes: list) -> list: async def get_default_end_user_budget( prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, ) -> Optional[LiteLLM_BudgetTable]: """ @@ -877,11 +878,12 @@ async def get_default_end_user_budget( cache_key = f"default_end_user_budget:{litellm.max_end_user_budget_id}" # Check cache first - cached_budget = await user_api_key_cache.async_get_cache(key=cache_key) + cached_budget = await user_api_key_cache.async_get_cache( + key=cache_key, + model_type=LiteLLM_BudgetTable, + ) if cached_budget is not None: - deserialized = CacheCodec.deserialize(cached_budget, LiteLLM_BudgetTable) - if deserialized is not None: - return deserialized + return cached_budget # Fetch from database try: @@ -899,7 +901,8 @@ async def get_default_end_user_budget( # Cache the budget for 60 seconds await user_api_key_cache.async_set_cache( key=cache_key, - value=CacheCodec.serialize(_budget_obj, model_type=LiteLLM_BudgetTable), + value=_budget_obj, + model_type=LiteLLM_BudgetTable, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) @@ -914,7 +917,7 @@ async def get_default_end_user_budget( async def get_team_member_default_budget( budget_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, ) -> Optional[LiteLLM_BudgetTable]: """ Fetches the team-level default per-member budget referenced by team.metadata["team_member_budget_id"]. @@ -970,7 +973,7 @@ async def get_team_member_default_budget( async def _apply_default_budget_to_end_user( end_user_obj: LiteLLM_EndUserTable, prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, ) -> LiteLLM_EndUserTable: """ @@ -1043,7 +1046,7 @@ def _check_end_user_budget( async def get_end_user_object( end_user_id: Optional[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, route: str, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, @@ -1074,22 +1077,24 @@ async def get_end_user_object( _key = "end_user_id:{}".format(end_user_id) # Check cache first - cached_user_obj = await user_api_key_cache.async_get_cache(key=_key) + cached_user_obj = await user_api_key_cache.async_get_cache( + key=_key, + model_type=LiteLLM_EndUserTable, + ) if cached_user_obj is not None: - return_obj = CacheCodec.deserialize(cached_user_obj, LiteLLM_EndUserTable) - if return_obj is not None: - # Apply default budget if needed - return_obj = await _apply_default_budget_to_end_user( - end_user_obj=return_obj, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - ) + return_obj = cached_user_obj + # Apply default budget if needed + return_obj = await _apply_default_budget_to_end_user( + end_user_obj=return_obj, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + ) - # Check budget limits - _check_end_user_budget(end_user_obj=return_obj, route=route) + # Check budget limits + _check_end_user_budget(end_user_obj=return_obj, route=route) - return return_obj + return return_obj # Fetch from database try: @@ -1115,7 +1120,8 @@ async def get_end_user_object( # Save to cache await user_api_key_cache.async_set_cache( key="end_user_id:{}".format(end_user_id), - value=CacheCodec.serialize(_response, model_type=LiteLLM_EndUserTable), + value=_response, + model_type=LiteLLM_EndUserTable, ) # Check budget limits @@ -1133,7 +1139,7 @@ async def get_end_user_object( async def get_tag_objects_batch( tag_names: List[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Dict[str, LiteLLM_TagTable]: @@ -1166,13 +1172,12 @@ async def get_tag_objects_batch( # Try to get all tags from cache first for tag_name in tag_names: cache_key = f"tag:{tag_name}" - cached_tag = await user_api_key_cache.async_get_cache(key=cache_key) + cached_tag = await user_api_key_cache.async_get_cache( + key=cache_key, + model_type=LiteLLM_TagTable, + ) if cached_tag is not None: - deserialized_tag = CacheCodec.deserialize(cached_tag, LiteLLM_TagTable) - if deserialized_tag is not None: - tag_objects[tag_name] = deserialized_tag - else: - uncached_tags.append(tag_name) + tag_objects[tag_name] = cached_tag else: uncached_tags.append(tag_name) @@ -1191,7 +1196,8 @@ async def get_tag_objects_batch( _tag_obj = LiteLLM_TagTable(**db_tag.dict()) await user_api_key_cache.async_set_cache( key=cache_key, - value=CacheCodec.serialize(_tag_obj, model_type=LiteLLM_TagTable), + value=_tag_obj, + model_type=LiteLLM_TagTable, ) tag_objects[tag_name] = _tag_obj except Exception as e: @@ -1204,7 +1210,7 @@ async def get_tag_objects_batch( async def get_tag_object( tag_name: Optional[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Optional[LiteLLM_TagTable]: @@ -1243,7 +1249,7 @@ async def get_team_membership( user_id: str, team_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Optional["LiteLLM_TeamMembership"]: @@ -1263,13 +1269,12 @@ async def get_team_membership( _key = "team_membership:{}:{}".format(user_id, team_id) # check if in cache - cached_membership_obj = await user_api_key_cache.async_get_cache(key=_key) + cached_membership_obj = await user_api_key_cache.async_get_cache( + key=_key, + model_type=LiteLLM_TeamMembership, + ) if cached_membership_obj is not None: - deserialized_membership = CacheCodec.deserialize( - cached_membership_obj, LiteLLM_TeamMembership - ) - if deserialized_membership is not None: - return deserialized_membership + return cached_membership_obj # else, check db try: @@ -1284,7 +1289,8 @@ async def get_team_membership( _response = LiteLLM_TeamMembership(**response.dict()) await user_api_key_cache.async_set_cache( key=_key, - value=CacheCodec.serialize(_response, model_type=LiteLLM_TeamMembership), + value=_response, + model_type=LiteLLM_TeamMembership, ) return _response @@ -1453,7 +1459,7 @@ async def _get_fuzzy_user_object( async def get_user_object( user_id: Optional[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, user_id_upsert: bool, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, @@ -1472,13 +1478,12 @@ async def get_user_object( # check if in cache if not check_db_only: - cached_user_obj = await user_api_key_cache.async_get_cache(key=user_id) + cached_user_obj = await user_api_key_cache.async_get_cache( + key=user_id, + model_type=LiteLLM_UserTable, + ) if cached_user_obj is not None: - deserialized_user = CacheCodec.deserialize( - cached_user_obj, LiteLLM_UserTable - ) - if deserialized_user is not None: - return deserialized_user + return cached_user_obj # else, check db if prisma_client is None: raise Exception("No db connected") @@ -1540,7 +1545,8 @@ async def get_user_object( # save the user object to cache await user_api_key_cache.async_set_cache( key=user_id, - value=CacheCodec.serialize(_response, model_type=LiteLLM_UserTable), + value=_response, + model_type=LiteLLM_UserTable, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) @@ -1562,22 +1568,20 @@ async def get_user_object( async def _cache_management_object( key: str, value: Union[BaseModel, Dict[str, Any]], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging], *, model_type: Type[BaseModel], ): """ - Persist management objects to DualCache (in-memory + optional Redis). + Persist management objects via ``UserApiKeyCache`` (in-memory + optional Redis). - Values must be JSON-serializable for the Redis path (``json.dumps``). Payloads are - encoded with ``CacheCodec.serialize(..., model_type=...)`` so writes match reads - via ``CacheCodec.deserialize(..., model_type)``. + ``UserApiKeyCache`` serializes with ``model_type`` so Redis and in-memory stay aligned. """ - cache_payload = CacheCodec.serialize(value, model_type=model_type) await user_api_key_cache.async_set_cache( key=key, - value=cache_payload, + value=value, + model_type=model_type, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) @@ -1585,7 +1589,7 @@ async def _cache_management_object( async def _cache_team_object( team_id: str, team_table: LiteLLM_TeamTableCachedObj, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging], ): key = "team_id:{}".format(team_id) @@ -1605,7 +1609,7 @@ async def _cache_team_object( async def _cache_key_object( hashed_token: str, user_api_key_obj: UserAPIKeyAuth, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging], ): key = hashed_token @@ -1624,7 +1628,7 @@ async def _cache_key_object( async def _delete_cache_key_object( hashed_token: str, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging], ): key = hashed_token @@ -1672,7 +1676,7 @@ async def _get_team_object_from_db(team_id: str, prisma_client: PrismaClient): async def _get_team_object_from_user_api_key_cache( team_id: str, prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, last_db_access_time: LimitedSizeOrderedDict, db_cache_expiry: int, proxy_logging_obj: Optional[ProxyLogging], @@ -1733,35 +1737,38 @@ async def _get_team_object_from_user_api_key_cache( async def _get_team_object_from_cache( key: str, proxy_logging_obj: Optional[ProxyLogging], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], ) -> Optional[LiteLLM_TeamTableCachedObj]: - cached_team_obj: Optional[LiteLLM_TeamTableCachedObj] = None - - ## CHECK REDIS CACHE ## + ## INTERNAL USAGE CACHE (plain DualCache) — checked before UserApiKeyCache stores ## if ( proxy_logging_obj is not None and proxy_logging_obj.internal_usage_cache.dual_cache ): - cached_team_obj = ( + cached_raw = ( await proxy_logging_obj.internal_usage_cache.dual_cache.async_get_cache( key=key, parent_otel_span=parent_otel_span ) ) + if cached_raw is not None: + from_internal = CacheCodec.deserialize( + cached_raw, LiteLLM_TeamTableCachedObj + ) + if from_internal is not None: + return from_internal - if cached_team_obj is None: - cached_team_obj = await user_api_key_cache.async_get_cache(key=key) - - if cached_team_obj is None: - return None - - return CacheCodec.deserialize(cached_team_obj, LiteLLM_TeamTableCachedObj) + decoded = await user_api_key_cache.async_get_cache( + key=key, + parent_otel_span=parent_otel_span, + model_type=LiteLLM_TeamTableCachedObj, + ) + return decoded async def get_team_object( team_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, check_cache_only: Optional[bool] = None, @@ -1827,23 +1834,21 @@ async def get_team_object( async def _cache_access_object( access_group_id: str, access_group_table: LiteLLM_AccessGroupTable, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging] = None, ): key = "access_group_id:{}".format(access_group_id) - cache_payload = CacheCodec.serialize( - access_group_table, model_type=LiteLLM_AccessGroupTable - ) await user_api_key_cache.async_set_cache( key=key, - value=cache_payload, + value=access_group_table, + model_type=LiteLLM_AccessGroupTable, ttl=DEFAULT_ACCESS_GROUP_CACHE_TTL, ) async def _delete_cache_access_object( access_group_id: str, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging] = None, ): key = "access_group_id:{}".format(access_group_id) @@ -1861,7 +1866,7 @@ async def _delete_cache_access_object( async def get_access_object( access_group_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> LiteLLM_AccessGroupTable: """ @@ -1883,10 +1888,12 @@ async def get_access_object( key = "access_group_id:{}".format(access_group_id) - cached_access_obj = await user_api_key_cache.async_get_cache(key=key) - deserialized = CacheCodec.deserialize(cached_access_obj, LiteLLM_AccessGroupTable) - if deserialized is not None: - return deserialized + cached_access_obj = await user_api_key_cache.async_get_cache( + key=key, + model_type=LiteLLM_AccessGroupTable, + ) + if cached_access_obj is not None: + return cached_access_obj # Not in cache - fetch from DB try: @@ -1932,7 +1939,7 @@ async def get_access_object( async def get_team_object_by_alias( team_alias: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional["Span"] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> LiteLLM_TeamTableCachedObj: @@ -2011,19 +2018,18 @@ async def get_team_object_by_alias( ) # Cache the result by both alias and team_id - _serialized_team = CacheCodec.serialize( - team_obj, model_type=LiteLLM_TeamTableCachedObj - ) await user_api_key_cache.async_set_cache( key=cache_key, - value=_serialized_team, + value=team_obj, + model_type=LiteLLM_TeamTableCachedObj, ttl=DEFAULT_IN_MEMORY_TTL, ) # Also cache by team_id for consistency team_id_cache_key = "team_id:{}".format(team_obj.team_id) await user_api_key_cache.async_set_cache( key=team_id_cache_key, - value=_serialized_team, + value=team_obj, + model_type=LiteLLM_TeamTableCachedObj, ttl=DEFAULT_IN_MEMORY_TTL, ) @@ -2045,7 +2051,7 @@ async def get_team_object_by_alias( async def get_org_object_by_alias( org_alias: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional["Span"] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Optional[LiteLLM_OrganizationTable]: @@ -2072,13 +2078,12 @@ async def get_org_object_by_alias( # Check cache first (keyed by alias) cache_key = "org_alias:{}".format(org_alias) - cached_org_obj = await user_api_key_cache.async_get_cache(key=cache_key) + cached_org_obj = await user_api_key_cache.async_get_cache( + key=cache_key, + model_type=LiteLLM_OrganizationTable, + ) if cached_org_obj is not None: - deserialized_org = CacheCodec.deserialize( - cached_org_obj, LiteLLM_OrganizationTable - ) - if deserialized_org is not None: - return deserialized_org + return cached_org_obj # Query database by organization_alias try: @@ -2105,19 +2110,18 @@ async def get_org_object_by_alias( org = orgs[0] org_obj = LiteLLM_OrganizationTable(**org.model_dump()) - _serialized_org = CacheCodec.serialize( - org_obj, model_type=LiteLLM_OrganizationTable - ) # Cache the result await user_api_key_cache.async_set_cache( key=cache_key, - value=_serialized_org, + value=org_obj, + model_type=LiteLLM_OrganizationTable, ttl=DEFAULT_IN_MEMORY_TTL, ) # Also cache by org_id for consistency await user_api_key_cache.async_set_cache( key="org_id:{}".format(org_obj.organization_id), - value=_serialized_org, + value=org_obj, + model_type=LiteLLM_OrganizationTable, ttl=DEFAULT_IN_MEMORY_TTL, ) @@ -2321,7 +2325,7 @@ async def get_jwt_key_mapping_object( async def get_key_object( hashed_token: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, check_cache_only: Optional[bool] = None, @@ -2341,11 +2345,12 @@ async def get_key_object( # Same flow as before: use cache only when we have a hit we can turn into UserAPIKeyAuth # (dict from Redis / model_dump, or UserAPIKeyAuth from in-memory). Otherwise fall through to DB. - cached_key_obj = await user_api_key_cache.async_get_cache(key=key) - if cached_key_obj is not None: - user_api_key_auth = CacheCodec.deserialize(cached_key_obj, UserAPIKeyAuth) - if user_api_key_auth is not None: - return user_api_key_auth + user_api_key_auth = await user_api_key_cache.async_get_cache( + key=key, + model_type=UserAPIKeyAuth, + ) + if user_api_key_auth is not None: + return user_api_key_auth if check_cache_only: raise Exception( @@ -2402,7 +2407,7 @@ async def get_key_object( async def get_object_permission( object_permission_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Optional[LiteLLM_ObjectPermissionTable]: @@ -2418,13 +2423,12 @@ async def get_object_permission( # check if in cache key = "object_permission_id:{}".format(object_permission_id) - cached_obj_permission = await user_api_key_cache.async_get_cache(key=key) - if cached_obj_permission is not None: - deserialized_perm = CacheCodec.deserialize( - cached_obj_permission, LiteLLM_ObjectPermissionTable - ) - if deserialized_perm is not None: - return deserialized_perm + deserialized_perm = await user_api_key_cache.async_get_cache( + key=key, + model_type=LiteLLM_ObjectPermissionTable, + ) + if deserialized_perm is not None: + return deserialized_perm # else, check db try: @@ -2438,9 +2442,8 @@ async def get_object_permission( _perm_obj = LiteLLM_ObjectPermissionTable(**response.dict()) await user_api_key_cache.async_set_cache( key=key, - value=CacheCodec.serialize( - _perm_obj, model_type=LiteLLM_ObjectPermissionTable - ), + value=_perm_obj, + model_type=LiteLLM_ObjectPermissionTable, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) @@ -2453,7 +2456,7 @@ async def get_object_permission( async def get_managed_vector_store_rows_by_uuids( uuids: List[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> List[LiteLLM_ManagedVectorStoresTable]: @@ -2473,15 +2476,12 @@ async def get_managed_vector_store_rows_by_uuids( for uuid in uuids: key = "managed_vector_store_id:{}".format(uuid) - cached = await user_api_key_cache.async_get_cache(key=key) - if cached is not None: - deserialized_vs = CacheCodec.deserialize( - cached, LiteLLM_ManagedVectorStoresTable - ) - if deserialized_vs is not None: - result.append(deserialized_vs) - else: - cache_misses.append(uuid) + deserialized_vs = await user_api_key_cache.async_get_cache( + key=key, + model_type=LiteLLM_ManagedVectorStoresTable, + ) + if deserialized_vs is not None: + result.append(deserialized_vs) else: cache_misses.append(uuid) @@ -2507,9 +2507,8 @@ async def get_managed_vector_store_rows_by_uuids( key = "managed_vector_store_id:{}".format(cached_obj.vector_store_id) await user_api_key_cache.async_set_cache( key=key, - value=CacheCodec.serialize( - cached_obj, model_type=LiteLLM_ManagedVectorStoresTable - ), + value=cached_obj, + model_type=LiteLLM_ManagedVectorStoresTable, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) result.append(cached_obj) @@ -2521,7 +2520,7 @@ async def get_managed_vector_store_rows_by_uuids( async def get_org_object( org_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, include_budget_table: bool = False, @@ -2552,13 +2551,12 @@ async def get_org_object( cache_key = "org_id:{}:with_budget".format(org_id) # check if in cache - cached_org_obj = await user_api_key_cache.async_get_cache(key=cache_key) - if cached_org_obj is not None: - deserialized_org = CacheCodec.deserialize( - cached_org_obj, LiteLLM_OrganizationTable - ) - if deserialized_org is not None: - return deserialized_org + deserialized_org = await user_api_key_cache.async_get_cache( + key=cache_key, + model_type=LiteLLM_OrganizationTable, + ) + if deserialized_org is not None: + return deserialized_org # else, check db try: query_kwargs: Dict[str, Any] = {"where": {"organization_id": org_id}} @@ -2576,7 +2574,8 @@ async def get_org_object( # Cache the result await user_api_key_cache.async_set_cache( key=cache_key, - value=CacheCodec.serialize(_org_obj, model_type=LiteLLM_OrganizationTable), + value=_org_obj, + model_type=LiteLLM_OrganizationTable, ttl=DEFAULT_IN_MEMORY_TTL, ) @@ -2593,7 +2592,7 @@ async def _get_resources_from_access_groups( "access_model_names", "access_mcp_server_ids", "access_agent_ids" ], prisma_client: Optional[PrismaClient] = None, - user_api_key_cache: Optional[DualCache] = None, + user_api_key_cache: Optional[UserApiKeyCache] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> List[str]: """ @@ -2651,7 +2650,7 @@ async def _get_resources_from_access_groups( async def _get_models_from_access_groups( access_group_ids: List[str], prisma_client: Optional[PrismaClient] = None, - user_api_key_cache: Optional[DualCache] = None, + user_api_key_cache: Optional[UserApiKeyCache] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> List[str]: """ @@ -2670,7 +2669,7 @@ async def _get_models_from_access_groups( async def _get_mcp_server_ids_from_access_groups( access_group_ids: List[str], prisma_client: Optional[PrismaClient] = None, - user_api_key_cache: Optional[DualCache] = None, + user_api_key_cache: Optional[UserApiKeyCache] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> List[str]: """ @@ -2689,7 +2688,7 @@ async def _get_mcp_server_ids_from_access_groups( async def _get_agent_ids_from_access_groups( access_group_ids: List[str], prisma_client: Optional[PrismaClient] = None, - user_api_key_cache: Optional[DualCache] = None, + user_api_key_cache: Optional[UserApiKeyCache] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> List[str]: """ @@ -3303,7 +3302,7 @@ async def _check_team_member_budget( user_object: Optional[LiteLLM_UserTable], valid_token: Optional[UserAPIKeyAuth], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ): """Check if team member is over their max budget within the team.""" @@ -3370,7 +3369,7 @@ async def _check_team_member_model_access( valid_token: UserAPIKeyAuth, llm_router: Optional[Router], prisma_client: Optional["PrismaClient"], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ) -> None: """ @@ -3677,7 +3676,7 @@ async def _project_soft_budget_check( async def get_project_object( project_id: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Optional[LiteLLM_ProjectTableCachedObj]: """ @@ -3692,13 +3691,12 @@ async def get_project_object( # Check cache first cache_key = "project_id:{}".format(project_id) - cached_obj = await user_api_key_cache.async_get_cache(key=cache_key) - if cached_obj is not None: - deserialized_project = CacheCodec.deserialize( - cached_obj, LiteLLM_ProjectTableCachedObj - ) - if deserialized_project is not None: - return deserialized_project + deserialized_project = await user_api_key_cache.async_get_cache( + key=cache_key, + model_type=LiteLLM_ProjectTableCachedObj, + ) + if deserialized_project is not None: + return deserialized_project # Fetch from DB project_row = await prisma_client.db.litellm_projecttable.find_unique( @@ -3727,7 +3725,7 @@ async def _organization_max_budget_check( valid_token: Optional[UserAPIKeyAuth], team_object: Optional[LiteLLM_TeamTable], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ): """ @@ -3821,7 +3819,7 @@ async def _organization_max_budget_check( async def _tag_max_budget_check( request_body: dict, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, valid_token: Optional[UserAPIKeyAuth], ): diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index fb86485ffad..9e7c0abd983 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -20,7 +20,6 @@ import jwt from jwt.api_jwk import PyJWK from litellm._logging import verbose_proxy_logger -from litellm.caching.caching import DualCache from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value from litellm.llms.custom_httpx.httpx_handler import HTTPHandler @@ -46,7 +45,7 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.auth_checks import can_team_access_model from litellm.proxy.auth.route_checks import RouteChecks -from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.utils import PrismaClient, ProxyLogging from .auth_checks import ( @@ -74,7 +73,7 @@ class JWTHandler: """ prisma_client: Optional[PrismaClient] - user_api_key_cache: DualCache + user_api_key_cache: UserApiKeyCache # Supported algos: https://pyjwt.readthedocs.io/en/stable/algorithms.html # "Warning: Make sure not to mix symmetric and asymmetric algorithms that interpret # the key in different ways (e.g. HS* and RS*)." @@ -100,7 +99,7 @@ class JWTHandler: def update_environment( self, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, litellm_jwtauth: LiteLLM_JWTAuth, leeway: int = 0, ) -> None: @@ -953,7 +952,7 @@ class JWTAuthManager: jwt_handler: JWTHandler, jwt_valid_token: dict, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, ) -> Tuple[Optional[str], Optional[LiteLLM_TeamTable]]: @@ -1046,7 +1045,7 @@ class JWTAuthManager: route: str, jwt_handler: JWTHandler, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, ) -> Tuple[Optional[str], Optional[LiteLLM_TeamTable]]: @@ -1134,7 +1133,7 @@ class JWTAuthManager: valid_user_email: Optional[bool], jwt_handler: JWTHandler, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, route: str, @@ -1350,7 +1349,7 @@ class JWTAuthManager: jwt_valid_token: dict, user_object: Optional[LiteLLM_UserTable], prisma_client: Optional[PrismaClient], - user_api_key_cache: Optional[DualCache] = None, + user_api_key_cache: Optional[UserApiKeyCache] = None, ) -> None: """ Sync user role and team memberships with JWT claims @@ -1378,9 +1377,8 @@ class JWTAuthManager: if user_api_key_cache is not None: await user_api_key_cache.async_set_cache( key=user_object.user_id, - value=CacheCodec.serialize( - user_object, model_type=LiteLLM_UserTable - ), + value=user_object, + model_type=LiteLLM_UserTable, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) @@ -1403,9 +1401,8 @@ class JWTAuthManager: if user_api_key_cache is not None: await user_api_key_cache.async_set_cache( key=user_object.user_id, - value=CacheCodec.serialize( - user_object, model_type=LiteLLM_UserTable - ), + value=user_object, + model_type=LiteLLM_UserTable, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, ) return None @@ -1417,7 +1414,7 @@ class JWTAuthManager: request_headers: Optional[dict], jwt_handler: JWTHandler, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, ) -> None: @@ -1464,7 +1461,7 @@ class JWTAuthManager: general_settings: dict, route: str, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, request_headers: Optional[dict] = None, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index f77fdc47398..30f78a50813 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -20,7 +20,6 @@ from fastapi.security.api_key import APIKeyHeader import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging -from litellm.caching import DualCache from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value from litellm.proxy._types import * @@ -58,7 +57,7 @@ from litellm.proxy.auth.oauth2_check import Oauth2Handler from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator -from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, _safe_get_request_headers, @@ -328,7 +327,7 @@ _global_spend_coordinator = EventDrivenCacheCoordinator(log_prefix="[GLOBAL SPEN async def _fetch_global_spend_with_event_coordination( cache_key: str, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, prisma_client: PrismaClient, ) -> Optional[float]: """ @@ -351,7 +350,7 @@ async def _fetch_global_spend_with_event_coordination( async def get_global_proxy_spend( litellm_proxy_admin_name: str, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, prisma_client: Optional[PrismaClient], token: str, proxy_logging_obj: ProxyLogging, @@ -504,7 +503,7 @@ async def _resolve_jwt_to_virtual_key( jwt_claims: dict, jwt_handler: JWTHandler, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, ) -> Optional[UserAPIKeyAuth]: @@ -1315,13 +1314,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 if prisma_client is not None: _cache_key = f"{valid_token.team_id}_{valid_token.user_id}" - _cached_member = await user_api_key_cache.async_get_cache( - key=_cache_key - ) - team_member_info: Optional[LiteLLM_TeamMembership] = ( - CacheCodec.deserialize(_cached_member, LiteLLM_TeamMembership) - if _cached_member is not None - else None + team_member_info = await user_api_key_cache.async_get_cache( + key=_cache_key, + model_type=LiteLLM_TeamMembership, ) if team_member_info is None: # read from DB @@ -1342,10 +1337,8 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 ) await user_api_key_cache.async_set_cache( key=_cache_key, - value=CacheCodec.serialize( - team_member_info, - model_type=LiteLLM_TeamMembership, - ), + value=team_member_info, + model_type=LiteLLM_TeamMembership, ttl=5, ) @@ -1497,9 +1490,8 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 team_cache_key = f"team_id:{valid_token.team_id}" await user_api_key_cache.async_set_cache( key=team_cache_key, - value=CacheCodec.serialize( - _team_obj, model_type=LiteLLM_TeamTableCachedObj - ), + value=_team_obj, + model_type=LiteLLM_TeamTableCachedObj, ) # save team table in cache - used for tpm/rpm limiting - tpm_rpm_limiter.py # Fetch project object if key belongs to a project diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py new file mode 100644 index 00000000000..2eebc291bb0 --- /dev/null +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -0,0 +1,156 @@ +from __future__ import annotations + +from typing import Any, Optional, Type, TypeVar, Union, cast, overload + +from pydantic import BaseModel + +from litellm._logging import verbose_proxy_logger +from litellm.caching.dual_cache import DualCache +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec + +T = TypeVar("T", bound=BaseModel) + + +class UserApiKeyCache(DualCache): + """ + DualCache wrapper for UserAPIKeyAuth-like payloads. + + Stores a Redis-safe JSON payload in BOTH in-memory and Redis to avoid + "memory returns BaseModel, Redis returns dict" format drift. + + When ``model_type`` is provided: + - writes are serialized via ``CacheCodec.serialize(..., model_type=...)`` + - reads are deserialized via ``CacheCodec.deserialize(..., model_type)`` + and return ``Optional[T]``: the model on success, ``None`` on cache miss + **or** if the cached payload fails validation (schema drift). On + validation failure after a cache hit, an error line is emitted via + ``verbose_proxy_logger``. + + When ``model_type`` is omitted, the interface behaves like ``DualCache``: + raw cached payload is returned (dict/str/etc.). + + ``async_set_cache_pipeline`` applies the same untyped Codec pass as omitting + ``model_type`` on ``async_set_cache`` (so ``BaseModel`` rows are dumped before Redis). + """ + + # Overloads: `model_type` must be a real parameter (not only via **kwargs) so + # the untyped branch cannot match calls that pass `model_type=...`. + + @overload + def get_cache( + self, + key, + parent_otel_span=None, + local_only: bool = False, + ) -> Any: ... + + @overload + def get_cache( + self, + key, + parent_otel_span=None, + local_only: bool = False, + *, + model_type: Type[T], + ) -> Optional[T]: ... + + def set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] + model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + payload = CacheCodec.serialize(value, model_type=model_type) + return super().set_cache(key=key, value=payload, local_only=local_only, **kwargs) + + async def async_set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] + model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + payload = CacheCodec.serialize(value, model_type=model_type) + return await super().async_set_cache( + key=key, value=payload, local_only=local_only, **kwargs + ) + + async def async_set_cache_pipeline( # type: ignore[override] + self, cache_list: list, local_only: bool = False, **kwargs + ) -> None: + """ + Batch writes with the same Codec boundary as ``async_set_cache`` without + ``model_type``: ``BaseModel`` values become JSON-safe dicts; dicts/scalars unchanged. + """ + normalized = [ + (key, CacheCodec.serialize(value, model_type=None)) + for key, value in cache_list + ] + return await super().async_set_cache_pipeline( + cache_list=normalized, local_only=local_only, **kwargs + ) + + @overload + async def async_get_cache( + self, + key, + parent_otel_span=None, + local_only: bool = False, + ) -> Any: ... + + @overload + async def async_get_cache( + self, + key, + parent_otel_span=None, + local_only: bool = False, + *, + model_type: Type[T], + ) -> Optional[T]: ... + + def get_cache( # type: ignore[override] + self, + key, + parent_otel_span=None, + local_only: bool = False, + model_type: Optional[Type[BaseModel]] = None, + **kwargs, + ) -> Union[Any, Optional[BaseModel]]: + if model_type is None and "model_type" in kwargs: + model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + cached = super().get_cache( + key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs + ) + if model_type is None: + return cached + if cached is None: + return None + decoded = CacheCodec.deserialize(cached, model_type=model_type) + if decoded is None: + verbose_proxy_logger.error( + "UserApiKeyCache.get_cache failed to deserialize cached value for " + "key=%r model_type=%s", + key, + getattr(model_type, "__name__", str(model_type)), + ) + return None + return decoded + + async def async_get_cache( # type: ignore[override] + self, + key, + parent_otel_span=None, + local_only: bool = False, + model_type: Optional[Type[BaseModel]] = None, + **kwargs, + ) -> Union[Any, Optional[BaseModel]]: + if model_type is None and "model_type" in kwargs: + model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + cached = await super().async_get_cache( + key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs + ) + if model_type is None: + return cached + if cached is None: + return None + decoded = CacheCodec.deserialize(cached, model_type=model_type) + if decoded is None: + verbose_proxy_logger.error( + "UserApiKeyCache.async_get_cache failed to deserialize cached value for " + "key=%r model_type=%s", + key, + getattr(model_type, "__name__", str(model_type)), + ) + return None + return decoded diff --git a/litellm/proxy/management_endpoints/access_group_endpoints.py b/litellm/proxy/management_endpoints/access_group_endpoints.py index 30fed096232..ceaef20a8d0 100644 --- a/litellm/proxy/management_endpoints/access_group_endpoints.py +++ b/litellm/proxy/management_endpoints/access_group_endpoints.py @@ -9,7 +9,6 @@ from litellm.proxy._types import ( LitellmUserRoles, UserAPIKeyAuth, ) -from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec from litellm.proxy.auth.auth_checks import ( _cache_access_object, _cache_key_object, @@ -237,10 +236,10 @@ async def _patch_key_caches_add_access_group( ) -> None: """Patch cached key objects to include access_group_id.""" for token in key_tokens: - raw = await user_api_key_cache.async_get_cache(key=token) - if raw is None: - continue - cached_key = CacheCodec.deserialize(raw, UserAPIKeyAuth) + cached_key = await user_api_key_cache.async_get_cache( + key=token, + model_type=UserAPIKeyAuth, + ) if cached_key is None: continue if cached_key.access_group_ids is None: @@ -267,10 +266,10 @@ async def _patch_key_caches_remove_access_group( ) -> None: """Patch cached key objects to remove access_group_id.""" for token in key_tokens: - raw = await user_api_key_cache.async_get_cache(key=token) - if raw is None: - continue - cached_key = CacheCodec.deserialize(raw, UserAPIKeyAuth) + cached_key = await user_api_key_cache.async_get_cache( + key=token, + model_type=UserAPIKeyAuth, + ) if cached_key is not None and cached_key.access_group_ids: cached_key.access_group_ids = [ ag for ag in cached_key.access_group_ids if ag != access_group_id diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 0362c3ff6b2..c00f7269bbd 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -39,7 +39,6 @@ from fastapi.responses import RedirectResponse import litellm from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid -from litellm.caching import DualCache from litellm.constants import ( LITELLM_UI_SESSION_DURATION, MAX_SPENDLOG_ROWS_TO_QUERY, @@ -70,7 +69,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_user_object -from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.auth.auth_utils import _has_user_setup_sso from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -1051,7 +1050,7 @@ async def get_existing_user_info_from_db( user_id: Optional[str], user_email: Optional[str], prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ) -> Optional[LiteLLM_UserTable]: try: @@ -1075,7 +1074,7 @@ async def get_existing_user_info_from_db( async def get_user_info_from_db( result: Union[CustomOpenID, OpenID, dict], prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, user_email: Optional[str], user_defined_values: Optional[SSOUserDefinedValues], @@ -1195,7 +1194,7 @@ async def _sync_user_role_from_jwt_role_map( received_response: Optional[dict], user_info: Optional[Union[LiteLLM_UserTable, NewUserResponse]], prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, user_defined_values: Optional[SSOUserDefinedValues], ) -> None: """ @@ -1234,7 +1233,8 @@ async def _sync_user_role_from_jwt_role_map( user_info.user_role = mapped_role.value await user_api_key_cache.async_set_cache( key=user_info.user_id, - value=CacheCodec.serialize(user_info, model_type=LiteLLM_UserTable), + value=user_info, + model_type=LiteLLM_UserTable, ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 8d199d2c092..865a46a5478 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -209,6 +209,7 @@ from litellm import Router from litellm._logging import verbose_proxy_logger, verbose_router_logger from litellm.caching.caching import DualCache, RedisCache from litellm.caching.redis_cluster_cache import RedisClusterCache +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.constants import ( _REALTIME_BODY_CACHE_SIZE, APSCHEDULER_COALESCE, @@ -1598,7 +1599,7 @@ prisma_client: Optional[PrismaClient] = None shared_aiohttp_session: Optional["ClientSession"] = ( None # Global shared session for connection reuse ) -user_api_key_cache = DualCache( +user_api_key_cache = UserApiKeyCache( default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value ) spend_counter_cache = DualCache( @@ -2041,19 +2042,13 @@ async def update_cache( # noqa: PLR0915 else: hashed_token = token verbose_proxy_logger.debug("_update_key_cache: hashed_token=%s", hashed_token) - cached_key = await user_api_key_cache.async_get_cache(key=hashed_token) - verbose_proxy_logger.debug( - f"_update_key_cache: existing_spend_obj={cached_key}" + existing_spend_obj = await user_api_key_cache.async_get_cache( + key=hashed_token, model_type=UserAPIKeyAuth + ) + verbose_proxy_logger.debug( + f"_update_key_cache: existing_spend_obj={existing_spend_obj}" ) - if cached_key is None: - return - existing_spend_obj = CacheCodec.deserialize(cached_key, UserAPIKeyAuth) if existing_spend_obj is None: - verbose_proxy_logger.warning( - "_update_key_cache: unexpected cached key type %s for hashed_token=%s; skipping spend update", - type(cached_key).__name__, - hashed_token, - ) return existing_spend = existing_spend_obj.spend or 0.0 @@ -2114,14 +2109,10 @@ async def update_cache( # noqa: PLR0915 existing_team_member_spend + response_cost ) - # Update the cost column for the given token (dict for Redis pipeline json.dumps) + # Existing spend_obj is mutated; UserApiKeyCache.async_set_cache_pipeline turns + # BaseModel values into dicts for Redis (same Codec path as async_set_cache). existing_spend_obj.spend = new_spend - values_to_update_in_cache.append( - ( - hashed_token, - CacheCodec.serialize(existing_spend_obj, model_type=UserAPIKeyAuth), - ) - ) + values_to_update_in_cache.append((hashed_token, existing_spend_obj)) ### UPDATE USER SPEND ### async def _update_user_cache(): diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 8612d243c41..37646373a03 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -922,19 +922,21 @@ async def test_get_tag_objects_batch(): # Simulate 5 tags: 2 cached, 3 uncached tag_names = ["cached-1", "uncached-1", "cached-2", "uncached-2", "uncached-3"] - # Mock cached tags - cached_tag_1 = { - "tag_name": "cached-1", - "spend": 10.0, - "models": [], - "litellm_budget_table": None, - } - cached_tag_2 = { - "tag_name": "cached-2", - "spend": 20.0, - "models": [], - "litellm_budget_table": None, - } + # Mock cached tags — must be LiteLLM_TagTable instances: the mocked async_get_cache + # bypasses UserApiKeyCache deserialization, so returning plain dicts would flow through + # as dict (production returns models after Codec.deserialize inside the cache). + cached_tag_1 = LiteLLM_TagTable( + tag_name="cached-1", + spend=10.0, + models=[], + litellm_budget_table=None, + ) + cached_tag_2 = LiteLLM_TagTable( + tag_name="cached-2", + spend=20.0, + models=[], + litellm_budget_table=None, + ) # Mock DB response for uncached tags uncached_tag_1 = MagicMock() @@ -980,13 +982,13 @@ async def test_get_tag_objects_batch(): ) # Mock cache behavior - return cached tags, None for uncached - async def mock_get_cache(key): + async def mock_get_cache(*args, **kwargs): + key = kwargs.get("key") if key == "tag:cached-1": return cached_tag_1 - elif key == "tag:cached-2": + if key == "tag:cached-2": return cached_tag_2 - else: - return None + return None mock_cache.async_get_cache = AsyncMock(side_effect=mock_get_cache) mock_cache.async_set_cache = AsyncMock() diff --git a/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py b/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py new file mode 100644 index 00000000000..cff331d9b19 --- /dev/null +++ b/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py @@ -0,0 +1,113 @@ +import json +from typing import Any + +import pytest + +from litellm.caching.redis_cache import RedisCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + +class FakeRedisCache(RedisCache): + """ + In-memory fake that enforces the UserApiKeyCache Redis payload contract. + + For user_api_key_cache entries we expect Redis to store a JSON object (dict) + produced by `CacheCodec.serialize(..., model_type=...)`. + + This fake: + - raises TypeError if the value is not a dict + - raises TypeError if the dict is not JSON-serializable + """ + + def __init__(self): # noqa: super().__init__ skipped intentionally + self._store: dict[str, str] = {} + + def set_cache(self, key: str, value: Any, **kwargs): # type: ignore[override] + if not isinstance(value, dict): + raise TypeError("FakeRedisCache only accepts dict payloads") + self._store[key] = json.dumps(value) + return True + + def get_cache(self, key: str, **kwargs): # type: ignore[override] + raw = self._store.get(key) + if raw is None: + return None + return json.loads(raw) + + async def async_set_cache(self, key: str, value: Any, **kwargs): # type: ignore[override] + if not isinstance(value, dict): + raise TypeError("FakeRedisCache only accepts dict payloads") + self._store[key] = json.dumps(value) + return True + + async def async_get_cache(self, key: str, **kwargs): # type: ignore[override] + raw = self._store.get(key) + if raw is None: + return None + return json.loads(raw) + + def delete_cache(self, key: str): # type: ignore[override] + self._store.pop(key, None) + + async def async_delete_cache(self, key: str): # type: ignore[override] + self._store.pop(key, None) + + +def _make_key_obj(token: str = "tok") -> UserAPIKeyAuth: + # Minimal object (UserAPIKeyAuth inherits token from base view). + return UserAPIKeyAuth(token=token) + + +class TestUserApiKeyCache: + @pytest.mark.asyncio + async def test_async_set_typed_stores_serialized_payload_in_memory_and_redis(self): + cache = UserApiKeyCache(redis_cache=FakeRedisCache()) + obj = _make_key_obj("abc") + + await cache.async_set_cache("k", obj, model_type=UserAPIKeyAuth) + + # In-memory hit should still be raw dict (not BaseModel) because wrapper + # stores the serialized payload into both layers. + raw = await cache.in_memory_cache.async_get_cache("k") # type: ignore[union-attr] + assert isinstance(raw, dict) + assert raw["token"] == "abc" + + # Redis should also hold the same serialized dict + redis_raw = await cache.redis_cache.async_get_cache("k") # type: ignore[union-attr] + assert redis_raw == raw + + @pytest.mark.asyncio + async def test_async_get_typed_returns_model_on_valid_hit(self): + cache = UserApiKeyCache(redis_cache=FakeRedisCache()) + await cache.async_set_cache("k", {"token": "abc"}, model_type=UserAPIKeyAuth) + + value = await cache.async_get_cache("k", model_type=UserAPIKeyAuth) + assert value is not None + assert isinstance(value, UserAPIKeyAuth) + assert value.token == "abc" + + @pytest.mark.asyncio + async def test_async_get_typed_returns_none_on_validation_failure_after_hit(self): + cache = UserApiKeyCache(redis_cache=FakeRedisCache()) + + # Bypass UserApiKeyCache.serialize: CacheCodec rejects non-dict cached values + # for dict-based models (deserialize returns None). + await cache.in_memory_cache.async_set_cache( + key="k", value="invalid-payload-not-a-dict" + ) + + value = await cache.async_get_cache("k", model_type=UserAPIKeyAuth) + assert value is None + + def test_fake_redis_cache_rejects_non_json_serializable_values(self): + fake = FakeRedisCache() + + class NotSerializable: + pass + + with pytest.raises(TypeError): + fake.set_cache("k", NotSerializable()) + + with pytest.raises(TypeError): + fake.set_cache("k2", {"ok": NotSerializable()}) diff --git a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py index 261222be9a2..d0cb5ec5465 100644 --- a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py +++ b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py @@ -6,6 +6,7 @@ the flag is explicitly set to True, and leaves it in-memory-only otherwise. """ from contextlib import contextmanager +import json from unittest.mock import MagicMock, patch import pytest @@ -29,7 +30,18 @@ class _FakeRedisCache(RedisCache): """ def __init__(self): # noqa: super().__init__ skipped intentionally - pass + self._store = {} + + def set_cache(self, key, value, **kwargs): # type: ignore[override] + # Enforce Redis JSON-serializable payload contract. + self._store[key] = json.dumps(value) + return True + + def get_cache(self, key, **kwargs): # type: ignore[override] + raw = self._store.get(key) + if raw is None: + return None + return json.loads(raw) @contextmanager From 354ad40d15f504dd84ad7c116f561c274287d215 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Tue, 28 Apr 2026 19:30:56 +0000 Subject: [PATCH 29/75] fix DualCache import error in CI and update docstrings for cache coordinator --- litellm/proxy/auth/handle_jwt.py | 2 ++ .../proxy/common_utils/cache_coordinator.py | 22 ++++++++++++++++--- 2 files changed, 21 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 9e7c0abd983..8cb871c9221 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -6,6 +6,8 @@ Currently only supports admin. JWT token must have 'litellm_proxy_admin' in scope. """ +from __future__ import annotations + import fnmatch import hashlib import os diff --git a/litellm/proxy/common_utils/cache_coordinator.py b/litellm/proxy/common_utils/cache_coordinator.py index 24da9450ab8..92d9916098d 100644 --- a/litellm/proxy/common_utils/cache_coordinator.py +++ b/litellm/proxy/common_utils/cache_coordinator.py @@ -20,11 +20,27 @@ T = TypeVar("T") class AsyncCacheProtocol(Protocol): - """Protocol for cache backends used by EventDrivenCacheCoordinator.""" + """Protocol for cache backends used by EventDrivenCacheCoordinator. - async def async_get_cache(self, key: str, **kwargs: Any) -> Any: ... + Matches ``DualCache`` / ``UserApiKeyCache`` call shapes (explicit optional params + before ``**kwargs``), not only ``(key, **kwargs)``, so overloads validate. + """ - async def async_set_cache(self, key: str, value: Any, **kwargs: Any) -> Any: ... + async def async_get_cache( + self, + key: str, + parent_otel_span: Any = None, + local_only: bool = False, + **kwargs: Any, + ) -> Any: ... + + async def async_set_cache( + self, + key: str, + value: Any, + local_only: bool = False, + **kwargs: Any, + ) -> Any: ... class EventDrivenCacheCoordinator: From 84b6bd60af93158dd7b5219315920222869eb518 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Tue, 28 Apr 2026 21:08:46 +0000 Subject: [PATCH 30/75] update test cases to match new behaviour. The earlier test cases assumed the cache stores a pydantic object --- tests/proxy_unit_tests/test_auth_checks.py | 11 +++- .../test_user_api_key_auth.py | 7 ++- .../test_access_group_endpoints.py | 57 ++++++++++++++----- 3 files changed, 58 insertions(+), 17 deletions(-) diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index 86cd5c0c413..5636a55c95a 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -16,6 +16,7 @@ import httpx from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_checks import get_end_user_object from litellm.caching.caching import DualCache +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy._types import ( LiteLLM_EndUserTable, LiteLLM_BudgetTable, @@ -48,9 +49,15 @@ async def test_get_end_user_object(customer_spend, customer_budget): litellm_budget_table=_budget, blocked=False, ) - _cache = DualCache() + # UserApiKeyCache applies model_type on get/set; plain DualCache returns raw dicts + # and breaks get_end_user_object's typed async_get_cache path. + _cache = UserApiKeyCache() _key = "end_user_id:{}".format(end_user_id) - _cache.set_cache(key=_key, value=end_user_obj.model_dump()) + await _cache.async_set_cache( + key=_key, + value=end_user_obj, + model_type=LiteLLM_EndUserTable, + ) try: await get_end_user_object( end_user_id=end_user_id, diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index f847b8aa656..4e0bdf074f2 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -268,7 +268,12 @@ async def test_aaauser_personal_budgets(key_ownership): test_user_cache = getattr(litellm.proxy.proxy_server, "user_api_key_cache") - assert test_user_cache.get_cache(key=hash_token(user_key)) == valid_token + assert ( + test_user_cache.get_cache( + key=hash_token(user_key), model_type=UserAPIKeyAuth + ) + == valid_token + ) try: await user_api_key_auth(request=request, api_key="Bearer " + user_key) diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py index cd2eb789589..016e10859b6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py @@ -738,15 +738,32 @@ def test_delete_access_group_patches_cached_team_and_key( return_value=None ) - # Build cached key object (returned from user_api_key_cache) - if key_cache_group_ids is not None: - cached_key = UserAPIKeyAuth( - token="hashed-key-1", - access_group_ids=list(key_cache_group_ids), + # user_api_key_cache is queried both for teams (fallback after dual_cache) and + # hashed keys — return the right stub per ``key``. A single AsyncMock(return_value=key) + # would wrongly serve the key blob for ``team_id:team-1`` and trigger team patching. + # Use a synchronous side_effect (not async def): AsyncMock awaits coroutine side_effects + # inconsistently across Python/unittest versions; sync returns are awaited as immediate results. + def user_cache_get_side_effect(*args, **kwargs): + cache_key = ( + kwargs.get("key") if "key" in kwargs else (args[0] if args else None) ) - mock_cache.async_get_cache = AsyncMock(return_value=cached_key) - else: - mock_cache.async_get_cache = AsyncMock(return_value=None) + if cache_key == "team_id:team-1": + if team_cache_group_ids is None: + return None + return LiteLLM_TeamTableCachedObj( + team_id="team-1", + access_group_ids=list(team_cache_group_ids), + ) + if cache_key == "hashed-key-1": + if key_cache_group_ids is None: + return None + return UserAPIKeyAuth( + token="hashed-key-1", + access_group_ids=list(key_cache_group_ids), + ) + return None + + mock_cache.async_get_cache = AsyncMock(side_effect=user_cache_get_side_effect) resp = client.delete("/v1/access_group/ag-to-delete") assert resp.status_code == 204 @@ -803,7 +820,7 @@ def test_delete_access_group_patches_cached_team_and_key( def test_delete_access_group_patches_key_cached_as_dict(client_and_mocks): - """Delete correctly patches a key cached as a raw dict (not UserAPIKeyAuth).""" + """Delete patches key cache — mock returns UserAPIKeyAuth (what UserApiKeyCache emits after deserialize).""" client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging = ( client_and_mocks ) @@ -826,12 +843,24 @@ def test_delete_access_group_patches_key_cached_as_dict(client_and_mocks): return_value=None ) - # Key cached as a plain dict (as can happen with Redis serialization) + # Serialized shape from Redis dict; UserApiKeyCache.async_get_cache(model_type=...) yields a model — simulate that. + cached_key_payload = { + "token": "hashed-key-dict", + "access_group_ids": ["ag-to-delete", "ag-other"], + } + + def user_cache_get_dict_when_key_matches(*args, **kwargs): + cache_key = ( + kwargs.get("key") if "key" in kwargs else (args[0] if args else None) + ) + if cache_key == "team_id:team-1": + return None + if cache_key == "hashed-key-dict": + return UserAPIKeyAuth.model_validate(cached_key_payload) + return None + mock_cache.async_get_cache = AsyncMock( - return_value={ - "token": "hashed-key-dict", - "access_group_ids": ["ag-to-delete", "ag-other"], - } + side_effect=user_cache_get_dict_when_key_matches ) resp = client.delete("/v1/access_group/ag-to-delete") From 18996326ef0450139d4c20b29366230b249bbc80 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Tue, 28 Apr 2026 21:17:53 +0000 Subject: [PATCH 31/75] update test cases to fix handle_jwt test cases --- tests/test_litellm/proxy/auth/test_handle_jwt.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index cdb9ae4d9ab..a5764f40809 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -403,9 +403,9 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_role_change(): mock_cache.async_set_cache.assert_called_once() call_kwargs = mock_cache.async_set_cache.call_args assert call_kwargs.kwargs["key"] == "u1" - assert ( - call_kwargs.kwargs["value"]["user_role"] == LitellmUserRoles.PROXY_ADMIN.value - ) + assert isinstance(call_kwargs.kwargs["value"], LiteLLM_UserTable) + assert call_kwargs.kwargs["value"].user_role == LitellmUserRoles.PROXY_ADMIN.value + assert call_kwargs.kwargs["model_type"] == LiteLLM_UserTable @pytest.mark.asyncio @@ -450,7 +450,9 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_team_change(): mock_cache.async_set_cache.assert_called_once() call_kwargs = mock_cache.async_set_cache.call_args assert call_kwargs.kwargs["key"] == "u1" - assert set(call_kwargs.kwargs["value"]["teams"]) == {"team1", "team2"} + assert isinstance(call_kwargs.kwargs["value"], LiteLLM_UserTable) + assert set(call_kwargs.kwargs["value"].teams) == {"team1", "team2"} + assert call_kwargs.kwargs["model_type"] == LiteLLM_UserTable @pytest.mark.asyncio From 13ffb9aed5c4bb2f5aa99f07a235aa07a344f0d4 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Tue, 28 Apr 2026 21:27:55 +0000 Subject: [PATCH 32/75] update auth metric increment comment from greptile --- litellm/proxy/auth/auth_metrics.py | 14 +++++++++++--- 1 file changed, 11 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/auth/auth_metrics.py b/litellm/proxy/auth/auth_metrics.py index 0b463cc920f..7284ba42736 100644 --- a/litellm/proxy/auth/auth_metrics.py +++ b/litellm/proxy/auth/auth_metrics.py @@ -13,6 +13,7 @@ Usage:: """ from litellm._logging import verbose_proxy_logger +from litellm.types.integrations.prometheus import UserAPIKeyLabelValues class AuthMetrics: @@ -42,9 +43,16 @@ class AuthMetrics: try: prom = AuthMetrics._get_prom() if prom is not None: - prom.litellm_auth_combined_view_queries_total.labels( - hashed_api_key=hashed_token - ).inc() + # Counter labelnames include ``hashed_api_key`` plus any + # ``custom_prometheus_metadata_labels`` / ``custom_prometheus_tags`` + # (see ``PrometheusMetricLabels.get_labels``). Use the same + # ``_inc_labeled_counter`` + ``prometheus_label_factory`` path as + # other metrics so label cardinality always matches registration. + prom._inc_labeled_counter( + prom.litellm_auth_combined_view_queries_total, + "litellm_auth_combined_view_queries_total", + UserAPIKeyLabelValues(hashed_api_key=hashed_token), + ) except Exception as e: verbose_proxy_logger.debug( "AuthMetrics.inc_combined_view_query: failed to increment counter: %s", From 90396431cb6e1f6d9d4256fb168e6925e2203c3d Mon Sep 17 00:00:00 2001 From: harish-berri Date: Tue, 28 Apr 2026 21:54:36 +0000 Subject: [PATCH 33/75] lint styling update --- litellm/proxy/proxy_server.py | 16 ++++++++++------ 1 file changed, 10 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 865a46a5478..f54df469d95 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2129,7 +2129,9 @@ async def update_cache( # noqa: PLR0915 if cached_user is None: # do nothing if there is no cache value return - existing_spend_obj = CacheCodec.deserialize(cached_user, LiteLLM_UserTable) + existing_spend_obj = CacheCodec.deserialize( + cached_user, LiteLLM_UserTable + ) if existing_spend_obj is None: return verbose_proxy_logger.debug( @@ -3606,21 +3608,23 @@ class ProxyConfig: verbose_proxy_logger.critical( "LITELLM_MASTER_KEY is not set! All requests will be treated as INTERNAL_USER with no admin access. Set LITELLM_MASTER_KEY for production use." ) - ### USER API KEY CACHE IN-MEMORY TTL ### + ### USER API KEY CACHE TTL (in-memory + Redis when Redis auth sharing is enabled) ### user_api_key_cache_ttl = general_settings.get( "user_api_key_cache_ttl", None ) if user_api_key_cache_ttl is not None: + ttl = float(user_api_key_cache_ttl) + # Mirror TTL on Redis as well when ``litellm_settings.enable_redis_auth_cache`` + # attaches Redis to ``user_api_key_cache``; otherwise DualCache misses in + # memory fall back to a key that outlasts ``user_api_key_cache_ttl``. user_api_key_cache.update_cache_ttl( - default_in_memory_ttl=float(user_api_key_cache_ttl), - default_redis_ttl=None, # user_api_key_cache uses in-memory TTL only; Redis not configured for key lookups + default_in_memory_ttl=ttl, + default_redis_ttl=ttl, ) ### PKCE MULTI-INSTANCE PREREQUISITE CHECK ### # PKCE verifiers are stored in redis_usage_cache when available so they can # be read back by any instance (not just the one that started the auth flow). - # user_api_key_cache is intentionally left in-memory-only to avoid routing - # all API-key lookups through Redis. use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true" if use_pkce and redis_usage_cache is None: global _pkce_no_redis_warning_emitted From 33d9190747203486cbc3fa7a8a702e78c5e20527 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Wed, 29 Apr 2026 00:27:05 +0000 Subject: [PATCH 34/75] correct team_id key bug --- litellm/proxy/auth/user_api_key_auth.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 30f78a50813..4c9861b540b 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1485,11 +1485,8 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 # Only cache when the key is a real team_id (non-team keys must not use key=None). if valid_token.team_id is not None and _team_obj is not None: - # Match get_team_object / spend counters: "team_id:{id}". Serialize for Redis - # (json.dumps) — same as _cache_team_object in auth_checks. - team_cache_key = f"team_id:{valid_token.team_id}" await user_api_key_cache.async_set_cache( - key=team_cache_key, + key=valid_token.team_id, value=_team_obj, model_type=LiteLLM_TeamTableCachedObj, ) # save team table in cache - used for tpm/rpm limiting - tpm_rpm_limiter.py From b6e13cb70abbf64fde83dfa3eba9c324e730a5e2 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Wed, 29 Apr 2026 01:20:07 +0000 Subject: [PATCH 35/75] fix linting --- litellm/proxy/common_utils/user_api_key_cache.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index 2eebc291bb0..7c42aa545d0 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -57,7 +57,9 @@ class UserApiKeyCache(DualCache): def set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) payload = CacheCodec.serialize(value, model_type=model_type) - return super().set_cache(key=key, value=payload, local_only=local_only, **kwargs) + return super().set_cache( + key=key, value=payload, local_only=local_only, **kwargs + ) async def async_set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) From bff26f6c014f2095d4a6447217aab61b44bb2582 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Wed, 29 Apr 2026 01:39:27 +0000 Subject: [PATCH 36/75] fix unused imports --- litellm/proxy/auth/auth_checks.py | 1 - 1 file changed, 1 deletion(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 8bd8eb794e0..333684eae56 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -19,7 +19,6 @@ from pydantic import BaseModel import litellm from litellm._logging import verbose_proxy_logger -from litellm.caching.caching import DualCache from litellm.caching.dual_cache import LimitedSizeOrderedDict from litellm.constants import ( CLI_JWT_EXPIRATION_HOURS, From 25cf097da8cd7f9ca9d3b46ccbfe306b51dcb3c6 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Wed, 29 Apr 2026 20:14:30 +0000 Subject: [PATCH 37/75] add test(tag-routing): prevent header regex bypass for strict plain tags. Add tests to validate the condition improve the conditional readability by naming the plain-tag check explicitly. --- litellm/router_strategy/tag_based_routing.py | 3 +- .../test_router_tag_routing.py | 28 +++++++++++++++++++ 2 files changed, 30 insertions(+), 1 deletion(-) diff --git a/litellm/router_strategy/tag_based_routing.py b/litellm/router_strategy/tag_based_routing.py index 0163f3bbd4f..07143af38a2 100644 --- a/litellm/router_strategy/tag_based_routing.py +++ b/litellm/router_strategy/tag_based_routing.py @@ -106,7 +106,8 @@ def _match_deployment( # check either didn't run (no request tags) or failed (step 1 returned # None). Block the regex path so it cannot circumvent the operator's # strict-tag policy. - strict_tag_check_failed = not match_any and bool(deployment_tags) + deployment_has_plain_tags = deployment_tags is not None and len(deployment_tags) > 0 + strict_tag_check_failed = not match_any and deployment_has_plain_tags if deployment_tag_regex and header_strings and not strict_tag_check_failed: regex_match = _is_valid_deployment_tag_regex( deployment_tag_regex, header_strings diff --git a/tests/test_litellm/router_strategy/test_router_tag_routing.py b/tests/test_litellm/router_strategy/test_router_tag_routing.py index 4424c68f1d9..a6e39ec3c0a 100644 --- a/tests/test_litellm/router_strategy/test_router_tag_routing.py +++ b/tests/test_litellm/router_strategy/test_router_tag_routing.py @@ -346,6 +346,34 @@ def test_tag_routing_with_list_of_tags_match_all(): assert not is_valid_deployment_tag(["default"], ["teamA"], match_any=False) +def test_strict_tag_routing_without_request_tags_blocks_header_regex_fallback(): + """ + When tag_filtering_match_any=False, deployments with plain tags must require + those request tags before header regex can match. A spoofed User-Agent must + not route to a tagged deployment when the request has no tags. + """ + from litellm.router_strategy.tag_based_routing import _match_deployment + + deployment = { + "model_name": "restricted-model", + "litellm_params": { + "model": "gpt-4o", + "tags": ["internal"], + "tag_regex": ["^User-Agent: internal-tool"], + }, + } + + assert ( + _match_deployment( + deployment=deployment, + request_tags=None, + header_strings=["User-Agent: internal-tool"], + match_any=False, + ) + is None + ) + + @pytest.mark.asyncio() async def test_router_free_paid_tier_with_responses_api(): """ From 1277cbe454d1df41cd9ca0e94fab518a3e3124c1 Mon Sep 17 00:00:00 2001 From: shivam Date: Wed, 29 Apr 2026 15:50:55 -0700 Subject: [PATCH 38/75] Add health status pagination controls Made-with: Cursor --- .../ModelsAndEndpointsView.tsx | 39 +++++- .../HealthCheckComponent.test.tsx | 53 ++++++- .../model_dashboard/HealthCheckComponent.tsx | 131 +++++++++++++----- 3 files changed, 174 insertions(+), 49 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx index 514ae673d06..12d90c40f90 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx @@ -62,6 +62,8 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te const [selectedModelId, setSelectedModelId] = useState(null); const [selectedTeamId, setSelectedTeamId] = useState(null); const [selectedTabIndex, setSelectedTabIndex] = useState(0); + const [healthCurrentPage, setHealthCurrentPage] = useState(1); + const healthPageSize = 50; const [showMissingProviderBanner, setShowMissingProviderBanner] = useState(() => { if (typeof window !== "undefined") { return localStorage.getItem("hideMissingProviderBanner") !== "true"; @@ -71,6 +73,10 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te const queryClient = useQueryClient(); const { data: modelDataResponse, isLoading: isLoadingModels, refetch: refetchModels } = useModelsInfo(); + const { data: healthModelDataResponse, isLoading: isLoadingHealthModels } = useModelsInfo( + healthCurrentPage, + healthPageSize, + ); const { data: modelCostMapData, isLoading: isLoadingModelCostMap } = useModelCostMap(); const { data: credentialsResponse, isLoading: isLoadingCredentials } = useCredentials(); const credentialsList = credentialsResponse?.credentials || []; @@ -104,12 +110,12 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te return modelDataResponse.data.map((model: any) => model.model_name); }, [modelDataResponse?.data]); - const allModelIdsOnProxy = useMemo(() => { - if (!modelDataResponse?.data) return []; - return modelDataResponse.data + const healthModelIdsOnProxy = useMemo(() => { + if (!healthModelDataResponse?.data) return []; + return healthModelDataResponse.data .map((model: any) => model.model_info?.id) .filter((id: string | undefined): id is string => Boolean(id)); - }, [modelDataResponse?.data]); + }, [healthModelDataResponse?.data]); const getProviderFromModel = (model: string) => { if (modelCostMapData !== null && modelCostMapData !== undefined) { @@ -125,6 +131,20 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te return transformModelData(modelDataResponse, getProviderFromModel); }, [modelDataResponse?.data, getProviderFromModel]); + const processedHealthModelData = useMemo(() => { + if (!healthModelDataResponse?.data) return { data: [] }; + return transformModelData(healthModelDataResponse, getProviderFromModel); + }, [healthModelDataResponse?.data, getProviderFromModel]); + + const healthPaginationMeta = useMemo(() => { + return { + total_count: healthModelDataResponse?.total_count ?? 0, + current_page: healthModelDataResponse?.current_page ?? healthCurrentPage, + total_pages: healthModelDataResponse?.total_pages ?? 1, + size: healthModelDataResponse?.size ?? healthPageSize, + }; + }, [healthModelDataResponse, healthCurrentPage, healthPageSize]); + const isProxyAdmin = userRole && isProxyAdminRole(userRole); const isInternalUser = userRole && internalUserRoles.includes(userRole); const isUserTeamAdmin = userID && isUserTeamAdminForAnyTeam(teams, userID); @@ -166,7 +186,7 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te const handleRefreshClick = () => { const currentDate = new Date(); - setLastRefreshed(currentDate.toLocaleTimeString([], { hour: '2-digit', minute: '2-digit' })); + setLastRefreshed(currentDate.toLocaleTimeString([], { hour: "2-digit", minute: "2-digit" })); queryClient.invalidateQueries({ queryKey: ["models", "list"] }); refetchModels(); }; @@ -441,11 +461,16 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te { expect(mockIndividualModelHealthCheckCall).not.toHaveBeenCalledWith("token-123", "gpt-4"); }); + it("should show pagination controls and request the next page", async () => { + const onPageChange = vi.fn(); + const modelData = { + data: [ + { + model_name: "gpt-4", + model_info: { id: "deployment-1" }, + litellm_model_name: "gpt-4", + }, + ], + }; + + render( + , + ); + + expect(screen.getByTestId("health-results-count")).toHaveTextContent("Showing 1 - 50 of 75 results"); + + await act(async () => { + screen.getByRole("button", { name: "Next" }).click(); + }); + + expect(onPageChange).toHaveBeenCalledWith(2); + }); + describe("latest_health_checks keyed by model id", () => { it("should show status from latest_health_checks when keys match model ids", async () => { const modelData = { data: [ - { - model_name: "gpt-4", - model_info: { id: "id-alpha" }, - litellm_model_name: "gpt-4", + { + model_name: "gpt-4", + model_info: { id: "id-alpha" }, + litellm_model_name: "gpt-4", }, - { - model_name: "gpt-4", - model_info: { id: "id-beta" }, + { + model_name: "gpt-4", + model_info: { id: "id-beta" }, litellm_model_name: "gpt-4", }, ], diff --git a/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.tsx b/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.tsx index db71733b5ba..c652e9a02ba 100644 --- a/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.tsx +++ b/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.tsx @@ -26,6 +26,16 @@ interface HealthCheckComponentProps { getDisplayModelName: (model: any) => string; setSelectedModelId?: (modelId: string) => void; teams?: Team[] | null; + isLoading?: boolean; + paginationMeta?: { + total_count: number; + current_page: number; + total_pages: number; + size: number; + }; + currentPage?: number; + pageSize?: number; + onPageChange?: (page: number) => void; } const HealthCheckComponent: React.FC = ({ @@ -35,6 +45,11 @@ const HealthCheckComponent: React.FC = ({ getDisplayModelName, setSelectedModelId, teams, + isLoading = false, + paginationMeta, + currentPage = 1, + pageSize = 50, + onPageChange, }) => { const [modelHealthStatuses, setModelHealthStatuses] = useState<{ [key: string]: HealthStatus }>({}); const [selectedModelsForHealth, setSelectedModelsForHealth] = useState([]); @@ -95,19 +110,19 @@ const HealthCheckComponent: React.FC = ({ const fullError = checkData.error_message || undefined; healthStatusMap[modelId] = { - status: checkData.status || "unknown", - lastCheck: checkData.checked_at ? new Date(checkData.checked_at).toLocaleString() : "None", - lastSuccess: - checkData.status === "healthy" - ? checkData.checked_at - ? new Date(checkData.checked_at).toLocaleString() - : "None" - : "None", - loading: false, - error: fullError ? extractMeaningfulError(fullError) : undefined, - fullError: fullError, - successResponse: checkData.status === "healthy" ? checkData : undefined, - }; + status: checkData.status || "unknown", + lastCheck: checkData.checked_at ? new Date(checkData.checked_at).toLocaleString() : "None", + lastSuccess: + checkData.status === "healthy" + ? checkData.checked_at + ? new Date(checkData.checked_at).toLocaleString() + : "None" + : "None", + loading: false, + error: fullError ? extractMeaningfulError(fullError) : undefined, + fullError: fullError, + successResponse: checkData.status === "healthy" ? checkData : undefined, + }; }); } } catch (healthError) { @@ -448,6 +463,12 @@ const HealthCheckComponent: React.FC = ({ } }; + const handlePageChange = (page: number) => { + setSelectedModelsForHealth([]); + setAllModelsSelected(false); + onPageChange?.(page); + }; + const getStatusBadge = (status: string) => { switch (status) { case "healthy": @@ -490,6 +511,34 @@ const HealthCheckComponent: React.FC = ({ setSelectedSuccessDetails(null); }; + const healthTableData = (modelData?.data ?? []).map((model: any) => { + const modelId = model.model_info?.id; + const healthStatus = modelId ? modelHealthStatuses[modelId] : null; + const status = healthStatus || { + status: "none", + lastCheck: "None", + loading: false, + }; + return { + model_name: model.model_name, + model_info: model.model_info, + provider: model.provider, + litellm_model_name: model.litellm_model_name, + health_status: status.status, + last_check: status.lastCheck, + last_success: status.lastSuccess || "None", + health_loading: status.loading, + health_error: status.error, + health_full_error: status.fullError, + }; + }); + + const totalCount = paginationMeta?.total_count ?? healthTableData.length; + const totalPages = paginationMeta?.total_pages ?? 1; + const resultsStart = totalCount > 0 ? (currentPage - 1) * pageSize + 1 : 0; + const resultsEnd = Math.min(currentPage * pageSize, totalCount); + const shouldShowPagination = Boolean(paginationMeta && onPageChange); + return (
@@ -522,6 +571,38 @@ const HealthCheckComponent: React.FC = ({
+ {shouldShowPagination && ( +
+ + {totalCount > 0 + ? `Showing ${resultsStart} - ${resultsEnd} of ${totalCount} results` + : "Showing 0 results"} + + +
+ + +
+
+ )} = ({ setSelectedModelId, teams, )} - data={modelData.data.map((model: any) => { - const modelId = model.model_info?.id; - const healthStatus = modelId ? modelHealthStatuses[modelId] : null; - const status = healthStatus || { - status: "none", - lastCheck: "None", - loading: false, - }; - return { - model_name: model.model_name, - model_info: model.model_info, - provider: model.provider, - litellm_model_name: model.litellm_model_name, - health_status: status.status, - last_check: status.lastCheck, - last_success: status.lastSuccess || "None", - health_loading: status.loading, - health_error: status.error, - health_full_error: status.fullError, - }; - })} - isLoading={false} + data={healthTableData} + isLoading={isLoading} />
From 0b9d06a50948e4aac9337f714aca715120d47cdb Mon Sep 17 00:00:00 2001 From: shivam Date: Wed, 29 Apr 2026 16:25:57 -0700 Subject: [PATCH 39/75] Fix model e2e result count locator Made-with: Cursor --- .../e2e_tests/tests/modelsPage/addModel.spec.ts | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts b/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts index 8834724f76b..0f11bb5b8db 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts +++ b/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts @@ -141,7 +141,7 @@ test.describe("Add Model", () => { await page.waitForTimeout(1000); // Verify the model appears in the results count (not "Showing 0 results") - await expect(page.getByText(/Showing \d+ - \d+ of \d+ results/)).toBeVisible({ timeout: 15_000 }); + await expect(page.getByTestId("models-results-count")).toBeVisible({ timeout: 15_000 }); // Verify the model name appears in the table body const tableBody = page.locator("table tbody"); @@ -181,7 +181,7 @@ test.describe("Add Model", () => { await page.waitForTimeout(1000); // Verify the model appears in the results count (not "Showing 0 results") - await expect(page.getByText(/Showing \d+ - \d+ of \d+ results/)).toBeVisible({ timeout: 15_000 }); + await expect(page.getByTestId("models-results-count")).toBeVisible({ timeout: 15_000 }); // Verify the wildcard model appears in the table body (wildcard models show as "cohere/*") const tableBody = page.locator("table tbody"); From a7a3f0a19d69b1744b9b5c7c2a6b86b0339f8f24 Mon Sep 17 00:00:00 2001 From: shivam Date: Wed, 29 Apr 2026 16:49:25 -0700 Subject: [PATCH 40/75] Address health pagination review feedback Made-with: Cursor --- .../e2e_tests/tests/modelsPage/addModel.spec.ts | 8 ++++++-- .../models-and-endpoints/ModelsAndEndpointsView.tsx | 11 ++++++----- .../model_dashboard/HealthCheckComponent.tsx | 11 +++++++---- 3 files changed, 19 insertions(+), 11 deletions(-) diff --git a/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts b/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts index 0f11bb5b8db..c3bd8489027 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts +++ b/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts @@ -141,7 +141,9 @@ test.describe("Add Model", () => { await page.waitForTimeout(1000); // Verify the model appears in the results count (not "Showing 0 results") - await expect(page.getByTestId("models-results-count")).toBeVisible({ timeout: 15_000 }); + await expect(page.getByTestId("models-results-count")).toHaveText(/Showing \d+ - \d+ of \d+ results/, { + timeout: 15_000, + }); // Verify the model name appears in the table body const tableBody = page.locator("table tbody"); @@ -181,7 +183,9 @@ test.describe("Add Model", () => { await page.waitForTimeout(1000); // Verify the model appears in the results count (not "Showing 0 results") - await expect(page.getByTestId("models-results-count")).toBeVisible({ timeout: 15_000 }); + await expect(page.getByTestId("models-results-count")).toHaveText(/Showing \d+ - \d+ of \d+ results/, { + timeout: 15_000, + }); // Verify the wildcard model appears in the table body (wildcard models show as "cohere/*") const tableBody = page.locator("table tbody"); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx index 12d90c40f90..c51df7f8144 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx @@ -46,6 +46,8 @@ interface GlobalRetryPolicyObject { [retryPolicyKey: string]: number; } +const HEALTH_PAGE_SIZE = 50; + const ModelsAndEndpointsView: React.FC = ({ premiumUser, teams }) => { const { accessToken, token, userRole, userId: userID } = useAuthorized(); const [addModelForm] = Form.useForm(); @@ -63,7 +65,6 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te const [selectedTeamId, setSelectedTeamId] = useState(null); const [selectedTabIndex, setSelectedTabIndex] = useState(0); const [healthCurrentPage, setHealthCurrentPage] = useState(1); - const healthPageSize = 50; const [showMissingProviderBanner, setShowMissingProviderBanner] = useState(() => { if (typeof window !== "undefined") { return localStorage.getItem("hideMissingProviderBanner") !== "true"; @@ -75,7 +76,7 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te const { data: modelDataResponse, isLoading: isLoadingModels, refetch: refetchModels } = useModelsInfo(); const { data: healthModelDataResponse, isLoading: isLoadingHealthModels } = useModelsInfo( healthCurrentPage, - healthPageSize, + HEALTH_PAGE_SIZE, ); const { data: modelCostMapData, isLoading: isLoadingModelCostMap } = useModelCostMap(); const { data: credentialsResponse, isLoading: isLoadingCredentials } = useCredentials(); @@ -141,9 +142,9 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te total_count: healthModelDataResponse?.total_count ?? 0, current_page: healthModelDataResponse?.current_page ?? healthCurrentPage, total_pages: healthModelDataResponse?.total_pages ?? 1, - size: healthModelDataResponse?.size ?? healthPageSize, + size: healthModelDataResponse?.size ?? HEALTH_PAGE_SIZE, }; - }, [healthModelDataResponse, healthCurrentPage, healthPageSize]); + }, [healthModelDataResponse, healthCurrentPage]); const isProxyAdmin = userRole && isProxyAdminRole(userRole); const isInternalUser = userRole && internalUserRoles.includes(userRole); @@ -469,7 +470,7 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te isLoading={isLoadingHealthModels} paginationMeta={healthPaginationMeta} currentPage={healthCurrentPage} - pageSize={healthPageSize} + pageSize={HEALTH_PAGE_SIZE} onPageChange={setHealthCurrentPage} /> diff --git a/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.tsx b/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.tsx index c652e9a02ba..d9b7fdd149c 100644 --- a/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.tsx +++ b/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.tsx @@ -466,6 +466,7 @@ const HealthCheckComponent: React.FC = ({ const handlePageChange = (page: number) => { setSelectedModelsForHealth([]); setAllModelsSelected(false); + setModelHealthStatuses({}); onPageChange?.(page); }; @@ -533,11 +534,13 @@ const HealthCheckComponent: React.FC = ({ }; }); - const totalCount = paginationMeta?.total_count ?? healthTableData.length; - const totalPages = paginationMeta?.total_pages ?? 1; - const resultsStart = totalCount > 0 ? (currentPage - 1) * pageSize + 1 : 0; - const resultsEnd = Math.min(currentPage * pageSize, totalCount); const shouldShowPagination = Boolean(paginationMeta && onPageChange); + const totalCount = paginationMeta?.total_count ?? 0; + const totalPages = paginationMeta?.total_pages ?? 1; + const pageForDisplay = paginationMeta?.current_page ?? currentPage; + const pageSizeForDisplay = paginationMeta?.size ?? pageSize; + const resultsStart = shouldShowPagination && totalCount > 0 ? (pageForDisplay - 1) * pageSizeForDisplay + 1 : 0; + const resultsEnd = shouldShowPagination ? Math.min(pageForDisplay * pageSizeForDisplay, totalCount) : 0; return (
From 181e99b9962ae3998ac91338b45bfbcdad6be60c Mon Sep 17 00:00:00 2001 From: shivam Date: Wed, 29 Apr 2026 17:16:32 -0700 Subject: [PATCH 41/75] Fix health pagination review issues Made-with: Cursor --- .../ModelsAndEndpointsView.tsx | 1 + .../HealthCheckComponent.test.tsx | 39 +++++++++++-------- 2 files changed, 23 insertions(+), 17 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx index c51df7f8144..7c162e2056b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx @@ -188,6 +188,7 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te const handleRefreshClick = () => { const currentDate = new Date(); setLastRefreshed(currentDate.toLocaleTimeString([], { hour: "2-digit", minute: "2-digit" })); + setHealthCurrentPage(1); queryClient.invalidateQueries({ queryKey: ["models", "list"] }); refetchModels(); }; diff --git a/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.test.tsx b/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.test.tsx index 738dab9d54f..1a4c0ed9ff5 100644 --- a/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.test.tsx +++ b/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.test.tsx @@ -104,23 +104,28 @@ describe("HealthCheckComponent", () => { ], }; - render( - , - ); + await act(async () => { + render( + , + ); + }); + await act(async () => { + await new Promise((r) => setTimeout(r, 0)); + }); expect(screen.getByTestId("health-results-count")).toHaveTextContent("Showing 1 - 50 of 75 results"); From 2c4979e5e8dab3e6e1d64c12ba7c8f726d726f8e Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 30 Apr 2026 09:10:24 +0530 Subject: [PATCH 42/75] Fix review --- litellm/rerank_api/main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index 0dd18b6a5bf..e27585116ce 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -177,7 +177,7 @@ def rerank( # noqa: PLR0915 model=model, user=user, optional_params=dict(optional_rerank_params), - litellm_params=rerank_litellm_params, + litellm_params=dict(rerank_litellm_params), custom_llm_provider=_custom_llm_provider, ) From 3e540e1d81cc249a9e482a31f4e9a695e10d3cc9 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 30 Apr 2026 17:17:39 +0000 Subject: [PATCH 43/75] Fix container file proxy routing --- .../container_endpoints/handler_factory.py | 6 ++ .../test_azure_container_transformation.py | 88 ++++++++++++++++++- 2 files changed, 91 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/container_endpoints/handler_factory.py b/litellm/proxy/container_endpoints/handler_factory.py index 465d1fa9d79..0bf5fe5f547 100644 --- a/litellm/proxy/container_endpoints/handler_factory.py +++ b/litellm/proxy/container_endpoints/handler_factory.py @@ -19,6 +19,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_headers, get_custom_llm_provider_from_request_query, ) +from litellm.responses.utils import ResponsesAPIRequestUtils def _load_endpoints_config() -> Dict: @@ -230,6 +231,7 @@ async def _process_binary_request( return Response( content=content, + headers=dict(fastapi_response.headers), media_type=content_type, ) @@ -287,6 +289,10 @@ async def _process_multipart_upload_request( or get_custom_llm_provider_from_request_query(request=request) or "openai" ) + decoded_container_id = ResponsesAPIRequestUtils._decode_container_id(container_id) + decoded_provider = decoded_container_id.get("custom_llm_provider") + if decoded_provider and custom_llm_provider == "openai": + custom_llm_provider = decoded_provider data["container_id"] = container_id data["custom_llm_provider"] = custom_llm_provider diff --git a/tests/test_litellm/containers/test_azure_container_transformation.py b/tests/test_litellm/containers/test_azure_container_transformation.py index 94ab1ae6b3d..1ea598c54d7 100644 --- a/tests/test_litellm/containers/test_azure_container_transformation.py +++ b/tests/test_litellm/containers/test_azure_container_transformation.py @@ -547,7 +547,9 @@ class TestAzureContainerKnownFailureRegressions: captured["route_type"] = route_type return {"id": "cfile_abc"} - from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) monkeypatch.setattr( ProxyBaseLLMRequestProcessing, @@ -584,6 +586,7 @@ class TestAzureContainerKnownFailureRegressions: async def test_regression_binary_file_request_routes_through_proxy_processor( self, monkeypatch ): + from fastapi import Response from starlette.requests import Request from litellm.proxy.container_endpoints import handler_factory @@ -605,9 +608,12 @@ class TestAzureContainerKnownFailureRegressions: ): captured["data"] = self.data captured["route_type"] = route_type + fastapi_response.headers["x-litellm-call-id"] = "call-123" return b"csv-bytes" - from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) monkeypatch.setattr( ProxyBaseLLMRequestProcessing, @@ -624,7 +630,7 @@ class TestAzureContainerKnownFailureRegressions: "query_string": b"", } ) - fastapi_response = MagicMock() + fastapi_response = Response() response = await handler_factory._process_binary_request( request=request, @@ -640,3 +646,79 @@ class TestAzureContainerKnownFailureRegressions: assert captured["data"]["custom_llm_provider"] == "openai" assert response.status_code == 200 assert response.body == b"csv-bytes" + assert response.headers["x-litellm-call-id"] == "call-123" + + @pytest.mark.asyncio + async def test_regression_multipart_upload_request_uses_provider_from_managed_id( + self, monkeypatch + ): + from starlette.requests import Request + + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) + from litellm.proxy.common_utils import http_parsing_utils + from litellm.proxy.container_endpoints import handler_factory + + encoded_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="azure", + model_id="model_abc123", + container_id="cntr_123", + ) + captured = {} + + async def _mock_get_form_data(request): + return {"file": "ignored"} + + async def _mock_convert_upload_files_to_file_data(form_data): + return {"file": [("data.csv", b"csv-bytes", "text/csv")]} + + async def _mock_base_process_llm_request( + self, + request, + fastapi_response, + user_api_key_dict, + route_type, + **kwargs, + ): + captured["data"] = self.data + captured["route_type"] = route_type + return {"id": "cfile_abc"} + + monkeypatch.setattr( + http_parsing_utils, + "get_form_data", + _mock_get_form_data, + ) + monkeypatch.setattr( + http_parsing_utils, + "convert_upload_files_to_file_data", + _mock_convert_upload_files_to_file_data, + ) + monkeypatch.setattr( + ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + _mock_base_process_llm_request, + ) + + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/containers/id/files", + "headers": [], + "query_string": b"", + } + ) + + await handler_factory._process_multipart_upload_request( + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=MagicMock(), + route_type="aupload_container_file", + container_id=encoded_id, + ) + + assert captured["route_type"] == "aupload_container_file" + assert captured["data"]["container_id"] == encoded_id + assert captured["data"]["custom_llm_provider"] == "azure" From 02ee8e2e44f22c8e112321fda27199164aa9abf2 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 30 Apr 2026 17:39:50 +0000 Subject: [PATCH 44/75] Remove duplicate container provider decoding --- litellm/proxy/container_endpoints/handler_factory.py | 5 ----- 1 file changed, 5 deletions(-) diff --git a/litellm/proxy/container_endpoints/handler_factory.py b/litellm/proxy/container_endpoints/handler_factory.py index 0bf5fe5f547..33df72ceaeb 100644 --- a/litellm/proxy/container_endpoints/handler_factory.py +++ b/litellm/proxy/container_endpoints/handler_factory.py @@ -19,7 +19,6 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_headers, get_custom_llm_provider_from_request_query, ) -from litellm.responses.utils import ResponsesAPIRequestUtils def _load_endpoints_config() -> Dict: @@ -289,10 +288,6 @@ async def _process_multipart_upload_request( or get_custom_llm_provider_from_request_query(request=request) or "openai" ) - decoded_container_id = ResponsesAPIRequestUtils._decode_container_id(container_id) - decoded_provider = decoded_container_id.get("custom_llm_provider") - if decoded_provider and custom_llm_provider == "openai": - custom_llm_provider = decoded_provider data["container_id"] = container_id data["custom_llm_provider"] = custom_llm_provider From 2af55422f6a33e713741e26ff19ed067150185c8 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 30 Apr 2026 17:50:01 +0000 Subject: [PATCH 45/75] Fix azure container multipart provider assertion --- .../containers/test_azure_container_transformation.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test_litellm/containers/test_azure_container_transformation.py b/tests/test_litellm/containers/test_azure_container_transformation.py index 1ea598c54d7..45fa23bcb6e 100644 --- a/tests/test_litellm/containers/test_azure_container_transformation.py +++ b/tests/test_litellm/containers/test_azure_container_transformation.py @@ -721,4 +721,4 @@ class TestAzureContainerKnownFailureRegressions: assert captured["route_type"] == "aupload_container_file" assert captured["data"]["container_id"] == encoded_id - assert captured["data"]["custom_llm_provider"] == "azure" + assert captured["data"]["custom_llm_provider"] == "openai" From 3bbb5c7fd74d1f4c08abd05c2a8ccf7e5c20d37f Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 30 Apr 2026 17:56:50 +0000 Subject: [PATCH 46/75] Remove dead _transform_response_for_json_mode wrapper The wrapper had no production callers after transform_parsed_response was refactored to call _resolve_json_mode_non_streaming directly. Updated the parametrized test to call the underlying method. --- litellm/llms/anthropic/chat/transformation.py | 11 ----------- tests/llm_translation/test_anthropic_completion.py | 2 +- 2 files changed, 1 insertion(+), 12 deletions(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 9aaaa39fd7a..2c5d7901ba2 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1587,17 +1587,6 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): filtered_tools = [t for i, t in enumerate(tool_calls) if i not in json_indices] return None, filtered_tools, extra_content - def _transform_response_for_json_mode( - self, - json_mode: Optional[bool], - tool_calls: List[ChatCompletionToolCallChunk], - ) -> Optional[LitellmMessage]: - replacement, _, _ = self._resolve_json_mode_non_streaming( - json_mode=json_mode, - tool_calls=tool_calls, - ) - return replacement - def extract_response_content(self, completion_response: dict) -> Tuple[ str, Optional[List[Any]], diff --git a/tests/llm_translation/test_anthropic_completion.py b/tests/llm_translation/test_anthropic_completion.py index fdf8c24ac9e..7b2b6bed6a4 100644 --- a/tests/llm_translation/test_anthropic_completion.py +++ b/tests/llm_translation/test_anthropic_completion.py @@ -870,7 +870,7 @@ from litellm.constants import RESPONSE_FORMAT_TOOL_NAME def test_anthropic_json_mode_and_tool_call_response( json_mode, tool_calls, expect_null_response ): - result = litellm.AnthropicConfig()._transform_response_for_json_mode( + result, _, _ = litellm.AnthropicConfig()._resolve_json_mode_non_streaming( json_mode=json_mode, tool_calls=tool_calls, ) From b5df9d9778d1f9bbd34bdae9a2f4631267dded9e Mon Sep 17 00:00:00 2001 From: mateo-berri Date: Thu, 30 Apr 2026 18:24:38 +0000 Subject: [PATCH 47/75] test(vertex_ai): add e2e tests for rerank userLabels propagation Cover the full litellm.rerank()/arerank() path with HTTP mocked, asserting metadata.requester_metadata reaches the Discovery Engine :rank body as userLabels (and stays absent when no metadata is set). Catches plumbing regressions that unit tests on transform_rerank_request alone would miss. --- .../test_vertex_ai_rerank_userlabels_e2e.py | 182 ++++++++++++++++++ 1 file changed, 182 insertions(+) create mode 100644 tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_userlabels_e2e.py diff --git a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_userlabels_e2e.py b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_userlabels_e2e.py new file mode 100644 index 00000000000..0b28fa9abc9 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_userlabels_e2e.py @@ -0,0 +1,182 @@ +""" +End-to-end tests for Vertex AI rerank `userLabels` propagation. + +These tests go through the full `litellm.rerank()` call path with the HTTP +layer mocked, so they catch plumbing bugs (e.g. `litellm_params` losing +`metadata` between the rerank entrypoint and the Vertex transform) that +unit tests on `VertexAIRerankConfig.transform_rerank_request` miss. +""" + +import asyncio +import json +import os +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +import litellm +import litellm.llms.vertex_ai.rerank.transformation + + +def _extract_body(call_kwargs): + """The rerank handler sends `data=json.dumps(...)`, not `json=...`.""" + if "json" in call_kwargs and call_kwargs["json"] is not None: + return call_kwargs["json"] + raw = call_kwargs.get("data") + if isinstance(raw, (bytes, bytearray)): + raw = raw.decode("utf-8") + return json.loads(raw) if raw else None + + +def _make_mock_rank_response(): + mock_response = MagicMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.json.return_value = { + "records": [ + {"id": "0", "score": 0.9, "title": "doc 0", "content": "hello"}, + {"id": "1", "score": 0.1, "title": "doc 1", "content": "world"}, + ] + } + mock_response.text = '{"records": []}' + return mock_response + + +def _make_async_mock_rank_response(): + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.json = MagicMock( + return_value={ + "records": [ + {"id": "0", "score": 0.9, "title": "doc 0", "content": "hello"}, + ] + } + ) + mock_response.text = '{"records": []}' + return mock_response + + +@pytest.fixture +def clean_vertex_env(): + saved = {} + for var in ( + "GOOGLE_APPLICATION_CREDENTIALS", + "GOOGLE_CLOUD_PROJECT", + "VERTEXAI_PROJECT", + "VERTEXAI_CREDENTIALS", + "VERTEX_AI_CREDENTIALS", + "VERTEX_PROJECT", + "VERTEX_LOCATION", + "VERTEX_AI_PROJECT", + ): + if var in os.environ: + saved[var] = os.environ.pop(var) + yield + for var, value in saved.items(): + os.environ[var] = value + + +def _patch_vertex_auth(): + return patch.object( + litellm.llms.vertex_ai.rerank.transformation.VertexAIRerankConfig, + "_ensure_access_token", + return_value=("test-access-token", "test-project-2049"), + ) + + +def test_rerank_userlabels_propagates_from_metadata_sync(clean_vertex_env): + """ + `litellm.rerank(metadata={"requester_metadata": {...}})` must end up as + `userLabels` on the Discovery Engine `:rank` request body. + """ + captured = {} + + def fake_post(*args, **kwargs): + captured["body"] = _extract_body(kwargs) + return _make_mock_rank_response() + + with ( + _patch_vertex_auth(), + patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + side_effect=fake_post, + ), + ): + litellm.rerank( + model="vertex_ai/semantic-ranker-default@latest", + query="what is gemini?", + documents=["hello", "world"], + vertex_project="test-project-2049", + vertex_credentials='{"type": "service_account"}', + metadata={"requester_metadata": {"team": "platform", "env": "prod"}}, + ) + + body = captured["body"] + assert body is not None, "expected POST body to be captured" + assert "userLabels" in body, ( + "Vertex rerank request body is missing `userLabels` — metadata was " + "lost between litellm.rerank() and transform_rerank_request. " + f"body keys: {sorted(body.keys())}" + ) + assert body["userLabels"] == {"team": "platform", "env": "prod"} + + +def test_rerank_userlabels_propagates_from_metadata_async(clean_vertex_env): + """Same as the sync test, but through `litellm.arerank`.""" + captured = {} + + async def fake_post(*args, **kwargs): + captured["body"] = _extract_body(kwargs) + return _make_async_mock_rank_response() + + with ( + _patch_vertex_auth(), + patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + side_effect=fake_post, + ), + ): + asyncio.run( + litellm.arerank( + model="vertex_ai/semantic-ranker-default@latest", + query="what is gemini?", + documents=["hello", "world"], + vertex_project="test-project-2049", + vertex_credentials='{"type": "service_account"}', + metadata={"requester_metadata": {"team": "platform"}}, + ) + ) + + body = captured["body"] + assert body is not None + assert body.get("userLabels") == {"team": "platform"} + + +def test_rerank_userlabels_absent_when_no_metadata(clean_vertex_env): + """No metadata → no `userLabels` key (don't send empty maps).""" + captured = {} + + def fake_post(*args, **kwargs): + captured["body"] = _extract_body(kwargs) + return _make_mock_rank_response() + + with ( + _patch_vertex_auth(), + patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + side_effect=fake_post, + ), + ): + litellm.rerank( + model="vertex_ai/semantic-ranker-default@latest", + query="what is gemini?", + documents=["hello", "world"], + vertex_project="test-project-2049", + vertex_credentials='{"type": "service_account"}', + ) + + body = captured["body"] + assert body is not None + assert "userLabels" not in body From 21c7864d752ba990aa7ca1b6f2432e12a4bae9a4 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 30 Apr 2026 18:44:03 +0000 Subject: [PATCH 48/75] Fix Vertex label metadata fallback --- litellm/llms/vertex_ai/common_utils.py | 2 +- .../test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py | 3 +++ 2 files changed, 4 insertions(+), 1 deletion(-) diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index 4f69e3bbad4..fae175612b1 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -69,7 +69,7 @@ def pop_vertex_request_labels( raw = optional_params.pop("labels") if isinstance(raw, dict): labels = {k: v for k, v in raw.items() if isinstance(v, str)} - if labels is None: + if not labels: labels = vertex_request_labels_from_litellm_params(litellm_params) return labels if labels else None diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py index 01c11e15543..7cb3faf6177 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py @@ -1482,6 +1482,9 @@ def test_pop_vertex_request_labels_prefers_explicit_labels_then_metadata(): optional2: dict = {} assert pop_vertex_request_labels(optional2, litellm_params) == {"team": "x"} + optional3 = {"labels": {"team": 123}} + assert pop_vertex_request_labels(optional3, litellm_params) == {"team": "x"} + def test_pop_vertex_request_labels_uses_litellm_metadata_when_metadata_absent(): optional: dict = {} From 3e73381449270a76541c7e59253a75e3116f134b Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 30 Apr 2026 18:57:03 +0000 Subject: [PATCH 49/75] fix(router): always unwrap managed container_id even when model_id is empty Managed cntr_... IDs can be encoded with an empty model_id (e.g. streaming responses without router metadata, or target_model_names=[]). The previous guard only unwrapped when model_id was truthy, so the raw cntr_... token leaked to the upstream provider, which rejects it. Always swap in decoded["response_id"] when it differs from the input, and keep the model_id check only for deciding whether to fan out via _ageneric_api_call_with_fallbacks. --- litellm/router.py | 4 ++- .../test_router_endpoints.py | 33 +++++++++++++++++++ 2 files changed, 36 insertions(+), 1 deletion(-) diff --git a/litellm/router.py b/litellm/router.py index a11281bb151..7449492e89f 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -5275,9 +5275,11 @@ class Router: container_id = kwargs.get("container_id") if isinstance(container_id, str): decoded = ResponsesAPIRequestUtils._decode_container_id(container_id) + original_id = decoded.get("response_id", container_id) + if original_id != container_id: + kwargs["container_id"] = original_id model_id = decoded.get("model_id") if model_id: - kwargs["container_id"] = decoded.get("response_id", container_id) kwargs["model"] = model_id decoded_provider = decoded.get("custom_llm_provider") if decoded_provider and kwargs.get("custom_llm_provider") == "openai": diff --git a/tests/router_unit_tests/test_router_endpoints.py b/tests/router_unit_tests/test_router_endpoints.py index 93bb84dfe35..19d7155e4aa 100644 --- a/tests/router_unit_tests/test_router_endpoints.py +++ b/tests/router_unit_tests/test_router_endpoints.py @@ -1170,3 +1170,36 @@ async def test_init_containers_api_endpoints_managed_id_routes_via_generic_fallb assert call_kw["container_id"] == "cfile_upstream_abc" assert call_kw["file_id"] == "cfile_xyz" assert call_kw["custom_llm_provider"] == "azure" + + +@pytest.mark.asyncio +async def test_init_containers_api_endpoints_managed_id_without_model_id_unwraps(): + """ + Managed ``cntr_`` IDs may be encoded with an empty ``model_id`` (e.g. when a + streaming response had no router metadata). The router must still unwrap the + managed ID before calling the upstream provider — otherwise the raw + ``cntr_...`` token leaks downstream and the provider rejects it. + """ + from litellm.responses.utils import ResponsesAPIRequestUtils + + router = Router(model_list=[]) + mock_original_function = AsyncMock(return_value={"ok": True}) + + managed_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="openai", + model_id=None, + container_id="cfile_upstream_abc", + ) + + await router._init_containers_api_endpoints( + original_function=mock_original_function, + custom_llm_provider="openai", + container_id=managed_id, + file_id="cfile_xyz", + ) + + mock_original_function.assert_called_once() + call_kw = mock_original_function.call_args.kwargs + assert call_kw["container_id"] == "cfile_upstream_abc" + assert call_kw["file_id"] == "cfile_xyz" + assert call_kw["custom_llm_provider"] == "openai" From 54f762eeebf773c8b5783f3b0c24502c60b707d8 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Thu, 30 Apr 2026 19:05:04 +0000 Subject: [PATCH 50/75] remove unrelated debugger tooling --- .vscode/launch.json | 24 ------------------------ 1 file changed, 24 deletions(-) delete mode 100644 .vscode/launch.json diff --git a/.vscode/launch.json b/.vscode/launch.json deleted file mode 100644 index e48948319b0..00000000000 --- a/.vscode/launch.json +++ /dev/null @@ -1,24 +0,0 @@ -{ - "version": "0.2.0", - "configurations": [ - { - "name": "LiteLLM proxy (debug)", - "type": "debugpy", - "request": "launch", - "program": "${workspaceFolder}/litellm/proxy/proxy_cli.py", - "cwd": "${workspaceFolder}", - "args": [ - "--host", - "0.0.0.0", - "--config", - "dev_config.yaml", - "--port", - "4000", - "--num_workers", - "1" - ], - "console": "integratedTerminal", - "justMyCode": false - } - ] -} \ No newline at end of file From 4dc950a24252f48bf3ac7af8c1738c5f320441c6 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Thu, 30 Apr 2026 19:11:29 +0000 Subject: [PATCH 51/75] fix linting --- litellm/proxy/auth/handle_jwt.py | 2 +- litellm/proxy/auth/user_api_key_auth.py | 1 - 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index c55cd7bb05e..71411bed7fd 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -1460,7 +1460,7 @@ class JWTAuthManager: user_object: Optional[LiteLLM_UserTable], user_id: Optional[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, parent_otel_span: Optional[Span], proxy_logging_obj: ProxyLogging, team_id_upsert: Optional[bool], diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index a94abbc9d42..d87bd9e7d59 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -20,7 +20,6 @@ from fastapi.security.api_key import APIKeyHeader import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging -from litellm.caching import DualCache from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value From 46ba48e69a9020c8ae8984cba420bd89842a8fd2 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 30 Apr 2026 19:51:02 +0000 Subject: [PATCH 52/75] fix(router): apply decoded provider for managed container IDs without model_id A managed cntr_ ID can encode a non-OpenAI provider (e.g. azure) with an empty model_id when streaming events have no router model_info.id. The provider override was nested inside 'if model_id:', so such IDs unwrapped the container_id but kept custom_llm_provider='openai', routing the request to the wrong upstream. Hoist the override out of the model_id guard. --- litellm/router.py | 6 ++-- .../test_router_endpoints.py | 33 +++++++++++++++++++ 2 files changed, 36 insertions(+), 3 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index 7449492e89f..676a9e6e205 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -5278,12 +5278,12 @@ class Router: original_id = decoded.get("response_id", container_id) if original_id != container_id: kwargs["container_id"] = original_id + decoded_provider = decoded.get("custom_llm_provider") + if decoded_provider and kwargs.get("custom_llm_provider") == "openai": + kwargs["custom_llm_provider"] = decoded_provider model_id = decoded.get("model_id") if model_id: kwargs["model"] = model_id - decoded_provider = decoded.get("custom_llm_provider") - if decoded_provider and kwargs.get("custom_llm_provider") == "openai": - kwargs["custom_llm_provider"] = decoded_provider return await self._ageneric_api_call_with_fallbacks( original_function=original_function, **kwargs, diff --git a/tests/router_unit_tests/test_router_endpoints.py b/tests/router_unit_tests/test_router_endpoints.py index 19d7155e4aa..0ce2dec9b56 100644 --- a/tests/router_unit_tests/test_router_endpoints.py +++ b/tests/router_unit_tests/test_router_endpoints.py @@ -1203,3 +1203,36 @@ async def test_init_containers_api_endpoints_managed_id_without_model_id_unwraps assert call_kw["container_id"] == "cfile_upstream_abc" assert call_kw["file_id"] == "cfile_xyz" assert call_kw["custom_llm_provider"] == "openai" + + +@pytest.mark.asyncio +async def test_init_containers_api_endpoints_managed_id_without_model_id_applies_decoded_provider(): + """ + A managed ``cntr_`` ID can encode a non-OpenAI provider (e.g. ``azure``) with + an empty ``model_id`` (streaming events without router ``model_info.id``). + The router must still apply the decoded provider so the request routes to + the correct upstream — not stay on the default ``openai``. + """ + from litellm.responses.utils import ResponsesAPIRequestUtils + + router = Router(model_list=[]) + mock_original_function = AsyncMock(return_value={"ok": True}) + + managed_id = ResponsesAPIRequestUtils._build_container_id( + custom_llm_provider="azure", + model_id=None, + container_id="cfile_upstream_abc", + ) + + await router._init_containers_api_endpoints( + original_function=mock_original_function, + custom_llm_provider="openai", + container_id=managed_id, + file_id="cfile_xyz", + ) + + mock_original_function.assert_called_once() + call_kw = mock_original_function.call_args.kwargs + assert call_kw["container_id"] == "cfile_upstream_abc" + assert call_kw["file_id"] == "cfile_xyz" + assert call_kw["custom_llm_provider"] == "azure" From 1f72ae5442b1198c16d607173b9ebee2eb8ae363 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Thu, 30 Apr 2026 22:06:33 +0000 Subject: [PATCH 53/75] fix: add type hint ignore for user_api_key_cache and update docstring for EventDrivenCacheCoordinator --- litellm/proxy/auth/user_api_key_auth.py | 2 +- litellm/proxy/common_utils/cache_coordinator.py | 5 ++++- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index d87bd9e7d59..b13e6cebf97 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -345,7 +345,7 @@ async def _fetch_global_spend_with_event_coordination( return await _global_spend_coordinator.get_or_load( cache_key=cache_key, - cache=user_api_key_cache, + cache=user_api_key_cache, # pyright: ignore[reportArgumentType] load_fn=_load_global_spend, ) diff --git a/litellm/proxy/common_utils/cache_coordinator.py b/litellm/proxy/common_utils/cache_coordinator.py index 92d9916098d..fc46ab6d78e 100644 --- a/litellm/proxy/common_utils/cache_coordinator.py +++ b/litellm/proxy/common_utils/cache_coordinator.py @@ -52,6 +52,9 @@ class EventDrivenCacheCoordinator: - Other requests: wait for the signal, then read from cache. Create one instance per resource (e.g. one for global spend, one for feature flags). + + Args: + log_prefix: Prefix for debug log messages. """ def __init__(self, log_prefix: str = "[CACHE]"): @@ -150,7 +153,7 @@ class EventDrivenCacheCoordinator: elapsed_ms, value, ) - + await cache.async_set_cache(key=cache_key, value=value) if self._log_prefix: verbose_proxy_logger.debug("%s Result cached", self._log_prefix) From 76fb9d3c966516ca0a64a64a05a3687661b7c622 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Thu, 30 Apr 2026 22:13:11 +0000 Subject: [PATCH 54/75] refactor: remove AuthMetrics and related combined_view query metrics This commit deletes the AuthMetrics class and its associated methods, which were responsible for tracking combined_view SQL query metrics. The PrometheusLogger integration has been updated to remove references to these metrics, streamlining the codebase. Additionally, minor whitespace adjustments were made in the cache coordinator for consistency. --- litellm/integrations/prometheus.py | 13 ---- litellm/proxy/auth/auth_checks.py | 2 - litellm/proxy/auth/auth_metrics.py | 60 ------------------- .../proxy/common_utils/cache_coordinator.py | 2 +- litellm/types/integrations/prometheus.py | 7 --- 5 files changed, 1 insertion(+), 83 deletions(-) delete mode 100644 litellm/proxy/auth/auth_metrics.py diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 1a899cf0e83..d9e57ee7cee 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -266,19 +266,6 @@ class PrometheusLogger(CustomLogger): # LiteLLM Virtual API KEY metrics ######################################## - # Auth DB load diagnostic: count direct combined_view SQL queries. - # Each increment means a virtual-key cache miss that hit the DB. - # Useful for validating that enable_redis_auth_cache is working. - self.litellm_auth_combined_view_queries_total = self._counter_factory( - "litellm_auth_combined_view_queries_total", - "Number of times the combined_view SQL query was issued for virtual-key auth. " - "Each count is a cache miss that hit the database. Use to validate " - "enable_redis_auth_cache is reducing DB load.", - labelnames=self.get_labels_for_metric( - "litellm_auth_combined_view_queries_total" - ), - ) - # Remaining MODEL RPM limit for API Key self.litellm_remaining_api_key_requests_for_model = self._gauge_factory( "litellm_remaining_api_key_requests_for_model", diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 9266a917c5f..113a8f538c0 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -59,7 +59,6 @@ from litellm.proxy._types import ( SpecialModelNames, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_metrics import AuthMetrics from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.guardrails.tool_name_extraction import ( @@ -2262,7 +2261,6 @@ async def _fetch_key_object_from_db_with_reconnect( Fetch key object from DB and retry once if a DB connection error can be healed. """ try: - AuthMetrics.inc_combined_view_query(hashed_token) return await prisma_client.get_data( token=hashed_token, table_name="combined_view", diff --git a/litellm/proxy/auth/auth_metrics.py b/litellm/proxy/auth/auth_metrics.py deleted file mode 100644 index 7284ba42736..00000000000 --- a/litellm/proxy/auth/auth_metrics.py +++ /dev/null @@ -1,60 +0,0 @@ -""" -Prometheus metric helpers for the auth layer. - -All metrics are thin wrappers around the shared ``PrometheusLogger`` instance so -that every counter follows the same registration path (``_counter_factory``, -label-filter config) as the rest of LiteLLM's metrics. - -Usage:: - - from litellm.proxy.auth.auth_metrics import AuthMetrics - - AuthMetrics.inc_combined_view_query(hashed_token="sk-xxx") -""" - -from litellm._logging import verbose_proxy_logger -from litellm.types.integrations.prometheus import UserAPIKeyLabelValues - - -class AuthMetrics: - """Static helpers for incrementing auth-layer Prometheus counters.""" - - @staticmethod - def _get_prom(): - """Return the active PrometheusLogger, or None if Prometheus is not configured.""" - try: - from litellm.router_utils.cooldown_callbacks import ( - _get_prometheus_logger_from_callbacks, - ) - - return _get_prometheus_logger_from_callbacks() - except Exception: - return None - - @staticmethod - def inc_combined_view_query(hashed_token: str) -> None: - """ - Increment ``litellm_auth_combined_view_queries_total``. - - Called once per virtual-key DB lookup (combined_view query). Each - increment represents a cache miss that hit the database — use this to - validate that ``enable_redis_auth_cache`` is reducing DB load. - """ - try: - prom = AuthMetrics._get_prom() - if prom is not None: - # Counter labelnames include ``hashed_api_key`` plus any - # ``custom_prometheus_metadata_labels`` / ``custom_prometheus_tags`` - # (see ``PrometheusMetricLabels.get_labels``). Use the same - # ``_inc_labeled_counter`` + ``prometheus_label_factory`` path as - # other metrics so label cardinality always matches registration. - prom._inc_labeled_counter( - prom.litellm_auth_combined_view_queries_total, - "litellm_auth_combined_view_queries_total", - UserAPIKeyLabelValues(hashed_api_key=hashed_token), - ) - except Exception as e: - verbose_proxy_logger.debug( - "AuthMetrics.inc_combined_view_query: failed to increment counter: %s", - e, - ) diff --git a/litellm/proxy/common_utils/cache_coordinator.py b/litellm/proxy/common_utils/cache_coordinator.py index fc46ab6d78e..abb0402d3b9 100644 --- a/litellm/proxy/common_utils/cache_coordinator.py +++ b/litellm/proxy/common_utils/cache_coordinator.py @@ -153,7 +153,7 @@ class EventDrivenCacheCoordinator: elapsed_ms, value, ) - + await cache.async_set_cache(key=cache_key, value=value) if self._log_prefix: verbose_proxy_logger.debug("%s Result cached", self._log_prefix) diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index a4573535f72..43a287f29bc 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -228,8 +228,6 @@ DEFINED_PROMETHEUS_METRICS = Literal[ "litellm_guardrail_latency_seconds", "litellm_guardrail_errors_total", "litellm_guardrail_requests_total", - # Auth DB diagnostic metrics - "litellm_auth_combined_view_queries_total", # Cache metrics "litellm_cache_hits_metric", "litellm_cache_misses_metric", @@ -309,11 +307,6 @@ class PrometheusMetricLabels: litellm_guardrail_errors_total: List[str] = [] litellm_guardrail_requests_total: List[str] = [] - # Auth DB diagnostic - label by key so you can see which virtual key causes DB hits - litellm_auth_combined_view_queries_total = [ - UserAPIKeyLabelNames.API_KEY_HASH.value, - ] - litellm_proxy_total_requests_metric = [ UserAPIKeyLabelNames.END_USER.value, UserAPIKeyLabelNames.API_KEY_HASH.value, From 3bff192c2f55a0df04d48f0c2bc7512f38eb959d Mon Sep 17 00:00:00 2001 From: harish-berri Date: Thu, 30 Apr 2026 22:35:23 +0000 Subject: [PATCH 55/75] fix linting --- .../proxy/common_utils/user_api_key_cache.py | 36 ++++++++++--------- .../key_management_endpoints.py | 20 +++++------ litellm/proxy/proxy_server.py | 8 ++--- 3 files changed, 34 insertions(+), 30 deletions(-) diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index 7c42aa545d0..c6b830371ba 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -39,20 +39,22 @@ class UserApiKeyCache(DualCache): @overload def get_cache( self, - key, - parent_otel_span=None, + key: Any, + parent_otel_span: Any = None, local_only: bool = False, - ) -> Any: ... + *, + model_type: Type[T], + **kwargs: Any, + ) -> Optional[T]: ... @overload def get_cache( self, - key, - parent_otel_span=None, + key: Any, + parent_otel_span: Any = None, local_only: bool = False, - *, - model_type: Type[T], - ) -> Optional[T]: ... + **kwargs: Any, + ) -> Any: ... def set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) @@ -86,20 +88,22 @@ class UserApiKeyCache(DualCache): @overload async def async_get_cache( self, - key, - parent_otel_span=None, + key: Any, + parent_otel_span: Any = None, local_only: bool = False, - ) -> Any: ... + *, + model_type: Type[T], + **kwargs: Any, + ) -> Optional[T]: ... @overload async def async_get_cache( self, - key, - parent_otel_span=None, + key: Any, + parent_otel_span: Any = None, local_only: bool = False, - *, - model_type: Type[T], - ) -> Optional[T]: ... + **kwargs: Any, + ) -> Any: ... def get_cache( # type: ignore[override] self, diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 8129fb0de5e..ae38b389b46 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -27,7 +27,7 @@ from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, s import litellm from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid -from litellm.caching import DualCache +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.constants import ( LENGTH_OF_LITELLM_GENERATED_KEY, LITELLM_PROXY_ADMIN_NAME, @@ -1058,7 +1058,7 @@ async def _check_project_key_limits( project_id: str, data: Union[GenerateKeyRequest, UpdateKeyRequest], prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, ) -> None: """ Validate that key's models and budget respect its project's limits. @@ -1833,7 +1833,7 @@ async def _process_single_key_update( user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str], prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: Any, llm_router: Optional[Router], user_custom_key_update: Optional[Callable] = None, @@ -3297,7 +3297,7 @@ async def _team_key_deletion_check( user_api_key_dict: UserAPIKeyAuth, key_info: LiteLLM_VerificationToken, prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, ): is_team_key = _is_team_key(data=key_info) @@ -3340,7 +3340,7 @@ async def _team_key_deletion_check( async def can_modify_verification_token( key_info: LiteLLM_VerificationToken, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, ) -> bool: @@ -3414,7 +3414,7 @@ async def can_modify_verification_token( async def delete_verification_tokens( tokens: List, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str] = None, ) -> Tuple[Optional[Dict], List[LiteLLM_VerificationToken]]: @@ -3604,7 +3604,7 @@ async def _persist_deleted_verification_tokens( async def delete_key_aliases( key_aliases: List[str], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str] = None, @@ -3850,7 +3850,7 @@ async def _execute_virtual_key_regeneration( data: Optional[RegenerateKeyRequest], user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ) -> GenerateKeyResponse: """Generate new token, update DB, invalidate cache, and return response.""" @@ -4140,7 +4140,7 @@ async def _check_proxy_or_team_admin_for_key( key_in_db: LiteLLM_VerificationToken, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, ) -> None: if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: return @@ -5161,7 +5161,7 @@ async def _check_key_admin_access( user_api_key_dict: UserAPIKeyAuth, hashed_token: str, prisma_client: Any, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, route: str, ) -> None: """ diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 342e1c4ab37..06fdef295d8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1613,7 +1613,7 @@ prisma_client: Optional[PrismaClient] = None shared_aiohttp_session: Optional["ClientSession"] = ( None # Global shared session for connection reuse ) -user_api_key_cache = UserApiKeyCache( +user_api_key_cache: UserApiKeyCache = UserApiKeyCache( default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value ) spend_counter_cache = DualCache( @@ -6318,7 +6318,7 @@ class ProxyStartupEvent: cls, general_settings: dict, prisma_client: Optional[PrismaClient], - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, ): """Initialize JWT auth on startup""" if general_settings.get("litellm_jwtauth", None) is not None: @@ -6367,7 +6367,7 @@ class ProxyStartupEvent: async def _warm_global_spend_cache( cls, litellm_proxy_admin_name: str, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, prisma_client: PrismaClient, ) -> None: """Warm global spend cache once at startup to reduce impact of first wave of requests.""" @@ -7007,7 +7007,7 @@ class ProxyStartupEvent: cls, database_url: Optional[str], proxy_logging_obj: ProxyLogging, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, ) -> Optional[PrismaClient]: """ - Sets up prisma client From 6fb1de9763a27be2a1ca41d96177d79fbd870329 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Thu, 30 Apr 2026 22:40:52 +0000 Subject: [PATCH 56/75] refactor: replace DualCache with UserApiKeyCache in multiple modules This commit updates the codebase to replace instances of DualCache with UserApiKeyCache in various files, including utils, expired_ui_session_key_cleanup_manager, and team_member_permission_checks. Additionally, it enhances the UserApiKeyCache class with new methods for cache management, improving type safety and consistency across the application. --- .../expired_ui_session_key_cleanup_manager.py | 4 +- .../proxy/common_utils/user_api_key_cache.py | 104 +++++++++--------- .../team_member_permission_checks.py | 4 +- litellm/proxy/utils.py | 5 +- 4 files changed, 59 insertions(+), 58 deletions(-) diff --git a/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py b/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py index c25d8533128..67a24567461 100644 --- a/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py +++ b/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py @@ -8,7 +8,7 @@ from datetime import datetime, timezone from typing import Any, Dict, List, Optional from litellm._logging import verbose_proxy_logger -from litellm.caching import DualCache +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.constants import ( EXPIRED_UI_SESSION_KEY_CLEANUP_JOB_NAME, LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE, @@ -31,7 +31,7 @@ class ExpiredUISessionKeyCleanupManager: def __init__( self, prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, pod_lock_manager=None, ): self.prisma_client = prisma_client diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index c6b830371ba..914be364579 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -31,11 +31,11 @@ class UserApiKeyCache(DualCache): ``async_set_cache_pipeline`` applies the same untyped Codec pass as omitting ``model_type`` on ``async_set_cache`` (so ``BaseModel`` rows are dumped before Redis). + + ``get_cache`` / ``async_get_cache`` overloads and implementations must be contiguous + (no other methods in between) so mypy resolves ``@overload`` + implementation correctly. """ - # Overloads: `model_type` must be a real parameter (not only via **kwargs) so - # the untyped branch cannot match calls that pass `model_type=...`. - @overload def get_cache( self, @@ -56,55 +56,6 @@ class UserApiKeyCache(DualCache): **kwargs: Any, ) -> Any: ... - def set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] - model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) - payload = CacheCodec.serialize(value, model_type=model_type) - return super().set_cache( - key=key, value=payload, local_only=local_only, **kwargs - ) - - async def async_set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] - model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) - payload = CacheCodec.serialize(value, model_type=model_type) - return await super().async_set_cache( - key=key, value=payload, local_only=local_only, **kwargs - ) - - async def async_set_cache_pipeline( # type: ignore[override] - self, cache_list: list, local_only: bool = False, **kwargs - ) -> None: - """ - Batch writes with the same Codec boundary as ``async_set_cache`` without - ``model_type``: ``BaseModel`` values become JSON-safe dicts; dicts/scalars unchanged. - """ - normalized = [ - (key, CacheCodec.serialize(value, model_type=None)) - for key, value in cache_list - ] - return await super().async_set_cache_pipeline( - cache_list=normalized, local_only=local_only, **kwargs - ) - - @overload - async def async_get_cache( - self, - key: Any, - parent_otel_span: Any = None, - local_only: bool = False, - *, - model_type: Type[T], - **kwargs: Any, - ) -> Optional[T]: ... - - @overload - async def async_get_cache( - self, - key: Any, - parent_otel_span: Any = None, - local_only: bool = False, - **kwargs: Any, - ) -> Any: ... - def get_cache( # type: ignore[override] self, key, @@ -133,6 +84,26 @@ class UserApiKeyCache(DualCache): return None return decoded + @overload + async def async_get_cache( + self, + key: Any, + parent_otel_span: Any = None, + local_only: bool = False, + *, + model_type: Type[T], + **kwargs: Any, + ) -> Optional[T]: ... + + @overload + async def async_get_cache( + self, + key: Any, + parent_otel_span: Any = None, + local_only: bool = False, + **kwargs: Any, + ) -> Any: ... + async def async_get_cache( # type: ignore[override] self, key, @@ -160,3 +131,32 @@ class UserApiKeyCache(DualCache): ) return None return decoded + + def set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] + model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + payload = CacheCodec.serialize(value, model_type=model_type) + return super().set_cache( + key=key, value=payload, local_only=local_only, **kwargs + ) + + async def async_set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override] + model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None)) + payload = CacheCodec.serialize(value, model_type=model_type) + return await super().async_set_cache( + key=key, value=payload, local_only=local_only, **kwargs + ) + + async def async_set_cache_pipeline( # type: ignore[override] + self, cache_list: list, local_only: bool = False, **kwargs + ) -> None: + """ + Batch writes with the same Codec boundary as ``async_set_cache`` without + ``model_type``: ``BaseModel`` values become JSON-safe dicts; dicts/scalars unchanged. + """ + normalized = [ + (key, CacheCodec.serialize(value, model_type=None)) + for key, value in cache_list + ] + return await super().async_set_cache_pipeline( + cache_list=normalized, local_only=local_only, **kwargs + ) diff --git a/litellm/proxy/management_helpers/team_member_permission_checks.py b/litellm/proxy/management_helpers/team_member_permission_checks.py index e035168ca00..50339210a6e 100644 --- a/litellm/proxy/management_helpers/team_member_permission_checks.py +++ b/litellm/proxy/management_helpers/team_member_permission_checks.py @@ -1,6 +1,5 @@ from typing import List, Optional -from litellm.caching import DualCache from litellm.proxy._types import ( KeyManagementRoutes, LiteLLM_TeamTableCachedObj, @@ -12,6 +11,7 @@ from litellm.proxy._types import ( ProxyException, UserAPIKeyAuth, ) +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.utils import PrismaClient @@ -65,7 +65,7 @@ class TeamMemberPermissionChecks: user_api_key_dict: UserAPIKeyAuth, route: KeyManagementRoutes, prisma_client: PrismaClient, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, existing_key_row: LiteLLM_VerificationToken, ): """ diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index d2dfa177515..8c5fce84099 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -101,6 +101,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.route_checks import RouteChecks +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.db.create_views import ( create_missing_views, should_create_missing_views, @@ -340,7 +341,7 @@ class ProxyLogging: def __init__( self, - user_api_key_cache: DualCache, + user_api_key_cache: UserApiKeyCache, premium_user: bool = False, ): ## INITIALIZE LITELLM CALLBACKS ## @@ -5715,7 +5716,7 @@ async def get_available_models_for_user( include_model_access_groups: bool = False, only_model_access_groups: bool = False, return_wildcard_routes: bool = False, - user_api_key_cache: Optional["DualCache"] = None, + user_api_key_cache: Optional["UserApiKeyCache"] = None, ) -> List[str]: """ Get the list of models available to a user based on their API key and team permissions. From 8671ec636b9d855bcad630c29a087aa8e6769307 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Fri, 1 May 2026 00:29:13 +0000 Subject: [PATCH 57/75] fix import error --- litellm/proxy/management_endpoints/ui_sso.py | 1 + .../common_utils/test_user_api_key_cache.py | 106 ++++++++++++++++++ 2 files changed, 107 insertions(+) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 7b4daf14a92..c1458dc37f4 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -39,6 +39,7 @@ from fastapi import APIRouter, Depends, Header, HTTPException, Request, status from fastapi.responses import RedirectResponse import litellm +from litellm.caching.dual_cache import DualCache from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.constants import ( diff --git a/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py b/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py index cff331d9b19..8667348d223 100644 --- a/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py +++ b/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py @@ -3,9 +3,23 @@ from typing import Any import pytest +from litellm.caching.in_memory_cache import InMemoryCache from litellm.caching.redis_cache import RedisCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.proxy_server import UserAPIKeyCacheTTLEnum + + +class CapturingInMemoryCache(InMemoryCache): + """Records ``ttl`` passed into ``set_cache`` (what DualCache injects).""" + + def __init__(self) -> None: + super().__init__() + self.last_ttl: Any = None + + def set_cache(self, key, value, **kwargs): # type: ignore[override] + self.last_ttl = kwargs.get("ttl") + super().set_cache(key, value, **kwargs) class FakeRedisCache(RedisCache): @@ -18,14 +32,18 @@ class FakeRedisCache(RedisCache): This fake: - raises TypeError if the value is not a dict - raises TypeError if the dict is not JSON-serializable + + Records the ``ttl`` kwarg DualCache forwards on each Redis write for tests. """ def __init__(self): # noqa: super().__init__ skipped intentionally self._store: dict[str, str] = {} + self.last_ttl: Any = None def set_cache(self, key: str, value: Any, **kwargs): # type: ignore[override] if not isinstance(value, dict): raise TypeError("FakeRedisCache only accepts dict payloads") + self.last_ttl = kwargs.get("ttl") self._store[key] = json.dumps(value) return True @@ -38,6 +56,7 @@ class FakeRedisCache(RedisCache): async def async_set_cache(self, key: str, value: Any, **kwargs): # type: ignore[override] if not isinstance(value, dict): raise TypeError("FakeRedisCache only accepts dict payloads") + self.last_ttl = kwargs.get("ttl") self._store[key] = json.dumps(value) return True @@ -60,6 +79,93 @@ def _make_key_obj(token: str = "tok") -> UserAPIKeyAuth: class TestUserApiKeyCache: + @pytest.mark.asyncio + async def test_async_set_in_memory_gets_enum_default_when_user_api_key_cache_ttl_omitted( + self, + ): + """ + If ``general_settings.user_api_key_cache_ttl`` is absent, the proxy never + calls ``update_cache_ttl``; ``user_api_key_cache`` keeps + ``default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl``. + DualCache must forward that as the in-memory ``ttl`` kwarg on each set. + """ + mem = CapturingInMemoryCache() + cache = UserApiKeyCache( + in_memory_cache=mem, + redis_cache=FakeRedisCache(), + default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value, + ) + await cache.async_set_cache( + "k", + _make_key_obj("t"), + model_type=UserAPIKeyAuth, + ) + expected = UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value + assert mem.last_ttl == expected + + def test_sync_set_in_memory_gets_enum_default_when_user_api_key_cache_ttl_omitted( + self, + ): + mem = CapturingInMemoryCache() + cache = UserApiKeyCache( + in_memory_cache=mem, + redis_cache=FakeRedisCache(), + default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value, + ) + cache.set_cache("sk", _make_key_obj("s"), model_type=UserAPIKeyAuth) + assert mem.last_ttl == UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value + + @pytest.mark.asyncio + async def test_async_set_forwards_default_in_memory_ttl_to_redis_layer(self): + """ + DualCache injects missing ``ttl`` from ``default_in_memory_ttl`` into kwargs + before calling ``redis_cache.async_set_cache`` — Redis should receive the same + TTL as memory (matches proxy defaults: enum 60s). + """ + fake = FakeRedisCache() + cache = UserApiKeyCache( + redis_cache=fake, + default_in_memory_ttl=60, + ) + + await cache.async_set_cache( + key="ttl-key", + value=_make_key_obj("ttl-tok"), + model_type=UserAPIKeyAuth, + ) + + assert fake.last_ttl == 60 + + @pytest.mark.asyncio + async def test_async_set_explicit_ttl_override_reaches_redis(self): + fake = FakeRedisCache() + cache = UserApiKeyCache( + redis_cache=fake, + default_in_memory_ttl=60, + ) + + await cache.async_set_cache( + key="k", + value=_make_key_obj("x"), + model_type=UserAPIKeyAuth, + ttl=900, + ) + + assert fake.last_ttl == 900 + + def test_sync_set_forwards_default_in_memory_ttl_to_redis_layer(self): + fake = FakeRedisCache() + cache = UserApiKeyCache( + redis_cache=fake, + default_in_memory_ttl=45, + ) + cache.set_cache( + "sk", + _make_key_obj("sync"), + model_type=UserAPIKeyAuth, + ) + assert fake.last_ttl == 45 + @pytest.mark.asyncio async def test_async_set_typed_stores_serialized_payload_in_memory_and_redis(self): cache = UserApiKeyCache(redis_cache=FakeRedisCache()) From 896e1024643e20f609e618df4141b3368aca3a0d Mon Sep 17 00:00:00 2001 From: harish-berri Date: Fri, 1 May 2026 00:32:10 +0000 Subject: [PATCH 58/75] fix linting --- litellm/proxy/management_endpoints/ui_sso.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index c1458dc37f4..9dfc67370fe 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -76,8 +76,10 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_user_object from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache -from litellm.proxy.auth.auth_utils import _has_user_setup_sso -from litellm.proxy.auth.auth_utils import _get_request_ip_address, _has_user_setup_sso +from litellm.proxy.auth.auth_utils import ( + _get_request_ip_address, + _has_user_setup_sso, +) from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.admin_ui_utils import ( From e72eac917609ccb34adced41f4c5dae8a627eb97 Mon Sep 17 00:00:00 2001 From: shivam Date: Thu, 30 Apr 2026 18:03:43 -0700 Subject: [PATCH 59/75] Fix add_model_file_id_mappings when router returns single deployment dict When model_info.id equals model_name (common for batch models), the router resolves via has_model_id and returns one deployment dict instead of a list. The dict branch incorrectly iterated deployment keys (model_name, litellm_params, model_info), producing non-string values that broke LiteLLM_ManagedFileTable validation on managed file upload. Normalize list vs dict by wrapping single deployments and extracting model_info.id for each response pair. Add regression tests including the batch model id == model_name case. Made-with: Cursor --- litellm/router_utils/common_utils.py | 27 +++-- .../test_router_utils_common_utils.py | 110 ++++++++++++++++++ 2 files changed, 127 insertions(+), 10 deletions(-) diff --git a/litellm/router_utils/common_utils.py b/litellm/router_utils/common_utils.py index bef42e23848..f6da26ccd7f 100644 --- a/litellm/router_utils/common_utils.py +++ b/litellm/router_utils/common_utils.py @@ -23,21 +23,28 @@ def add_model_file_id_mappings( healthy_deployments: Union[List[Dict], Dict], responses: List["OpenAIFileObject"] ) -> dict: """ - Create a mapping of model name to file id + Create a mapping of model id to file id { "model_id": "file_id", "model_id": "file_id", } + + `healthy_deployments` may be either a list of deployment dicts (multiple + matched deployments) or a single deployment dict (when the router resolved + a specific deployment, e.g. because the requested model matched a + `model_info.id`). Both shapes must be handled by extracting + `model_info.id` from each deployment. """ - model_file_id_mapping = {} - if isinstance(healthy_deployments, list): - for deployment, response in zip(healthy_deployments, responses): - model_file_id_mapping[deployment.get("model_info", {}).get("id")] = ( - response.id - ) - elif isinstance(healthy_deployments, dict): - for model_id, file_id in healthy_deployments.items(): - model_file_id_mapping[model_id] = file_id + model_file_id_mapping: Dict[str, str] = {} + deployments_list: List[Dict] = ( + healthy_deployments + if isinstance(healthy_deployments, list) + else [healthy_deployments] + ) + for deployment, response in zip(deployments_list, responses): + model_id = deployment.get("model_info", {}).get("id") + if model_id is not None: + model_file_id_mapping[model_id] = response.id return model_file_id_mapping diff --git a/tests/test_litellm/router_utils/test_router_utils_common_utils.py b/tests/test_litellm/router_utils/test_router_utils_common_utils.py index 02241d4bc92..465c6669ceb 100644 --- a/tests/test_litellm/router_utils/test_router_utils_common_utils.py +++ b/tests/test_litellm/router_utils/test_router_utils_common_utils.py @@ -6,6 +6,7 @@ import pytest from litellm import Router from litellm.router_utils.common_utils import ( _deployment_supports_web_search, + add_model_file_id_mappings, filter_team_based_models, filter_web_search_deployments, ) @@ -362,3 +363,112 @@ def test_invalidate_model_group_info_cache(): # Invalidate and verify cache is cleared router._invalidate_model_group_info_cache() assert router._cached_get_model_group_info.cache_info().currsize == 0 + + +class TestAddModelFileIdMappings: + """Test cases for add_model_file_id_mappings. + + The router may pass either a list of deployment dicts (multiple matched + deployments) or a single deployment dict (when a specific deployment was + resolved, e.g. because the requested model matched a `model_info.id`). + Both shapes must produce a `{model_id: file_id}` mapping by extracting + `model_info.id` from each deployment. + """ + + @staticmethod + def _make_response(file_id: str): + response = Mock() + response.id = file_id + return response + + def test_should_map_each_deployment_id_when_given_list(self): + deployments = [ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + "model_info": {"id": "deployment-1"}, + }, + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + "model_info": {"id": "deployment-2"}, + }, + ] + responses = [self._make_response("file-1"), self._make_response("file-2")] + + result = add_model_file_id_mappings(deployments, responses) + + assert result == {"deployment-1": "file-1", "deployment-2": "file-2"} + + def test_should_extract_model_info_id_when_given_single_deployment_dict(self): + """Regression test: when `_common_checks_available_deployment` resolves + a specific deployment (returned as a dict, not a list), the function + must still extract `model_info.id` rather than iterate over the + deployment's own keys (`model_name`, `litellm_params`, `model_info`). + """ + deployment = { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "sk-test"}, + "model_info": {"id": "deployment-1", "mode": "chat"}, + } + responses = [self._make_response("file-1")] + + result = add_model_file_id_mappings(deployment, responses) + + assert result == {"deployment-1": "file-1"} + assert all(isinstance(v, str) for v in result.values()) + + def test_should_handle_batch_model_when_id_matches_model_name(self): + """Regression test for the batch-model case: when `model_info.id` is + intentionally set equal to `model_name`, the router resolves a single + deployment via `has_model_id` and returns it as a dict. The mapping + must contain only `{id: file_id}` with string values so the resulting + `LiteLLM_ManagedFileTable` Pydantic validation passes. + """ + deployment = { + "model_name": "openai/openai/gpt-5.5-batch", + "litellm_params": { + "model": "openai/gpt-5.5", + "api_key": "sk-test", + "tpm": 40000000, + "rpm": 15000, + }, + "model_info": { + "id": "openai/openai/gpt-5.5-batch", + "mode": "batch", + "base_model": "gpt-5.5", + "access_groups": ["default-models"], + }, + } + responses = [self._make_response("file-batch-1")] + + result = add_model_file_id_mappings(deployment, responses) + + # Bug case would have produced keys ["model_name", "litellm_params", + # "model_info"] with non-string values. + assert result == {"openai/openai/gpt-5.5-batch": "file-batch-1"} + assert "litellm_params" not in result + assert "model_info" not in result + + def test_should_skip_deployment_when_model_info_id_missing(self): + deployments = [ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + "model_info": {}, + }, + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + "model_info": {"id": "deployment-2"}, + }, + ] + responses = [self._make_response("file-1"), self._make_response("file-2")] + + result = add_model_file_id_mappings(deployments, responses) + + assert result == {"deployment-2": "file-2"} + + def test_should_return_empty_mapping_when_given_empty_list(self): + result = add_model_file_id_mappings([], []) + assert result == {} From dd549d9c50198895a734dec2c97b39882fd62658 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 30 Apr 2026 21:38:05 -0700 Subject: [PATCH 60/75] =?UTF-8?q?bump:=20version=200.4.69=20=E2=86=92=200.?= =?UTF-8?q?4.70?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- litellm-proxy-extras/pyproject.toml | 4 ++-- pyproject.toml | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index 41c78296fbb..b8710da3438 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-proxy-extras" -version = "0.4.69" +version = "0.4.70" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.4.69" +version = "0.4.70" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-proxy-extras==", diff --git a/pyproject.toml b/pyproject.toml index 657632d69e8..0ef0a993dd5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -52,7 +52,7 @@ proxy = [ "azure-identity==1.25.2", "azure-storage-blob==12.28.0", "mcp==1.26.0", - "litellm-proxy-extras==0.4.69", + "litellm-proxy-extras==0.4.70", "litellm-enterprise==0.1.39", "RestrictedPython==8.1", "rich==13.9.4", From 6da13efcec80af72de57aaa7e37fbf63c762d273 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 30 Apr 2026 21:40:09 -0700 Subject: [PATCH 61/75] uv lock --- uv.lock | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/uv.lock b/uv.lock index f837e2b5eff..04a9e73fed8 100644 --- a/uv.lock +++ b/uv.lock @@ -3418,7 +3418,7 @@ source = { editable = "enterprise" } [[package]] name = "litellm-proxy-extras" -version = "0.4.69" +version = "0.4.70" source = { editable = "litellm-proxy-extras" } [[package]] From 8947a74e13e27257cee26a30293f9ce0b5e5d399 Mon Sep 17 00:00:00 2001 From: Noah <46087295+ohnoah@users.noreply.github.com> Date: Tue, 7 Apr 2026 05:28:21 +0100 Subject: [PATCH 62/75] fix(cache): persist and replay streamed Responses API requests (#24580) * fix(cache): persist and replay streamed Responses API requests * Add focused coverage for streamed responses cache * Cover streamed responses cache helper branches * Exercise streamed responses cache edge branches --- litellm/caching/caching.py | 3 +- litellm/caching/caching_handler.py | 106 ++- litellm/responses/streaming_iterator.py | 738 ++++++++++++--- litellm/types/llms/openai.py | 5 +- .../test_responses_hooks.py | 888 +++++++++++++++++- tests/local_testing/test_caching_handler.py | 399 +++++++- .../test_responses_stream_cache_keys.py | 141 +++ 7 files changed, 2105 insertions(+), 175 deletions(-) create mode 100644 tests/local_testing/test_responses_stream_cache_keys.py diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index ce1bc26c5e0..11733ce4cee 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -432,9 +432,10 @@ class Cache: str: The final hashed cache key with the redis namespace. """ dynamic_cache_control: DynamicCacheControl = kwargs.get("cache", {}) + metadata = kwargs.get("metadata") or {} namespace = ( dynamic_cache_control.get("namespace") - or kwargs.get("metadata", {}).get("redis_namespace") + or metadata.get("redis_namespace") or self.namespace ) if namespace: diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 7d514e648fe..02207a207ed 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -85,6 +85,19 @@ class CachingHandlerResponse(BaseModel): in_memory_cache_obj = InMemoryCache() +_RESPONSES_STREAMING_CALLBACK_CALL_TYPES = { + CallTypes.aresponses.value, + CallTypes.responses.value, +} + + +def _should_defer_streaming_cache_hit_callbacks( + *, call_type: str, kwargs: Dict[str, Any] +) -> bool: + return ( + kwargs.get("stream", False) is True + and call_type in _RESPONSES_STREAMING_CALLBACK_CALL_TYPES + ) class LLMCachingHandler: @@ -99,6 +112,7 @@ class LLMCachingHandler: self.async_streaming_chunks: List[ModelResponse] = [] self.sync_streaming_chunks: List[ModelResponse] = [] self.request_kwargs = request_kwargs + self.preset_cache_key: Optional[str] = None self.original_function = original_function self.start_time = start_time if litellm.cache is not None and isinstance(litellm.cache.cache, RedisCache): @@ -206,7 +220,10 @@ class LLMCachingHandler: custom_llm_provider=kwargs.get("custom_llm_provider", None), args=args, ) - if kwargs.get("stream", False) is False: + if not _should_defer_streaming_cache_hit_callbacks( + call_type=call_type, + kwargs=kwargs, + ): # LOG SUCCESS self._async_log_cache_hit_on_callbacks( logging_obj=logging_obj, @@ -215,11 +232,12 @@ class LLMCachingHandler: end_time=end_time, cache_hit=cache_hit, ) - cache_key = litellm.cache.get_cache_key(**kwargs) - if ( - isinstance(cached_result, BaseModel) - or isinstance(cached_result, CustomStreamWrapper) - ) and hasattr(cached_result, "_hidden_params"): + cache_key = ( + self.preset_cache_key + or self.request_kwargs.get("cache_key") + or litellm.cache.get_cache_key(**self.request_kwargs) + ) + if hasattr(cached_result, "_hidden_params"): cached_result._hidden_params["cache_key"] = cache_key # type: ignore return CachingHandlerResponse(cached_result=cached_result) elif ( @@ -265,8 +283,6 @@ class LLMCachingHandler: kwargs: Dict[str, Any], args: Optional[Tuple[Any, ...]] = None, ) -> CachingHandlerResponse: - from litellm.utils import CustomStreamWrapper - cached_result: Optional[Any] = None # Check if caching should be performed BEFORE doing expensive kwargs copy @@ -282,6 +298,11 @@ class LLMCachingHandler: args, ) ) + if new_kwargs.get("metadata") is None: + new_kwargs.pop("metadata", None) + if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs: + new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs) + self.request_kwargs = new_kwargs print_verbose("Checking Sync Cache") cached_result = litellm.cache.get_cache(**new_kwargs) if cached_result is not None: @@ -322,17 +343,22 @@ class LLMCachingHandler: is_async=False, ) - logging_obj.handle_sync_success_callbacks_for_async_calls( - result=cached_result, - start_time=start_time, - end_time=end_time, - cache_hit=cache_hit, + if not _should_defer_streaming_cache_hit_callbacks( + call_type=call_type, + kwargs=kwargs, + ): + logging_obj.handle_sync_success_callbacks_for_async_calls( + result=cached_result, + start_time=start_time, + end_time=end_time, + cache_hit=cache_hit, + ) + cache_key = ( + self.preset_cache_key + or self.request_kwargs.get("cache_key") + or litellm.cache.get_cache_key(**self.request_kwargs) ) - cache_key = litellm.cache.get_cache_key(**kwargs) - if ( - isinstance(cached_result, BaseModel) - or isinstance(cached_result, CustomStreamWrapper) - ) and hasattr(cached_result, "_hidden_params"): + if hasattr(cached_result, "_hidden_params"): cached_result._hidden_params["cache_key"] = cache_key # type: ignore return CachingHandlerResponse(cached_result=cached_result) return CachingHandlerResponse(cached_result=cached_result) @@ -686,6 +712,11 @@ class LLMCachingHandler: args, ) ) + if new_kwargs.get("metadata") is None: + new_kwargs.pop("metadata", None) + if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs: + new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs) + self.request_kwargs = new_kwargs cached_result: Optional[Any] = None if call_type == CallTypes.aembedding.value: if isinstance(new_kwargs["input"], str): @@ -710,14 +741,26 @@ class LLMCachingHandler: if all(result is None for result in cached_result): cached_result = None else: + request_kwargs = new_kwargs.copy() + request_cache_key = request_kwargs.pop("cache_key", None) if litellm.cache._supports_async() is True: ## check if dual cache is supported ## + self.preset_cache_key = ( + request_cache_key or litellm.cache.get_cache_key(**request_kwargs) + ) cached_result = await litellm.cache.async_get_cache( - dynamic_cache_object=self.dual_cache, **new_kwargs + dynamic_cache_object=self.dual_cache, + cache_key=self.preset_cache_key, + **request_kwargs, ) else: # fallback for caches that don't support async + self.preset_cache_key = ( + request_cache_key or litellm.cache.get_cache_key(**request_kwargs) + ) cached_result = litellm.cache.get_cache( - dynamic_cache_object=self.dual_cache, **new_kwargs + dynamic_cache_object=self.dual_cache, + cache_key=self.preset_cache_key, + **request_kwargs, ) return cached_result @@ -825,8 +868,27 @@ class LLMCachingHandler: elif (call_type == "aresponses" or call_type == "responses") and isinstance( cached_result, dict ): - # Convert cached dict back to ResponsesAPIResponse object - cached_result = ResponsesAPIResponse(**cached_result) + from litellm.responses.streaming_iterator import ( + CachedResponsesAPIStreamingIterator, + ) + + response_obj = ResponsesAPIResponse(**cached_result) + if ( + hasattr(response_obj, "_hidden_params") + and response_obj._hidden_params is not None + and isinstance(response_obj._hidden_params, dict) + ): + response_obj._hidden_params["cache_hit"] = True + + if kwargs.get("stream", False) is True: + cached_result = CachedResponsesAPIStreamingIterator( + response=response_obj, + logging_obj=logging_obj, + request_data=kwargs, + call_type=call_type, + ) + else: + cached_result = response_obj if ( hasattr(cached_result, "_hidden_params") diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 145ec3a641a..1902a347215 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -1,9 +1,12 @@ +from __future__ import annotations + import asyncio import json import time import traceback from datetime import datetime -from typing import Any, Dict, List, Optional +from functools import lru_cache +from typing import Any, Dict, List, Literal, Optional import httpx @@ -22,19 +25,25 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import ( from litellm.litellm_core_utils.thread_pool_executor import executor from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.responses.utils import ResponsesAPIRequestUtils -from litellm.types.llms.openai import ( - OutputTextDeltaEvent, - ResponseAPIUsage, - ResponseCompletedEvent, - ResponsesAPIRequestParams, - ResponsesAPIResponse, - ResponsesAPIStreamEvents, - ResponsesAPIStreamingResponse, -) from litellm.types.utils import CallTypes from litellm.utils import CustomStreamWrapper, async_post_call_success_deployment_hook +@lru_cache(maxsize=1) +def _get_openai_response_types(): + from litellm.types.llms import openai as openai_types + + return openai_types + + +def _log_background_task_failure(task: "asyncio.Task[Any]", *, task_name: str) -> None: + if task.cancelled(): + return + exception = task.exception() + if exception is not None: + verbose_logger.error("%s failed: %s", task_name, exception) + + class BaseResponsesAPIStreamingIterator: """ Base class for streaming iterators that process responses from the Responses API. @@ -46,7 +55,7 @@ class BaseResponsesAPIStreamingIterator: self, response: httpx.Response, model: str, - responses_api_provider_config: BaseResponsesAPIConfig, + responses_api_provider_config: Optional[BaseResponsesAPIConfig], logging_obj: LiteLLMLoggingObj, litellm_metadata: Optional[Dict[str, Any]] = None, custom_llm_provider: Optional[str] = None, @@ -58,9 +67,13 @@ class BaseResponsesAPIStreamingIterator: self.logging_obj = logging_obj self.finished = False self.responses_api_provider_config = responses_api_provider_config - self.completed_response: Optional[ResponsesAPIStreamingResponse] = None + self.completed_response: Optional[Any] = None self.start_time = getattr(logging_obj, "start_time", datetime.now()) self._failure_handled = False # Track if failure handler has been called + self._completed_response_cached = False + self._completed_response_logged = False + self._completed_response_cache_hit: Optional[bool] = None + self._persist_completed_response_before_logging = True self._stream_created_time: float = time.time() # track request context for hooks @@ -101,7 +114,7 @@ class BaseResponsesAPIStreamingIterator: llm_provider=self.custom_llm_provider or "", ) - def _process_chunk(self, chunk) -> Optional[ResponsesAPIStreamingResponse]: + def _process_chunk(self, chunk) -> Optional[Any]: """Process a single chunk of data from the stream""" if not chunk: return None @@ -122,6 +135,10 @@ class BaseResponsesAPIStreamingIterator: # Format as ResponsesAPIStreamingResponse if isinstance(parsed_chunk, dict): + if self.responses_api_provider_config is None: + raise ValueError( + "responses_api_provider_config is required to process live streaming chunks" + ) openai_responses_api_chunk = ( self.responses_api_provider_config.transform_streaming_response( model=self.model, @@ -195,10 +212,11 @@ class BaseResponsesAPIStreamingIterator: if self.litellm_metadata and self.litellm_metadata.get( "encrypted_content_affinity_enabled" ): + openai_types = _get_openai_response_types() event_type = getattr(openai_responses_api_chunk, "type", None) if event_type in ( - ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, - ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, ): item = getattr(openai_responses_api_chunk, "item", None) if item: @@ -219,10 +237,11 @@ class BaseResponsesAPIStreamingIterator: # Store the completed response (also for incomplete/failed so logging still fires) _chunk_type = getattr(openai_responses_api_chunk, "type", None) + openai_types = _get_openai_response_types() if openai_responses_api_chunk and _chunk_type in ( - ResponsesAPIStreamEvents.RESPONSE_COMPLETED, - ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, - ResponsesAPIStreamEvents.RESPONSE_FAILED, + openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, + openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED, ): self.completed_response = openai_responses_api_chunk # Add cost to usage object if include_cost_in_streaming_usage is True @@ -230,11 +249,11 @@ class BaseResponsesAPIStreamingIterator: litellm.include_cost_in_streaming_usage and self.logging_obj is not None ): - response_obj: Optional[ResponsesAPIResponse] = getattr( + response_obj: Optional[Any] = getattr( openai_responses_api_chunk, "response", None ) if response_obj: - usage_obj: Optional[ResponseAPIUsage] = getattr( + usage_obj: Optional[Any] = getattr( response_obj, "usage", None ) if usage_obj is not None: @@ -247,9 +266,13 @@ class BaseResponsesAPIStreamingIterator: if cost is not None: setattr(usage_obj, "cost", cost) except Exception: + # Best-effort usage cost annotation should not break stream replay. pass - if _chunk_type == ResponsesAPIStreamEvents.RESPONSE_FAILED: + if ( + _chunk_type + == openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED + ): self._handle_logging_failed_response() else: self._handle_logging_completed_response() @@ -266,6 +289,59 @@ class BaseResponsesAPIStreamingIterator: self._handle_failure(e) raise + def _log_completed_response(self, *, is_async: bool) -> None: + if self._completed_response_logged: + return + self._completed_response_logged = True + + if self._persist_completed_response_before_logging: + self._persist_completed_response_to_cache(is_async=is_async) + + # Create a copy for logging to avoid modifying the response object that will be returned to the user + # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens) + # to chat completion format (prompt_tokens/completion_tokens) for internal logging + # Use model_dump + model_validate instead of deepcopy to avoid pickle errors with + # Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192) + logging_response = self.completed_response + if self.completed_response is not None and hasattr( + self.completed_response, "model_dump" + ): + try: + logging_response = type(self.completed_response).model_validate( + self.completed_response.model_dump() + ) + except Exception: + # Fallback to original if serialization fails + pass + + end_time = datetime.now() + if is_async: + asyncio.create_task( + self.logging_obj.async_success_handler( + result=logging_response, + start_time=self.start_time, + end_time=end_time, + cache_hit=self._completed_response_cache_hit, + ) + ) + else: + run_async_function( + async_function=self.logging_obj.async_success_handler, + result=logging_response, + start_time=self.start_time, + end_time=end_time, + cache_hit=self._completed_response_cache_hit, + ) + + executor.submit( + self.logging_obj.success_handler, + result=logging_response, + cache_hit=self._completed_response_cache_hit, + start_time=self.start_time, + end_time=end_time, + ) + self._run_post_success_hooks(end_time=end_time) + def _handle_logging_completed_response(self): """Base implementation - should be overridden by subclasses""" pass @@ -296,6 +372,88 @@ class BaseResponsesAPIStreamingIterator: ) self._handle_failure(exception) + def _get_completed_response_object(self) -> Optional[Any]: + openai_types = _get_openai_response_types() + completed_response = self.completed_response + if isinstance(completed_response, openai_types.ResponsesAPIResponse): + return completed_response + + response_obj = getattr(completed_response, "response", None) + if isinstance(response_obj, openai_types.ResponsesAPIResponse): + return response_obj + + return None + + def _persist_completed_response_to_cache(self, *, is_async: bool) -> None: + if self._completed_response_cached: + return + + completed_response = self.completed_response + openai_types = _get_openai_response_types() + if ( + getattr(completed_response, "type", None) + != openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ): + return + + response_obj = self._get_completed_response_object() + if response_obj is None: + return + + caching_handler = getattr(self.logging_obj, "_llm_caching_handler", None) + if caching_handler is None: + return + + request_kwargs = getattr(caching_handler, "request_kwargs", None) + if ( + not isinstance(request_kwargs, dict) + or request_kwargs.get("stream") is not True + ): + return + request_kwargs = request_kwargs.copy() + preset_cache_key = getattr(caching_handler, "preset_cache_key", None) + request_cache_key = request_kwargs.pop("cache_key", None) + if preset_cache_key is None: + preset_cache_key = request_cache_key + if request_kwargs.get("metadata") is None: + request_kwargs.pop("metadata", None) + request_kwargs.pop("custom_llm_provider", None) + if preset_cache_key is not None: + request_kwargs["cache_key"] = preset_cache_key + + if not caching_handler._should_store_result_in_cache( + original_function=caching_handler.original_function, + kwargs=request_kwargs, + ): + return + + if litellm.cache is None: + return + + cached_response = response_obj.model_dump_json() + if is_async: + cache_write_task = asyncio.create_task( + litellm.cache.async_add_cache( + cached_response, + dynamic_cache_object=getattr(caching_handler, "dual_cache", None), + **request_kwargs, + ) + ) + cache_write_task.add_done_callback( + lambda task: _log_background_task_failure( + task, + task_name="Responses stream cache write", + ) + ) + else: + litellm.cache.add_cache( + cached_response, + dynamic_cache_object=getattr(caching_handler, "dual_cache", None), + **request_kwargs, + ) + + self._completed_response_cached = True + async def _call_post_streaming_deployment_hook(self, chunk): """ Allow callbacks to modify streaming chunks before returning (parity with chat). @@ -480,7 +638,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def __aiter__(self): return self - async def __anext__(self) -> ResponsesAPIStreamingResponse: + async def __anext__(self) -> Any: try: self._check_max_streaming_duration() while True: @@ -520,40 +678,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def _handle_logging_completed_response(self): """Handle logging for completed responses in async context""" - # Create a copy for logging to avoid modifying the response object that will be returned to the user - # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens) - # to chat completion format (prompt_tokens/completion_tokens) for internal logging - # Use model_dump + model_validate instead of deepcopy to avoid pickle errors with - # Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192) - logging_response = self.completed_response - if self.completed_response is not None and hasattr( - self.completed_response, "model_dump" - ): - try: - logging_response = type(self.completed_response).model_validate( - self.completed_response.model_dump() - ) - except Exception: - # Fallback to original if serialization fails - pass - - asyncio.create_task( - self.logging_obj.async_success_handler( - result=logging_response, - start_time=self.start_time, - end_time=datetime.now(), - cache_hit=None, - ) - ) - - executor.submit( - self.logging_obj.success_handler, - result=logging_response, - cache_hit=None, - start_time=self.start_time, - end_time=datetime.now(), - ) - self._run_post_success_hooks(end_time=datetime.now()) + self._log_completed_response(is_async=True) class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): @@ -627,39 +752,7 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def _handle_logging_completed_response(self): """Handle logging for completed responses in sync context""" - # Create a copy for logging to avoid modifying the response object that will be returned to the user - # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens) - # to chat completion format (prompt_tokens/completion_tokens) for internal logging - # Use model_dump + model_validate instead of deepcopy to avoid pickle errors with - # Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192) - logging_response = self.completed_response - if self.completed_response is not None and hasattr( - self.completed_response, "model_dump" - ): - try: - logging_response = type(self.completed_response).model_validate( - self.completed_response.model_dump() - ) - except Exception: - # Fallback to original if serialization fails - pass - - run_async_function( - async_function=self.logging_obj.async_success_handler, - result=logging_response, - start_time=self.start_time, - end_time=datetime.now(), - cache_hit=None, - ) - - executor.submit( - self.logging_obj.success_handler, - result=logging_response, - cache_hit=None, - start_time=self.start_time, - end_time=datetime.now(), - ) - self._run_post_success_hooks(end_time=datetime.now()) + self._log_completed_response(is_async=False) class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): @@ -683,90 +776,441 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): request_data: Optional[Dict[str, Any]] = None, call_type: Optional[str] = None, ): - super().__init__( - response=response, + transformed = responses_api_provider_config.transform_response_api_response( model=model, - responses_api_provider_config=responses_api_provider_config, + raw_response=response, + logging_obj=logging_obj, + ) + super().__init__( + response=httpx.Response(200), + model=model, + responses_api_provider_config=None, logging_obj=logging_obj, litellm_metadata=litellm_metadata, custom_llm_provider=custom_llm_provider, request_data=request_data, call_type=call_type, ) + self._set_events_from_response(transformed=transformed, logging_obj=logging_obj) - # one-time transform - transformed = ( - self.responses_api_provider_config.transform_response_api_response( - model=self.model, - raw_response=response, - logging_obj=logging_obj, - ) + def _set_events_from_response( + self, + transformed: Any, + logging_obj: LiteLLMLoggingObj, + ) -> None: + self._events = _build_synthetic_response_events( + transformed=transformed, + logging_obj=logging_obj, + chunk_size=self.CHUNK_SIZE, ) - full_text = self._collect_text(transformed) - - # build a list of 5‑char delta events - deltas = [ - OutputTextDeltaEvent( - type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, - delta=full_text[i : i + self.CHUNK_SIZE], - item_id=transformed.id, - output_index=0, - content_index=0, - ) - for i in range(0, len(full_text), self.CHUNK_SIZE) - ] - - # Add cost to usage object if include_cost_in_streaming_usage is True - if litellm.include_cost_in_streaming_usage and logging_obj is not None: - usage_obj: Optional[ResponseAPIUsage] = getattr(transformed, "usage", None) - if usage_obj is not None: - try: - cost: Optional[float] = logging_obj._response_cost_calculator( - result=transformed - ) - if cost is not None: - setattr(usage_obj, "cost", cost) - except Exception: - # If cost calculation fails, continue without cost - pass - - # append the completed event - self._events = deltas + [ - ResponseCompletedEvent( - type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, - response=transformed, - ) - ] self._idx = 0 + self.completed_response = self._events[-1] def __aiter__(self): return self - async def __anext__(self) -> ResponsesAPIStreamingResponse: + async def __anext__(self) -> Any: if self._idx >= len(self._events): raise StopAsyncIteration evt = self._events[self._idx] self._idx += 1 + openai_types = _get_openai_response_types() + if ( + getattr(evt, "type", None) + == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ): + self.completed_response = evt + self._log_completed_response(is_async=True) return evt def __iter__(self): return self - def __next__(self) -> ResponsesAPIStreamingResponse: + def __next__(self) -> Any: if self._idx >= len(self._events): raise StopIteration evt = self._events[self._idx] self._idx += 1 + openai_types = _get_openai_response_types() + if ( + getattr(evt, "type", None) + == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ): + self.completed_response = evt + self._log_completed_response(is_async=False) return evt - def _collect_text(self, resp: ResponsesAPIResponse) -> str: - out = "" - for out_item in resp.output: - item_type = getattr(out_item, "type", None) - if item_type == "message": - for c in getattr(out_item, "content", []): - out += c.text - return out + +class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): + def __init__( + self, + response: Any, + logging_obj: LiteLLMLoggingObj, + request_data: Optional[Dict[str, Any]] = None, + call_type: Optional[str] = None, + ): + BaseResponsesAPIStreamingIterator.__init__( + self, + response=httpx.Response(200), + model=getattr(response, "model", ""), + responses_api_provider_config=None, + logging_obj=logging_obj, + litellm_metadata=None, + custom_llm_provider="cached_response", + request_data=request_data, + call_type=call_type, + ) + self._completed_response_cache_hit = True + self._persist_completed_response_before_logging = False + self._events: List[Any] = [] + self._idx = 0 + self._set_events_from_response(transformed=response, logging_obj=logging_obj) + + def _set_events_from_response( + self, + transformed: Any, + logging_obj: LiteLLMLoggingObj, + ) -> None: + self._events = _build_synthetic_response_events( + transformed=transformed, + logging_obj=logging_obj, + chunk_size=MockResponsesAPIStreamingIterator.CHUNK_SIZE, + ) + self._idx = 0 + self.completed_response = self._events[-1] + + def __aiter__(self): + return self + + async def __anext__(self) -> Any: + if self._idx >= len(self._events): + raise StopAsyncIteration + evt = self._events[self._idx] + self._idx += 1 + openai_types = _get_openai_response_types() + if ( + getattr(evt, "type", None) + == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ): + self.completed_response = evt + self._log_completed_response(is_async=True) + return evt + + def __iter__(self): + return self + + def __next__(self) -> Any: + if self._idx >= len(self._events): + raise StopIteration + evt = self._events[self._idx] + self._idx += 1 + openai_types = _get_openai_response_types() + if ( + getattr(evt, "type", None) + == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ): + self.completed_response = evt + self._log_completed_response(is_async=False) + return evt + + +def _dump_response_object(obj: Any) -> Dict[str, Any]: + if hasattr(obj, "model_dump"): + return obj.model_dump() + if isinstance(obj, dict): + return obj + return {} + + +def _build_response_status_event( + event_type: Literal[ + "response.created", + "response.in_progress", + ], + transformed: Any, +) -> Any: + openai_types = _get_openai_response_types() + in_progress_response = transformed.model_copy( + deep=True, + update={"status": "in_progress", "output": []}, + ) + if event_type == openai_types.ResponsesAPIStreamEvents.RESPONSE_CREATED: + return openai_types.ResponseCreatedEvent( + type=event_type, response=in_progress_response + ) + return openai_types.ResponseInProgressEvent( + type=event_type, response=in_progress_response + ) + + +def _build_content_part_done_event( + *, + item_id: str, + output_index: int, + content_index: int, + part_payload: Dict[str, Any], +) -> Optional[Any]: + openai_types = _get_openai_response_types() + part_type = part_payload.get("type") + part: Any + if part_type == "output_text": + annotations = [ + openai_types.BaseLiteLLMOpenAIResponseObject(**annotation) + for annotation in part_payload.get("annotations", []) or [] + ] + part = openai_types.ContentPartDonePartOutputText( + type="output_text", + text=str(part_payload.get("text") or ""), + annotations=annotations, + logprobs=part_payload.get("logprobs"), + ) + elif part_type == "refusal": + part = openai_types.ContentPartDonePartRefusal( + type="refusal", + refusal=str(part_payload.get("refusal") or ""), + ) + elif part_type == "reasoning_text": + part = openai_types.ContentPartDonePartReasoningText( + type="reasoning_text", + reasoning=str(part_payload.get("reasoning") or ""), + ) + else: + return None + + return openai_types.ContentPartDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.CONTENT_PART_DONE, + item_id=item_id, + output_index=output_index, + content_index=content_index, + part=part, + ) + + +def _add_text_like_part_events( + *, + events: List[Any], + item_id: str, + output_index: int, + content_index: int, + part_payload: Dict[str, Any], + chunk_size: int, +) -> None: + openai_types = _get_openai_response_types() + part_type = part_payload.get("type") + if part_type == "output_text": + text = str(part_payload.get("text") or "") + for i in range(0, len(text), chunk_size): + events.append( + openai_types.OutputTextDeltaEvent( + type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id=item_id, + output_index=output_index, + content_index=content_index, + delta=text[i : i + chunk_size], + ) + ) + for annotation_index, annotation in enumerate( + part_payload.get("annotations", []) or [] + ): + events.append( + openai_types.OutputTextAnnotationAddedEvent( + type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED, + item_id=item_id, + output_index=output_index, + content_index=content_index, + annotation_index=annotation_index, + annotation=annotation, + ) + ) + events.append( + openai_types.OutputTextDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE, + item_id=item_id, + output_index=output_index, + content_index=content_index, + text=text, + ) + ) + elif part_type == "refusal": + refusal = str(part_payload.get("refusal") or "") + for i in range(0, len(refusal), chunk_size): + events.append( + openai_types.RefusalDeltaEvent( + type=openai_types.ResponsesAPIStreamEvents.REFUSAL_DELTA, + item_id=item_id, + output_index=output_index, + content_index=content_index, + delta=refusal[i : i + chunk_size], + ) + ) + events.append( + openai_types.RefusalDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.REFUSAL_DONE, + item_id=item_id, + output_index=output_index, + content_index=content_index, + refusal=refusal, + ) + ) + + +def _build_synthetic_response_events( + *, + transformed: Any, + logging_obj: LiteLLMLoggingObj, + chunk_size: int, +) -> List[Any]: + openai_types = _get_openai_response_types() + if litellm.include_cost_in_streaming_usage and logging_obj is not None: + usage_obj: Optional[Any] = getattr(transformed, "usage", None) + if usage_obj is not None: + try: + cost: Optional[float] = logging_obj._response_cost_calculator( + result=transformed + ) + if cost is not None: + setattr(usage_obj, "cost", cost) + except Exception: + pass + + events: List[Any] = [ + _build_response_status_event( + openai_types.ResponsesAPIStreamEvents.RESPONSE_CREATED, transformed + ), + _build_response_status_event( + openai_types.ResponsesAPIStreamEvents.RESPONSE_IN_PROGRESS, transformed + ), + ] + + sequence_number = 0 + for output_index, output_item in enumerate( + getattr(transformed, "output", []) or [] + ): + output_item_payload = _dump_response_object(output_item) + item_id = str(output_item_payload.get("id") or transformed.id) + item_type = output_item_payload.get("type") + + events.append( + openai_types.OutputItemAddedEvent( + type=openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + output_index=output_index, + item=openai_types.BaseLiteLLMOpenAIResponseObject( + **output_item_payload + ), + ) + ) + + if item_type == "message": + for content_index, part in enumerate( + output_item_payload.get("content", []) or [] + ): + part_payload = _dump_response_object(part) + events.append( + openai_types.ContentPartAddedEvent( + type=openai_types.ResponsesAPIStreamEvents.CONTENT_PART_ADDED, + item_id=item_id, + output_index=output_index, + content_index=content_index, + part=openai_types.BaseLiteLLMOpenAIResponseObject( + **part_payload + ), + ) + ) + _add_text_like_part_events( + events=events, + item_id=item_id, + output_index=output_index, + content_index=content_index, + part_payload=part_payload, + chunk_size=chunk_size, + ) + done_event = _build_content_part_done_event( + item_id=item_id, + output_index=output_index, + content_index=content_index, + part_payload=part_payload, + ) + if done_event is not None: + events.append(done_event) + elif item_type == "function_call": + arguments = str(output_item_payload.get("arguments") or "") + for i in range(0, len(arguments), chunk_size): + events.append( + openai_types.FunctionCallArgumentsDeltaEvent( + type=openai_types.ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, + item_id=item_id, + output_index=output_index, + delta=arguments[i : i + chunk_size], + ) + ) + events.append( + openai_types.FunctionCallArgumentsDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE, + item_id=item_id, + output_index=output_index, + arguments=arguments, + ) + ) + elif item_type == "reasoning": + for summary_index, summary in enumerate( + output_item_payload.get("summary", []) or [] + ): + summary_payload = _dump_response_object(summary) + summary_text = str(summary_payload.get("text") or "") + for i in range(0, len(summary_text), chunk_size): + events.append( + openai_types.ReasoningSummaryTextDeltaEvent( + type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA, + item_id=item_id, + output_index=output_index, + summary_index=summary_index, + delta=summary_text[i : i + chunk_size], + ) + ) + sequence_number += 1 + events.append( + openai_types.ReasoningSummaryTextDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DONE, + item_id=item_id, + output_index=output_index, + sequence_number=sequence_number, + summary_index=summary_index, + text=summary_text, + ) + ) + sequence_number += 1 + events.append( + openai_types.ReasoningSummaryPartDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_PART_DONE, + item_id=item_id, + output_index=output_index, + sequence_number=sequence_number, + summary_index=summary_index, + part=openai_types.BaseLiteLLMOpenAIResponseObject( + **summary_payload + ), + ) + ) + + sequence_number += 1 + events.append( + openai_types.OutputItemDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + output_index=output_index, + sequence_number=sequence_number, + item=openai_types.BaseLiteLLMOpenAIResponseObject( + **output_item_payload + ), + ) + ) + + events.append( + openai_types.ResponseCompletedEvent( + type=openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=transformed, + ) + ) + return events # --------------------------------------------------------------------------- @@ -951,8 +1395,8 @@ class ResponsesWebSocketStreaming: # --------------------------------------------------------------------------- _RESPONSE_CREATE_PARAMS: frozenset = ( - ResponsesAPIRequestParams.__required_keys__ - | ResponsesAPIRequestParams.__optional_keys__ + _get_openai_response_types().ResponsesAPIRequestParams.__required_keys__ + | _get_openai_response_types().ResponsesAPIRequestParams.__optional_keys__ ) _MANAGED_WS_SKIP_KWARGS: frozenset = frozenset( @@ -1085,7 +1529,7 @@ class ManagedResponsesWebSocketHandler: @staticmethod def _extract_output_messages( - completed_event: Dict[str, Any] + completed_event: Dict[str, Any], ) -> List[Dict[str, Any]]: """ Convert the output items in a ``response.completed`` event into diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 2fd0c4ea970..986ec39f3bb 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1482,6 +1482,7 @@ class ReasoningSummaryTextDeltaEvent(BaseLiteLLMOpenAIResponseObject): type: Literal[ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA] item_id: str output_index: int + summary_index: int = 0 delta: str @@ -1490,7 +1491,7 @@ class ReasoningSummaryTextDoneEvent(BaseLiteLLMOpenAIResponseObject): item_id: str output_index: int sequence_number: int - summary_index: int + summary_index: int = 0 text: str @@ -1499,7 +1500,7 @@ class ReasoningSummaryPartDoneEvent(BaseLiteLLMOpenAIResponseObject): item_id: str output_index: int sequence_number: int - summary_index: int + summary_index: int = 0 part: BaseLiteLLMOpenAIResponseObject diff --git a/tests/llm_responses_api_testing/test_responses_hooks.py b/tests/llm_responses_api_testing/test_responses_hooks.py index 3227fecdfb2..8933ec8f3c5 100644 --- a/tests/llm_responses_api_testing/test_responses_hooks.py +++ b/tests/llm_responses_api_testing/test_responses_hooks.py @@ -1,6 +1,9 @@ import asyncio +from contextlib import suppress from datetime import datetime +import json from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock import httpx import pytest @@ -8,8 +11,17 @@ import pytest import litellm from litellm.integrations.custom_logger import CustomLogger from litellm.responses import streaming_iterator as streaming_module -from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator -from litellm.types.llms.openai import ResponsesAPIStreamEvents +from litellm.responses.streaming_iterator import ( + CachedResponsesAPIStreamingIterator, + MockResponsesAPIStreamingIterator, + ResponsesAPIStreamingIterator, + SyncResponsesAPIStreamingIterator, +) +from litellm.types.llms.openai import ( + ResponseCompletedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, +) from litellm.types.utils import CallTypes @@ -19,15 +31,19 @@ class _FakeLoggingObj: self.async_success_calls = 0 self.failure_calls = 0 self.async_failure_calls = 0 + self.last_success_kwargs = None + self.last_async_success_kwargs = None self.start_time = datetime.now() self.model_call_details = {"litellm_params": {}} # Signature alignment with Logging handlers def success_handler(self, *args, **kwargs): self.success_calls += 1 + self.last_success_kwargs = kwargs async def async_success_handler(self, *args, **kwargs): self.async_success_calls += 1 + self.last_async_success_kwargs = kwargs def failure_handler(self, *args, **kwargs): self.failure_calls += 1 @@ -36,6 +52,115 @@ class _FakeLoggingObj: self.async_failure_calls += 1 +def _make_completed_response(response_id: str = "resp_test") -> ResponseCompletedEvent: + return ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=ResponsesAPIResponse( + id=response_id, + created_at=int(datetime.now().timestamp()), + status="completed", + model="test-model", + object="response", + output=[ + { + "type": "message", + "id": f"msg_{response_id}", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "cached streamed response", + "annotations": [], + } + ], + } + ], + ), + ) + + +@pytest.mark.asyncio +async def test_log_background_task_failure_logs_task_exceptions(monkeypatch): + error_logger = MagicMock() + monkeypatch.setattr(streaming_module.verbose_logger, "error", error_logger) + + async def _boom(): + raise RuntimeError("boom") + + task = asyncio.create_task(_boom()) + with suppress(RuntimeError): + await task + + streaming_module._log_background_task_failure(task, task_name="cache write") + + error_logger.assert_called_once() + assert error_logger.call_args.args == ( + "%s failed: %s", + "cache write", + task.exception(), + ) + + +@pytest.mark.asyncio +async def test_log_background_task_failure_ignores_cancelled_tasks(monkeypatch): + error_logger = MagicMock() + monkeypatch.setattr(streaming_module.verbose_logger, "error", error_logger) + + task = asyncio.create_task(asyncio.sleep(1)) + task.cancel() + with suppress(asyncio.CancelledError): + await task + + streaming_module._log_background_task_failure(task, task_name="cache write") + + error_logger.assert_not_called() + + +def test_content_part_done_event_supports_refusal_and_reasoning_text(): + refusal_event = streaming_module._build_content_part_done_event( + item_id="msg_1", + output_index=0, + content_index=0, + part_payload={"type": "refusal", "refusal": "no"}, + ) + reasoning_event = streaming_module._build_content_part_done_event( + item_id="msg_1", + output_index=0, + content_index=1, + part_payload={"type": "reasoning_text", "reasoning": "because"}, + ) + unsupported_event = streaming_module._build_content_part_done_event( + item_id="msg_1", + output_index=0, + content_index=2, + part_payload={"type": "image"}, + ) + + assert refusal_event.part.type == "refusal" + assert refusal_event.part.refusal == "no" + assert reasoning_event.part.type == "reasoning_text" + assert reasoning_event.part.reasoning == "because" + assert unsupported_event is None + + +def test_dump_response_object_handles_model_and_unknown_values(): + response = ResponsesAPIResponse( + id="resp_dump", + created_at=int(datetime.now().timestamp()), + status="completed", + model="gpt-4.1-mini", + object="response", + output=[], + ) + + assert streaming_module._dump_response_object(response)["id"] == "resp_dump" + assert streaming_module._dump_response_object({"type": "message"}) == { + "type": "message" + } + assert streaming_module._dump_response_object(object()) == {} + + @pytest.mark.asyncio async def test_responses_streaming_triggers_hooks(monkeypatch): """ @@ -167,3 +292,762 @@ async def test_responses_streaming_failure_triggers_failure_handlers(): await asyncio.sleep(0.2) assert logging_obj.failure_calls >= 1 assert logging_obj.async_failure_calls >= 1 + + +def test_process_chunk_requires_provider_config(): + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=None, + logging_obj=_FakeLoggingObj(), + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + + with pytest.raises(ValueError, match="responses_api_provider_config is required"): + iterator._process_chunk(json.dumps({"type": "response.completed"})) + + +def test_process_chunk_wraps_encrypted_content_with_model_id(): + openai_types = streaming_module._get_openai_response_types() + + class _EncryptedConfig: + def transform_streaming_response(self, **kwargs): + return openai_types.OutputItemAddedEvent( + type=openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + output_index=0, + item=openai_types.BaseLiteLLMOpenAIResponseObject( + id="rs_123", + type="reasoning", + encrypted_content="ciphertext", + ), + ) + + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_EncryptedConfig(), + logging_obj=_FakeLoggingObj(), + litellm_metadata={ + "encrypted_content_affinity_enabled": True, + "model_info": {"id": "model-123"}, + }, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + + event = iterator._process_chunk(json.dumps({"type": "response.output_item.added"})) + + assert event.item.encrypted_content.startswith("litellm_enc:") + assert event.item.encrypted_content.endswith(";ciphertext") + + +def test_process_chunk_completed_response_updates_id_and_usage_cost(monkeypatch): + original_include_cost = litellm.include_cost_in_streaming_usage + litellm.include_cost_in_streaming_usage = True + openai_types = streaming_module._get_openai_response_types() + + class _CompletedConfig: + def transform_streaming_response(self, **kwargs): + return openai_types.ResponseCompletedEvent( + type=openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=ResponsesAPIResponse( + id="resp_live", + created_at=int(datetime.now().timestamp()), + status="completed", + model="test-model", + object="response", + output=[], + usage=openai_types.ResponseAPIUsage( + input_tokens=1, + output_tokens=2, + total_tokens=3, + ), + ), + ) + + logging_obj = _FakeLoggingObj() + logging_obj._response_cost_calculator = MagicMock(return_value=1.23) + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_CompletedConfig(), + logging_obj=logging_obj, + litellm_metadata={"model_info": {"id": "model-123"}}, + custom_llm_provider="openai", + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + completion_handler = MagicMock() + monkeypatch.setattr( + iterator, "_handle_logging_completed_response", completion_handler + ) + + try: + event = iterator._process_chunk(json.dumps({"type": "response.completed"})) + finally: + litellm.include_cost_in_streaming_usage = original_include_cost + + assert iterator.completed_response is event + assert event.response.id != "resp_live" + assert event.response.id.startswith("resp_") + assert event.response.usage.cost == 1.23 + completion_handler.assert_called_once() + + +def test_process_chunk_failed_response_triggers_failure_logging(monkeypatch): + openai_types = streaming_module._get_openai_response_types() + + class _FailedConfig: + def transform_streaming_response(self, **kwargs): + return openai_types.ResponseFailedEvent( + type=openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED, + response=ResponsesAPIResponse( + id="resp_failed", + created_at=int(datetime.now().timestamp()), + status="failed", + model="test-model", + object="response", + output=[], + error={"message": "provider failed"}, + ), + ) + + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_FailedConfig(), + logging_obj=_FakeLoggingObj(), + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + failure_handler = MagicMock() + monkeypatch.setattr(iterator, "_handle_logging_failed_response", failure_handler) + + event = iterator._process_chunk(json.dumps({"type": "response.failed"})) + + assert iterator.completed_response is event + failure_handler.assert_called_once() + + +@pytest.mark.asyncio +async def test_handle_logging_failed_response_uses_response_error_message(): + openai_types = streaming_module._get_openai_response_types() + logging_obj = _FakeLoggingObj() + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + iterator.completed_response = openai_types.ResponseFailedEvent( + type=openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED, + response=ResponsesAPIResponse( + id="resp_failed_real", + created_at=int(datetime.now().timestamp()), + status="failed", + model="test-model", + object="response", + output=[], + error={"message": "provider failed"}, + ), + ) + + iterator._handle_logging_failed_response() + await asyncio.sleep(0.2) + + assert logging_obj.failure_calls == 1 + assert logging_obj.async_failure_calls == 1 + + +def test_process_chunk_returns_none_for_invalid_json_and_non_dict_payload(): + class _NoopConfig: + def transform_streaming_response(self, **kwargs): + raise AssertionError("should not be called") + + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_NoopConfig(), + logging_obj=_FakeLoggingObj(), + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + + assert iterator._process_chunk("not-json") is None + assert iterator._process_chunk(json.dumps(["not", "a", "dict"])) is None + + +def test_process_chunk_cost_annotation_failure_is_nonfatal(monkeypatch): + original_include_cost = litellm.include_cost_in_streaming_usage + litellm.include_cost_in_streaming_usage = True + openai_types = streaming_module._get_openai_response_types() + + class _CompletedConfig: + def transform_streaming_response(self, **kwargs): + return openai_types.ResponseCompletedEvent( + type=openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=ResponsesAPIResponse( + id="resp_cost_failure", + created_at=int(datetime.now().timestamp()), + status="completed", + model="test-model", + object="response", + output=[], + usage=openai_types.ResponseAPIUsage( + input_tokens=1, + output_tokens=2, + total_tokens=3, + ), + ), + ) + + logging_obj = _FakeLoggingObj() + logging_obj._response_cost_calculator = MagicMock(side_effect=RuntimeError("boom")) + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_CompletedConfig(), + logging_obj=logging_obj, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + completion_handler = MagicMock() + monkeypatch.setattr( + iterator, "_handle_logging_completed_response", completion_handler + ) + + try: + event = iterator._process_chunk(json.dumps({"type": "response.completed"})) + finally: + litellm.include_cost_in_streaming_usage = original_include_cost + + assert iterator.completed_response is event + assert event.response.usage.cost is None + completion_handler.assert_called_once() + + +def test_get_completed_response_object_accepts_direct_response(): + logging_obj = _FakeLoggingObj() + iterator = SyncResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + direct_response = _make_completed_response("resp_direct").response + iterator.completed_response = direct_response + + assert iterator._get_completed_response_object() is direct_response + + +@pytest.mark.asyncio +async def test_responses_streaming_completed_event_persists_async_cache(): + logging_obj = _FakeLoggingObj() + original_cache = litellm.cache + litellm.cache = SimpleNamespace( + async_add_cache=AsyncMock(), + add_cache=MagicMock(), + ) + caching_handler = SimpleNamespace( + request_kwargs={ + "model": "test-model", + "input": "hello", + "stream": True, + "caching": True, + "cache_key": "stale-request-cache-key", + "metadata": None, + "custom_llm_provider": "openai", + }, + preset_cache_key="responses-stream-cache-key", + original_function=litellm.aresponses, + async_set_cache=AsyncMock(), + _should_store_result_in_cache=lambda original_function, kwargs: True, + ) + logging_obj._llm_caching_handler = caching_handler + + iterator = ResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data=caching_handler.request_kwargs, + call_type=CallTypes.aresponses.value, + ) + iterator.completed_response = _make_completed_response() + + iterator._handle_logging_completed_response() + await asyncio.sleep(0.2) + + litellm.cache.async_add_cache.assert_called_once() + assert litellm.cache.async_add_cache.call_args.kwargs["stream"] is True + assert ( + litellm.cache.async_add_cache.call_args.kwargs["cache_key"] + == "responses-stream-cache-key" + ) + assert "metadata" not in litellm.cache.async_add_cache.call_args.kwargs + assert "custom_llm_provider" not in litellm.cache.async_add_cache.call_args.kwargs + assert ( + json.loads(litellm.cache.async_add_cache.call_args.args[0])["id"] + == iterator.completed_response.response.id + ) + litellm.cache = original_cache + + +def test_responses_streaming_completed_event_persists_sync_cache(): + logging_obj = _FakeLoggingObj() + original_cache = litellm.cache + litellm.cache = SimpleNamespace( + async_add_cache=AsyncMock(), + add_cache=MagicMock(), + ) + caching_handler = SimpleNamespace( + request_kwargs={ + "model": "test-model", + "input": "hello", + "stream": True, + "caching": True, + "cache_key": "stale-request-cache-key", + "metadata": None, + "custom_llm_provider": "openai", + }, + preset_cache_key="responses-stream-cache-key", + original_function=litellm.responses, + sync_set_cache=MagicMock(), + _should_store_result_in_cache=lambda original_function, kwargs: True, + ) + logging_obj._llm_caching_handler = caching_handler + + iterator = SyncResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data=caching_handler.request_kwargs, + call_type=CallTypes.responses.value, + ) + iterator.completed_response = _make_completed_response("resp_sync") + + iterator._handle_logging_completed_response() + + litellm.cache.add_cache.assert_called_once() + assert litellm.cache.add_cache.call_args.kwargs["stream"] is True + assert ( + litellm.cache.add_cache.call_args.kwargs["cache_key"] + == "responses-stream-cache-key" + ) + assert "metadata" not in litellm.cache.add_cache.call_args.kwargs + assert "custom_llm_provider" not in litellm.cache.add_cache.call_args.kwargs + assert ( + json.loads(litellm.cache.add_cache.call_args.args[0])["id"] + == iterator.completed_response.response.id + ) + litellm.cache = original_cache + + +def test_log_completed_response_sync_direct_path(monkeypatch): + hook_calls = {"post_call": 0, "metadata": 0} + + async def fake_post_call(request_data, response, call_type): + hook_calls["post_call"] += 1 + + def fake_update_metadata(**kwargs): + hook_calls["metadata"] += 1 + + monkeypatch.setattr( + streaming_module, + "async_post_call_success_deployment_hook", + fake_post_call, + ) + monkeypatch.setattr( + streaming_module, + "update_response_metadata", + fake_update_metadata, + ) + + logging_obj = _FakeLoggingObj() + iterator = SyncResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + iterator._persist_completed_response_before_logging = False + iterator.completed_response = _make_completed_response("resp_log_sync") + + iterator._log_completed_response(is_async=False) + asyncio.run(asyncio.sleep(0.2)) + + assert logging_obj.success_calls == 1 + assert logging_obj.async_success_calls == 1 + assert hook_calls["post_call"] == 1 + assert hook_calls["metadata"] == 1 + + +def test_log_completed_response_falls_back_when_model_validate_fails(monkeypatch): + class _BadSerializableResponse: + @classmethod + def model_validate(cls, value): + raise RuntimeError("nope") + + def model_dump(self): + return {"id": "bad"} + + logging_obj = _FakeLoggingObj() + iterator = SyncResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + iterator._persist_completed_response_before_logging = False + iterator.completed_response = _BadSerializableResponse() + monkeypatch.setattr(iterator, "_run_post_success_hooks", MagicMock()) + + iterator._log_completed_response(is_async=False) + asyncio.run(asyncio.sleep(0.2)) + + assert logging_obj.success_calls == 1 + assert logging_obj.async_success_calls == 1 + + +@pytest.mark.parametrize( + "scenario", + [ + "already_cached", + "not_completed", + "missing_caching_handler", + "not_streaming", + "store_disabled", + "missing_cache_backend", + ], +) +def test_persist_completed_response_to_cache_guard_branches(monkeypatch, scenario): + logging_obj = _FakeLoggingObj() + iterator = SyncResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=SimpleNamespace(), + logging_obj=logging_obj, + request_data={"foo": "bar"}, + call_type=CallTypes.responses.value, + ) + openai_types = streaming_module._get_openai_response_types() + completed_event = _make_completed_response("resp_guard") + iterator.completed_response = completed_event + + if scenario == "already_cached": + iterator._completed_response_cached = True + elif scenario == "not_completed": + iterator.completed_response = openai_types.ResponseIncompleteEvent( + type=openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, + response=completed_event.response, + ) + elif scenario == "missing_caching_handler": + logging_obj._llm_caching_handler = None + else: + logging_obj._llm_caching_handler = SimpleNamespace( + request_kwargs={ + "model": "test-model", + "input": "hello", + "stream": scenario != "not_streaming", + "cache_key": "request-cache-key", + "metadata": None, + "custom_llm_provider": "openai", + }, + preset_cache_key=None, + original_function=litellm.responses, + dual_cache=None, + _should_store_result_in_cache=lambda original_function, kwargs: ( + scenario != "store_disabled" + ), + ) + if scenario == "missing_cache_backend": + monkeypatch.setattr(streaming_module.litellm, "cache", None) + else: + monkeypatch.setattr( + streaming_module.litellm, + "cache", + SimpleNamespace(add_cache=MagicMock(), async_add_cache=AsyncMock()), + ) + + iterator._persist_completed_response_to_cache(is_async=False) + + expected_cached_flag = scenario == "already_cached" + assert iterator._completed_response_cached is expected_cached_flag + + +def test_build_synthetic_response_events_covers_annotations_function_calls_and_refusals(): + original_include_cost = litellm.include_cost_in_streaming_usage + litellm.include_cost_in_streaming_usage = True + logging_obj = _FakeLoggingObj() + logging_obj._response_cost_calculator = MagicMock(side_effect=RuntimeError("boom")) + transformed = ResponsesAPIResponse( + id="resp_events", + created_at=int(datetime.now().timestamp()), + status="completed", + model="gpt-4.1-mini", + object="response", + output=[ + { + "type": "message", + "id": "msg_events", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "hello world", + "annotations": [{"type": "file_citation", "file_id": "file_1"}], + }, + { + "type": "refusal", + "refusal": "no thanks", + }, + ], + }, + { + "type": "function_call", + "id": "fc_events", + "call_id": "call_123", + "name": "lookup", + "arguments": '{"id":1}', + }, + ], + ) + + try: + events = streaming_module._build_synthetic_response_events( + transformed=transformed, + logging_obj=logging_obj, + chunk_size=5, + ) + finally: + litellm.include_cost_in_streaming_usage = original_include_cost + + event_types = [ + event.type.value if hasattr(event.type, "value") else str(event.type) + for event in events + ] + + assert "response.output_text.annotation.added" in event_types + assert "response.refusal.delta" in event_types + assert "response.refusal.done" in event_types + assert "response.function_call_arguments.delta" in event_types + assert "response.function_call_arguments.done" in event_types + assert event_types[-1] == "response.completed" + + +@pytest.mark.asyncio +async def test_mock_responses_streaming_iterator_async_iteration_logs_completion( + monkeypatch, +): + hook_calls = {"post_call": 0, "metadata": 0} + + async def fake_post_call(request_data, response, call_type): + hook_calls["post_call"] += 1 + + def fake_update_metadata(**kwargs): + hook_calls["metadata"] += 1 + + monkeypatch.setattr( + streaming_module, + "async_post_call_success_deployment_hook", + fake_post_call, + ) + monkeypatch.setattr( + streaming_module, + "update_response_metadata", + fake_update_metadata, + ) + + class _MockTransformConfig: + def transform_response_api_response(self, **kwargs): + return _make_completed_response("resp_mock").response + + logging_obj = _FakeLoggingObj() + + iterator = MockResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_MockTransformConfig(), + logging_obj=logging_obj, + request_data={"model": "test-model", "stream": True}, + call_type=CallTypes.responses.value, + ) + + streamed_events = [event async for event in iterator] + await asyncio.sleep(0.2) + + assert streamed_events[0].type == ResponsesAPIStreamEvents.RESPONSE_CREATED + assert streamed_events[-1].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert logging_obj.success_calls == 1 + assert logging_obj.async_success_calls == 1 + assert hook_calls["post_call"] == 1 + assert hook_calls["metadata"] == 1 + + +def test_mock_responses_streaming_iterator_sync_iteration_logs_completion(monkeypatch): + hook_calls = {"post_call": 0, "metadata": 0} + + async def fake_post_call(request_data, response, call_type): + hook_calls["post_call"] += 1 + + def fake_update_metadata(**kwargs): + hook_calls["metadata"] += 1 + + monkeypatch.setattr( + streaming_module, + "async_post_call_success_deployment_hook", + fake_post_call, + ) + monkeypatch.setattr( + streaming_module, + "update_response_metadata", + fake_update_metadata, + ) + + class _MockTransformConfig: + def transform_response_api_response(self, **kwargs): + return _make_completed_response("resp_mock_sync").response + + logging_obj = _FakeLoggingObj() + iterator = MockResponsesAPIStreamingIterator( + response=httpx.Response(200), + model="test-model", + responses_api_provider_config=_MockTransformConfig(), + logging_obj=logging_obj, + request_data={"model": "test-model", "stream": True}, + call_type=CallTypes.responses.value, + ) + + streamed_events = list(iterator) + asyncio.run(asyncio.sleep(0.2)) + + assert streamed_events[0].type == ResponsesAPIStreamEvents.RESPONSE_CREATED + assert streamed_events[-1].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert logging_obj.success_calls == 1 + assert logging_obj.async_success_calls == 1 + assert hook_calls["post_call"] == 1 + assert hook_calls["metadata"] == 1 + + +@pytest.mark.asyncio +async def test_cached_responses_stream_async_hit_triggers_success_callbacks( + monkeypatch, +): + hook_calls = {"post_call": 0, "metadata": 0} + + async def fake_post_call(request_data, response, call_type): + hook_calls["post_call"] += 1 + + def fake_update_metadata(**kwargs): + hook_calls["metadata"] += 1 + + monkeypatch.setattr( + streaming_module, + "async_post_call_success_deployment_hook", + fake_post_call, + ) + monkeypatch.setattr( + streaming_module, + "update_response_metadata", + fake_update_metadata, + ) + + logging_obj = _FakeLoggingObj() + original_cache = litellm.cache + litellm.cache = SimpleNamespace( + async_add_cache=AsyncMock(), + add_cache=MagicMock(), + ) + logging_obj._llm_caching_handler = SimpleNamespace( + request_kwargs={"model": "test-model", "input": "hello", "stream": True}, + preset_cache_key="responses-stream-cache-key", + original_function=litellm.aresponses, + _should_store_result_in_cache=lambda original_function, kwargs: True, + ) + + iterator = CachedResponsesAPIStreamingIterator( + response=_make_completed_response("resp_cached_async").response, + logging_obj=logging_obj, + request_data={"model": "test-model", "input": "hello", "stream": True}, + call_type=CallTypes.aresponses.value, + ) + + streamed_events = [event async for event in iterator] + await asyncio.sleep(0.2) + + assert streamed_events[-1].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert logging_obj.success_calls == 1 + assert logging_obj.async_success_calls == 1 + assert logging_obj.last_success_kwargs["cache_hit"] is True + assert logging_obj.last_async_success_kwargs["cache_hit"] is True + assert hook_calls["post_call"] == 1 + assert hook_calls["metadata"] == 1 + litellm.cache.async_add_cache.assert_not_called() + litellm.cache.add_cache.assert_not_called() + litellm.cache = original_cache + + +def test_cached_responses_stream_sync_hit_triggers_success_callbacks(monkeypatch): + hook_calls = {"post_call": 0, "metadata": 0} + + async def fake_post_call(request_data, response, call_type): + hook_calls["post_call"] += 1 + + def fake_update_metadata(**kwargs): + hook_calls["metadata"] += 1 + + monkeypatch.setattr( + streaming_module, + "async_post_call_success_deployment_hook", + fake_post_call, + ) + monkeypatch.setattr( + streaming_module, + "update_response_metadata", + fake_update_metadata, + ) + + logging_obj = _FakeLoggingObj() + original_cache = litellm.cache + litellm.cache = SimpleNamespace( + async_add_cache=AsyncMock(), + add_cache=MagicMock(), + ) + logging_obj._llm_caching_handler = SimpleNamespace( + request_kwargs={"model": "test-model", "input": "hello", "stream": True}, + preset_cache_key="responses-stream-cache-key", + original_function=litellm.responses, + _should_store_result_in_cache=lambda original_function, kwargs: True, + ) + + iterator = CachedResponsesAPIStreamingIterator( + response=_make_completed_response("resp_cached_sync").response, + logging_obj=logging_obj, + request_data={"model": "test-model", "input": "hello", "stream": True}, + call_type=CallTypes.responses.value, + ) + + streamed_events = list(iterator) + asyncio.run(asyncio.sleep(0.2)) + + assert streamed_events[-1].type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert logging_obj.success_calls == 1 + assert logging_obj.async_success_calls == 1 + assert logging_obj.last_success_kwargs["cache_hit"] is True + assert logging_obj.last_async_success_kwargs["cache_hit"] is True + assert hook_calls["post_call"] == 1 + assert hook_calls["metadata"] == 1 + litellm.cache.async_add_cache.assert_not_called() + litellm.cache.add_cache.assert_not_called() + litellm.cache = original_cache diff --git a/tests/local_testing/test_caching_handler.py b/tests/local_testing/test_caching_handler.py index 806f72bfde8..0ce9938e710 100644 --- a/tests/local_testing/test_caching_handler.py +++ b/tests/local_testing/test_caching_handler.py @@ -19,9 +19,14 @@ import pytest import litellm from litellm import aembedding, completion, embedding, aresponses, responses from litellm.caching.caching import Cache +from litellm.responses.streaming_iterator import CachedResponsesAPIStreamingIterator from unittest.mock import AsyncMock, patch, MagicMock -from litellm.caching.caching_handler import LLMCachingHandler, CachingHandlerResponse +from litellm.caching.caching_handler import ( + LLMCachingHandler, + CachingHandlerResponse, + _should_defer_streaming_cache_hit_callbacks, +) from litellm.caching.caching import LiteLLMCacheType from litellm.types.utils import CallTypes from litellm.types.rerank import RerankResponse @@ -627,6 +632,55 @@ async def test_async_responses_api_caching(): assert cached_response.cached_result._hidden_params["cache_hit"] == True +@pytest.mark.asyncio +async def test_async_get_cache_updates_request_kwargs_for_streaming_responses(): + """ + Ensure streamed responses retain the normalized lookup kwargs so a later + cache write can reuse the exact cache key from the read path. + """ + setup_cache() + + caching_handler = LLMCachingHandler( + original_function=aresponses, + request_kwargs={"stale": True}, + start_time=datetime.now(), + ) + + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.aresponses.value, + model="gpt-4o", + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + + kwargs = { + "model": "gpt-4o", + "input": "hello", + "stream": True, + "caching": True, + } + + await caching_handler._async_get_cache( + model="gpt-4o", + original_function=aresponses, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.aresponses.value, + kwargs=kwargs, + ) + + assert "stale" not in caching_handler.request_kwargs + assert caching_handler.request_kwargs["model"] == "gpt-4o" + assert caching_handler.request_kwargs["input"] == "hello" + assert caching_handler.request_kwargs["stream"] is True + assert caching_handler.request_kwargs["cache_key"] == litellm.cache.get_cache_key( + **caching_handler.request_kwargs + ) + + def test_sync_responses_api_caching(): """ Test that synchronous responses API calls are properly cached and retrieved. @@ -769,6 +823,349 @@ def test_convert_cached_responses_api_result_to_model_response(): assert len(result.output) == 1 +def test_sync_get_cache_does_not_eagerly_log_streaming_responses_hits(): + litellm.set_verbose = True + setup_cache() + caching_handler = LLMCachingHandler( + original_function=responses, request_kwargs={}, start_time=datetime.now() + ) + + original_model = "gpt-4o" + responses_api_response = ResponsesAPIResponse( + id="resp_stream_sync_hit", + created_at=int(time.time()), + status="completed", + model=original_model, + object="response", + output=[ + { + "type": "message", + "id": "msg_stream_sync_hit", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "Sync streamed cache hit response.", + "annotations": [], + } + ], + } + ], + ) + + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.responses.value, + model=original_model, + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock() + + kwargs = { + "model": original_model, + "input": "Tell me a cached story", + "stream": True, + "caching": True, + } + + caching_handler.sync_set_cache(result=responses_api_response, kwargs=kwargs) + time.sleep(0.2) + + cached_response = caching_handler._sync_get_cache( + model=original_model, + original_function=responses, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.responses.value, + kwargs=kwargs, + ) + + assert cached_response.cached_result is not None + assert isinstance( + cached_response.cached_result, CachedResponsesAPIStreamingIterator + ) + logging_obj.handle_sync_success_callbacks_for_async_calls.assert_not_called() + + +def test_sync_get_cache_still_eagerly_logs_streaming_completion_hits(): + litellm.set_verbose = True + setup_cache() + caching_handler = LLMCachingHandler( + original_function=completion, request_kwargs={}, start_time=datetime.now() + ) + + original_model = "gpt-4o" + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.completion.value, + model=original_model, + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock() + + kwargs = { + "model": original_model, + "messages": [{"role": "user", "content": "Tell me a cached joke"}], + "stream": True, + "caching": True, + } + + caching_handler.sync_set_cache(result=chat_completion_response, kwargs=kwargs) + time.sleep(0.2) + + cached_response = caching_handler._sync_get_cache( + model=original_model, + original_function=completion, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.completion.value, + kwargs=kwargs, + ) + + assert cached_response.cached_result is not None + logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once() + + +def test_should_defer_streaming_cache_hit_callbacks_only_for_responses_streams(): + assert ( + _should_defer_streaming_cache_hit_callbacks( + call_type=CallTypes.responses.value, + kwargs={"stream": True}, + ) + is True + ) + assert ( + _should_defer_streaming_cache_hit_callbacks( + call_type=CallTypes.aresponses.value, + kwargs={"stream": True}, + ) + is True + ) + assert ( + _should_defer_streaming_cache_hit_callbacks( + call_type=CallTypes.completion.value, + kwargs={"stream": True}, + ) + is False + ) + assert ( + _should_defer_streaming_cache_hit_callbacks( + call_type=CallTypes.responses.value, + kwargs={"stream": False}, + ) + is False + ) + + +@pytest.mark.asyncio +async def test_async_get_cache_still_eagerly_logs_streaming_completion_hits(): + litellm.set_verbose = True + setup_cache() + caching_handler = LLMCachingHandler( + original_function=completion, request_kwargs={}, start_time=datetime.now() + ) + + original_model = "gpt-4o" + kwargs = { + "model": original_model, + "messages": [{"role": "user", "content": "Tell me a cached joke"}], + "stream": True, + "caching": True, + } + + await caching_handler.async_set_cache( + result=chat_completion_response, + original_function=litellm.acompletion, + kwargs=kwargs, + ) + await asyncio.sleep(0.2) + + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.acompletion.value, + model=original_model, + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + caching_handler._async_log_cache_hit_on_callbacks = MagicMock() + + cached_response = await caching_handler._async_get_cache( + model=original_model, + original_function=litellm.acompletion, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.acompletion.value, + kwargs=kwargs, + ) + + assert cached_response is not None + assert cached_response.cached_result is not None + caching_handler._async_log_cache_hit_on_callbacks.assert_called_once() + + +def test_convert_cached_streaming_responses_result_to_iterator(): + """ + Test that cached streaming Responses results are replayed through a synthetic + streaming iterator instead of being returned as a full response object. + """ + caching_handler = LLMCachingHandler( + original_function=responses, request_kwargs={}, start_time=datetime.now() + ) + + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.responses.value, + model="gpt-4o", + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + + cached_result = { + "id": "resp_stream_cache_test", + "created_at": int(time.time()), + "status": "completed", + "model": "gpt-4o", + "object": "response", + "output": [ + { + "type": "message", + "id": "msg_stream_cache_test", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "Streaming cache replay test.", + "annotations": [], + } + ], + } + ], + } + + result = caching_handler._convert_cached_result_to_model_response( + cached_result=cached_result, + call_type=CallTypes.responses.value, + kwargs={"model": "gpt-4o", "input": "test", "stream": True}, + logging_obj=logging_obj, + model="gpt-4o", + args=(), + ) + + assert isinstance(result, CachedResponsesAPIStreamingIterator) + assert result.completed_response is not None + assert result.completed_response.response.id == cached_result["id"] + + streamed_events = list(result) + assert streamed_events[0].type == "response.created" + assert streamed_events[1].type == "response.in_progress" + assert streamed_events[2].type == "response.output_item.added" + assert streamed_events[3].type == "response.content_part.added" + assert streamed_events[-4].type == "response.output_text.done" + assert streamed_events[-3].type == "response.content_part.done" + assert streamed_events[-2].type == "response.output_item.done" + assert streamed_events[-1].type == "response.completed" + assert streamed_events[-1].response.id == cached_result["id"] + assert streamed_events[-1].response.output[0].content[0].text == ( + "Streaming cache replay test." + ) + + +def test_convert_cached_streaming_reasoning_result_to_iterator(): + caching_handler = LLMCachingHandler( + original_function=responses, request_kwargs={}, start_time=datetime.now() + ) + + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.responses.value, + model="gpt-4o", + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + + cached_result = { + "id": "resp_stream_reasoning_cache_test", + "created_at": int(time.time()), + "status": "completed", + "model": "gpt-4o", + "object": "response", + "output": [ + { + "type": "reasoning", + "id": "rs_stream_cache_test", + "summary": [ + { + "type": "summary_text", + "text": "Cached reasoning summary.", + } + ], + } + ], + } + + result = caching_handler._convert_cached_result_to_model_response( + cached_result=cached_result, + call_type=CallTypes.responses.value, + kwargs={"model": "gpt-4o", "input": "test", "stream": True}, + logging_obj=logging_obj, + model="gpt-4o", + args=(), + ) + + assert isinstance(result, CachedResponsesAPIStreamingIterator) + + streamed_events = list(result) + streamed_event_types = [ + event.type.value if hasattr(event.type, "value") else str(event.type) + for event in streamed_events + ] + + assert streamed_event_types[:3] == [ + "response.created", + "response.in_progress", + "response.output_item.added", + ] + assert streamed_event_types[-4:] == [ + "response.reasoning_summary_text.done", + "response.reasoning_summary_part.done", + "response.output_item.done", + "response.completed", + ] + assert streamed_event_types.count("response.reasoning_summary_text.delta") >= 1 + + delta_events = [ + event + for event in streamed_events + if (event.type.value if hasattr(event.type, "value") else str(event.type)) + == "response.reasoning_summary_text.delta" + ] + text_done_event = streamed_events[-4] + part_done_event = streamed_events[-3] + output_item_done_event = streamed_events[-2] + + assert all(delta_event.summary_index == 0 for delta_event in delta_events) + assert text_done_event.text == "Cached reasoning summary." + assert text_done_event.summary_index == 0 + assert part_done_event.part.type == "summary_text" + assert part_done_event.part.text == "Cached reasoning summary." + assert output_item_done_event.item.type == "reasoning" + assert output_item_done_event.item.summary[0]["text"] == "Cached reasoning summary." + + @pytest.mark.asyncio async def test_responses_api_cache_with_different_inputs(): """ diff --git a/tests/local_testing/test_responses_stream_cache_keys.py b/tests/local_testing/test_responses_stream_cache_keys.py new file mode 100644 index 00000000000..5637028f550 --- /dev/null +++ b/tests/local_testing/test_responses_stream_cache_keys.py @@ -0,0 +1,141 @@ +from datetime import datetime +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm +from litellm import aresponses +from litellm._uuid import uuid +from litellm.caching.caching_handler import LLMCachingHandler +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging +from litellm.types.llms import openai as openai_types +from litellm.types.utils import CallTypes + + +@pytest.mark.asyncio +async def test_async_get_cache_reuses_preset_cache_key_for_responses(): + caching_handler = LLMCachingHandler( + original_function=aresponses, + request_kwargs={}, + start_time=datetime.now(), + ) + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.aresponses.value, + model="gpt-4.1-mini", + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + + original_cache = litellm.cache + mock_cache = MagicMock() + mock_cache.supported_call_types = [CallTypes.aresponses.value] + mock_cache._supports_async.return_value = True + mock_cache.get_cache_key.return_value = "responses-stream-cache-key" + mock_cache.async_get_cache = AsyncMock(return_value=None) + litellm.cache = mock_cache + + kwargs = { + "model": "gpt-4.1-mini", + "input": "hello", + "stream": True, + "litellm_params": {}, + } + await caching_handler._async_get_cache( + model="gpt-4.1-mini", + original_function=aresponses, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.aresponses.value, + kwargs=kwargs, + ) + + assert caching_handler.preset_cache_key == "responses-stream-cache-key" + mock_cache.async_get_cache.assert_awaited_once() + assert ( + mock_cache.async_get_cache.call_args.kwargs["cache_key"] + == "responses-stream-cache-key" + ) + + litellm.cache = original_cache + + +@pytest.mark.asyncio +async def test_async_get_cache_falls_back_to_sync_cache_for_responses(): + caching_handler = LLMCachingHandler( + original_function=aresponses, + request_kwargs={}, + start_time=datetime.now(), + ) + logging_obj = LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=CallTypes.aresponses.value, + model="gpt-4.1-mini", + messages=[], + function_id=str(uuid.uuid4()), + stream=True, + start_time=datetime.now(), + ) + + original_cache = litellm.cache + mock_cache = MagicMock() + mock_cache.supported_call_types = [CallTypes.aresponses.value] + mock_cache._supports_async.return_value = False + mock_cache.get_cache_key.return_value = "responses-stream-cache-key" + mock_cache.get_cache.return_value = None + litellm.cache = mock_cache + + kwargs = { + "model": "gpt-4.1-mini", + "input": "hello", + "stream": True, + "litellm_params": {}, + } + await caching_handler._async_get_cache( + model="gpt-4.1-mini", + original_function=aresponses, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.aresponses.value, + kwargs=kwargs, + ) + + assert caching_handler.preset_cache_key == "responses-stream-cache-key" + mock_cache.get_cache.assert_called_once() + assert mock_cache.get_cache.call_args.kwargs["cache_key"] == ( + "responses-stream-cache-key" + ) + + litellm.cache = original_cache + + +def test_reasoning_summary_events_default_summary_index(): + delta_event = openai_types.ReasoningSummaryTextDeltaEvent( + type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA, + item_id="rs_1", + output_index=0, + delta="abc", + ) + text_done_event = openai_types.ReasoningSummaryTextDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DONE, + item_id="rs_1", + output_index=0, + sequence_number=1, + text="abc", + ) + part_done_event = openai_types.ReasoningSummaryPartDoneEvent( + type=openai_types.ResponsesAPIStreamEvents.REASONING_SUMMARY_PART_DONE, + item_id="rs_1", + output_index=0, + sequence_number=2, + part=openai_types.BaseLiteLLMOpenAIResponseObject( + type="summary_text", + text="abc", + ), + ) + + assert delta_event.summary_index == 0 + assert text_done_event.summary_index == 0 + assert part_done_event.summary_index == 0 From 262782f5b225a0d4bc5308af8818b947c6096894 Mon Sep 17 00:00:00 2001 From: Gowtham Raj Date: Tue, 7 Apr 2026 10:01:29 +0530 Subject: [PATCH 63/75] fix: resolve access group names in _filter_models_by_team_id (#25224) The /v2/model/info endpoint (used by the UI's Models + Endpoints page) was not resolving access group names when filtering models by team. When a team has models: ["Group-A"] where "Group-A" is an access group, _filter_models_by_team_id() passed it as a literal model name to get_model_list(), which found no deployments with that name. This caused the UI to show all models instead of only team-accessible ones. The request-time auth path (model_in_access_group in auth_checks.py) correctly resolves access groups via get_model_access_groups(). This fix applies the same resolution in _filter_models_by_team_id() for both the in-memory router lookup and the database fallback query. Tests added: - test_filter_resolves_access_group_names - test_filter_resolves_mix_of_access_groups_and_literal_names - test_filter_excludes_models_from_other_access_group - test_filter_db_fallback_receives_resolved_model_names --- litellm/proxy/proxy_server.py | 20 +- ...test_filter_models_by_team_access_group.py | 236 ++++++++++++++++++ 2 files changed, 254 insertions(+), 2 deletions(-) create mode 100644 tests/test_litellm/proxy/test_filter_models_by_team_access_group.py diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6cba6a3e96b..0668880f908 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -10324,6 +10324,7 @@ async def _filter_models_by_team_id( # Get models accessible to this team (similar to _add_team_models_to_all_models) team_accessible_model_ids: Set[str] = set() + access_groups = llm_router.get_model_access_groups() if llm_router else {} if ( not team_object.models # empty list = all model access @@ -10347,8 +10348,16 @@ async def _filter_models_by_team_id( if can_add_model: team_accessible_model_ids.add(model_id) else: - # Team has access to specific models + # Team has access to specific models — resolve access group names first + resolved_model_names: Set[str] = set() for model_name in team_object.models: + if model_name in access_groups: + # This is an access group name — expand to member model names + resolved_model_names.update(access_groups[model_name]) + else: + resolved_model_names.add(model_name) + + for model_name in resolved_model_names: _models = ( llm_router.get_model_list(model_name=model_name, team_id=team_id) if llm_router @@ -10368,8 +10377,15 @@ async def _filter_models_by_team_id( and SpecialModelNames.all_proxy_models.value not in team_object.models ): # Team has specific models - check database for those model names + # Resolve access group names to actual model names for the DB query + _resolved_names: list[str] = [] + for _m in team_object.models: + if _m in access_groups: + _resolved_names.extend(access_groups[_m]) + else: + _resolved_names.append(_m) db_models = await prisma_client.db.litellm_proxymodeltable.find_many( - where={"model_name": {"in": team_object.models}} + where={"model_name": {"in": _resolved_names}} ) for db_model in db_models: model_id = db_model.model_id diff --git a/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py b/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py new file mode 100644 index 00000000000..2d8a9f30c1b --- /dev/null +++ b/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py @@ -0,0 +1,236 @@ +""" +Tests for _filter_models_by_team_id resolving access group names. + +Verifies that when a team's `models` field contains an access group name +(e.g., "Group-A"), the filter resolves it to the member model names before +looking up deployments — matching the behavior of the auth path in +auth_checks.py:model_in_access_group(). +""" + +import os +import sys +from unittest.mock import AsyncMock, MagicMock + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.proxy.proxy_server import _filter_models_by_team_id + + +def _make_model(model_name: str, model_id: str, access_groups: list[str] = None): + """Helper to build a model dict matching the router's format.""" + return { + "model_name": model_name, + "litellm_params": {"model": model_name}, + "model_info": { + "id": model_id, + "access_groups": access_groups or [], + }, + } + + +def _make_team(models: list[str], team_id: str = "team_alpha"): + """Helper to build a mock team DB object.""" + mock = MagicMock() + mock.model_dump.return_value = { + "team_id": team_id, + "team_alias": "Team Alpha", + "models": models, + "max_budget": None, + "spend": 0.0, + "blocked": False, + "members_with_roles": [], + "metadata": {}, + } + return mock + + +@pytest.mark.asyncio +async def test_filter_resolves_access_group_names(): + """ + When team.models contains an access group name, _filter_models_by_team_id + should resolve it to the member models and return only those deployments. + """ + # Models on the proxy + gpt4o = _make_model("gpt-4o", "id-1", ["Group-A"]) + gpt5 = _make_model("gpt-5", "id-2", ["Group-A"]) + claude = _make_model("claude-3", "id-3", ["Group-B"]) + + all_models = [gpt4o, gpt5, claude] + + # Router mock + mock_router = MagicMock() + # get_model_access_groups returns {group_name: [model_names]} + mock_router.get_model_access_groups.return_value = { + "Group-A": ["gpt-4o", "gpt-5"], + "Group-B": ["claude-3"], + } + + # get_model_list returns deployments matching a model_name + def fake_get_model_list(model_name=None, team_id=None): + return [m for m in all_models if m["model_name"] == model_name] + + mock_router.get_model_list = MagicMock(side_effect=fake_get_model_list) + + # Team has models: ["Group-A"] — an access group name, not a literal model + team_db = _make_team(models=["Group-A"]) + + # Prisma mock + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + + result = await _filter_models_by_team_id( + all_models=all_models, + team_id="team_alpha", + prisma_client=mock_prisma, + llm_router=mock_router, + ) + + result_ids = {m["model_info"]["id"] for m in result} + # Should include gpt-4o and gpt-5 (Group-A), but NOT claude-3 (Group-B) + assert result_ids == { + "id-1", + "id-2", + }, f"Expected Group-A models only, got {result_ids}" + + # Verify DB fallback query received resolved model names, not access group name + call_kwargs = mock_prisma.db.litellm_proxymodeltable.find_many.call_args[1] + assert set(call_kwargs["where"]["model_name"]["in"]) == { + "gpt-4o", + "gpt-5", + }, "find_many should receive resolved model names, not the access group name" + + +@pytest.mark.asyncio +async def test_filter_resolves_mix_of_access_groups_and_literal_names(): + """ + When team.models contains both an access group name and a literal model name, + both should be resolved correctly. + """ + gpt4o = _make_model("gpt-4o", "id-1", ["Group-A"]) + gpt5 = _make_model("gpt-5", "id-2", ["Group-A"]) + claude = _make_model("claude-3", "id-3", ["Group-B"]) + mistral = _make_model("mistral-large", "id-4", []) # no access group + + all_models = [gpt4o, gpt5, claude, mistral] + + mock_router = MagicMock() + mock_router.get_model_access_groups.return_value = { + "Group-A": ["gpt-4o", "gpt-5"], + "Group-B": ["claude-3"], + } + + def fake_get_model_list(model_name=None, team_id=None): + return [m for m in all_models if m["model_name"] == model_name] + + mock_router.get_model_list = MagicMock(side_effect=fake_get_model_list) + + # Team has access to Group-A (access group) + mistral-large (literal name) + team_db = _make_team(models=["Group-A", "mistral-large"]) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + + result = await _filter_models_by_team_id( + all_models=all_models, + team_id="team_alpha", + prisma_client=mock_prisma, + llm_router=mock_router, + ) + + result_ids = {m["model_info"]["id"] for m in result} + # Group-A models + mistral-large, but NOT claude-3 + assert result_ids == { + "id-1", + "id-2", + "id-4", + }, f"Expected Group-A + mistral-large, got {result_ids}" + + +@pytest.mark.asyncio +async def test_filter_excludes_models_from_other_access_group(): + """ + Models belonging only to a different access group must not appear in results. + """ + gpt4o = _make_model("gpt-4o", "id-1", ["Group-A"]) + claude = _make_model("claude-3", "id-3", ["Group-B"]) + llama = _make_model("llama-4", "id-4", ["Group-B"]) + + all_models = [gpt4o, claude, llama] + + mock_router = MagicMock() + mock_router.get_model_access_groups.return_value = { + "Group-A": ["gpt-4o"], + "Group-B": ["claude-3", "llama-4"], + } + + def fake_get_model_list(model_name=None, team_id=None): + return [m for m in all_models if m["model_name"] == model_name] + + mock_router.get_model_list = MagicMock(side_effect=fake_get_model_list) + + team_db = _make_team(models=["Group-A"]) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + + result = await _filter_models_by_team_id( + all_models=all_models, + team_id="team_alpha", + prisma_client=mock_prisma, + llm_router=mock_router, + ) + + result_names = {m["model_name"] for m in result} + assert "claude-3" not in result_names, "Group-B model should not be accessible" + assert "llama-4" not in result_names, "Group-B model should not be accessible" + assert "gpt-4o" in result_names, "Group-A model should be accessible" + + +@pytest.mark.asyncio +async def test_filter_db_fallback_receives_resolved_model_names(): + """ + When get_model_list returns no results (forcing the DB fallback path), + the DB query should receive resolved model names, not the raw access group name. + """ + gpt4o = _make_model("gpt-4o", "id-1", ["Group-A"]) + all_models = [gpt4o] + + mock_router = MagicMock() + mock_router.get_model_access_groups.return_value = { + "Group-A": ["gpt-4o", "gpt-5"], + } + # get_model_list returns nothing — forces reliance on the DB fallback + mock_router.get_model_list = MagicMock(return_value=[]) + + team_db = _make_team(models=["Group-A"]) + + # DB returns a model that the router didn't find + mock_db_model = MagicMock() + mock_db_model.model_id = "id-db-1" + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock( + return_value=[mock_db_model] + ) + + result = await _filter_models_by_team_id( + all_models=all_models, + team_id="team_alpha", + prisma_client=mock_prisma, + llm_router=mock_router, + ) + + # Verify DB query received resolved names, not "Group-A" + call_kwargs = mock_prisma.db.litellm_proxymodeltable.find_many.call_args[1] + queried_names = set(call_kwargs["where"]["model_name"]["in"]) + assert queried_names == { + "gpt-4o", + "gpt-5", + }, f"DB query should receive resolved model names, got {queried_names}" + assert "Group-A" not in queried_names, "Raw access group name should not be in DB query" From f49c91ea922f6ce7ee849a0571e9134ff167be40 Mon Sep 17 00:00:00 2001 From: shubham-arora-clear Date: Tue, 7 Apr 2026 10:03:41 +0530 Subject: [PATCH 64/75] fix(bedrock): handle document content blocks in Converse API message conversion (#24644) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(bedrock): handle document content blocks in Converse API message conversion Document content blocks (used for PDF support) were silently dropped during message conversion for Bedrock's Converse API. The content block processing loop only handled text, image_url, and file types — document blocks were skipped without warning, causing the model to respond as if no document was provided. Adds document block handling in three locations: - Sync user message processing (_bedrock_converse_messages_pt) - Async user message processing (_bedrock_converse_messages_pt_async) - Tool result conversion (_convert_to_bedrock_tool_call_result) Fixes #24641 Co-Authored-By: Claude Opus 4.6 (1M context) * fix: use _validate_format for proper MIME type to Bedrock format mapping Address Greptile review: naive media_type.split("/")[1] produced invalid Bedrock format names for complex MIME types (e.g. OOXML → docx, text/plain → txt, text/markdown → md). Now reuses BedrockImageProcessor._validate_format which handles all MIME types correctly via mimetypes + fallback. Also fixes test assertions to expect correct Bedrock format values and adds text/plain and text/markdown test cases. Co-Authored-By: Claude Opus 4.6 (1M context) * fix: reject non-base64 document sources with a clear error URL-type document sources (e.g. {"type": "url", "url": "..."}) would crash with an opaque KeyError on missing 'media_type'. Guard at the top of _process_document_message and raise a clear ValueError since Bedrock Converse only supports base64-encoded document sources. Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Claude Opus 4.6 (1M context) --- .../prompt_templates/factory.py | 59 ++++- ...llm_core_utils_prompt_templates_factory.py | 238 +++++++----------- 2 files changed, 149 insertions(+), 148 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 3a83162fb20..4ef08b0a01f 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -4042,9 +4042,13 @@ def _convert_to_bedrock_tool_call_result( tool_result_content_blocks.append( BedrockToolResultContentBlock(image=_block["image"]) ) - elif "document" in _block: - tool_result_content_blocks.append( - BedrockToolResultContentBlock(document=_block["document"]) + elif content["type"] == "document": + _doc_block = BedrockConverseMessagesProcessor._process_document_message( + content + ) + if "document" in _doc_block: + tool_result_content_blocks.append( + BedrockToolResultContentBlock(document=_doc_block["document"]) ) else: verbose_logger.warning( @@ -4090,7 +4094,6 @@ def _convert_to_bedrock_tool_call_result( "%s for file tool-result block %s; dropping.", list(_file_block.keys()), content, - ) message.get("name", "") id = str(message.get("tool_call_id", str(uuid.uuid4()))) @@ -4582,6 +4585,11 @@ class BedrockConverseMessagesProcessor: message=cast(ChatCompletionFileObject, element) ) _parts.append(_part) + elif element["type"] == "document": + _part = BedrockConverseMessagesProcessor._process_document_message( + element + ) + _parts.append(_part) _cache_point_block = ( litellm.AmazonConverseConfig()._get_cache_point_block( message_block=cast( @@ -4864,6 +4872,44 @@ class BedrockConverseMessagesProcessor: image_url=cast(str, file_id or file_data), format=format ) + @staticmethod + def _process_document_message(element: dict) -> BedrockContentBlock: + """Convert a document content block to a Bedrock DocumentBlock. + + Handles the Anthropic-style document format: + {"type": "document", "source": {"type": "base64", "media_type": "application/pdf", "data": "..."}} + """ + source = element["source"] + source_type = source.get("type") + if source_type != "base64": + raise ValueError( + f"Bedrock Converse only supports base64-encoded document sources, got '{source_type}'. " + "Please convert the document to base64 before sending to Bedrock." + ) + media_type: str = source["media_type"] + data: str = source["data"] + doc_format = BedrockImageProcessor._validate_format( + mime_type=media_type, image_format=media_type.split("/")[1] + ) + + # Deterministic name using the same hashing pattern as _create_bedrock_block + HASH_SAMPLE_BYTES = 64 * 1024 + normalized = "".join(data.split()).encode("utf-8") + sample = normalized[:HASH_SAMPLE_BYTES] + hasher = hashlib.sha256() + hasher.update(sample) + hasher.update(str(len(normalized)).encode("utf-8")) + content_hash = hasher.hexdigest()[:16] + document_name = f"Document_{content_hash}_{doc_format}" + + return BedrockContentBlock( + document=BedrockDocumentBlock( + source=BedrockSourceBlock(bytes=data), + format=doc_format, + name=document_name, + ) + ) + @staticmethod def add_thinking_blocks_to_assistant_content( thinking_blocks: List[BedrockContentBlock], @@ -4961,6 +5007,11 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 ) ) _parts.append(_part) + elif element["type"] == "document": + _part = BedrockConverseMessagesProcessor._process_document_message( + element + ) + _parts.append(_part) _cache_point_block = ( litellm.AmazonConverseConfig()._get_cache_point_block( message_block=cast( diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index d424cd8599f..33bdf169ccd 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -9,8 +9,10 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( BAD_MESSAGE_ERROR_STR, BedrockConverseMessagesProcessor, BedrockImageProcessor, - anthropic_messages_pt, + _bedrock_converse_messages_pt, _convert_to_bedrock_tool_call_invoke, + _convert_to_bedrock_tool_call_result, + anthropic_messages_pt, convert_to_gemini_tool_call_result, ollama_pt, sanitize_messages_for_tool_calling, @@ -2505,157 +2507,105 @@ def test_convert_to_anthropic_tool_result_openai_file_pdf_becomes_document(): ], } - result = convert_to_anthropic_tool_result(message) + result = _convert_to_bedrock_tool_call_result(message) - assert result["type"] == "tool_result" - assert result["tool_use_id"] == "toolu_pdf_1" - content = result["content"] - assert isinstance(content, list) and len(content) == 1 - block = content[0] - assert block["type"] == "document" - assert block["source"]["type"] == "base64" - assert block["source"]["media_type"] == "application/pdf" - assert block["source"]["data"] == pdf_b64 + tool_result = result["toolResult"] + assert len(tool_result["content"]) == 1 + assert "document" in tool_result["content"][0] + assert tool_result["content"][0]["document"]["format"] == "pdf" + assert tool_result["content"][0]["document"]["source"]["bytes"] == "dGVzdA==" -def test_convert_to_anthropic_tool_result_image_url_pdf_data_uri_becomes_document(): - """ - Regression: a PDF sent as an `image_url` data URI on the tool-result path - must translate to an Anthropic document block (not an image block — Anthropic - rejects image blocks whose media_type is a non-image like application/pdf). - """ - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_anthropic_tool_result, - ) +def test_bedrock_converse_messages_pt_document_various_formats(): + """Test that various document media types produce the correct format value.""" + test_cases = [ + ("application/pdf", "pdf"), + ("text/csv", "csv"), + ("text/html", "html"), + ("text/plain", "txt"), + ("text/markdown", "md"), + ( + "application/vnd.openxmlformats-officedocument.wordprocessingml.document", + "docx", + ), + ] - pdf_b64 = "JVBERi0xLjQKJeLjz9MK" - message = { - "tool_call_id": "toolu_pdf_img_1", - "role": "tool", - "name": "fetch_document", - "content": [ + for media_type, expected_format in test_cases: + messages = [ { - "type": "image_url", - "image_url": { - "url": f"data:application/pdf;base64,{pdf_b64}", + "role": "user", + "content": [ + { + "type": "document", + "source": { + "type": "base64", + "media_type": media_type, + "data": "dGVzdA==", + }, + }, + ], + } + ] + + result = _bedrock_converse_messages_pt( + messages, "anthropic.claude-sonnet-4-6", "bedrock" + ) + + doc_block = result[0]["content"][0] + assert doc_block["document"]["format"] == expected_format, ( + f"Expected format '{expected_format}' for media_type '{media_type}', " + f"got '{doc_block['document']['format']}'" + ) + + +def test_bedrock_converse_messages_pt_document_deterministic_name(): + """Test that the same document data always produces the same name.""" + messages = [ + { + "role": "user", + "content": [ + { + "type": "document", + "source": { + "type": "base64", + "media_type": "application/pdf", + "data": "dGVzdA==", + }, }, - }, - ], - } + ], + } + ] - result = convert_to_anthropic_tool_result(message) - - content = result["content"] - assert isinstance(content, list) and len(content) == 1 - block = content[0] - assert block["type"] == "document" - assert block["source"]["media_type"] == "application/pdf" - assert block["source"]["data"] == pdf_b64 - - -def test_convert_to_anthropic_tool_result_image_url_unsupported_mime_stays_image_path(): - """ - An `image_url` data URI whose mime is neither application/pdf nor text/plain - (e.g. application/json) must NOT be routed through the document path. Anthropic - only accepts application/pdf and text/plain as base64 document media_types — - anything else would produce a document block the API rejects. The old - (pre-fix) behavior was to wrap such data as an image block, which also - fails but stays on the image code path; preserve that failure mode rather - than switching to a document path that is equally broken. - """ - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_anthropic_tool_result, + result1 = _bedrock_converse_messages_pt( + messages, "anthropic.claude-sonnet-4-6", "bedrock" + ) + result2 = _bedrock_converse_messages_pt( + messages, "anthropic.claude-sonnet-4-6", "bedrock" ) - message = { - "tool_call_id": "toolu_json_1", - "role": "tool", - "name": "fetch_json", - "content": [ - { - "type": "image_url", - "image_url": { - "url": "data:application/json;base64,eyJrIjoidiJ9", + name1 = result1[0]["content"][0]["document"]["name"] + name2 = result2[0]["content"][0]["document"]["name"] + assert name1 == name2 + + +def test_bedrock_converse_messages_pt_document_rejects_url_source(): + """Test that a URL-type document source raises a clear error instead of KeyError.""" + messages = [ + { + "role": "user", + "content": [ + { + "type": "document", + "source": { + "type": "url", + "url": "https://example.com/doc.pdf", + }, }, - }, - ], - } + ], + } + ] - result = convert_to_anthropic_tool_result(message) - - content = result["content"] - assert isinstance(content, list) and len(content) == 1 - block = content[0] - assert block["type"] == "image", ( - f"unsupported mime {block.get('source', {}).get('media_type')!r} " - f"should not be routed to document path; got {block}" - ) - - -def test_convert_to_anthropic_tool_result_image_url_text_plain_data_uri_becomes_document(): - """ - text/plain is one of the two mimes Anthropic accepts as a base64 document - media_type. Confirm it routes through the document path so tightening the - gate to {application/pdf, text/plain} (not "application/*") covers both. - """ - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_anthropic_tool_result, - ) - - txt_b64 = "aGVsbG8=" # "hello" - message = { - "tool_call_id": "toolu_txt_1", - "role": "tool", - "name": "fetch_text", - "content": [ - { - "type": "image_url", - "image_url": { - "url": f"data:text/plain;base64,{txt_b64}", - }, - }, - ], - } - - result = convert_to_anthropic_tool_result(message) - - content = result["content"] - assert isinstance(content, list) and len(content) == 1 - block = content[0] - assert block["type"] == "document" - assert block["source"]["media_type"] == "text/plain" - assert block["source"]["data"] == txt_b64 - - -def test_convert_to_anthropic_tool_result_image_url_png_still_becomes_image(): - """ - Regression: image_url with a real image mime type must continue to translate - to an Anthropic image block. Locks in existing behavior after the - data-URI-mime-type branching for PDFs. - """ - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_anthropic_tool_result, - ) - - png_b64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR4nGNgYGBgAAAABQABXvMqOgAAAABJRU5ErkJggg==" - message = { - "tool_call_id": "toolu_png_1", - "role": "tool", - "name": "fetch_image", - "content": [ - { - "type": "image_url", - "image_url": { - "url": f"data:image/png;base64,{png_b64}", - }, - }, - ], - } - - result = convert_to_anthropic_tool_result(message) - - content = result["content"] - assert isinstance(content, list) and len(content) == 1 - block = content[0] - assert block["type"] == "image" - assert block["source"]["media_type"] == "image/png" + with pytest.raises(ValueError, match="only supports base64-encoded"): + _bedrock_converse_messages_pt( + messages, "anthropic.claude-sonnet-4-6", "bedrock" + ) From ec38f2b17beef24551339f04b282e1605da0515d Mon Sep 17 00:00:00 2001 From: Baqiao Date: Mon, 6 Apr 2026 21:35:28 -0700 Subject: [PATCH 65/75] feat(xai): add parallel_tool_calls to supported params (#25106) --- litellm/llms/xai/chat/transformation.py | 1 + .../llms/xai/test_xai_chat_transformation.py | 39 +++++++++++++++++++ 2 files changed, 40 insertions(+) create mode 100644 tests/test_litellm/llms/xai/test_xai_chat_transformation.py diff --git a/litellm/llms/xai/chat/transformation.py b/litellm/llms/xai/chat/transformation.py index bfa55105a6c..64b4a545acb 100644 --- a/litellm/llms/xai/chat/transformation.py +++ b/litellm/llms/xai/chat/transformation.py @@ -43,6 +43,7 @@ class XAIChatConfig(OpenAIGPTConfig): "logprobs", "max_tokens", "n", + "parallel_tool_calls", "presence_penalty", "response_format", "seed", diff --git a/tests/test_litellm/llms/xai/test_xai_chat_transformation.py b/tests/test_litellm/llms/xai/test_xai_chat_transformation.py new file mode 100644 index 00000000000..5a236de900e --- /dev/null +++ b/tests/test_litellm/llms/xai/test_xai_chat_transformation.py @@ -0,0 +1,39 @@ +import os +import sys + +sys.path.insert( + 0, os.path.abspath("../../../..") +) # Adds the parent directory to the system path + +from litellm.llms.xai.chat.transformation import XAIChatConfig + + +class TestXAIParallelToolCalls: + """Test suite for XAI parallel tool calls functionality.""" + + def test_get_supported_openai_params_includes_parallel_tool_calls(self): + """Test that parallel_tool_calls is in supported parameters.""" + config = XAIChatConfig() + supported_params = config.get_supported_openai_params( + "xai/grok-4.20" + ) + assert "parallel_tool_calls" in supported_params + + def test_transform_request_preserves_parallel_tool_calls(self): + """Test that transform_request preserves parallel_tool_calls parameter.""" + config = XAIChatConfig() + + messages = [{"role": "user", "content": "What's the weather like?"}] + optional_params = {"parallel_tool_calls": True} + + result = config.transform_request( + model="xai/grok-4.20", + messages=messages, + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + assert result.get("parallel_tool_calls") is True + assert len(result["messages"]) == 1 + assert result["messages"][0]["role"] == "user" From d2346b525912a5c006c70bbf32b2cceeaba0c8bb Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 1 May 2026 12:09:19 +0530 Subject: [PATCH 66/75] Fix prompt factory --- .../litellm_core_utils/prompt_templates/factory.py | 11 ++++------- 1 file changed, 4 insertions(+), 7 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 4ef08b0a01f..ba840bc3d89 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -4042,13 +4042,9 @@ def _convert_to_bedrock_tool_call_result( tool_result_content_blocks.append( BedrockToolResultContentBlock(image=_block["image"]) ) - elif content["type"] == "document": - _doc_block = BedrockConverseMessagesProcessor._process_document_message( - content - ) - if "document" in _doc_block: - tool_result_content_blocks.append( - BedrockToolResultContentBlock(document=_doc_block["document"]) + elif "document" in _block: + tool_result_content_blocks.append( + BedrockToolResultContentBlock(document=_block["document"]) ) else: verbose_logger.warning( @@ -4094,6 +4090,7 @@ def _convert_to_bedrock_tool_call_result( "%s for file tool-result block %s; dropping.", list(_file_block.keys()), content, + ) message.get("name", "") id = str(message.get("tool_call_id", str(uuid.uuid4()))) From 94162916a81e37eaaa97dc21614d6b038e54f9a5 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 1 May 2026 12:14:53 +0530 Subject: [PATCH 67/75] Fix lint and ruff --- litellm/proxy/proxy_server.py | 186 +++++++++++++----------- litellm/responses/streaming_iterator.py | 1 + 2 files changed, 101 insertions(+), 86 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0668880f908..1c286b7f2e4 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -10285,6 +10285,101 @@ def _paginate_models_response( } +def _team_models_resolve_to_names( + team_models: List[str], access_groups: Dict[str, Any] +) -> List[str]: + """Expand team model entries (including access group names) to concrete model names.""" + resolved: List[str] = [] + for name in team_models: + if name in access_groups: + resolved.extend(access_groups[name]) + else: + resolved.append(name) + return resolved + + +async def _load_team_object_for_model_filter( + team_id: str, prisma_client: PrismaClient +) -> Optional[LiteLLM_TeamTable]: + """Load team row from DB; returns None if missing or on error.""" + try: + team_db_object = await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id} + ) + if team_db_object is None: + verbose_proxy_logger.warning(f"Team {team_id} not found in database") + return None + return LiteLLM_TeamTable(**team_db_object.model_dump()) + except Exception as e: + verbose_proxy_logger.exception(f"Error fetching team {team_id}: {str(e)}") + return None + + +async def _gather_team_accessible_model_ids( + team_object: LiteLLM_TeamTable, + team_id: str, + prisma_client: PrismaClient, + llm_router: Router, +) -> Set[str]: + """Collect model IDs the team can use from router config and DB.""" + team_accessible_model_ids: Set[str] = set() + access_groups = llm_router.get_model_access_groups() if llm_router else {} + + if ( + not team_object.models + or SpecialModelNames.all_proxy_models.value in team_object.models + ): + model_list = llm_router.get_model_list() if llm_router else [] + if model_list is not None: + for model in model_list: + model_id = model.get("model_info", {}).get("id", None) + if model_id is None: + continue + team_model_id = model.get("model_info", {}).get("team_id", None) + if team_model_id is None or team_model_id == team_id: + team_accessible_model_ids.add(model_id) + else: + resolved_model_names: Set[str] = set() + for model_name in team_object.models: + if model_name in access_groups: + resolved_model_names.update(access_groups[model_name]) + else: + resolved_model_names.add(model_name) + + for model_name in resolved_model_names: + _models = ( + llm_router.get_model_list(model_name=model_name, team_id=team_id) + if llm_router + else [] + ) + if _models is not None: + for model in _models: + model_id = model.get("model_info", {}).get("id", None) + if model_id is not None: + team_accessible_model_ids.add(model_id) + + try: + if ( + team_object.models + and SpecialModelNames.all_proxy_models.value not in team_object.models + ): + _resolved_names = _team_models_resolve_to_names( + team_object.models, access_groups + ) + db_models = await prisma_client.db.litellm_proxymodeltable.find_many( + where={"model_name": {"in": _resolved_names}} + ) + for db_model in db_models: + if db_model.model_id: + team_accessible_model_ids.add(db_model.model_id) + except Exception as e: + verbose_proxy_logger.debug( + f"Error querying database models for team {team_id}: {str(e)}" + ) + + return team_accessible_model_ids + + async def _filter_models_by_team_id( all_models: List[Dict[str, Any]], team_id: str, @@ -10307,94 +10402,13 @@ async def _filter_models_by_team_id( Returns: Filtered list of models """ - # Get team from database - try: - team_db_object = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id} - ) - if team_db_object is None: - verbose_proxy_logger.warning(f"Team {team_id} not found in database") - # If team doesn't exist, return empty list - return [] - - team_object = LiteLLM_TeamTable(**team_db_object.model_dump()) - except Exception as e: - verbose_proxy_logger.exception(f"Error fetching team {team_id}: {str(e)}") + team_object = await _load_team_object_for_model_filter(team_id, prisma_client) + if team_object is None: return [] - # Get models accessible to this team (similar to _add_team_models_to_all_models) - team_accessible_model_ids: Set[str] = set() - access_groups = llm_router.get_model_access_groups() if llm_router else {} - - if ( - not team_object.models # empty list = all model access - or SpecialModelNames.all_proxy_models.value in team_object.models - ): - # Team has access to all models - model_list = llm_router.get_model_list() if llm_router else [] - if model_list is not None: - for model in model_list: - model_id = model.get("model_info", {}).get("id", None) - if model_id is None: - continue - # if team model id set, check if team id matches - team_model_id = model.get("model_info", {}).get("team_id", None) - can_add_model = False - if team_model_id is None: - can_add_model = True - elif team_model_id == team_id: - can_add_model = True - - if can_add_model: - team_accessible_model_ids.add(model_id) - else: - # Team has access to specific models — resolve access group names first - resolved_model_names: Set[str] = set() - for model_name in team_object.models: - if model_name in access_groups: - # This is an access group name — expand to member model names - resolved_model_names.update(access_groups[model_name]) - else: - resolved_model_names.add(model_name) - - for model_name in resolved_model_names: - _models = ( - llm_router.get_model_list(model_name=model_name, team_id=team_id) - if llm_router - else [] - ) - if _models is not None: - for model in _models: - model_id = model.get("model_info", {}).get("id", None) - if model_id is not None: - team_accessible_model_ids.add(model_id) - - # Also search database for models accessible to this team - # This complements the config search done above - try: - if ( - team_object.models - and SpecialModelNames.all_proxy_models.value not in team_object.models - ): - # Team has specific models - check database for those model names - # Resolve access group names to actual model names for the DB query - _resolved_names: list[str] = [] - for _m in team_object.models: - if _m in access_groups: - _resolved_names.extend(access_groups[_m]) - else: - _resolved_names.append(_m) - db_models = await prisma_client.db.litellm_proxymodeltable.find_many( - where={"model_name": {"in": _resolved_names}} - ) - for db_model in db_models: - model_id = db_model.model_id - if model_id: - team_accessible_model_ids.add(model_id) - except Exception as e: - verbose_proxy_logger.debug( - f"Error querying database models for team {team_id}: {str(e)}" - ) + team_accessible_model_ids = await _gather_team_accessible_model_ids( + team_object, team_id, prisma_client, llm_router + ) # Filter models based on direct_access or access_via_team_ids # Models are already enriched with these fields before this function is called diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 1902a347215..da8da1b486f 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -25,6 +25,7 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import ( from litellm.litellm_core_utils.thread_pool_executor import executor from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.responses.utils import ResponsesAPIRequestUtils +from litellm.types.llms.openai import ResponsesAPIStreamEvents from litellm.types.utils import CallTypes from litellm.utils import CustomStreamWrapper, async_post_call_success_deployment_hook From 900ef454c83097452884000ea3185520a07b4794 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 1 May 2026 15:35:30 +0530 Subject: [PATCH 68/75] test: fix Bedrock PDF tool-result bytes assertion in factory test The test supplies a minimal PDF base64 payload but expected the wrong constant (base64 for "test"). Assert against the same pdf_b64 value and drop the unused import. Made-with: Cursor --- .../test_litellm_core_utils_prompt_templates_factory.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 33bdf169ccd..27a3ddb553d 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -2487,10 +2487,6 @@ def test_convert_to_anthropic_tool_result_openai_file_pdf_becomes_document(): inside the tool_result content. Reuses anthropic_process_openai_file_message, which already handles this for user messages. """ - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_anthropic_tool_result, - ) - pdf_b64 = "JVBERi0xLjQKJeLjz9MK" message = { "tool_call_id": "toolu_pdf_1", @@ -2513,7 +2509,7 @@ def test_convert_to_anthropic_tool_result_openai_file_pdf_becomes_document(): assert len(tool_result["content"]) == 1 assert "document" in tool_result["content"][0] assert tool_result["content"][0]["document"]["format"] == "pdf" - assert tool_result["content"][0]["document"]["source"]["bytes"] == "dGVzdA==" + assert tool_result["content"][0]["document"]["source"]["bytes"] == pdf_b64 def test_bedrock_converse_messages_pt_document_various_formats(): From 15288f3ae7d57177365ca37c0459532060f0fe23 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 1 May 2026 15:41:17 +0530 Subject: [PATCH 69/75] test: include response key in response.completed chunk for ID hook test _base_process_chunk only encodes response IDs when parsed_chunk contains a top-level "response" key. Align test_process_chunk_completed_response_ updates_id_and_usage_cost with that contract and test_base_responses_api_streaming_iterator. Made-with: Cursor --- tests/llm_responses_api_testing/test_responses_hooks.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/tests/llm_responses_api_testing/test_responses_hooks.py b/tests/llm_responses_api_testing/test_responses_hooks.py index 8933ec8f3c5..3799a0b9121 100644 --- a/tests/llm_responses_api_testing/test_responses_hooks.py +++ b/tests/llm_responses_api_testing/test_responses_hooks.py @@ -384,7 +384,13 @@ def test_process_chunk_completed_response_updates_id_and_usage_cost(monkeypatch) ) try: - event = iterator._process_chunk(json.dumps({"type": "response.completed"})) + # Chunk must include a top-level "response" key so BaseResponsesAPIStreamingIterator + # runs _update_responses_api_response_id_with_model_id (see streaming_iterator.py). + event = iterator._process_chunk( + json.dumps( + {"type": "response.completed", "response": {"id": "resp_live"}} + ) + ) finally: litellm.include_cost_in_streaming_usage = original_include_cost From e0398cade78ad10cb431b1c9a429d43b7133223f Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 1 May 2026 17:03:32 +0530 Subject: [PATCH 70/75] fix(caching): defer streaming cache-hit callbacks for all stream=True Success handlers already run when CustomStreamWrapper or CachedResponsesAPIStreamingIterator finishes replay. Logging at cache-hit time for acompletion/completion streaming duplicated spend and callbacks. Align tests with deferred behavior. Made-with: Cursor --- litellm/caching/caching_handler.py | 31 ++++++++------------- tests/local_testing/test_caching_handler.py | 24 +++++----------- 2 files changed, 19 insertions(+), 36 deletions(-) diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 02207a207ed..3cf1d911d7f 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -85,19 +85,18 @@ class CachingHandlerResponse(BaseModel): in_memory_cache_obj = InMemoryCache() -_RESPONSES_STREAMING_CALLBACK_CALL_TYPES = { - CallTypes.aresponses.value, - CallTypes.responses.value, -} -def _should_defer_streaming_cache_hit_callbacks( - *, call_type: str, kwargs: Dict[str, Any] -) -> bool: - return ( - kwargs.get("stream", False) is True - and call_type in _RESPONSES_STREAMING_CALLBACK_CALL_TYPES - ) +def _should_defer_streaming_cache_hit_callbacks(*, kwargs: Dict[str, Any]) -> bool: + """ + When stream=True, do not run success callbacks at cache-hit time. + + Cached chat/text completion replay uses CustomStreamWrapper; cached Responses + replay uses CachedResponsesAPIStreamingIterator. Both invoke logging success + handlers when the stream finishes; firing them here too would double-count + spend and callback records. + """ + return kwargs.get("stream", False) is True class LLMCachingHandler: @@ -220,10 +219,7 @@ class LLMCachingHandler: custom_llm_provider=kwargs.get("custom_llm_provider", None), args=args, ) - if not _should_defer_streaming_cache_hit_callbacks( - call_type=call_type, - kwargs=kwargs, - ): + if not _should_defer_streaming_cache_hit_callbacks(kwargs=kwargs): # LOG SUCCESS self._async_log_cache_hit_on_callbacks( logging_obj=logging_obj, @@ -343,10 +339,7 @@ class LLMCachingHandler: is_async=False, ) - if not _should_defer_streaming_cache_hit_callbacks( - call_type=call_type, - kwargs=kwargs, - ): + if not _should_defer_streaming_cache_hit_callbacks(kwargs=kwargs): logging_obj.handle_sync_success_callbacks_for_async_calls( result=cached_result, start_time=start_time, diff --git a/tests/local_testing/test_caching_handler.py b/tests/local_testing/test_caching_handler.py index 0ce9938e710..2b6712cbaa3 100644 --- a/tests/local_testing/test_caching_handler.py +++ b/tests/local_testing/test_caching_handler.py @@ -891,7 +891,7 @@ def test_sync_get_cache_does_not_eagerly_log_streaming_responses_hits(): logging_obj.handle_sync_success_callbacks_for_async_calls.assert_not_called() -def test_sync_get_cache_still_eagerly_logs_streaming_completion_hits(): +def test_sync_get_cache_defers_streaming_completion_hit_callbacks(): litellm.set_verbose = True setup_cache() caching_handler = LLMCachingHandler( @@ -930,42 +930,32 @@ def test_sync_get_cache_still_eagerly_logs_streaming_completion_hits(): ) assert cached_response.cached_result is not None - logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once() + logging_obj.handle_sync_success_callbacks_for_async_calls.assert_not_called() -def test_should_defer_streaming_cache_hit_callbacks_only_for_responses_streams(): +def test_should_defer_streaming_cache_hit_callbacks_for_any_streaming_request(): assert ( _should_defer_streaming_cache_hit_callbacks( - call_type=CallTypes.responses.value, kwargs={"stream": True}, ) is True ) assert ( _should_defer_streaming_cache_hit_callbacks( - call_type=CallTypes.aresponses.value, - kwargs={"stream": True}, - ) - is True - ) - assert ( - _should_defer_streaming_cache_hit_callbacks( - call_type=CallTypes.completion.value, - kwargs={"stream": True}, + kwargs={"stream": False}, ) is False ) assert ( _should_defer_streaming_cache_hit_callbacks( - call_type=CallTypes.responses.value, - kwargs={"stream": False}, + kwargs={}, ) is False ) @pytest.mark.asyncio -async def test_async_get_cache_still_eagerly_logs_streaming_completion_hits(): +async def test_async_get_cache_defers_streaming_completion_hit_callbacks(): litellm.set_verbose = True setup_cache() caching_handler = LLMCachingHandler( @@ -1009,7 +999,7 @@ async def test_async_get_cache_still_eagerly_logs_streaming_completion_hits(): assert cached_response is not None assert cached_response.cached_result is not None - caching_handler._async_log_cache_hit_on_callbacks.assert_called_once() + caching_handler._async_log_cache_hit_on_callbacks.assert_not_called() def test_convert_cached_streaming_responses_result_to_iterator(): From 475220ecc65236236b0c3c0550845d6f4adf4186 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 1 May 2026 17:46:20 +0530 Subject: [PATCH 71/75] Fix greptile review --- litellm/proxy/container_endpoints/handler_factory.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/litellm/proxy/container_endpoints/handler_factory.py b/litellm/proxy/container_endpoints/handler_factory.py index 33df72ceaeb..794051e90f8 100644 --- a/litellm/proxy/container_endpoints/handler_factory.py +++ b/litellm/proxy/container_endpoints/handler_factory.py @@ -228,6 +228,12 @@ async def _process_binary_request( elif ".pdf" in file_id_lower: content_type = "application/pdf" + if not isinstance(content, bytes): + raise TypeError( + "aretrieve_container_file_content expected bytes, got " + f"{type(content).__name__}" + ) + return Response( content=content, headers=dict(fastapi_response.headers), From 19813527fa6b8f2d6d7d869074fd933bce42b1cd Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 20 Apr 2026 10:18:34 +0530 Subject: [PATCH 72/75] feat(vertex_ai): Model Garden OpenAPI for publisher model ids - Route publisher/model ids (e.g. xai/grok) to .../endpoints/openapi; keep model in JSON body - Add model_prices keys for vertex_ai/openai/xai/grok-* - Document xAI Grok on vertex_partner (aligned with GPT-OSS) - Add tests for create_vertex_url and body-model heuristic Made-with: Cursor --- .../vertex_ai/vertex_model_garden/main.py | 23 ++++++- model_prices_and_context_window.json | 66 +++++++++++++++++++ .../test_vertex_model_garden_openapi.py | 41 ++++++++++++ 3 files changed, 128 insertions(+), 2 deletions(-) create mode 100644 tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py diff --git a/litellm/llms/vertex_ai/vertex_model_garden/main.py b/litellm/llms/vertex_ai/vertex_model_garden/main.py index c37bb449ecf..7240d9dce57 100644 --- a/litellm/llms/vertex_ai/vertex_model_garden/main.py +++ b/litellm/llms/vertex_ai/vertex_model_garden/main.py @@ -27,6 +27,17 @@ from ..common_utils import VertexAIError, get_vertex_base_model_name from ..vertex_llm_base import VertexBase +def _vertex_model_garden_model_id_in_json_body(model: str) -> bool: + """ + Vertex catalog / publisher models are addressed as publisher/model (e.g. + xai/grok-4.1-fast-reasoning) on the shared OpenAPI URL, with the id in the JSON body. + + Deployed Model Garden endpoints are typically a single segment (often numeric) + and use .../endpoints/{ENDPOINT_ID}/chat/completions with an empty model field. + """ + return "/" in model + + def create_vertex_url( vertex_location: str, vertex_project: str, @@ -34,8 +45,13 @@ def create_vertex_url( model: str, api_base: Optional[str] = None, ) -> str: - """Return the base url for the vertex garden models""" + """Return the api base for vertex model garden (without /chat/completions).""" base_url = get_vertex_base_url(vertex_location) + if _vertex_model_garden_model_id_in_json_body(model): + return ( + f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}" + "/endpoints/openapi" + ) return f"{base_url}/v1beta1/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}" @@ -129,7 +145,10 @@ class VertexAIModelGardenModels(VertexBase): vertex_location=vertex_location or "us-central1", vertex_api_version="v1beta1", ) - model = "" + # Publisher/catalog models: model id must be sent in the JSON body (OpenAPI route). + # Single-segment endpoint ids: model is encoded in the URL path; body model stays empty. + if not _vertex_model_garden_model_id_in_json_body(model): + model = "" return openai_like_chat_completions.completion( model=model, messages=messages, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index bbe13442d63..a3047118fe3 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -33391,6 +33391,72 @@ "source": "https://console.cloud.google.com/vertex-ai/publishers/openai/model-garden/gpt-oss-120b-maas", "supports_reasoning": true }, + "vertex_ai/openai/xai/grok-4.1-fast-non-reasoning": { + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "vertex_ai/openai/xai/grok-4.1-fast-reasoning": { + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "vertex_ai/openai/xai/grok-4.20-non-reasoning": { + "cache_read_input_token_cost": 2e-07, + "input_cost_per_token": 2e-06, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 6e-06, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "vertex_ai/openai/xai/grok-4.20-reasoning": { + "cache_read_input_token_cost": 2e-07, + "input_cost_per_token": 2e-06, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 6e-06, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, "vertex_ai/qwen/qwen3-235b-a22b-instruct-2507-maas": { "input_cost_per_token": 2.5e-07, "litellm_provider": "vertex_ai-qwen_models", diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py b/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py new file mode 100644 index 00000000000..91261b63252 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py @@ -0,0 +1,41 @@ +"""Vertex Model Garden: OpenAPI base URL for publisher/model ids vs per-endpoint path.""" + +import pytest + +from litellm.llms.vertex_ai.vertex_model_garden.main import ( + _vertex_model_garden_model_id_in_json_body, + create_vertex_url, +) + + +@pytest.mark.parametrize( + "model,expect_openapi_base", + [ + ("xai/grok-4.1-fast-reasoning", True), + ("openai/foo/bar", True), + ("5464397967697903616", False), + ("gpt-oss-20b-maas", False), + ], +) +def test_create_vertex_url_openapi_vs_deployed_endpoint( + model: str, expect_openapi_base: bool +) -> None: + url = create_vertex_url( + vertex_location="us-central1", + vertex_project="my-project", + stream=False, + model=model, + ) + if expect_openapi_base: + assert "/v1/projects/my-project/locations/us-central1/endpoints/openapi" in url + else: + assert ( + "/v1beta1/projects/my-project/locations/us-central1/endpoints/" + f"{model}" in url + ) + assert "openapi" not in url + + +def test_model_id_in_json_body_heuristic() -> None: + assert _vertex_model_garden_model_id_in_json_body("xai/grok-4.1-fast-reasoning") is True + assert _vertex_model_garden_model_id_in_json_body("5464397967697903616") is False From 71e2b6fcade5fe3388e62cefbbdfc59a22522173 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 1 May 2026 18:04:55 +0530 Subject: [PATCH 73/75] Fix reviews --- litellm/llms/vertex_ai/common_utils.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index fae175612b1..c72160f7d0a 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -97,7 +97,7 @@ def get_vertex_ai_model_route( Determine which handler to use for a Vertex AI model based on the model name. Args: - model: The model name (e.g., "llama3-405b", "gemini-pro", "gemma/gemma-3-12b-it", "openai/gpt-oss-120b") + model: The model name (e.g., "llama3-405b", "gemini-pro", "gemma/gemma-3-12b-it", "xai/grok-4.1-fast-non-reasoning") litellm_params: Optional litellm parameters dict that may contain base_model for routing Returns: @@ -113,7 +113,7 @@ def get_vertex_ai_model_route( >>> get_vertex_ai_model_route("gemma/gemma-3-12b-it") VertexAIModelRoute.GEMMA - >>> get_vertex_ai_model_route("openai/gpt-oss-120b") + >>> get_vertex_ai_model_route("xai/grok-4.1-fast-non-reasoning") VertexAIModelRoute.MODEL_GARDEN >>> get_vertex_ai_model_route("1234567890", {"api_base": "http://10.96.32.8"}) @@ -149,8 +149,11 @@ def get_vertex_ai_model_route( if "gemma/" in model: return VertexAIModelRoute.GEMMA - # Check for model garden openai models - if "openai" in model: + # Check for model garden OpenAI-compatible publisher models. + # Examples: + # - openai/gpt-oss-120b-maas + # - xai/grok-4.1-fast-non-reasoning + if "openai" in model or model.startswith("xai/"): return VertexAIModelRoute.MODEL_GARDEN # Check for gemini models @@ -256,8 +259,8 @@ def get_vertex_base_model_name(model: str) -> str: >>> get_vertex_base_model_name("gemma/gemma-3-12b-it") "gemma-3-12b-it" - >>> get_vertex_base_model_name("openai/gpt-oss-120b") - "gpt-oss-120b" + >>> get_vertex_base_model_name("xai/grok-4.1-fast-non-reasoning") + "grok-4.1-fast-non-reasoning" >>> get_vertex_base_model_name("1234567890") "1234567890" From e656b2a47b2a539811a1747dd8d6416455b782fe Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 1 May 2026 18:07:33 +0530 Subject: [PATCH 74/75] correct model map --- ...odel_prices_and_context_window_backup.json | 66 +++++++++++++++++++ model_prices_and_context_window.json | 8 +-- 2 files changed, 70 insertions(+), 4 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 13a45fd1650..a1e3e42a9c5 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -33337,6 +33337,72 @@ "source": "https://console.cloud.google.com/vertex-ai/publishers/openai/model-garden/gpt-oss-120b-maas", "supports_reasoning": true }, + "vertex_ai/xai/grok-4.1-fast-non-reasoning": { + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "vertex_ai/xai/grok-4.1-fast-reasoning": { + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "vertex_ai/xai/grok-4.20-non-reasoning": { + "cache_read_input_token_cost": 2e-07, + "input_cost_per_token": 2e-06, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 6e-06, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "vertex_ai/xai/grok-4.20-reasoning": { + "cache_read_input_token_cost": 2e-07, + "input_cost_per_token": 2e-06, + "litellm_provider": "vertex_ai", + "max_input_tokens": 2000000, + "max_output_tokens": 2000000, + "max_tokens": 2000000, + "mode": "chat", + "output_cost_per_token": 6e-06, + "source": "https://docs.x.ai/docs/models (Vertex AI Model Garden)", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, "vertex_ai/qwen/qwen3-235b-a22b-instruct-2507-maas": { "input_cost_per_token": 2.5e-07, "litellm_provider": "vertex_ai-qwen_models", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index a3047118fe3..d21a4bf11d6 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -33391,7 +33391,7 @@ "source": "https://console.cloud.google.com/vertex-ai/publishers/openai/model-garden/gpt-oss-120b-maas", "supports_reasoning": true }, - "vertex_ai/openai/xai/grok-4.1-fast-non-reasoning": { + "vertex_ai/xai/grok-4.1-fast-non-reasoning": { "cache_read_input_token_cost": 5e-08, "input_cost_per_token": 2e-07, "litellm_provider": "vertex_ai", @@ -33407,7 +33407,7 @@ "supports_vision": true, "supports_web_search": true }, - "vertex_ai/openai/xai/grok-4.1-fast-reasoning": { + "vertex_ai/xai/grok-4.1-fast-reasoning": { "cache_read_input_token_cost": 5e-08, "input_cost_per_token": 2e-07, "litellm_provider": "vertex_ai", @@ -33424,7 +33424,7 @@ "supports_vision": true, "supports_web_search": true }, - "vertex_ai/openai/xai/grok-4.20-non-reasoning": { + "vertex_ai/xai/grok-4.20-non-reasoning": { "cache_read_input_token_cost": 2e-07, "input_cost_per_token": 2e-06, "litellm_provider": "vertex_ai", @@ -33440,7 +33440,7 @@ "supports_vision": true, "supports_web_search": true }, - "vertex_ai/openai/xai/grok-4.20-reasoning": { + "vertex_ai/xai/grok-4.20-reasoning": { "cache_read_input_token_cost": 2e-07, "input_cost_per_token": 2e-06, "litellm_provider": "vertex_ai", From 8300657af90acff4cae4c5a89a6c27cbae029580 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 1 May 2026 22:56:10 +0530 Subject: [PATCH 75/75] fix(mcp): preserve oauth2 m2m auth for tools routes (#26871) * Fix tool/list M2M creds issue * Fix tool call creds issue * Fix greptile review * Fix lint * Fix lint * Fix lint * Fix lint --- .../mcp_server/mcp_server_manager.py | 69 ++- .../proxy/_experimental/mcp_server/server.py | 48 ++- .../mcp_server/test_mcp_hook_extra_headers.py | 100 +++++ .../mcp_server/test_mcp_server.py | 394 ++++++++++++++++++ 4 files changed, 600 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index f96350500db..9923c3ce4bf 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -169,6 +169,37 @@ def _deserialize_json_dict(data: Any) -> Optional[Dict[str, str]]: class MCPServerManager: _STDIO_ENV_TEMPLATE_PATTERN = re.compile(r"^\$\{(X-[^}]+)\}$") + @staticmethod + def _resolve_oauth2_flow( + *, + auth_type: Optional[MCPAuthType], + oauth2_flow: Optional[str], + token_url: Optional[str], + authorization_url: Optional[str], + client_id: Optional[str], + client_secret: Optional[str], + ) -> Optional[Literal["client_credentials", "authorization_code"]]: + """Infer oauth2_flow for legacy records that omit the field. + + DB rows created before oauth2_flow support may have OAuth2 client + credentials + token_url but a null oauth2_flow. Treat these as M2M, + unless authorization_url is present (interactive OAuth). + """ + if oauth2_flow in ("client_credentials", "authorization_code"): + return cast( + Literal["client_credentials", "authorization_code"], oauth2_flow + ) + if oauth2_flow: + # Ignore unknown/untyped values and continue legacy inference. + return None + if auth_type != MCPAuth.oauth2: + return None + if authorization_url: + return None + if token_url and client_id and client_secret: + return "client_credentials" + return None + def __init__(self): self.registry: Dict[str, MCPServer] = {} self.config_mcp_servers: Dict[str, MCPServer] = {} @@ -342,7 +373,14 @@ class MCPServerManager: # oauth specific fields client_id=server_config.get("client_id", None), client_secret=server_config.get("client_secret", None), - oauth2_flow=server_config.get("oauth2_flow", None), + oauth2_flow=self._resolve_oauth2_flow( + auth_type=auth_type, + oauth2_flow=server_config.get("oauth2_flow", None), + token_url=resolved_token_url, + authorization_url=resolved_authorization_url, + client_id=server_config.get("client_id", None), + client_secret=server_config.get("client_secret", None), + ), scopes=resolved_scopes, authorization_url=resolved_authorization_url, token_url=resolved_token_url, @@ -679,7 +717,17 @@ class MCPServerManager: client_id=client_id_value or getattr(mcp_server, "client_id", None), client_secret=client_secret_value or getattr(mcp_server, "client_secret", None), - oauth2_flow=getattr(mcp_server, "oauth2_flow", None), + oauth2_flow=self._resolve_oauth2_flow( + auth_type=auth_type, + oauth2_flow=getattr(mcp_server, "oauth2_flow", None), + token_url=mcp_server.token_url + or getattr(mcp_oauth_metadata, "token_url", None), + authorization_url=mcp_server.authorization_url + or getattr(mcp_oauth_metadata, "authorization_url", None), + client_id=client_id_value or getattr(mcp_server, "client_id", None), + client_secret=client_secret_value + or getattr(mcp_server, "client_secret", None), + ), scopes=resolved_scopes, authorization_url=mcp_server.authorization_url or getattr(mcp_oauth_metadata, "authorization_url", None), @@ -2426,7 +2474,7 @@ class MCPServerManager: ) ) - async def _call_regular_mcp_tool( + async def _call_regular_mcp_tool( # noqa: PLR0915 self, mcp_server: MCPServer, original_tool_name: str, @@ -2489,7 +2537,11 @@ class MCPServerManager: # oauth2 headers extra_headers: Optional[Dict[str, str]] = None if mcp_server.auth_type == MCPAuth.oauth2: - extra_headers = oauth2_headers + if mcp_server.has_client_credentials: + # For M2M OAuth servers, Authorization must come from token fetch. + extra_headers = None + else: + extra_headers = oauth2_headers if mcp_server.extra_headers and raw_headers: if extra_headers is None: @@ -2501,6 +2553,11 @@ class MCPServerManager: for header in mcp_server.extra_headers: if not isinstance(header, str): continue + if ( + mcp_server.has_client_credentials + and header.lower() == "authorization" + ): + continue header_value = normalized_raw_headers.get(header.lower()) if header_value is None: continue @@ -2536,6 +2593,10 @@ class MCPServerManager: ) extra_headers.update(hook_extra_headers) + # Reset to None if no headers were actually added + if extra_headers is not None and len(extra_headers) == 0: + extra_headers = None + stdio_env = self._build_stdio_env(mcp_server, raw_headers) client = await self._create_mcp_client( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index ae6055217b8..abb4b5cfa6f 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -153,6 +153,7 @@ if MCP_AVAILABLE: MCPAuthenticatedUser, ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, global_mcp_server_manager, ) from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( @@ -900,6 +901,20 @@ if MCP_AVAILABLE: allowed_mcp_server_id ) if mcp_server is not None: + # Apply oauth2_flow resolution for legacy DB rows where it may be NULL + resolved_flow = MCPServerManager._resolve_oauth2_flow( + auth_type=mcp_server.auth_type, + oauth2_flow=mcp_server.oauth2_flow, + token_url=mcp_server.token_url, + authorization_url=mcp_server.authorization_url, + client_id=mcp_server.client_id, + client_secret=mcp_server.client_secret, + ) + if resolved_flow and resolved_flow != mcp_server.oauth2_flow: + # Create a new instance with the resolved flow for this request + mcp_server = mcp_server.model_copy( + update={"oauth2_flow": resolved_flow} + ) allowed_mcp_servers.append(mcp_server) if mcp_servers is not None: @@ -1100,8 +1115,13 @@ if MCP_AVAILABLE: extra_headers: Optional[Dict[str, str]] = None if server.auth_type == MCPAuth.oauth2: - # Copy to avoid mutating the original dict (important for parallel fetching) - extra_headers = oauth2_headers.copy() if oauth2_headers else None + # For OAuth2 M2M servers, upstream Authorization must come from + # client_credentials token fetch, never from caller headers. + if server.has_client_credentials: + extra_headers = None + else: + # Copy to avoid mutating the original dict (important for parallel fetching) + extra_headers = oauth2_headers.copy() if oauth2_headers else None if server.extra_headers and raw_headers: if extra_headers is None: @@ -1114,11 +1134,17 @@ if MCP_AVAILABLE: for header in server.extra_headers: if not isinstance(header, str): continue + if server.has_client_credentials and header.lower() == "authorization": + continue header_value = normalized_raw_headers.get(header.lower()) if header_value is None: continue extra_headers[header] = header_value + # Reset to None if no headers were actually added + if extra_headers is not None and len(extra_headers) == 0: + extra_headers = None + if server_auth_header is None: server_auth_header = mcp_auth_header @@ -1377,11 +1403,19 @@ if MCP_AVAILABLE: spend_meta["per_server_tool_counts"] = per_server_tool_counts end_time = datetime.now() - await litellm_logging_obj.async_success_handler( - result=all_tools, - start_time=list_tools_start_time, - end_time=end_time, - ) + try: + await litellm_logging_obj.async_success_handler( + result=all_tools, + start_time=list_tools_start_time, + end_time=end_time, + ) + except Exception as log_exc: + # list_tools responses must not be dropped due to non-blocking + # observability/serialization failures. + verbose_logger.warning( + "MCP list_tools success logging failed (continuing): %s", + log_exc, + ) verbose_logger.info( f"Successfully fetched {len(all_tools)} tools total from all MCP servers" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index 84c556b8ddc..649a08e8744 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -673,6 +673,106 @@ class TestHookHeaderMergePriority: assert headers["X-OAuth"] == "yes" assert headers["X-Trace-Id"] == "trace-123" + @pytest.mark.asyncio + async def test_m2m_oauth2_does_not_forward_litellm_caller_authorization(self): + """M2M must not put caller Bearer (LiteLLM API key) into extra_headers (#23652).""" + manager = MCPServerManager() + server = MCPServer( + server_id="test-id", + name="Test Server", + server_name="test_server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://auth.example.com/token", + ) + + captured_extra_headers: Dict[str, Any] = {} + + async def fake_create_mcp_client( + server, mcp_auth_header=None, extra_headers=None, stdio_env=None + ): + captured_extra_headers["value"] = extra_headers + mock_client = MagicMock() + mock_client.call_tool = AsyncMock(return_value=MagicMock()) + return mock_client + + with patch.object( + manager, "_create_mcp_client", side_effect=fake_create_mcp_client + ): + with patch.object(manager, "_build_stdio_env", return_value=None): + try: + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="test_tool", + arguments={"key": "val"}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers={"Authorization": "Bearer sk-1234"}, + raw_headers={"authorization": "Bearer sk-1234"}, + proxy_logging_obj=None, + hook_extra_headers=None, + ) + except Exception: + pass + + assert captured_extra_headers.get("value") is None + + @pytest.mark.asyncio + async def test_m2m_oauth2_skips_authorization_in_configured_extra_headers(self): + """M2M must not take Authorization from raw_headers even if extra_headers lists it.""" + manager = MCPServerManager() + server = MCPServer( + server_id="test-id", + name="Test Server", + server_name="test_server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://auth.example.com/token", + extra_headers=["Authorization", "X-Custom"], + ) + + captured_extra_headers: Dict[str, Any] = {} + + async def fake_create_mcp_client( + server, mcp_auth_header=None, extra_headers=None, stdio_env=None + ): + captured_extra_headers["value"] = extra_headers + mock_client = MagicMock() + mock_client.call_tool = AsyncMock(return_value=MagicMock()) + return mock_client + + with patch.object( + manager, "_create_mcp_client", side_effect=fake_create_mcp_client + ): + with patch.object(manager, "_build_stdio_env", return_value=None): + try: + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="test_tool", + arguments={"key": "val"}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers={"Authorization": "Bearer sk-1234"}, + raw_headers={ + "authorization": "Bearer sk-1234", + "x-custom": "from-client", + }, + proxy_logging_obj=None, + hook_extra_headers=None, + ) + except Exception: + pass + + headers = captured_extra_headers.get("value") or {} + assert "Authorization" not in headers + assert headers.get("X-Custom") == "from-client" + class TestUserAPIKeyAuthJwtClaims: """Tests that UserAPIKeyAuth correctly carries jwt_claims.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 9df6408b0d7..06f95159c08 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -17,6 +17,7 @@ from litellm.proxy._types import ( MCPTransport, UserAPIKeyAuth, ) +from litellm.types.mcp import MCPAuth from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -135,6 +136,152 @@ def test_prepare_mcp_server_headers_case_insensitive_extra_headers(): assert extra_headers == {"Authorization": "Bearer token"} +def test_prepare_mcp_server_headers_oauth2_m2m_omits_litellm_caller_authorization(): + """M2M OAuth must not put caller Bearer (LiteLLM API key) into extra_headers (#23652).""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = MCPServer( + server_id="m2m-server", + name="m2m", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://auth.example.com/token", + ) + caller_key = {"Authorization": "Bearer sk-litellm-caller"} + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers=caller_key, + raw_headers=None, + ) + + assert server_auth_header is None + assert extra_headers is None + + +def test_prepare_mcp_server_headers_oauth2_interactive_copies_oauth2_headers(): + """Interactive OAuth still forwards the user's OAuth token in extra_headers.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + user_oauth = {"Authorization": "Bearer upstream-user-token"} + + server = MCPServer( + server_id="3lo-server", + name="3lo", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow=None, + ) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers=user_oauth, + raw_headers=None, + ) + + assert server_auth_header is None + assert extra_headers == user_oauth + + +def test_prepare_mcp_server_headers_m2m_skips_authorization_from_raw_extra_headers(): + """M2M must not merge caller Authorization from raw_headers when extra_headers lists it.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = MCPServer( + server_id="m2m-raw", + name="m2m", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://auth.example.com/token", + extra_headers=["Authorization", "X-Custom"], + ) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers={"Authorization": "Bearer sk-1234"}, + raw_headers={ + "authorization": "Bearer sk-1234", + "x-custom": "trace", + }, + ) + + assert server_auth_header is None + assert extra_headers is not None + assert "Authorization" not in extra_headers + assert extra_headers.get("X-Custom") == "trace" + + +@pytest.mark.asyncio +async def test_call_tool_m2m_skips_authorization_headers(): + """M2M call_tool must not forward caller Authorization in oauth2/raw headers.""" + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + except ImportError: + pytest.skip("MCP server not available") + + manager = MCPServerManager() + server = MCPServer( + server_id="m2m-call-tool", + name="m2m-call-tool", + server_name="m2m-call-tool", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://auth.example.com/token", + client_id="cid", + client_secret="csecret", + extra_headers=["Authorization", "X-Custom"], + ) + + mock_client = MagicMock() + mock_client.call_tool = AsyncMock(return_value=MagicMock()) + + with patch.object( + manager, "_create_mcp_client", new=AsyncMock(return_value=mock_client) + ) as create_client_mock: + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="echo", + arguments={"message": "hello"}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers={"Authorization": "Bearer sk-1234"}, + raw_headers={"authorization": "Bearer sk-1234", "x-custom": "trace"}, + proxy_logging_obj=None, + ) + + create_kwargs = create_client_mock.await_args.kwargs + extra_headers = create_kwargs["extra_headers"] or {} + assert "Authorization" not in extra_headers + assert extra_headers.get("X-Custom") == "trace" + + @pytest.mark.asyncio async def test_get_prompts_from_mcp_servers_success(): try: @@ -2288,6 +2435,79 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab assert spend_meta["per_server_tool_counts"]["server_a"] == 1 +@pytest.mark.asyncio +async def test_get_tools_from_mcp_servers_returns_tools_when_success_logging_fails(): + """ + Regression test: list_tools should still return fetched tools even if + async_success_handler raises (e.g. serialization errors in logging path). + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_tools_from_mcp_servers, + ) + from litellm.proxy._types import UserAPIKeyAuth + except ImportError: + pytest.skip("MCP server not available") + + user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") + + server_a = MagicMock(name="server_a_obj") + server_a.name = "server_a" + server_a.alias = "server_a" + server_a.server_name = "server_a" + server_a.server_id = "a" + server_a.auth_type = None + server_a.extra_headers = None + + tool_1 = MagicMock() + tool_1.name = "server_a-tool_1" + + dummy_logging_obj = MagicMock() + dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}} + dummy_logging_obj.async_success_handler = AsyncMock( + side_effect=TypeError("Object of type Tool is not JSON serializable") + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server_a]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + return_value=(None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + side_effect=lambda tools, _server: tools, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + new=AsyncMock(side_effect=lambda tools, **_: tools), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.function_setup", + return_value=(dummy_logging_obj, None), + ), + ): + mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) + + tools = await _get_tools_from_mcp_servers( + user_api_key_auth=user_auth, + mcp_auth_header=None, + mcp_servers=["server_a"], + mcp_server_auth_headers=None, + log_list_tools_to_spendlogs=True, + list_tools_log_source="mcp_protocol", + ) + + assert tools == [tool_1] + dummy_logging_obj.async_success_handler.assert_awaited_once() + + def test_tool_name_matches_case_insensitive(): """Test that _tool_name_matches performs case-insensitive comparison. @@ -2719,3 +2939,177 @@ class TestGatewayCreateInitializationOptions: _mcp_gateway_initialize_instructions.reset(tok) opts = server.create_initialization_options() assert getattr(opts, "instructions", None) is None + + +@pytest.mark.asyncio +async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow(): + """ + P1 Regression: list_tools path must apply _resolve_oauth2_flow to legacy DB + rows where oauth2_flow is NULL but M2M credentials are present. + + Without this fix, has_client_credentials returns False and the caller's + Authorization header is forwarded upstream instead of being blocked. + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_tools_from_mcp_servers, + ) + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.mcp import MCPAuth + except ImportError: + pytest.skip("MCP server not available") + + user_auth = UserAPIKeyAuth(api_key="sk-1234", user_id="test-user") + + # Simulate a legacy DB row: OAuth2 with M2M credentials but oauth2_flow=None + legacy_server = MagicMock(name="legacy_m2m_server") + legacy_server.name = "legacy_m2m" + legacy_server.alias = "legacy_m2m" + legacy_server.server_name = "legacy_m2m" + legacy_server.server_id = "legacy-m2m-id" + legacy_server.auth_type = MCPAuth.oauth2 + legacy_server.oauth2_flow = None # Legacy: field not set in DB + legacy_server.token_url = "https://oauth.example.com/token" + legacy_server.authorization_url = None + legacy_server.client_id = "client-id" + legacy_server.client_secret = "client-secret" + legacy_server.extra_headers = None + legacy_server.has_client_credentials = False # This is the bug: should be True + legacy_server.model_copy = MagicMock( + side_effect=lambda update: MCPServer( + server_id=legacy_server.server_id, + name=legacy_server.name, + transport=MCPTransport.http, + auth_type=legacy_server.auth_type, + oauth2_flow=update.get("oauth2_flow", legacy_server.oauth2_flow), + token_url=legacy_server.token_url, + authorization_url=legacy_server.authorization_url, + client_id=legacy_server.client_id, + client_secret=legacy_server.client_secret, + ) + ) + + tool_1 = MagicMock() + tool_1.name = "legacy_m2m-tool" + + captured_extra_headers = None + + async def capture_extra_headers(*args, **kwargs): + nonlocal captured_extra_headers + captured_extra_headers = kwargs.get("extra_headers") + return [tool_1] + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + side_effect=lambda tools, _server: tools, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + new=AsyncMock(side_effect=lambda tools, **_: tools), + ), + ): + mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["legacy-m2m-id"]) + mock_manager.get_mcp_server_by_id = MagicMock(return_value=legacy_server) + mock_manager.filter_server_ids_by_ip_with_info = MagicMock( + return_value=(["legacy-m2m-id"], 0) + ) + mock_manager._get_tools_from_server = AsyncMock( + side_effect=capture_extra_headers + ) + + tools = await _get_tools_from_mcp_servers( + user_api_key_auth=user_auth, + mcp_auth_header=None, + mcp_servers=["legacy_m2m"], + mcp_server_auth_headers=None, + oauth2_headers={"Authorization": "Bearer sk-1234"}, # Caller's token + ) + + # With P1 fix: _get_allowed_mcp_servers applies _resolve_oauth2_flow, + # so has_client_credentials becomes True and extra_headers should be None + # (caller's Authorization blocked) + assert captured_extra_headers is None, ( + "P1 security issue: caller's Authorization header was forwarded to M2M server. " + "Expected None, got: " + str(captured_extra_headers) + ) + assert tools == [tool_1] + + +@pytest.mark.asyncio +async def test_call_tool_empty_extra_headers_returns_none(): + """ + P2 Regression: When all configured extra_headers are filtered out (e.g. + Authorization for M2M), the resulting extra_headers should be None, not {}. + + Downstream code that checks `if extra_headers is None` will behave + differently if an empty dict is passed instead. + """ + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + from litellm.types.mcp import MCPAuth + except ImportError: + pytest.skip("MCP server not available") + + manager = MCPServerManager() + + # M2M server with only Authorization in extra_headers + m2m_server = MCPServer( + server_id="m2m-srv", + name="m2m_test", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://oauth.example.com/token", + client_id="client-id", + client_secret="client-secret", + extra_headers=["Authorization"], # Will be filtered out for M2M + ) + + raw_headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"} + + captured_extra_headers = None + + async def capture_create_mcp_client(*args, **kwargs): + nonlocal captured_extra_headers + captured_extra_headers = kwargs.get("extra_headers") + # Return a mock client + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock(return_value=MagicMock(content=[])) + return mock_client + + with ( + patch.object( + manager, + "_create_mcp_client", + side_effect=capture_create_mcp_client, + ), + patch.object( + manager, + "get_mcp_server_by_id", + return_value=m2m_server, + ), + ): + try: + await manager._call_regular_mcp_tool( + mcp_server=m2m_server, + original_tool_name="test_tool", + arguments={}, + mcp_auth_header=None, + oauth2_headers=None, + raw_headers=raw_headers, + ) + except Exception: + pass # We only care about the captured headers + + # With P2 fix: extra_headers should be None (not {}) when all headers filtered + assert captured_extra_headers is None, ( + "P2 API consistency issue: expected None for empty extra_headers, got: " + + str(captured_extra_headers) + ) +