From 948b05da8646ba15f26bc2355f60d1b17c5ace33 Mon Sep 17 00:00:00 2001 From: tanjiro <56165694+NANDINI-star@users.noreply.github.com> Date: Wed, 27 Aug 2025 16:33:19 +0900 Subject: [PATCH 01/40] added badge --- .../src/components/view_users/user_info_view.tsx | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/src/components/view_users/user_info_view.tsx b/ui/litellm-dashboard/src/components/view_users/user_info_view.tsx index d416df2d9b0..823baaf64ef 100644 --- a/ui/litellm-dashboard/src/components/view_users/user_info_view.tsx +++ b/ui/litellm-dashboard/src/components/view_users/user_info_view.tsx @@ -284,7 +284,17 @@ export default function UserInfoView({ Teams
- {userData.teams?.length || 0} teams + {userData.teams?.length && userData.teams?.length > 0 ? ( +
+ {userData.teams?.map((team, index) => ( + + {team.team_alias} + + ))} +
+ ) : ( + No teams + )}
From 2d0a57a719f0e755ddf31aeb89f7e4cd8815d7de Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=9D=8E=E6=B5=B7=E5=B3=B0?= Date: Thu, 28 Aug 2025 15:05:11 +0800 Subject: [PATCH 02/40] Add Volcengine embedding module with handler and transformation logic - Implemented VolcEngineEmbeddingHandler for synchronous and asynchronous embedding requests. - Created VolcEngineEmbeddingConfig for transforming requests and responses to/from Volcengine format. - Added integration tests for embedding functionality, covering various scenarios including error handling and parameter validation. - Established test structure for Volcengine embedding, ensuring compliance with LiteLLM testing patterns. - Included comprehensive tests for parameter mapping, request transformation, and response handling. --- docs/my-website/docs/providers/volcano.md | 59 ++- litellm/__init__.py | 2 +- litellm/llms/volcengine/__init__.py | 25 + .../chat/transformation.py} | 12 +- litellm/llms/volcengine/common_utils.py | 62 +++ litellm/llms/volcengine/embedding/__init__.py | 8 + litellm/llms/volcengine/embedding/handler.py | 208 ++++++++ .../volcengine/embedding/transformation.py | 245 ++++++++++ litellm/main.py | 47 +- litellm/router.py | 7 +- model_prices_and_context_window.json | 60 +++ .../test_volcengine_embedding.py | 262 ++++++++++ .../test_litellm/llms/volcengine/__init__.py | 1 + .../llms/volcengine/embedding/__init__.py | 1 + .../embedding/test_volcengine_embedding.py | 450 ++++++++++++++++++ .../llms/{ => volcengine}/test_volcengine.py | 2 +- tests/test_litellm/test_utils.py | 2 +- 17 files changed, 1438 insertions(+), 15 deletions(-) create mode 100644 litellm/llms/volcengine/__init__.py rename litellm/llms/{volcengine.py => volcengine/chat/transformation.py} (91%) create mode 100644 litellm/llms/volcengine/common_utils.py create mode 100644 litellm/llms/volcengine/embedding/__init__.py create mode 100644 litellm/llms/volcengine/embedding/handler.py create mode 100644 litellm/llms/volcengine/embedding/transformation.py create mode 100644 tests/llm_translation/test_volcengine_embedding.py create mode 100644 tests/test_litellm/llms/volcengine/__init__.py create mode 100644 tests/test_litellm/llms/volcengine/embedding/__init__.py create mode 100644 tests/test_litellm/llms/volcengine/embedding/test_volcengine_embedding.py rename tests/test_litellm/llms/{ => volcengine}/test_volcengine.py (97%) diff --git a/docs/my-website/docs/providers/volcano.md b/docs/my-website/docs/providers/volcano.md index 1742a43d819..efd1e02b60b 100644 --- a/docs/my-website/docs/providers/volcano.md +++ b/docs/my-website/docs/providers/volcano.md @@ -3,7 +3,7 @@ https://www.volcengine.com/docs/82379/1263482 :::tip -**We support ALL Volcengine NIM models, just set `model=volcengine/` as a prefix when sending litellm requests** +**We support ALL Volcengine models including Chat and Embeddings, just set `model=volcengine/` as a prefix when sending litellm requests** ::: @@ -11,6 +11,8 @@ https://www.volcengine.com/docs/82379/1263482 ```python # env variable os.environ['VOLCENGINE_API_KEY'] +# or +os.environ['ARK_API_KEY'] ``` ## Sample Usage @@ -64,9 +66,42 @@ for chunk in response: print(chunk) ``` +## Sample Usage - Embedding +```python +from litellm import embedding +import os -## Supported Models - πŸ’₯ ALL Volcengine NIM Models Supported! -We support ALL `volcengine` models, just set `volcengine/` as a prefix when sending completion requests +os.environ['VOLCENGINE_API_KEY'] = "" +response = embedding( + model="volcengine/doubao-embedding-text-240715", + input=["hello world", "good morning"] +) +print(response) +``` + +### Supported Embedding Models +- `doubao-embedding-large` (2048 dimensions) +- `doubao-embedding-large-text-250515` (2048 dimensions) +- `doubao-embedding-large-text-240915` (4096 dimensions) +- `doubao-embedding` (2560 dimensions) +- `doubao-embedding-text-240715` (2560 dimensions) + +### Embedding Parameters +```python +from litellm import embedding + +response = embedding( + model="volcengine/doubao-embedding-text-240715", + input=["sample text"], + encoding_format="float", # optional: "float" (default), "base64" + user="user-123", # optional: user identifier for tracking +) +``` + +## Supported Models - πŸ’₯ ALL Volcengine Models Supported! +We support ALL `volcengine` models for both chat completions and embeddings: +- **Chat Models**: Set `volcengine/` as a prefix when sending completion requests +- **Embedding Models**: Use the specific model names listed above (e.g., `volcengine/doubao-embedding-text-240715`) ## Sample Usage - LiteLLM Proxy @@ -74,14 +109,21 @@ We support ALL `volcengine` models, just set `volcengine/` as a ```yaml model_list: + # Chat model - model_name: volcengine-model litellm_params: model: volcengine/ api_key: os.environ/VOLCENGINE_API_KEY + # Embedding model + - model_name: volcengine-embedding + litellm_params: + model: volcengine/doubao-embedding-text-240715 + api_key: os.environ/VOLCENGINE_API_KEY ``` ### Send Request +#### Chat Completion ```shell curl --location 'http://localhost:4000/chat/completions' \ --header 'Authorization: Bearer sk-1234' \ @@ -95,4 +137,15 @@ curl --location 'http://localhost:4000/chat/completions' \ } ] }' +``` + +#### Embedding +```shell +curl --location 'http://localhost:4000/embeddings' \ + --header 'Authorization: Bearer sk-1234' \ + --header 'Content-Type: application/json' \ + --data '{ + "model": "volcengine-embedding", + "input": ["hello world", "good morning"] +}' ``` \ No newline at end of file diff --git a/litellm/__init__.py b/litellm/__init__.py index c405d3cdeb2..2416d1ee0d6 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1215,7 +1215,7 @@ from .llms.jina_ai.embedding.transformation import JinaAIEmbeddingConfig from .llms.xai.chat.transformation import XAIChatConfig from .llms.xai.common_utils import XAIModelInfo from .llms.aiml.chat.transformation import AIMLChatConfig -from .llms.volcengine import VolcEngineConfig +from .llms.volcengine.chat.transformation import VolcEngineChatConfig as VolcEngineConfig from .llms.codestral.completion.transformation import CodestralTextCompletionConfig from .llms.azure.azure import ( AzureOpenAIError, diff --git a/litellm/llms/volcengine/__init__.py b/litellm/llms/volcengine/__init__.py new file mode 100644 index 00000000000..0be9a4f428c --- /dev/null +++ b/litellm/llms/volcengine/__init__.py @@ -0,0 +1,25 @@ +""" +Volcengine LLM Provider +Support for Volcengine (ByteDance) chat and embedding models +""" + +from .chat.transformation import VolcEngineChatConfig +from .embedding import VolcEngineEmbeddingHandler, VolcEngineEmbeddingConfig +from .common_utils import ( + VolcEngineError, + get_volcengine_base_url, + get_volcengine_headers, +) + +# For backward compatibility, keep the old class name +VolcEngineConfig = VolcEngineChatConfig + +__all__ = [ + "VolcEngineChatConfig", + "VolcEngineConfig", # backward compatibility + "VolcEngineEmbeddingHandler", + "VolcEngineEmbeddingConfig", + "VolcEngineError", + "get_volcengine_base_url", + "get_volcengine_headers", +] diff --git a/litellm/llms/volcengine.py b/litellm/llms/volcengine/chat/transformation.py similarity index 91% rename from litellm/llms/volcengine.py rename to litellm/llms/volcengine/chat/transformation.py index c878aaf933c..216570a1aba 100644 --- a/litellm/llms/volcengine.py +++ b/litellm/llms/volcengine/chat/transformation.py @@ -3,7 +3,7 @@ from typing import Optional, Union from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig -class VolcEngineConfig(OpenAILikeChatConfig): +class VolcEngineChatConfig(OpenAILikeChatConfig): frequency_penalty: Optional[int] = None function_call: Optional[Union[str, dict]] = None functions: Optional[list] = None @@ -82,17 +82,19 @@ class VolcEngineConfig(OpenAILikeChatConfig): if "thinking" in optional_params: thinking_value = optional_params.pop("thinking") - + # Handle disabled thinking case - don't add to extra_body if disabled if ( - thinking_value is not None - and isinstance(thinking_value, dict) + thinking_value is not None + and isinstance(thinking_value, dict) and thinking_value.get("type") == "disabled" ): # Skip adding thinking parameter when it's disabled pass else: # Add thinking parameter to extra_body for all other cases - optional_params.setdefault("extra_body", {})["thinking"] = thinking_value + optional_params.setdefault("extra_body", {})[ + "thinking" + ] = thinking_value return optional_params diff --git a/litellm/llms/volcengine/common_utils.py b/litellm/llms/volcengine/common_utils.py new file mode 100644 index 00000000000..0c8d3daebdc --- /dev/null +++ b/litellm/llms/volcengine/common_utils.py @@ -0,0 +1,62 @@ +""" +Common utilities for Volcengine LLM provider +""" + +from typing import Optional + +import httpx + +from litellm.llms.base_llm.chat.transformation import BaseLLMException + + +class VolcEngineError(BaseLLMException): + """ + Custom exception class for Volcengine provider errors. + """ + + def __init__( + self, status_code: int, message: str, headers: Optional[httpx.Headers] = None + ): + self.status_code = status_code + self.message = message + self.headers = headers or httpx.Headers() + super().__init__( + status_code=status_code, message=message, headers=dict(self.headers) + ) + + +def get_volcengine_base_url(api_base: Optional[str] = None) -> str: + """ + Get the base URL for Volcengine API calls. + + Args: + api_base: Optional custom API base URL + + Returns: + The base URL to use for API calls + """ + if api_base: + return api_base + return "https://ark.cn-beijing.volces.com" + + +def get_volcengine_headers(api_key: str, extra_headers: Optional[dict] = None) -> dict: + """ + Get headers for Volcengine API calls. + + Args: + api_key: The API key for authentication + extra_headers: Optional additional headers + + Returns: + Dictionary of headers + """ + headers = { + "Content-Type": "application/json", + "Authorization": f"Bearer {api_key}", + } + + if extra_headers: + headers.update(extra_headers) + + return headers diff --git a/litellm/llms/volcengine/embedding/__init__.py b/litellm/llms/volcengine/embedding/__init__.py new file mode 100644 index 00000000000..6063e88b740 --- /dev/null +++ b/litellm/llms/volcengine/embedding/__init__.py @@ -0,0 +1,8 @@ +""" +Volcengine Embedding Module +""" + +from .handler import VolcEngineEmbeddingHandler +from .transformation import VolcEngineEmbeddingConfig + +__all__ = ["VolcEngineEmbeddingHandler", "VolcEngineEmbeddingConfig"] diff --git a/litellm/llms/volcengine/embedding/handler.py b/litellm/llms/volcengine/embedding/handler.py new file mode 100644 index 00000000000..e29f920afe8 --- /dev/null +++ b/litellm/llms/volcengine/embedding/handler.py @@ -0,0 +1,208 @@ +""" +Volcengine Embedding Handler +Handles embedding requests to Volcengine's embedding API +""" + +from typing import Dict, List, Optional, Union, Any + +import httpx +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler +from litellm.types.utils import EmbeddingResponse +import litellm + +from .transformation import VolcEngineEmbeddingConfig +from ..common_utils import VolcEngineError + + +class VolcEngineEmbeddingHandler: + """Handler for Volcengine embedding API calls""" + + def __init__(self): + self.config = VolcEngineEmbeddingConfig() + + def _convert_to_litellm_response(self, transformed_response: Dict, model: str, input: Union[str, List[str]]) -> EmbeddingResponse: + """Convert transformed response to LiteLLM EmbeddingResponse""" + model_response = EmbeddingResponse() + model_response.object = transformed_response.get("object", "list") + model_response.data = transformed_response.get("data", []) + model_response.model = transformed_response.get("model", model) + + # Set usage information + usage_data = transformed_response.get("usage", {}) + if usage_data: + model_response.usage = litellm.Usage( + prompt_tokens=usage_data.get("prompt_tokens", 0), + completion_tokens=0, + total_tokens=usage_data.get("total_tokens", usage_data.get("prompt_tokens", 0)), + prompt_tokens_details=None, + completion_tokens_details=None, + ) + + return model_response + + def embedding( + self, + model: str, + input: Union[str, List[str]], + api_key: str, + api_base: Optional[str] = None, + encoding_format: Optional[str] = "float", + user: Optional[str] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: Optional[Dict[str, str]] = None, + litellm_logging_obj: Optional[LiteLLMLoggingObj] = None, + **kwargs, + ) -> EmbeddingResponse: + """ + Synchronous embedding call to Volcengine API. + + Args: + model: Volcengine model ID (e.g., "doubao-embedding-text-240715") + input: Text or list of texts to embed + api_key: Volcengine API key + api_base: Optional custom API base URL + encoding_format: Response format (float, base64, null) + user: Optional user identifier + timeout: Request timeout + extra_headers: Optional additional headers + litellm_logging_obj: Optional logging object + **kwargs: Additional parameters + + Returns: + EmbeddingResponse object + """ + # Transform request to Volcengine format + request_data = self.config.transform_request( + model=model, + input=input, + api_key=api_key, + api_base=api_base, + encoding_format=encoding_format, + user=user, + extra_headers=extra_headers, + **kwargs, + ) + + # Make HTTP request + try: + client = HTTPHandler(timeout=timeout) + response = client.post( + url=request_data["url"], + headers=request_data["headers"], + json=request_data["data"], + ) + except Exception as e: + raise VolcEngineError( + status_code=500, + message=f"Network error during embedding request: {str(e)}", + ) + + # Handle HTTP errors + if response.status_code != 200: + error_message = f"Volcengine embedding request failed with status {response.status_code}" + try: + error_details = response.json() + if "error" in error_details: + error_message += f": {error_details['error']}" + elif "message" in error_details: + error_message += f": {error_details['message']}" + except Exception: + error_message += f": {response.text}" + + raise VolcEngineError( + status_code=response.status_code, + message=error_message, + headers=response.headers, + ) + + # Transform response to OpenAI format + transformed_response = self.config.transform_response( + response=response, model=model, input=input, encoding=encoding_format + ) + + # Convert to LiteLLM EmbeddingResponse + return self._convert_to_litellm_response(transformed_response, model, input) + + async def async_embedding( + self, + model: str, + input: Union[str, List[str]], + api_key: str, + api_base: Optional[str] = None, + encoding_format: Optional[str] = "float", + user: Optional[str] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + extra_headers: Optional[Dict[str, str]] = None, + litellm_logging_obj: Optional[LiteLLMLoggingObj] = None, + **kwargs, + ) -> EmbeddingResponse: + """ + Asynchronous embedding call to Volcengine API. + + Args: + model: Volcengine model ID (e.g., "doubao-embedding-text-240715") + input: Text or list of texts to embed + api_key: Volcengine API key + api_base: Optional custom API base URL + encoding_format: Response format (float, base64, null) + user: Optional user identifier + timeout: Request timeout + extra_headers: Optional additional headers + litellm_logging_obj: Optional logging object + **kwargs: Additional parameters + + Returns: + EmbeddingResponse object + """ + # Transform request to Volcengine format + request_data = self.config.transform_request( + model=model, + input=input, + api_key=api_key, + api_base=api_base, + encoding_format=encoding_format, + user=user, + extra_headers=extra_headers, + **kwargs, + ) + + # Make async HTTP request + try: + client = AsyncHTTPHandler(timeout=timeout) + response = await client.post( + url=request_data["url"], + headers=request_data["headers"], + json=request_data["data"], + ) + except Exception as e: + raise VolcEngineError( + status_code=500, + message=f"Network error during embedding request: {str(e)}", + ) + + # Handle HTTP errors + if response.status_code != 200: + error_message = f"Volcengine embedding request failed with status {response.status_code}" + try: + error_details = response.json() + if "error" in error_details: + error_message += f": {error_details['error']}" + elif "message" in error_details: + error_message += f": {error_details['message']}" + except Exception: + error_message += f": {response.text}" + + raise VolcEngineError( + status_code=response.status_code, + message=error_message, + headers=response.headers, + ) + + # Transform response to OpenAI format + transformed_response = self.config.transform_response( + response=response, model=model, input=input, encoding=encoding_format + ) + + # Convert to LiteLLM EmbeddingResponse + return self._convert_to_litellm_response(transformed_response, model, input) diff --git a/litellm/llms/volcengine/embedding/transformation.py b/litellm/llms/volcengine/embedding/transformation.py new file mode 100644 index 00000000000..ba2f07a4945 --- /dev/null +++ b/litellm/llms/volcengine/embedding/transformation.py @@ -0,0 +1,245 @@ +""" +Volcengine Embedding Transformation +Transforms OpenAI embedding requests to Volcengine format +""" + +from typing import List, Optional, Union, Dict, Any +import httpx +from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues +from litellm.types.utils import EmbeddingResponse +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from ..common_utils import get_volcengine_base_url, get_volcengine_headers + + +class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): + """ + Configuration class for Volcengine embedding models. + Reference: https://ark.cn-beijing.volces.com/api/v3/embeddings + """ + + def __init__( + self, + encoding_format: Optional[str] = None, + ) -> None: + locals_ = locals().copy() + for key, value in locals_.items(): + if key != "self" and value is not None: + setattr(self.__class__, key, value) + + @classmethod + def get_config(cls): + return super().get_config() + + def get_supported_openai_params(self, model: str) -> List[str]: + """ + Get the list of OpenAI parameters supported by Volcengine embedding models. + + Args: + model: The model name + + Returns: + List of supported parameter names + """ + return [ + "encoding_format", + "user", + "extra_headers", + ] + + def map_openai_params( + self, + non_default_params: Dict[str, Any], + optional_params: Dict[str, Any], + model: str, + drop_params: bool, + ) -> Dict[str, Any]: + """ + Map OpenAI embedding parameters to Volcengine format. + + Args: + non_default_params: Parameters that are not default values + optional_params: Optional parameters dict to update + model: The model name + drop_params: Whether to drop unsupported parameters + + Returns: + Updated optional_params dict + """ + for param, value in non_default_params.items(): + if param == "encoding_format": + # Volcengine supports: float, base64, null + if value in ["float", "base64", None]: + optional_params["encoding_format"] = value + else: + if not drop_params: + raise ValueError( + f"Unsupported encoding_format: {value}. Volcengine supports: float, base64, null" + ) + elif param == "user": + # Keep user parameter as-is + optional_params["user"] = value + elif param in self.get_supported_openai_params(model): + optional_params[param] = value + elif not drop_params: + raise ValueError(f"Unsupported parameter for Volcengine: {param}") + + return optional_params + + def transform_request( + self, + model: str, + input: Union[str, List[str]], + api_key: str, + api_base: Optional[str] = None, + encoding_format: Optional[str] = "float", + user: Optional[str] = None, + extra_headers: Optional[Dict[str, str]] = None, + **kwargs, + ) -> Dict[str, Any]: + """ + Transform OpenAI embedding request to Volcengine format. + + Args: + model: Model ID (e.g., "doubao-embedding-text-240715") + input: Text or list of texts to embed + api_key: Volcengine API key + api_base: Optional custom API base URL + encoding_format: Response format (float, base64, null) + user: Optional user identifier + extra_headers: Optional additional headers + **kwargs: Additional parameters + + Returns: + Dict containing url, headers, and data for the request + """ + # Get base URL + base_url = get_volcengine_base_url(api_base) + # Avoid duplicate /api/v3 if base_url already contains it + if base_url.endswith("/api/v3"): + url = f"{base_url}/embeddings" + else: + url = f"{base_url}/api/v3/embeddings" + + # Get headers + headers = get_volcengine_headers(api_key, extra_headers) + + # Prepare request data + data = { + "model": model, + "input": input if isinstance(input, list) else [input], + } + + # Add optional parameters + if encoding_format is not None: + data["encoding_format"] = encoding_format + + return { + "url": url, + "headers": headers, + "data": data, + } + + def transform_response( + self, + response: httpx.Response, + model: str, + input: Union[str, List[str]], + encoding: Optional[str] = None, + ) -> Dict[str, Any]: + """ + Transform Volcengine embedding response to OpenAI format. + + Args: + response: The HTTP response from Volcengine + model: The model used + input: The input that was embedded + encoding: The encoding format requested + + Returns: + OpenAI-compatible embedding response + """ + try: + response_json = response.json() + except Exception as e: + raise ValueError(f"Failed to parse Volcengine response as JSON: {str(e)}") + + # Volcengine response format matches OpenAI format closely + # Just need to ensure all required fields are present + transformed_response = { + "object": "list", + "data": response_json.get("data", []), + "model": response_json.get("model", model), + "usage": response_json.get("usage", {}), + } + + # Add id if present + if "id" in response_json: + transformed_response["id"] = response_json["id"] + + return transformed_response + + def transform_embedding_request( + self, + model: str, + input: AllEmbeddingInputValues, + optional_params: dict, + headers: dict, + ) -> dict: + """Transform embedding request to Volcengine format""" + # Use existing transform_request method + return self.transform_request( + model=model, + input=input, + api_key="", # api_key will be in headers + **optional_params, + ) + + def transform_embedding_response( + self, + model: str, + raw_response: httpx.Response, + model_response: EmbeddingResponse, + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str], + request_data: dict, + optional_params: dict, + litellm_params: dict, + ) -> EmbeddingResponse: + """Transform Volcengine response to EmbeddingResponse""" + # Use existing transform_response method + transformed_response = self.transform_response( + response=raw_response, + model=model, + input=request_data.get("input", []), + ) + + # Create EmbeddingResponse from transformed data + return EmbeddingResponse(**transformed_response) + + def validate_environment( + self, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + """Validate environment and return headers""" + # Get Volcengine headers + volcengine_headers = get_volcengine_headers(api_key) + return {**headers, **volcengine_headers} + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> BaseLLMException: + """Get error class for Volcengine errors""" + from ..common_utils import VolcEngineError + return VolcEngineError( + status_code=status_code, + message=error_message, + headers=headers, + ) diff --git a/litellm/main.py b/litellm/main.py index 6102fe3ccce..776e81a0110 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -183,6 +183,7 @@ from .llms.vertex_ai.text_to_speech.text_to_speech_handler import VertexTextToSp from .llms.vertex_ai.vertex_ai_partner_models.main import VertexAIPartnerModels from .llms.vertex_ai.vertex_embeddings.embedding_handler import VertexEmbedding from .llms.vertex_ai.vertex_model_garden.main import VertexAIModelGardenModels +from .llms.volcengine.embedding.handler import VolcEngineEmbeddingHandler from .llms.vllm.completion import handler as vllm_handler from .llms.watsonx.chat.handler import WatsonXChatHandler from .llms.watsonx.common_utils import IBMWatsonXMixin @@ -500,7 +501,7 @@ async def acompletion( } if custom_llm_provider is None: _, custom_llm_provider, _, _ = get_llm_provider( - model=model, api_base=completion_kwargs.get("base_url", None) + model=model, custom_llm_provider=custom_llm_provider, api_base=completion_kwargs.get("base_url", None) ) fallbacks = fallbacks or litellm.model_fallbacks @@ -3582,7 +3583,7 @@ async def aembedding(*args, **kwargs) -> EmbeddingResponse: model = args[0] if len(args) > 0 else kwargs["model"] ### PASS ARGS TO Embedding ### kwargs["aembedding"] = True - custom_llm_provider = None + custom_llm_provider = kwargs.get("custom_llm_provider", None) try: # Use a partial function to pass your keyword arguments func = partial(embedding, *args, **kwargs) @@ -3592,7 +3593,7 @@ async def aembedding(*args, **kwargs) -> EmbeddingResponse: func_with_context = partial(ctx.run, func) _, custom_llm_provider, _, _ = get_llm_provider( - model=model, api_base=kwargs.get("api_base", None) + model=model, custom_llm_provider=custom_llm_provider, api_base=kwargs.get("api_base", None) ) # Await normally @@ -4414,6 +4415,46 @@ def embedding( # noqa: PLR0915 client=client, aembedding=aembedding, ) + elif custom_llm_provider == "volcengine": + api_key = ( + api_key + or litellm.api_key + or get_secret_str("ARK_API_KEY") + or get_secret_str("VOLCENGINE_API_KEY") + ) + if api_key is None: + raise ValueError( + "Missing API key for Volcengine. Set ARK_API_KEY or VOLCENGINE_API_KEY environment variable or pass api_key parameter." + ) + + handler = VolcEngineEmbeddingHandler() + + if aembedding: + response = handler.async_embedding( + model=model, + input=input, + api_key=api_key, + api_base=api_base, + encoding_format=optional_params.get("encoding_format", "float"), + user=optional_params.get("user"), + timeout=timeout, + extra_headers=optional_params.get("extra_headers"), + litellm_logging_obj=logging, + **optional_params, + ) + else: + response = handler.embedding( + model=model, + input=input, + api_key=api_key, + api_base=api_base, + encoding_format=optional_params.get("encoding_format", "float"), + user=optional_params.get("user"), + timeout=timeout, + extra_headers=optional_params.get("extra_headers"), + litellm_logging_obj=logging, + **optional_params, + ) elif custom_llm_provider in litellm._custom_providers: custom_handler: Optional[CustomLLM] = None for item in litellm.custom_provider_map: diff --git a/litellm/router.py b/litellm/router.py index 190d19598c3..1ed95ee7b29 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -5658,6 +5658,11 @@ class Router: ) if supported_openai_params is None: supported_openai_params = [] + + # Get mode from database model_info if available, otherwise default to "chat" + db_model_info = model.get("model_info", {}) + mode = db_model_info.get("mode", "chat") + model_info = ModelMapInfo( key=model_group, max_tokens=None, @@ -5666,7 +5671,7 @@ class Router: input_cost_per_token=0, output_cost_per_token=0, litellm_provider=llm_provider, - mode="chat", + mode=mode, supported_openai_params=supported_openai_params, supports_system_messages=None, ) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index a29d03f2c6c..d7903c9a0a7 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -19746,5 +19746,65 @@ "metadata": { "notes": "DALL-E 2 via AI/ML API - Reliable text-to-image generation" } + }, + "doubao-embedding-large": { + "max_tokens": 4096, + "max_input_tokens": 4096, + "output_vector_size": 2048, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "mode": "embedding", + "metadata": { + "notes": "Volcengine Doubao embedding model - large version with 2048 dimensions" + } + }, + "doubao-embedding-large-text-250515": { + "max_tokens": 4096, + "max_input_tokens": 4096, + "output_vector_size": 2048, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "mode": "embedding", + "metadata": { + "notes": "Volcengine Doubao embedding model - text-250515 version with 2048 dimensions" + } + }, + "doubao-embedding-large-text-240915": { + "max_tokens": 4096, + "max_input_tokens": 4096, + "output_vector_size": 4096, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "mode": "embedding", + "metadata": { + "notes": "Volcengine Doubao embedding model - text-240915 version with 4096 dimensions" + } + }, + "doubao-embedding": { + "max_tokens": 4096, + "max_input_tokens": 4096, + "output_vector_size": 2560, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "mode": "embedding", + "metadata": { + "notes": "Volcengine Doubao embedding model - standard version with 2560 dimensions" + } + }, + "doubao-embedding-text-240715": { + "max_tokens": 4096, + "max_input_tokens": 4096, + "output_vector_size": 2560, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "mode": "embedding", + "metadata": { + "notes": "Volcengine Doubao embedding model - text-240715 version with 2560 dimensions" + } } } diff --git a/tests/llm_translation/test_volcengine_embedding.py b/tests/llm_translation/test_volcengine_embedding.py new file mode 100644 index 00000000000..9503d91f3c8 --- /dev/null +++ b/tests/llm_translation/test_volcengine_embedding.py @@ -0,0 +1,262 @@ +""" +Integration tests for Volcengine embedding following LiteLLM testing patterns +Based on the BaseLLMEmbeddingTest framework +""" + +import os +import sys +from unittest.mock import MagicMock, patch +import pytest + +# Add parent directory to path for imports +sys.path.insert(0, os.path.abspath("../..")) + +from base_embedding_unit_tests import BaseLLMEmbeddingTest +import litellm +from litellm.types.utils import EmbeddingResponse + + +class TestVolcEngineEmbedding(BaseLLMEmbeddingTest): + """Test Volcengine embedding integration following LiteLLM patterns""" + + def get_custom_llm_provider(self) -> litellm.LlmProviders: + return litellm.LlmProviders.VOLCENGINE + + def get_base_embedding_call_args(self) -> dict: + return { + "model": "volcengine/doubao-embedding-text-240715", + } + + @pytest.mark.asyncio() + @pytest.mark.parametrize("sync_mode", [True, False]) + async def test_basic_embedding(self, sync_mode): + """Test basic embedding functionality with realistic response""" + litellm.set_verbose = True + embedding_call_args = self.get_base_embedding_call_args() + + # Mock the embedding functions to avoid actual API calls + with patch("litellm.embedding") as mock_embedding, patch("litellm.aembedding") as mock_aembedding: + # Create realistic Volcengine response + mock_response = MagicMock() + mock_response.model = "doubao-embedding-text-240715" + mock_response.object = "list" + mock_response.data = [ + { + "object": "embedding", + "embedding": [0.1, 0.2, 0.3] + [0.01 * i for i in range(1021)], # 1024-dim embedding + "index": 0 + }, + { + "object": "embedding", + "embedding": [0.4, 0.5, 0.6] + [0.02 * i for i in range(1021)], # 1024-dim embedding + "index": 1 + } + ] + mock_response.usage.prompt_tokens = 2 + mock_response.usage.total_tokens = 2 + + mock_embedding.return_value = mock_response + mock_aembedding.return_value = mock_response + + # Test sync mode + if sync_mode is True: + response = litellm.embedding( + **embedding_call_args, + input=["hello", "world"], + ) + + # Verify response structure matches Volcengine format + assert response.model == "doubao-embedding-text-240715" + assert response.object == "list" + assert len(response.data) == 2 + assert len(response.data[0]["embedding"]) == 1024 + assert response.usage.total_tokens > 0 + + # Test async mode + else: + response = await litellm.aembedding( + **embedding_call_args, + input=["hello", "world"], + ) + + # Verify response structure + assert response.model == "doubao-embedding-text-240715" + assert response.object == "list" + assert len(response.data) == 2 + assert len(response.data[0]["embedding"]) == 1024 + assert response.usage.total_tokens > 0 + + +def test_volcengine_embedding_with_encoding_formats(): + """Test Volcengine embedding with different encoding formats""" + + test_cases = [ + {"encoding_format": "float"}, + {"encoding_format": "base64"}, + {"encoding_format": None}, # Default + ] + + for params in test_cases: + with patch("litellm.embedding") as mock_embedding: + # Create mock response based on encoding format + mock_response = MagicMock() + mock_response.model = "doubao-embedding-text-240715" + mock_response.object = "list" + + if params["encoding_format"] == "base64": + # Simulate base64 encoded embeddings + mock_response.data = [ + { + "object": "embedding", + "embedding": "c29tZS1iYXNlNjQtZW5jb2RlZC1lbWJlZGRpbmc=", # base64 encoded + "index": 0 + } + ] + else: + # Float embeddings (default) + mock_response.data = [ + { + "object": "embedding", + "embedding": [0.1, 0.2, 0.3, -0.1] * 256, # 1024 dimensions + "index": 0 + } + ] + + mock_response.usage.prompt_tokens = 3 + mock_response.usage.total_tokens = 3 + mock_embedding.return_value = mock_response + + # Test the call + litellm.embedding( + model="volcengine/doubao-embedding-text-240715", + input=["test text"], + **params + ) + + # Verify the call was made with correct parameters + mock_embedding.assert_called_once() + call_args = mock_embedding.call_args + assert call_args[1]["model"] == "volcengine/doubao-embedding-text-240715" + assert call_args[1]["input"] == ["test text"] + + if params["encoding_format"] is not None: + assert call_args[1]["encoding_format"] == params["encoding_format"] + + +def test_volcengine_embedding_with_user_parameter(): + """Test Volcengine embedding with user parameter for tracking""" + + with patch("litellm.embedding") as mock_embedding: + mock_response = MagicMock() + mock_response.model = "doubao-embedding-text-240715" + mock_response.object = "list" + mock_response.data = [ + { + "object": "embedding", + "embedding": [0.1] * 1024, + "index": 0 + } + ] + mock_response.usage.prompt_tokens = 5 + mock_response.usage.total_tokens = 5 + mock_embedding.return_value = mock_response + + # Test with user parameter + litellm.embedding( + model="volcengine/doubao-embedding-text-240715", + input=["user tracking test"], + user="test-user-12345" + ) + + # Verify user parameter was passed + mock_embedding.assert_called_once() + call_args = mock_embedding.call_args + assert call_args[1]["user"] == "test-user-12345" + + +def test_volcengine_embedding_error_scenarios(): + """Test Volcengine embedding error handling in integration context""" + + error_scenarios = [ + # Invalid model name + { + "model": "volcengine/invalid-model-name", + "expected_error_pattern": "model" + }, + # Invalid encoding format + { + "model": "volcengine/doubao-embedding-text-240715", + "encoding_format": "invalid_format", + "expected_error_pattern": "encoding_format" + } + ] + + for scenario in error_scenarios: + with patch("litellm.embedding") as mock_embedding: + # Configure mock to raise appropriate errors + if "invalid-model" in scenario.get("model", ""): + mock_embedding.side_effect = Exception("Model not found") + elif scenario.get("encoding_format") == "invalid_format": + mock_embedding.side_effect = ValueError("Unsupported encoding_format") + + # Test that errors are properly raised + with pytest.raises(Exception) as exc_info: + test_params = {k: v for k, v in scenario.items() if k != "expected_error_pattern"} + litellm.embedding( + input=["test"], + **test_params + ) + + # Verify error message contains expected pattern + assert scenario["expected_error_pattern"].lower() in str(exc_info.value).lower() + + +def test_volcengine_embedding_with_multiple_inputs(): + """Test Volcengine embedding with various input lengths and types""" + + test_inputs = [ + # Single short text + ["hello"], + # Multiple short texts + ["hello", "world", "test"], + # Mixed length texts + ["short", "This is a much longer text that should be handled properly by the embedding service"], + # Unicode content + ["ζ΅‹θ―•δΈ­ζ–‡ζ–‡ζœ¬", "Test English text", "ζ··εˆθ―­θ¨€ mixed language"], + # Many inputs (batch processing) + [f"Test sentence number {i}" for i in range(10)] + ] + + for test_input in test_inputs: + with patch("litellm.embedding") as mock_embedding: + # Create proportional mock response + mock_response = MagicMock() + mock_response.model = "doubao-embedding-text-240715" + mock_response.object = "list" + mock_response.data = [ + { + "object": "embedding", + "embedding": [0.1 * (i + 1)] * 1024, # Unique embedding per input + "index": i + } + for i in range(len(test_input)) + ] + mock_response.usage.prompt_tokens = len(test_input) * 5 # Realistic token estimate + mock_response.usage.total_tokens = len(test_input) * 5 + mock_embedding.return_value = mock_response + + # Test the call + response = litellm.embedding( + model="volcengine/doubao-embedding-text-240715", + input=test_input + ) + + # Verify response matches input count + assert len(response.data) == len(test_input) + for i, embedding_data in enumerate(response.data): + assert embedding_data["index"] == i + assert len(embedding_data["embedding"]) == 1024 + + +if __name__ == "__main__": + pytest.main([__file__]) \ No newline at end of file diff --git a/tests/test_litellm/llms/volcengine/__init__.py b/tests/test_litellm/llms/volcengine/__init__.py new file mode 100644 index 00000000000..6ac3aa6b71a --- /dev/null +++ b/tests/test_litellm/llms/volcengine/__init__.py @@ -0,0 +1 @@ +# Volcengine tests \ No newline at end of file diff --git a/tests/test_litellm/llms/volcengine/embedding/__init__.py b/tests/test_litellm/llms/volcengine/embedding/__init__.py new file mode 100644 index 00000000000..bb087ba3563 --- /dev/null +++ b/tests/test_litellm/llms/volcengine/embedding/__init__.py @@ -0,0 +1 @@ +# Volcengine embedding tests \ No newline at end of file diff --git a/tests/test_litellm/llms/volcengine/embedding/test_volcengine_embedding.py b/tests/test_litellm/llms/volcengine/embedding/test_volcengine_embedding.py new file mode 100644 index 00000000000..f2f143b5b99 --- /dev/null +++ b/tests/test_litellm/llms/volcengine/embedding/test_volcengine_embedding.py @@ -0,0 +1,450 @@ +""" +Improved tests for Volcengine Embedding functionality +Tests real business logic without excessive mocking +""" + +import pytest +import json +import httpx +from unittest.mock import Mock, patch, MagicMock +from typing import List, Dict, Any + +from litellm.llms.volcengine.embedding import VolcEngineEmbeddingHandler, VolcEngineEmbeddingConfig +from litellm.llms.volcengine.common_utils import VolcEngineError +from litellm.types.utils import EmbeddingResponse +from litellm.types.llms.openai import AllEmbeddingInputValues + + +class TestVolcEngineEmbeddingConfigBusinessLogic: + """Test real business logic of VolcEngineEmbeddingConfig without excessive mocking""" + + def setup_method(self): + """Setup test fixtures""" + self.config = VolcEngineEmbeddingConfig() + self.model = "doubao-embedding-text-240715" + self.api_key = "test-api-key-12345" + + def test_supported_params_completeness(self): + """Test that all required parameters are supported""" + params = self.config.get_supported_openai_params(self.model) + + # Verify essential parameters are supported + required_params = ["encoding_format", "user", "extra_headers"] + for param in required_params: + assert param in params, f"Required parameter '{param}' not supported" + + def test_parameter_mapping_with_valid_values(self): + """Test parameter mapping with various valid values""" + test_cases = [ + # Standard float encoding + {"encoding_format": "float", "user": "test-user"}, + # Base64 encoding + {"encoding_format": "base64", "user": "batch-user"}, + # None encoding (default) + {"encoding_format": None, "user": "api-user"}, + # Only user parameter + {"user": "minimal-user"}, + ] + + for test_params in test_cases: + result = self.config.map_openai_params( + non_default_params=test_params, + optional_params={}, + model=self.model, + drop_params=False + ) + + # Verify all valid parameters are preserved + for key, value in test_params.items(): + if value is not None: + assert result[key] == value, f"Parameter {key} not mapped correctly" + + def test_parameter_mapping_with_invalid_encoding(self): + """Test proper error handling for invalid encoding formats""" + invalid_encodings = ["int32", "binary", "invalid_format", 123, []] + + for invalid_encoding in invalid_encodings: + with pytest.raises(ValueError) as exc_info: + self.config.map_openai_params( + non_default_params={"encoding_format": invalid_encoding}, + optional_params={}, + model=self.model, + drop_params=False + ) + + assert "Unsupported encoding_format" in str(exc_info.value) + assert str(invalid_encoding) in str(exc_info.value) + + def test_parameter_dropping_behavior(self): + """Test parameter dropping when drop_params=True""" + invalid_params = { + "encoding_format": "invalid_format", + "unsupported_param": "value", + "another_invalid": 123 + } + + result = self.config.map_openai_params( + non_default_params=invalid_params, + optional_params={}, + model=self.model, + drop_params=True + ) + + # Should drop all invalid parameters + for param in invalid_params.keys(): + assert param not in result, f"Invalid parameter {param} was not dropped" + + def test_request_transformation_structure(self): + """Test request transformation produces correct structure""" + test_inputs = [ + # Single string input + "Hello world", + # Multiple strings + ["Hello", "World", "Test"], + # Mixed content + ["Short", "This is a longer text for testing purposes"], + ] + + for input_data in test_inputs: + result = self.config.transform_request( + model=self.model, + input=input_data, + api_key=self.api_key, + encoding_format="float" + ) + + # Verify structure + assert "url" in result + assert "headers" in result + assert "data" in result + + # Verify URL + assert result["url"] == "https://ark.cn-beijing.volces.com/api/v3/embeddings" + + # Verify headers + headers = result["headers"] + assert headers["Authorization"] == f"Bearer {self.api_key}" + assert headers["Content-Type"] == "application/json" + + # Verify data + data = result["data"] + assert data["model"] == self.model + assert data["encoding_format"] == "float" + + # Input should always be a list + if isinstance(input_data, str): + assert data["input"] == [input_data] + else: + assert data["input"] == input_data + + def test_response_transformation_with_real_data(self): + """Test response transformation with realistic Volcengine response data""" + # Simulate real Volcengine API response + volcengine_responses = [ + # Single embedding response + { + "id": "cmpl-123456789", + "object": "list", + "model": "doubao-embedding-text-240715", + "data": [ + { + "object": "embedding", + "index": 0, + "embedding": [0.1, -0.2, 0.3, 0.4, -0.5] * 100 # Realistic embedding size + } + ], + "usage": { + "prompt_tokens": 5, + "total_tokens": 5 + } + }, + # Multiple embeddings response + { + "id": "cmpl-987654321", + "object": "list", + "model": "doubao-embedding-text-240715", + "data": [ + { + "object": "embedding", + "index": 0, + "embedding": [0.1, 0.2, 0.3] * 256 + }, + { + "object": "embedding", + "index": 1, + "embedding": [0.4, 0.5, 0.6] * 256 + } + ], + "usage": { + "prompt_tokens": 12, + "total_tokens": 12 + } + } + ] + + for response_data in volcengine_responses: + mock_response = Mock(spec=httpx.Response) + mock_response.json.return_value = response_data + + result = self.config.transform_response( + response=mock_response, + model=self.model, + input=["test input"], + ) + + # Verify transformation preserves important data + assert result["object"] == "list" + assert result["model"] == response_data["model"] + assert len(result["data"]) == len(response_data["data"]) + assert result["usage"] == response_data["usage"] + + # Verify embedding data integrity + for i, embedding_item in enumerate(result["data"]): + original_item = response_data["data"][i] + assert embedding_item["object"] == "embedding" + assert embedding_item["index"] == original_item["index"] + assert len(embedding_item["embedding"]) == len(original_item["embedding"]) + + def test_response_transformation_with_error_data(self): + """Test response transformation handles error response formats correctly""" + # Test that transform_response can handle both success and error response structures + + # Success response (should work) + success_response = { + "id": "cmpl-123", + "object": "list", + "model": "doubao-embedding-text-240715", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}], + "usage": {"prompt_tokens": 2, "total_tokens": 2} + } + + mock_response = Mock(spec=httpx.Response) + mock_response.json.return_value = success_response + + result = self.config.transform_response( + response=mock_response, + model=self.model, + input=["test"], + ) + + # Should successfully transform + assert result["object"] == "list" + assert result["model"] == "doubao-embedding-text-240715" + + # Error response (should still transform but with empty/missing data) + error_response = { + "error": { + "message": "Rate limit exceeded", + "type": "rate_limit_error" + } + } + + mock_response.json.return_value = error_response + + result = self.config.transform_response( + response=mock_response, + model=self.model, + input=["test"], + ) + + # Should handle missing fields gracefully + assert result["object"] == "list" # default value + assert result["data"] == [] # default empty data + assert result["usage"] == {} # default empty usage + + +class TestVolcEngineEmbeddingHandlerBusinessLogic: + """Test VolcEngineEmbeddingHandler with focus on business logic""" + + def setup_method(self): + self.handler = VolcEngineEmbeddingHandler() + self.model = "doubao-embedding-text-240715" + self.api_key = "test-api-key-12345" + + def test_response_conversion_to_litellm_format(self): + """Test conversion of Volcengine response to LiteLLM EmbeddingResponse""" + volcengine_response = { + "id": "emb-123", + "object": "list", + "model": self.model, + "data": [ + { + "object": "embedding", + "index": 0, + "embedding": [0.1, 0.2, 0.3, -0.1, -0.2] * 200 # 1000-dimensional embedding + } + ], + "usage": { + "prompt_tokens": 8, + "total_tokens": 8 + } + } + + result = self.handler._convert_to_litellm_response( + volcengine_response, + self.model, + ["test input"] + ) + + # Verify result is proper EmbeddingResponse + assert isinstance(result, EmbeddingResponse) + assert result.object == "list" + assert result.model == self.model + assert len(result.data) == 1 + assert len(result.data[0]["embedding"]) == 1000 + + # Verify usage information + assert result.usage.prompt_tokens == 8 + assert result.usage.total_tokens == 8 + assert result.usage.completion_tokens == 0 + + def test_network_error_handling_without_mocking_business_logic(self): + """Test network error handling preserves business logic""" + + # Test with actual VolcEngineError class + with pytest.raises(VolcEngineError) as exc_info: + # This would raise a network error in real scenario + error = VolcEngineError( + status_code=500, + message="Network error during embedding request: Connection timeout" + ) + raise error + + # Verify error contains meaningful information + assert exc_info.value.status_code == 500 + assert "Network error during embedding request" in str(exc_info.value.message) + assert "Connection timeout" in str(exc_info.value.message) + + def test_input_validation_and_preprocessing(self): + """Test input validation and preprocessing logic""" + test_cases = [ + # String input should be converted to list + ("single string", ["single string"]), + # List input should remain list + (["multiple", "strings"], ["multiple", "strings"]), + # Empty string handling + ("", [""]), + # Unicode handling + ("ζ΅‹θ―•δΈ­ζ–‡", ["ζ΅‹θ―•δΈ­ζ–‡"]), + # Special characters + ("Special chars: @#$%^&*()", ["Special chars: @#$%^&*()"]), + ] + + for input_data, expected_output in test_cases: + # Test the actual transformation logic + config = VolcEngineEmbeddingConfig() + result = config.transform_request( + model=self.model, + input=input_data, + api_key=self.api_key, + ) + + assert result["data"]["input"] == expected_output + + +class TestVolcEngineEmbeddingIntegration: + """Integration tests that test the full pipeline with minimal mocking""" + + def setup_method(self): + self.handler = VolcEngineEmbeddingHandler() + self.model = "doubao-embedding-text-240715" + self.api_key = "test-api-key-12345" + + def test_full_request_response_cycle(self): + """Test the complete request-response cycle with realistic data""" + + # Create a realistic Volcengine response + realistic_response_data = { + "id": "cmpl-uqkvlQyYK7bGYrRHQ0eXlWi6", + "object": "list", + "model": "doubao-embedding-text-240715", + "data": [ + { + "object": "embedding", + "index": 0, + "embedding": [0.0023064255] + [0.1 * (i % 10 - 5) for i in range(1023)] # Realistic 1024-dim embedding + }, + { + "object": "embedding", + "index": 1, + "embedding": [-0.0038562391] + [0.05 * (i % 20 - 10) for i in range(1023)] + } + ], + "usage": { + "prompt_tokens": 6, + "total_tokens": 6 + } + } + + mock_response = Mock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.json.return_value = realistic_response_data + + # Only mock the HTTP call, not the business logic + with patch('litellm.llms.volcengine.embedding.handler.HTTPHandler') as mock_handler: + mock_client = Mock() + mock_client.post.return_value = mock_response + mock_handler.return_value = mock_client + + # Test the actual embedding call + result = self.handler.embedding( + model=self.model, + input=["Hello world", "Test embedding"], + api_key=self.api_key, + encoding_format="float" + ) + + # Verify the HTTP request was made correctly (this tests integration) + mock_client.post.assert_called_once() + call_args = mock_client.post.call_args + + # Verify request structure + assert call_args.kwargs["url"] == "https://ark.cn-beijing.volces.com/api/v3/embeddings" + assert call_args.kwargs["headers"]["Authorization"] == f"Bearer {self.api_key}" + + request_data = call_args.kwargs["json"] + assert request_data["model"] == self.model + assert request_data["input"] == ["Hello world", "Test embedding"] + assert request_data["encoding_format"] == "float" + + # Verify the response processing (real business logic) + assert isinstance(result, EmbeddingResponse) + assert result.model == self.model + assert len(result.data) == 2 + assert len(result.data[0]["embedding"]) == 1024 + assert len(result.data[1]["embedding"]) == 1024 + assert result.usage.prompt_tokens == 6 + + def test_parameter_validation_integration(self): + """Test parameter validation in the full integration context""" + + # Test with various parameter combinations that should work + valid_param_sets = [ + {"encoding_format": "float"}, + {"encoding_format": "base64"}, + {"user": "test-user-123"}, + {"encoding_format": "float", "user": "test-user"}, + {"extra_headers": {"Custom-Header": "value"}}, + ] + + for params in valid_param_sets: + # Only create the request, don't execute (avoids HTTP call) + config = VolcEngineEmbeddingConfig() + try: + result = config.transform_request( + model=self.model, + input=["test"], + api_key=self.api_key, + **params + ) + # Verify structure is correct + assert "url" in result + assert "headers" in result + assert "data" in result + + except Exception as e: + pytest.fail(f"Valid parameters {params} caused error: {e}") + + +if __name__ == "__main__": + pytest.main([__file__]) \ No newline at end of file diff --git a/tests/test_litellm/llms/test_volcengine.py b/tests/test_litellm/llms/volcengine/test_volcengine.py similarity index 97% rename from tests/test_litellm/llms/test_volcengine.py rename to tests/test_litellm/llms/volcengine/test_volcengine.py index 9db91217c28..59317914192 100644 --- a/tests/test_litellm/llms/test_volcengine.py +++ b/tests/test_litellm/llms/volcengine/test_volcengine.py @@ -4,7 +4,7 @@ from unittest.mock import MagicMock, patch from pydantic import BaseModel -from litellm.llms.volcengine import VolcEngineConfig +from litellm.llms.volcengine.chat.transformation import VolcEngineChatConfig as VolcEngineConfig from litellm.utils import get_optional_params diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index bd39fbfc9c4..4cbcb25ca3f 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -170,7 +170,7 @@ def test_all_model_configs(): drop_params=False, ) == {"max_tokens": 10} - from litellm.llms.volcengine import VolcEngineConfig + from litellm.llms.volcengine.chat.transformation import VolcEngineChatConfig as VolcEngineConfig assert "max_completion_tokens" in VolcEngineConfig().get_supported_openai_params( model="llama3" From b9ff636763add15c54d495c0426a1dda08e65016 Mon Sep 17 00:00:00 2001 From: onlylhf <27225745+onlylhf@users.noreply.github.com> Date: Thu, 28 Aug 2025 15:51:22 +0800 Subject: [PATCH 03/40] Optimize import statements and remove any unused type prompts --- litellm/llms/volcengine/embedding/handler.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/llms/volcengine/embedding/handler.py b/litellm/llms/volcengine/embedding/handler.py index e29f920afe8..961495e72f1 100644 --- a/litellm/llms/volcengine/embedding/handler.py +++ b/litellm/llms/volcengine/embedding/handler.py @@ -3,7 +3,7 @@ Volcengine Embedding Handler Handles embedding requests to Volcengine's embedding API """ -from typing import Dict, List, Optional, Union, Any +from typing import Dict, List, Optional, Union import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj From 7333060fb0469f12f8d51000f88232675c110b52 Mon Sep 17 00:00:00 2001 From: TomuHirata Date: Thu, 28 Aug 2025 18:43:43 +0900 Subject: [PATCH 04/40] feat(databricks): add anthropic citation support --- docs/my-website/docs/providers/databricks.md | 5 ++ .../llms/databricks/chat/transformation.py | 28 ++++++++++ litellm/types/llms/databricks.py | 5 +- .../test_databricks_chat_transformation.py | 53 ++++++++++++++++++- 4 files changed, 88 insertions(+), 3 deletions(-) diff --git a/docs/my-website/docs/providers/databricks.md b/docs/my-website/docs/providers/databricks.md index 8631cbfdad9..921b06a17b7 100644 --- a/docs/my-website/docs/providers/databricks.md +++ b/docs/my-website/docs/providers/databricks.md @@ -282,6 +282,11 @@ ModelResponse( ) ``` +### Citations + +Anthropic models served through Databricks can return citation metadata. LiteLLM +exposes these via `response.choices[0].message.provider_specific_fields["citations"]`. + ### Pass `thinking` to Anthropic models You can also pass the `thinking` parameter to Anthropic models. diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index 908419f7193..5600d5c6426 100644 --- a/litellm/llms/databricks/chat/transformation.py +++ b/litellm/llms/databricks/chat/transformation.py @@ -379,6 +379,21 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): thinking_blocks.append(thinking_block) return reasoning_content, thinking_blocks + @staticmethod + def extract_citations( + content: Optional[AllDatabricksContentValues], + ) -> Optional[List[Any]]: + if content is None: + return None + citations: Optional[List[Any]] = None + if isinstance(content, list): + for item in content: + if item.get("citations") is not None: + if citations is None: + citations = [] + citations.append(item["citations"]) + return citations + def _transform_dbrx_choices( self, choices: List[DatabricksChoice], json_mode: Optional[bool] = None ) -> List[Choices]: @@ -427,12 +442,19 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): choice["message"].get("content") ) + citations = DatabricksConfig.extract_citations( + choice["message"].get("content") + ) + translated_message = Message( role="assistant", content=content_str, reasoning_content=reasoning_content, thinking_blocks=thinking_blocks, tool_calls=choice["message"].get("tool_calls"), + provider_specific_fields={"citations": citations} + if citations is not None + else None, ) if finish_reason is None: @@ -561,6 +583,12 @@ class DatabricksChatResponseIterator(BaseModelResponseIterator): for _tc in tool_calls: if _tc.get("function", {}).get("arguments") == "{}": _tc["function"]["arguments"] = "" # avoid invalid json + citation = choice["delta"].get("citation") + if citation is not None: + choice["delta"].setdefault("provider_specific_fields", {})[ + "citation" + ] = citation + choice["delta"].pop("citation", None) # extract the content str content_str = DatabricksConfig.extract_content_str( choice["delta"].get("content") diff --git a/litellm/types/llms/databricks.py b/litellm/types/llms/databricks.py index bb59b692ef7..37151408161 100644 --- a/litellm/types/llms/databricks.py +++ b/litellm/types/llms/databricks.py @@ -1,5 +1,5 @@ import json -from typing import Any, List, Literal, Optional, TypedDict, Union +from typing import Any, Dict, List, Literal, Optional, TypedDict, Union from pydantic import BaseModel from typing_extensions import ( @@ -24,9 +24,10 @@ class GenericStreamingChunk(TypedDict, total=False): usage: Optional[BaseModel] -class DatabricksTextContent(TypedDict): +class DatabricksTextContent(TypedDict, total=False): type: Literal["text"] text: Required[str] + citations: Optional[List[Dict[str, Any]]] class DatabricksReasoningSummary(TypedDict): diff --git a/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py b/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py index fc44d44aba9..d61f826e89b 100644 --- a/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py +++ b/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py @@ -10,7 +10,10 @@ sys.path.insert( ) # Adds the parent directory to the system path from unittest.mock import MagicMock, patch -from litellm.llms.databricks.chat.transformation import DatabricksConfig +from litellm.llms.databricks.chat.transformation import ( + DatabricksChatResponseIterator, + DatabricksConfig, +) def test_transform_choices(): @@ -90,3 +93,51 @@ def test_transform_choices_without_signature(): thinking_block = choices[0].message.thinking_blocks[0] assert thinking_block["type"] == "thinking" assert thinking_block["thinking"] == "i'm thinking without signature." + + +def test_transform_choices_with_citations(): + config = DatabricksConfig() + databricks_choices = [ + { + "message": { + "role": "assistant", + "content": [ + { + "type": "text", + "text": "Paris", + "citations": [{"source": "wiki"}], + } + ], + }, + "index": 0, + "finish_reason": "stop", + } + ] + + choices = config._transform_dbrx_choices(choices=databricks_choices) + + assert choices[0].message.provider_specific_fields == { + "citations": [[{"source": "wiki"}]] + } + + +def test_chunk_parser_with_citation(): + iterator = DatabricksChatResponseIterator(None, sync_stream=True) + chunk = { + "id": "1", + "object": "chat.completion.chunk", + "created": 0, + "model": "test", + "choices": [ + { + "delta": {"citation": {"source": "wiki"}}, + "index": 0, + "finish_reason": None, + } + ], + } + + parsed = iterator.chunk_parser(chunk) + assert parsed.choices[0].delta.provider_specific_fields == { + "citation": {"source": "wiki"} + } From 38a1dbd13a549967f4fe4f8895934810b2a4ebab Mon Sep 17 00:00:00 2001 From: TomuHirata Date: Thu, 28 Aug 2025 22:30:27 +0900 Subject: [PATCH 05/40] fix(databricks): include citations in reasoning content type --- litellm/types/llms/databricks.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/litellm/types/llms/databricks.py b/litellm/types/llms/databricks.py index 37151408161..112427c6b56 100644 --- a/litellm/types/llms/databricks.py +++ b/litellm/types/llms/databricks.py @@ -36,9 +36,10 @@ class DatabricksReasoningSummary(TypedDict): signature: str -class DatabricksReasoningContent(TypedDict): +class DatabricksReasoningContent(TypedDict, total=False): type: Literal["reasoning"] - summary: List[DatabricksReasoningSummary] + summary: Required[List[DatabricksReasoningSummary]] + citations: Optional[List[Dict[str, Any]]] AllDatabricksContentListValues = Union[ From d83c420d484b99e8ecf068f1c6bf446dbece80d8 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Sun, 31 Aug 2025 16:29:19 +0900 Subject: [PATCH 06/40] feat: Add guardrail for the Anthropic API endpoint --- litellm/integrations/custom_guardrail.py | 7 +- .../proxy/anthropic_endpoints/endpoints.py | 42 ++++++-- .../integrations/test_custom_guardrail.py | 100 +++++++++++++++++- 3 files changed, 133 insertions(+), 16 deletions(-) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 501185b207e..1ca45f907e1 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -119,11 +119,8 @@ class CustomGuardrail(CustomLogger): """ if "guardrails" in data: return data["guardrails"] - metadata = data.get("metadata") or {} - requested_guardrails = metadata.get("guardrails") or [] - if requested_guardrails: - return requested_guardrails - return requested_guardrails + metadata = data.get("litellm_metadata") or data.get("metadata", {}) + return metadata.get("guardrails") or [] def _guardrail_is_in_requested_guardrails( self, diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index a10a39a6a57..2de5ec1ee12 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -90,6 +90,17 @@ async def anthropic_response( # noqa: PLR0915 user_api_key_dict=user_api_key_dict, data=data, call_type="text_completion" ) + tasks = [] + tasks.append( + proxy_logging_obj.during_call_hook( + data=data, + user_api_key_dict=user_api_key_dict, + call_type=ProxyBaseLLMRequestProcessing._get_pre_call_type( + route_type="anthropic_messages" # type: ignore + ), + ) + ) + ### ROUTE THE REQUESTs ### router_model_names = llm_router.model_names if llm_router is not None else [] @@ -97,23 +108,21 @@ async def anthropic_response( # noqa: PLR0915 if ( llm_router is not None and data["model"] in router_model_names ): # model in router model list - llm_response = asyncio.create_task(llm_router.aanthropic_messages(**data)) + llm_coro = llm_router.aanthropic_messages(**data) elif ( llm_router is not None and llm_router.model_group_alias is not None and data["model"] in llm_router.model_group_alias ): # model set in model_group_alias - llm_response = asyncio.create_task(llm_router.aanthropic_messages(**data)) + llm_coro = llm_router.aanthropic_messages(**data) elif ( llm_router is not None and data["model"] in llm_router.deployment_names ): # model in router deployments, calling a specific deployment on the router - llm_response = asyncio.create_task( - llm_router.aanthropic_messages(**data, specific_deployment=True) - ) + llm_coro = llm_router.aanthropic_messages(**data, specific_deployment=True) elif ( llm_router is not None and data["model"] in llm_router.get_model_ids() ): # model in router model list - llm_response = asyncio.create_task(llm_router.aanthropic_messages(**data)) + llm_coro = llm_router.aanthropic_messages(**data) elif ( llm_router is not None and data["model"] not in router_model_names @@ -122,9 +131,9 @@ async def anthropic_response( # noqa: PLR0915 or len(llm_router.pattern_router.patterns) > 0 ) ): # model in router deployments, calling a specific deployment on the router - llm_response = asyncio.create_task(llm_router.aanthropic_messages(**data)) + llm_coro = llm_router.aanthropic_messages(**data) elif user_model is not None: # `litellm --model ` - llm_response = asyncio.create_task(litellm.anthropic_messages(**data)) + llm_coro = litellm.anthropic_messages(**data) else: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -134,8 +143,16 @@ async def anthropic_response( # noqa: PLR0915 }, ) - # Await the llm_response task - response = await llm_response + tasks.append(llm_coro) + + # wait for call to end + llm_responses = asyncio.gather( + *tasks + ) # run the moderation check in parallel to the actual llm api call + + responses = await llm_responses + + response = responses[1] hidden_params = getattr(response, "_hidden_params", {}) or {} model_id = hidden_params.get("model_id", None) or "" @@ -183,6 +200,11 @@ async def anthropic_response( # noqa: PLR0915 headers=dict(fastapi_response.headers), ) + ### CALL HOOKS ### - modify outgoing data + response = await proxy_logging_obj.post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=response # type: ignore + ) + verbose_proxy_logger.info("\nResponse from Litellm:\n{}".format(response)) return response except Exception as e: diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index e71b68ab934..182e0134928 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -1,4 +1,4 @@ -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock import pytest @@ -82,3 +82,101 @@ class TestCustomGuardrailDeploymentHook: # Verify messages were updated in result assert result["messages"] == mock_result["messages"] assert result["messages"] != original_messages + + +class TestCustomGuardrailShouldRunGuardrail: + + def test_should_run_guardrail_with_litellm_metadata(self): + """Test that should_run_guardrail works with litellm_metadata pattern""" + from litellm.types.guardrails import GuardrailEventHooks + + custom_guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=False, + event_hook=GuardrailEventHooks.pre_call + ) + + # Test with guardrails in litellm_metadata + data = { + "model": "gpt-3.5-turbo", + "litellm_metadata": { + "guardrails": ["test_guardrail"] + } + } + + result = custom_guardrail.should_run_guardrail( + data=data, event_type=GuardrailEventHooks.pre_call + ) + + assert result is True + + def test_should_run_guardrail_with_metadata(self): + """Test that should_run_guardrail works with metadata pattern""" + from litellm.types.guardrails import GuardrailEventHooks + + custom_guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=False, + event_hook=GuardrailEventHooks.pre_call + ) + + # Test with guardrails in metadata + data = { + "model": "gpt-3.5-turbo", + "metadata": { + "guardrails": ["test_guardrail"] + } + } + + result = custom_guardrail.should_run_guardrail( + data=data, event_type=GuardrailEventHooks.pre_call + ) + + assert result is True + + def test_should_run_guardrail_with_root_level_guardrails(self): + """Test that should_run_guardrail works with root level guardrails""" + from litellm.types.guardrails import GuardrailEventHooks + + custom_guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=False, + event_hook=GuardrailEventHooks.pre_call + ) + + # Test with guardrails at root level + data = { + "model": "gpt-3.5-turbo", + "guardrails": ["test_guardrail"] + } + + result = custom_guardrail.should_run_guardrail( + data=data, event_type=GuardrailEventHooks.pre_call + ) + + assert result is True + + + def test_should_run_guardrail_no_matching_guardrail(self): + """Test that should_run_guardrail returns False when guardrail name doesn't match""" + from litellm.types.guardrails import GuardrailEventHooks + + custom_guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=False, + event_hook=GuardrailEventHooks.pre_call + ) + + # Test with different guardrail name + data = { + "model": "gpt-3.5-turbo", + "litellm_metadata": { + "guardrails": ["different_guardrail"] + } + } + + result = custom_guardrail.should_run_guardrail( + data=data, event_type=GuardrailEventHooks.pre_call + ) + + assert result is False From bfed4e0a6a5e70f44bc0ad0505677b717069ff00 Mon Sep 17 00:00:00 2001 From: onlylhf <27225745+onlylhf@users.noreply.github.com> Date: Mon, 1 Sep 2025 13:18:19 +0800 Subject: [PATCH 07/40] Add a complete URL generation method for embedding Volcengine API and optimize request and response processing logic; Delete redundant test files and refactor integration testing to improve readability and maintainability. --- .../volcengine/embedding/transformation.py | 70 ++- litellm/main.py | 51 +- litellm/utils.py | 6 + .../embedding/test_volcengine_embedding.py | 450 ------------------ .../volcengine}/test_volcengine_embedding.py | 4 +- 5 files changed, 84 insertions(+), 497 deletions(-) delete mode 100644 tests/test_litellm/llms/volcengine/embedding/test_volcengine_embedding.py rename tests/{llm_translation => test_litellm/llms/volcengine}/test_volcengine_embedding.py (98%) diff --git a/litellm/llms/volcengine/embedding/transformation.py b/litellm/llms/volcengine/embedding/transformation.py index ba2f07a4945..87e89626218 100644 --- a/litellm/llms/volcengine/embedding/transformation.py +++ b/litellm/llms/volcengine/embedding/transformation.py @@ -48,6 +48,36 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): "extra_headers", ] + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + """ + Get the complete URL for volcengine embedding API calls. + + Args: + api_base: Optional custom API base URL + api_key: API key (not used for URL construction) + model: Model name (not used for URL construction) + optional_params: Optional parameters (not used for URL construction) + litellm_params: LiteLLM parameters (not used for URL construction) + stream: Stream parameter (not used for URL construction) + + Returns: + Complete URL for the embedding API endpoint + """ + base_url = get_volcengine_base_url(api_base) + # Construct the complete URL with /embeddings endpoint + if base_url.endswith("/api/v3"): + return f"{base_url}/embeddings" + else: + return f"{base_url}/api/v3/embeddings" + def map_openai_params( self, non_default_params: Dict[str, Any], @@ -114,13 +144,14 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): Returns: Dict containing url, headers, and data for the request """ - # Get base URL - base_url = get_volcengine_base_url(api_base) - # Avoid duplicate /api/v3 if base_url already contains it - if base_url.endswith("/api/v3"): - url = f"{base_url}/embeddings" - else: - url = f"{base_url}/api/v3/embeddings" + # Get complete URL using the centralized method + url = self.get_complete_url( + api_base=api_base, + api_key=api_key, + model=model, + optional_params={}, + litellm_params={}, + ) # Get headers headers = get_volcengine_headers(api_key, extra_headers) @@ -188,13 +219,24 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): headers: dict, ) -> dict: """Transform embedding request to Volcengine format""" - # Use existing transform_request method - return self.transform_request( - model=model, - input=input, - api_key="", # api_key will be in headers - **optional_params, - ) + # Prepare request data (only the JSON body, not the full request) + data = { + "model": model, + "input": input if isinstance(input, list) else [input], + } + + # Add optional parameters from optional_params + if "encoding_format" in optional_params: + encoding_format = optional_params["encoding_format"] + if encoding_format is not None: + data["encoding_format"] = encoding_format + + if "user" in optional_params: + user = optional_params["user"] + if user is not None: + data["user"] = user + + return data def transform_embedding_response( self, diff --git a/litellm/main.py b/litellm/main.py index 776e81a0110..71daf1c970e 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -183,7 +183,6 @@ from .llms.vertex_ai.text_to_speech.text_to_speech_handler import VertexTextToSp from .llms.vertex_ai.vertex_ai_partner_models.main import VertexAIPartnerModels from .llms.vertex_ai.vertex_embeddings.embedding_handler import VertexEmbedding from .llms.vertex_ai.vertex_model_garden.main import VertexAIModelGardenModels -from .llms.volcengine.embedding.handler import VolcEngineEmbeddingHandler from .llms.vllm.completion import handler as vllm_handler from .llms.watsonx.chat.handler import WatsonXChatHandler from .llms.watsonx.common_utils import IBMWatsonXMixin @@ -4416,45 +4415,35 @@ def embedding( # noqa: PLR0915 aembedding=aembedding, ) elif custom_llm_provider == "volcengine": - api_key = ( + volcengine_key = ( api_key or litellm.api_key or get_secret_str("ARK_API_KEY") or get_secret_str("VOLCENGINE_API_KEY") ) - if api_key is None: + if volcengine_key is None: raise ValueError( "Missing API key for Volcengine. Set ARK_API_KEY or VOLCENGINE_API_KEY environment variable or pass api_key parameter." ) - - handler = VolcEngineEmbeddingHandler() - - if aembedding: - response = handler.async_embedding( - model=model, - input=input, - api_key=api_key, - api_base=api_base, - encoding_format=optional_params.get("encoding_format", "float"), - user=optional_params.get("user"), - timeout=timeout, - extra_headers=optional_params.get("extra_headers"), - litellm_logging_obj=logging, - **optional_params, - ) + if extra_headers is not None and isinstance(extra_headers, dict): + headers = extra_headers else: - response = handler.embedding( - model=model, - input=input, - api_key=api_key, - api_base=api_base, - encoding_format=optional_params.get("encoding_format", "float"), - user=optional_params.get("user"), - timeout=timeout, - extra_headers=optional_params.get("extra_headers"), - litellm_logging_obj=logging, - **optional_params, - ) + headers = {} + response = base_llm_http_handler.embedding( + model=model, + input=input, + timeout=timeout, + custom_llm_provider=custom_llm_provider, + logging_obj=logging, + api_base=api_base, + optional_params=optional_params, + litellm_params={}, + model_response=EmbeddingResponse(), + api_key=volcengine_key, + client=client, + aembedding=aembedding, + headers=headers, + ) elif custom_llm_provider in litellm._custom_providers: custom_handler: Optional[CustomLLM] = None for item in litellm.custom_provider_map: diff --git a/litellm/utils.py b/litellm/utils.py index aa3c00735ec..f3f9c49cb39 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7082,6 +7082,12 @@ class ProviderConfigManager: ) return JinaAIEmbeddingConfig() + elif litellm.LlmProviders.VOLCENGINE == provider: + from litellm.llms.volcengine.embedding.transformation import ( + VolcEngineEmbeddingConfig, + ) + + return VolcEngineEmbeddingConfig() return None @staticmethod diff --git a/tests/test_litellm/llms/volcengine/embedding/test_volcengine_embedding.py b/tests/test_litellm/llms/volcengine/embedding/test_volcengine_embedding.py deleted file mode 100644 index f2f143b5b99..00000000000 --- a/tests/test_litellm/llms/volcengine/embedding/test_volcengine_embedding.py +++ /dev/null @@ -1,450 +0,0 @@ -""" -Improved tests for Volcengine Embedding functionality -Tests real business logic without excessive mocking -""" - -import pytest -import json -import httpx -from unittest.mock import Mock, patch, MagicMock -from typing import List, Dict, Any - -from litellm.llms.volcengine.embedding import VolcEngineEmbeddingHandler, VolcEngineEmbeddingConfig -from litellm.llms.volcengine.common_utils import VolcEngineError -from litellm.types.utils import EmbeddingResponse -from litellm.types.llms.openai import AllEmbeddingInputValues - - -class TestVolcEngineEmbeddingConfigBusinessLogic: - """Test real business logic of VolcEngineEmbeddingConfig without excessive mocking""" - - def setup_method(self): - """Setup test fixtures""" - self.config = VolcEngineEmbeddingConfig() - self.model = "doubao-embedding-text-240715" - self.api_key = "test-api-key-12345" - - def test_supported_params_completeness(self): - """Test that all required parameters are supported""" - params = self.config.get_supported_openai_params(self.model) - - # Verify essential parameters are supported - required_params = ["encoding_format", "user", "extra_headers"] - for param in required_params: - assert param in params, f"Required parameter '{param}' not supported" - - def test_parameter_mapping_with_valid_values(self): - """Test parameter mapping with various valid values""" - test_cases = [ - # Standard float encoding - {"encoding_format": "float", "user": "test-user"}, - # Base64 encoding - {"encoding_format": "base64", "user": "batch-user"}, - # None encoding (default) - {"encoding_format": None, "user": "api-user"}, - # Only user parameter - {"user": "minimal-user"}, - ] - - for test_params in test_cases: - result = self.config.map_openai_params( - non_default_params=test_params, - optional_params={}, - model=self.model, - drop_params=False - ) - - # Verify all valid parameters are preserved - for key, value in test_params.items(): - if value is not None: - assert result[key] == value, f"Parameter {key} not mapped correctly" - - def test_parameter_mapping_with_invalid_encoding(self): - """Test proper error handling for invalid encoding formats""" - invalid_encodings = ["int32", "binary", "invalid_format", 123, []] - - for invalid_encoding in invalid_encodings: - with pytest.raises(ValueError) as exc_info: - self.config.map_openai_params( - non_default_params={"encoding_format": invalid_encoding}, - optional_params={}, - model=self.model, - drop_params=False - ) - - assert "Unsupported encoding_format" in str(exc_info.value) - assert str(invalid_encoding) in str(exc_info.value) - - def test_parameter_dropping_behavior(self): - """Test parameter dropping when drop_params=True""" - invalid_params = { - "encoding_format": "invalid_format", - "unsupported_param": "value", - "another_invalid": 123 - } - - result = self.config.map_openai_params( - non_default_params=invalid_params, - optional_params={}, - model=self.model, - drop_params=True - ) - - # Should drop all invalid parameters - for param in invalid_params.keys(): - assert param not in result, f"Invalid parameter {param} was not dropped" - - def test_request_transformation_structure(self): - """Test request transformation produces correct structure""" - test_inputs = [ - # Single string input - "Hello world", - # Multiple strings - ["Hello", "World", "Test"], - # Mixed content - ["Short", "This is a longer text for testing purposes"], - ] - - for input_data in test_inputs: - result = self.config.transform_request( - model=self.model, - input=input_data, - api_key=self.api_key, - encoding_format="float" - ) - - # Verify structure - assert "url" in result - assert "headers" in result - assert "data" in result - - # Verify URL - assert result["url"] == "https://ark.cn-beijing.volces.com/api/v3/embeddings" - - # Verify headers - headers = result["headers"] - assert headers["Authorization"] == f"Bearer {self.api_key}" - assert headers["Content-Type"] == "application/json" - - # Verify data - data = result["data"] - assert data["model"] == self.model - assert data["encoding_format"] == "float" - - # Input should always be a list - if isinstance(input_data, str): - assert data["input"] == [input_data] - else: - assert data["input"] == input_data - - def test_response_transformation_with_real_data(self): - """Test response transformation with realistic Volcengine response data""" - # Simulate real Volcengine API response - volcengine_responses = [ - # Single embedding response - { - "id": "cmpl-123456789", - "object": "list", - "model": "doubao-embedding-text-240715", - "data": [ - { - "object": "embedding", - "index": 0, - "embedding": [0.1, -0.2, 0.3, 0.4, -0.5] * 100 # Realistic embedding size - } - ], - "usage": { - "prompt_tokens": 5, - "total_tokens": 5 - } - }, - # Multiple embeddings response - { - "id": "cmpl-987654321", - "object": "list", - "model": "doubao-embedding-text-240715", - "data": [ - { - "object": "embedding", - "index": 0, - "embedding": [0.1, 0.2, 0.3] * 256 - }, - { - "object": "embedding", - "index": 1, - "embedding": [0.4, 0.5, 0.6] * 256 - } - ], - "usage": { - "prompt_tokens": 12, - "total_tokens": 12 - } - } - ] - - for response_data in volcengine_responses: - mock_response = Mock(spec=httpx.Response) - mock_response.json.return_value = response_data - - result = self.config.transform_response( - response=mock_response, - model=self.model, - input=["test input"], - ) - - # Verify transformation preserves important data - assert result["object"] == "list" - assert result["model"] == response_data["model"] - assert len(result["data"]) == len(response_data["data"]) - assert result["usage"] == response_data["usage"] - - # Verify embedding data integrity - for i, embedding_item in enumerate(result["data"]): - original_item = response_data["data"][i] - assert embedding_item["object"] == "embedding" - assert embedding_item["index"] == original_item["index"] - assert len(embedding_item["embedding"]) == len(original_item["embedding"]) - - def test_response_transformation_with_error_data(self): - """Test response transformation handles error response formats correctly""" - # Test that transform_response can handle both success and error response structures - - # Success response (should work) - success_response = { - "id": "cmpl-123", - "object": "list", - "model": "doubao-embedding-text-240715", - "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}], - "usage": {"prompt_tokens": 2, "total_tokens": 2} - } - - mock_response = Mock(spec=httpx.Response) - mock_response.json.return_value = success_response - - result = self.config.transform_response( - response=mock_response, - model=self.model, - input=["test"], - ) - - # Should successfully transform - assert result["object"] == "list" - assert result["model"] == "doubao-embedding-text-240715" - - # Error response (should still transform but with empty/missing data) - error_response = { - "error": { - "message": "Rate limit exceeded", - "type": "rate_limit_error" - } - } - - mock_response.json.return_value = error_response - - result = self.config.transform_response( - response=mock_response, - model=self.model, - input=["test"], - ) - - # Should handle missing fields gracefully - assert result["object"] == "list" # default value - assert result["data"] == [] # default empty data - assert result["usage"] == {} # default empty usage - - -class TestVolcEngineEmbeddingHandlerBusinessLogic: - """Test VolcEngineEmbeddingHandler with focus on business logic""" - - def setup_method(self): - self.handler = VolcEngineEmbeddingHandler() - self.model = "doubao-embedding-text-240715" - self.api_key = "test-api-key-12345" - - def test_response_conversion_to_litellm_format(self): - """Test conversion of Volcengine response to LiteLLM EmbeddingResponse""" - volcengine_response = { - "id": "emb-123", - "object": "list", - "model": self.model, - "data": [ - { - "object": "embedding", - "index": 0, - "embedding": [0.1, 0.2, 0.3, -0.1, -0.2] * 200 # 1000-dimensional embedding - } - ], - "usage": { - "prompt_tokens": 8, - "total_tokens": 8 - } - } - - result = self.handler._convert_to_litellm_response( - volcengine_response, - self.model, - ["test input"] - ) - - # Verify result is proper EmbeddingResponse - assert isinstance(result, EmbeddingResponse) - assert result.object == "list" - assert result.model == self.model - assert len(result.data) == 1 - assert len(result.data[0]["embedding"]) == 1000 - - # Verify usage information - assert result.usage.prompt_tokens == 8 - assert result.usage.total_tokens == 8 - assert result.usage.completion_tokens == 0 - - def test_network_error_handling_without_mocking_business_logic(self): - """Test network error handling preserves business logic""" - - # Test with actual VolcEngineError class - with pytest.raises(VolcEngineError) as exc_info: - # This would raise a network error in real scenario - error = VolcEngineError( - status_code=500, - message="Network error during embedding request: Connection timeout" - ) - raise error - - # Verify error contains meaningful information - assert exc_info.value.status_code == 500 - assert "Network error during embedding request" in str(exc_info.value.message) - assert "Connection timeout" in str(exc_info.value.message) - - def test_input_validation_and_preprocessing(self): - """Test input validation and preprocessing logic""" - test_cases = [ - # String input should be converted to list - ("single string", ["single string"]), - # List input should remain list - (["multiple", "strings"], ["multiple", "strings"]), - # Empty string handling - ("", [""]), - # Unicode handling - ("ζ΅‹θ―•δΈ­ζ–‡", ["ζ΅‹θ―•δΈ­ζ–‡"]), - # Special characters - ("Special chars: @#$%^&*()", ["Special chars: @#$%^&*()"]), - ] - - for input_data, expected_output in test_cases: - # Test the actual transformation logic - config = VolcEngineEmbeddingConfig() - result = config.transform_request( - model=self.model, - input=input_data, - api_key=self.api_key, - ) - - assert result["data"]["input"] == expected_output - - -class TestVolcEngineEmbeddingIntegration: - """Integration tests that test the full pipeline with minimal mocking""" - - def setup_method(self): - self.handler = VolcEngineEmbeddingHandler() - self.model = "doubao-embedding-text-240715" - self.api_key = "test-api-key-12345" - - def test_full_request_response_cycle(self): - """Test the complete request-response cycle with realistic data""" - - # Create a realistic Volcengine response - realistic_response_data = { - "id": "cmpl-uqkvlQyYK7bGYrRHQ0eXlWi6", - "object": "list", - "model": "doubao-embedding-text-240715", - "data": [ - { - "object": "embedding", - "index": 0, - "embedding": [0.0023064255] + [0.1 * (i % 10 - 5) for i in range(1023)] # Realistic 1024-dim embedding - }, - { - "object": "embedding", - "index": 1, - "embedding": [-0.0038562391] + [0.05 * (i % 20 - 10) for i in range(1023)] - } - ], - "usage": { - "prompt_tokens": 6, - "total_tokens": 6 - } - } - - mock_response = Mock(spec=httpx.Response) - mock_response.status_code = 200 - mock_response.json.return_value = realistic_response_data - - # Only mock the HTTP call, not the business logic - with patch('litellm.llms.volcengine.embedding.handler.HTTPHandler') as mock_handler: - mock_client = Mock() - mock_client.post.return_value = mock_response - mock_handler.return_value = mock_client - - # Test the actual embedding call - result = self.handler.embedding( - model=self.model, - input=["Hello world", "Test embedding"], - api_key=self.api_key, - encoding_format="float" - ) - - # Verify the HTTP request was made correctly (this tests integration) - mock_client.post.assert_called_once() - call_args = mock_client.post.call_args - - # Verify request structure - assert call_args.kwargs["url"] == "https://ark.cn-beijing.volces.com/api/v3/embeddings" - assert call_args.kwargs["headers"]["Authorization"] == f"Bearer {self.api_key}" - - request_data = call_args.kwargs["json"] - assert request_data["model"] == self.model - assert request_data["input"] == ["Hello world", "Test embedding"] - assert request_data["encoding_format"] == "float" - - # Verify the response processing (real business logic) - assert isinstance(result, EmbeddingResponse) - assert result.model == self.model - assert len(result.data) == 2 - assert len(result.data[0]["embedding"]) == 1024 - assert len(result.data[1]["embedding"]) == 1024 - assert result.usage.prompt_tokens == 6 - - def test_parameter_validation_integration(self): - """Test parameter validation in the full integration context""" - - # Test with various parameter combinations that should work - valid_param_sets = [ - {"encoding_format": "float"}, - {"encoding_format": "base64"}, - {"user": "test-user-123"}, - {"encoding_format": "float", "user": "test-user"}, - {"extra_headers": {"Custom-Header": "value"}}, - ] - - for params in valid_param_sets: - # Only create the request, don't execute (avoids HTTP call) - config = VolcEngineEmbeddingConfig() - try: - result = config.transform_request( - model=self.model, - input=["test"], - api_key=self.api_key, - **params - ) - # Verify structure is correct - assert "url" in result - assert "headers" in result - assert "data" in result - - except Exception as e: - pytest.fail(f"Valid parameters {params} caused error: {e}") - - -if __name__ == "__main__": - pytest.main([__file__]) \ No newline at end of file diff --git a/tests/llm_translation/test_volcengine_embedding.py b/tests/test_litellm/llms/volcengine/test_volcengine_embedding.py similarity index 98% rename from tests/llm_translation/test_volcengine_embedding.py rename to tests/test_litellm/llms/volcengine/test_volcengine_embedding.py index 9503d91f3c8..3be7f6ca8d4 100644 --- a/tests/llm_translation/test_volcengine_embedding.py +++ b/tests/test_litellm/llms/volcengine/test_volcengine_embedding.py @@ -9,9 +9,9 @@ from unittest.mock import MagicMock, patch import pytest # Add parent directory to path for imports -sys.path.insert(0, os.path.abspath("../..")) +sys.path.insert(0, os.path.abspath("../../../../..")) -from base_embedding_unit_tests import BaseLLMEmbeddingTest +from tests.llm_translation.base_embedding_unit_tests import BaseLLMEmbeddingTest import litellm from litellm.types.utils import EmbeddingResponse From e312c235334cbc2b240066e6dc049800dfe0d684 Mon Sep 17 00:00:00 2001 From: onlylhf <27225745+onlylhf@users.noreply.github.com> Date: Mon, 1 Sep 2025 13:29:06 +0800 Subject: [PATCH 08/40] Refactoring: Remove the transform-REquest and transform-REsponse methods, and directly implement the response transformation logic in transform_ embedding-REsponse; Enhance environment validation to ensure the validity of api_key --- .../volcengine/embedding/transformation.py | 120 ++++-------------- 1 file changed, 22 insertions(+), 98 deletions(-) diff --git a/litellm/llms/volcengine/embedding/transformation.py b/litellm/llms/volcengine/embedding/transformation.py index 87e89626218..20747b76725 100644 --- a/litellm/llms/volcengine/embedding/transformation.py +++ b/litellm/llms/volcengine/embedding/transformation.py @@ -117,99 +117,7 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): return optional_params - def transform_request( - self, - model: str, - input: Union[str, List[str]], - api_key: str, - api_base: Optional[str] = None, - encoding_format: Optional[str] = "float", - user: Optional[str] = None, - extra_headers: Optional[Dict[str, str]] = None, - **kwargs, - ) -> Dict[str, Any]: - """ - Transform OpenAI embedding request to Volcengine format. - Args: - model: Model ID (e.g., "doubao-embedding-text-240715") - input: Text or list of texts to embed - api_key: Volcengine API key - api_base: Optional custom API base URL - encoding_format: Response format (float, base64, null) - user: Optional user identifier - extra_headers: Optional additional headers - **kwargs: Additional parameters - - Returns: - Dict containing url, headers, and data for the request - """ - # Get complete URL using the centralized method - url = self.get_complete_url( - api_base=api_base, - api_key=api_key, - model=model, - optional_params={}, - litellm_params={}, - ) - - # Get headers - headers = get_volcengine_headers(api_key, extra_headers) - - # Prepare request data - data = { - "model": model, - "input": input if isinstance(input, list) else [input], - } - - # Add optional parameters - if encoding_format is not None: - data["encoding_format"] = encoding_format - - return { - "url": url, - "headers": headers, - "data": data, - } - - def transform_response( - self, - response: httpx.Response, - model: str, - input: Union[str, List[str]], - encoding: Optional[str] = None, - ) -> Dict[str, Any]: - """ - Transform Volcengine embedding response to OpenAI format. - - Args: - response: The HTTP response from Volcengine - model: The model used - input: The input that was embedded - encoding: The encoding format requested - - Returns: - OpenAI-compatible embedding response - """ - try: - response_json = response.json() - except Exception as e: - raise ValueError(f"Failed to parse Volcengine response as JSON: {str(e)}") - - # Volcengine response format matches OpenAI format closely - # Just need to ensure all required fields are present - transformed_response = { - "object": "list", - "data": response_json.get("data", []), - "model": response_json.get("model", model), - "usage": response_json.get("usage", {}), - } - - # Add id if present - if "id" in response_json: - transformed_response["id"] = response_json["id"] - - return transformed_response def transform_embedding_request( self, @@ -250,12 +158,23 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): litellm_params: dict, ) -> EmbeddingResponse: """Transform Volcengine response to EmbeddingResponse""" - # Use existing transform_response method - transformed_response = self.transform_response( - response=raw_response, - model=model, - input=request_data.get("input", []), - ) + try: + response_json = raw_response.json() + except Exception as e: + raise ValueError(f"Failed to parse Volcengine response as JSON: {str(e)}") + + # Volcengine response format matches OpenAI format closely + # Just need to ensure all required fields are present + transformed_response = { + "object": "list", + "data": response_json.get("data", []), + "model": response_json.get("model", model), + "usage": response_json.get("usage", {}), + } + + # Add id if present + if "id" in response_json: + transformed_response["id"] = response_json["id"] # Create EmbeddingResponse from transformed data return EmbeddingResponse(**transformed_response) @@ -272,6 +191,8 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): ) -> dict: """Validate environment and return headers""" # Get Volcengine headers + if api_key is None: + raise ValueError("api_key is required for Volcengine authentication") volcengine_headers = get_volcengine_headers(api_key) return {**headers, **volcengine_headers} @@ -280,6 +201,9 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): ) -> BaseLLMException: """Get error class for Volcengine errors""" from ..common_utils import VolcEngineError + # Convert dict to httpx.Headers if needed + if isinstance(headers, dict): + headers = httpx.Headers(headers) return VolcEngineError( status_code=status_code, message=error_message, From ca6d77b479b771b53ea0cf0107719db06b06165c Mon Sep 17 00:00:00 2001 From: TomeHirata Date: Mon, 1 Sep 2025 16:41:35 +0900 Subject: [PATCH 09/40] fix citation field name --- .../llms/databricks/chat/transformation.py | 28 +++++---- .../test_databricks_chat_transformation.py | 57 +++++++++++++++++-- 2 files changed, 69 insertions(+), 16 deletions(-) diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index 5600d5c6426..9330b019235 100644 --- a/litellm/llms/databricks/chat/transformation.py +++ b/litellm/llms/databricks/chat/transformation.py @@ -26,7 +26,6 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo _should_convert_tool_call_to_json_mode, ) from litellm.litellm_core_utils.prompt_templates.common_utils import ( - handle_messages_with_content_list_to_str_conversion, strip_name_from_messages, ) from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator @@ -301,7 +300,6 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: """ Databricks does not support: - - content in list format. - 'name' in user message. """ new_messages = [] @@ -311,7 +309,6 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): else: _message = message new_messages.append(_message) - new_messages = handle_messages_with_content_list_to_str_conversion(new_messages) new_messages = strip_name_from_messages(new_messages) if is_async: @@ -388,10 +385,16 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): citations: Optional[List[Any]] = None if isinstance(content, list): for item in content: + text = item.get("text", None) if item.get("citations") is not None: if citations is None: citations = [] - citations.append(item["citations"]) + citations.append( + [ + {**citation, "supported_text": text} + for citation in item["citations"] + ] + ) return citations def _transform_dbrx_choices( @@ -583,12 +586,17 @@ class DatabricksChatResponseIterator(BaseModelResponseIterator): for _tc in tool_calls: if _tc.get("function", {}).get("arguments") == "{}": _tc["function"]["arguments"] = "" # avoid invalid json - citation = choice["delta"].get("citation") - if citation is not None: - choice["delta"].setdefault("provider_specific_fields", {})[ - "citation" - ] = citation - choice["delta"].pop("citation", None) + if isinstance(choice["delta"]["content"], list) and ( + content := choice["delta"]["content"] + ): + if citations := content[0].get("citations"): + # TODO: Databricks delta does not include supported text or chunk type. + # Add either here once Databricks supports it to enable citation linkage. + choice["delta"].setdefault("provider_specific_fields", {})[ + "citation" + ] = citations[ + 0 + ] # Databricks Content item always has citation as a list of list # extract the content str content_str = DatabricksConfig.extract_content_str( choice["delta"].get("content") diff --git a/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py b/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py index d61f826e89b..51a2e971c09 100644 --- a/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py +++ b/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py @@ -88,7 +88,7 @@ def test_transform_choices_without_signature(): assert choices[0].message.reasoning_content == "i'm thinking without signature." assert choices[0].message.thinking_blocks is not None assert len(choices[0].message.thinking_blocks) == 1 - + # Verify the thinking block was created successfully without signature thinking_block = choices[0].message.thinking_blocks[0] assert thinking_block["type"] == "thinking" @@ -104,8 +104,17 @@ def test_transform_choices_with_citations(): "content": [ { "type": "text", - "text": "Paris", - "citations": [{"source": "wiki"}], + "text": "Blue", + "citations": [ + { + "type": "char_location", + "cited_text": "The sky is blue.", + "document_index": 0, + "document_title": "My Document", + "start_char_index": 0, + "end_char_index": 50, + } + ], } ], }, @@ -117,7 +126,19 @@ def test_transform_choices_with_citations(): choices = config._transform_dbrx_choices(choices=databricks_choices) assert choices[0].message.provider_specific_fields == { - "citations": [[{"source": "wiki"}]] + "citations": [ + [ + { + "type": "char_location", + "cited_text": "The sky is blue.", + "document_index": 0, + "document_title": "My Document", + "start_char_index": 0, + "end_char_index": 50, + "supported_text": "Blue", + } + ] + ] } @@ -130,7 +151,24 @@ def test_chunk_parser_with_citation(): "model": "test", "choices": [ { - "delta": {"citation": {"source": "wiki"}}, + "delta": { + "content": [ + { + "type": "text", + "text": "", + "citations": [ + { + "type": "char_location", + "cited_text": "The sky is blue.", + "document_index": 0, + "document_title": "My Document", + "start_char_index": 0, + "end_char_index": 50, + } + ], + } + ], + }, "index": 0, "finish_reason": None, } @@ -139,5 +177,12 @@ def test_chunk_parser_with_citation(): parsed = iterator.chunk_parser(chunk) assert parsed.choices[0].delta.provider_specific_fields == { - "citation": {"source": "wiki"} + "citation": { + "type": "char_location", + "cited_text": "The sky is blue.", + "document_index": 0, + "document_title": "My Document", + "start_char_index": 0, + "end_char_index": 50, + } } From cf676e7aeff459330c5df604a636955d494aff45 Mon Sep 17 00:00:00 2001 From: TomeHirata Date: Mon, 1 Sep 2025 17:51:24 +0900 Subject: [PATCH 10/40] fix mypy --- litellm/llms/databricks/chat/transformation.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index 9330b019235..d3df5bbf361 100644 --- a/litellm/llms/databricks/chat/transformation.py +++ b/litellm/llms/databricks/chat/transformation.py @@ -382,20 +382,18 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): ) -> Optional[List[Any]]: if content is None: return None - citations: Optional[List[Any]] = None + citations = [] if isinstance(content, list): for item in content: text = item.get("text", None) - if item.get("citations") is not None: - if citations is None: - citations = [] + if citations_item := item.get("citations"): citations.append( [ {**citation, "supported_text": text} - for citation in item["citations"] + for citation in citations_item ] ) - return citations + return citations or None def _transform_dbrx_choices( self, choices: List[DatabricksChoice], json_mode: Optional[bool] = None From 51f44c0419bf338c8104d08a13ec62835d45a5b7 Mon Sep 17 00:00:00 2001 From: tanjiro <56165694+NANDINI-star@users.noreply.github.com> Date: Tue, 2 Sep 2025 19:26:09 +0900 Subject: [PATCH 11/40] limit to 20 teams and make it expandable after that --- .../components/view_users/user_info_view.tsx | 133 +++++++++++++----- 1 file changed, 99 insertions(+), 34 deletions(-) diff --git a/ui/litellm-dashboard/src/components/view_users/user_info_view.tsx b/ui/litellm-dashboard/src/components/view_users/user_info_view.tsx index 823baaf64ef..59d2b8e3867 100644 --- a/ui/litellm-dashboard/src/components/view_users/user_info_view.tsx +++ b/ui/litellm-dashboard/src/components/view_users/user_info_view.tsx @@ -1,6 +1,6 @@ import React, { useState } from "react" import { Card, Text, Button, Grid, Col, Tab, TabList, TabGroup, TabPanel, TabPanels, Title, Badge } from "@tremor/react" -import { ArrowLeftIcon, TrashIcon, RefreshIcon } from "@heroicons/react/outline" +import { ArrowLeftIcon, TrashIcon, RefreshIcon, ChevronDownIcon, ChevronUpIcon } from "@heroicons/react/outline" import { userInfoCall, userDeleteCall, @@ -14,9 +14,9 @@ import { rolesWithWriteAccess } from "../../utils/roles" import { UserEditView } from "../user_edit_view" import OnboardingModal, { InvitationLink } from "../onboarding_link" import { formatNumberWithCommas, copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils" -import { CopyIcon, CheckIcon } from "lucide-react"; -import NotificationsManager from "../molecules/notifications_manager"; -import { getBudgetDurationLabel } from "../common_components/budget_duration_dropdown"; +import { CopyIcon, CheckIcon } from "lucide-react" +import NotificationsManager from "../molecules/notifications_manager" +import { getBudgetDurationLabel } from "../common_components/budget_duration_dropdown" interface UserInfoViewProps { userId: string @@ -57,16 +57,17 @@ export default function UserInfoView({ initialTab = 0, startInEditMode = false, }: UserInfoViewProps) { - const [userData, setUserData] = useState(null); - const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); - const [isLoading, setIsLoading] = useState(true); - const [isEditing, setIsEditing] = useState(startInEditMode); - const [userModels, setUserModels] = useState([]); - const [isInvitationLinkModalVisible, setIsInvitationLinkModalVisible] = useState(false); - const [invitationLinkData, setInvitationLinkData] = useState(null); - const [baseUrl, setBaseUrl] = useState(null); - const [activeTab, setActiveTab] = useState(initialTab); - const [copiedStates, setCopiedStates] = useState>({}); + const [userData, setUserData] = useState(null) + const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false) + const [isLoading, setIsLoading] = useState(true) + const [isEditing, setIsEditing] = useState(startInEditMode) + const [userModels, setUserModels] = useState([]) + const [isInvitationLinkModalVisible, setIsInvitationLinkModalVisible] = useState(false) + const [invitationLinkData, setInvitationLinkData] = useState(null) + const [baseUrl, setBaseUrl] = useState(null) + const [activeTab, setActiveTab] = useState(initialTab) + const [copiedStates, setCopiedStates] = useState>({}) + const [isTeamsExpanded, setIsTeamsExpanded] = useState(false) React.useEffect(() => { setBaseUrl(getProxyBaseUrl()) @@ -175,14 +176,14 @@ export default function UserInfoView({ } const copyToClipboard = async (text: string, key: string) => { - const success = await utilCopyToClipboard(text); + const success = await utilCopyToClipboard(text) if (success) { - setCopiedStates((prev) => ({ ...prev, [key]: true })); + setCopiedStates((prev) => ({ ...prev, [key]: true })) setTimeout(() => { - setCopiedStates((prev) => ({ ...prev, [key]: false })); - }, 2000); + setCopiedStates((prev) => ({ ...prev, [key]: false })) + }, 2000) } - }; + } return (
@@ -200,9 +201,9 @@ export default function UserInfoView({ icon={copiedStates["user-id"] ? : } onClick={() => copyToClipboard(userData.user_id, "user-id")} className={`left-2 z-10 transition-all duration-200 ${ - copiedStates["user-id"] - ? 'text-green-600 bg-green-50 border-green-200' - : 'text-gray-500 hover:text-gray-700 hover:bg-gray-100' + copiedStates["user-id"] + ? "text-green-600 bg-green-50 border-green-200" + : "text-gray-500 hover:text-gray-700 hover:bg-gray-100" }`} />
@@ -282,15 +283,44 @@ export default function UserInfoView({ - Teams +
+ Teams + {userData.teams?.length && userData.teams?.length > 20 && ( + + )} +
{userData.teams?.length && userData.teams?.length > 0 ? (
- {userData.teams?.map((team, index) => ( - + {userData.teams?.slice(0, isTeamsExpanded ? userData.teams.length : 20).map((team, index) => ( + {team.team_alias} ))} + {!isTeamsExpanded && userData.teams?.length > 20 && ( +
+ + +{userData.teams.length - 20} more + +
+
+ {userData.teams?.slice(20).map((team, index) => ( +
+ {team.team_alias} +
+ ))} +
+
+
+
+ )}
) : ( No teams @@ -354,9 +384,9 @@ export default function UserInfoView({ icon={copiedStates["user-id"] ? : } onClick={() => copyToClipboard(userData.user_id, "user-id")} className={`left-2 z-10 transition-all duration-200 ${ - copiedStates["user-id"] - ? 'text-green-600 bg-green-50 border-green-200' - : 'text-gray-500 hover:text-gray-700 hover:bg-gray-100' + copiedStates["user-id"] + ? "text-green-600 bg-green-50 border-green-200" + : "text-gray-500 hover:text-gray-700 hover:bg-gray-100" }`} />
@@ -391,14 +421,49 @@ export default function UserInfoView({
- Teams +
+ Teams + {userData.teams?.length && userData.teams?.length > 20 && ( + + )} +
{userData.teams?.length && userData.teams?.length > 0 ? ( - userData.teams?.map((team, index) => ( - - {team.team_alias || team.team_id} - - )) + <> + {userData.teams?.slice(0, isTeamsExpanded ? userData.teams.length : 20).map((team, index) => ( + + {team.team_alias || team.team_id} + + ))} + {!isTeamsExpanded && userData.teams?.length > 20 && ( +
+ + +{userData.teams.length - 20} more + +
+
+ {userData.teams?.slice(20).map((team, index) => ( +
+ {team.team_alias || team.team_id} +
+ ))} +
+
+
+
+ )} + ) : ( No teams )} From 80970951c500ec1a61bfa8546d9fd3d472a08c9f Mon Sep 17 00:00:00 2001 From: tanjiro <56165694+NANDINI-star@users.noreply.github.com> Date: Tue, 2 Sep 2025 19:36:06 +0900 Subject: [PATCH 12/40] fix ui for expandable badge --- .../components/view_users/user_info_view.tsx | 90 +++++++------------ 1 file changed, 33 insertions(+), 57 deletions(-) diff --git a/ui/litellm-dashboard/src/components/view_users/user_info_view.tsx b/ui/litellm-dashboard/src/components/view_users/user_info_view.tsx index 59d2b8e3867..c36bde78a7d 100644 --- a/ui/litellm-dashboard/src/components/view_users/user_info_view.tsx +++ b/ui/litellm-dashboard/src/components/view_users/user_info_view.tsx @@ -1,6 +1,6 @@ import React, { useState } from "react" import { Card, Text, Button, Grid, Col, Tab, TabList, TabGroup, TabPanel, TabPanels, Title, Badge } from "@tremor/react" -import { ArrowLeftIcon, TrashIcon, RefreshIcon, ChevronDownIcon, ChevronUpIcon } from "@heroicons/react/outline" +import { ArrowLeftIcon, TrashIcon, RefreshIcon } from "@heroicons/react/outline" import { userInfoCall, userDeleteCall, @@ -283,19 +283,7 @@ export default function UserInfoView({ -
- Teams - {userData.teams?.length && userData.teams?.length > 20 && ( - - )} -
+ Teams
{userData.teams?.length && userData.teams?.length > 0 ? (
@@ -305,21 +293,22 @@ export default function UserInfoView({ ))} {!isTeamsExpanded && userData.teams?.length > 20 && ( -
- - +{userData.teams.length - 20} more - -
-
- {userData.teams?.slice(20).map((team, index) => ( -
- {team.team_alias} -
- ))} -
-
-
-
+ setIsTeamsExpanded(true)} + > + +{userData.teams.length - 20} more + + )} + {isTeamsExpanded && userData.teams?.length > 20 && ( + setIsTeamsExpanded(false)} + > + Show Less + )}
) : ( @@ -421,19 +410,7 @@ export default function UserInfoView({
-
- Teams - {userData.teams?.length && userData.teams?.length > 20 && ( - - )} -
+ Teams
{userData.teams?.length && userData.teams?.length > 0 ? ( <> @@ -447,21 +424,20 @@ export default function UserInfoView({ ))} {!isTeamsExpanded && userData.teams?.length > 20 && ( -
- - +{userData.teams.length - 20} more - -
-
- {userData.teams?.slice(20).map((team, index) => ( -
- {team.team_alias || team.team_id} -
- ))} -
-
-
-
+ setIsTeamsExpanded(true)} + > + +{userData.teams.length - 20} more + + )} + {isTeamsExpanded && userData.teams?.length > 20 && ( + setIsTeamsExpanded(false)} + > + Show Less + )} ) : ( From 447016817cfc5f5af57eb7e5e4550d3523b143de Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Thu, 4 Sep 2025 07:11:44 +0900 Subject: [PATCH 13/40] fix: Call guardrail during stream processing --- litellm/proxy/common_request_processing.py | 14 +++++--------- 1 file changed, 5 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index f4d794d94bc..d41da69b6dc 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -108,7 +108,6 @@ async def create_streaming_response( final_status_code = default_status_code try: - # Handle coroutine that returns a generator if asyncio.iscoroutine(generator): generator = await generator @@ -117,7 +116,6 @@ async def create_streaming_response( first_chunk_value = await generator.__anext__() if first_chunk_value is not None: - try: error_code_from_chunk = await _parse_event_data_for_error( first_chunk_value @@ -131,7 +129,6 @@ async def create_streaming_response( verbose_proxy_logger.debug(f"Error parsing first chunk value: {e}") except StopAsyncIteration: - # Generator was empty. Default status async def empty_gen() -> AsyncGenerator[str, None]: if False: @@ -144,7 +141,6 @@ async def create_streaming_response( status_code=default_status_code, ) except Exception as e: - # Unexpected error consuming first chunk. verbose_proxy_logger.exception( f"Error consuming first chunk from generator: {e}" @@ -167,7 +163,6 @@ async def create_streaming_response( with tracer.trace(DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE): yield first_chunk_value async for chunk in generator: - with tracer.trace(DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE): yield chunk @@ -460,7 +455,6 @@ class ProxyBaseLLMRequestProcessing: ) or self._is_streaming_response( response ): # use generate_responses to stream responses - custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers( user_api_key_dict=user_api_key_dict, call_id=logging_obj.litellm_call_id, @@ -478,7 +472,6 @@ class ProxyBaseLLMRequestProcessing: if route_type == "allm_passthrough_route": # Check if response is an async generator if self._is_streaming_response(response): - if asyncio.iscoroutine(response): generator = await response else: @@ -499,7 +492,6 @@ class ProxyBaseLLMRequestProcessing: headers=custom_headers, ) else: - selected_data_generator = select_data_generator( response=response, user_api_key_dict=user_api_key_dict, @@ -738,7 +730,11 @@ class ProxyBaseLLMRequestProcessing: verbose_proxy_logger.debug("inside generator") try: str_so_far = "" - async for chunk in response: + async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=response, + request_data=request_data, + ): verbose_proxy_logger.debug( "async_data_generator: received streaming chunk - {}".format(chunk) ) From 006ffea98f5c6223c5e0c2c5eca2312df1714ecf Mon Sep 17 00:00:00 2001 From: Eitan1112 <52412573+Eitan1112@users.noreply.github.com> Date: Thu, 4 Sep 2025 17:53:40 +0300 Subject: [PATCH 14/40] Add additionalProperties to vertex ai Schema definition Add additionalProperties field to vertex ai Schema TypedDict --- litellm/types/llms/vertex_ai.py | 1 + 1 file changed, 1 insertion(+) diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 1b74ee25803..c3027504dff 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -113,6 +113,7 @@ class Schema(TypedDict, total=False): pattern: str example: Any anyOf: List["Schema"] + additionalProperties: bool class FunctionDeclaration(TypedDict, total=False): From da136fa07b85cd27c87d49147e9727c28b7c7a9d Mon Sep 17 00:00:00 2001 From: Eitan1112 <52412573+Eitan1112@users.noreply.github.com> Date: Thu, 4 Sep 2025 18:05:53 +0300 Subject: [PATCH 15/40] Change additionalProperties type to Any This is aligned with "default" which is also `Any`, and both in vertex ai docs: https://cloud.google.com/vertex-ai/docs/reference/rest/v1/projects.locations.cachedContents#Schema are both with 'value' type --- litellm/types/llms/vertex_ai.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index c3027504dff..625a76b6789 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -113,7 +113,7 @@ class Schema(TypedDict, total=False): pattern: str example: Any anyOf: List["Schema"] - additionalProperties: bool + additionalProperties: Any class FunctionDeclaration(TypedDict, total=False): From 99eceb8835a2faa1795ba5d885480c8d5f958497 Mon Sep 17 00:00:00 2001 From: tobias-mayr Date: Thu, 4 Sep 2025 18:29:06 +0100 Subject: [PATCH 16/40] feat: Add support for reasoning_effort='minimal' for Gemini models - Add DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET constant (128 tokens) - Update Gemini transformation to handle 'minimal' reasoning_effort - Maps 'minimal' to 128 tokens (Gemini's minimum thinking budget) - Maintains backward compatibility with existing reasoning_effort values - Fixes issue where Gemini API rejected 0 token thinking budget --- litellm/constants.py | 3 + .../exception_mapping_utils.py | 76 +++++++++++++++++-- .../vertex_and_google_ai_studio_gemini.py | 8 +- 3 files changed, 79 insertions(+), 8 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 21e30bef32b..9f55d2a94ef 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -51,6 +51,9 @@ SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD = int( DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET = int( os.getenv("DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET", 0) ) +DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET = int( + os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET", 128) +) DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET = int( os.getenv("DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET", 1024) ) diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index 25ae0269ab3..ad6b3dcaeb4 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -24,6 +24,55 @@ from ..exceptions import ( ) +def _is_operational_404(original_exception) -> bool: + """ + Determine if a 404 status code represents an operational issue rather than a missing model. + + Args: + original_exception: The exception with status_code 404 + + Returns: + True if this is an operational issue (rate limiting, cooldowns, etc.) + False if this is actually a missing model + """ + try: + # Import here to avoid circular imports + from litellm.types.router import RouterErrors + + # Check for known operational error patterns + error_message = str(original_exception).lower() + + # Check for router-specific operational errors + operational_patterns = [ + RouterErrors.no_deployments_available.value.lower(), + "no deployments available", + "no healthy deployment available", + "no healthy deployments available", + "deployment over user-defined ratelimit", + "crossed budget", + "cooldown", + "rate limit exceeded", + "too many requests" + ] + + for pattern in operational_patterns: + if pattern in error_message: + return True + + # Check if this is a RouterRateLimitError (which indicates operational issues) + if hasattr(original_exception, '__class__'): + exception_class_name = original_exception.__class__.__name__ + if "RouterRateLimitError" in exception_class_name: + return True + + return False + + except Exception: + # If we can't determine, default to treating it as a missing model + # This is safer than potentially hiding real model not found errors + return False + + class ExceptionCheckers: """ Helper class for checking various error conditions in exception strings. @@ -462,13 +511,26 @@ def exception_type( # type: ignore # noqa: PLR0915 ) elif original_exception.status_code == 404: exception_mapping_worked = True - raise NotFoundError( - message=f"NotFoundError: {exception_provider} - {message}", - model=model, - llm_provider=custom_llm_provider, - response=getattr(original_exception, "response", None), - litellm_debug_info=extra_information, - ) + # Check if this is actually a "model not found" vs operational issue + if _is_operational_404(original_exception): + # This is operational (rate limiting, cooldowns), not a missing model + # The proxy will map this to 429 status code, which is correct + raise litellm.ServiceUnavailableError( + message=f"ServiceUnavailableError: {exception_provider} - {message}", + model=model, + llm_provider=custom_llm_provider, + response=getattr(original_exception, "response", None), + litellm_debug_info=extra_information, + ) + else: + # This is actually a missing model + raise NotFoundError( + message=f"NotFoundError: {exception_provider} - {message}", + model=model, + llm_provider=custom_llm_provider, + response=getattr(original_exception, "response", None), + litellm_debug_info=extra_information, + ) elif original_exception.status_code == 408: exception_mapping_worked = True raise Timeout( diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 37470a6ee09..4da99204165 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -30,6 +30,7 @@ from litellm.constants import ( DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, + DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET, ) from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.custom_httpx.http_handler import ( @@ -423,7 +424,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def _map_reasoning_effort_to_thinking_budget( reasoning_effort: str, ) -> GeminiThinkingConfig: - if reasoning_effort == "low": + if reasoning_effort == "minimal": + return { + "thinkingBudget": DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET, + "includeThoughts": True, + } + elif reasoning_effort == "low": return { "thinkingBudget": DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET, "includeThoughts": True, From 0ade6cceff672a0a9001f0f185649b8c76c9c45d Mon Sep 17 00:00:00 2001 From: tobias-mayr Date: Thu, 4 Sep 2025 18:40:46 +0100 Subject: [PATCH 17/40] remove old code --- .../exception_mapping_utils.py | 82 +++---------------- 1 file changed, 10 insertions(+), 72 deletions(-) diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index ad6b3dcaeb4..f02c862f0fa 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -24,55 +24,6 @@ from ..exceptions import ( ) -def _is_operational_404(original_exception) -> bool: - """ - Determine if a 404 status code represents an operational issue rather than a missing model. - - Args: - original_exception: The exception with status_code 404 - - Returns: - True if this is an operational issue (rate limiting, cooldowns, etc.) - False if this is actually a missing model - """ - try: - # Import here to avoid circular imports - from litellm.types.router import RouterErrors - - # Check for known operational error patterns - error_message = str(original_exception).lower() - - # Check for router-specific operational errors - operational_patterns = [ - RouterErrors.no_deployments_available.value.lower(), - "no deployments available", - "no healthy deployment available", - "no healthy deployments available", - "deployment over user-defined ratelimit", - "crossed budget", - "cooldown", - "rate limit exceeded", - "too many requests" - ] - - for pattern in operational_patterns: - if pattern in error_message: - return True - - # Check if this is a RouterRateLimitError (which indicates operational issues) - if hasattr(original_exception, '__class__'): - exception_class_name = original_exception.__class__.__name__ - if "RouterRateLimitError" in exception_class_name: - return True - - return False - - except Exception: - # If we can't determine, default to treating it as a missing model - # This is safer than potentially hiding real model not found errors - return False - - class ExceptionCheckers: """ Helper class for checking various error conditions in exception strings. @@ -91,16 +42,16 @@ class ExceptionCheckers: """ if not isinstance(error_str, str): return False - + if "429" in error_str or "rate limit" in error_str.lower(): return True - + ####################################### # Mistral API returns this error string ######################################### if "service tier capacity exceeded" in error_str.lower(): return True - + return False @staticmethod @@ -511,26 +462,13 @@ def exception_type( # type: ignore # noqa: PLR0915 ) elif original_exception.status_code == 404: exception_mapping_worked = True - # Check if this is actually a "model not found" vs operational issue - if _is_operational_404(original_exception): - # This is operational (rate limiting, cooldowns), not a missing model - # The proxy will map this to 429 status code, which is correct - raise litellm.ServiceUnavailableError( - message=f"ServiceUnavailableError: {exception_provider} - {message}", - model=model, - llm_provider=custom_llm_provider, - response=getattr(original_exception, "response", None), - litellm_debug_info=extra_information, - ) - else: - # This is actually a missing model - raise NotFoundError( - message=f"NotFoundError: {exception_provider} - {message}", - model=model, - llm_provider=custom_llm_provider, - response=getattr(original_exception, "response", None), - litellm_debug_info=extra_information, - ) + raise NotFoundError( + message=f"NotFoundError: {exception_provider} - {message}", + model=model, + llm_provider=custom_llm_provider, + response=getattr(original_exception, "response", None), + litellm_debug_info=extra_information, + ) elif original_exception.status_code == 408: exception_mapping_worked = True raise Timeout( From 2d30b55964324ee43d3c89e7b981d23aa5380c66 Mon Sep 17 00:00:00 2001 From: tobias-mayr Date: Thu, 4 Sep 2025 18:59:37 +0100 Subject: [PATCH 18/40] distinguish between gemini models --- litellm/constants.py | 22 +++++++++++++++---- .../vertex_and_google_ai_studio_gemini.py | 22 +++++++++++++++++-- 2 files changed, 38 insertions(+), 6 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 9f55d2a94ef..25f26639494 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -51,9 +51,23 @@ SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD = int( DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET = int( os.getenv("DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET", 0) ) + +# Gemini model-specific minimal thinking budget constants +DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH = int( + os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH", 1) +) +DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO = int( + os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO", 128) +) +DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE = int( + os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE", 512) +) + +# Generic fallback for unknown models DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET = int( os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET", 128) ) + DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET = int( os.getenv("DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET", 1024) ) @@ -830,7 +844,7 @@ known_tokenizer_config = { "add_eos_token": False, "bos_token": { "__type": "AddedToken", - "content": "<|begin▁of▁sentence|>", + "content": " "lstrip": False, "normalized": True, "rstrip": False, @@ -839,7 +853,7 @@ known_tokenizer_config = { "clean_up_tokenization_spaces": False, "eos_token": { "__type": "AddedToken", - "content": "<|end▁of▁sentence|>", + "content": " "lstrip": False, "normalized": True, "rstrip": False, @@ -849,7 +863,7 @@ known_tokenizer_config = { "model_max_length": 16384, "pad_token": { "__type": "AddedToken", - "content": "<|end▁of▁sentence|>", + "content": " "lstrip": False, "normalized": True, "rstrip": False, @@ -858,7 +872,7 @@ known_tokenizer_config = { "sp_model_kwargs": {}, "unk_token": None, "tokenizer_class": "LlamaTokenizerFast", - "chat_template": "{% if not add_generation_prompt is defined %}{% set add_generation_prompt = false %}{% endif %}{% set ns = namespace(is_first=false, is_tool=false, is_output_first=true, system_prompt='') %}{%- for message in messages %}{%- if message['role'] == 'system' %}{% set ns.system_prompt = message['content'] %}{%- endif %}{%- endfor %}{{bos_token}}{{ns.system_prompt}}{%- for message in messages %}{%- if message['role'] == 'user' %}{%- set ns.is_tool = false -%}{{'<|User|>' + message['content']}}{%- endif %}{%- if message['role'] == 'assistant' and message['content'] is none %}{%- set ns.is_tool = false -%}{%- for tool in message['tool_calls']%}{%- if not ns.is_first %}{{'<|Assistant|><|tool▁calls▁begin|><|tool▁call▁begin|>' + tool['type'] + '<|tool▁sep|>' + tool['function']['name'] + '\\n' + '```json' + '\\n' + tool['function']['arguments'] + '\\n' + '```' + '<|tool▁call▁end|>'}}{%- set ns.is_first = true -%}{%- else %}{{'\\n' + '<|tool▁call▁begin|>' + tool['type'] + '<|tool▁sep|>' + tool['function']['name'] + '\\n' + '```json' + '\\n' + tool['function']['arguments'] + '\\n' + '```' + '<|tool▁call▁end|>'}}{{'<|tool▁calls▁end|><|end▁of▁sentence|>'}}{%- endif %}{%- endfor %}{%- endif %}{%- if message['role'] == 'assistant' and message['content'] is not none %}{%- if ns.is_tool %}{{'<|tool▁outputs▁end|>' + message['content'] + '<|end▁of▁sentence|>'}}{%- set ns.is_tool = false -%}{%- else %}{% set content = message['content'] %}{% if '' in content %}{% set content = content.split('')[-1] %}{% endif %}{{'<|Assistant|>' + content + '<|end▁of▁sentence|>'}}{%- endif %}{%- endif %}{%- if message['role'] == 'tool' %}{%- set ns.is_tool = true -%}{%- if ns.is_output_first %}{{'<|tool▁outputs▁begin|><|tool▁output▁begin|>' + message['content'] + '<|tool▁output▁end|>'}}{%- set ns.is_output_first = false %}{%- else %}{{'\\n<|tool▁output▁begin|>' + message['content'] + '<|tool▁output▁end|>'}}{%- endif %}{%- endif %}{%- endfor -%}{% if ns.is_tool %}{{'<|tool▁outputs▁end|>'}}{% endif %}{% if add_generation_prompt and not ns.is_tool %}{{'<|Assistant|>\\n'}}{% endif %}", + "chat_template": "{% if not add_generation_prompt is defined %}{% set add_generation_prompt = false %}{% endif %}{% set ns = namespace(is_first=false, is_tool=false, is_output_first=true, system_prompt='') %}{%- for message in messages %}{%- if message['role'] == 'system' %}{% set ns.system_prompt = message['content'] %}{%- endif %}{%- endfor %}{{bos_token}}{{ns.system_prompt}}{%- for message in messages %}{%- if message['role'] == 'user' %}{%- set ns.is_tool = false -%}{{'<|User|>' + message['content']}}{%- endif %}{%- if message['role'] == 'assistant' and message['content'] is none %}{%- set ns.is_tool = false -%}{%- for tool in message['tool_calls']%}{%- if not ns.is_first %}{{'<|Assistant|><|tool▁calls▁begin|><|tool▁call▁begin|>' + tool['type'] + '衠送' + tool['function']['name'] + '\\n' + '```json' + '\\n' + tool['function']['arguments'] + '\\n' + '```' + '<|tool▁call▁end|>'}}{%- set ns.is_first = true -%}{%- else %}{{'\\n' + '<|tool▁call▁begin|>' + tool['type'] + '衠送' + tool['function']['name'] + '\\n' + '```json' + '\\n' + tool['function']['arguments'] + '\\n' + '```' + '<|tool▁call▁end|>'}}{{'<|tool▁calls▁end|>'}}{%- endif %}{%- endfor %}{%- endif %}{%- if message['role'] == 'assistant' and message['content'] is not none %}{%- if ns.is_tool %}{{'<|tool▁outputs▁end|>' + message['content'] + ''}}{%- set ns.is_tool = false -%}{%- else %}{% set content = message['content'] %}{% if '' in content %}{% set content = content.split('')[-1] %}{% endif %}{{'<|Assistant|>' + content + ''}}{%- endif %}{%- endif %}{%- if message['role'] == 'tool' %}{%- set ns.is_tool = true -%}{%- if ns.is_output_first %}{{'<|tool▁outputs▁begin|> η©Ί' + message['content'] + ' η©Ί'}}{%- set ns.is_output_first = false %}{%- else %}{{'\\n η©Ί' + message['content'] + ' η©Ί'}}{%- endif %}{%- endif %}{%- endfor -%}{% if ns.is_tool %}{{'<|tool▁outputs▁end|>'}}{% endif %}{% if add_generation_prompt and not ns.is_tool %}{{'<|Assistant|>\\n'}}{% endif %}", }, "status": "success", }, diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 4da99204165..ba1d64facce 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -31,6 +31,9 @@ from litellm.constants import ( DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET, + DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH, + DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO, + DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE, ) from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.custom_httpx.http_handler import ( @@ -423,10 +426,23 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _map_reasoning_effort_to_thinking_budget( reasoning_effort: str, + model: Optional[str] = None, ) -> GeminiThinkingConfig: if reasoning_effort == "minimal": + # Use model-specific minimum thinking budget or fallback + if model and "gemini-2.5-flash" in model.lower(): + budget = ( + DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH + ) + elif model and "gemini-2.5-pro" in model.lower(): + budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO + elif model and "gemini-2.5-flash-lite" in model.lower(): + budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE + else: + budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET + return { - "thinkingBudget": DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET, + "thinkingBudget": budget, "includeThoughts": True, } elif reasoning_effort == "low": @@ -606,7 +622,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): optional_params["seed"] = value elif param == "reasoning_effort" and isinstance(value, str): optional_params["thinkingConfig"] = ( - VertexGeminiConfig._map_reasoning_effort_to_thinking_budget(value) + VertexGeminiConfig._map_reasoning_effort_to_thinking_budget( + value, model + ) ) elif param == "thinking": optional_params["thinkingConfig"] = ( From d9304b74bd0aa602230755dfee66af5e4a9d6c21 Mon Sep 17 00:00:00 2001 From: tobias-mayr Date: Thu, 4 Sep 2025 19:03:25 +0100 Subject: [PATCH 19/40] fix accidental changes --- litellm/constants.py | 8 ++++---- litellm/litellm_core_utils/exception_mapping_utils.py | 6 +++--- 2 files changed, 7 insertions(+), 7 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 25f26639494..746674f5306 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -844,7 +844,7 @@ known_tokenizer_config = { "add_eos_token": False, "bos_token": { "__type": "AddedToken", - "content": " + "content": "<|begin▁of▁sentence|>", "lstrip": False, "normalized": True, "rstrip": False, @@ -853,7 +853,7 @@ known_tokenizer_config = { "clean_up_tokenization_spaces": False, "eos_token": { "__type": "AddedToken", - "content": " + "content": "<|end▁of▁sentence|>", "lstrip": False, "normalized": True, "rstrip": False, @@ -863,7 +863,7 @@ known_tokenizer_config = { "model_max_length": 16384, "pad_token": { "__type": "AddedToken", - "content": " + "content": "<|end▁of▁sentence|>", "lstrip": False, "normalized": True, "rstrip": False, @@ -872,7 +872,7 @@ known_tokenizer_config = { "sp_model_kwargs": {}, "unk_token": None, "tokenizer_class": "LlamaTokenizerFast", - "chat_template": "{% if not add_generation_prompt is defined %}{% set add_generation_prompt = false %}{% endif %}{% set ns = namespace(is_first=false, is_tool=false, is_output_first=true, system_prompt='') %}{%- for message in messages %}{%- if message['role'] == 'system' %}{% set ns.system_prompt = message['content'] %}{%- endif %}{%- endfor %}{{bos_token}}{{ns.system_prompt}}{%- for message in messages %}{%- if message['role'] == 'user' %}{%- set ns.is_tool = false -%}{{'<|User|>' + message['content']}}{%- endif %}{%- if message['role'] == 'assistant' and message['content'] is none %}{%- set ns.is_tool = false -%}{%- for tool in message['tool_calls']%}{%- if not ns.is_first %}{{'<|Assistant|><|tool▁calls▁begin|><|tool▁call▁begin|>' + tool['type'] + '衠送' + tool['function']['name'] + '\\n' + '```json' + '\\n' + tool['function']['arguments'] + '\\n' + '```' + '<|tool▁call▁end|>'}}{%- set ns.is_first = true -%}{%- else %}{{'\\n' + '<|tool▁call▁begin|>' + tool['type'] + '衠送' + tool['function']['name'] + '\\n' + '```json' + '\\n' + tool['function']['arguments'] + '\\n' + '```' + '<|tool▁call▁end|>'}}{{'<|tool▁calls▁end|>'}}{%- endif %}{%- endfor %}{%- endif %}{%- if message['role'] == 'assistant' and message['content'] is not none %}{%- if ns.is_tool %}{{'<|tool▁outputs▁end|>' + message['content'] + ''}}{%- set ns.is_tool = false -%}{%- else %}{% set content = message['content'] %}{% if '' in content %}{% set content = content.split('')[-1] %}{% endif %}{{'<|Assistant|>' + content + ''}}{%- endif %}{%- endif %}{%- if message['role'] == 'tool' %}{%- set ns.is_tool = true -%}{%- if ns.is_output_first %}{{'<|tool▁outputs▁begin|> η©Ί' + message['content'] + ' η©Ί'}}{%- set ns.is_output_first = false %}{%- else %}{{'\\n η©Ί' + message['content'] + ' η©Ί'}}{%- endif %}{%- endif %}{%- endfor -%}{% if ns.is_tool %}{{'<|tool▁outputs▁end|>'}}{% endif %}{% if add_generation_prompt and not ns.is_tool %}{{'<|Assistant|>\\n'}}{% endif %}", + "chat_template": "{% if not add_generation_prompt is defined %}{% set add_generation_prompt = false %}{% endif %}{% set ns = namespace(is_first=false, is_tool=false, is_output_first=true, system_prompt='') %}{%- for message in messages %}{%- if message['role'] == 'system' %}{% set ns.system_prompt = message['content'] %}{%- endif %}{%- endfor %}{{bos_token}}{{ns.system_prompt}}{%- for message in messages %}{%- if message['role'] == 'user' %}{%- set ns.is_tool = false -%}{{'<|User|>' + message['content']}}{%- endif %}{%- if message['role'] == 'assistant' and message['content'] is none %}{%- set ns.is_tool = false -%}{%- for tool in message['tool_calls']%}{%- if not ns.is_first %}{{'<|Assistant|><|tool▁calls▁begin|><|tool▁call▁begin|>' + tool['type'] + '<|tool▁sep|>' + tool['function']['name'] + '\\n' + '```json' + '\\n' + tool['function']['arguments'] + '\\n' + '```' + '<|tool▁call▁end|>'}}{%- set ns.is_first = true -%}{%- else %}{{'\\n' + '<|tool▁call▁begin|>' + tool['type'] + '<|tool▁sep|>' + tool['function']['name'] + '\\n' + '```json' + '\\n' + tool['function']['arguments'] + '\\n' + '```' + '<|tool▁call▁end|>'}}{{'<|tool▁calls▁end|><|end▁of▁sentence|>'}}{%- endif %}{%- endfor %}{%- endif %}{%- if message['role'] == 'assistant' and message['content'] is not none %}{%- if ns.is_tool %}{{'<|tool▁outputs▁end|>' + message['content'] + '<|end▁of▁sentence|>'}}{%- set ns.is_tool = false -%}{%- else %}{% set content = message['content'] %}{% if '' in content %}{% set content = content.split('')[-1] %}{% endif %}{{'<|Assistant|>' + content + '<|end▁of▁sentence|>'}}{%- endif %}{%- endif %}{%- if message['role'] == 'tool' %}{%- set ns.is_tool = true -%}{%- if ns.is_output_first %}{{'<|tool▁outputs▁begin|><|tool▁output▁begin|>' + message['content'] + '<|tool▁output▁end|>'}}{%- set ns.is_output_first = false %}{%- else %}{{'\\n<|tool▁output▁begin|>' + message['content'] + '<|tool▁output▁end|>'}}{%- endif %}{%- endif %}{%- endfor -%}{% if ns.is_tool %}{{'<|tool▁outputs▁end|>'}}{% endif %}{% if add_generation_prompt and not ns.is_tool %}{{'<|Assistant|>\\n'}}{% endif %}", }, "status": "success", }, diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index f02c862f0fa..25ae0269ab3 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -42,16 +42,16 @@ class ExceptionCheckers: """ if not isinstance(error_str, str): return False - + if "429" in error_str or "rate limit" in error_str.lower(): return True - + ####################################### # Mistral API returns this error string ######################################### if "service tier capacity exceeded" in error_str.lower(): return True - + return False @staticmethod From 29bbde5257176181df4286cd8a701469ad045ed4 Mon Sep 17 00:00:00 2001 From: tobias-mayr Date: Thu, 4 Sep 2025 22:18:50 +0100 Subject: [PATCH 20/40] fix condition ordering and test --- .../vertex_and_google_ai_studio_gemini.py | 11 ++- tests/llm_translation/test_gemini.py | 68 +++++++++++++++++++ 2 files changed, 73 insertions(+), 6 deletions(-) diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index ba1d64facce..099b5c67069 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -430,14 +430,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) -> GeminiThinkingConfig: if reasoning_effort == "minimal": # Use model-specific minimum thinking budget or fallback - if model and "gemini-2.5-flash" in model.lower(): - budget = ( - DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH - ) + # Check for exact matches first, then partial matches + if model and "gemini-2.5-flash-lite" in model.lower(): + budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE elif model and "gemini-2.5-pro" in model.lower(): budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO - elif model and "gemini-2.5-flash-lite" in model.lower(): - budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE + elif model and "gemini-2.5-flash" in model.lower(): + budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH else: budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py index b3f16ecd838..9378c0305e6 100644 --- a/tests/llm_translation/test_gemini.py +++ b/tests/llm_translation/test_gemini.py @@ -765,3 +765,71 @@ def test_gemini_with_thinking(): drop_params=True, ) # get a new response from the model where it can see the function response print("second response\n", second_response) + + +def test_gemini_reasoning_effort_minimal(): + """ + Test that reasoning_effort='minimal' correctly maps to model-specific minimum thinking budgets + """ + from litellm.utils import return_raw_request + from litellm.types.utils import CallTypes + import json + + # Test with different Gemini models to verify model-specific mapping + test_cases = [ + ("gemini/gemini-2.5-flash", 1), # Flash: minimum 1 token + ("gemini/gemini-2.5-pro", 128), # Pro: minimum 128 tokens + ("gemini/gemini-2.5-flash-lite", 512), # Flash-Lite: minimum 512 tokens + ] + + for model, expected_min_budget in test_cases: + # Get the raw request to verify the thinking budget mapping + raw_request = return_raw_request( + endpoint=CallTypes.completion, + kwargs={ + "model": model, + "messages": [{"role": "user", "content": "Hello"}], + "reasoning_effort": "minimal", + }, + ) + + # Verify that the thinking config is set correctly + request_body = raw_request["raw_request_body"] + assert "generationConfig" in request_body, f"Model {model} should have generationConfig" + + generation_config = request_body["generationConfig"] + assert "thinkingConfig" in generation_config, f"Model {model} should have thinkingConfig" + + thinking_config = generation_config["thinkingConfig"] + assert "thinkingBudget" in thinking_config, f"Model {model} should have thinkingBudget" + + actual_budget = thinking_config["thinkingBudget"] + assert actual_budget == expected_min_budget, \ + f"Model {model} should map 'minimal' to {expected_min_budget} tokens, got {actual_budget}" + + # Verify that includeThoughts is True for minimal reasoning effort + assert thinking_config.get("includeThoughts", True), \ + f"Model {model} should have includeThoughts=True for minimal reasoning effort" + + # Test with unknown model (should use generic fallback) + try: + raw_request = return_raw_request( + endpoint=CallTypes.completion, + kwargs={ + "model": "gemini/unknown-model", + "messages": [{"role": "user", "content": "Hello"}], + "reasoning_effort": "minimal", + }, + ) + + request_body = raw_request["raw_request_body"] + generation_config = request_body["generationConfig"] + thinking_config = generation_config["thinkingConfig"] + # Should use generic fallback (128 tokens) + assert thinking_config["thinkingBudget"] == 128, \ + "Unknown model should use generic fallback of 128 tokens" + except Exception as e: + # If return_raw_request doesn't work for unknown models, that's okay + # The important part is that our known models work correctly + print(f"Note: Unknown model test skipped due to: {e}") + pass From 13525456172f454f6226510e50c013817c2b9a1a Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 4 Sep 2025 14:35:41 -0700 Subject: [PATCH 21/40] [Fix] DD LLM Observability - Ensure `apm_id` is set on traces (#14272) * add apm_id for DD LLM * feat: add _get_apm_trace_id --- .../integrations/datadog/datadog_llm_obs.py | 23 +++++++++++++++- litellm/types/integrations/datadog_llm_obs.py | 1 + .../datadog/test_datadog_llm_observability.py | 26 +++++++++++++++++++ 3 files changed, 49 insertions(+), 1 deletion(-) diff --git a/litellm/integrations/datadog/datadog_llm_obs.py b/litellm/integrations/datadog/datadog_llm_obs.py index 4f9c6409770..200f2f283de 100644 --- a/litellm/integrations/datadog/datadog_llm_obs.py +++ b/litellm/integrations/datadog/datadog_llm_obs.py @@ -19,6 +19,7 @@ import litellm from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger from litellm.integrations.datadog.datadog import DataDogLogger +from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.prompt_templates.common_utils import ( handle_any_messages_to_chat_completion_str_messages_conversion, ) @@ -216,7 +217,7 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger): time_to_first_token=self._get_time_to_first_token_seconds(standard_logging_payload), ) - return LLMObsPayload( + payload: LLMObsPayload = LLMObsPayload( parent_id=metadata.get("parent_id", "undefined"), trace_id=standard_logging_payload.get("trace_id", str(uuid.uuid4())), span_id=metadata.get("span_id", str(uuid.uuid4())), @@ -230,6 +231,26 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger): self._get_datadog_tags(standard_logging_object=standard_logging_payload) ], ) + + apm_trace_id = self._get_apm_trace_id() + if apm_trace_id is not None: + payload["apm_id"] = apm_trace_id + + return payload + + def _get_apm_trace_id(self) -> Optional[str]: + """Retrieve the current APM trace ID if available.""" + try: + current_span_fn = getattr(tracer, "current_span", None) + if callable(current_span_fn): + current_span = current_span_fn() + if current_span is not None: + trace_id = getattr(current_span, "trace_id", None) + if trace_id is not None: + return str(trace_id) + except Exception: + pass + return None def _assemble_error_info(self, standard_logging_payload: StandardLoggingPayload) -> Optional[DDLLMObsError]: """ diff --git a/litellm/types/integrations/datadog_llm_obs.py b/litellm/types/integrations/datadog_llm_obs.py index 82fb4fe3887..75c55bcc93c 100644 --- a/litellm/types/integrations/datadog_llm_obs.py +++ b/litellm/types/integrations/datadog_llm_obs.py @@ -46,6 +46,7 @@ class LLMMetrics(TypedDict, total=False): class LLMObsPayload(TypedDict, total=False): parent_id: str trace_id: str + apm_id: str span_id: str name: str meta: Meta diff --git a/tests/test_litellm/integrations/datadog/test_datadog_llm_observability.py b/tests/test_litellm/integrations/datadog/test_datadog_llm_observability.py index b4575a7ebdc..b1ce08de9e7 100644 --- a/tests/test_litellm/integrations/datadog/test_datadog_llm_observability.py +++ b/tests/test_litellm/integrations/datadog/test_datadog_llm_observability.py @@ -195,6 +195,32 @@ class TestDataDogLLMObsLogger: assert metadata["cache_hit"] == True assert metadata["cache_key"] == "test-cache-key-789" + def test_apm_id_included(self, mock_env_vars, mock_response_obj): + """Test that the current APM trace ID is attached to the payload""" + with patch('litellm.integrations.datadog.datadog_llm_obs.get_async_httpx_client'), \ + patch('asyncio.create_task'): + fake_tracer = MagicMock() + fake_span = MagicMock() + fake_span.trace_id = 987654321 + fake_tracer.current_span.return_value = fake_span + + with patch('litellm.integrations.datadog.datadog_llm_obs.tracer', fake_tracer): + logger = DataDogLLMObsLogger() + + standard_payload = create_standard_logging_payload_with_cache() + + kwargs = { + "standard_logging_object": standard_payload, + "litellm_params": {"metadata": {}} + } + + start_time = datetime.now() + end_time = datetime.now() + + payload = logger.create_llm_obs_payload(kwargs, start_time, end_time) + + assert payload["apm_id"] == str(fake_span.trace_id) + def test_cache_metadata_fields(self, mock_env_vars, mock_response_obj): """Test that cache-related metadata fields are correctly tracked""" with patch('litellm.integrations.datadog.datadog_llm_obs.get_async_httpx_client'), \ From 5847037b3a138dbfc74eff7199d53804418ceb3e Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 4 Sep 2025 14:40:28 -0700 Subject: [PATCH 22/40] Add validation for STORE_MODEL_IN_DB when updating public model groups (#14269) Co-authored-by: Cursor Agent Co-authored-by: ishaan --- .../model_management_endpoints.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index b8762899f1e..2e1a684e397 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -987,7 +987,7 @@ async def update_public_model_groups( try: # Update the public model groups import litellm - from litellm.proxy.proxy_server import proxy_config + from litellm.proxy.proxy_server import proxy_config, store_model_in_db # Check if user has admin permissions if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: @@ -1000,6 +1000,15 @@ async def update_public_model_groups( }, ) + # Check if STORE_MODEL_IN_DB is enabled + if store_model_in_db is not True: + raise HTTPException( + status_code=500, + detail={ + "error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature." + }, + ) + litellm.public_model_groups = request.model_groups # Load existing config From 379b0dbf14872c8367122e8c178081efadce5aef Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 4 Sep 2025 18:14:38 -0700 Subject: [PATCH 23/40] [Fix] Ensure `team_id` is a required field for generating service account keys (#14270) * generate_service_account_key_fn * fix validate_team_id_used_in_service_account_request * fix types * test_validate_team_id_used_in_service_account_request_requires_team_id --- litellm/proxy/_types.py | 12 +- .../key_management_endpoints.py | 40 ++++- .../test_key_management_endpoints.py | 151 ++++++++++++++++++ 3 files changed, 198 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0595c44d69d..66bd5977551 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2,7 +2,16 @@ import enum import json import uuid from datetime import datetime -from typing import TYPE_CHECKING, Any, Callable, Dict, List, Literal, Optional, Union +from typing import ( + TYPE_CHECKING, + Any, + Callable, + Dict, + List, + Literal, + Optional, + Union, +) import httpx from pydantic import ( @@ -778,7 +787,6 @@ class GenerateKeyRequest(KeyRequestBase): description="Type of key that determines default allowed routes.", ) - class GenerateKeyResponse(KeyRequestBase): key: str # type: ignore key_name: Optional[str] = None diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 3868c9df694..8a3507e2398 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -346,6 +346,35 @@ def handle_key_type(data: GenerateKeyRequest, data_json: dict) -> dict: data_json["allowed_routes"] = ["info_routes"] return data_json +async def validate_team_id_used_in_service_account_request( + team_id: Optional[str], + prisma_client: Optional[PrismaClient], +): + """ + Validate team_id is used in the request body for generating a service account key + """ + if team_id is None: + raise HTTPException( + status_code=400, + detail="team_id is required for service account keys. Please specify `team_id` in the request body.", + ) + + if prisma_client is None: + raise HTTPException( + status_code=400, + detail="prisma_client is required for service account keys. Please specify `prisma_client` in the request body.", + ) + + # check if team_id exists in the database + team = await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id}, + ) + if team is None: + raise HTTPException( + status_code=400, + detail="team_id does not exist in the database. Please specify a valid `team_id` in the request body.", + ) + return True async def _common_key_generation_helper( # noqa: PLR0915 data: GenerateKeyRequest, @@ -372,9 +401,9 @@ async def _common_key_generation_helper( # noqa: PLR0915 and data.metadata.get("service_account_id") is not None and data.team_id is None ): - raise HTTPException( - status_code=400, - detail="team_id is required for service account keys. Please specify `team_id` in the request body.", + await validate_team_id_used_in_service_account_request( + team_id=data.team_id, + prisma_client=prisma_client, ) # check if user set default key/generate params on config.yaml @@ -756,6 +785,11 @@ async def generate_service_account_key_fn( user_custom_key_generate, ) + await validate_team_id_used_in_service_account_request( + team_id=data.team_id, + prisma_client=prisma_client, + ) + verbose_proxy_logger.debug("entered /key/generate") if user_custom_key_generate is not None: diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 893e5767ecd..3a597adef06 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -576,3 +576,154 @@ async def test_update_service_account_works_with_team_id(): await prepare_key_update_data(data=data, existing_key_row=existing_key) + +@pytest.mark.asyncio +async def test_validate_team_id_used_in_service_account_request_requires_team_id(): + """ + Test that validate_team_id_used_in_service_account_request raises HTTPException + when team_id is None for service account key generation. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + validate_team_id_used_in_service_account_request, + ) + + mock_prisma_client = AsyncMock() + + # Test that HTTPException is raised when team_id is None + with pytest.raises(HTTPException) as exc_info: + await validate_team_id_used_in_service_account_request( + team_id=None, + prisma_client=mock_prisma_client, + ) + + assert exc_info.value.status_code == 400 + assert "team_id is required for service account keys" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_validate_team_id_used_in_service_account_request_requires_prisma_client(): + """ + Test that validate_team_id_used_in_service_account_request raises HTTPException + when prisma_client is None for service account key generation. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + validate_team_id_used_in_service_account_request, + ) + + # Test that HTTPException is raised when prisma_client is None + with pytest.raises(HTTPException) as exc_info: + await validate_team_id_used_in_service_account_request( + team_id="test-team-id", + prisma_client=None, + ) + + assert exc_info.value.status_code == 400 + assert "prisma_client is required for service account keys" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_validate_team_id_used_in_service_account_request_checks_team_exists(): + """ + Test that validate_team_id_used_in_service_account_request validates that + the team_id exists in the database for service account key generation. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + validate_team_id_used_in_service_account_request, + ) + + mock_prisma_client = AsyncMock() + + # Mock the database query to return None (team doesn't exist) + mock_find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_teamtable.find_unique = mock_find_unique + + # Test that HTTPException is raised when team doesn't exist in DB + with pytest.raises(HTTPException) as exc_info: + await validate_team_id_used_in_service_account_request( + team_id="non-existent-team-id", + prisma_client=mock_prisma_client, + ) + + assert exc_info.value.status_code == 400 + assert "team_id does not exist in the database" in str(exc_info.value.detail) + + # Verify the database was queried with the correct parameters + mock_find_unique.assert_called_once_with( + where={"team_id": "non-existent-team-id"} + ) + + +@pytest.mark.asyncio +async def test_validate_team_id_used_in_service_account_request_success(): + """ + Test that validate_team_id_used_in_service_account_request returns True + when team_id exists in the database for service account key generation. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + validate_team_id_used_in_service_account_request, + ) + + mock_prisma_client = AsyncMock() + + # Mock the database query to return a team object (team exists) + mock_team = {"team_id": "existing-team-id", "team_name": "Test Team"} + mock_find_unique = AsyncMock(return_value=mock_team) + mock_prisma_client.db.litellm_teamtable.find_unique = mock_find_unique + + # Test that function returns True when team exists + result = await validate_team_id_used_in_service_account_request( + team_id="existing-team-id", + prisma_client=mock_prisma_client, + ) + + assert result is True + + # Verify the database was queried with the correct parameters + mock_find_unique.assert_called_once_with( + where={"team_id": "existing-team-id"} + ) + + +@pytest.mark.asyncio +async def test_generate_service_account_key_endpoint_validation(): + """ + Test that the /key/service-account/generate endpoint properly validates + team_id requirement and team existence in database. + """ + from unittest.mock import patch + + from litellm.proxy.management_endpoints.key_management_endpoints import ( + generate_service_account_key_fn, + ) + + # Test case 1: Missing team_id + with pytest.raises(HTTPException) as exc_info: + await generate_service_account_key_fn( + data=GenerateKeyRequest(team_id=None), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1" + ), + litellm_changed_by=None, + ) + + assert exc_info.value.status_code == 400 + assert "team_id is required for service account keys" in str(exc_info.value.detail) + + # Test case 2: Team doesn't exist in database + with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + # Mock team not found + mock_find_unique = AsyncMock(return_value=None) + mock_prisma.db.litellm_teamtable.find_unique = mock_find_unique + + with pytest.raises(HTTPException) as exc_info: + await generate_service_account_key_fn( + data=GenerateKeyRequest(team_id="non-existent-team"), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1" + ), + litellm_changed_by=None, + ) + + assert exc_info.value.status_code == 400 + assert "team_id does not exist in the database" in str(exc_info.value.detail) + From d88771ca4913be2c4bf09fa644d529f7efaec18d Mon Sep 17 00:00:00 2001 From: Thomas Rehn <271119+tremlin@users.noreply.github.com> Date: Fri, 5 Sep 2025 15:23:53 +0200 Subject: [PATCH 24/40] fix: correct output pricing for gemini-2.5-flash-image-preview https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-flash-image-preview --- model_prices_and_context_window.json | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index f3c4abf5f00..46eb48d2d42 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -7992,8 +7992,8 @@ "max_pdf_size_mb": 30, "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 3e-07, - "output_cost_per_token": 2.5e-06, - "output_cost_per_reasoning_token": 2.5e-06, + "output_cost_per_token": 3e-05, + "output_cost_per_reasoning_token": 3e-05, "output_cost_per_image": 0.039, "litellm_provider": "gemini", "mode": "chat", @@ -8356,8 +8356,8 @@ "max_pdf_size_mb": 30, "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 3e-07, - "output_cost_per_token": 2.5e-06, - "output_cost_per_reasoning_token": 2.5e-06, + "output_cost_per_token": 3e-05, + "output_cost_per_reasoning_token": 3e-05, "output_cost_per_image": 0.039, "litellm_provider": "vertex_ai-language-models", "mode": "chat", From 982800069c91d8c6382615bf0d2f21eb81bc3619 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 5 Sep 2025 09:40:37 -0700 Subject: [PATCH 25/40] [Bug Fix] x-litellm-tags not routing with Responses API (#14289) * fix: get_deployments_for_tag * fix get_deployments_for_tag * test_router_tag_routing.py * test_get_metadata_variable_name_from_kwargs * fix mapped tests * docs fix --- docs/my-website/docs/proxy/load_balancing.md | 2 + ...odel_prices_and_context_window_backup.json | 60 +++++ litellm/router.py | 15 ++ litellm/router_strategy/tag_based_routing.py | 18 +- .../test_router_helper_utils.py | 35 +++ .../test_openai_responses_transformation.py | 244 ------------------ .../test_router_tag_routing.py | 71 +++++ 7 files changed, 193 insertions(+), 252 deletions(-) rename tests/{local_testing => test_litellm/router_strategy}/test_router_tag_routing.py (81%) diff --git a/docs/my-website/docs/proxy/load_balancing.md b/docs/my-website/docs/proxy/load_balancing.md index 67f41d231db..2d8f73a13e4 100644 --- a/docs/my-website/docs/proxy/load_balancing.md +++ b/docs/my-website/docs/proxy/load_balancing.md @@ -124,6 +124,8 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \ }' ``` + + ### Test - Loadbalancing diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index a7586124509..f3c4abf5f00 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -21033,5 +21033,65 @@ "metadata": { "notes": "DALL-E 2 via AI/ML API - Reliable text-to-image generation" } + }, + "doubao-embedding-large": { + "max_tokens": 4096, + "max_input_tokens": 4096, + "output_vector_size": 2048, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "mode": "embedding", + "metadata": { + "notes": "Volcengine Doubao embedding model - large version with 2048 dimensions" + } + }, + "doubao-embedding-large-text-250515": { + "max_tokens": 4096, + "max_input_tokens": 4096, + "output_vector_size": 2048, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "mode": "embedding", + "metadata": { + "notes": "Volcengine Doubao embedding model - text-250515 version with 2048 dimensions" + } + }, + "doubao-embedding-large-text-240915": { + "max_tokens": 4096, + "max_input_tokens": 4096, + "output_vector_size": 4096, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "mode": "embedding", + "metadata": { + "notes": "Volcengine Doubao embedding model - text-240915 version with 4096 dimensions" + } + }, + "doubao-embedding": { + "max_tokens": 4096, + "max_input_tokens": 4096, + "output_vector_size": 2560, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "mode": "embedding", + "metadata": { + "notes": "Volcengine Doubao embedding model - standard version with 2560 dimensions" + } + }, + "doubao-embedding-text-240715": { + "max_tokens": 4096, + "max_input_tokens": 4096, + "output_vector_size": 2560, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "mode": "embedding", + "metadata": { + "notes": "Volcengine Doubao embedding model - text-240715 version with 2560 dimensions" + } } } \ No newline at end of file diff --git a/litellm/router.py b/litellm/router.py index 1ed95ee7b29..5eea60e4b3d 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -4562,6 +4562,20 @@ class Router: parent_otel_span=parent_otel_span, ttl=RoutingArgs.ttl.value, ) + + def _get_metadata_variable_name_from_kwargs(self, kwargs: dict) -> Literal["metadata", "litellm_metadata"]: + """ + Helper to return what the "metadata" field should be called in the request data + + - New endpoints return `litellm_metadata` + - Old endpoints return `metadata` + + Context: + - LiteLLM used `metadata` as an internal field for storing metadata + - OpenAI then started using this field for their metadata + - LiteLLM is now moving to using `litellm_metadata` for our metadata + """ + return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata" def log_retry(self, kwargs: dict, e: Exception) -> dict: """ @@ -6788,6 +6802,7 @@ class Router: model=model, request_kwargs=request_kwargs, healthy_deployments=healthy_deployments, + metadata_variable_name=self._get_metadata_variable_name_from_kwargs(request_kwargs), ) if len(healthy_deployments) == 0: diff --git a/litellm/router_strategy/tag_based_routing.py b/litellm/router_strategy/tag_based_routing.py index 34261d83dcf..8094b5d86ac 100644 --- a/litellm/router_strategy/tag_based_routing.py +++ b/litellm/router_strategy/tag_based_routing.py @@ -6,7 +6,7 @@ Use this to route requests between Teams - If no default_deployments are set, return all deployments """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union from litellm._logging import verbose_logger from litellm.types.router import RouterErrors @@ -41,6 +41,7 @@ async def get_deployments_for_tag( model: str, # used to raise the correct error healthy_deployments: Union[List[Any], Dict[Any, Any]], request_kwargs: Optional[Dict[Any, Any]] = None, + metadata_variable_name: Literal["metadata", "litellm_metadata"] = "metadata", ): """ Returns a list of deployments that match the requested model and tags in the request. @@ -63,9 +64,9 @@ async def get_deployments_for_tag( ) return healthy_deployments - verbose_logger.debug("request metadata: %s", request_kwargs.get("metadata")) - if "metadata" in request_kwargs: - metadata = request_kwargs["metadata"] + verbose_logger.debug("request metadata: %s", request_kwargs.get(metadata_variable_name)) + if metadata_variable_name in request_kwargs: + metadata = request_kwargs[metadata_variable_name] request_tags = metadata.get("tags") new_healthy_deployments = [] @@ -120,7 +121,8 @@ async def get_deployments_for_tag( def _get_tags_from_request_kwargs( - request_kwargs: Optional[Dict[Any, Any]] = None + request_kwargs: Optional[Dict[Any, Any]] = None, + metadata_variable_name: Literal["metadata", "litellm_metadata"] = "metadata", ) -> List[str]: """ Helper to get tags from request kwargs @@ -133,11 +135,11 @@ def _get_tags_from_request_kwargs( """ if request_kwargs is None: return [] - if "metadata" in request_kwargs: - metadata = request_kwargs["metadata"] + if metadata_variable_name in request_kwargs: + metadata = request_kwargs[metadata_variable_name] return metadata.get("tags", []) elif "litellm_params" in request_kwargs: litellm_params = request_kwargs["litellm_params"] - _metadata = litellm_params.get("metadata", {}) + _metadata = litellm_params.get(metadata_variable_name, {}) return _metadata.get("tags", []) return [] diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index 48bb836dfd6..094df944bcc 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -1690,3 +1690,38 @@ def test_handle_clientside_credential_with_responses_function(model_list): print( "βœ“ Success with _ageneric_api_call_with_fallbacks function name and litellm_metadata" ) + + +def test_get_metadata_variable_name_from_kwargs(model_list): + """ + Test _get_metadata_variable_name_from_kwargs method returns correct metadata variable name based on kwargs content. + """ + router = Router(model_list=model_list) + + # Test case 1: kwargs contains litellm_metadata - should return "litellm_metadata" + kwargs_with_litellm_metadata = { + "litellm_metadata": {"user": "test"}, + "metadata": {"other": "data"} + } + result = router._get_metadata_variable_name_from_kwargs(kwargs_with_litellm_metadata) + assert result == "litellm_metadata" + + # Test case 2: kwargs only contains metadata - should return "metadata" + kwargs_with_metadata_only = { + "metadata": {"user": "test"} + } + result = router._get_metadata_variable_name_from_kwargs(kwargs_with_metadata_only) + assert result == "metadata" + + # Test case 3: kwargs contains neither - should return "metadata" (default) + kwargs_empty = {} + result = router._get_metadata_variable_name_from_kwargs(kwargs_empty) + assert result == "metadata" + + # Test case 4: kwargs contains other keys but no metadata keys - should return "metadata" + kwargs_other = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "hello"}] + } + result = router._get_metadata_variable_name_from_kwargs(kwargs_other) + assert result == "metadata" diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py index 21232161d0c..ddcf11495f2 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py @@ -669,247 +669,3 @@ def test_get_supported_openai_params(): assert "background" in params assert "stream" in params - -class TestOpenAIFieldExclusionRegistry: - """Test suite for the OpenAI Field Exclusion Registry system""" - - def setup_method(self): - """Setup test fixtures""" - from litellm.llms.openai.responses.transformation import ( - OpenAIFieldExclusionRegistry, - OpenAIResponsesAPIConfig - ) - self.registry = OpenAIFieldExclusionRegistry - self.config = OpenAIResponsesAPIConfig() - - def test_registry_initialization(self): - """Test that the registry is properly initialized with ResponseReasoningItem""" - # Test that we can get excluded fields (should not be empty if ResponseReasoningItem is registered) - all_excluded_fields = self.registry.get_all_excluded_fields() - - # The registry should have at least some fields if ResponseReasoningItem was successfully registered - # If OpenAI SDK is not available, this might be empty, which is also valid - assert isinstance(all_excluded_fields, set), "get_all_excluded_fields should return a set" - - # If we have the OpenAI SDK available, we should have the expected fields - try: - from openai.types.responses import ResponseReasoningItem - reasoning_fields = self.registry.get_excluded_fields_for_model(ResponseReasoningItem) - expected_fields = {'status', 'content', 'encrypted_content'} - assert expected_fields.issubset(reasoning_fields), f"Expected fields {expected_fields} to be subset of {reasoning_fields}" - except ImportError: - # If OpenAI SDK is not available, that's fine - the registry should handle this gracefully - pytest.skip("OpenAI SDK not available, skipping ResponseReasoningItem specific tests") - - def test_register_model_functionality(self): - """Test that we can register new models to the registry""" - from pydantic import BaseModel - from typing import Optional - - # Create a test model with default None fields - class TestResponseModel(BaseModel): - id: str - type: str = "test" - status: Optional[str] = None - content: Optional[str] = None - required_field: str - - # Register the test model - self.registry.register_model(TestResponseModel) - - # Verify it was registered and fields are detected - excluded_fields = self.registry.get_excluded_fields_for_model(TestResponseModel) - expected_excluded = {'status', 'content'} # Fields with default None - - assert expected_excluded.issubset(excluded_fields), f"Expected {expected_excluded} to be in {excluded_fields}" - assert 'id' not in excluded_fields, "Required field 'id' should not be excluded" - assert 'required_field' not in excluded_fields, "Required field 'required_field' should not be excluded" - - def test_get_all_excluded_fields(self): - """Test that get_all_excluded_fields aggregates fields from all registered models""" - all_fields_before = self.registry.get_all_excluded_fields() - - # Create and register a test model - from pydantic import BaseModel - from typing import Optional - - class AnotherTestModel(BaseModel): - id: str - unique_field: Optional[str] = None - - self.registry.register_model(AnotherTestModel) - - all_fields_after = self.registry.get_all_excluded_fields() - - # The new fields should be included - assert 'unique_field' in all_fields_after, "New model's excluded field should be included" - assert len(all_fields_after) >= len(all_fields_before), "Should have at least as many fields as before" - - def test_convenience_registration_method(self): - """Test the convenience method for registering models""" - from pydantic import BaseModel - from typing import Optional - - class ConvenienceTestModel(BaseModel): - id: str - convenience_field: Optional[str] = None - - # Use the convenience method - self.config.register_model_for_field_exclusion(ConvenienceTestModel) - - # Verify it was registered - excluded_fields = self.registry.get_excluded_fields_for_model(ConvenienceTestModel) - assert 'convenience_field' in excluded_fields, "Field should be excluded after registration" - - def test_field_filtering_with_registry(self): - """Test that the field filtering works correctly with the registry""" - - # Test data that matches the structure of ResponseReasoningItem - test_input = [ - { - "role": "user", - "content": "test message" - }, - { - "id": "reasoning-123", - "type": "reasoning", - "status": None, # Should be filtered out - "content": None, # Should be filtered out - "encrypted_content": None, # Should be filtered out - "summary": [{"text": "This reasoning shows...", "type": "summary_text"}], - "role": "assistant" - }, - { - "id": "message-456", - "type": "message", - "status": "completed", # Should be preserved (not None) - "content": "Hello! How can I help?", # Should be preserved (not None) - "role": "assistant" - } - ] - - # Process the input through the validation - result = self.config._validate_input_param(test_input) - - # Verify the structure - assert len(result) == 3, "Should have 3 items" - - # Check the reasoning item (index 1) - reasoning_item = result[1] - assert reasoning_item["type"] == "reasoning" - assert reasoning_item["id"] == "reasoning-123" - assert "summary" in reasoning_item, "summary field should be preserved" - assert "role" in reasoning_item, "role field should be preserved" - - # These fields should be filtered out if they are in the registry - all_excluded_fields = self.registry.get_all_excluded_fields() - if 'status' in all_excluded_fields: - assert "status" not in reasoning_item, "status field should be filtered out" - if 'content' in all_excluded_fields: - assert "content" not in reasoning_item, "content field should be filtered out" - if 'encrypted_content' in all_excluded_fields: - assert "encrypted_content" not in reasoning_item, "encrypted_content field should be filtered out" - - # Check the message item (index 2) - non-None values should be preserved - message_item = result[2] - assert message_item["type"] == "message" - assert message_item["status"] == "completed", "Non-None status should be preserved" - assert message_item["content"] == "Hello! How can I help?", "Non-None content should be preserved" - - def test_field_filtering_with_empty_registry(self): - """Test that filtering works gracefully when no models are registered""" - # Create a fresh registry for this test - from litellm.llms.openai.responses.transformation import OpenAIFieldExclusionRegistry - - # Save the current state - original_models = OpenAIFieldExclusionRegistry._MODELS_REQUIRING_EXCLUSION.copy() - - try: - # Clear the registry - OpenAIFieldExclusionRegistry._MODELS_REQUIRING_EXCLUSION.clear() - - # Test data - test_input = [{ - "id": "test-123", - "status": None, - "content": None, - "other_field": "should be preserved" - }] - - # Process the input - result = self.config._validate_input_param(test_input) - - # With empty registry, nothing should be filtered (all fields preserved) - assert len(result) == 1 - item = result[0] - assert "status" in item, "With empty registry, status should be preserved" - assert "content" in item, "With empty registry, content should be preserved" - assert item["other_field"] == "should be preserved" - - finally: - # Restore the original state - OpenAIFieldExclusionRegistry._MODELS_REQUIRING_EXCLUSION = original_models - - def test_pydantic_v1_v2_compatibility(self): - """Test that the registry works with both Pydantic v1 and v2""" - from pydantic import BaseModel - from typing import Optional - - class CompatibilityTestModel(BaseModel): - id: str - optional_field: Optional[str] = None - required_field: str = "default" - - # Register the model - self.registry.register_model(CompatibilityTestModel) - - # Get excluded fields - excluded_fields = self.registry.get_excluded_fields_for_model(CompatibilityTestModel) - - # Should work regardless of Pydantic version - assert isinstance(excluded_fields, set), "Should return a set" - assert 'optional_field' in excluded_fields, "Field with default None should be excluded" - - # Test that the model fields are accessible (works in both v1 and v2) - model_fields = getattr(CompatibilityTestModel, "model_fields", None) - if model_fields is None: - model_fields = getattr(CompatibilityTestModel, "__fields__", {}) - assert len(model_fields) > 0, "Should be able to access model fields" - - def test_non_registered_model_returns_empty_set(self): - """Test that non-registered models return empty excluded fields""" - from pydantic import BaseModel - - class UnregisteredModel(BaseModel): - id: str - some_field: str = None - - # Don't register this model - excluded_fields = self.registry.get_excluded_fields_for_model(UnregisteredModel) - - assert excluded_fields == set(), "Non-registered model should return empty set" - - @pytest.mark.parametrize("field_value", [None, "", 0, False, []]) - def test_only_none_values_are_filtered(self, field_value): - """Test that only None values are filtered, not other falsy values""" - test_input = [{ - "id": "test-123", - "status": field_value, - "content": "actual content", - "other_field": "preserved" - }] - - result = self.config._validate_input_param(test_input) - item = result[0] - - if field_value is None: - # Only None should be filtered (if status is in the registry) - all_excluded_fields = self.registry.get_all_excluded_fields() - if 'status' in all_excluded_fields: - assert "status" not in item, f"None value should be filtered out" - else: - assert item["status"] is None, f"If not in registry, None should be preserved" - else: - # Other falsy values should be preserved - assert "status" in item, f"Non-None value {field_value} should be preserved" - assert item["status"] == field_value, f"Value should be exactly {field_value}" diff --git a/tests/local_testing/test_router_tag_routing.py b/tests/test_litellm/router_strategy/test_router_tag_routing.py similarity index 81% rename from tests/local_testing/test_router_tag_routing.py rename to tests/test_litellm/router_strategy/test_router_tag_routing.py index 87cf2261a67..e78a16c6212 100644 --- a/tests/local_testing/test_router_tag_routing.py +++ b/tests/test_litellm/router_strategy/test_router_tag_routing.py @@ -63,6 +63,7 @@ async def test_router_free_paid_tier(): model="gpt-4", messages=[{"role": "user", "content": "Tell me a joke."}], metadata={"tags": ["free"]}, + mock_response="Tell me a joke.", ) print("Response: ", response) @@ -78,6 +79,7 @@ async def test_router_free_paid_tier(): model="gpt-4", messages=[{"role": "user", "content": "Tell me a joke."}], metadata={"tags": ["paid"]}, + mock_response="Tell me a joke.", ) print("Response: ", response) @@ -136,6 +138,7 @@ async def test_router_free_paid_tier_embeddings(): model="gpt-4", input="Tell me a joke.", metadata={"tags": ["free"]}, + mock_response=[1, 2, 3], ) print("Response: ", response) @@ -151,6 +154,7 @@ async def test_router_free_paid_tier_embeddings(): model="gpt-4", input="Tell me a joke.", metadata={"tags": ["paid"]}, + mock_response=[1, 2, 3], ) print("Response: ", response) @@ -205,6 +209,7 @@ async def test_default_tagged_deployments(): response = await router.acompletion( model="gpt-4", messages=[{"role": "user", "content": "Tell me a joke."}], + mock_response="Tell me a joke.", ) print("Response: ", response) @@ -220,6 +225,7 @@ async def test_default_tagged_deployments(): model="gpt-4", messages=[{"role": "user", "content": "Tell me a joke."}], metadata={"tags": ["default"]}, + mock_response="Tell me a joke.", ) print("Response: ", response) @@ -235,6 +241,7 @@ async def test_default_tagged_deployments(): model="gpt-4", messages=[{"role": "user", "content": "Tell me a joke."}], metadata={"tags": ["invalid-tag"]}, + mock_response="Tell me a joke.", ) print("Response: ", response) @@ -292,6 +299,7 @@ async def test_error_from_tag_routing(): model="gpt-4", messages=[{"role": "user", "content": "Tell me a joke."}], metadata={"tags": ["paid"]}, + mock_response="Tell me a joke.", ) pytest.fail("this should have failed - expected it to fail") @@ -315,3 +323,66 @@ def test_tag_routing_with_list_of_tags(): assert not is_valid_deployment_tag(["teamA", "teamB"], ["teamC"]) assert not is_valid_deployment_tag(["teamA", "teamB"], []) assert not is_valid_deployment_tag(["default"], ["teamA"]) + + +@pytest.mark.asyncio() +async def test_router_free_paid_tier_with_responses_api(): + """ + Pass list of orgs in 1 model definition, + expect a unique deployment for each to be created + """ + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": { + "model": "gpt-4o", + "api_base": "https://exampleopenaiendpoint-production.up.railway.app/", + "tags": ["free"], + }, + "model_info": {"id": "very-cheap-model"}, + }, + { + "model_name": "gpt-4", + "litellm_params": { + "model": "gpt-4o-mini", + "api_base": "https://exampleopenaiendpoint-production.up.railway.app/", + "tags": ["paid"], + }, + "model_info": {"id": "very-expensive-model"}, + }, + ], + enable_tag_filtering=True, + ) + + for _ in range(5): + # this should pick model with id == very-cheap-model + response = await router.aresponses( + model="gpt-4", + input="Tell me a joke.", + litellm_metadata={"tags": ["free"]}, + mock_response="Tell me a joke.", + ) + + print("Response: ", response) + + response_extra_info = response._hidden_params + print("response_extra_info: ", response_extra_info) + + assert response_extra_info["model_id"] == "very-cheap-model" + + for _ in range(5): + # this should pick model with id == very-cheap-model + response = await router.aresponses( + model="gpt-4", + input="Tell me a joke.", + litellm_metadata={"tags": ["paid"]}, + mock_response="Tell me a joke.", + ) + + print("Response: ", response) + + response_extra_info = response._hidden_params + print("response_extra_info: ", response_extra_info) + + assert response_extra_info["model_id"] == "very-expensive-model" \ No newline at end of file From 0a60390521db09006de4c8425f75e9bb50d25073 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 5 Sep 2025 10:04:07 -0700 Subject: [PATCH 26/40] Revert "[Feat] LiteLLM CloudZero Integration updates - using LiteLLM_SpendLogs Table (#12922)" This reverts commit e3b752d3dc9522e35932c77afb3139bf10603bd6. --- litellm/integrations/cloudzero/cloudzero.py | 139 ++++--------- litellm/integrations/cloudzero/database.py | 119 ++++++++--- litellm/integrations/cloudzero/transform.py | 193 ++++++------------ .../spend_tracking/cloudzero_endpoints.py | 6 +- .../integrations/cloudzero/test_transform.py | 183 +++++++++++++++++ 5 files changed, 376 insertions(+), 264 deletions(-) create mode 100644 tests/test_litellm/integrations/cloudzero/test_transform.py diff --git a/litellm/integrations/cloudzero/cloudzero.py b/litellm/integrations/cloudzero/cloudzero.py index 85aa1679732..ab1de17b9f2 100644 --- a/litellm/integrations/cloudzero/cloudzero.py +++ b/litellm/integrations/cloudzero/cloudzero.py @@ -1,6 +1,4 @@ -import asyncio import os -from datetime import datetime, timedelta from typing import Optional from litellm._logging import verbose_logger @@ -30,17 +28,16 @@ class CloudZeroLogger(CustomLogger): self.connection_id = connection_id or os.getenv("CLOUDZERO_CONNECTION_ID") self.timezone = timezone or os.getenv("CLOUDZERO_TIMEZONE", "UTC") - async def export_usage_data(self, target_hour: datetime, limit: Optional[int] = 1000, operation: str = "replace_hourly"): + async def export_usage_data(self, limit: Optional[int] = None, operation: str = "replace_hourly"): """ - Exports the usage data for a specific hour to CloudZero. + Exports the usage data to CloudZero. - - Reads spend logs from the DB for the specified hour + - Reads data from the DB - Transforms the data to the CloudZero format - Sends the data to CloudZero Args: - target_hour: The specific hour to export data for - limit: Optional limit on number of records to export (default: 1000) + limit: Optional limit on number of records to export operation: CloudZero operation type ("replace_hourly" or "sum") """ try: @@ -52,11 +49,23 @@ class CloudZeroLogger(CustomLogger): "CloudZero configuration missing. Please set CLOUDZERO_API_KEY and CLOUDZERO_CONNECTION_ID environment variables." ) - # Fetch and transform data using helper - cbf_data = await self._fetch_cbf_data_for_hour(target_hour, limit) + # Initialize database connection and load data + database = LiteLLMDatabase() + verbose_logger.debug("CloudZero Logger: Loading usage data from database") + data = await database.get_usage_data(limit=limit) + + if data.is_empty(): + verbose_logger.info("CloudZero Logger: No usage data found to export") + return + + verbose_logger.debug(f"CloudZero Logger: Processing {len(data)} records") + + # Transform data to CloudZero CBF format + transformer = CBFTransformer() + cbf_data = transformer.transform(data) if cbf_data.is_empty(): - verbose_logger.info("CloudZero Logger: No usage data found to export") + verbose_logger.warning("CloudZero Logger: No valid data after transformation") return # Send data to CloudZero @@ -75,53 +84,33 @@ class CloudZeroLogger(CustomLogger): verbose_logger.error(f"CloudZero Logger: Error exporting usage data: {str(e)}") raise - async def _fetch_cbf_data_for_hour(self, target_hour: datetime, limit: Optional[int] = 1000): + async def dry_run_export_usage_data(self, limit: Optional[int] = 10000): """ - Helper method to fetch usage data for a specific hour and transform it to CloudZero CBF format. + Only prints the data that would be exported to CloudZero. Args: - target_hour: The specific hour to fetch data for - limit: Optional limit on number of records to fetch (default: 1000) - - Returns: - CBF formatted data ready for CloudZero ingestion - """ - # Initialize database connection and load data - database = LiteLLMDatabase() - verbose_logger.debug(f"CloudZero Logger: Loading spend logs for hour {target_hour}") - data = await database.get_usage_data_for_hour(target_hour=target_hour, limit=limit) - - if data.is_empty(): - verbose_logger.info("CloudZero Logger: No usage data found for the specified hour") - return data # Return empty data - - verbose_logger.debug(f"CloudZero Logger: Processing {len(data)} records") - - # Transform data to CloudZero CBF format - transformer = CBFTransformer() - cbf_data = transformer.transform(data) - - if cbf_data.is_empty(): - verbose_logger.warning("CloudZero Logger: No valid data after transformation") - - return cbf_data - - async def dry_run_export_usage_data(self, target_hour: datetime, limit: Optional[int] = 1000): - """ - Only prints the spend logs data for a specific hour that would be exported to CloudZero. - - Args: - target_hour: The specific hour to export data for - limit: Limit number of records to display (default: 1000) + limit: Limit number of records to display (default: 10000) """ try: verbose_logger.debug("CloudZero Logger: Starting dry run export") - # Fetch and transform data using helper - cbf_data = await self._fetch_cbf_data_for_hour(target_hour, limit) + # Initialize database connection and load data + database = LiteLLMDatabase() + verbose_logger.debug("CloudZero Logger: Loading usage data for dry run") + data = await database.get_usage_data(limit=limit) + + if data.is_empty(): + verbose_logger.warning("CloudZero Dry Run: No usage data found") + return + + verbose_logger.debug(f"CloudZero Dry Run: Processing {len(data)} records...") + + # Transform data to CloudZero CBF format + transformer = CBFTransformer() + cbf_data = transformer.transform(data) if cbf_data.is_empty(): - verbose_logger.warning("CloudZero Dry Run: No usage data found") + verbose_logger.warning("CloudZero Dry Run: No valid data after transformation") return # Display the transformed data on screen @@ -198,56 +187,4 @@ class CloudZeroLogger(CustomLogger): console.print(f" Unique Accounts: {unique_accounts}") console.print(f" Unique Services: {unique_services}") - console.print("\n[dim]πŸ’‘ This is the CloudZero CBF format ready for AnyCost ingestion[/dim]") - - async def init_background_job(self, redis_cache=None): - """ - Initialize a background job that exports usage data every hour. - Uses PodLockManager to ensure only one instance runs the export at a time. - - Args: - redis_cache: Redis cache instance for pod locking - """ - from litellm.proxy.db.db_transaction_queue.pod_lock_manager import ( - PodLockManager, - ) - - lock_manager = PodLockManager(redis_cache=redis_cache) - cronjob_id = "cloudzero_hourly_export" - - async def hourly_export_task(): - while True: - try: - # Calculate the previous completed hour - now = datetime.utcnow() - target_hour = now.replace(minute=0, second=0, microsecond=0) - # Export data for the previous hour to ensure all data is available - target_hour = target_hour - timedelta(hours=1) - - # Try to acquire lock - lock_acquired = await lock_manager.acquire_lock(cronjob_id) - - if lock_acquired: - try: - verbose_logger.info(f"CloudZero Background Job: Starting export for hour {target_hour}") - await self.export_usage_data(target_hour) - verbose_logger.info(f"CloudZero Background Job: Completed export for hour {target_hour}") - finally: - # Always release the lock - await lock_manager.release_lock(cronjob_id) - else: - verbose_logger.debug("CloudZero Background Job: Another instance is already running the export") - - # Wait until the next hour - next_hour = (datetime.utcnow() + timedelta(hours=1)).replace(minute=0, second=0, microsecond=0) - sleep_seconds = (next_hour - datetime.utcnow()).total_seconds() - await asyncio.sleep(sleep_seconds) - - except Exception as e: - verbose_logger.error(f"CloudZero Background Job: Error in hourly export task: {str(e)}") - # Sleep for 5 minutes before retrying on error - await asyncio.sleep(300) - - # Start the background task - asyncio.create_task(hourly_export_task()) - verbose_logger.debug("CloudZero Background Job: Initialized hourly export task") \ No newline at end of file + console.print("\n[dim]πŸ’‘ This is the CloudZero CBF format ready for AnyCost ingestion[/dim]") \ No newline at end of file diff --git a/litellm/integrations/cloudzero/database.py b/litellm/integrations/cloudzero/database.py index 6d12c5cfbd9..73a5c28e038 100644 --- a/litellm/integrations/cloudzero/database.py +++ b/litellm/integrations/cloudzero/database.py @@ -12,14 +12,12 @@ # See the License for the specific language governing permissions and # limitations under the License. # -# CHANGELOG: 2025-07-23 - Added support for using LiteLLM_SpendLogs table for CBF mapping (ishaan-jaff) # CHANGELOG: 2025-01-19 - Refactored to use daily spend tables for proper CBF mapping (erik.peterson) # CHANGELOG: 2025-01-19 - Migrated from pandas to polars for database operations (erik.peterson) # CHANGELOG: 2025-01-19 - Initial database module for LiteLLM data extraction (erik.peterson) """Database connection and data extraction for LiteLLM.""" -from datetime import datetime, timedelta from typing import Any, Dict, Optional import polars as pl @@ -37,60 +35,123 @@ class LiteLLMDatabase: ) return prisma_client - async def get_usage_data_for_hour(self, target_hour: datetime, limit: Optional[int] = 1000) -> pl.DataFrame: - """Retrieve spend logs for a specific hour from LiteLLM_SpendLogs table with batching.""" + async def get_usage_data(self, limit: Optional[int] = None) -> pl.DataFrame: + """Retrieve consolidated usage data from LiteLLM daily spend tables.""" client = self._ensure_prisma_client() - # Calculate hour range - hour_start = target_hour.replace(minute=0, second=0, microsecond=0) - hour_end = hour_start + timedelta(hours=1) - - # Convert datetime objects to ISO format strings for PostgreSQL compatibility - hour_start_str = hour_start.isoformat() - hour_end_str = hour_end.isoformat() - - # Query to get spend logs for the specific hour + # Union query to combine user, team, and tag spend data query = """ - SELECT * - FROM "LiteLLM_SpendLogs" - WHERE "startTime" >= $1::timestamp - AND "startTime" < $2::timestamp - ORDER BY "startTime" ASC + WITH consolidated_spend AS ( + -- User spend data + SELECT + id, + date, + user_id as entity_id, + 'user' as entity_type, + api_key, + model, + model_group, + custom_llm_provider, + prompt_tokens, + completion_tokens, + spend, + api_requests, + successful_requests, + failed_requests, + cache_creation_input_tokens, + cache_read_input_tokens, + created_at, + updated_at + FROM "LiteLLM_DailyUserSpend" + + UNION ALL + + -- Team spend data + SELECT + id, + date, + team_id as entity_id, + 'team' as entity_type, + api_key, + model, + model_group, + custom_llm_provider, + prompt_tokens, + completion_tokens, + spend, + api_requests, + successful_requests, + failed_requests, + cache_creation_input_tokens, + cache_read_input_tokens, + created_at, + updated_at + FROM "LiteLLM_DailyTeamSpend" + + UNION ALL + + -- Tag spend data + SELECT + id, + date, + tag as entity_id, + 'tag' as entity_type, + api_key, + model, + model_group, + custom_llm_provider, + prompt_tokens, + completion_tokens, + spend, + api_requests, + successful_requests, + failed_requests, + cache_creation_input_tokens, + cache_read_input_tokens, + created_at, + updated_at + FROM "LiteLLM_DailyTagSpend" + ) + SELECT * FROM consolidated_spend + ORDER BY date DESC, created_at DESC """ if limit: query += f" LIMIT {limit}" try: - db_response = await client.db.query_raw(query, hour_start_str, hour_end_str) + db_response = await client.db.query_raw(query) # Convert the response to polars DataFrame - return pl.DataFrame(db_response) if db_response else pl.DataFrame() + return pl.DataFrame(db_response) except Exception as e: - raise Exception(f"Error retrieving spend logs for hour {target_hour}: {str(e)}") - + raise Exception(f"Error retrieving usage data: {str(e)}") async def get_table_info(self) -> Dict[str, Any]: - """Get information about the LiteLLM_SpendLogs table.""" + """Get information about the consolidated daily spend tables.""" client = self._ensure_prisma_client() try: - # Get row count from SpendLogs table - spend_logs_count = await self._get_table_row_count('LiteLLM_SpendLogs') + # Get combined row count from both tables + user_count = await self._get_table_row_count('LiteLLM_DailyUserSpend') + team_count = await self._get_table_row_count('LiteLLM_DailyTeamSpend') + tag_count = await self._get_table_row_count('LiteLLM_DailyTagSpend') - # Get column structure from spend logs table + # Get column structure from user spend table (representative) query = """ SELECT column_name, data_type, is_nullable FROM information_schema.columns - WHERE table_name = 'LiteLLM_SpendLogs' + WHERE table_name = 'LiteLLM_DailyUserSpend' ORDER BY ordinal_position; """ columns_response = await client.db.query_raw(query) return { 'columns': columns_response, - 'row_count': spend_logs_count, + 'row_count': user_count + team_count + tag_count, 'table_breakdown': { - 'spend_logs': spend_logs_count + 'user_spend': user_count, + 'team_spend': team_count, + 'tag_spend': tag_count } } except Exception as e: diff --git a/litellm/integrations/cloudzero/transform.py b/litellm/integrations/cloudzero/transform.py index 7091ea26b95..c8aba5dbe66 100644 --- a/litellm/integrations/cloudzero/transform.py +++ b/litellm/integrations/cloudzero/transform.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. # -# CHANGELOG: 2025-01-19 - Updated CBF transformation for LiteLLM_SpendLogs with hourly aggregation and team_id focus (ishaan-jaff) +# CHANGELOG: 2025-01-19 - Updated CBF transformation for daily spend tables and proper CloudZero mapping (erik.peterson) # CHANGELOG: 2025-01-19 - Migrated from pandas to polars for data transformation (erik.peterson) # CHANGELOG: 2025-01-19 - Initial CBF transformation module (erik.peterson) @@ -35,160 +35,90 @@ class CBFTransformer: self.czrn_generator = CZRNGenerator() def transform(self, data: pl.DataFrame) -> pl.DataFrame: - """Transform LiteLLM SpendLogs data to hourly aggregated CBF format.""" + """Transform LiteLLM data to CBF format, dropping records with zero successful_requests or invalid CZRNs.""" if data.is_empty(): return pl.DataFrame() - # Filter out records with zero spend or invalid team_id + # Filter out records with zero successful_requests first original_count = len(data) - filtered_data = data.filter( - (pl.col('spend') > 0) & - (pl.col('team_id').is_not_null()) & - (pl.col('team_id') != "") - ) - filtered_count = len(filtered_data) - zero_spend_dropped = original_count - filtered_count + if 'successful_requests' in data.columns: + filtered_data = data.filter(pl.col('successful_requests') > 0) + zero_requests_dropped = original_count - len(filtered_data) + else: + filtered_data = data + zero_requests_dropped = 0 - if filtered_data.is_empty(): - from rich.console import Console - console = Console() - console.print(f"[yellow]⚠️ Dropped all {original_count:,} records due to zero spend or missing team_id[/yellow]") - return pl.DataFrame() - - # Aggregate data to hourly level - hourly_aggregated = self._aggregate_to_hourly(filtered_data) - - # Transform aggregated data to CBF format cbf_data = [] czrn_dropped_count = 0 - - for row in hourly_aggregated.iter_rows(named=True): + filtered_count = len(filtered_data) + + for row in filtered_data.iter_rows(named=True): try: cbf_record = self._create_cbf_record(row) + # Only include the record if CZRN generation was successful cbf_data.append(cbf_record) except Exception: # Skip records that fail CZRN generation czrn_dropped_count += 1 continue - # Print summary of transformations + # Print summary of dropped records if any from rich.console import Console console = Console() - if zero_spend_dropped > 0: - console.print(f"[yellow]⚠️ Dropped {zero_spend_dropped:,} of {original_count:,} records with zero spend or missing team_id[/yellow]") + if zero_requests_dropped > 0: + console.print(f"[yellow]⚠️ Dropped {zero_requests_dropped:,} of {original_count:,} records with zero successful_requests[/yellow]") if czrn_dropped_count > 0: - console.print(f"[yellow]⚠️ Dropped {czrn_dropped_count:,} of {len(hourly_aggregated):,} aggregated records due to invalid CZRNs[/yellow]") + console.print(f"[yellow]⚠️ Dropped {czrn_dropped_count:,} of {filtered_count:,} filtered records due to invalid CZRNs[/yellow]") if len(cbf_data) > 0: - console.print(f"[green]βœ“ Successfully transformed {len(cbf_data):,} hourly aggregated records[/green]") + console.print(f"[green]βœ“ Successfully transformed {len(cbf_data):,} records[/green]") return pl.DataFrame(cbf_data) - def _aggregate_to_hourly(self, data: pl.DataFrame) -> pl.DataFrame: - """Aggregate spend logs to hourly level by team_id, key_name, model, and tags.""" - - # Extract hour from startTime, skip tags and metadata for now - data_with_hour = data.with_columns([ - pl.col('startTime').str.to_datetime().dt.truncate('1h').alias('usage_hour'), - pl.lit([]).cast(pl.List(pl.String)).alias('parsed_tags'), # Empty tags list for now - pl.lit("").alias('key_name') # Empty key name for now - ]) - - # Skip tag explosion for now - just add a null tag column - all_data = data_with_hour.with_columns([ - pl.lit(None, dtype=pl.String).alias('tag') - ]) - - # Group by hour, team_id, key_name, model, provider, and tag - aggregated = all_data.group_by([ - 'usage_hour', - 'team_id', - 'key_name', - 'model', - 'model_group', - 'custom_llm_provider', - 'tag' - ]).agg([ - pl.col('spend').sum().alias('total_spend'), - pl.col('total_tokens').sum().alias('total_tokens'), - pl.col('prompt_tokens').sum().alias('total_prompt_tokens'), - pl.col('completion_tokens').sum().alias('total_completion_tokens'), - pl.col('request_id').count().alias('request_count'), - pl.col('api_key').first().alias('api_key_sample'), # Keep one for reference - pl.col('status').filter(pl.col('status') == 'success').count().alias('successful_requests'), - pl.col('status').filter(pl.col('status') != 'success').count().alias('failed_requests') - ]) - return aggregated - - def _create_cbf_record(self, row: dict[str, Any]) -> CBFRecord: - """Create a single CBF record from aggregated hourly spend data.""" + """Create a single CBF record from LiteLLM daily spend row.""" - # Helper function to extract scalar values from polars data - def extract_scalar(value): - if hasattr(value, 'item') and not isinstance(value, (str, int, float, bool)): - return value.item() if value is not None else None - return value + # Parse date (daily spend tables use date strings like '2025-04-19') + usage_date = self._parse_date(row.get('date')) - # Use the aggregated hour as usage time - usage_time = self._parse_datetime(extract_scalar(row.get('usage_hour'))) - - # Use team_id as the primary entity_id - entity_id = str(extract_scalar(row.get('team_id', ''))) - key_name = str(extract_scalar(row.get('key_name', ''))) - model = str(extract_scalar(row.get('model', ''))) - model_group = str(extract_scalar(row.get('model_group', ''))) - provider = str(extract_scalar(row.get('custom_llm_provider', ''))) - tag = extract_scalar(row.get('tag')) - - # Calculate aggregated metrics - total_spend = float(extract_scalar(row.get('total_spend', 0.0)) or 0.0) - total_tokens = int(extract_scalar(row.get('total_tokens', 0)) or 0) - total_prompt_tokens = int(extract_scalar(row.get('total_prompt_tokens', 0)) or 0) - total_completion_tokens = int(extract_scalar(row.get('total_completion_tokens', 0)) or 0) - request_count = int(extract_scalar(row.get('request_count', 0)) or 0) - successful_requests = int(extract_scalar(row.get('successful_requests', 0)) or 0) - failed_requests = int(extract_scalar(row.get('failed_requests', 0)) or 0) + # Calculate total tokens + prompt_tokens = int(row.get('prompt_tokens', 0)) + completion_tokens = int(row.get('completion_tokens', 0)) + total_tokens = prompt_tokens + completion_tokens # Create CloudZero Resource Name (CZRN) as resource_id - # Create a mock row for CZRN generation with team_id as entity_id - czrn_row = { - 'entity_id': entity_id, - 'entity_type': 'team', - 'model': model, - 'custom_llm_provider': provider, - 'api_key': str(extract_scalar(row.get('api_key_sample', ''))) - } - resource_id = self.czrn_generator.create_from_litellm_data(czrn_row) + resource_id = self.czrn_generator.create_from_litellm_data(row) + + # Build dimensions for CloudZero + entity_id = str(row.get('entity_id', '')) + model = str(row.get('model', '')) + api_key_hash = str(row.get('api_key', ''))[:8] # First 8 chars for identification - # Build dimensions for CloudZero tracking dimensions = { - 'entity_type': 'team', + 'entity_type': str(row.get('entity_type', '')), # 'user' or 'team' 'entity_id': entity_id, - 'key_name': key_name, 'model': model, - 'model_group': model_group, - 'provider': provider, - 'request_count': str(request_count), - 'successful_requests': str(successful_requests), - 'failed_requests': str(failed_requests), + 'model_group': str(row.get('model_group', '')), + 'provider': str(row.get('custom_llm_provider', '')), + 'api_key_prefix': api_key_hash, + 'api_requests': str(row.get('api_requests', 0)), + 'successful_requests': str(row.get('successful_requests', 0)), + 'failed_requests': str(row.get('failed_requests', 0)), + 'cache_creation_tokens': str(row.get('cache_creation_input_tokens', 0)), + 'cache_read_tokens': str(row.get('cache_read_input_tokens', 0)), } - - # Add tag if present - if tag is not None and str(tag) not in ['', 'null', 'None']: - dimensions['tag'] = str(tag) # Extract CZRN components to populate corresponding CBF columns czrn_components = self.czrn_generator.extract_components(resource_id) - service_type, provider_czrn, region, owner_account_id, resource_type, cloud_local_id = czrn_components + service_type, provider, region, owner_account_id, resource_type, cloud_local_id = czrn_components # CloudZero CBF format with proper column names cbf_record = { # Required CBF fields - 'time/usage_start': usage_time.isoformat() if usage_time else None, # Required: ISO-formatted UTC datetime - 'cost/cost': total_spend, # Required: billed cost + 'time/usage_start': usage_date.isoformat() if usage_date else None, # Required: ISO-formatted UTC datetime + 'cost/cost': float(row.get('spend', 0.0)), # Required: billed cost 'resource/id': resource_id, # Required when resource tags are present # Usage metrics for token consumption @@ -206,41 +136,42 @@ class CBFTransformer: } # Add CZRN components that don't have direct CBF column mappings as resource tags - cbf_record['resource/tag:provider'] = provider_czrn # CZRN provider component + cbf_record['resource/tag:provider'] = provider # CZRN provider component cbf_record['resource/tag:model'] = cloud_local_id # CZRN cloud-local-id component (model) # Add resource tags for all dimensions (using resource/tag: format) for key, value in dimensions.items(): - # Ensure value is a scalar and not empty - if hasattr(value, 'item') and not isinstance(value, str): - value = value.item() if value is not None else None - if value is not None and str(value) not in ['', 'N/A', 'None', 'null']: # Only add non-empty tags + if value and value != 'N/A': # Only add non-empty tags cbf_record[f'resource/tag:{key}'] = str(value) # Add token breakdown as resource tags for analysis - if total_prompt_tokens > 0: - cbf_record['resource/tag:prompt_tokens'] = str(total_prompt_tokens) - if total_completion_tokens > 0: - cbf_record['resource/tag:completion_tokens'] = str(total_completion_tokens) + if prompt_tokens > 0: + cbf_record['resource/tag:prompt_tokens'] = str(prompt_tokens) + if completion_tokens > 0: + cbf_record['resource/tag:completion_tokens'] = str(completion_tokens) if total_tokens > 0: cbf_record['resource/tag:total_tokens'] = str(total_tokens) return CBFRecord(cbf_record) - def _parse_datetime(self, datetime_obj) -> Optional[datetime]: - """Parse datetime object to ensure proper format.""" - if datetime_obj is None: + def _parse_date(self, date_str) -> Optional[datetime]: + """Parse date string from daily spend tables (e.g., '2025-04-19').""" + if date_str is None: return None - if isinstance(datetime_obj, datetime): - return datetime_obj + if isinstance(date_str, datetime): + return date_str - if isinstance(datetime_obj, str): + if isinstance(date_str, str): try: - # Try to parse ISO format - return pl.Series([datetime_obj]).str.to_datetime().item() + # Parse date string and set to midnight UTC for daily aggregation + return pl.Series([date_str]).str.to_datetime("%Y-%m-%d").item() except Exception: - return None + try: + # Fallback: try ISO format parsing + return pl.Series([date_str]).str.to_datetime().item() + except Exception: + return None return None diff --git a/litellm/proxy/spend_tracking/cloudzero_endpoints.py b/litellm/proxy/spend_tracking/cloudzero_endpoints.py index 67de202aa7a..08f801c6468 100644 --- a/litellm/proxy/spend_tracking/cloudzero_endpoints.py +++ b/litellm/proxy/spend_tracking/cloudzero_endpoints.py @@ -302,7 +302,7 @@ async def init_cloudzero_background_job(): ) # Initialize the background job - await logger.init_background_job() + #await logger.init_background_job() _cloudzero_background_job_initialized = True verbose_proxy_logger.info("CloudZero background job initialized successfully") @@ -430,7 +430,7 @@ async def cloudzero_dry_run_export( try: # Import and initialize CloudZero logger with credentials - from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger + from litellm.integrations.cloudzero.ll2cz.cloudzero import CloudZeroLogger # Initialize logger with credentials directly logger = CloudZeroLogger() @@ -490,7 +490,7 @@ async def cloudzero_export( settings = await _get_cloudzero_settings() # Import and initialize CloudZero logger with credentials - from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger + from litellm.integrations.cloudzero.ll2cz.cloudzero import CloudZeroLogger # Initialize logger with credentials directly logger = CloudZeroLogger( diff --git a/tests/test_litellm/integrations/cloudzero/test_transform.py b/tests/test_litellm/integrations/cloudzero/test_transform.py new file mode 100644 index 00000000000..1f4db10cab8 --- /dev/null +++ b/tests/test_litellm/integrations/cloudzero/test_transform.py @@ -0,0 +1,183 @@ +import os +import sys +from datetime import datetime +from unittest.mock import MagicMock, patch + +import polars as pl +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +from litellm.integrations.cloudzero.transform import CBFTransformer +from litellm.types.integrations.cloudzero import CBFRecord + + +class TestCBFTransformer: + """Test suite for CBFTransformer class.""" + + def test_init(self): + """Test CBFTransformer initialization.""" + transformer = CBFTransformer() + assert hasattr(transformer, 'czrn_generator') + assert transformer.czrn_generator is not None + + def test_transform_empty_dataframe(self): + """Test transform method with empty DataFrame.""" + transformer = CBFTransformer() + empty_df = pl.DataFrame() + + result = transformer.transform(empty_df) + + assert result.is_empty() + assert isinstance(result, pl.DataFrame) + + def test_transform_with_zero_successful_requests(self): + """Test transform method filters out records with zero successful_requests.""" + transformer = CBFTransformer() + data = pl.DataFrame({ + 'date': ['2025-01-19'], + 'successful_requests': [0], + 'spend': [10.0], + 'entity_id': ['test_entity'], + 'model': ['gpt-4'] + }) + + result = transformer.transform(data) + + assert result.is_empty() + + def test_transform_with_valid_data(self): + """Test transform method with valid data.""" + transformer = CBFTransformer() + with patch.object(transformer, '_create_cbf_record') as mock_create: + mock_create.return_value = CBFRecord({'test': 'data'}) + + data = pl.DataFrame({ + 'date': ['2025-01-19'], + 'successful_requests': [5], + 'spend': [10.0], + 'entity_id': ['test_entity'], + 'model': ['gpt-4'] + }) + + result = transformer.transform(data) + + assert len(result) == 1 + mock_create.assert_called_once() + + def test_transform_handles_czrn_generation_failures(self): + """Test transform method handles CZRN generation failures gracefully.""" + transformer = CBFTransformer() + with patch.object(transformer, '_create_cbf_record') as mock_create: + mock_create.side_effect = Exception("CZRN generation failed") + + data = pl.DataFrame({ + 'date': ['2025-01-19'], + 'successful_requests': [5], + 'spend': [10.0], + 'entity_id': ['test_entity'], + 'model': ['gpt-4'] + }) + + result = transformer.transform(data) + + assert result.is_empty() + + def test_create_cbf_record(self): + """Test _create_cbf_record method with valid row data.""" + transformer = CBFTransformer() + with patch.object(transformer.czrn_generator, 'create_from_litellm_data') as mock_czrn, \ + patch.object(transformer.czrn_generator, 'extract_components') as mock_extract: + + mock_czrn.return_value = 'test-czrn' + mock_extract.return_value = ('service', 'provider', 'region', 'account', 'resource', 'local_id') + + row = { + 'date': '2025-01-19', + 'spend': 10.5, + 'prompt_tokens': 100, + 'completion_tokens': 50, + 'entity_id': 'test_entity', + 'model': 'gpt-4', + 'entity_type': 'user', + 'model_group': 'openai', + 'custom_llm_provider': 'openai', + 'api_key': 'sk-test123', + 'api_requests': 5, + 'successful_requests': 5, + 'failed_requests': 0 + } + + result = transformer._create_cbf_record(row) + + assert isinstance(result, CBFRecord) + assert result['cost/cost'] == 10.5 + assert result['usage/amount'] == 150 # 100 + 50 + assert result['usage/units'] == 'tokens' + assert result['resource/id'] == 'test-czrn' + + def test_create_cbf_record_minimal_data(self): + """Test _create_cbf_record method with minimal row data.""" + transformer = CBFTransformer() + with patch.object(transformer.czrn_generator, 'create_from_litellm_data') as mock_czrn, \ + patch.object(transformer.czrn_generator, 'extract_components') as mock_extract: + + mock_czrn.return_value = 'test-czrn' + mock_extract.return_value = ('service', 'provider', 'region', 'account', 'resource', 'local_id') + + row = { + 'date': '2025-01-19', + 'spend': 0.0 + } + + result = transformer._create_cbf_record(row) + + assert isinstance(result, CBFRecord) + assert result['cost/cost'] == 0.0 + assert result['usage/amount'] == 0 # no tokens + assert result['usage/units'] == 'tokens' + + def test_parse_date_with_valid_string(self): + """Test _parse_date method with valid date string.""" + transformer = CBFTransformer() + + result = transformer._parse_date('2025-01-19') + + assert isinstance(result, datetime) + assert result.year == 2025 + assert result.month == 1 + assert result.day == 19 + + def test_parse_date_with_datetime_object(self): + """Test _parse_date method with datetime object.""" + transformer = CBFTransformer() + dt = datetime(2025, 1, 19) + + result = transformer._parse_date(dt) + + assert result == dt + + def test_parse_date_with_none(self): + """Test _parse_date method with None.""" + transformer = CBFTransformer() + + result = transformer._parse_date(None) + + assert result is None + + def test_parse_date_with_invalid_string(self): + """Test _parse_date method with invalid date string.""" + transformer = CBFTransformer() + + result = transformer._parse_date('invalid-date') + + assert result is None + + def test_parse_date_with_iso_format(self): + """Test _parse_date method with ISO format string.""" + transformer = CBFTransformer() + + result = transformer._parse_date('2025-01-19T10:30:00Z') + + assert isinstance(result, datetime) + assert result.year == 2025 \ No newline at end of file From c051ab5b5afc719a753333b77fa1dcc498360fbc Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Fri, 5 Sep 2025 10:21:16 -0700 Subject: [PATCH 27/40] refactor: remove unused function --- litellm/llms/volcengine/__init__.py | 3 +- litellm/llms/volcengine/embedding/__init__.py | 3 +- litellm/llms/volcengine/embedding/handler.py | 208 ------------------ 3 files changed, 2 insertions(+), 212 deletions(-) delete mode 100644 litellm/llms/volcengine/embedding/handler.py diff --git a/litellm/llms/volcengine/__init__.py b/litellm/llms/volcengine/__init__.py index 0be9a4f428c..0887937bed5 100644 --- a/litellm/llms/volcengine/__init__.py +++ b/litellm/llms/volcengine/__init__.py @@ -4,12 +4,12 @@ Support for Volcengine (ByteDance) chat and embedding models """ from .chat.transformation import VolcEngineChatConfig -from .embedding import VolcEngineEmbeddingHandler, VolcEngineEmbeddingConfig from .common_utils import ( VolcEngineError, get_volcengine_base_url, get_volcengine_headers, ) +from .embedding import VolcEngineEmbeddingConfig # For backward compatibility, keep the old class name VolcEngineConfig = VolcEngineChatConfig @@ -17,7 +17,6 @@ VolcEngineConfig = VolcEngineChatConfig __all__ = [ "VolcEngineChatConfig", "VolcEngineConfig", # backward compatibility - "VolcEngineEmbeddingHandler", "VolcEngineEmbeddingConfig", "VolcEngineError", "get_volcengine_base_url", diff --git a/litellm/llms/volcengine/embedding/__init__.py b/litellm/llms/volcengine/embedding/__init__.py index 6063e88b740..7b3efc4f961 100644 --- a/litellm/llms/volcengine/embedding/__init__.py +++ b/litellm/llms/volcengine/embedding/__init__.py @@ -2,7 +2,6 @@ Volcengine Embedding Module """ -from .handler import VolcEngineEmbeddingHandler from .transformation import VolcEngineEmbeddingConfig -__all__ = ["VolcEngineEmbeddingHandler", "VolcEngineEmbeddingConfig"] +__all__ = ["VolcEngineEmbeddingConfig"] diff --git a/litellm/llms/volcengine/embedding/handler.py b/litellm/llms/volcengine/embedding/handler.py deleted file mode 100644 index 961495e72f1..00000000000 --- a/litellm/llms/volcengine/embedding/handler.py +++ /dev/null @@ -1,208 +0,0 @@ -""" -Volcengine Embedding Handler -Handles embedding requests to Volcengine's embedding API -""" - -from typing import Dict, List, Optional, Union - -import httpx -from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler -from litellm.types.utils import EmbeddingResponse -import litellm - -from .transformation import VolcEngineEmbeddingConfig -from ..common_utils import VolcEngineError - - -class VolcEngineEmbeddingHandler: - """Handler for Volcengine embedding API calls""" - - def __init__(self): - self.config = VolcEngineEmbeddingConfig() - - def _convert_to_litellm_response(self, transformed_response: Dict, model: str, input: Union[str, List[str]]) -> EmbeddingResponse: - """Convert transformed response to LiteLLM EmbeddingResponse""" - model_response = EmbeddingResponse() - model_response.object = transformed_response.get("object", "list") - model_response.data = transformed_response.get("data", []) - model_response.model = transformed_response.get("model", model) - - # Set usage information - usage_data = transformed_response.get("usage", {}) - if usage_data: - model_response.usage = litellm.Usage( - prompt_tokens=usage_data.get("prompt_tokens", 0), - completion_tokens=0, - total_tokens=usage_data.get("total_tokens", usage_data.get("prompt_tokens", 0)), - prompt_tokens_details=None, - completion_tokens_details=None, - ) - - return model_response - - def embedding( - self, - model: str, - input: Union[str, List[str]], - api_key: str, - api_base: Optional[str] = None, - encoding_format: Optional[str] = "float", - user: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - extra_headers: Optional[Dict[str, str]] = None, - litellm_logging_obj: Optional[LiteLLMLoggingObj] = None, - **kwargs, - ) -> EmbeddingResponse: - """ - Synchronous embedding call to Volcengine API. - - Args: - model: Volcengine model ID (e.g., "doubao-embedding-text-240715") - input: Text or list of texts to embed - api_key: Volcengine API key - api_base: Optional custom API base URL - encoding_format: Response format (float, base64, null) - user: Optional user identifier - timeout: Request timeout - extra_headers: Optional additional headers - litellm_logging_obj: Optional logging object - **kwargs: Additional parameters - - Returns: - EmbeddingResponse object - """ - # Transform request to Volcengine format - request_data = self.config.transform_request( - model=model, - input=input, - api_key=api_key, - api_base=api_base, - encoding_format=encoding_format, - user=user, - extra_headers=extra_headers, - **kwargs, - ) - - # Make HTTP request - try: - client = HTTPHandler(timeout=timeout) - response = client.post( - url=request_data["url"], - headers=request_data["headers"], - json=request_data["data"], - ) - except Exception as e: - raise VolcEngineError( - status_code=500, - message=f"Network error during embedding request: {str(e)}", - ) - - # Handle HTTP errors - if response.status_code != 200: - error_message = f"Volcengine embedding request failed with status {response.status_code}" - try: - error_details = response.json() - if "error" in error_details: - error_message += f": {error_details['error']}" - elif "message" in error_details: - error_message += f": {error_details['message']}" - except Exception: - error_message += f": {response.text}" - - raise VolcEngineError( - status_code=response.status_code, - message=error_message, - headers=response.headers, - ) - - # Transform response to OpenAI format - transformed_response = self.config.transform_response( - response=response, model=model, input=input, encoding=encoding_format - ) - - # Convert to LiteLLM EmbeddingResponse - return self._convert_to_litellm_response(transformed_response, model, input) - - async def async_embedding( - self, - model: str, - input: Union[str, List[str]], - api_key: str, - api_base: Optional[str] = None, - encoding_format: Optional[str] = "float", - user: Optional[str] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - extra_headers: Optional[Dict[str, str]] = None, - litellm_logging_obj: Optional[LiteLLMLoggingObj] = None, - **kwargs, - ) -> EmbeddingResponse: - """ - Asynchronous embedding call to Volcengine API. - - Args: - model: Volcengine model ID (e.g., "doubao-embedding-text-240715") - input: Text or list of texts to embed - api_key: Volcengine API key - api_base: Optional custom API base URL - encoding_format: Response format (float, base64, null) - user: Optional user identifier - timeout: Request timeout - extra_headers: Optional additional headers - litellm_logging_obj: Optional logging object - **kwargs: Additional parameters - - Returns: - EmbeddingResponse object - """ - # Transform request to Volcengine format - request_data = self.config.transform_request( - model=model, - input=input, - api_key=api_key, - api_base=api_base, - encoding_format=encoding_format, - user=user, - extra_headers=extra_headers, - **kwargs, - ) - - # Make async HTTP request - try: - client = AsyncHTTPHandler(timeout=timeout) - response = await client.post( - url=request_data["url"], - headers=request_data["headers"], - json=request_data["data"], - ) - except Exception as e: - raise VolcEngineError( - status_code=500, - message=f"Network error during embedding request: {str(e)}", - ) - - # Handle HTTP errors - if response.status_code != 200: - error_message = f"Volcengine embedding request failed with status {response.status_code}" - try: - error_details = response.json() - if "error" in error_details: - error_message += f": {error_details['error']}" - elif "message" in error_details: - error_message += f": {error_details['message']}" - except Exception: - error_message += f": {response.text}" - - raise VolcEngineError( - status_code=response.status_code, - message=error_message, - headers=response.headers, - ) - - # Transform response to OpenAI format - transformed_response = self.config.transform_response( - response=response, model=model, input=input, encoding=encoding_format - ) - - # Convert to LiteLLM EmbeddingResponse - return self._convert_to_litellm_response(transformed_response, model, input) From 31f806f7d021c25a1502c9846566b90623061cb0 Mon Sep 17 00:00:00 2001 From: Pierre-Emmanuel MERCIER <77622864+btpemercier@users.noreply.github.com> Date: Fri, 5 Sep 2025 19:35:11 +0200 Subject: [PATCH 28/40] feat: add redis ssl and username support (#11319) --- docs/my-website/docs/proxy/caching.md | 2 + litellm/_redis.py | 19 +++-- tests/test_litellm/test_redis.py | 109 ++++++++++++++++++++++++++ 3 files changed, 124 insertions(+), 6 deletions(-) create mode 100644 tests/test_litellm/test_redis.py diff --git a/docs/my-website/docs/proxy/caching.md b/docs/my-website/docs/proxy/caching.md index 1fb7385f689..49f0e199436 100644 --- a/docs/my-website/docs/proxy/caching.md +++ b/docs/my-website/docs/proxy/caching.md @@ -278,6 +278,8 @@ Set either `REDIS_URL` or the `REDIS_HOST` in your os environment, to enable cac REDIS_HOST = "" # REDIS_HOST='redis-18841.c274.us-east-1-3.ec2.cloud.redislabs.com' REDIS_PORT = "" # REDIS_PORT='18841' REDIS_PASSWORD = "" # REDIS_PASSWORD='liteLlmIsAmazing' + REDIS_USERNAME = "" # REDIS_USERNAME='my-redis-username' [OPTIONAL] if your redis server requires a username + REDIS_SSL = "True" # REDIS_SSL='True' to enable SSL by default is False ``` **Additional kwargs** diff --git a/litellm/_redis.py b/litellm/_redis.py index 8371ef5bbc7..8b64fe3dad9 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -174,14 +174,21 @@ def get_redis_url_from_environment(): raise ValueError( "Either 'REDIS_URL' or both 'REDIS_HOST' and 'REDIS_PORT' must be specified for Redis." ) - - if "REDIS_PASSWORD" in os.environ: - redis_password = f":{os.environ['REDIS_PASSWORD']}@" + + if "REDIS_SSL" in os.environ and os.environ["REDIS_SSL"].lower() == "true": + redis_protocol = "rediss" else: - redis_password = "" - + redis_protocol = "redis" + + # Build authentication part of URL + auth_part = "" + if "REDIS_USERNAME" in os.environ and "REDIS_PASSWORD" in os.environ: + auth_part = f"{os.environ['REDIS_USERNAME']}:{os.environ['REDIS_PASSWORD']}@" + elif "REDIS_PASSWORD" in os.environ: + auth_part = f"{os.environ['REDIS_PASSWORD']}@" + return ( - f"redis://{redis_password}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}" + f"{redis_protocol}://{auth_part}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}" ) diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py new file mode 100644 index 00000000000..991126c2fef --- /dev/null +++ b/tests/test_litellm/test_redis.py @@ -0,0 +1,109 @@ +from litellm._redis import get_redis_url_from_environment +import os +import pytest + +def test_get_redis_url_from_environment_single_url(monkeypatch): + """Test when REDIS_URL is directly provided""" + # Set the environment variable + monkeypatch.setenv("REDIS_URL", "redis://redis-server:6379/0") + + # Call the function to get the Redis URL + redis_url = get_redis_url_from_environment() + + # Assert that the returned URL matches the expected value + assert redis_url == "redis://redis-server:6379/0" + +def test_get_redis_url_from_environment_host_port(monkeypatch): + """Test when REDIS_HOST and REDIS_PORT are provided""" + # Set the environment variables + monkeypatch.setenv("REDIS_HOST", "redis-server") + monkeypatch.setenv("REDIS_PORT", "6379") + + # Call the function to get the Redis URL + redis_url = get_redis_url_from_environment() + + # Assert that the returned URL matches the expected value + assert redis_url == "redis://redis-server:6379" + +def test_get_redis_url_from_environment_with_ssl(monkeypatch): + """Test when SSL is enabled""" + # Set the environment variables + monkeypatch.setenv("REDIS_HOST", "redis-server") + monkeypatch.setenv("REDIS_PORT", "6379") + monkeypatch.setenv("REDIS_SSL", "true") + + # Call the function to get the Redis URL + redis_url = get_redis_url_from_environment() + + # Assert that the returned URL uses rediss:// protocol + assert redis_url == "rediss://redis-server:6379" + +def test_get_redis_url_from_environment_with_username_password(monkeypatch): + """Test when username and password are provided""" + # Set the environment variables + monkeypatch.setenv("REDIS_HOST", "redis-server") + monkeypatch.setenv("REDIS_PORT", "6379") + monkeypatch.setenv("REDIS_USERNAME", "user") + monkeypatch.setenv("REDIS_PASSWORD", "password") + + # Call the function to get the Redis URL + redis_url = get_redis_url_from_environment() + + # Assert that the returned URL includes username:password@ + assert redis_url == "redis://user:password@redis-server:6379" + +def test_get_redis_url_from_environment_with_password_only(monkeypatch): + """Test when only password is provided""" + # Set the environment variables + monkeypatch.setenv("REDIS_HOST", "redis-server") + monkeypatch.setenv("REDIS_PORT", "6379") + monkeypatch.setenv("REDIS_PASSWORD", "password") + + # Call the function to get the Redis URL + redis_url = get_redis_url_from_environment() + + # Assert that the returned URL includes :password@ + assert redis_url == "redis://password@redis-server:6379" + +def test_get_redis_url_from_environment_with_all_options(monkeypatch): + """Test when all options are provided""" + # Set the environment variables + monkeypatch.setenv("REDIS_HOST", "redis-server") + monkeypatch.setenv("REDIS_PORT", "6379") + monkeypatch.setenv("REDIS_USERNAME", "user") + monkeypatch.setenv("REDIS_PASSWORD", "password") + monkeypatch.setenv("REDIS_SSL", "true") + + # Call the function to get the Redis URL + redis_url = get_redis_url_from_environment() + + # Assert that the returned URL includes all components + assert redis_url == "rediss://user:password@redis-server:6379" + +def test_get_redis_url_from_environment_missing_host_port(monkeypatch): + """Test error when required variables are missing""" + # Make sure these environment variables don't exist + monkeypatch.delenv("REDIS_URL", raising=False) + monkeypatch.delenv("REDIS_HOST", raising=False) + monkeypatch.delenv("REDIS_PORT", raising=False) + + # Call the function and expect a ValueError + with pytest.raises(ValueError) as excinfo: + get_redis_url_from_environment() + + # Check the error message + assert "Either 'REDIS_URL' or both 'REDIS_HOST' and 'REDIS_PORT' must be specified" in str(excinfo.value) + +def test_get_redis_url_from_environment_missing_port(monkeypatch): + """Test error when only REDIS_HOST is provided but REDIS_PORT is missing""" + # Make sure REDIS_URL doesn't exist and set only REDIS_HOST + monkeypatch.delenv("REDIS_URL", raising=False) + monkeypatch.delenv("REDIS_PORT", raising=False) + monkeypatch.setenv("REDIS_HOST", "redis-server") + + # Call the function and expect a ValueError + with pytest.raises(ValueError) as excinfo: + get_redis_url_from_environment() + + # Check the error message + assert "Either 'REDIS_URL' or both 'REDIS_HOST' and 'REDIS_PORT' must be specified" in str(excinfo.value) From 07ba3ff036812bbbbc7b6df32a4f44343c0d9e45 Mon Sep 17 00:00:00 2001 From: Sameer Kankute <135028480+kankute-sameer@users.noreply.github.com> Date: Sat, 6 Sep 2025 00:55:49 +0530 Subject: [PATCH 29/40] [Feat] Add pass through image gen and image editing on OpenAI (#14292) * add pass through image gen and image editing on OpenAI * fix lint --- litellm/litellm_core_utils/litellm_logging.py | 8 + .../openai_passthrough_logging_handler.py | 260 +++++++++++++--- ...test_openai_passthrough_logging_handler.py | 286 ++++++++++++++++++ 3 files changed, 512 insertions(+), 42 deletions(-) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 7bc7702684d..397858060de 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -1165,6 +1165,14 @@ class Logging(LiteLLMLoggingBaseClass): used for consistent cost calculation across response headers + logging integrations. """ + # Check if response_cost is already calculated and stored in model_call_details + # This is used by passthrough endpoints that calculate costs manually + if ( + hasattr(self, "model_call_details") + and self.model_call_details.get("response_cost") is not None + ): + return self.model_call_details["response_cost"] + if isinstance(result, BaseModel) and hasattr(result, "_hidden_params"): hidden_params = getattr(result, "_hidden_params", {}) if ( diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py index dd772ffa502..d230023a231 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py @@ -29,7 +29,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import ( EndpointType, PassthroughStandardLoggingPayload, ) -from litellm.types.utils import LlmProviders +from litellm.types.utils import LlmProviders, PassthroughCallTypes from litellm.utils import ModelResponse, TextCompletionResponse @@ -62,6 +62,36 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): and "/v1/chat/completions" in parsed_url.path ) + @staticmethod + def is_openai_image_generation_route(url_route: str) -> bool: + """Check if the URL route is an OpenAI image generation endpoint.""" + if not url_route: + return False + parsed_url = urlparse(url_route) + return bool( + parsed_url.hostname + and ( + "api.openai.com" in parsed_url.hostname + or "openai.azure.com" in parsed_url.hostname + ) + and "/v1/images/generations" in parsed_url.path + ) + + @staticmethod + def is_openai_image_editing_route(url_route: str) -> bool: + """Check if the URL route is an OpenAI image editing endpoint.""" + if not url_route: + return False + parsed_url = urlparse(url_route) + return bool( + parsed_url.hostname + and ( + "api.openai.com" in parsed_url.hostname + or "openai.azure.com" in parsed_url.hostname + ) + and "/v1/images/edits" in parsed_url.path + ) + @staticmethod def _get_user_from_metadata( passthrough_logging_payload: PassthroughStandardLoggingPayload, @@ -73,7 +103,79 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): return None @staticmethod - def openai_passthrough_handler( + def _calculate_image_generation_cost( + model: str, + response_body: dict, + request_body: dict, + ) -> float: + """Calculate cost for OpenAI image generation.""" + try: + # Extract parameters from request + n = request_body.get("n", 1) + try: + n = int(n) + except Exception: + n = 1 + size = request_body.get("size", "1024x1024") + quality = request_body.get("quality", None) + + # Use LiteLLM's default image cost calculator + from litellm.cost_calculator import default_image_cost_calculator + + cost = default_image_cost_calculator( + model=model, + custom_llm_provider="openai", + quality=quality, + n=n, + size=size, + optional_params=request_body, + ) + + return cost + except Exception as e: + verbose_proxy_logger.warning( + f"Error calculating image generation cost: {str(e)}" + ) + return 0.0 + + @staticmethod + def _calculate_image_editing_cost( + model: str, + response_body: dict, + request_body: dict, + ) -> float: + """Calculate cost for OpenAI image editing.""" + try: + # Extract parameters from request + n = request_body.get("n", 1) + # Image edit typically uses multipart/form-data (because of files), so all fields arrive as strings (e.g., n = "1"). + try: + n = int(n) + except Exception: + n = 1 + size = request_body.get("size", "1024x1024") + + # Use LiteLLM's default image cost calculator + from litellm.cost_calculator import default_image_cost_calculator + + cost = default_image_cost_calculator( + model=model, + custom_llm_provider="openai", + quality=None, # Image editing doesn't have quality parameter + n=n, + size=size, + optional_params=request_body, + ) + + return cost + except Exception as e: + verbose_proxy_logger.warning( + f"Error calculating image editing cost: {str(e)}" + ) + return 0.0 + + @staticmethod + def openai_passthrough_handler( # noqa: PLR0915 httpx_response: httpx.Response, response_body: dict, logging_obj: LiteLLMLoggingObj, @@ -86,13 +188,21 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): **kwargs, ) -> PassThroughEndpointLoggingTypedDict: """ - Handle OpenAI passthrough logging with cost tracking for chat completions. + Handle OpenAI passthrough logging with cost tracking for chat completions, image generation, and image editing. """ - # Only handle chat completions endpoints - if not OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route( - url_route - ): - # For non-chat-completions endpoints, use the base handler without cost tracking + # Check if this is a supported endpoint for cost tracking + is_chat_completions = ( + OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route(url_route) + ) + is_image_generation = ( + OpenAIPassthroughLoggingHandler.is_openai_image_generation_route(url_route) + ) + is_image_editing = ( + OpenAIPassthroughLoggingHandler.is_openai_image_editing_route(url_route) + ) + + if not (is_chat_completions or is_image_generation or is_image_editing): + # For unsupported endpoints, use the base handler without cost tracking base_handler = OpenAIPassthroughLoggingHandler() return base_handler.passthrough_chat_handler( httpx_response=httpx_response, @@ -128,31 +238,89 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): ) try: - # Transform the response to LiteLLM format for cost calculation - provider_config = OpenAIPassthroughLoggingHandler.get_provider_config( - model=model - ) - litellm_model_response: ModelResponse = provider_config.transform_response( - raw_response=httpx_response, - model_response=litellm.ModelResponse(), - model=model, - messages=request_body.get("messages", []), - logging_obj=logging_obj, - optional_params=request_body.get("optional_params", {}), - api_key="", - request_data=request_body, - encoding=litellm.encoding, - json_mode=request_body.get("response_format", {}).get("type") - == "json_object", - litellm_params={}, - ) + response_cost = 0.0 + litellm_model_response = None - # Calculate cost using LiteLLM's cost calculator - response_cost = litellm.completion_cost( - completion_response=litellm_model_response, - model=model, - custom_llm_provider="openai", - ) + if is_chat_completions: + # Handle chat completions with existing logic + provider_config = OpenAIPassthroughLoggingHandler.get_provider_config( + model=model + ) + litellm_model_response = provider_config.transform_response( + raw_response=httpx_response, + model_response=litellm.ModelResponse(), + model=model, + messages=request_body.get("messages", []), + logging_obj=logging_obj, + optional_params=request_body.get("optional_params", {}), + api_key="", + request_data=request_body, + encoding=litellm.encoding, + json_mode=request_body.get("response_format", {}).get("type") + == "json_object", + litellm_params={}, + ) + + # Calculate cost using LiteLLM's cost calculator + response_cost = litellm.completion_cost( + completion_response=litellm_model_response, + model=model, + custom_llm_provider="openai", + ) + elif is_image_generation: + # Handle image generation cost calculation + response_cost = ( + OpenAIPassthroughLoggingHandler._calculate_image_generation_cost( + model=model, + response_body=response_body, + request_body=request_body, + ) + ) + # Mark call type for downstream image-aware logic/metrics + try: + logging_obj.call_type = ( + PassthroughCallTypes.passthrough_image_generation.value + ) + except Exception: + pass + # Create a simple response object for logging + from litellm.types.utils import ImageResponse + + litellm_model_response = ImageResponse( + data=response_body.get("data", []), + model=model, + ) + # Set the calculated cost in _hidden_params to prevent recalculation + if not hasattr(litellm_model_response, "_hidden_params"): + litellm_model_response._hidden_params = {} + litellm_model_response._hidden_params["response_cost"] = response_cost + elif is_image_editing: + # Handle image editing cost calculation + response_cost = ( + OpenAIPassthroughLoggingHandler._calculate_image_editing_cost( + model=model, + response_body=response_body, + request_body=request_body, + ) + ) + # Mark call type for downstream image-aware logic/metrics + try: + logging_obj.call_type = ( + PassthroughCallTypes.passthrough_image_generation.value + ) + except Exception: + pass + # Create a simple response object for logging + from litellm.types.utils import ImageResponse + + litellm_model_response = ImageResponse( + data=response_body.get("data", []), + model=model, + ) + # Set the calculated cost in _hidden_params to prevent recalculation + if not hasattr(litellm_model_response, "_hidden_params"): + litellm_model_response._hidden_params = {} + litellm_model_response._hidden_params["response_cost"] = response_cost # Update kwargs with cost information kwargs["response_cost"] = response_cost @@ -174,26 +342,34 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): ) # Create standard logging object - get_standard_logging_object_payload( - kwargs=kwargs, - init_response_obj=litellm_model_response, - start_time=start_time, - end_time=end_time, - logging_obj=logging_obj, - status="success", - ) + if litellm_model_response is not None: + get_standard_logging_object_payload( + kwargs=kwargs, + init_response_obj=litellm_model_response, + start_time=start_time, + end_time=end_time, + logging_obj=logging_obj, + status="success", + ) # Update logging object with cost information logging_obj.model_call_details["model"] = model logging_obj.model_call_details["custom_llm_provider"] = "openai" logging_obj.model_call_details["response_cost"] = response_cost + endpoint_type = ( + "chat_completions" + if is_chat_completions + else "image_generation" + if is_image_generation + else "image_editing" + ) verbose_proxy_logger.debug( - f"OpenAI passthrough cost tracking - Model: {model}, Cost: ${response_cost:.6f}" + f"OpenAI passthrough cost tracking - Endpoint: {endpoint_type}, Model: {model}, Cost: ${response_cost:.6f}" ) return { - "result": litellm_model_response, + "result": litellm_model_response or response_body, "kwargs": kwargs, } diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py index 6d5e80910ba..6f808c9759c 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py @@ -105,6 +105,30 @@ class TestOpenAIPassthroughLoggingHandler: assert OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route("https://api.anthropic.com/v1/messages") == False assert OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route("") == False + def test_is_openai_image_generation_route(self): + """Test OpenAI image generation route detection""" + # Positive cases + assert OpenAIPassthroughLoggingHandler.is_openai_image_generation_route("https://api.openai.com/v1/images/generations") == True + assert OpenAIPassthroughLoggingHandler.is_openai_image_generation_route("https://openai.azure.com/v1/images/generations") == True + + # Negative cases + assert OpenAIPassthroughLoggingHandler.is_openai_image_generation_route("https://api.openai.com/v1/chat/completions") == False + assert OpenAIPassthroughLoggingHandler.is_openai_image_generation_route("https://api.openai.com/v1/images/edits") == False + assert OpenAIPassthroughLoggingHandler.is_openai_image_generation_route("http://localhost:4000/openai/v1/images/generations") == False + assert OpenAIPassthroughLoggingHandler.is_openai_image_generation_route("") == False + + def test_is_openai_image_editing_route(self): + """Test OpenAI image editing route detection""" + # Positive cases + assert OpenAIPassthroughLoggingHandler.is_openai_image_editing_route("https://api.openai.com/v1/images/edits") == True + assert OpenAIPassthroughLoggingHandler.is_openai_image_editing_route("https://openai.azure.com/v1/images/edits") == True + + # Negative cases + assert OpenAIPassthroughLoggingHandler.is_openai_image_editing_route("https://api.openai.com/v1/chat/completions") == False + assert OpenAIPassthroughLoggingHandler.is_openai_image_editing_route("https://api.openai.com/v1/images/generations") == False + assert OpenAIPassthroughLoggingHandler.is_openai_image_editing_route("http://localhost:4000/openai/v1/images/edits") == False + assert OpenAIPassthroughLoggingHandler.is_openai_image_editing_route("") == False + @patch('litellm.completion_cost') @patch('litellm.litellm_core_utils.litellm_logging.get_standard_logging_object_payload') def test_openai_passthrough_handler_success(self, mock_get_standard_logging, mock_completion_cost): @@ -349,6 +373,34 @@ class TestOpenAIPassthroughIntegration: def setup_method(self): """Set up test fixtures""" self.handler = PassThroughEndpointLogging() + self.start_time = datetime.now() + self.end_time = datetime.now() + + def _create_mock_logging_obj(self) -> LiteLLMLoggingObj: + """Create a mock logging object""" + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {} + return mock_logging_obj + + def _create_mock_httpx_response(self, response_data: dict = None) -> httpx.Response: + """Create a mock httpx response""" + if response_data is None: + response_data = {"id": "test", "choices": [{"message": {"content": "Hello"}}]} + + mock_response = MagicMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.text = json.dumps(response_data) + mock_response.json.return_value = response_data + mock_response.headers = {"content-type": "application/json"} + return mock_response + + def _create_passthrough_logging_payload(self, user: str = "test_user") -> PassthroughStandardLoggingPayload: + """Create a mock passthrough logging payload""" + return PassthroughStandardLoggingPayload( + url="https://api.openai.com/v1/chat/completions", + request_body={"model": "gpt-4o", "messages": [{"role": "user", "content": "Hello"}]}, + request_method="POST", + ) def test_is_openai_route_detection(self): """Test OpenAI route detection in the main success handler""" @@ -446,6 +498,240 @@ class TestOpenAIPassthroughIntegration: # Assert - Should call the base handler, not our OpenAI handler self.handler._handle_logging.assert_called_once() + @patch('litellm.cost_calculator.default_image_cost_calculator') + def test_calculate_image_generation_cost(self, mock_image_cost_calculator): + """Test image generation cost calculation""" + # Arrange + mock_image_cost_calculator.return_value = 0.040 + model = "dall-e-3" + response_body = { + "data": [ + { + "url": "https://example.com/image1.png", + "revised_prompt": "A beautiful sunset over the ocean" + } + ] + } + request_body = { + "model": "dall-e-3", + "prompt": "A beautiful sunset over the ocean", + "n": 1, + "size": "1024x1024", + "quality": "standard" + } + + # Act + cost = OpenAIPassthroughLoggingHandler._calculate_image_generation_cost( + model=model, + response_body=response_body, + request_body=request_body, + ) + + # Assert + assert cost == 0.040 + mock_image_cost_calculator.assert_called_once_with( + model=model, + custom_llm_provider="openai", + quality="standard", + n=1, + size="1024x1024", + optional_params=request_body, + ) + + @patch('litellm.cost_calculator.default_image_cost_calculator') + def test_calculate_image_editing_cost(self, mock_image_cost_calculator): + """Test image editing cost calculation""" + # Arrange + mock_image_cost_calculator.return_value = 0.020 + model = "dall-e-2" + response_body = { + "data": [ + { + "url": "https://example.com/edited_image.png", + "revised_prompt": "A beautiful sunset over the ocean with added clouds" + } + ] + } + request_body = { + "model": "dall-e-2", + "prompt": "Add clouds to the sky", + "n": 1, + "size": "1024x1024" + } + + # Act + cost = OpenAIPassthroughLoggingHandler._calculate_image_editing_cost( + model=model, + response_body=response_body, + request_body=request_body, + ) + + # Assert + assert cost == 0.020 + mock_image_cost_calculator.assert_called_once_with( + model=model, + custom_llm_provider="openai", + quality=None, # Image editing doesn't have quality parameter + n=1, + size="1024x1024", + optional_params=request_body, + ) + + def test_cost_calculation_preservation(self): + """Test that manually calculated costs are preserved and not overridden.""" + # Create a logging object + logging_obj = LiteLLMLoggingObj( + model="dall-e-3", + messages=[{"role": "user", "content": "Generate an image"}], + stream=False, + call_type="pass_through_endpoint", + start_time=self.start_time, + litellm_call_id="test_123", + function_id="test_fn", + ) + + # Set a manually calculated cost in model_call_details + test_cost = 0.040000 + logging_obj.model_call_details["response_cost"] = test_cost + logging_obj.model_call_details["model"] = "dall-e-3" + logging_obj.model_call_details["custom_llm_provider"] = "openai" + + # Create an ImageResponse with cost in _hidden_params + from litellm.types.utils import ImageResponse + image_response = ImageResponse( + data=[{"url": "https://example.com/image.png"}], + model="dall-e-3", + ) + image_response._hidden_params = {"response_cost": test_cost} + + # Test the _response_cost_calculator method + calculated_cost = logging_obj._response_cost_calculator(result=image_response) + + assert calculated_cost == test_cost, f"Expected {test_cost}, got {calculated_cost}" + + @patch('litellm.cost_calculator.default_image_cost_calculator') + def test_openai_passthrough_handler_image_generation(self, mock_image_cost_calculator): + """Test successful cost tracking for OpenAI image generation""" + # Arrange + mock_image_cost_calculator.return_value = 0.040 + + mock_image_response = { + "data": [ + { + "url": "https://example.com/image1.png", + "revised_prompt": "A beautiful sunset over the ocean" + } + ] + } + + mock_httpx_response = self._create_mock_httpx_response(mock_image_response) + mock_logging_obj = self._create_mock_logging_obj() + passthrough_payload = self._create_passthrough_logging_payload() + + kwargs = { + "passthrough_logging_payload": passthrough_payload, + "model": "dall-e-3", + } + + request_body = { + "model": "dall-e-3", + "prompt": "A beautiful sunset over the ocean", + "n": 1, + "size": "1024x1024", + "quality": "standard" + } + + # Act + result = OpenAIPassthroughLoggingHandler.openai_passthrough_handler( + httpx_response=mock_httpx_response, + response_body=mock_image_response, + logging_obj=mock_logging_obj, + url_route="https://api.openai.com/v1/images/generations", + result="", + start_time=self.start_time, + end_time=self.end_time, + cache_hit=False, + request_body=request_body, + **kwargs + ) + + # Assert + assert result is not None + assert "result" in result + assert "kwargs" in result + assert result["kwargs"]["response_cost"] == 0.040 + assert result["kwargs"]["model"] == "dall-e-3" + assert result["kwargs"]["custom_llm_provider"] == "openai" + + # Verify cost calculation was called + mock_image_cost_calculator.assert_called_once() + + # Verify logging object was updated + assert mock_logging_obj.model_call_details["response_cost"] == 0.040 + assert mock_logging_obj.model_call_details["model"] == "dall-e-3" + assert mock_logging_obj.model_call_details["custom_llm_provider"] == "openai" + + @patch('litellm.cost_calculator.default_image_cost_calculator') + def test_openai_passthrough_handler_image_editing(self, mock_image_cost_calculator): + """Test successful cost tracking for OpenAI image editing""" + # Arrange + mock_image_cost_calculator.return_value = 0.020 + + mock_image_response = { + "data": [ + { + "url": "https://example.com/edited_image.png", + "revised_prompt": "A beautiful sunset over the ocean with added clouds" + } + ] + } + + mock_httpx_response = self._create_mock_httpx_response(mock_image_response) + mock_logging_obj = self._create_mock_logging_obj() + passthrough_payload = self._create_passthrough_logging_payload() + + kwargs = { + "passthrough_logging_payload": passthrough_payload, + "model": "dall-e-2", + } + + request_body = { + "model": "dall-e-2", + "prompt": "Add clouds to the sky", + "n": 1, + "size": "1024x1024" + } + + # Act + result = OpenAIPassthroughLoggingHandler.openai_passthrough_handler( + httpx_response=mock_httpx_response, + response_body=mock_image_response, + logging_obj=mock_logging_obj, + url_route="https://api.openai.com/v1/images/edits", + result="", + start_time=self.start_time, + end_time=self.end_time, + cache_hit=False, + request_body=request_body, + **kwargs + ) + + # Assert + assert result is not None + assert "result" in result + assert "kwargs" in result + assert result["kwargs"]["response_cost"] == 0.020 + assert result["kwargs"]["model"] == "dall-e-2" + assert result["kwargs"]["custom_llm_provider"] == "openai" + + # Verify cost calculation was called + mock_image_cost_calculator.assert_called_once() + + # Verify logging object was updated + assert mock_logging_obj.model_call_details["response_cost"] == 0.020 + assert mock_logging_obj.model_call_details["model"] == "dall-e-2" + assert mock_logging_obj.model_call_details["custom_llm_provider"] == "openai" + if __name__ == "__main__": pytest.main([__file__]) From 5310bba35bf9f7784d2f04d210210afc5d43a88d Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 5 Sep 2025 21:29:41 -0700 Subject: [PATCH 30/40] [Feat] Litellm x CloudZero Integration - Cost Tracking (#14296) * fix: just pull LiteLLM_DailyUserSpend * get the team_id from user daily spend table * cloudzero_dry_run_export * fix CZ endpoints * trace entity_id * fix: get_usage_data * fix get_usage_data * fix _create_cbf_record * fix get_usage_data * ensure start and end time is used for exporting data * fix init_cloudzero_background_job * fix CloudZeroExportRequest * fix initialize_cloudzero_export_job * fix initialize_cloudzero_export_job * allow init with env + config.yaml for cloudzero * fix: init CZ through config.yaml * fix DRY run on CZ * TestCloudZeroDryRunEndpoint * fix: CLOUDZERO_EXPORT_INTERVAL_MINUTES * fix init_cloudzero_background_job * fix exporting data * fix transform * stash cloudzero docs * docs: CloudZero * ruff fix * fix rendering key alias * fix polars --- .circleci/config.yml | 1 + .../docs/observability/cloudzero.md | 209 ++++++++++++++++++ litellm/__init__.py | 1 + litellm/constants.py | 5 + litellm/integrations/cloudzero/cloudzero.py | 176 ++++++++++++++- .../cloudzero/cz_resource_names.py | 9 +- litellm/integrations/cloudzero/database.py | 145 +++++------- litellm/integrations/cloudzero/transform.py | 21 +- litellm/litellm_core_utils/litellm_logging.py | 14 +- litellm/proxy/proxy_config.yaml | 2 + litellm/proxy/proxy_server.py | 16 +- .../spend_tracking/cloudzero_endpoints.py | 146 ++++++------ litellm/types/proxy/cloudzero_endpoints.py | 7 +- .../cloudzero/test_dry_run_endpoint.py | 163 ++++++++++++++ 14 files changed, 719 insertions(+), 196 deletions(-) create mode 100644 docs/my-website/docs/observability/cloudzero.md create mode 100644 tests/test_litellm/integrations/cloudzero/test_dry_run_endpoint.py diff --git a/.circleci/config.yml b/.circleci/config.yml index 7debc582915..2c2a2b6d6d3 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -1292,6 +1292,7 @@ jobs: pip install "tokenizers==0.20.0" pip install "uvloop==0.21.0" pip install "fastuuid==0.12.0" + pip install "polars==1.31.0" pip install jsonschema - setup_litellm_enterprise_pip - run: diff --git a/docs/my-website/docs/observability/cloudzero.md b/docs/my-website/docs/observability/cloudzero.md new file mode 100644 index 00000000000..f213ef64e13 --- /dev/null +++ b/docs/my-website/docs/observability/cloudzero.md @@ -0,0 +1,209 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# CloudZero Integration + +LiteLLM provides an integration with CloudZero's AnyCost API, allowing you to export your LLM usage data to CloudZero for cost tracking analysis. + +## Overview + +| Property | Details | +|----------|---------| +| Description | Export LiteLLM usage data to CloudZero AnyCost API for cost tracking and analysis | +| callback name | `cloudzero`| +| Supported Operations | β€’ Automatic hourly data export
β€’ Manual data export
β€’ Dry run testing
β€’ Cost and token usage tracking | +| Data Format | CloudZero Billing Format (CBF) with proper resource tagging | +| Export Frequency | Hourly (configurable via `CLOUDZERO_EXPORT_INTERVAL_MINUTES`) | + +## Environment Variables + +| Variable | Required | Description | Example | +|----------|----------|-------------|---------| +| `CLOUDZERO_API_KEY` | Yes | Your CloudZero API key | `cz_api_xxxxxxxxxx` | +| `CLOUDZERO_CONNECTION_ID` | Yes | CloudZero connection ID for data submission | `conn_xxxxxxxxxx` | +| `CLOUDZERO_TIMEZONE` | No | Timezone for date handling (default: UTC) | `America/New_York` | +| `CLOUDZERO_EXPORT_INTERVAL_MINUTES` | No | Export frequency in minutes (default: 60) | `60` | + +## Setup + +### End to End Video Walkthrough +This video walks through the entire process of setting up LiteLLM with CloudZero integration and viewing LiteLLM exported usage data in CloudZero. + + + +### Step 1: Configure Environment Variables + +Set your CloudZero credentials in your environment: + +```bash +export CLOUDZERO_API_KEY="cz_api_xxxxxxxxxx" +export CLOUDZERO_CONNECTION_ID="conn_xxxxxxxxxx" +export CLOUDZERO_TIMEZONE="UTC" # Optional, defaults to UTC +``` + +### Step 2: Enable CloudZero Integration + +Add the CloudZero callback to your LiteLLM configuration YAML file: + + +```yaml +model_list: + - model_name: gpt-4o + litellm_params: + model: openai/gpt-4o + api_key: sk-xxxxxxx + +litellm_settings: + callbacks: ["cloudzero"] # Enable CloudZero integration +``` + +### Step 3: Start LiteLLM Proxy + +Start your LiteLLM proxy with the configuration: + +```bash +litellm --config /path/to/config.yaml +``` + +## Testing Your Setup + +### Dry Run Export + +Call the dry run endpoint to test your CloudZero configuration without sending data to CloudZero. This endpoint will not send any data to CloudZero, but will return the data that would be exported. + +```bash +curl -X POST "http://localhost:4000/cloudzero/dry-run" \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer sk-1234" \ + -d '{ + "limit": 10 + }' | jq +``` + +**Expected Response:** +```json +{ + "message": "CloudZero dry run export completed successfully.", + "status": "success", + "dry_run_data": { + "usage_data": [...], + "cbf_data": [...], + "summary": { + "total_cost": 0.05, + "total_tokens": 1250, + "total_records": 10 + } + } +} +``` + +### Manual Export + +Call the export endpoint to send data immediately to CloudZero. We suggest setting a small `limit` to test the export. This will only export the last 10 records to CloudZero. Note: Cloudzero can take up to 15 minutes to process the exported data. + +```bash +curl -X POST "http://localhost:4000/cloudzero/export" \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer sk-1234" \ + -d '{ + "limit": 10 + }' | jq +``` + +**Expected Response:** +```json +{ + "message": "CloudZero export completed successfully", + "status": "success" +} +``` + +## Data Export Details + +### Automatic Export Schedule + +- **Frequency**: Every 60 minutes (configurable via `CLOUDZERO_EXPORT_INTERVAL_MINUTES`) +- **Data Processing**: LiteLLM automatically processes and exports usage data hourly +- **CloudZero Processing**: CloudZero typically takes 10-15 minutes to process data from LiteLLM + +### Data Format + +LiteLLM exports data in CloudZero Billing Format (CBF) with the following structure: + +```json +{ + "time/usage_start": "2024-01-15T14:00:00Z", + "cost/cost": 0.002, + "usage/amount": 150, + "usage/units": "tokens", + "resource/id": "czrn:litellm:openai:cross-region:team-123:llm-usage:gpt-4o", + "resource/service": "litellm", + "resource/account": "team-123", + "resource/region": "cross-region", + "resource/usage_family": "llm-usage", + "resource/tag:provider": "openai", + "resource/tag:model": "gpt-4o", + "resource/tag:prompt_tokens": "100", + "resource/tag:completion_tokens": "50" +} +``` + +### Resource Tagging + +LiteLLM automatically creates comprehensive resource tags for cost attribution: + +- **Provider Tags**: `openai`, `anthropic`, `azure`, etc. +- **Model Tags**: Specific model names like `gpt-4o`, `claude-3-sonnet` +- **Team/User Tags**: Team IDs and user IDs for cost allocation +- **Token Breakdown**: Separate tracking of prompt and completion tokens +- **Usage Metrics**: Total tokens consumed per request + +## Advanced Configuration + +### Custom Export Frequency + +Change the export frequency (not recommended to go below 60 minutes): + +```bash +export CLOUDZERO_EXPORT_INTERVAL_MINUTES=120 # Export every 2 hours +``` + +### Custom Time Range Export + +Export data for a specific time range: + +```bash +curl -X POST "http://localhost:4000/cloudzero/export" \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer sk-1234" \ + -d '{ + "start_time_utc": "2024-01-15T00:00:00Z", + "end_time_utc": "2024-01-15T23:59:59Z", + "operation": "replace_hourly" + }' | jq +``` + +## Troubleshooting + +### Common Issues + +1. **Missing Credentials Error** + ``` + CloudZero configuration missing. Please set CLOUDZERO_API_KEY and CLOUDZERO_CONNECTION_ID environment variables. + ``` + **Solution**: Ensure both environment variables are set with valid values. + +2. **Connection Issues** + - Verify your CloudZero API key is valid + - Check that the connection ID exists in your CloudZero account + - Ensure your proxy has internet access to reach CloudZero's API + +3. **No Data in CloudZero** + - CloudZero can take 10-15 minutes to process data + - Check that your LiteLLM proxy is generating usage data + - Use the dry-run endpoint to verify data is being formatted correctly + +## Related Links + +- [CloudZero Documentation](https://docs.cloudzero.com/) +- [CloudZero AnyCost API](https://docs.cloudzero.com/reference/anycost-api) diff --git a/litellm/__init__.py b/litellm/__init__.py index da5eb9d1b68..0ebea89941a 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -146,6 +146,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "aws_sqs", "vector_store_pre_call_hook", "dotprompt", + "cloudzero", ] configured_cold_storage_logger: Optional[_custom_logger_compatible_callbacks_literal] = None logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None diff --git a/litellm/constants.py b/litellm/constants.py index 21e30bef32b..089e73fc3b4 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -873,6 +873,9 @@ AZURE_STORAGE_MSFT_VERSION = "2019-07-07" PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES = int( os.getenv("PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES", 5) ) +CLOUDZERO_EXPORT_INTERVAL_MINUTES = int( + os.getenv("CLOUDZERO_EXPORT_INTERVAL_MINUTES", 60) +) MCP_TOOL_NAME_PREFIX = "mcp_tool" MAXIMUM_TRACEBACK_LINES_TO_LOG = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100)) @@ -927,6 +930,8 @@ LITELLM_CLI_SESSION_TOKEN_PREFIX = "litellm-session-token" ########################### DB CRON JOB NAMES ########################### DB_SPEND_UPDATE_JOB_NAME = "db_spend_update_job" PROMETHEUS_EMIT_BUDGET_METRICS_JOB_NAME = "prometheus_emit_budget_metrics" +CLOUDZERO_EXPORT_USAGE_DATA_JOB_NAME = "cloudzero_export_usage_data" +CLOUDZERO_MAX_FETCHED_DATA_RECORDS = int(os.getenv("CLOUDZERO_MAX_FETCHED_DATA_RECORDS", 50000)) SPEND_LOG_CLEANUP_JOB_NAME = "spend_log_cleanup" SPEND_LOG_RUN_LOOPS = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500)) SPEND_LOG_CLEANUP_BATCH_SIZE = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000)) diff --git a/litellm/integrations/cloudzero/cloudzero.py b/litellm/integrations/cloudzero/cloudzero.py index ab1de17b9f2..727dabc0945 100644 --- a/litellm/integrations/cloudzero/cloudzero.py +++ b/litellm/integrations/cloudzero/cloudzero.py @@ -1,6 +1,8 @@ import os -from typing import Optional +from datetime import datetime +from typing import TYPE_CHECKING, Any, List, Optional, cast +import litellm from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger @@ -8,6 +10,11 @@ from .cz_stream_api import CloudZeroStreamer from .database import LiteLLMDatabase from .transform import CBFTransformer +if TYPE_CHECKING: + from apscheduler.schedulers.asyncio import AsyncIOScheduler +else: + AsyncIOScheduler = Any + class CloudZeroLogger(CustomLogger): """ @@ -27,8 +34,66 @@ class CloudZeroLogger(CustomLogger): self.api_key = api_key or os.getenv("CLOUDZERO_API_KEY") self.connection_id = connection_id or os.getenv("CLOUDZERO_CONNECTION_ID") self.timezone = timezone or os.getenv("CLOUDZERO_TIMEZONE", "UTC") + verbose_logger.debug(f"CloudZero Logger initialized with connection ID: {self.connection_id}, timezone: {self.timezone}") - async def export_usage_data(self, limit: Optional[int] = None, operation: str = "replace_hourly"): + async def initialize_cloudzero_export_job(self): + """ + Handler for initializing CloudZero export job. + + Runs when CloudZero logger starts up. + + - If redis cache is available, we use the pod lock manager to acquire a lock and export the data. + - Ensures only one pod exports the data at a time. + - If redis cache is not available, we export the data directly. + """ + from litellm.constants import ( + CLOUDZERO_EXPORT_USAGE_DATA_JOB_NAME, + ) + from litellm.proxy.proxy_server import proxy_logging_obj + pod_lock_manager = proxy_logging_obj.db_spend_update_writer.pod_lock_manager + + # if using redis, ensure only one pod exports the data at a time + if pod_lock_manager and pod_lock_manager.redis_cache: + if await pod_lock_manager.acquire_lock( + cronjob_id=CLOUDZERO_EXPORT_USAGE_DATA_JOB_NAME + ): + try: + await self._hourly_usage_data_export() + finally: + await pod_lock_manager.release_lock( + cronjob_id=CLOUDZERO_EXPORT_USAGE_DATA_JOB_NAME + ) + else: + # if not using redis, export the data directly + await self._hourly_usage_data_export() + + async def _hourly_usage_data_export(self): + """ + Exports the hourly usage data to CloudZero. + + Start time: 1 hour ago + End time: current time + """ + from datetime import timedelta, timezone + + from litellm.constants import CLOUDZERO_MAX_FETCHED_DATA_RECORDS + current_time_utc = datetime.now(timezone.utc) + one_hour_ago_utc = current_time_utc - timedelta(hours=1) + await self.export_usage_data( + limit=CLOUDZERO_MAX_FETCHED_DATA_RECORDS, + operation="replace_hourly", + start_time_utc=one_hour_ago_utc, + end_time_utc=current_time_utc + ) + + + async def export_usage_data( + self, + limit: Optional[int] = None, + operation: str = "replace_hourly", + start_time_utc: Optional[datetime] = None, + end_time_utc: Optional[datetime] = None + ): """ Exports the usage data to CloudZero. @@ -52,7 +117,11 @@ class CloudZeroLogger(CustomLogger): # Initialize database connection and load data database = LiteLLMDatabase() verbose_logger.debug("CloudZero Logger: Loading usage data from database") - data = await database.get_usage_data(limit=limit) + data = await database.get_usage_data( + limit=limit, + start_time_utc=start_time_utc, + end_time_utc=end_time_utc + ) if data.is_empty(): verbose_logger.info("CloudZero Logger: No usage data found to export") @@ -86,10 +155,13 @@ class CloudZeroLogger(CustomLogger): async def dry_run_export_usage_data(self, limit: Optional[int] = 10000): """ - Only prints the data that would be exported to CloudZero. + Returns the data that would be exported to CloudZero without actually sending it. Args: limit: Limit number of records to display (default: 10000) + + Returns: + dict: Contains usage_data, cbf_data, and summary statistics """ try: verbose_logger.debug("CloudZero Logger: Starting dry run export") @@ -101,23 +173,64 @@ class CloudZeroLogger(CustomLogger): if data.is_empty(): verbose_logger.warning("CloudZero Dry Run: No usage data found") - return + return { + "usage_data": [], + "cbf_data": [], + "summary": { + "total_records": 0, + "total_cost": 0, + "total_tokens": 0, + "unique_accounts": 0, + "unique_services": 0 + } + } verbose_logger.debug(f"CloudZero Dry Run: Processing {len(data)} records...") + # Convert usage data to dict format for response + usage_data_sample = data.head(50).to_dicts() # Return first 50 rows + # Transform data to CloudZero CBF format transformer = CBFTransformer() cbf_data = transformer.transform(data) if cbf_data.is_empty(): verbose_logger.warning("CloudZero Dry Run: No valid data after transformation") - return + return { + "usage_data": usage_data_sample, + "cbf_data": [], + "summary": { + "total_records": len(usage_data_sample), + "total_cost": sum(row.get('spend', 0) for row in usage_data_sample), + "total_tokens": sum(row.get('prompt_tokens', 0) + row.get('completion_tokens', 0) for row in usage_data_sample), + "unique_accounts": 0, + "unique_services": 0 + } + } - # Display the transformed data on screen - self._display_cbf_data_on_screen(cbf_data) + # Convert CBF data to dict format for response + cbf_data_dict = cbf_data.to_dicts() + + # Calculate summary statistics + total_cost = sum(record.get('cost/cost', 0) for record in cbf_data_dict) + unique_accounts = len(set(record.get('resource/account', '') for record in cbf_data_dict if record.get('resource/account'))) + unique_services = len(set(record.get('resource/service', '') for record in cbf_data_dict if record.get('resource/service'))) + total_tokens = sum(record.get('usage/amount', 0) for record in cbf_data_dict) verbose_logger.info(f"CloudZero Logger: Dry run completed for {len(cbf_data)} records") + return { + "usage_data": usage_data_sample, + "cbf_data": cbf_data_dict, + "summary": { + "total_records": len(cbf_data_dict), + "total_cost": total_cost, + "total_tokens": total_tokens, + "unique_accounts": unique_accounts, + "unique_services": unique_services + } + } + except Exception as e: verbose_logger.error(f"CloudZero Logger: Error in dry run export: {str(e)}") verbose_logger.error(f"CloudZero Dry Run Error: {str(e)}") @@ -144,6 +257,11 @@ class CloudZeroLogger(CustomLogger): cbf_table = Table(show_header=True, header_style="bold cyan", box=SIMPLE, padding=(0, 1)) cbf_table.add_column("time/usage_start", style="blue", no_wrap=False) cbf_table.add_column("cost/cost", style="green", justify="right", no_wrap=False) + cbf_table.add_column("entity_type", style="magenta", justify="right", no_wrap=False) + cbf_table.add_column("entity_id", style="magenta", justify="right", no_wrap=False) + cbf_table.add_column("team_id", style="cyan", no_wrap=False) + cbf_table.add_column("team_alias", style="cyan", no_wrap=False) + cbf_table.add_column("api_key_alias", style="yellow", no_wrap=False) cbf_table.add_column("usage/amount", style="yellow", justify="right", no_wrap=False) cbf_table.add_column("resource/id", style="magenta", no_wrap=False) cbf_table.add_column("resource/service", style="cyan", no_wrap=False) @@ -159,10 +277,20 @@ class CloudZeroLogger(CustomLogger): resource_service = str(record.get('resource/service', 'N/A')) resource_account = str(record.get('resource/account', 'N/A')) resource_region = str(record.get('resource/region', 'N/A')) + entity_type = str(record.get('entity_type', 'N/A')) + entity_id = str(record.get('entity_id', 'N/A')) + team_id = str(record.get('resource/tag:team_id', 'N/A')) + team_alias = str(record.get('resource/tag:team_alias', 'N/A')) + api_key_alias = str(record.get('resource/tag:api_key_alias', 'N/A')) cbf_table.add_row( time_usage_start, cost_cost, + entity_type, + entity_id, + team_id, + team_alias, + api_key_alias, usage_amount, resource_id, resource_service, @@ -187,4 +315,34 @@ class CloudZeroLogger(CustomLogger): console.print(f" Unique Accounts: {unique_accounts}") console.print(f" Unique Services: {unique_services}") - console.print("\n[dim]πŸ’‘ This is the CloudZero CBF format ready for AnyCost ingestion[/dim]") \ No newline at end of file + console.print("\n[dim]πŸ’‘ This is the CloudZero CBF format ready for AnyCost ingestion[/dim]") + + @staticmethod + async def init_cloudzero_background_job(scheduler: AsyncIOScheduler): + """ + Initialize the CloudZero background job. + + Starts the background job that exports the usage data to CloudZero every hour. + """ + from litellm.constants import CLOUDZERO_EXPORT_INTERVAL_MINUTES + from litellm.integrations.custom_logger import CustomLogger + + + prometheus_loggers: List[CustomLogger] = ( + litellm.logging_callback_manager.get_custom_loggers_for_type( + callback_type=CloudZeroLogger + ) + ) + # we need to get the initialized prometheus logger instance(s) and call logger.initialize_remaining_budget_metrics() on them + verbose_logger.debug("found %s cloudzero loggers", len(prometheus_loggers)) + if len(prometheus_loggers) > 0: + cloudzero_logger = cast(CloudZeroLogger, prometheus_loggers[0]) + verbose_logger.debug( + "Initializing remaining budget metrics as a cron job executing every %s minutes" + % CLOUDZERO_EXPORT_INTERVAL_MINUTES + ) + scheduler.add_job( + cloudzero_logger.initialize_cloudzero_export_job, + "interval", + minutes=CLOUDZERO_EXPORT_INTERVAL_MINUTES + ) \ No newline at end of file diff --git a/litellm/integrations/cloudzero/cz_resource_names.py b/litellm/integrations/cloudzero/cz_resource_names.py index 44147f9c210..f1098d20381 100644 --- a/litellm/integrations/cloudzero/cz_resource_names.py +++ b/litellm/integrations/cloudzero/cz_resource_names.py @@ -17,11 +17,16 @@ """CloudZero Resource Names (CZRN) generation and validation for LiteLLM resources.""" import re +from enum import Enum from typing import Any, cast import litellm +class CZEntityType(str, Enum): + TEAM = "team" + + class CZRNGenerator: """Generate CloudZero Resource Names (CZRNs) for LiteLLM resources.""" @@ -49,8 +54,8 @@ class CZRNGenerator: region = 'cross-region' # Use the actual entity_id (team_id or user_id) as the owner account - entity_id = row.get('entity_id', 'unknown') - owner_account_id = self._normalize_component(entity_id) + team_id = row.get('team_id', 'unknown') + owner_account_id = self._normalize_component(team_id) resource_type = 'llm-usage' diff --git a/litellm/integrations/cloudzero/database.py b/litellm/integrations/cloudzero/database.py index 73a5c28e038..71b4125ed75 100644 --- a/litellm/integrations/cloudzero/database.py +++ b/litellm/integrations/cloudzero/database.py @@ -18,6 +18,7 @@ """Database connection and data extraction for LiteLLM.""" +from datetime import datetime from typing import Any, Dict, Optional import polars as pl @@ -35,85 +36,54 @@ class LiteLLMDatabase: ) return prisma_client - async def get_usage_data(self, limit: Optional[int] = None) -> pl.DataFrame: - """Retrieve consolidated usage data from LiteLLM daily spend tables.""" + async def get_usage_data( + self, + limit: Optional[int] = None, + start_time_utc: Optional[datetime] = None, + end_time_utc: Optional[datetime] = None + ) -> pl.DataFrame: + """Retrieve usage data from LiteLLM daily user spend table.""" client = self._ensure_prisma_client() - # Union query to combine user, team, and tag spend data - query = """ - WITH consolidated_spend AS ( - -- User spend data - SELECT - id, - date, - user_id as entity_id, - 'user' as entity_type, - api_key, - model, - model_group, - custom_llm_provider, - prompt_tokens, - completion_tokens, - spend, - api_requests, - successful_requests, - failed_requests, - cache_creation_input_tokens, - cache_read_input_tokens, - created_at, - updated_at - FROM "LiteLLM_DailyUserSpend" - - UNION ALL - - -- Team spend data - SELECT - id, - date, - team_id as entity_id, - 'team' as entity_type, - api_key, - model, - model_group, - custom_llm_provider, - prompt_tokens, - completion_tokens, - spend, - api_requests, - successful_requests, - failed_requests, - cache_creation_input_tokens, - cache_read_input_tokens, - created_at, - updated_at - FROM "LiteLLM_DailyTeamSpend" - - UNION ALL - - -- Tag spend data - SELECT - id, - date, - tag as entity_id, - 'tag' as entity_type, - api_key, - model, - model_group, - custom_llm_provider, - prompt_tokens, - completion_tokens, - spend, - api_requests, - successful_requests, - failed_requests, - cache_creation_input_tokens, - cache_read_input_tokens, - created_at, - updated_at - FROM "LiteLLM_DailyTagSpend" - ) - SELECT * FROM consolidated_spend - ORDER BY date DESC, created_at DESC + # Build WHERE clause for time filtering + where_conditions = [] + if start_time_utc: + where_conditions.append(f"dus.created_at >= '{start_time_utc.isoformat()}'") + if end_time_utc: + where_conditions.append(f"dus.created_at <= '{end_time_utc.isoformat()}'") + + where_clause = "" + if where_conditions: + where_clause = "WHERE " + " AND ".join(where_conditions) + + # Query to get user spend data with team information + query = f""" + SELECT + dus.id, + dus.date, + dus.user_id, + dus.api_key, + dus.model, + dus.model_group, + dus.custom_llm_provider, + dus.prompt_tokens, + dus.completion_tokens, + dus.spend, + dus.api_requests, + dus.successful_requests, + dus.failed_requests, + dus.cache_creation_input_tokens, + dus.cache_read_input_tokens, + dus.created_at, + dus.updated_at, + vt.team_id, + vt.key_alias as api_key_alias, + tt.team_alias + FROM "LiteLLM_DailyUserSpend" dus + LEFT JOIN "LiteLLM_VerificationToken" vt ON dus.api_key = vt.token + LEFT JOIN "LiteLLM_TeamTable" tt ON vt.team_id = tt.team_id + {where_clause} + ORDER BY dus.date DESC, dus.created_at DESC """ if limit: @@ -121,22 +91,21 @@ class LiteLLMDatabase: try: db_response = await client.db.query_raw(query) - # Convert the response to polars DataFrame - return pl.DataFrame(db_response) + # Convert the response to polars DataFrame with full schema inference + # This prevents schema mismatch errors when data types vary across rows + return pl.DataFrame(db_response, infer_schema_length=None) except Exception as e: raise Exception(f"Error retrieving usage data: {str(e)}") async def get_table_info(self) -> Dict[str, Any]: - """Get information about the consolidated daily spend tables.""" + """Get information about the daily user spend table.""" client = self._ensure_prisma_client() try: - # Get combined row count from both tables + # Get row count from user spend table user_count = await self._get_table_row_count('LiteLLM_DailyUserSpend') - team_count = await self._get_table_row_count('LiteLLM_DailyTeamSpend') - tag_count = await self._get_table_row_count('LiteLLM_DailyTagSpend') - # Get column structure from user spend table (representative) + # Get column structure from user spend table query = """ SELECT column_name, data_type, is_nullable FROM information_schema.columns @@ -147,12 +116,8 @@ class LiteLLMDatabase: return { 'columns': columns_response, - 'row_count': user_count + team_count + tag_count, - 'table_breakdown': { - 'user_spend': user_count, - 'team_spend': team_count, - 'tag_spend': tag_count - } + 'row_count': user_count, + 'table_name': 'LiteLLM_DailyUserSpend' } except Exception as e: raise Exception(f"Error getting table info: {str(e)}") diff --git a/litellm/integrations/cloudzero/transform.py b/litellm/integrations/cloudzero/transform.py index c8aba5dbe66..e0263295388 100644 --- a/litellm/integrations/cloudzero/transform.py +++ b/litellm/integrations/cloudzero/transform.py @@ -24,7 +24,7 @@ from typing import Any, Optional import polars as pl from ...types.integrations.cloudzero import CBFRecord -from .cz_resource_names import CZRNGenerator +from .cz_resource_names import CZEntityType, CZRNGenerator class CBFTransformer: @@ -92,17 +92,26 @@ class CBFTransformer: resource_id = self.czrn_generator.create_from_litellm_data(row) # Build dimensions for CloudZero - entity_id = str(row.get('entity_id', '')) model = str(row.get('model', '')) api_key_hash = str(row.get('api_key', ''))[:8] # First 8 chars for identification - + + # Handle team information with fallbacks + team_id = row.get('team_id') + team_alias = row.get('team_alias') + + # Use team_alias if available, otherwise team_id, otherwise fallback to 'unknown' + entity_id = str(team_alias) if team_alias else (str(team_id) if team_id else 'unknown') + dimensions = { - 'entity_type': str(row.get('entity_type', '')), # 'user' or 'team' + 'entity_type': CZEntityType.TEAM.value, 'entity_id': entity_id, + 'team_id': str(team_id) if team_id else 'unknown', + 'team_alias': str(team_alias) if team_alias else 'unknown', 'model': model, 'model_group': str(row.get('model_group', '')), 'provider': str(row.get('custom_llm_provider', '')), 'api_key_prefix': api_key_hash, + 'api_key_alias': str(row.get('api_key_alias', '')), 'api_requests': str(row.get('api_requests', 0)), 'successful_requests': str(row.get('successful_requests', 0)), 'failed_requests': str(row.get('failed_requests', 0)), @@ -138,10 +147,10 @@ class CBFTransformer: # Add CZRN components that don't have direct CBF column mappings as resource tags cbf_record['resource/tag:provider'] = provider # CZRN provider component cbf_record['resource/tag:model'] = cloud_local_id # CZRN cloud-local-id component (model) - + # Add resource tags for all dimensions (using resource/tag: format) for key, value in dimensions.items(): - if value and value != 'N/A': # Only add non-empty tags + if value and value != 'N/A' and value != 'unknown': # Only add meaningful tags cbf_record[f'resource/tag:{key}'] = str(value) # Add token breakdown as resource tags for analysis diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 397858060de..7134f52c95a 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -3369,7 +3369,14 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 galileo_logger = GalileoObserve() _in_memory_loggers.append(galileo_logger) return galileo_logger # type: ignore - + elif logging_integration == "cloudzero": + from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger + for callback in _in_memory_loggers: + if isinstance(callback, CloudZeroLogger): + return callback # type: ignore + cloudzero_logger = CloudZeroLogger() + _in_memory_loggers.append(cloudzero_logger) + return cloudzero_logger # type: ignore elif logging_integration == "deepeval": for callback in _in_memory_loggers: if isinstance(callback, DeepEvalLogger): @@ -3589,6 +3596,11 @@ def get_custom_logger_compatible_class( # noqa: PLR0915 for callback in _in_memory_loggers: if isinstance(callback, GalileoObserve): return callback + elif logging_integration == "cloudzero": + from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger + for callback in _in_memory_loggers: + if isinstance(callback, CloudZeroLogger): + return callback elif logging_integration == "deepeval": for callback in _in_memory_loggers: if isinstance(callback, DeepEvalLogger): diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 72c69a28e95..7ee09105254 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -3,3 +3,5 @@ model_list: litellm_params: model: openai/* api_base: https://exampleopenaiendpoint-production-0ee2.up.railway.app/ +litellm_settings: + callbacks: ["cloudzero"] \ No newline at end of file diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 547aaf50788..9f1566b2e00 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -248,7 +248,9 @@ from litellm.proxy.management_endpoints.customer_endpoints import ( from litellm.proxy.management_endpoints.internal_user_endpoints import ( router as internal_user_router, ) -from litellm.proxy.management_endpoints.internal_user_endpoints import user_update +from litellm.proxy.management_endpoints.internal_user_endpoints import ( + user_update, +) from litellm.proxy.management_endpoints.key_management_endpoints import ( delete_verification_tokens, duration_in_seconds, @@ -295,7 +297,9 @@ from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMi from litellm.proxy.openai_files_endpoints.files_endpoints import ( router as openai_files_router, ) -from litellm.proxy.openai_files_endpoints.files_endpoints import set_files_config +from litellm.proxy.openai_files_endpoints.files_endpoints import ( + set_files_config, +) from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( passthrough_endpoint_router, ) @@ -3807,13 +3811,13 @@ class ProxyStartupEvent: ######################################################## # CloudZero Background Job ######################################################## + from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger from litellm.proxy.spend_tracking.cloudzero_endpoints import ( - init_cloudzero_background_job, - is_cloudzero_setup_in_db, + is_cloudzero_setup, ) - if await is_cloudzero_setup_in_db(): - await init_cloudzero_background_job() + if await is_cloudzero_setup(): + await CloudZeroLogger.init_cloudzero_background_job(scheduler=scheduler) ######################################################## # Prometheus Background Job diff --git a/litellm/proxy/spend_tracking/cloudzero_endpoints.py b/litellm/proxy/spend_tracking/cloudzero_endpoints.py index 08f801c6468..502537cb70f 100644 --- a/litellm/proxy/spend_tracking/cloudzero_endpoints.py +++ b/litellm/proxy/spend_tracking/cloudzero_endpoints.py @@ -82,14 +82,8 @@ async def _get_cloudzero_settings(): cloudzero_config = await prisma_client.db.litellm_config.find_first( where={"param_name": "cloudzero_settings"} ) - - if not cloudzero_config or not cloudzero_config.param_value: - raise HTTPException( - status_code=400, - detail={ - "error": "CloudZero settings not configured. Please run /cloudzero/init first." - }, - ) + if cloudzero_config is None: + return {} settings = dict(cloudzero_config.param_value) @@ -257,62 +251,6 @@ async def update_cloudzero_settings( _cloudzero_background_job_initialized = False -async def init_cloudzero_background_job(): - """ - Initialize CloudZero background job if not already initialized. - This should be called from the proxy server startup. - """ - global _cloudzero_background_job_initialized - - if _cloudzero_background_job_initialized: - verbose_proxy_logger.debug( - "CloudZero background job already initialized, skipping" - ) - return - - try: - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - verbose_proxy_logger.warning( - "Prisma client not available, skipping CloudZero background job initialization" - ) - return - - # Get CloudZero settings from database - cloudzero_config = await prisma_client.db.litellm_config.find_first( - where={"param_name": "cloudzero_settings"} - ) - - if not cloudzero_config or not cloudzero_config.param_value: - verbose_proxy_logger.debug( - "CloudZero settings not configured, skipping background job initialization" - ) - return - - settings = dict(cloudzero_config.param_value) - - # Initialize CloudZero logger with credentials - from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger - - logger = CloudZeroLogger( - api_key=settings["api_key"], - connection_id=settings["connection_id"], - timezone=settings["timezone"], - ) - - # Initialize the background job - #await logger.init_background_job() - - _cloudzero_background_job_initialized = True - verbose_proxy_logger.info("CloudZero background job initialized successfully") - - except Exception as e: - verbose_proxy_logger.error( - f"Error initializing CloudZero background job: {str(e)}" - ) - - async def is_cloudzero_setup_in_db() -> bool: """ Check if CloudZero is setup in the database. @@ -343,6 +281,47 @@ async def is_cloudzero_setup_in_db() -> bool: return False +def is_cloudzero_setup_in_config() -> bool: + """ + Check if CloudZero is setup in config.yaml or environment variables. + + CloudZero is considered setup in config if: + - "cloudzero" is in the callbacks list in config.yaml, OR + Returns: + bool: True if CloudZero is configured, False otherwise + """ + import litellm + return "cloudzero" in litellm.callbacks + + +async def is_cloudzero_setup() -> bool: + """ + Check if CloudZero is setup in either config.yaml/env vars OR database. + + CloudZero is considered setup if: + - CloudZero is configured in config.yaml callbacks, OR + - CloudZero environment variables are set, OR + - CloudZero settings exist in the database + + Returns: + bool: True if CloudZero is configured anywhere, False otherwise + """ + try: + # Check config.yaml/environment variables first + if is_cloudzero_setup_in_config(): + return True + + # Check database as fallback + if await is_cloudzero_setup_in_db(): + return True + + return False + + except Exception as e: + verbose_proxy_logger.error(f"Error checking CloudZero setup: {str(e)}") + return False + + @router.post( "/cloudzero/init", tags=["CloudZero"], @@ -383,9 +362,6 @@ async def init_cloudzero_settings( verbose_proxy_logger.info("CloudZero settings initialized successfully") - # Initialize background job after settings are saved - await init_cloudzero_background_job() - return CloudZeroInitResponse( message="CloudZero settings initialized successfully", status="success" ) @@ -412,15 +388,18 @@ async def cloudzero_dry_run_export( Perform a dry run export using the CloudZero logger. This endpoint uses the CloudZero logger to perform a dry run export, - which displays the data that would be exported without actually sending it to CloudZero. + which returns the data that would be exported without actually sending it to CloudZero. Parameters: - limit: Optional limit on number of records to process (default: 10000) + Returns: + - usage_data: Sample of the raw usage data (first 50 records) + - cbf_data: CloudZero CBF formatted data ready for export + - summary: Statistics including total cost, tokens, and record counts + Only admin users can perform CloudZero exports. """ - from datetime import datetime - # Validation if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: raise HTTPException( @@ -430,19 +409,21 @@ async def cloudzero_dry_run_export( try: # Import and initialize CloudZero logger with credentials - from litellm.integrations.cloudzero.ll2cz.cloudzero import CloudZeroLogger + from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger # Initialize logger with credentials directly logger = CloudZeroLogger() - await logger.dry_run_export_usage_data( - target_hour=datetime.utcnow(), limit=request.limit + dry_run_result = await logger.dry_run_export_usage_data( + limit=request.limit ) verbose_proxy_logger.info("CloudZero dry run export completed successfully") return CloudZeroExportResponse( - message="CloudZero dry run export completed successfully. Check logs for output.", + message="CloudZero dry run export completed successfully.", status="success", + dry_run_data=dry_run_result, + summary=dry_run_result.get("summary") if dry_run_result else None, ) except Exception as e: @@ -477,7 +458,6 @@ async def cloudzero_export( Only admin users can perform CloudZero exports. """ - from datetime import datetime if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: raise HTTPException( @@ -490,24 +470,28 @@ async def cloudzero_export( settings = await _get_cloudzero_settings() # Import and initialize CloudZero logger with credentials - from litellm.integrations.cloudzero.ll2cz.cloudzero import CloudZeroLogger + from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger # Initialize logger with credentials directly logger = CloudZeroLogger( - api_key=settings["api_key"], - connection_id=settings["connection_id"], - timezone=settings["timezone"], + api_key=settings.get("api_key"), + connection_id=settings.get("connection_id"), + timezone=settings.get("timezone"), ) await logger.export_usage_data( - target_hour=datetime.utcnow(), limit=request.limit, operation=request.operation, + start_time_utc=request.start_time_utc, + end_time_utc=request.end_time_utc, ) verbose_proxy_logger.info("CloudZero export completed successfully") return CloudZeroExportResponse( - message="CloudZero export completed successfully", status="success" + message="CloudZero export completed successfully", + status="success", + dry_run_data=None, + summary=None ) except Exception as e: diff --git a/litellm/types/proxy/cloudzero_endpoints.py b/litellm/types/proxy/cloudzero_endpoints.py index f7f63233d4d..1d909bf7f8c 100644 --- a/litellm/types/proxy/cloudzero_endpoints.py +++ b/litellm/types/proxy/cloudzero_endpoints.py @@ -2,7 +2,8 @@ CloudZero endpoint types for LiteLLM Proxy """ -from typing import Optional +from datetime import datetime +from typing import Any, Dict, List, Optional from pydantic import BaseModel, Field @@ -27,6 +28,8 @@ class CloudZeroExportRequest(BaseModel): limit: Optional[int] = Field(None, description="Optional limit on number of records to export") operation: str = Field(default="replace_hourly", description="CloudZero operation type (replace_hourly or sum)") + start_time_utc: Optional[datetime] = Field(None, description="Start time for data export in UTC") + end_time_utc: Optional[datetime] = Field(None, description="End time for data export in UTC") class CloudZeroExportResponse(BaseModel): @@ -35,6 +38,8 @@ class CloudZeroExportResponse(BaseModel): message: str status: str records_exported: Optional[int] = None + dry_run_data: Optional[Dict[str, Any]] = Field(None, description="Dry run data including usage data and CBF transformed data") + summary: Optional[Dict[str, Any]] = Field(None, description="Summary statistics for dry run") class CloudZeroSettingsView(BaseModel): diff --git a/tests/test_litellm/integrations/cloudzero/test_dry_run_endpoint.py b/tests/test_litellm/integrations/cloudzero/test_dry_run_endpoint.py new file mode 100644 index 00000000000..9a31a140aa8 --- /dev/null +++ b/tests/test_litellm/integrations/cloudzero/test_dry_run_endpoint.py @@ -0,0 +1,163 @@ +""" +Test the CloudZero dry run endpoint functionality +""" +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import polars as pl +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger + + +class TestCloudZeroDryRunEndpoint: + """Test suite for CloudZero dry run endpoint functionality.""" + + @pytest.mark.asyncio + async def test_dry_run_export_usage_data_returns_data(self): + """ + Test that dry_run_export_usage_data returns expected data structure + instead of just logging to console. + """ + logger = CloudZeroLogger() + + # Mock database data + mock_usage_data = pl.DataFrame({ + 'date': ['2025-01-19', '2025-01-20'], + 'model': ['gpt-4', 'gpt-3.5-turbo'], + 'custom_llm_provider': ['openai', 'openai'], + 'team_id': ['team1', 'team2'], + 'team_alias': ['Team One', 'Team Two'], + 'api_key_alias': ['key1', 'key2'], + 'prompt_tokens': [100, 200], + 'completion_tokens': [50, 100], + 'spend': [0.01, 0.02], + 'successful_requests': [1, 2] + }) + + # Mock CBF transformed data + mock_cbf_data = pl.DataFrame({ + 'time/usage_start': ['2025-01-19T00:00:00Z', '2025-01-20T00:00:00Z'], + 'cost/cost': [0.01, 0.02], + 'usage/amount': [150, 300], + 'resource/service': ['openai', 'openai'], + 'resource/account': ['litellm', 'litellm'], + 'resource/region': ['us-east-1', 'us-east-1'], + 'resource/id': ['gpt-4', 'gpt-3.5-turbo'], + 'entity_type': ['user', 'user'], + 'entity_id': ['team1', 'team2'], + 'resource/tag:team_id': ['team1', 'team2'], + 'resource/tag:team_alias': ['Team One', 'Team Two'], + 'resource/tag:api_key_alias': ['key1', 'key2'] + }) + + with patch('litellm.integrations.cloudzero.cloudzero.LiteLLMDatabase') as mock_db_class, \ + patch('litellm.integrations.cloudzero.cloudzero.CBFTransformer') as mock_transformer_class: + + # Setup mocks + mock_db = AsyncMock() + mock_db.get_usage_data.return_value = mock_usage_data + mock_db_class.return_value = mock_db + + mock_transformer = MagicMock() + mock_transformer.transform.return_value = mock_cbf_data + mock_transformer_class.return_value = mock_transformer + + # Call the method + result = await logger.dry_run_export_usage_data(limit=1000) + + # Verify the result structure + assert isinstance(result, dict) + assert 'usage_data' in result + assert 'cbf_data' in result + assert 'summary' in result + + # Verify usage_data + assert isinstance(result['usage_data'], list) + assert len(result['usage_data']) == 2 + assert result['usage_data'][0]['model'] == 'gpt-4' + assert result['usage_data'][1]['model'] == 'gpt-3.5-turbo' + + # Verify cbf_data + assert isinstance(result['cbf_data'], list) + assert len(result['cbf_data']) == 2 + assert result['cbf_data'][0]['cost/cost'] == 0.01 + assert result['cbf_data'][1]['cost/cost'] == 0.02 + + # Verify summary + summary = result['summary'] + assert summary['total_records'] == 2 + assert summary['total_cost'] == 0.03 + assert summary['total_tokens'] == 450 # 150 + 300 + assert summary['unique_accounts'] == 1 + assert summary['unique_services'] == 1 + + @pytest.mark.asyncio + async def test_dry_run_export_usage_data_empty_data(self): + """ + Test that dry_run_export_usage_data handles empty data gracefully. + """ + logger = CloudZeroLogger() + + # Mock empty database data + mock_empty_data = pl.DataFrame() + + with patch('litellm.integrations.cloudzero.cloudzero.LiteLLMDatabase') as mock_db_class: + + # Setup mocks + mock_db = AsyncMock() + mock_db.get_usage_data.return_value = mock_empty_data + mock_db_class.return_value = mock_db + + # Call the method + result = await logger.dry_run_export_usage_data(limit=1000) + + # Verify the result structure for empty data + assert isinstance(result, dict) + assert result['usage_data'] == [] + assert result['cbf_data'] == [] + assert result['summary']['total_records'] == 0 + assert result['summary']['total_cost'] == 0 + assert result['summary']['total_tokens'] == 0 + + @pytest.mark.asyncio + async def test_dry_run_export_usage_data_cbf_transformation_failure(self): + """ + Test that dry_run_export_usage_data handles CBF transformation failure gracefully. + """ + logger = CloudZeroLogger() + + # Mock database data + mock_usage_data = pl.DataFrame({ + 'date': ['2025-01-19'], + 'model': ['gpt-4'], + 'spend': [0.01], + 'successful_requests': [1] + }) + + # Mock empty CBF data (transformation failed) + mock_empty_cbf_data = pl.DataFrame() + + with patch('litellm.integrations.cloudzero.cloudzero.LiteLLMDatabase') as mock_db_class, \ + patch('litellm.integrations.cloudzero.cloudzero.CBFTransformer') as mock_transformer_class: + + # Setup mocks + mock_db = AsyncMock() + mock_db.get_usage_data.return_value = mock_usage_data + mock_db_class.return_value = mock_db + + mock_transformer = MagicMock() + mock_transformer.transform.return_value = mock_empty_cbf_data + mock_transformer_class.return_value = mock_transformer + + # Call the method + result = await logger.dry_run_export_usage_data(limit=1000) + + # Verify the result handles CBF transformation failure + assert isinstance(result, dict) + assert len(result['usage_data']) == 1 # Usage data should still be present + assert result['cbf_data'] == [] # CBF data should be empty + assert result['summary']['total_cost'] == 0.01 # Should calculate from usage data From e27c4c98c0f41027967c064f4a8c4c99524975a1 Mon Sep 17 00:00:00 2001 From: eycjur Date: Sat, 6 Sep 2025 21:11:53 +0900 Subject: [PATCH 31/40] Added conditional branch for gpt-oss --- .../bedrock/chat/converse_transformation.py | 26 +++++++++++++------ 1 file changed, 18 insertions(+), 8 deletions(-) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 273b12c9c39..9e885b3f839 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -164,7 +164,9 @@ class AmazonConverseConfig(BaseConfig): # only anthropic and mistral support tool choice config. otherwise (E.g. cohere) will fail the call - https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_ToolChoice.html supported_params.append("tool_choice") - if ( + if "gpt-oss" in model: + supported_params.append("reasoning_effort") + elif ( "claude-3-7" in model or "claude-sonnet-4" in model or "claude-opus-4" in model @@ -319,7 +321,6 @@ class AmazonConverseConfig(BaseConfig): return computer_use_tools, regular_tools - def _create_json_tool_call_for_response_format( self, json_schema: Optional[dict] = None, @@ -462,13 +463,22 @@ class AmazonConverseConfig(BaseConfig): if param == "thinking": optional_params["thinking"] = value elif param == "reasoning_effort" and isinstance(value, str): - optional_params["thinking"] = AnthropicConfig._map_reasoning_effort( - value - ) + if "gpt-oss" in model: + # GPT-OSS models: keep reasoning_effort as-is + # It will be passed through to additionalModelRequestFields + optional_params["reasoning_effort"] = value + continue + else: + # Anthropic and other models: convert to thinking parameter + optional_params["thinking"] = AnthropicConfig._map_reasoning_effort( + value + ) - self.update_optional_params_with_thinking_tokens( - non_default_params=non_default_params, optional_params=optional_params - ) + # Only update thinking tokens for non-GPT-OSS models + if not ("gpt-oss" in model): + self.update_optional_params_with_thinking_tokens( + non_default_params=non_default_params, optional_params=optional_params + ) return optional_params From 58cf72ef5e15e49669e10a43095d119a3bfeab3a Mon Sep 17 00:00:00 2001 From: eycjur Date: Sat, 6 Sep 2025 21:12:18 +0900 Subject: [PATCH 32/40] add test --- tests/llm_translation/test_bedrock_gpt_oss.py | 26 +++++++++++++++++++ 1 file changed, 26 insertions(+) diff --git a/tests/llm_translation/test_bedrock_gpt_oss.py b/tests/llm_translation/test_bedrock_gpt_oss.py index 61bce04e2d0..9487abfbc77 100644 --- a/tests/llm_translation/test_bedrock_gpt_oss.py +++ b/tests/llm_translation/test_bedrock_gpt_oss.py @@ -2,11 +2,13 @@ from base_llm_unit_tests import BaseLLMChatTest import pytest import sys import os +from unittest.mock import patch, MagicMock sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path import litellm +from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig class TestBedrockGPTOSS(BaseLLMChatTest): @@ -25,3 +27,27 @@ class TestBedrockGPTOSS(BaseLLMChatTest): Remove override once we have access to Bedrock prompt caching """ pass + + @pytest.mark.parametrize("model", [ + "bedrock/openai.gpt-oss-20b-1:0", + "bedrock/openai.gpt-oss-120b-1:0", + ]) + def test_reasoning_effort_transformation_gpt_oss(self, model): + """Test that reasoning_effort is handled correctly for GPT-OSS models.""" + config = AmazonConverseConfig() + + # Test GPT-OSS model - should keep reasoning_effort as-is + non_default_params = {"reasoning_effort": "low"} + optional_params = {} + + result = config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=model, + drop_params=False, + ) + + # GPT-OSS should have reasoning_effort in result, not thinking + assert "reasoning_effort" in result + assert result["reasoning_effort"] == "low" + assert "thinking" not in result From b472bf6aef848b804d15db4d60e5fb96ee9df7aa Mon Sep 17 00:00:00 2001 From: eycjur Date: Sat, 6 Sep 2025 21:15:59 +0900 Subject: [PATCH 33/40] update docs --- docs/my-website/docs/providers/bedrock.md | 2 +- docs/my-website/docs/reasoning_content.md | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index 1356ec1744e..c191b742268 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -467,7 +467,7 @@ print(f"\nResponse: {resp}") ## Usage - 'thinking' / 'reasoning content' -This is currently only supported for Anthropic's Claude 3.7 Sonnet + Deepseek R1. +This is currently only supported for Anthropic's Claude 3.7 Sonnet + Deepseek R1 + GPT-OSS models. Works on v1.61.20+. diff --git a/docs/my-website/docs/reasoning_content.md b/docs/my-website/docs/reasoning_content.md index 5ddb5aefd47..12db17325d4 100644 --- a/docs/my-website/docs/reasoning_content.md +++ b/docs/my-website/docs/reasoning_content.md @@ -12,7 +12,7 @@ Requires LiteLLM v1.63.0+ Supported Providers: - Deepseek (`deepseek/`) - Anthropic API (`anthropic/`) -- Bedrock (Anthropic + Deepseek) (`bedrock/`) +- Bedrock (Anthropic + Deepseek + GPT-OSS) (`bedrock/`) - Vertex AI (Anthropic) (`vertexai/`) - OpenRouter (`openrouter/`) - XAI (`xai/`) From 6eb1b40336b86a99bbcec297b18e5ac2b44f0032 Mon Sep 17 00:00:00 2001 From: eycjur Date: Sat, 6 Sep 2025 21:30:08 +0900 Subject: [PATCH 34/40] refactor --- litellm/llms/bedrock/chat/converse_transformation.py | 11 +++++------ 1 file changed, 5 insertions(+), 6 deletions(-) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 9e885b3f839..080ec05576e 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -466,13 +466,12 @@ class AmazonConverseConfig(BaseConfig): if "gpt-oss" in model: # GPT-OSS models: keep reasoning_effort as-is # It will be passed through to additionalModelRequestFields - optional_params["reasoning_effort"] = value continue - else: - # Anthropic and other models: convert to thinking parameter - optional_params["thinking"] = AnthropicConfig._map_reasoning_effort( - value - ) + + # Anthropic and other models: convert to thinking parameter + optional_params["thinking"] = AnthropicConfig._map_reasoning_effort( + value + ) # Only update thinking tokens for non-GPT-OSS models if not ("gpt-oss" in model): From 67315d8727324466f84a7840b2fe594ed737d542 Mon Sep 17 00:00:00 2001 From: eycjur Date: Sat, 6 Sep 2025 21:43:30 +0900 Subject: [PATCH 35/40] fix ci --- litellm/llms/bedrock/chat/converse_transformation.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 080ec05576e..88b65132138 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -474,7 +474,7 @@ class AmazonConverseConfig(BaseConfig): ) # Only update thinking tokens for non-GPT-OSS models - if not ("gpt-oss" in model): + if "gpt-oss" not in model: self.update_optional_params_with_thinking_tokens( non_default_params=non_default_params, optional_params=optional_params ) From 51de2ebb64a1dead05fa89965968a34579c41a97 Mon Sep 17 00:00:00 2001 From: katsuhiro muto <63308909+eycjur@users.noreply.github.com> Date: Sun, 7 Sep 2025 00:58:51 +0900 Subject: [PATCH 36/40] [Feat]Cancel upstream on client disconnect (#14295) * cancel upstream on client disconnect * add comments * add test * set timeout in constraints.py * Guard against missing 'type' key * update dependency to fix uvicorn bugs --- litellm/constants.py | 3 ++ litellm/proxy/common_request_processing.py | 41 +++++++++++++++- litellm/proxy/proxy_server.py | 27 ----------- poetry.lock | 12 ++--- pyproject.toml | 2 +- requirements.txt | 2 +- .../test_client_disconnection.py | 47 +++++++++++++++++++ 7 files changed, 97 insertions(+), 37 deletions(-) create mode 100644 tests/proxy_unit_tests/test_client_disconnection.py diff --git a/litellm/constants.py b/litellm/constants.py index 089e73fc3b4..bcf394c1832 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -893,6 +893,9 @@ MAX_SPENDLOG_ROWS_TO_QUERY = int( DEFAULT_SOFT_BUDGET = float( os.getenv("DEFAULT_SOFT_BUDGET", 50.0) ) # by default all litellm proxy keys have a soft budget of 50.0 +DEFAULT_CLIENT_DISCONNECT_CHECK_TIMEOUT_SECONDS = int( + os.getenv("DEFAULT_CLIENT_DISCONNECT_CHECK_TIMEOUT_SECONDS", 600) +) # 10 minutes timeout for client disconnect checking in proxy # makes it clear this is a rate limit error for a litellm virtual key RATE_LIMIT_ERROR_MESSAGE_FOR_VIRTUAL_KEY = "LiteLLM Virtual Key user_api_key_hash" diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index e900975f1cc..a3a9c2cffc0 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1,6 +1,7 @@ import asyncio import json import logging +import time import traceback from datetime import datetime from typing import ( @@ -24,6 +25,7 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.constants import ( DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE, + DEFAULT_CLIENT_DISCONNECT_CHECK_TIMEOUT_SECONDS, STREAM_SSE_DATA_PREFIX, ) from litellm.litellm_core_utils.dd_tracing import tracer @@ -175,6 +177,29 @@ async def create_streaming_response( ) +async def _check_request_disconnection(request: Request, llm_api_call_task): + """ + Asynchronously checks if the request is disconnected at regular intervals. + If the request is disconnected + - cancel the litellm.router task + + Parameters: + - request: Request: The request object to check for disconnection. + Returns: + - None + """ + + # only run this function for configured timeout -> if these don't get cancelled -> we don't want the server to have many while loops + start_time = time.time() + while time.time() - start_time < DEFAULT_CLIENT_DISCONNECT_CHECK_TIMEOUT_SECONDS: + await asyncio.sleep(1) + message = await request.receive() + if message.get("type") == "http.disconnect": + # cancel the LLM API Call task if any passed - this is passed from individual providers + # Example OpenAI, Azure, VertexAI etc + llm_api_call_task.cancel() + return + class ProxyBaseLLMRequestProcessing: def __init__(self, data: dict): self.data = data @@ -425,12 +450,24 @@ class ProxyBaseLLMRequestProcessing: ) tasks.append(llm_call) - # wait for call to end llm_responses = asyncio.gather( *tasks ) # run the moderation check in parallel to the actual llm api call - responses = await llm_responses + # Execute the task to detect disconnection + disconnect_task = asyncio.create_task(_check_request_disconnection(request, llm_responses)) + + try: + # wait for call to end + # Note: In the case of streaming, processing does not wait here, so disconnection detection is performed in StreamingResponse. + responses = await llm_responses + disconnect_task.cancel() + + except asyncio.CancelledError: + raise HTTPException( + status_code=499, + detail="Client disconnected the request", + ) response = responses[1] diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9f1566b2e00..e15d5401374 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -997,33 +997,6 @@ db_writer_client: Optional[AsyncHTTPHandler] = None ### logger ### -async def check_request_disconnection(request: Request, llm_api_call_task): - """ - Asynchronously checks if the request is disconnected at regular intervals. - If the request is disconnected - - cancel the litellm.router task - - raises an HTTPException with status code 499 and detail "Client disconnected the request". - - Parameters: - - request: Request: The request object to check for disconnection. - Returns: - - None - """ - - # only run this function for 10 mins -> if these don't get cancelled -> we don't want the server to have many while loops - start_time = time.time() - while time.time() - start_time < 600: - await asyncio.sleep(1) - if await request.is_disconnected(): - # cancel the LLM API Call task if any passed - this is passed from individual providers - # Example OpenAI, Azure, VertexAI etc - llm_api_call_task.cancel() - - raise HTTPException( - status_code=499, - detail="Client disconnected the request", - ) - def _resolve_typed_dict_type(typ): """Resolve the actual TypedDict class from a potentially wrapped type.""" diff --git a/poetry.lock b/poetry.lock index 29d1a877087..0ab437aec25 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1,4 +1,4 @@ -# This file is automatically @generated by Poetry 2.1.2 and should not be changed by hand. +# This file is automatically @generated by Poetry 2.1.4 and should not be changed by hand. [[package]] name = "aiohappyeyeballs" @@ -6122,15 +6122,15 @@ zstd = ["zstandard (>=0.18.0)"] [[package]] name = "uvicorn" -version = "0.29.0" +version = "0.32.1" description = "The lightning-fast ASGI server." optional = true python-versions = ">=3.8" groups = ["main"] markers = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"proxy\") or extra == \"proxy\"" files = [ - {file = "uvicorn-0.29.0-py3-none-any.whl", hash = "sha256:2c2aac7ff4f4365c206fd773a39bf4ebd1047c238f8b8268ad996829323473de"}, - {file = "uvicorn-0.29.0.tar.gz", hash = "sha256:6a69214c0b6a087462412670b3ef21224fa48cae0e452b5883e8e8bdfdd11dd0"}, + {file = "uvicorn-0.32.1-py3-none-any.whl", hash = "sha256:82ad92fd58da0d12af7482ecdb5f2470a04c9c9a53ced65b9bbb4a205377602e"}, + {file = "uvicorn-0.32.1.tar.gz", hash = "sha256:ee9519c246a72b1c084cea8d3b44ed6026e78a4a309cbedae9c37e4cb9fbb175"}, ] [package.dependencies] @@ -6139,7 +6139,7 @@ h11 = ">=0.8" typing-extensions = {version = ">=4.0", markers = "python_version < \"3.11\""} [package.extras] -standard = ["colorama (>=0.4) ; sys_platform == \"win32\"", "httptools (>=0.5.0)", "python-dotenv (>=0.13)", "pyyaml (>=5.1)", "uvloop (>=0.14.0,!=0.15.0,!=0.15.1) ; sys_platform != \"win32\" and sys_platform != \"cygwin\" and platform_python_implementation != \"PyPy\"", "watchfiles (>=0.13)", "websockets (>=10.4)"] +standard = ["colorama (>=0.4) ; sys_platform == \"win32\"", "httptools (>=0.6.3)", "python-dotenv (>=0.13)", "pyyaml (>=5.1)", "uvloop (>=0.14.0,!=0.15.0,!=0.15.1) ; sys_platform != \"win32\" and sys_platform != \"cygwin\" and platform_python_implementation != \"PyPy\"", "watchfiles (>=0.13)", "websockets (>=10.4)"] [[package]] name = "uvloop" @@ -6576,4 +6576,4 @@ utils = ["numpydoc"] [metadata] lock-version = "2.1" python-versions = ">=3.8.1,<4.0, !=3.9.7" -content-hash = "f41e6359109c5c52dba2a28f301b04030d865265f408974082b390bf45568a01" +content-hash = "e48cc445bc012e020a9e311942e46833dda587b70a630a04bfc08b629746fe56" diff --git a/pyproject.toml b/pyproject.toml index 9f5d876cf2c..b1b11f5d21d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -34,7 +34,7 @@ pydantic = "^2.5.0" jsonschema = "^4.22.0" numpydoc = {version = "*", optional = true} # used in utils.py -uvicorn = {version = "^0.29.0", optional = true} +uvicorn = {version = "^0.32.0", optional = true} uvloop = {version = "^0.21.0", optional = true, markers="sys_platform != 'win32'"} gunicorn = {version = "^23.0.0", optional = true} fastapi = {version = "^0.115.5", optional = true} diff --git a/requirements.txt b/requirements.txt index 2d31819dc5b..9b858e08a03 100644 --- a/requirements.txt +++ b/requirements.txt @@ -5,7 +5,7 @@ openai==1.99.5 # openai req. fastapi==0.115.5 # server dep backoff==2.2.1 # server dep pyyaml==6.0.2 # server dep -uvicorn==0.29.0 # server dep +uvicorn==0.32.0 # server dep gunicorn==23.0.0 # server dep fastuuid==0.12.0 # for uuid4 uvloop==0.21.0 # uvicorn dep, gives us much better performance under load diff --git a/tests/proxy_unit_tests/test_client_disconnection.py b/tests/proxy_unit_tests/test_client_disconnection.py new file mode 100644 index 00000000000..d894d7ad015 --- /dev/null +++ b/tests/proxy_unit_tests/test_client_disconnection.py @@ -0,0 +1,47 @@ +""" +Test client disconnection detection functionality. +""" +import asyncio +import pytest +from unittest.mock import AsyncMock + +from litellm.proxy.common_request_processing import _check_request_disconnection + + +@pytest.mark.asyncio +async def test_check_request_disconnection_with_disconnect(): + """Test that _check_request_disconnection cancels task when client disconnects.""" + mock_request = AsyncMock() + mock_request.receive.side_effect = [ + {"type": "http.request"}, # First call + {"type": "http.disconnect"} # Second call - disconnect + ] + + mock_llm_task = AsyncMock() + + await _check_request_disconnection(mock_request, mock_llm_task) + + mock_llm_task.cancel.assert_called_once() + + +@pytest.mark.asyncio +async def test_check_request_disconnection_no_disconnect(): + """Test that _check_request_disconnection handles normal requests.""" + mock_request = AsyncMock() + mock_request.receive.return_value = {"type": "http.request"} + + mock_llm_task = AsyncMock() + + # This will timeout after 600 seconds, but we don't need to wait + # Just test that it doesn't crash immediately + task = asyncio.create_task(_check_request_disconnection(mock_request, mock_llm_task)) + await asyncio.sleep(0.1) # Let it run briefly + task.cancel() + + try: + await task + except asyncio.CancelledError: + pass + + # Task should not be cancelled during normal operation + mock_llm_task.cancel.assert_not_called() \ No newline at end of file From 3478c53c6045cfd31579dd97f3029755345f924e Mon Sep 17 00:00:00 2001 From: Duc Tran Date: Sat, 6 Sep 2025 23:06:04 +0700 Subject: [PATCH 37/40] Update constants.py (#14242) --- litellm/constants.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/constants.py b/litellm/constants.py index bcf394c1832..ce485fc7264 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -15,7 +15,7 @@ DEFAULT_SQS_FLUSH_INTERVAL_SECONDS = int( os.getenv("DEFAULT_SQS_FLUSH_INTERVAL_SECONDS", 10) ) DEFAULT_NUM_WORKERS_LITELLM_PROXY = int( - os.getenv("DEFAULT_NUM_WORKERS_LITELLM_PROXY", 4) + os.getenv("DEFAULT_NUM_WORKERS_LITELLM_PROXY", os.cpu_count() or 4) ) DEFAULT_SQS_BATCH_SIZE = int(os.getenv("DEFAULT_SQS_BATCH_SIZE", 512)) SQS_SEND_MESSAGE_ACTION = "SendMessage" From 0cb01d60278cb370e56ee619aeea9320cac96770 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 6 Sep 2025 09:06:43 -0700 Subject: [PATCH 38/40] Fix: Include model name in Azure base_model error (#14294) Co-authored-by: Cursor Agent Co-authored-by: ishaan --- litellm/router.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/router.py b/litellm/router.py index 5eea60e4b3d..6255c2fdf92 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -5465,7 +5465,7 @@ class Router: ## SET MODEL TO 'model=' - if base_model is None + not azure if custom_llm_provider == "azure" and base_model is None: verbose_router_logger.error( - "Could not identify azure model. Set azure 'base_model' for accurate max tokens, cost tracking, etc.- https://docs.litellm.ai/docs/proxy/cost_tracking#spend-tracking-for-azure-openai-models" + f"Could not identify azure model '{_model}'. Set azure 'base_model' for accurate max tokens, cost tracking, etc.- https://docs.litellm.ai/docs/proxy/cost_tracking#spend-tracking-for-azure-openai-models" ) elif custom_llm_provider != "azure": model = _model From cb117647fce2551f5241fe27241e7b23b0b25a68 Mon Sep 17 00:00:00 2001 From: Mubashir Osmani Date: Sat, 6 Sep 2025 12:09:22 -0400 Subject: [PATCH 39/40] [docs]: added loom for claude code (#14223) * added loom for claude code * docs: add web search models * added new loom --- docs/my-website/docs/completion/web_search.md | 59 ++++++++++++++++++- .../docs/tutorials/claude_responses_api.md | 13 ++++ 2 files changed, 71 insertions(+), 1 deletion(-) diff --git a/docs/my-website/docs/completion/web_search.md b/docs/my-website/docs/completion/web_search.md index fe49be852a7..262e3fc4f9c 100644 --- a/docs/my-website/docs/completion/web_search.md +++ b/docs/my-website/docs/completion/web_search.md @@ -8,10 +8,25 @@ Use web search with litellm | Feature | Details | |---------|---------| | Supported Endpoints | - `/chat/completions`
- `/responses` | -| Supported Providers | `openai`, `xai`, `vertex_ai`, `gemini`, `perplexity` | +| Supported Providers | `openai`, `xai`, `vertex_ai`, `anthropic`, `gemini`, `perplexity` | | LiteLLM Cost Tracking | βœ… Supported | | LiteLLM Version | `v1.71.0+` | +## Which Search Engine is Used? + +Each provider uses their own search backend: + +| Provider | Search Engine | Notes | +|----------|---------------|-------| +| **OpenAI** (`gpt-4o-search-preview`) | OpenAI's internal search | Real-time web data | +| **xAI** (`grok-3`) | xAI's search + X/Twitter | Real-time social media data | +| **Google AI/Vertex** (`gemini-2.0-flash`) | **Google Search** | Uses actual Google search results | +| **Anthropic** (`claude-3-5-sonnet`) | Anthropic's web search | Real-time web data | +| **Perplexity** | Perplexity's search engine | AI-powered search and reasoning | + +:::info +**Anthropic Web Search Models**: Claude models that support web search: `claude-3-5-sonnet-latest`, `claude-3-5-sonnet-20241022`, `claude-3-5-haiku-latest`, `claude-3-5-haiku-20241022`, `claude-3-7-sonnet-20250219` +::: ## `/chat/completions` (litellm.completion) @@ -56,6 +71,12 @@ model_list: model: xai/grok-3 api_key: os.environ/XAI_API_KEY + # Anthropic + - model_name: claude-3-5-sonnet-latest + litellm_params: + model: anthropic/claude-3-5-sonnet-latest + api_key: os.environ/ANTHROPIC_API_KEY + # VertexAI - model_name: gemini-2-flash litellm_params: @@ -143,6 +164,31 @@ response = completion( ) ``` +**Anthropic (using web_search_options)** +```python showLineNumbers +from litellm import completion + +# Customize search context size for Anthropic +response = completion( + model="anthropic/claude-3-5-sonnet-latest", + messages=[ + { + "role": "user", + "content": "What was a positive news story from today?", + } + ], + web_search_options={ + "search_context_size": "medium", # Options: "low", "medium" (default), "high" + "user_location": { + "type": "approximate", + "approximate": { + "city": "San Francisco", + }, + } + } +) +``` + **VertexAI/Gemini (using web_search_options)** ```python showLineNumbers from litellm import completion @@ -375,6 +421,9 @@ assert litellm.supports_web_search(model="openai/gpt-4o-search-preview") == True # Check xAI models assert litellm.supports_web_search(model="xai/grok-3") == True +# Check Anthropic models +assert litellm.supports_web_search(model="anthropic/claude-3-5-sonnet-latest") == True + # Check VertexAI models assert litellm.supports_web_search(model="gemini-2.0-flash") == True @@ -405,6 +454,14 @@ model_list: model_info: supports_web_search: True + # Anthropic + - model_name: claude-3-5-sonnet-latest + litellm_params: + model: anthropic/claude-3-5-sonnet-latest + api_key: os.environ/ANTHROPIC_API_KEY + model_info: + supports_web_search: True + # VertexAI - model_name: gemini-2-flash litellm_params: diff --git a/docs/my-website/docs/tutorials/claude_responses_api.md b/docs/my-website/docs/tutorials/claude_responses_api.md index 09b352a7663..a333faee5d2 100644 --- a/docs/my-website/docs/tutorials/claude_responses_api.md +++ b/docs/my-website/docs/tutorials/claude_responses_api.md @@ -12,6 +12,13 @@ This tutorial is based on [Anthropic's official LiteLLM configuration documentat ::: +
+ +### LiteLLM x Claude Code + + + + ## Prerequisites - [Claude Code](https://docs.anthropic.com/en/docs/claude-code/overview) installed @@ -83,11 +90,17 @@ curl -X POST http://0.0.0.0:4000/v1/messages \ Configure Claude Code to use LiteLLM's unified endpoint: +Either a virtual key / master key can be used here + ```bash export ANTHROPIC_BASE_URL="http://0.0.0.0:4000" export ANTHROPIC_AUTH_TOKEN="$LITELLM_MASTER_KEY" ``` +:::tip +LITELLM_MASTER_KEY gives claude access to all proxy models, whereas a virtual key would be limited to the models set in UI +::: + #### Method 2: Provider-specific Pass-through Endpoint Alternatively, use the Anthropic pass-through endpoint: From 29e410b04c099c54c97cbc6b2160d9b0ee4a26d2 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 6 Sep 2025 09:12:30 -0700 Subject: [PATCH 40/40] docs Video Walkthrough claude code --- docs/my-website/docs/tutorials/claude_responses_api.md | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/docs/my-website/docs/tutorials/claude_responses_api.md b/docs/my-website/docs/tutorials/claude_responses_api.md index a333faee5d2..5000161a520 100644 --- a/docs/my-website/docs/tutorials/claude_responses_api.md +++ b/docs/my-website/docs/tutorials/claude_responses_api.md @@ -14,10 +14,9 @@ This tutorial is based on [Anthropic's official LiteLLM configuration documentat
-### LiteLLM x Claude Code - - +### Video Walkthrough + ## Prerequisites