From e63537c6c1c266b1fb171dbd77190eb298702f7d Mon Sep 17 00:00:00 2001 From: Nikita Timofeev Date: Fri, 16 Jan 2026 16:11:14 +0000 Subject: [PATCH 01/16] Fix: ensure function content is valid JSON for GigaChat --- litellm/llms/gigachat/chat/transformation.py | 56 +++++++--- tests/llm_translation/test_gigachat.py | 108 +++++++++++++------ 2 files changed, 116 insertions(+), 48 deletions(-) diff --git a/litellm/llms/gigachat/chat/transformation.py b/litellm/llms/gigachat/chat/transformation.py index 90cf67da6b2..ba14de1f65d 100644 --- a/litellm/llms/gigachat/chat/transformation.py +++ b/litellm/llms/gigachat/chat/transformation.py @@ -31,6 +31,16 @@ else: GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1" +def is_valid_json(value: str) -> bool: + """Checks whether the value passed is a valid serialized JSON string""" + try: + json.loads(value) + except json.JSONDecodeError: + return False + else: + return True + + class GigaChatError(BaseLLMException): """GigaChat API error.""" @@ -101,7 +111,11 @@ class GigaChatConfig(BaseConfig): Set up headers with OAuth token. """ # Get access token - credentials = api_key or get_secret_str("GIGACHAT_CREDENTIALS") or get_secret_str("GIGACHAT_API_KEY") + credentials = ( + api_key + or get_secret_str("GIGACHAT_CREDENTIALS") + or get_secret_str("GIGACHAT_API_KEY") + ) access_token = get_access_token(credentials=credentials) # Store credentials for image uploads @@ -193,11 +207,13 @@ class GigaChatConfig(BaseConfig): for tool in tools: if tool.get("type") == "function": func = tool.get("function", {}) - functions.append({ - "name": func.get("name", ""), - "description": func.get("description", ""), - "parameters": func.get("parameters", {}), - }) + functions.append( + { + "name": func.get("name", ""), + "description": func.get("description", ""), + "parameters": func.get("parameters", {}), + } + ) return functions def _map_tool_choice( @@ -281,8 +297,14 @@ class GigaChatConfig(BaseConfig): } # Add optional params - for key in ["temperature", "top_p", "max_tokens", "stream", - "repetition_penalty", "profanity_check"]: + for key in [ + "temperature", + "top_p", + "max_tokens", + "stream", + "repetition_penalty", + "profanity_check", + ]: if key in optional_params: request_data[key] = optional_params[key] @@ -314,7 +336,7 @@ class GigaChatConfig(BaseConfig): elif role == "tool": message["role"] = "function" content = message.get("content", "") - if not isinstance(content, str): + if not isinstance(content, str) or not is_valid_json(content): message["content"] = json.dumps(content, ensure_ascii=False) # Handle None content @@ -441,14 +463,16 @@ class GigaChatConfig(BaseConfig): # Convert to tool_calls format if isinstance(args, dict): args = json.dumps(args, ensure_ascii=False) - message_data["tool_calls"] = [{ - "id": f"call_{uuid.uuid4().hex[:24]}", - "type": "function", - "function": { - "name": func_call.get("name", ""), - "arguments": args, + message_data["tool_calls"] = [ + { + "id": f"call_{uuid.uuid4().hex[:24]}", + "type": "function", + "function": { + "name": func_call.get("name", ""), + "arguments": args, + }, } - }] + ] message_data.pop("function_call", None) finish_reason = "tool_calls" diff --git a/tests/llm_translation/test_gigachat.py b/tests/llm_translation/test_gigachat.py index dd8d56ff54a..b69a5428e42 100644 --- a/tests/llm_translation/test_gigachat.py +++ b/tests/llm_translation/test_gigachat.py @@ -5,9 +5,7 @@ Tests message transformation, parameter handling, and response transformation. Run with: pytest tests/llm_translation/test_gigachat.py -v """ -import json import pytest -from unittest.mock import Mock, MagicMock class TestGigaChatMessageTransformation: @@ -16,6 +14,7 @@ class TestGigaChatMessageTransformation: @pytest.fixture def config(self): from litellm.llms.gigachat.chat.transformation import GigaChatConfig + return GigaChatConfig() def test_simple_user_message(self, config): @@ -52,20 +51,46 @@ class TestGigaChatMessageTransformation: assert result[0]["role"] == "function" + def test_tool_content_convertation_non_string_value(self, config): + """Non string tool content should be serialized""" + messages = [{"role": "tool", "content": {"output": 42}}] + result = config._transform_messages(messages) + + assert result[0]["content"] == '{"output": 42}' + + def test_tool_content_convertation_json_string_value(self, config): + """JSON string tool content left unchanged""" + valid_json = '{"output": "red car"}' + messages = [{"role": "tool", "content": valid_json}] + result = config._transform_messages(messages) + + assert result[0]["content"] == valid_json + + def test_tool_content_convertation_random_string_value(self, config): + """Non JSON tool content should be serialized""" + messages = [{"role": "tool", "content": "random string"}] + result = config._transform_messages(messages) + + assert result[0]["content"] == '"random string"' + def test_tool_calls_to_function_call(self, config): """tool_calls should be converted to function_call""" - messages = [{ - "role": "assistant", - "content": "", - "tool_calls": [{ - "id": "call_123", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"city": "Moscow"}' - } - }] - }] + messages = [ + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_123", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"city": "Moscow"}', + }, + } + ], + } + ] result = config._transform_messages(messages) assert "function_call" in result[0] @@ -94,6 +119,7 @@ class TestGigaChatCollapseUserMessages: @pytest.fixture def config(self): from litellm.llms.gigachat.chat.transformation import GigaChatConfig + return GigaChatConfig() def test_no_collapse_single_message(self, config): @@ -136,23 +162,24 @@ class TestGigaChatToolsTransformation: @pytest.fixture def config(self): from litellm.llms.gigachat.chat.transformation import GigaChatConfig + return GigaChatConfig() def test_single_tool_conversion(self, config): """Single tool should be converted correctly""" - tools = [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather for a city", - "parameters": { - "type": "object", - "properties": { - "city": {"type": "string"} - } - } + tools = [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather for a city", + "parameters": { + "type": "object", + "properties": {"city": {"type": "string"}}, + }, + }, } - }] + ] result = config._convert_tools_to_functions(tools) assert len(result) == 1 @@ -162,8 +189,22 @@ class TestGigaChatToolsTransformation: def test_multiple_tools_conversion(self, config): """Multiple tools should all be converted""" tools = [ - {"type": "function", "function": {"name": "func1", "description": "First", "parameters": {"type": "object", "properties": {}}}}, - {"type": "function", "function": {"name": "func2", "description": "Second", "parameters": {"type": "object", "properties": {}}}}, + { + "type": "function", + "function": { + "name": "func1", + "description": "First", + "parameters": {"type": "object", "properties": {}}, + }, + }, + { + "type": "function", + "function": { + "name": "func2", + "description": "Second", + "parameters": {"type": "object", "properties": {}}, + }, + }, ] result = config._convert_tools_to_functions(tools) @@ -178,6 +219,7 @@ class TestGigaChatParamsTransformation: @pytest.fixture def config(self): from litellm.llms.gigachat.chat.transformation import GigaChatConfig + return GigaChatConfig() def test_temperature_zero_becomes_top_p_zero(self, config): @@ -229,10 +271,10 @@ class TestGigaChatParamsTransformation: "type": "object", "properties": { "name": {"type": "string"}, - "age": {"type": "integer"} - } - } - } + "age": {"type": "integer"}, + }, + }, + }, } } result = config.map_openai_params( @@ -283,6 +325,7 @@ class TestGigaChatTransformRequest: @pytest.fixture def config(self): from litellm.llms.gigachat.chat.transformation import GigaChatConfig + return GigaChatConfig() def test_basic_request(self, config): @@ -335,6 +378,7 @@ class TestGigaChatSupportedParams: @pytest.fixture def config(self): from litellm.llms.gigachat.chat.transformation import GigaChatConfig + return GigaChatConfig() def test_supported_params(self, config): From 1ac32992dec19bef4a35f9c403690085843ec04a Mon Sep 17 00:00:00 2001 From: Chesars Date: Fri, 23 Jan 2026 15:25:40 -0300 Subject: [PATCH 02/16] fix(oci): serialize imageUrl as object for OCI GenAI API OCI GenAI expects imageUrl to be an object with a 'url' property, not a plain string. This was causing 400 errors when sending images. Fixes #19589 --- litellm/llms/oci/chat/transformation.py | 3 +- litellm/types/llms/oci.py | 9 ++++- .../oci/chat/test_oci_chat_transformation.py | 37 ++++++++++++++++++- 3 files changed, 45 insertions(+), 4 deletions(-) diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index 7af7be2094a..84f39ef2525 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -32,6 +32,7 @@ from litellm.types.llms.oci import ( OCICompletionResponse, OCIContentPartUnion, OCIImageContentPart, + OCIImageUrl, OCIMessage, OCIRoles, OCIServingMode, @@ -1129,7 +1130,7 @@ def adapt_messages_to_generic_oci_standard_content_message( image_url = image_url.get("url") if not isinstance(image_url, str): raise Exception("Prop `image_url` must be a string or an object with a `url` property") - new_content.append(OCIImageContentPart(imageUrl=image_url)) + new_content.append(OCIImageContentPart(imageUrl=OCIImageUrl(url=image_url))) return OCIMessage( role=open_ai_to_generic_oci_role_map[role], diff --git a/litellm/types/llms/oci.py b/litellm/types/llms/oci.py index b9a82cc8b73..9a654bc0f6c 100644 --- a/litellm/types/llms/oci.py +++ b/litellm/types/llms/oci.py @@ -35,11 +35,18 @@ class OCITextContentPart(OCIContentPart): text: str +class OCIImageUrl(BaseModel): + """ImageUrl object for OCI API. See: https://docs.oracle.com/en-us/iaas/tools/python/latest/api/generative_ai_inference/models/oci.generative_ai_inference.models.ImageUrl.html""" + + url: str + detail: Optional[Literal["AUTO", "HIGH", "LOW"]] = None + + class OCIImageContentPart(OCIContentPart): """Image content part for the OCI API.""" type: Literal["IMAGE"] = "IMAGE" - imageUrl: str + imageUrl: OCIImageUrl OCIContentPartUnion = Union[OCITextContentPart, OCIImageContentPart] diff --git a/tests/litellm/llms/oci/chat/test_oci_chat_transformation.py b/tests/litellm/llms/oci/chat/test_oci_chat_transformation.py index f706a025a09..0a6c59d1b44 100644 --- a/tests/litellm/llms/oci/chat/test_oci_chat_transformation.py +++ b/tests/litellm/llms/oci/chat/test_oci_chat_transformation.py @@ -203,6 +203,7 @@ class TestOCIImageUrlTransformation: """Tests for OCI image_url format handling in multimodal messages. Fixes: https://github.com/BerriAI/litellm/issues/18270 + Fixes: https://github.com/BerriAI/litellm/issues/19589 """ def test_image_url_as_string(self): @@ -224,7 +225,8 @@ class TestOCIImageUrlTransformation: assert len(result) == 1 assert result[0].role == "USER" assert len(result[0].content) == 2 - assert result[0].content[1].imageUrl == "https://example.com/image.png" + # imageUrl is now an OCIImageUrl object with a 'url' property + assert result[0].content[1].imageUrl.url == "https://example.com/image.png" def test_image_url_as_openai_object(self): """Test that image_url as OpenAI-style object {"url": "..."} works.""" @@ -245,7 +247,38 @@ class TestOCIImageUrlTransformation: assert len(result) == 1 assert result[0].role == "USER" assert len(result[0].content) == 2 - assert result[0].content[1].imageUrl == "https://example.com/image.png" + # imageUrl is now an OCIImageUrl object with a 'url' property + assert result[0].content[1].imageUrl.url == "https://example.com/image.png" + + def test_image_url_serializes_as_object(self): + """Test that imageUrl serializes as {"url": "..."} for OCI API. + + Fixes: https://github.com/BerriAI/litellm/issues/19589 + OCI expects imageUrl to be an object with a 'url' property, not a plain string. + """ + from litellm.llms.oci.chat.transformation import adapt_messages_to_generic_oci_standard + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Describe this image."}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,ABC123"}}, + ], + } + ] + + result = adapt_messages_to_generic_oci_standard(messages) + image_part = result[0].content[1] + + # Serialize as OCI would receive it (with exclude_none=True) + serialized = image_part.model_dump(exclude_none=True) + + # Verify the structure matches OCI's expected format + assert serialized == { + "type": "IMAGE", + "imageUrl": {"url": "data:image/png;base64,ABC123"} + } def test_image_url_invalid_type_raises_error(self): """Test that invalid image_url type raises an error.""" From caa2c57619ad658c16af79706b7eaf970ef1b3e9 Mon Sep 17 00:00:00 2001 From: Ephrim Stanley Date: Mon, 26 Jan 2026 10:36:43 -0500 Subject: [PATCH 03/16] Fix /batches to return encoded ids (from managed objects table) --- .../proxy/hooks/managed_files.py | 87 +++++++ litellm/proxy/batches_endpoints/endpoints.py | 22 +- .../proxy/hooks/test_managed_files.py | 218 ++++++++++++++++++ 3 files changed, 322 insertions(+), 5 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 445d2b242b4..776706bdd1b 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -244,6 +244,93 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): return managed_object.created_by == user_id return True # don't raise error if managed object is not found + async def list_user_batches( + self, + user_api_key_dict: UserAPIKeyAuth, + limit: Optional[int] = None, + after: Optional[str] = None, + provider: Optional[str] = None, + target_model_names: Optional[str] = None, + ) -> Dict[str, Any]: + # Provider filtering is not supported for managed batches + # This is because the encoded object ids stored in the managed objects table do not contain the provider information + # To support provider filtering, we would need to store the provider information in the encoded object ids + if provider: + raise Exception( + "Filtering by 'provider' is not supported when using managed batches. " + "Use 'target_model_names' to filter by specific model names instead." + ) + + where_clause: Dict[str, Any] = {"file_purpose": "batch"} + + # Filter by user who created the batch + if user_api_key_dict.user_id: + where_clause["created_by"] = user_api_key_dict.user_id + + if after: + where_clause["id"] = {"gt": after} + + # Fetch more than needed to allow for post-fetch filtering + fetch_limit = limit or 20 + if target_model_names: + # Fetch extra to account for filtering + fetch_limit = max(fetch_limit * 3, 100) + + batches = await self.prisma_client.db.litellm_managedobjecttable.find_many( + where=where_clause, + take=fetch_limit, + order={"created_at": "desc"}, + ) + + # Parse target_model_names filter + target_models_filter: List[str] = [] + if target_model_names: + target_models_filter = [m.strip() for m in target_model_names.split(",") if m.strip()] + + batch_objects: List[LiteLLMBatch] = [] + for batch in batches: + try: + # Stop once we have enough after filtering + if len(batch_objects) >= (limit or 20): + break + + batch_data = json.loads(batch.file_object) if isinstance(batch.file_object, str) else batch.file_object + batch_obj = LiteLLMBatch(**batch_data) + batch_obj.id = batch.unified_object_id + + # If no target_model_names filter, add the batch to the list + if not target_models_filter: + batch_objects.append(batch_obj) + continue + + # Filter by target_model_names + decoded_id = _is_base64_encoded_unified_file_id(batch.unified_object_id) + model_id = None + if decoded_id: + model_id = get_model_id_from_unified_batch_id(decoded_id) + + # Skip batches without decodable IDs if filtering is requested + if not model_id: + continue + + if any(target.lower() in model_id.lower() for target in target_models_filter): + batch_objects.append(batch_obj) + continue + + except Exception as e: + verbose_logger.warning( + f"Failed to parse batch object {batch.unified_object_id}: {e}" + ) + continue + + return { + "object": "list", + "data": batch_objects, + "first_id": batch_objects[0].id if batch_objects else None, + "last_id": batch_objects[-1].id if batch_objects else None, + "has_more": len(batch_objects) == (limit or 20), + } + async def get_user_created_file_ids( self, user_api_key_dict: UserAPIKeyAuth, model_object_ids: List[str] ) -> List[OpenAIFileObject]: diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 086105042e8..078e21f9bb4 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -542,14 +542,26 @@ async def list_batches( route_type="alist_batches", ) - model_param = ( + # Try to use managed objects table for listing batches (returns encoded IDs) + managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") + if managed_files_obj is not None and hasattr(managed_files_obj, "list_user_batches"): + verbose_proxy_logger.debug( + "Using managed objects table for batch listing" + ) + response = await managed_files_obj.list_user_batches( + user_api_key_dict=user_api_key_dict, + limit=limit, + after=after, + provider=provider, + target_model_names=target_model_names, + llm_router=llm_router, + ) + elif (model_param := ( data.get("model") or request.query_params.get("model") or request.headers.get("x-litellm-model") - ) - - # SCENARIO 2: Use model-based routing from header/query/body - if model_param: + )): + # SCENARIO 2: Use model-based routing from header/query/body credentials = get_credentials_for_model( llm_router=llm_router, model_id=model_param, diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index 9a6e153a22b..e70e0640aa6 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -827,3 +827,221 @@ async def test_afile_retrieve_raises_error_for_non_managed_file(): ) assert "not found" in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_list_batches_from_managed_objects_table(): + from litellm.proxy._types import UserAPIKeyAuth + from openai.types.batch import BatchRequestCounts + + prisma_client = AsyncMock() + + batch_record_1 = MagicMock() + batch_record_1.unified_object_id = "unified-batch-id-1" + batch_record_1.file_object = json.dumps({ + "id": "batch_abc123", + "object": "batch", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "status": "completed", + "created_at": 1234567890, + "input_file_id": "file-input-1", + "request_counts": {"total": 1, "completed": 1, "failed": 0}, + }) + + batch_record_2 = MagicMock() + batch_record_2.unified_object_id = "unified-batch-id-2" + batch_record_2.file_object = json.dumps({ + "id": "batch_xyz789", + "object": "batch", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "status": "in_progress", + "created_at": 1234567891, + "input_file_id": "file-input-2", + "request_counts": {"total": 5, "completed": 2, "failed": 0}, + }) + + prisma_client.db.litellm_managedobjecttable.find_many.return_value = [ + batch_record_1, + batch_record_2, + ] + + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=prisma_client + ) + + result = await proxy_managed_files.list_user_batches( + user_api_key_dict=UserAPIKeyAuth(user_id="test-user"), + limit=10, + ) + + assert result["object"] == "list" + assert len(result["data"]) == 2 + assert result["data"][0].id == "unified-batch-id-1" + assert result["data"][1].id == "unified-batch-id-2" + assert result["first_id"] == "unified-batch-id-1" + assert result["last_id"] == "unified-batch-id-2" + + # Should filter by user_id (created_by) + prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with( + where={"file_purpose": "batch", "created_by": "test-user"}, + take=10, + order={"created_at": "desc"}, + ) + + +@pytest.mark.asyncio +async def test_list_batches_from_managed_objects_table_empty_list(): + from litellm.proxy._types import UserAPIKeyAuth + + prisma_client = AsyncMock() + prisma_client.db.litellm_managedobjecttable.find_many.return_value = [] + + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=prisma_client + ) + + result = await proxy_managed_files.list_user_batches( + user_api_key_dict=UserAPIKeyAuth(user_id="test-user"), + ) + + assert result["object"] == "list" + assert len(result["data"]) == 0 + assert result["first_id"] is None + assert result["last_id"] is None + assert result["has_more"] is False + + # Verify where clause includes created_by filter + # Default take is 20 when no limit is provided + prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with( + where={"file_purpose": "batch", "created_by": "test-user"}, + take=20, + order={"created_at": "desc"}, + ) + + +def _create_unified_batch_id(model_id: str, batch_id: str) -> str: + import base64 + unified_str = f"litellm_proxy;model_id:{model_id};llm_batch_id:{batch_id}" + return base64.urlsafe_b64encode(unified_str.encode()).decode().rstrip("=") + + +@pytest.mark.asyncio +async def test_list_batches_from_managed_objects_table_provider_filter_raises_exception(): + from litellm.proxy._types import UserAPIKeyAuth + + prisma_client = AsyncMock() + + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=prisma_client + ) + + # Filtering by provider should raise Exception + with pytest.raises(Exception) as exc_info: + await proxy_managed_files.list_user_batches( + user_api_key_dict=UserAPIKeyAuth(user_id="test-user"), + limit=10, + provider="openai", + ) + + assert str(exc_info.value) == ( + "Filtering by 'provider' is not supported when using managed batches." + ) + + # Verify find_many was NOT called since exception is raised before database query + prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called() + + +@pytest.mark.asyncio +async def test_list_batches_from_managed_objects_table_target_model_name_filter_raises_exception(): + from litellm.proxy._types import UserAPIKeyAuth + + prisma_client = AsyncMock() + + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=prisma_client + ) + + # Filtering by provider should raise Exception + with pytest.raises(Exception) as exc_info: + await proxy_managed_files.list_user_batches( + user_api_key_dict=UserAPIKeyAuth(user_id="test-user"), + limit=10, + target_model_names="gpt-4o,gpt-3.5", + ) + + assert str(exc_info.value) == ( + "Filtering by 'target_model_names' is not supported when using managed batches." + ) + + # Verify find_many was NOT called since exception is raised before database query + prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called() + +@pytest.mark.asyncio +async def test_list_batches_from_managed_objects_table_filters_by_created_by(): + from litellm.proxy._types import UserAPIKeyAuth + + prisma_client = AsyncMock() + + # Create batch for user1 + batch_user1 = MagicMock() + batch_user1.unified_object_id = "unified-batch-user1" + batch_user1.file_object = json.dumps({ + "id": "batch_user1_abc", + "object": "batch", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "status": "completed", + "created_at": 1234567890, + "input_file_id": "file-input-user1", + "request_counts": {"total": 1, "completed": 1, "failed": 0}, + }) + + # Create batch for user2 + batch_user2 = MagicMock() + batch_user2.unified_object_id = "unified-batch-user2" + batch_user2.file_object = json.dumps({ + "id": "batch_user2_xyz", + "object": "batch", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "status": "completed", + "created_at": 1234567891, + "input_file_id": "file-input-user2", + "request_counts": {"total": 2, "completed": 2, "failed": 0}, + }) + + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=prisma_client + ) + + # Query with user1's API key - should only return user1's batch + prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user1] + result_user1 = await proxy_managed_files.list_user_batches( + user_api_key_dict=UserAPIKeyAuth(user_id="user1"), + limit=10, + ) + + assert len(result_user1["data"]) == 1 + assert result_user1["data"][0].id == "unified-batch-user1" + prisma_client.db.litellm_managedobjecttable.find_many.assert_called_with( + where={"file_purpose": "batch", "created_by": "user1"}, + take=10, + order={"created_at": "desc"}, + ) + + # Query with user2's API key - should only return user2's batch + prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user2] + result_user2 = await proxy_managed_files.list_user_batches( + user_api_key_dict=UserAPIKeyAuth(user_id="user2"), + limit=10, + ) + + assert len(result_user2["data"]) == 1 + assert result_user2["data"][0].id == "unified-batch-user2" + prisma_client.db.litellm_managedobjecttable.find_many.assert_called_with( + where={"file_purpose": "batch", "created_by": "user2"}, + take=10, + order={"created_at": "desc"}, + ) \ No newline at end of file From 88280d9cca4b7fbd28fa25cf24881a639d2990af Mon Sep 17 00:00:00 2001 From: Ephrim Stanley Date: Mon, 26 Jan 2026 12:02:35 -0500 Subject: [PATCH 04/16] Fix batch creation to return the input file's expires_at attribute --- .../proxy/hooks/managed_files.py | 40 ++++++----------- .../proxy/hooks/test_managed_files.py | 43 ++++++++++++++++++- 2 files changed, 55 insertions(+), 28 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 776706bdd1b..dd0613f87ff 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -251,14 +251,22 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): after: Optional[str] = None, provider: Optional[str] = None, target_model_names: Optional[str] = None, + llm_router: Optional[Router] = None, ) -> Dict[str, Any]: # Provider filtering is not supported for managed batches # This is because the encoded object ids stored in the managed objects table do not contain the provider information # To support provider filtering, we would need to store the provider information in the encoded object ids if provider: raise Exception( - "Filtering by 'provider' is not supported when using managed batches. " - "Use 'target_model_names' to filter by specific model names instead." + "Filtering by 'provider' is not supported when using managed batches." + ) + + # Model name filtering is not supported for managed batches + # This is because the encoded object ids stored in the managed objects table do not contain the model name + # A hash of the model name + litellm_params for the model name is encoded as the model id. This is not sufficient to reliably map the target model names to the model ids. + if target_model_names: + raise Exception( + "Filtering by 'target_model_names' is not supported when using managed batches." ) where_clause: Dict[str, Any] = {"file_purpose": "batch"} @@ -281,12 +289,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): take=fetch_limit, order={"created_at": "desc"}, ) - - # Parse target_model_names filter - target_models_filter: List[str] = [] - if target_model_names: - target_models_filter = [m.strip() for m in target_model_names.split(",") if m.strip()] - + batch_objects: List[LiteLLMBatch] = [] for batch in batches: try: @@ -297,26 +300,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): batch_data = json.loads(batch.file_object) if isinstance(batch.file_object, str) else batch.file_object batch_obj = LiteLLMBatch(**batch_data) batch_obj.id = batch.unified_object_id + batch_objects.append(batch_obj) - # If no target_model_names filter, add the batch to the list - if not target_models_filter: - batch_objects.append(batch_obj) - continue - - # Filter by target_model_names - decoded_id = _is_base64_encoded_unified_file_id(batch.unified_object_id) - model_id = None - if decoded_id: - model_id = get_model_id_from_unified_batch_id(decoded_id) - - # Skip batches without decodable IDs if filtering is requested - if not model_id: - continue - - if any(target.lower() in model_id.lower() for target in target_models_filter): - batch_objects.append(batch_obj) - continue - except Exception as e: verbose_logger.warning( f"Failed to parse batch object {batch.unified_object_id}: {e}" @@ -760,6 +745,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): bytes=file_objects[0].bytes, filename=file_objects[0].filename, status="uploaded", + expires_at=file_objects[0].expires_at, ) return response diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index e70e0640aa6..4fa16066e4b 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -1044,4 +1044,45 @@ async def test_list_batches_from_managed_objects_table_filters_by_created_by(): where={"file_purpose": "batch", "created_by": "user2"}, take=10, order={"created_at": "desc"}, - ) \ No newline at end of file + ) + + +@pytest.mark.asyncio +async def test_return_unified_file_id_includes_expires_at(): + from litellm.types.llms.openai import OpenAIFileObject + + # Create a mock file object with expires_at set + file_object = OpenAIFileObject( + id="file-abc123", + object="file", + bytes=1234, + created_at=1234567890, + filename="test.jsonl", + purpose="batch", + status="uploaded", + expires_at=1234657890, + ) + file_object._hidden_params = {"model_id": "test-model-id"} + + create_file_request = { + "file": ("test.jsonl", b"test content", "application/jsonl"), + "purpose": "batch", + } + + internal_usage_cache = MagicMock() + + result = await _PROXY_LiteLLMManagedFiles.return_unified_file_id( + file_objects=[file_object], + create_file_request=create_file_request, + internal_usage_cache=internal_usage_cache, + litellm_parent_otel_span=None, + target_model_names_list=["gpt-4o"], + ) + + # Verify expires_at is passed through + assert result.expires_at == 1234657890 + assert result.purpose == "batch" + assert result.filename == "test.jsonl" + assert result.bytes == 1234 + assert result.created_at == 1234567890 + assert _is_base64_encoded_unified_file_id(result.id) \ No newline at end of file From 6a54dcfa934671e0b9d6698a80cc5dfb8c6ccf2b Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Mon, 26 Jan 2026 20:32:08 -0800 Subject: [PATCH 05/16] feat: Add model_id label to Prometheus metrics (#18048) (#19678) Co-authored-by: Cursor Agent --- docs/my-website/docs/proxy/prometheus.md | 8 ++--- litellm/integrations/prometheus.py | 12 ++++--- litellm/types/integrations/prometheus.py | 7 ++++ .../test_prometheus_logging_callbacks.py | 14 +++++++- .../integrations/test_prometheus_labels.py | 35 +++++++++++++------ 5 files changed, 55 insertions(+), 21 deletions(-) diff --git a/docs/my-website/docs/proxy/prometheus.md b/docs/my-website/docs/proxy/prometheus.md index cd2b3b68f37..d5a3466a84c 100644 --- a/docs/my-website/docs/proxy/prometheus.md +++ b/docs/my-website/docs/proxy/prometheus.md @@ -121,8 +121,8 @@ Use this to track overall LiteLLM Proxy usage. | Metric Name | Description | |----------------------|--------------------------------------| -| `litellm_proxy_failed_requests_metric` | Total number of failed responses from proxy - the client did not get a success response from litellm proxy. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "exception_status", "exception_class", "route"` | -| `litellm_proxy_total_requests_metric` | Total number of requests made to the proxy server - track number of client side requests. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "status_code", "user_email", "route"` | +| `litellm_proxy_failed_requests_metric` | Total number of failed responses from proxy - the client did not get a success response from litellm proxy. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "user_email", "exception_status", "exception_class", "route", "model_id"` | +| `litellm_proxy_total_requests_metric` | Total number of requests made to the proxy server - track number of client side requests. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "status_code", "user_email", "route", "model_id"` | ### Callback Logging Metrics @@ -191,10 +191,10 @@ Use this for LLM API Error monitoring and tracking remaining rate limits and tok | Metric Name | Description | |----------------------|--------------------------------------| -| `litellm_request_total_latency_metric` | Total latency (seconds) for a request to LiteLLM Proxy Server - tracked for labels "end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model" | +| `litellm_request_total_latency_metric` | Total latency (seconds) for a request to LiteLLM Proxy Server - tracked for labels "end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model", "model_id" | | `litellm_overhead_latency_metric` | Latency overhead (seconds) added by LiteLLM processing - tracked for labels "model_group", "api_provider", "api_base", "litellm_model_name", "hashed_api_key", "api_key_alias" | | `litellm_llm_api_latency_metric` | Latency (seconds) for just the LLM API call - tracked for labels "model", "hashed_api_key", "api_key_alias", "team", "team_alias", "requested_model", "end_user", "user" | -| `litellm_llm_api_time_to_first_token_metric` | Time to first token for LLM API call - tracked for labels `model`, `hashed_api_key`, `api_key_alias`, `team`, `team_alias` [Note: only emitted for streaming requests] | +| `litellm_llm_api_time_to_first_token_metric` | Time to first token for LLM API call - tracked for labels `model`, `hashed_api_key`, `api_key_alias`, `team`, `team_alias`, `requested_model`, `end_user`, `user`, `model_id` [Note: only emitted for streaming requests] | ## Tracking `end_user` on Prometheus diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index bafb0d88c82..8710df05b4c 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -1296,12 +1296,14 @@ class PrometheusLogger(CustomLogger): time_to_first_token_seconds is not None and kwargs.get("stream", False) is True # only emit for streaming requests ): + _ttft_labels = prometheus_label_factory( + supported_enum_labels=self.get_labels_for_metric( + metric_name="litellm_llm_api_time_to_first_token_metric" + ), + enum_values=enum_values, + ) self.litellm_llm_api_time_to_first_token_metric.labels( - model, - user_api_key, - user_api_key_alias, - user_api_team, - user_api_team_alias, + **_ttft_labels ).observe(time_to_first_token_seconds) else: verbose_logger.debug( diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index fd9b722287d..670f8215d89 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -217,6 +217,10 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.API_KEY_ALIAS.value, UserAPIKeyLabelNames.TEAM.value, UserAPIKeyLabelNames.TEAM_ALIAS.value, + UserAPIKeyLabelNames.REQUESTED_MODEL.value, + UserAPIKeyLabelNames.END_USER.value, + UserAPIKeyLabelNames.USER.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_request_total_latency_metric = [ @@ -228,6 +232,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.TEAM_ALIAS.value, UserAPIKeyLabelNames.USER.value, UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_request_queue_time_seconds = [ @@ -258,6 +263,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.STATUS_CODE.value, UserAPIKeyLabelNames.USER_EMAIL.value, UserAPIKeyLabelNames.ROUTE.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_proxy_failed_requests_metric = [ @@ -272,6 +278,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.EXCEPTION_STATUS.value, UserAPIKeyLabelNames.EXCEPTION_CLASS.value, UserAPIKeyLabelNames.ROUTE.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_deployment_latency_per_output_token = [ diff --git a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index 2419c61c25c..a479d1a9fc9 100644 --- a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -411,7 +411,15 @@ def test_set_latency_metrics(prometheus_logger): # completion_start_time - api_call_start_time prometheus_logger.litellm_llm_api_time_to_first_token_metric.labels.assert_called_once_with( - "gpt-3.5-turbo", "key1", "alias1", "team1", "team_alias1" + end_user=None, + user="test_user", + hashed_api_key="test_hash", + api_key_alias="test_alias", + team="test_team", + team_alias="test_team_alias", + requested_model="openai-gpt", + model="gpt-3.5-turbo", + model_id="model-123", ) prometheus_logger.litellm_llm_api_time_to_first_token_metric.labels().observe.assert_called_once_with( 0.5 @@ -442,6 +450,7 @@ def test_set_latency_metrics(prometheus_logger): team_alias="test_team_alias", requested_model="openai-gpt", model="gpt-3.5-turbo", + model_id="model-123", ) prometheus_logger.litellm_request_total_latency_metric.labels().observe.assert_called_once_with( 2.0 @@ -737,6 +746,7 @@ async def test_async_post_call_failure_hook(prometheus_logger): exception_status="429", exception_class="Openai.RateLimitError", route=user_api_key_dict.request_route, + model_id=None, ) prometheus_logger.litellm_proxy_failed_requests_metric.labels().inc.assert_called_once() @@ -752,6 +762,7 @@ async def test_async_post_call_failure_hook(prometheus_logger): status_code="429", user_email=None, route=user_api_key_dict.request_route, + model_id=None, ) prometheus_logger.litellm_proxy_total_requests_metric.labels().inc.assert_called_once() @@ -798,6 +809,7 @@ async def test_async_post_call_success_hook(prometheus_logger): status_code="200", user_email=None, route=user_api_key_dict.request_route, + model_id=None, ) prometheus_logger.litellm_proxy_total_requests_metric.labels().inc.assert_called_once() diff --git a/tests/test_litellm/integrations/test_prometheus_labels.py b/tests/test_litellm/integrations/test_prometheus_labels.py index c0b863ef6ee..d9a57df1617 100644 --- a/tests/test_litellm/integrations/test_prometheus_labels.py +++ b/tests/test_litellm/integrations/test_prometheus_labels.py @@ -70,6 +70,30 @@ def test_prometheus_metric_labels_structure(): print(f"✅ {metric_name} has proper label structure with user_email") +def test_model_id_in_required_metrics(): + """ + Test that model_id label is present in all the metrics that should have it: + - litellm_proxy_total_requests_metric + - litellm_proxy_failed_requests_metric + - litellm_request_total_latency_metric + - litellm_llm_api_time_to_first_token_metric + """ + model_id_label = UserAPIKeyLabelNames.MODEL_ID.value + + # Metrics that should have model_id + metrics_with_model_id = [ + "litellm_proxy_total_requests_metric", + "litellm_proxy_failed_requests_metric", + "litellm_request_total_latency_metric", + "litellm_llm_api_time_to_first_token_metric" + ] + + for metric_name in metrics_with_model_id: + labels = PrometheusMetricLabels.get_labels(metric_name) + assert model_id_label in labels, f"Metric {metric_name} should contain model_id label" + print(f"✅ {metric_name} contains model_id label") + + def test_route_normalization_for_responses_api(): """ Test that route normalization prevents high cardinality in Prometheus metrics @@ -217,14 +241,3 @@ def test_prometheus_metrics_use_normalized_routes(): print("✅ Prometheus metrics use normalized routes in labels") - -if __name__ == "__main__": - test_user_email_in_required_metrics() - test_user_email_label_exists() - test_prometheus_metric_labels_structure() - test_route_normalization_for_responses_api() - test_route_normalization_for_sub_routes() - test_route_normalization_preserves_static_routes() - test_route_normalization_other_dynamic_apis() - test_prometheus_metrics_use_normalized_routes() - print("\n✅ All prometheus label tests passed!") \ No newline at end of file From 0d45b010691578a9fafb4ac241193c17b8d10401 Mon Sep 17 00:00:00 2001 From: Cesar Garcia <128240629+Chesars@users.noreply.github.com> Date: Tue, 27 Jan 2026 01:36:10 -0300 Subject: [PATCH 06/16] fix(models): set gpt-5.2-codex mode to responses for Azure and OpenRouter (#19770) Fixes #19754 The gpt-5.2-codex model only supports the responses API, not chat completions. Updated azure/gpt-5.2-codex and openrouter/openai/gpt-5.2-codex entries to use mode: "responses" and supported_endpoints: ["/v1/responses"]. --- model_prices_and_context_window.json | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index d958ea4503a..244a679bc27 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3653,10 +3653,9 @@ "max_input_tokens": 128000, "max_output_tokens": 16384, "max_tokens": 16384, - "mode": "chat", + "mode": "responses", "output_cost_per_token": 1.4e-05, "supported_endpoints": [ - "/v1/chat/completions", "/v1/responses" ], "supported_modalities": [ @@ -23813,8 +23812,11 @@ "max_input_tokens": 400000, "max_output_tokens": 128000, "max_tokens": 128000, - "mode": "chat", + "mode": "responses", "output_cost_per_token": 1.4e-05, + "supported_endpoints": [ + "/v1/responses" + ], "supported_modalities": [ "text", "image" From b1968a8e3372773123872d534c0faaf92fb9ede7 Mon Sep 17 00:00:00 2001 From: Cesar Garcia <128240629+Chesars@users.noreply.github.com> Date: Tue, 27 Jan 2026 01:47:35 -0300 Subject: [PATCH 07/16] fix(responses): update local_vars with detected provider (#19782) (#19798) When using the responses API with provider-specific params (aws_*, vertex_*) without explicitly passing custom_llm_provider, the code crashed with: AttributeError: 'NoneType' object has no attribute 'startswith' Root cause: local_vars was captured via locals() before get_llm_provider() detected the provider from the model string (e.g., "bedrock/..."), so custom_llm_provider remained None when processing provider-specific params. Fix: Update local_vars["custom_llm_provider"] after get_llm_provider() call so the detected provider is available for param processing. Affected provider-specific params: - aws_* (aws_region_name, aws_access_key_id, etc.) for Bedrock/SageMaker - vertex_* (vertex_project, vertex_location, etc.) for Vertex AI --- litellm/responses/main.py | 10 +++++ .../responses/test_responses_utils.py | 43 +++++++++++++++++++ 2 files changed, 53 insertions(+) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 83c23a58500..b2c2493c812 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -434,6 +434,8 @@ async def aresponses( _, custom_llm_provider, _, _ = litellm.get_llm_provider( model=model, api_base=local_vars.get("base_url", None) ) + # Update local_vars with detected provider (fixes #19782) + local_vars["custom_llm_provider"] = custom_llm_provider func = partial( responses, @@ -583,6 +585,9 @@ def responses( api_key=litellm_params.api_key, ) + # Update local_vars with detected provider (fixes #19782) + local_vars["custom_llm_provider"] = custom_llm_provider + # Use dynamic credentials from get_llm_provider (e.g., when use_litellm_proxy=True) if dynamic_api_key is not None: litellm_params.api_key = dynamic_api_key @@ -1411,6 +1416,8 @@ async def acompact_responses( _, custom_llm_provider, _, _ = litellm.get_llm_provider( model=model, api_base=local_vars.get("base_url", None) ) + # Update local_vars with detected provider (fixes #19782) + local_vars["custom_llm_provider"] = custom_llm_provider func = partial( compact_responses, @@ -1498,6 +1505,9 @@ def compact_responses( api_key=litellm_params.api_key, ) + # Update local_vars with detected provider (fixes #19782) + local_vars["custom_llm_provider"] = custom_llm_provider + # Use dynamic credentials from get_llm_provider (e.g., when use_litellm_proxy=True) if dynamic_api_key is not None: litellm_params.api_key = dynamic_api_key diff --git a/tests/test_litellm/responses/test_responses_utils.py b/tests/test_litellm/responses/test_responses_utils.py index 3eb9c63e1be..8f7acb6c120 100644 --- a/tests/test_litellm/responses/test_responses_utils.py +++ b/tests/test_litellm/responses/test_responses_utils.py @@ -309,3 +309,46 @@ class TestResponseAPILoggingUtils: assert result.completion_tokens_details.reasoning_tokens == 30 assert result.completion_tokens_details.image_tokens == 100 assert result.completion_tokens_details.text_tokens == 70 + + +class TestResponsesAPIProviderSpecificParams: + """ + Tests for fix #19782: provider-specific params (aws_*, vertex_*) should work + without explicitly passing custom_llm_provider. + """ + + def test_provider_specific_params_no_crash_with_bedrock(self): + """Test that processing aws_* params with bedrock provider doesn't crash.""" + params = { + "temperature": 0.7, + "custom_llm_provider": "bedrock", + "kwargs": {"aws_region_name": "eu-central-1"}, + } + + # Should not raise any exception + result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params) + assert "temperature" in result + + def test_provider_specific_params_no_crash_with_openai(self): + """Test that processing aws_* params with openai provider doesn't crash.""" + params = { + "temperature": 0.7, + "custom_llm_provider": "openai", + "kwargs": {"aws_region_name": "eu-central-1"}, + } + + # Should not raise any exception + result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params) + assert "temperature" in result + + def test_provider_specific_params_no_crash_with_vertex_ai(self): + """Test that processing vertex_* params with vertex_ai provider doesn't crash.""" + params = { + "temperature": 0.7, + "custom_llm_provider": "vertex_ai", + "kwargs": {"vertex_project": "my-project"}, + } + + # Should not raise any exception + result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params) + assert "temperature" in result From 16f456ad822f08643d373c4c539df74fcc5dba72 Mon Sep 17 00:00:00 2001 From: Cesar Garcia <128240629+Chesars@users.noreply.github.com> Date: Tue, 27 Jan 2026 02:00:03 -0300 Subject: [PATCH 08/16] fix(azure): use generic cost calculator for audio token pricing (#19771) Azure audio models were charging audio output tokens at the text token rate instead of the correct audio token rate. This resulted in costs being ~6.65x lower than expected. The fix replaces Azure's custom cost calculation logic with the generic cost calculator that properly handles text, audio, cached, reasoning, and image tokens. Fixes #19764 --- litellm/llms/azure/cost_calculation.py | 37 ++++------ tests/test_litellm/test_cost_calculator.py | 84 ++++++++++++++++++++++ 2 files changed, 97 insertions(+), 24 deletions(-) diff --git a/litellm/llms/azure/cost_calculation.py b/litellm/llms/azure/cost_calculation.py index 96c58d95ff2..5b411095ea1 100644 --- a/litellm/llms/azure/cost_calculation.py +++ b/litellm/llms/azure/cost_calculation.py @@ -1,11 +1,12 @@ """ Helper util for handling azure openai-specific cost calculation -- e.g.: prompt caching +- e.g.: prompt caching, audio tokens """ from typing import Optional, Tuple from litellm._logging import verbose_logger +from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token from litellm.types.utils import Usage from litellm.utils import get_model_info @@ -18,34 +19,15 @@ def cost_per_token( Input: - model: str, the model name without provider prefix - - usage: LiteLLM Usage block, containing anthropic caching information + - usage: LiteLLM Usage block, containing caching and audio token information Returns: Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd """ ## GET MODEL INFO model_info = get_model_info(model=model, custom_llm_provider="azure") - cached_tokens: Optional[int] = None - ## CALCULATE INPUT COST - non_cached_text_tokens = usage.prompt_tokens - if usage.prompt_tokens_details and usage.prompt_tokens_details.cached_tokens: - cached_tokens = usage.prompt_tokens_details.cached_tokens - non_cached_text_tokens = non_cached_text_tokens - cached_tokens - prompt_cost: float = non_cached_text_tokens * model_info["input_cost_per_token"] - ## CALCULATE OUTPUT COST - completion_cost: float = ( - usage["completion_tokens"] * model_info["output_cost_per_token"] - ) - - ## Prompt Caching cost calculation - if model_info.get("cache_read_input_token_cost") is not None and cached_tokens: - # Note: We read ._cache_read_input_tokens from the Usage - since cost_calculator.py standardizes the cache read tokens on usage._cache_read_input_tokens - prompt_cost += cached_tokens * ( - model_info.get("cache_read_input_token_cost", 0) or 0 - ) - - ## Speech / Audio cost calculation + ## Speech / Audio cost calculation (cost per second for TTS models) if ( "output_cost_per_second" in model_info and model_info["output_cost_per_second"] is not None @@ -55,7 +37,14 @@ def cost_per_token( f"For model={model} - output_cost_per_second: {model_info.get('output_cost_per_second')}; response time: {response_time_ms}" ) ## COST PER SECOND ## - prompt_cost = 0 + prompt_cost = 0.0 completion_cost = model_info["output_cost_per_second"] * response_time_ms / 1000 + return prompt_cost, completion_cost - return prompt_cost, completion_cost + ## Use generic cost calculator for all other cases + ## This properly handles: text tokens, audio tokens, cached tokens, reasoning tokens, etc. + return generic_cost_per_token( + model=model, + usage=usage, + custom_llm_provider="azure", + ) diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 9d968d482c6..f727ca6f6c4 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -357,6 +357,90 @@ def test_azure_realtime_cost_calculator(): assert cost > 0 +def test_azure_audio_output_cost_calculation(): + """ + Test that Azure audio models correctly calculate costs for audio output tokens. + + Reproduces issue: https://github.com/BerriAI/litellm/issues/19764 + Audio tokens should be charged at output_cost_per_audio_token rate, + not at the text token rate (output_cost_per_token). + """ + from litellm.types.utils import ( + Choices, + CompletionTokensDetailsWrapper, + Message, + ) + + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + # Scenario from issue #19764: + # Input: 17 text tokens, 0 audio tokens + # Output: 110 text tokens, 482 audio tokens + usage_object = Usage( + prompt_tokens=17, + completion_tokens=592, # 110 text + 482 audio + total_tokens=609, + prompt_tokens_details=PromptTokensDetailsWrapper( + audio_tokens=0, + cached_tokens=0, + text_tokens=17, + image_tokens=0, + ), + completion_tokens_details=CompletionTokensDetailsWrapper( + audio_tokens=482, + reasoning_tokens=0, + text_tokens=110, + ), + ) + + completion = ModelResponse( + id="test-azure-audio-cost", + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message( + content="Test response", + role="assistant", + ), + ) + ], + created=1729282652, + model="azure/gpt-audio-2025-08-28", + object="chat.completion", + usage=usage_object, + ) + + cost = completion_cost(completion, model="azure/gpt-audio-2025-08-28") + + model_info = litellm.get_model_info("azure/gpt-audio-2025-08-28") + + # Calculate expected cost + expected_input_cost = ( + model_info["input_cost_per_token"] * 17 # text tokens + ) + expected_output_cost = ( + model_info["output_cost_per_token"] * 110 # text tokens + + model_info["output_cost_per_audio_token"] * 482 # audio tokens + ) + expected_total_cost = expected_input_cost + expected_output_cost + + # The bug was: all output tokens charged at text rate + wrong_output_cost = model_info["output_cost_per_token"] * 592 + wrong_total_cost = expected_input_cost + wrong_output_cost + + # Verify audio tokens are NOT charged at text rate (the bug) + assert abs(cost - wrong_total_cost) > 0.001, ( + "Bug: Audio tokens are being charged at text token rate" + ) + + # Verify cost matches + assert abs(cost - expected_total_cost) < 0.0000001, ( + f"Expected cost {expected_total_cost}, got {cost}" + ) + + def test_default_image_cost_calculator(monkeypatch): from litellm.cost_calculator import default_image_cost_calculator From e4a557d95fe74e3a08a79ce20c3fe83ad5a5888f Mon Sep 17 00:00:00 2001 From: Cesar Garcia <128240629+Chesars@users.noreply.github.com> Date: Tue, 27 Jan 2026 02:00:35 -0300 Subject: [PATCH 09/16] fix(xai): correct cached token cost calculation for xAI models (#19772) * fix(azure): use generic cost calculator for audio token pricing Azure audio models were charging audio output tokens at the text token rate instead of the correct audio token rate. This resulted in costs being ~6.65x lower than expected. The fix replaces Azure's custom cost calculation logic with the generic cost calculator that properly handles text, audio, cached, reasoning, and image tokens. Fixes #19764 * fix(xai): correct cached token cost calculation for xAI models - Fix double-counting issue where xAI reports text_tokens = prompt_tokens (including cached), causing tokens to be charged twice - Add cache_read_input_token_cost to xAI grok-3 and grok-3-mini model variants - Detection: when text_tokens + cached_tokens > prompt_tokens, recalculate text_tokens = prompt_tokens - cached_tokens xAI pricing (25% of input for cached): - grok-3 variants: $0.75/M cached (input $3/M) - grok-3-mini variants: $0.075/M cached (input $0.30/M) --- .../litellm_core_utils/llm_cost_calc/utils.py | 24 +++++++++++++++---- ...odel_prices_and_context_window_backup.json | 11 +++++++++ model_prices_and_context_window.json | 11 +++++++++ 3 files changed, 41 insertions(+), 5 deletions(-) diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 785976ed319..642cbbb7922 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -566,14 +566,28 @@ def generic_cost_per_token( # noqa: PLR0915 if usage.prompt_tokens_details: prompt_tokens_details = _parse_prompt_tokens_details(usage) - ## EDGE CASE - text tokens not set inside PromptTokensDetails + ## EDGE CASE - text tokens not set or includes cached tokens (double-counting) + ## Some providers (like xAI) report text_tokens = prompt_tokens (including cached) + ## We detect this when: text_tokens + cached_tokens + other > prompt_tokens + ## Ref: https://github.com/BerriAI/litellm/issues/19680, #14874, #14875 - if prompt_tokens_details["text_tokens"] == 0: + cache_hit = prompt_tokens_details["cache_hit_tokens"] + text_tokens = prompt_tokens_details["text_tokens"] + audio_tokens = prompt_tokens_details["audio_tokens"] + cache_creation = prompt_tokens_details["cache_creation_tokens"] + image_tokens = prompt_tokens_details["image_tokens"] + + # Check for double-counting: sum of details > prompt_tokens means overlap + total_details = text_tokens + cache_hit + audio_tokens + cache_creation + image_tokens + has_double_counting = cache_hit > 0 and total_details > usage.prompt_tokens + + if text_tokens == 0 or has_double_counting: text_tokens = ( usage.prompt_tokens - - prompt_tokens_details["cache_hit_tokens"] - - prompt_tokens_details["audio_tokens"] - - prompt_tokens_details["cache_creation_tokens"] + - cache_hit + - audio_tokens + - cache_creation + - image_tokens ) prompt_tokens_details["text_tokens"] = text_tokens diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index d958ea4503a..cf6b0e7e3f2 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -30382,6 +30382,7 @@ "supports_web_search": true }, "xai/grok-3": { + "cache_read_input_token_cost": 7.5e-07, "input_cost_per_token": 3e-06, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30396,6 +30397,7 @@ "supports_web_search": true }, "xai/grok-3-beta": { + "cache_read_input_token_cost": 7.5e-07, "input_cost_per_token": 3e-06, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30410,6 +30412,7 @@ "supports_web_search": true }, "xai/grok-3-fast-beta": { + "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30424,6 +30427,7 @@ "supports_web_search": true }, "xai/grok-3-fast-latest": { + "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30438,6 +30442,7 @@ "supports_web_search": true }, "xai/grok-3-latest": { + "cache_read_input_token_cost": 7.5e-07, "input_cost_per_token": 3e-06, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30452,6 +30457,7 @@ "supports_web_search": true }, "xai/grok-3-mini": { + "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 3e-07, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30467,6 +30473,7 @@ "supports_web_search": true }, "xai/grok-3-mini-beta": { + "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 3e-07, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30482,6 +30489,7 @@ "supports_web_search": true }, "xai/grok-3-mini-fast": { + "cache_read_input_token_cost": 1.5e-07, "input_cost_per_token": 6e-07, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30497,6 +30505,7 @@ "supports_web_search": true }, "xai/grok-3-mini-fast-beta": { + "cache_read_input_token_cost": 1.5e-07, "input_cost_per_token": 6e-07, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30512,6 +30521,7 @@ "supports_web_search": true }, "xai/grok-3-mini-fast-latest": { + "cache_read_input_token_cost": 1.5e-07, "input_cost_per_token": 6e-07, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30527,6 +30537,7 @@ "supports_web_search": true }, "xai/grok-3-mini-latest": { + "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 3e-07, "litellm_provider": "xai", "max_input_tokens": 131072, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 244a679bc27..d0a3bec9009 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -30384,6 +30384,7 @@ "supports_web_search": true }, "xai/grok-3": { + "cache_read_input_token_cost": 7.5e-07, "input_cost_per_token": 3e-06, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30398,6 +30399,7 @@ "supports_web_search": true }, "xai/grok-3-beta": { + "cache_read_input_token_cost": 7.5e-07, "input_cost_per_token": 3e-06, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30412,6 +30414,7 @@ "supports_web_search": true }, "xai/grok-3-fast-beta": { + "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30426,6 +30429,7 @@ "supports_web_search": true }, "xai/grok-3-fast-latest": { + "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30440,6 +30444,7 @@ "supports_web_search": true }, "xai/grok-3-latest": { + "cache_read_input_token_cost": 7.5e-07, "input_cost_per_token": 3e-06, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30454,6 +30459,7 @@ "supports_web_search": true }, "xai/grok-3-mini": { + "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 3e-07, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30469,6 +30475,7 @@ "supports_web_search": true }, "xai/grok-3-mini-beta": { + "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 3e-07, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30484,6 +30491,7 @@ "supports_web_search": true }, "xai/grok-3-mini-fast": { + "cache_read_input_token_cost": 1.5e-07, "input_cost_per_token": 6e-07, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30499,6 +30507,7 @@ "supports_web_search": true }, "xai/grok-3-mini-fast-beta": { + "cache_read_input_token_cost": 1.5e-07, "input_cost_per_token": 6e-07, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30514,6 +30523,7 @@ "supports_web_search": true }, "xai/grok-3-mini-fast-latest": { + "cache_read_input_token_cost": 1.5e-07, "input_cost_per_token": 6e-07, "litellm_provider": "xai", "max_input_tokens": 131072, @@ -30529,6 +30539,7 @@ "supports_web_search": true }, "xai/grok-3-mini-latest": { + "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 3e-07, "litellm_provider": "xai", "max_input_tokens": 131072, From 885a02e6c85442e76cef5b237c1f5efd051ebe3c Mon Sep 17 00:00:00 2001 From: Harshit Jain <48647625+Harshit28j@users.noreply.github.com> Date: Tue, 27 Jan 2026 11:38:17 +0530 Subject: [PATCH 10/16] fix: token calculations and refactor (#19696) --- litellm/cost_calculator.py | 56 +++-- .../litellm_core_utils/llm_cost_calc/utils.py | 16 +- tests/test_litellm/test_cost_calculator.py | 231 +++++++++++------- 3 files changed, 202 insertions(+), 101 deletions(-) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 1ab7e260a83..490f0288b00 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -23,7 +23,11 @@ from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import from litellm.litellm_core_utils.llm_cost_calc.utils import ( CostCalculatorUtils, _generic_cost_per_character, + _get_service_tier_cost_key, + _parse_prompt_tokens_details, + calculate_cost_component, generic_cost_per_token, + get_billable_input_tokens, select_cost_metric_for_model, ) from litellm.llms.anthropic.cost_calculation import ( @@ -431,12 +435,18 @@ def cost_per_token( # noqa: PLR0915 model=model, custom_llm_provider=custom_llm_provider ) - if model_info["input_cost_per_token"] > 0: - ## COST PER TOKEN ## - prompt_tokens_cost_usd_dollar = ( - model_info["input_cost_per_token"] * prompt_tokens + if ( + model_info.get("input_cost_per_token", 0) > 0 + or model_info.get("output_cost_per_token", 0) > 0 + ): + return generic_cost_per_token( + model=model, + usage=usage_block, + custom_llm_provider=custom_llm_provider, + service_tier=service_tier, ) - elif ( + + if ( model_info.get("input_cost_per_second", None) is not None and response_time_ms is not None ): @@ -451,11 +461,7 @@ def cost_per_token( # noqa: PLR0915 model_info["input_cost_per_second"] * response_time_ms / 1000 # type: ignore ) - if model_info["output_cost_per_token"] > 0: - completion_tokens_cost_usd_dollar = ( - model_info["output_cost_per_token"] * completion_tokens - ) - elif ( + if ( model_info.get("output_cost_per_second", None) is not None and response_time_ms is not None ): @@ -955,7 +961,10 @@ def completion_cost( # noqa: PLR0915 router_model_id=router_model_id, ) - potential_model_names = [selected_model, _get_response_model(completion_response)] + potential_model_names = [ + selected_model, + _get_response_model(completion_response), + ] if model is not None: potential_model_names.append(model) @@ -1710,10 +1719,16 @@ def default_image_cost_calculator( ) # Priority 1: Use per-image pricing if available (for gpt-image-1 and similar models) - if "input_cost_per_image" in cost_info and cost_info["input_cost_per_image"] is not None: + if ( + "input_cost_per_image" in cost_info + and cost_info["input_cost_per_image"] is not None + ): return cost_info["input_cost_per_image"] * n # Priority 2: Fall back to per-pixel pricing for backward compatibility - elif "input_cost_per_pixel" in cost_info and cost_info["input_cost_per_pixel"] is not None: + elif ( + "input_cost_per_pixel" in cost_info + and cost_info["input_cost_per_pixel"] is not None + ): return cost_info["input_cost_per_pixel"] * height * width * n else: raise Exception( @@ -1833,9 +1848,22 @@ def batch_cost_calculator( if input_cost_per_token_batches: total_prompt_cost = usage.prompt_tokens * input_cost_per_token_batches elif input_cost_per_token: + # Subtract cached tokens from prompt_tokens before calculating cost + # Fixes issue where cached tokens are being charged again total_prompt_cost = ( - usage.prompt_tokens * (input_cost_per_token) / 2 + get_billable_input_tokens(usage) * (input_cost_per_token) / 2 ) # batch cost is usually half of the regular token cost + + # Add cache read cost if applicable + details = _parse_prompt_tokens_details(usage) + cache_read_tokens = details["cache_hit_tokens"] + cache_read_cost_key = _get_service_tier_cost_key( + "cache_read_input_token_cost", None + ) + total_prompt_cost += ( + calculate_cost_component(model_info, cache_read_cost_key, cache_read_tokens) + / 2 + ) if output_cost_per_token_batches: total_completion_cost = usage.completion_tokens * output_cost_per_token_batches elif output_cost_per_token: diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 642cbbb7922..fe06641a389 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -23,6 +23,15 @@ def _is_above_128k(tokens: float) -> bool: return False +def get_billable_input_tokens(usage: Usage) -> int: + """ + Returns the number of billable input tokens. + Subtracts cached tokens from prompt tokens if applicable. + """ + details = _parse_prompt_tokens_details(usage) + return usage.prompt_tokens - details["cache_hit_tokens"] + + def select_cost_metric_for_model( model_info: ModelInfo, ) -> Literal["cost_per_character", "cost_per_token"]: @@ -190,7 +199,6 @@ def _get_token_base_cost( 1000 if "k" in threshold_str else 1 ) if usage.prompt_tokens > threshold: - prompt_base_cost = cast( float, _get_cost_per_unit(model_info, key, prompt_base_cost) ) @@ -633,7 +641,11 @@ def generic_cost_per_token( # noqa: PLR0915 # Calculate text tokens as remainder when we have a breakdown # This handles cases like OpenAI's reasoning models where text_tokens isn't provided text_tokens = max( - 0, usage.completion_tokens - reasoning_tokens - audio_tokens - image_tokens + 0, + usage.completion_tokens + - reasoning_tokens + - audio_tokens + - image_tokens, ) else: # No breakdown at all, all tokens are text tokens diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index f727ca6f6c4..e277baf0b0c 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -1,4 +1,3 @@ -import json import os import sys @@ -8,7 +7,6 @@ sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path -from unittest.mock import MagicMock, patch from pydantic import BaseModel @@ -77,7 +75,9 @@ def test_cost_calculator_with_usage(monkeypatch): prompt_tokens=120, completion_tokens=100, prompt_tokens_details=PromptTokensDetailsWrapper( - text_tokens=10, audio_tokens=90, image_tokens=20, + text_tokens=10, + audio_tokens=90, + image_tokens=20, ), ) mr = ModelResponse(usage=usage, model="gemini-2.0-flash-001") @@ -96,7 +96,9 @@ def test_cost_calculator_with_usage(monkeypatch): # Step 1: Test a model where input_cost_per_image_token is not set. # In this case the calculation should use input_cost_per_token as fallback. - assert model_info.get("input_cost_per_image_token") is None, "Test case expects that input_cost_per_image_token is not set" + assert ( + model_info.get("input_cost_per_image_token") is None + ), "Test case expects that input_cost_per_image_token is not set" expected_cost = ( usage.prompt_tokens_details.audio_tokens @@ -116,9 +118,7 @@ def test_cost_calculator_with_usage(monkeypatch): monkeypatch.setattr( litellm, "model_cost", - { - "gemini-2.0-flash-001": temp_model_info_object - }, + {"gemini-2.0-flash-001": temp_model_info_object}, ) # Invalidate caches after modifying litellm.model_cost @@ -138,8 +138,10 @@ def test_cost_calculator_with_usage(monkeypatch): expected_cost = ( usage.prompt_tokens_details.audio_tokens * temp_model_info_object["input_cost_per_audio_token"] - + usage.prompt_tokens_details.text_tokens * temp_model_info_object["input_cost_per_token"] - + usage.prompt_tokens_details.image_tokens * temp_model_info_object["input_cost_per_image_token"] + + usage.prompt_tokens_details.text_tokens + * temp_model_info_object["input_cost_per_token"] + + usage.prompt_tokens_details.image_tokens + * temp_model_info_object["input_cost_per_image_token"] + usage.completion_tokens * temp_model_info_object["output_cost_per_token"] ) @@ -331,8 +333,6 @@ def test_custom_pricing_with_router_model_id(): def test_azure_realtime_cost_calculator(): - from litellm import get_model_info - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") @@ -473,9 +473,7 @@ def test_cost_calculator_with_cache_creation(): from litellm import completion_cost from litellm.types.utils import ( Choices, - CompletionTokensDetailsWrapper, Message, - PromptTokensDetailsWrapper, Usage, ) @@ -531,7 +529,7 @@ def test_cost_calculator_with_cache_creation(): def test_bedrock_cost_calculator_comparison_with_without_cache(): """Test that Bedrock caching reduces costs compared to non-cached requests""" from litellm import completion_cost - from litellm.types.utils import Choices, Message, PromptTokensDetailsWrapper, Usage + from litellm.types.utils import Choices, Message, Usage # Response WITHOUT caching response_no_cache = ModelResponse( @@ -782,7 +780,7 @@ def test_log_context_cost_calculation(): f"DEBUG: Tiered input cost per token (>200k): ${input_cost_above_200k:.2e}" ) else: - print(f"DEBUG: No tiered input pricing available, using base pricing") + print("DEBUG: No tiered input pricing available, using base pricing") input_cost_above_200k = input_cost_per_token if output_cost_above_200k is not None: @@ -790,7 +788,7 @@ def test_log_context_cost_calculation(): f"DEBUG: Tiered output cost per token (>200k): ${output_cost_above_200k:.2e}" ) else: - print(f"DEBUG: No tiered output pricing available, using base pricing") + print("DEBUG: No tiered output pricing available, using base pricing") output_cost_above_200k = output_cost_per_token if cache_creation_above_200k is not None: @@ -798,7 +796,7 @@ def test_log_context_cost_calculation(): f"DEBUG: Tiered cache creation cost per token (>200k): ${cache_creation_above_200k:.2e}" ) else: - print(f"DEBUG: No tiered cache creation pricing available, using base pricing") + print("DEBUG: No tiered cache creation pricing available, using base pricing") cache_creation_above_200k = cache_creation_cost_per_token # Since we're above 200k tokens, we should use tiered pricing if available @@ -1007,7 +1005,7 @@ def test_cost_discount_vertex_ai(): expected_cost = cost_without_discount * 0.95 assert cost_with_discount == pytest.approx(expected_cost, rel=1e-9) - print(f"✓ Cost discount test passed:") + print("✓ Cost discount test passed:") print(f" - Original cost: ${cost_without_discount:.6f}") print(f" - Discounted cost (5% off): ${cost_with_discount:.6f}") print(f" - Savings: ${cost_without_discount - cost_with_discount:.6f}") @@ -1057,7 +1055,7 @@ def test_cost_discount_not_applied_to_other_providers(): # Costs should be the same (no discount applied to OpenAI) assert cost_with_selective_discount == cost_without_discount - print(f"✓ Selective discount test passed:") + print("✓ Selective discount test passed:") print(f" - OpenAI cost (no discount configured): ${cost_without_discount:.6f}") print(f" - Cost remains unchanged: ${cost_with_selective_discount:.6f}") @@ -1107,7 +1105,7 @@ def test_cost_margin_percentage(): expected_cost = cost_without_margin * 1.10 assert cost_with_margin == pytest.approx(expected_cost, rel=1e-9) - print(f"✓ Cost margin percentage test passed:") + print("✓ Cost margin percentage test passed:") print(f" - Original cost: ${cost_without_margin:.6f}") print(f" - Cost with margin (10%): ${cost_with_margin:.6f}") print(f" - Margin added: ${cost_with_margin - cost_without_margin:.6f}") @@ -1158,7 +1156,7 @@ def test_cost_margin_fixed_amount(): expected_cost = cost_without_margin + 0.001 assert cost_with_margin == pytest.approx(expected_cost, rel=1e-9) - print(f"✓ Cost margin fixed amount test passed:") + print("✓ Cost margin fixed amount test passed:") print(f" - Original cost: ${cost_without_margin:.6f}") print(f" - Cost with margin ($0.001): ${cost_with_margin:.6f}") print(f" - Margin added: ${cost_with_margin - cost_without_margin:.6f}") @@ -1193,7 +1191,9 @@ def test_cost_margin_combined(): ) # Set 8% margin + $0.0005 fixed for openai - litellm.cost_margin_config = {"openai": {"percentage": 0.08, "fixed_amount": 0.0005}} + litellm.cost_margin_config = { + "openai": {"percentage": 0.08, "fixed_amount": 0.0005} + } # Calculate cost with margin cost_with_margin = completion_cost( @@ -1209,7 +1209,7 @@ def test_cost_margin_combined(): expected_cost = cost_without_margin * 1.08 + 0.0005 assert cost_with_margin == pytest.approx(expected_cost, rel=1e-9) - print(f"✓ Cost margin combined test passed:") + print("✓ Cost margin combined test passed:") print(f" - Original cost: ${cost_without_margin:.6f}") print(f" - Cost with margin (8% + $0.0005): ${cost_with_margin:.6f}") print(f" - Margin added: ${cost_with_margin - cost_without_margin:.6f}") @@ -1260,7 +1260,7 @@ def test_cost_margin_global(): expected_cost = cost_without_margin * 1.05 assert cost_with_global_margin == pytest.approx(expected_cost, rel=1e-9) - print(f"✓ Cost margin global test passed:") + print("✓ Cost margin global test passed:") print(f" - Original cost: ${cost_without_margin:.6f}") print(f" - Cost with global margin (5%): ${cost_with_global_margin:.6f}") print(f" - Margin added: ${cost_with_global_margin - cost_without_margin:.6f}") @@ -1311,9 +1311,11 @@ def test_cost_margin_provider_overrides_global(): expected_cost = cost_without_margin * 1.10 # 10% from provider, not 5% from global assert cost_with_provider_margin == pytest.approx(expected_cost, rel=1e-9) - print(f"✓ Cost margin provider override test passed:") + print("✓ Cost margin provider override test passed:") print(f" - Original cost: ${cost_without_margin:.6f}") - print(f" - Cost with provider margin (10%, overrides 5% global): ${cost_with_provider_margin:.6f}") + print( + f" - Cost with provider margin (10%, overrides 5% global): ${cost_with_provider_margin:.6f}" + ) print(f" - Margin added: ${cost_with_provider_margin - cost_without_margin:.6f}") @@ -1367,7 +1369,7 @@ def test_cost_margin_with_discount(): expected_cost = base_cost * 0.95 * 1.10 assert cost_with_both == pytest.approx(expected_cost, rel=1e-9) - print(f"✓ Cost margin with discount test passed:") + print("✓ Cost margin with discount test passed:") print(f" - Base cost: ${base_cost:.6f}") print(f" - Cost with 5% discount + 10% margin: ${cost_with_both:.6f}") print(f" - Expected: ${expected_cost:.6f}") @@ -1436,14 +1438,10 @@ def test_completion_cost_extracts_service_tier_from_response(): # Test with gpt-5-nano which has flex pricing model = "gpt-5-nano" - + # Create usage object - usage = Usage( - prompt_tokens=1000, - completion_tokens=500, - total_tokens=1500 - ) - + usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500) + # Create ModelResponse with service_tier in the response object response_with_service_tier = ModelResponse( usage=usage, @@ -1451,34 +1449,36 @@ def test_completion_cost_extracts_service_tier_from_response(): ) # Set service_tier as an attribute on the response setattr(response_with_service_tier, "service_tier", "flex") - + # Test that flex pricing is used when service_tier is in response flex_cost = completion_cost( completion_response=response_with_service_tier, model=model, custom_llm_provider="openai", ) - + # Create ModelResponse without service_tier (should use standard pricing) response_without_service_tier = ModelResponse( usage=usage, model=model, ) - + # Test that standard pricing is used when service_tier is not in response standard_cost = completion_cost( completion_response=response_without_service_tier, model=model, custom_llm_provider="openai", ) - + # Flex should be approximately 50% of standard assert flex_cost > 0, "Flex cost should be greater than 0" assert standard_cost > 0, "Standard cost should be greater than 0" assert flex_cost < standard_cost, "Flex cost should be less than standard cost" - + flex_ratio = flex_cost / standard_cost - assert 0.45 <= flex_ratio <= 0.55, f"Flex pricing should be ~50% of standard, got {flex_ratio:.2f}" + assert ( + 0.45 <= flex_ratio <= 0.55 + ), f"Flex pricing should be ~50% of standard, got {flex_ratio:.2f}" def test_completion_cost_extracts_service_tier_from_usage(): @@ -1490,56 +1490,54 @@ def test_completion_cost_extracts_service_tier_from_usage(): # Test with gpt-5-nano which has flex pricing model = "gpt-5-nano" - + # Create usage object with service_tier usage_with_service_tier = Usage( - prompt_tokens=1000, - completion_tokens=500, - total_tokens=1500 + prompt_tokens=1000, completion_tokens=500, total_tokens=1500 ) # Set service_tier as an attribute on the usage object setattr(usage_with_service_tier, "service_tier", "flex") - + # Create ModelResponse with usage containing service_tier response = ModelResponse( usage=usage_with_service_tier, model=model, ) - + # Test that flex pricing is used when service_tier is in usage flex_cost = completion_cost( completion_response=response, model=model, custom_llm_provider="openai", ) - + # Create usage object without service_tier usage_without_service_tier = Usage( - prompt_tokens=1000, - completion_tokens=500, - total_tokens=1500 + prompt_tokens=1000, completion_tokens=500, total_tokens=1500 ) - + # Create ModelResponse with usage without service_tier response_standard = ModelResponse( usage=usage_without_service_tier, model=model, ) - + # Test that standard pricing is used when service_tier is not in usage standard_cost = completion_cost( completion_response=response_standard, model=model, custom_llm_provider="openai", ) - + # Flex should be approximately 50% of standard assert flex_cost > 0, "Flex cost should be greater than 0" assert standard_cost > 0, "Standard cost should be greater than 0" assert flex_cost < standard_cost, "Flex cost should be less than standard cost" - + flex_ratio = flex_cost / standard_cost - assert 0.45 <= flex_ratio <= 0.55, f"Flex pricing should be ~50% of standard, got {flex_ratio:.2f}" + assert ( + 0.45 <= flex_ratio <= 0.55 + ), f"Flex pricing should be ~50% of standard, got {flex_ratio:.2f}" def test_completion_cost_service_tier_priority(): @@ -1551,22 +1549,18 @@ def test_completion_cost_service_tier_priority(): # Test with gpt-5-nano which has flex pricing model = "gpt-5-nano" - + # Create usage object with service_tier="flex" - usage = Usage( - prompt_tokens=1000, - completion_tokens=500, - total_tokens=1500 - ) + usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500) setattr(usage, "service_tier", "flex") - + # Create response with service_tier="priority" response = ModelResponse( usage=usage, model=model, ) setattr(response, "service_tier", "priority") - + # Test that optional_params takes priority over response and usage cost_from_params = completion_cost( completion_response=response, @@ -1574,14 +1568,14 @@ def test_completion_cost_service_tier_priority(): custom_llm_provider="openai", optional_params={"service_tier": "flex"}, ) - + # Test that response takes priority over usage when optional_params is not provided - cost_from_response = completion_cost( + completion_cost( completion_response=response, model=model, custom_llm_provider="openai", ) - + # Test that usage is used when neither optional_params nor response have service_tier # Create a new response without service_tier attribute response_no_tier = ModelResponse( @@ -1589,25 +1583,27 @@ def test_completion_cost_service_tier_priority(): model=model, ) # Don't set service_tier on response, so it will fall back to usage - + cost_from_usage = completion_cost( completion_response=response_no_tier, model=model, custom_llm_provider="openai", ) - + # All should use flex pricing (from different sources) assert cost_from_params > 0, "Cost from params should be greater than 0" assert cost_from_usage > 0, "Cost from usage should be greater than 0" - + # Costs should be similar (all using flex) - assert abs(cost_from_params - cost_from_usage) < 1e-6, "Costs from params and usage should be similar (both flex)" + assert ( + abs(cost_from_params - cost_from_usage) < 1e-6 + ), "Costs from params and usage should be similar (both flex)" def test_gemini_cache_tokens_details_no_negative_values(): """ Test for Issue #18750: Negative text_tokens with Gemini caching - + When using Gemini with explicit caching, the response includes cacheTokensDetails which breaks down cached tokens by modality. This test ensures that: 1. text_tokens is never negative @@ -1628,41 +1624,47 @@ def test_gemini_cache_tokens_details_no_negative_values(): # Total tokens by modality (includes cached + non-cached) "promptTokensDetails": [ {"modality": "TEXT", "tokenCount": 9402}, - {"modality": "IMAGE", "tokenCount": 258} + {"modality": "IMAGE", "tokenCount": 258}, ], # Breakdown of cached tokens by modality "cacheTokensDetails": [ {"modality": "TEXT", "tokenCount": 9393}, - {"modality": "IMAGE", "tokenCount": 258} - ] + {"modality": "IMAGE", "tokenCount": 258}, + ], } } usage = VertexGeminiConfig._calculate_usage(completion_response) # Text tokens should be non-cached text only: 9402 - 9393 = 9 - assert usage.prompt_tokens_details.text_tokens == 9, \ - f"Expected text_tokens=9, got {usage.prompt_tokens_details.text_tokens}" + assert ( + usage.prompt_tokens_details.text_tokens == 9 + ), f"Expected text_tokens=9, got {usage.prompt_tokens_details.text_tokens}" # Image tokens should be non-cached image only: 258 - 258 = 0 - assert usage.prompt_tokens_details.image_tokens == 0, \ - f"Expected image_tokens=0, got {usage.prompt_tokens_details.image_tokens}" + assert ( + usage.prompt_tokens_details.image_tokens == 0 + ), f"Expected image_tokens=0, got {usage.prompt_tokens_details.image_tokens}" # Total cached should match - assert usage.prompt_tokens_details.cached_tokens == 9651, \ - f"Expected cached_tokens=9651, got {usage.prompt_tokens_details.cached_tokens}" + assert ( + usage.prompt_tokens_details.cached_tokens == 9651 + ), f"Expected cached_tokens=9651, got {usage.prompt_tokens_details.cached_tokens}" # MOST IMPORTANT: text_tokens should NEVER be negative - assert usage.prompt_tokens_details.text_tokens >= 0, \ - f"BUG: text_tokens is negative ({usage.prompt_tokens_details.text_tokens})! This was the issue in #18750" + assert ( + usage.prompt_tokens_details.text_tokens >= 0 + ), f"BUG: text_tokens is negative ({usage.prompt_tokens_details.text_tokens})! This was the issue in #18750" - print("✅ Issue #18750 fix verified: text_tokens is correctly calculated and non-negative") + print( + "✅ Issue #18750 fix verified: text_tokens is correctly calculated and non-negative" + ) def test_gemini_without_cache_tokens_details(): """ Test Gemini response without cacheTokensDetails (implicit caching or no cache) - + When cacheTokensDetails is not present, we should use promptTokensDetails as-is without subtracting anything. """ @@ -1677,7 +1679,7 @@ def test_gemini_without_cache_tokens_details(): "totalTokenCount": 279, "promptTokensDetails": [ {"modality": "TEXT", "tokenCount": 6}, - {"modality": "IMAGE", "tokenCount": 258} + {"modality": "IMAGE", "tokenCount": 258}, ] # No cacheTokensDetails } @@ -1691,3 +1693,62 @@ def test_gemini_without_cache_tokens_details(): assert usage.prompt_tokens_details.text_tokens >= 0 print("✅ Gemini without cacheTokensDetails works correctly") + + +def test_generic_provider_cached_token_cost(): + """ + Test that the generic cost calculator correctly handles cached tokens + for providers like z.ai/deepseek that are not explicitly handled. + """ + from litellm.cost_calculator import completion_cost + from litellm.types.utils import ModelResponse, PromptTokensDetailsWrapper, Usage + + # Setup model cost for a generic provider + # We use a name that will bypass complex provider mapping logic + model_name = "custom-cached-model" + litellm.model_cost[model_name] = { + "input_cost_per_token": 0.0000006, + "output_cost_per_token": 0.0000006, + "cache_read_input_token_cost": 0.0000001, + "litellm_provider": "openai", + } + + # Case 1: Standard nested cached tokens (prompt_tokens_details.cached_tokens) + usage = Usage( + prompt_tokens=10000, + completion_tokens=0, + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=9000), + ) + response = ModelResponse(usage=usage, model=model_name) + + cost = completion_cost( + completion_response=response, + model=model_name, + custom_llm_provider="openai", # Explicitly set provider to trigger generic path + ) + + # Expected: (1000 * 0.0000006) + (9000 * 0.0000001) = 0.0006 + 0.0009 = 0.0015 + expected_cost = 0.0015 + assert ( + abs(cost - expected_cost) < 1e-9 + ), f"Nested cache cost failed. Got {cost}, expected {expected_cost}" + + # Case 2: Top-level cached tokens (cache_read_input_tokens) + usage_top = Usage( + prompt_tokens=10000, + completion_tokens=0, + cache_read_input_tokens=9000, + ) + response_top = ModelResponse(usage=usage_top, model=model_name) + + cost_top = completion_cost( + completion_response=response_top, + model=model_name, + custom_llm_provider="openai", + ) + + assert ( + abs(cost_top - expected_cost) < 1e-9 + ), f"Top-level cache cost failed. Got {cost_top}, expected {expected_cost}" + + print("✅ Generic provider cached token cost verified") From a5bc98a18a0e512f2f479b7071404736446d6186 Mon Sep 17 00:00:00 2001 From: Harshit Jain <48647625+Harshit28j@users.noreply.github.com> Date: Tue, 27 Jan 2026 11:40:54 +0530 Subject: [PATCH 11/16] =?UTF-8?q?fix(prometheus):=20safely=20handle=20None?= =?UTF-8?q?=20metadata=20in=20logging=20to=20prevent=20At=E2=80=A6=20(#196?= =?UTF-8?q?91)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(prometheus): safely handle None metadata in logging to prevent AttributeError * fix: lint issues --- litellm/integrations/prometheus.py | 99 +++++++++++++++--------------- 1 file changed, 48 insertions(+), 51 deletions(-) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 8710df05b4c..af569b9ed69 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -891,7 +891,7 @@ class PrometheusLogger(CustomLogger): model = kwargs.get("model", "") litellm_params = kwargs.get("litellm_params", {}) or {} - _metadata = litellm_params.get("metadata", {}) + _metadata = litellm_params.get("metadata") or {} get_end_user_id_for_cost_tracking = _get_cached_end_user_id_for_cost_tracking() end_user_id = get_end_user_id_for_cost_tracking( @@ -1165,26 +1165,15 @@ class PrometheusLogger(CustomLogger): response_cost: float, user_id: Optional[str] = None, ): - _team_spend = litellm_params.get("metadata", {}).get( - "user_api_key_team_spend", None - ) - _team_max_budget = litellm_params.get("metadata", {}).get( - "user_api_key_team_max_budget", None - ) + _metadata = litellm_params.get("metadata") or {} + _team_spend = _metadata.get("user_api_key_team_spend", None) + _team_max_budget = _metadata.get("user_api_key_team_max_budget", None) - _api_key_spend = litellm_params.get("metadata", {}).get( - "user_api_key_spend", None - ) - _api_key_max_budget = litellm_params.get("metadata", {}).get( - "user_api_key_max_budget", None - ) + _api_key_spend = _metadata.get("user_api_key_spend", None) + _api_key_max_budget = _metadata.get("user_api_key_max_budget", None) - _user_spend = litellm_params.get("metadata", {}).get( - "user_api_key_user_spend", None - ) - _user_max_budget = litellm_params.get("metadata", {}).get( - "user_api_key_user_max_budget", None - ) + _user_spend = _metadata.get("user_api_key_user_spend", None) + _user_max_budget = _metadata.get("user_api_key_user_max_budget", None) await self._set_api_key_budget_metrics_after_api_request( user_api_key=user_api_key, @@ -1343,7 +1332,7 @@ class PrometheusLogger(CustomLogger): # request queue time (time from arrival to processing start) _litellm_params = kwargs.get("litellm_params", {}) or {} - queue_time_seconds = _litellm_params.get("metadata", {}).get( + queue_time_seconds = (_litellm_params.get("metadata") or {}).get( "queue_time_seconds" ) if queue_time_seconds is not None and queue_time_seconds >= 0: @@ -1367,14 +1356,14 @@ class PrometheusLogger(CustomLogger): standard_logging_payload: StandardLoggingPayload = kwargs.get( "standard_logging_object", {} ) - + if self._should_skip_metrics_for_invalid_key( kwargs=kwargs, standard_logging_payload=standard_logging_payload ): return - + model = kwargs.get("model", "") - + litellm_params = kwargs.get("litellm_params", {}) or {} get_end_user_id_for_cost_tracking = _get_cached_end_user_id_for_cost_tracking() @@ -1415,49 +1404,57 @@ class PrometheusLogger(CustomLogger): ) -> Optional[int]: """ Extract HTTP status code from various input formats for validation. - + This is a centralized helper to extract status code from different callback function signatures. Handles both ProxyException (uses 'code') and standard exceptions (uses 'status_code'). - + Args: kwargs: Dictionary potentially containing 'exception' key enum_values: Object with 'status_code' attribute exception: Exception object to extract status code from directly - + Returns: Status code as integer if found, None otherwise """ status_code = None - + # Try from enum_values first (most common in our callbacks) - if enum_values and hasattr(enum_values, "status_code") and enum_values.status_code: + if ( + enum_values + and hasattr(enum_values, "status_code") + and enum_values.status_code + ): try: status_code = int(enum_values.status_code) except (ValueError, TypeError): pass - + if not status_code and exception: # ProxyException uses 'code' attribute, other exceptions may use 'status_code' - status_code = getattr(exception, "status_code", None) or getattr(exception, "code", None) + status_code = getattr(exception, "status_code", None) or getattr( + exception, "code", None + ) if status_code is not None: try: status_code = int(status_code) except (ValueError, TypeError): status_code = None - + if not status_code and kwargs: exception_in_kwargs = kwargs.get("exception") if exception_in_kwargs: - status_code = getattr(exception_in_kwargs, "status_code", None) or getattr(exception_in_kwargs, "code", None) + status_code = getattr( + exception_in_kwargs, "status_code", None + ) or getattr(exception_in_kwargs, "code", None) if status_code is not None: try: status_code = int(status_code) except (ValueError, TypeError): status_code = None - + return status_code - + def _is_invalid_api_key_request( self, status_code: Optional[int], @@ -1465,23 +1462,23 @@ class PrometheusLogger(CustomLogger): ) -> bool: """ Determine if a request has an invalid API key based on status code and exception. - + This method prevents invalid authentication attempts from being recorded in Prometheus metrics. A 401 status code is the definitive indicator of authentication failure. Additionally, we check exception messages for authentication error patterns to catch cases where the exception hasn't been converted to a ProxyException yet. - + Args: status_code: HTTP status code (401 indicates authentication error) exception: Exception object to check for auth-related error messages - + Returns: True if the request has an invalid API key and metrics should be skipped, False otherwise """ if status_code == 401: return True - + # Handle cases where AssertionError is raised before conversion to ProxyException if exception is not None: exception_str = str(exception).lower() @@ -1494,9 +1491,9 @@ class PrometheusLogger(CustomLogger): ] if any(pattern in exception_str for pattern in auth_error_patterns): return True - + return False - + def _should_skip_metrics_for_invalid_key( self, kwargs: Optional[dict] = None, @@ -1507,18 +1504,18 @@ class PrometheusLogger(CustomLogger): ) -> bool: """ Determine if Prometheus metrics should be skipped for invalid API key requests. - + This is a centralized validation method that extracts status code and exception information from various callback function signatures and determines if the request represents an invalid API key attempt that should be filtered from metrics. - + Args: kwargs: Dictionary potentially containing exception and other data user_api_key_dict: User API key authentication object (currently unused) enum_values: Object with status_code attribute standard_logging_payload: Standard logging payload dictionary exception: Exception object to check directly - + Returns: True if metrics should be skipped (invalid key detected), False otherwise """ @@ -1527,17 +1524,17 @@ class PrometheusLogger(CustomLogger): enum_values=enum_values, exception=exception, ) - + if exception is None and kwargs: exception = kwargs.get("exception") - + if self._is_invalid_api_key_request(status_code, exception=exception): verbose_logger.debug( "Skipping Prometheus metrics for invalid API key request: " f"status_code={status_code}, exception={type(exception).__name__ if exception else None}" ) return True - + return False async def async_post_call_failure_hook( @@ -1686,7 +1683,7 @@ class PrometheusLogger(CustomLogger): exception = request_kwargs.get("exception", None) llm_provider = _litellm_params.get("custom_llm_provider", None) - + if self._should_skip_metrics_for_invalid_key( kwargs=request_kwargs, standard_logging_payload=standard_logging_payload, @@ -2414,8 +2411,8 @@ class PrometheusLogger(CustomLogger): self, user_api_team: Optional[str], user_api_team_alias: Optional[str], - team_spend: float, - team_max_budget: float, + team_spend: Optional[float], + team_max_budget: Optional[float], response_cost: float, ): """ @@ -2577,7 +2574,7 @@ class PrometheusLogger(CustomLogger): user_api_key: Optional[str], user_api_key_alias: Optional[str], response_cost: float, - key_max_budget: float, + key_max_budget: Optional[float], key_spend: Optional[float], ): if user_api_key: @@ -2594,7 +2591,7 @@ class PrometheusLogger(CustomLogger): self, user_api_key: str, user_api_key_alias: str, - key_max_budget: float, + key_max_budget: Optional[float], key_spend: Optional[float], response_cost: float, ) -> UserAPIKeyAuth: From fd2f1481610f03227d6a036749cc8009b18ec8d0 Mon Sep 17 00:00:00 2001 From: Harshit Jain <48647625+Harshit28j@users.noreply.github.com> Date: Tue, 27 Jan 2026 11:41:36 +0530 Subject: [PATCH 12/16] fix: resolve 'does not exist' migration errors as applied in setup_database (#19281) --- .../litellm_proxy_extras/utils.py | 109 +++++++++++------- .../test_litellm_proxy_extras_utils.py | 5 + 2 files changed, 74 insertions(+), 40 deletions(-) diff --git a/litellm-proxy-extras/litellm_proxy_extras/utils.py b/litellm-proxy-extras/litellm_proxy_extras/utils.py index 7ffbe95be13..f3155722187 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/utils.py +++ b/litellm-proxy-extras/litellm_proxy_extras/utils.py @@ -18,14 +18,15 @@ def str_to_bool(value: Optional[str]) -> bool: return value.lower() in ("true", "1", "t", "y", "yes") - def _get_prisma_env() -> dict: """Get environment variables for Prisma, handling offline mode if configured.""" prisma_env = os.environ.copy() if str_to_bool(os.getenv("PRISMA_OFFLINE_MODE")): # These env vars prevent Prisma from attempting downloads prisma_env["NPM_CONFIG_PREFER_OFFLINE"] = "true" - prisma_env["NPM_CONFIG_CACHE"] = os.getenv("NPM_CONFIG_CACHE", "/app/.cache/npm") + prisma_env["NPM_CONFIG_CACHE"] = os.getenv( + "NPM_CONFIG_CACHE", "/app/.cache/npm" + ) return prisma_env @@ -34,29 +35,28 @@ def _get_prisma_command() -> str: if str_to_bool(os.getenv("PRISMA_OFFLINE_MODE")): # Primary location where Prisma Python package installs the CLI default_cli_path = "/app/.cache/prisma-python/binaries/node_modules/.bin/prisma" - + # Check if custom path is provided (for flexibility) custom_cli_path = os.getenv("PRISMA_CLI_PATH") if custom_cli_path and os.path.exists(custom_cli_path): logger.info(f"Using custom Prisma CLI at {custom_cli_path}") return custom_cli_path - + # Check the default location if os.path.exists(default_cli_path): logger.info(f"Using cached Prisma CLI at {default_cli_path}") return default_cli_path - + # If not found, log warning and fall back logger.warning( f"Prisma CLI not found at {default_cli_path}. " "Falling back to Python wrapper (may attempt downloads)" ) - + # Fall back to the Python wrapper (will work in online mode) return "prisma" - class ProxyExtrasDBManager: @staticmethod def _get_prisma_dir() -> str: @@ -119,7 +119,7 @@ class ProxyExtrasDBManager: stdout=open(migration_file, "w"), check=True, timeout=30, - env=prisma_env + env=prisma_env, ) # 3. Mark the migration as applied since it represents current state @@ -134,7 +134,7 @@ class ProxyExtrasDBManager: ], check=True, timeout=30, - env=prisma_env + env=prisma_env, ) return True @@ -159,14 +159,20 @@ class ProxyExtrasDBManager: @staticmethod def _roll_back_migration(migration_name: str): """Mark a specific migration as rolled back""" - # Set up environment for offline mode if configured + # Set up environment for offline mode if configured prisma_env = _get_prisma_env() subprocess.run( - [_get_prisma_command(), "migrate", "resolve", "--rolled-back", migration_name], + [ + _get_prisma_command(), + "migrate", + "resolve", + "--rolled-back", + migration_name, + ], timeout=60, check=True, capture_output=True, - env=prisma_env + env=prisma_env, ) @staticmethod @@ -178,7 +184,7 @@ class ProxyExtrasDBManager: timeout=60, check=True, capture_output=True, - env=prisma_env + env=prisma_env, ) @staticmethod @@ -228,6 +234,8 @@ class ProxyExtrasDBManager: r"duplicate key value violates", r"relation .* already exists", r"constraint .* already exists", + r"does not exist", + r"Can't drop database.* because it doesn't exist", ] for pattern in idempotent_patterns: @@ -248,7 +256,7 @@ class ProxyExtrasDBManager: if not database_url: logger.error("DATABASE_URL not set") return - + diff_dir = ( Path(migrations_dir) / "migrations" @@ -283,7 +291,7 @@ class ProxyExtrasDBManager: check=True, timeout=60, stdout=f, - env=_get_prisma_env() + env=_get_prisma_env(), ) except subprocess.CalledProcessError as e: logger.warning(f"Failed to generate migration diff: {e.stderr}") @@ -313,7 +321,7 @@ class ProxyExtrasDBManager: check=True, capture_output=True, text=True, - env=_get_prisma_env() + env=_get_prisma_env(), ) logger.info(f"prisma db execute stdout: {result.stdout}") logger.info("✅ Migration diff applied successfully") @@ -331,12 +339,18 @@ class ProxyExtrasDBManager: try: logger.info(f"Resolving migration: {migration_name}") subprocess.run( - [_get_prisma_command(), "migrate", "resolve", "--applied", migration_name], + [ + _get_prisma_command(), + "migrate", + "resolve", + "--applied", + migration_name, + ], timeout=60, check=True, capture_output=True, text=True, - env=_get_prisma_env() + env=_get_prisma_env(), ) logger.debug(f"Resolved migration: {migration_name}") except subprocess.CalledProcessError as e: @@ -375,7 +389,7 @@ class ProxyExtrasDBManager: check=True, capture_output=True, text=True, - env=_get_prisma_env() + env=_get_prisma_env(), ) logger.info(f"prisma migrate deploy stdout: {result.stdout}") @@ -397,27 +411,42 @@ class ProxyExtrasDBManager: ) if migration_match: failed_migration = migration_match.group(1) - logger.info( - f"Found failed migration: {failed_migration}, marking as rolled back" - ) - # Mark the failed migration as rolled back - subprocess.run( - [ - _get_prisma_command(), - "migrate", - "resolve", - "--rolled-back", - failed_migration, - ], - timeout=60, - check=True, - capture_output=True, - text=True, - env=_get_prisma_env() - ) - logger.info( - f"✅ Migration {failed_migration} marked as rolled back... retrying" - ) + if ProxyExtrasDBManager._is_idempotent_error(e.stderr): + logger.info( + f"Migration {failed_migration} failed due to idempotent error (e.g., column already exists), resolving as applied" + ) + ProxyExtrasDBManager._roll_back_migration( + failed_migration + ) + ProxyExtrasDBManager._resolve_specific_migration( + failed_migration + ) + logger.info( + f"✅ Migration {failed_migration} resolved." + ) + return True + else: + logger.info( + f"Found failed migration: {failed_migration}, marking as rolled back" + ) + # Mark the failed migration as rolled back + subprocess.run( + [ + _get_prisma_command(), + "migrate", + "resolve", + "--rolled-back", + failed_migration, + ], + timeout=60, + check=True, + capture_output=True, + text=True, + env=_get_prisma_env(), + ) + logger.info( + f"✅ Migration {failed_migration} marked as rolled back... retrying" + ) elif ( "P3005" in e.stderr and "database schema is not empty" in e.stderr diff --git a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py b/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py index 7c151c80ae3..597c0845d43 100644 --- a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py +++ b/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py @@ -97,6 +97,11 @@ class TestIdempotentErrorDetection: error_message = "constraint 'fk_user_id' already exists" assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True + def test_is_idempotent_error_does_not_exist(self): + """Test detection of 'does not exist' error""" + error_message = "ERROR: index 'idx' does not exist" + assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True + def test_is_idempotent_error_case_insensitive(self): """Test that idempotent error detection is case insensitive""" error_message = "COLUMN 'ID' ALREADY EXISTS" From 3f99a91a47e5fc19840d8bd02e181e2de6981e64 Mon Sep 17 00:00:00 2001 From: Milan Date: Tue, 27 Jan 2026 22:05:04 +0200 Subject: [PATCH 13/16] initialize tiktoken environment at import time to support offline usage --- litellm/__init__.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/litellm/__init__.py b/litellm/__init__.py index e5c09702b9b..46124bf8363 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -80,6 +80,10 @@ import dotenv litellm_mode = os.getenv("LITELLM_MODE", "DEV") # "PRODUCTION", "DEV" if litellm_mode == "DEV": dotenv.load_dotenv() + +# Import default_encoding to ensure environment variables are initialized at import time +from litellm.litellm_core_utils import default_encoding # noqa: F401 + #################################################### if set_verbose: _turn_on_debug() From 0b8c10c48828d13317154a25886211f719e6e2f1 Mon Sep 17 00:00:00 2001 From: Xianzong Xie Date: Tue, 27 Jan 2026 15:54:46 -0800 Subject: [PATCH 14/16] Add native_background_mode to override polling_via_cache for specific models This follow-up to PR #16862 allows users to specify models that should use the native provider's background mode instead of polling via cache. Config example: litellm_settings: responses: background_mode: polling_via_cache: ["openai"] native_background_mode: ["o4-mini-deep-research"] ttl: 3600 When a model is in native_background_mode list, should_use_polling_for_request returns False, allowing the request to fall through to native provider handling. Committed-By-Agent: cursor --- litellm/proxy/proxy_server.py | 8 ++++++-- litellm/proxy/response_api_endpoints/endpoints.py | 2 ++ litellm/proxy/response_polling/polling_handler.py | 12 +++++++++++- 3 files changed, 19 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 183c25ed463..cf5d133bc33 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1224,6 +1224,7 @@ redis_usage_cache: Optional[ RedisCache ] = None # redis cache used for tracking spend, tpm/rpm limits polling_via_cache_enabled: Union[Literal["all"], List[str], bool] = False +native_background_mode: List[str] = [] # Models that should use native provider background mode instead of polling polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache user_custom_auth = None user_custom_key_generate = None @@ -2456,14 +2457,17 @@ class ProxyConfig: pass elif key == "responses": # Initialize global polling via cache settings - global polling_via_cache_enabled, polling_cache_ttl + global polling_via_cache_enabled, native_background_mode, polling_cache_ttl background_mode = value.get("background_mode", {}) polling_via_cache_enabled = background_mode.get( "polling_via_cache", False ) + native_background_mode = background_mode.get( + "native_background_mode", [] + ) polling_cache_ttl = background_mode.get("ttl", 3600) verbose_proxy_logger.debug( - f"{blue_color_code} Initialized polling via cache: enabled={polling_via_cache_enabled}, ttl={polling_cache_ttl}{reset_color_code}" + f"{blue_color_code} Initialized polling via cache: enabled={polling_via_cache_enabled}, native_background_mode={native_background_mode}, ttl={polling_cache_ttl}{reset_color_code}" ) elif key == "default_team_settings": for idx, team_setting in enumerate( diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index ec1bc5497bd..44e8c42b2c1 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -68,6 +68,7 @@ async def responses_api( _read_request_body, general_settings, llm_router, + native_background_mode, polling_cache_ttl, polling_via_cache_enabled, proxy_config, @@ -95,6 +96,7 @@ async def responses_api( redis_cache=redis_usage_cache, model=data.get("model", ""), llm_router=llm_router, + native_background_mode=native_background_mode, ) # If polling is enabled, use polling mode diff --git a/litellm/proxy/response_polling/polling_handler.py b/litellm/proxy/response_polling/polling_handler.py index c47578c8d7b..f0b850049bf 100644 --- a/litellm/proxy/response_polling/polling_handler.py +++ b/litellm/proxy/response_polling/polling_handler.py @@ -3,7 +3,7 @@ Response Polling Handler for Background Responses with Cache """ import json from datetime import datetime, timezone -from typing import Any, Dict, Optional +from typing import Any, Dict, List, Optional from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid4 @@ -257,6 +257,7 @@ def should_use_polling_for_request( redis_cache, # RedisCache or None model: str, llm_router, # Router instance or None + native_background_mode: Optional[List[str]] = None, # List of models that should use native background mode ) -> bool: """ Determine if polling via cache should be used for a request. @@ -267,6 +268,8 @@ def should_use_polling_for_request( redis_cache: Redis cache instance (required for polling) model: Model name from the request (e.g., "gpt-5" or "openai/gpt-4o") llm_router: LiteLLM router instance for looking up model deployments + native_background_mode: List of model names that should use native provider + background mode instead of polling via cache Returns: True if polling should be used, False otherwise @@ -275,6 +278,13 @@ def should_use_polling_for_request( if not (background_mode and polling_via_cache_enabled and redis_cache): return False + # Check if model is in native_background_mode list - these use native provider background mode + if native_background_mode and model in native_background_mode: + verbose_proxy_logger.debug( + f"Model {model} is in native_background_mode list, skipping polling via cache" + ) + return False + # "all" enables polling for all providers if polling_via_cache_enabled == "all": return True From f9eea06a371360972373fa880b87ccad01b756f3 Mon Sep 17 00:00:00 2001 From: Xianzong Xie Date: Tue, 27 Jan 2026 16:48:22 -0800 Subject: [PATCH 15/16] Add tests for native_background_mode feature Added 8 new unit tests for the native_background_mode feature: - test_polling_disabled_when_model_in_native_background_mode - test_polling_disabled_for_native_background_mode_with_provider_list - test_polling_enabled_when_model_not_in_native_background_mode - test_polling_enabled_when_native_background_mode_is_none - test_polling_enabled_when_native_background_mode_is_empty_list - test_native_background_mode_exact_match_required - test_native_background_mode_with_provider_prefix_in_request - test_native_background_mode_with_router_lookup Committed-By-Agent: cursor --- .../test_response_polling_handler.py | 136 ++++++++++++++++++ 1 file changed, 136 insertions(+) diff --git a/tests/proxy_unit_tests/test_response_polling_handler.py b/tests/proxy_unit_tests/test_response_polling_handler.py index cb4cd0efe57..26f8ac24adc 100644 --- a/tests/proxy_unit_tests/test_response_polling_handler.py +++ b/tests/proxy_unit_tests/test_response_polling_handler.py @@ -1005,6 +1005,142 @@ class TestPollingConditionChecks: assert result is False + # ==================== Native Background Mode Tests ==================== + + def test_polling_disabled_when_model_in_native_background_mode(self): + """Test that polling is disabled when model is in native_background_mode list""" + from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request + + result = should_use_polling_for_request( + background_mode=True, + polling_via_cache_enabled="all", + redis_cache=Mock(), + model="o4-mini-deep-research", + llm_router=None, + native_background_mode=["o4-mini-deep-research", "o3-deep-research"], + ) + + assert result is False + + def test_polling_disabled_for_native_background_mode_with_provider_list(self): + """Test that native_background_mode takes precedence even when provider matches""" + from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request + + result = should_use_polling_for_request( + background_mode=True, + polling_via_cache_enabled=["openai"], + redis_cache=Mock(), + model="openai/o4-mini-deep-research", + llm_router=None, + native_background_mode=["openai/o4-mini-deep-research"], + ) + + assert result is False + + def test_polling_enabled_when_model_not_in_native_background_mode(self): + """Test that polling is enabled when model is not in native_background_mode list""" + from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request + + result = should_use_polling_for_request( + background_mode=True, + polling_via_cache_enabled="all", + redis_cache=Mock(), + model="gpt-4o", + llm_router=None, + native_background_mode=["o4-mini-deep-research"], + ) + + assert result is True + + def test_polling_enabled_when_native_background_mode_is_none(self): + """Test that polling works normally when native_background_mode is None""" + from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request + + result = should_use_polling_for_request( + background_mode=True, + polling_via_cache_enabled="all", + redis_cache=Mock(), + model="gpt-4o", + llm_router=None, + native_background_mode=None, + ) + + assert result is True + + def test_polling_enabled_when_native_background_mode_is_empty_list(self): + """Test that polling works normally when native_background_mode is empty list""" + from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request + + result = should_use_polling_for_request( + background_mode=True, + polling_via_cache_enabled="all", + redis_cache=Mock(), + model="gpt-4o", + llm_router=None, + native_background_mode=[], + ) + + assert result is True + + def test_native_background_mode_exact_match_required(self): + """Test that native_background_mode uses exact model name matching""" + from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request + + # "o4-mini" should not match "o4-mini-deep-research" + result = should_use_polling_for_request( + background_mode=True, + polling_via_cache_enabled="all", + redis_cache=Mock(), + model="o4-mini", + llm_router=None, + native_background_mode=["o4-mini-deep-research"], + ) + + assert result is True + + def test_native_background_mode_with_provider_prefix_in_request(self): + """Test native_background_mode matching when request model has provider prefix""" + from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request + + # Model in native_background_mode without provider prefix + # Request comes in with provider prefix - should not match + result = should_use_polling_for_request( + background_mode=True, + polling_via_cache_enabled=["openai"], + redis_cache=Mock(), + model="openai/o4-mini-deep-research", + llm_router=None, + native_background_mode=["o4-mini-deep-research"], # Without prefix + ) + + # Should return True because "openai/o4-mini-deep-research" != "o4-mini-deep-research" + assert result is True + + def test_native_background_mode_with_router_lookup(self): + """Test that native_background_mode works with router-resolved models""" + from litellm.proxy.response_polling.polling_handler import should_use_polling_for_request + + mock_router = Mock() + mock_router.model_name_to_deployment_indices = {"deep-research": [0]} + mock_router.model_list = [ + { + "model_name": "deep-research", + "litellm_params": {"model": "openai/o4-mini-deep-research"} + } + ] + + # Model alias "deep-research" is in native_background_mode + result = should_use_polling_for_request( + background_mode=True, + polling_via_cache_enabled=["openai"], + redis_cache=Mock(), + model="deep-research", + llm_router=mock_router, + native_background_mode=["deep-research"], + ) + + assert result is False + class TestStreamingEventParsing: """ From e4cb28aa07846709b3e4818c31825af5623b1d55 Mon Sep 17 00:00:00 2001 From: lizhen Date: Wed, 28 Jan 2026 15:39:01 +0800 Subject: [PATCH 16/16] fix(anthropic): remove explicit cache_control null in tool_result content Fixes issue where tool_result content blocks include explicit 'cache_control': null which breaks some Anthropic API channels. Changes: - Only include cache_control field when explicitly set and not None - Prevents serialization of null values in tool_result text content - Maintains backward compatibility with existing cache_control usage Related issue: Anthropic tool_result conversion adds explicit null values that cause compatibility issues with certain API implementations. Co-Authored-By: Claude (claude-4.5-sonnet) --- .../prompt_templates/factory.py | 17 ++++++++++------- 1 file changed, 10 insertions(+), 7 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 30263543fc6..034089080f7 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1677,13 +1677,16 @@ def convert_to_anthropic_tool_result( ] = [] for content in content_list: if content["type"] == "text": - anthropic_content_list.append( - AnthropicMessagesToolResultContent( - type="text", - text=content["text"], - cache_control=content.get("cache_control", None), - ) - ) + # Only include cache_control if explicitly set and not None + # to avoid sending "cache_control": null which breaks some API channels + text_content: AnthropicMessagesToolResultContent = { + "type": "text", + "text": content["text"], + } + cache_control_value = content.get("cache_control") + if cache_control_value is not None: + text_content["cache_control"] = cache_control_value + anthropic_content_list.append(text_content) elif content["type"] == "image_url": format = ( content["image_url"].get("format")