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/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/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/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/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/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/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/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/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/`) diff --git a/docs/my-website/docs/tutorials/claude_responses_api.md b/docs/my-website/docs/tutorials/claude_responses_api.md index 09b352a7663..5000161a520 100644 --- a/docs/my-website/docs/tutorials/claude_responses_api.md +++ b/docs/my-website/docs/tutorials/claude_responses_api.md @@ -12,6 +12,12 @@ This tutorial is based on [Anthropic's official LiteLLM configuration documentat ::: +
+ +### Video Walkthrough + + + ## Prerequisites - [Claude Code](https://docs.anthropic.com/en/docs/claude-code/overview) installed @@ -83,11 +89,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: diff --git a/litellm/__init__.py b/litellm/__init__.py index d411ff0ad45..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 @@ -1196,7 +1197,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/_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/litellm/constants.py b/litellm/constants.py index 21e30bef32b..0dee5638ccb 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" @@ -51,6 +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) ) @@ -873,6 +890,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)) @@ -890,6 +910,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" @@ -927,6 +950,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 85aa1679732..727dabc0945 100644 --- a/litellm/integrations/cloudzero/cloudzero.py +++ b/litellm/integrations/cloudzero/cloudzero.py @@ -1,8 +1,8 @@ -import asyncio import os -from datetime import datetime, timedelta -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 @@ -10,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): """ @@ -29,18 +34,75 @@ 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, target_hour: datetime, limit: Optional[int] = 1000, operation: str = "replace_hourly"): + async def initialize_cloudzero_export_job(self): """ - Exports the usage data for a specific hour to CloudZero. + Handler for initializing CloudZero export job. - - Reads spend logs from the DB for the specified hour + 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. + + - 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 +114,27 @@ 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, + 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") + 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,60 +153,84 @@ 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. + Returns the data that would be exported to CloudZero without actually sending it. Args: - target_hour: The specific hour to fetch data for - limit: Optional limit on number of records to fetch (default: 1000) + limit: Limit number of records to display (default: 10000) 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) + dict: Contains usage_data, cbf_data, and summary statistics """ 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 { + "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 usage data found") - return + verbose_logger.warning("CloudZero Dry Run: No valid data after transformation") + 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)}") @@ -155,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) @@ -170,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, @@ -199,55 +316,33 @@ class CloudZeroLogger(CustomLogger): console.print(f" Unique Services: {unique_services}") 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. - async def init_background_job(self, redis_cache=None): + Starts the background job that exports the usage data to CloudZero every hour. """ - Initialize a background job that exports usage data every hour. - Uses PodLockManager to ensure only one instance runs the export at a time. + from litellm.constants import CLOUDZERO_EXPORT_INTERVAL_MINUTES + from litellm.integrations.custom_logger import CustomLogger - Args: - redis_cache: Redis cache instance for pod locking - """ - from litellm.proxy.db.db_transaction_queue.pod_lock_manager import ( - PodLockManager, + + prometheus_loggers: List[CustomLogger] = ( + litellm.logging_callback_manager.get_custom_loggers_for_type( + callback_type=CloudZeroLogger + ) ) - - 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 + # 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 6d12c5cfbd9..71b4125ed75 100644 --- a/litellm/integrations/cloudzero/database.py +++ b/litellm/integrations/cloudzero/database.py @@ -12,14 +12,13 @@ # 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 datetime import datetime from typing import Any, Dict, Optional import polars as pl @@ -37,61 +36,88 @@ 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, + 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() - # Calculate hour range - hour_start = target_hour.replace(minute=0, second=0, microsecond=0) - hour_end = hour_start + timedelta(hours=1) + # 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()}'") - # Convert datetime objects to ISO format strings for PostgreSQL compatibility - hour_start_str = hour_start.isoformat() - hour_end_str = hour_end.isoformat() + where_clause = "" + if where_conditions: + where_clause = "WHERE " + " AND ".join(where_conditions) - # Query to get spend logs for the specific hour - query = """ - SELECT * - FROM "LiteLLM_SpendLogs" - WHERE "startTime" >= $1::timestamp - AND "startTime" < $2::timestamp - ORDER BY "startTime" ASC + # 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: query += f" LIMIT {limit}" try: - db_response = await client.db.query_raw(query, hour_start_str, hour_end_str) - # Convert the response to polars DataFrame - return pl.DataFrame(db_response) if db_response else pl.DataFrame() + db_response = await client.db.query_raw(query) + # 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 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 daily user spend table.""" client = self._ensure_prisma_client() try: - # Get row count from SpendLogs table - spend_logs_count = await self._get_table_row_count('LiteLLM_SpendLogs') + # Get row count from user spend table + user_count = await self._get_table_row_count('LiteLLM_DailyUserSpend') - # Get column structure from spend logs table + # Get column structure from user spend table 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, - 'table_breakdown': { - 'spend_logs': spend_logs_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 7091ea26b95..e0263295388 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) @@ -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: @@ -35,160 +35,99 @@ 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 tracking - dimensions = { - 'entity_type': '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), - } + # Build dimensions for CloudZero + model = str(row.get('model', '')) + api_key_hash = str(row.get('api_key', ''))[:8] # First 8 chars for identification - # Add tag if present - if tag is not None and str(tag) not in ['', 'null', 'None']: - dimensions['tag'] = str(tag) + # 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': 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)), + 'cache_creation_tokens': str(row.get('cache_creation_input_tokens', 0)), + 'cache_read_tokens': str(row.get('cache_read_input_tokens', 0)), + } # 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 +145,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' and value != 'unknown': # Only add meaningful 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/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/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/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 7bc7702684d..7134f52c95a 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 ( @@ -3361,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): @@ -3581,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/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 273b12c9c39..88b65132138 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,21 @@ class AmazonConverseConfig(BaseConfig): if param == "thinking": optional_params["thinking"] = value elif param == "reasoning_effort" and isinstance(value, str): + if "gpt-oss" in model: + # GPT-OSS models: keep reasoning_effort as-is + # It will be passed through to additionalModelRequestFields + continue + + # 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 "gpt-oss" not in model: + self.update_optional_params_with_thinking_tokens( + non_default_params=non_default_params, optional_params=optional_params + ) return optional_params diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index 908419f7193..d3df5bbf361 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: @@ -379,6 +376,25 @@ 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 = [] + if isinstance(content, list): + for item in content: + text = item.get("text", None) + if citations_item := item.get("citations"): + citations.append( + [ + {**citation, "supported_text": text} + for citation in citations_item + ] + ) + return citations or None + def _transform_dbrx_choices( self, choices: List[DatabricksChoice], json_mode: Optional[bool] = None ) -> List[Choices]: @@ -427,12 +443,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 +584,17 @@ class DatabricksChatResponseIterator(BaseModelResponseIterator): for _tc in tool_calls: if _tc.get("function", {}).get("arguments") == "{}": _tc["function"]["arguments"] = "" # avoid invalid json + 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/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..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 @@ -30,6 +30,10 @@ 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, + 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 ( @@ -422,8 +426,25 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _map_reasoning_effort_to_thinking_budget( reasoning_effort: str, + model: Optional[str] = None, ) -> GeminiThinkingConfig: - if reasoning_effort == "low": + if reasoning_effort == "minimal": + # Use model-specific minimum thinking budget or fallback + # 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" in model.lower(): + budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH + else: + budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET + + return { + "thinkingBudget": budget, + "includeThoughts": True, + } + elif reasoning_effort == "low": return { "thinkingBudget": DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET, "includeThoughts": True, @@ -600,7 +621,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"] = ( diff --git a/litellm/llms/volcengine/__init__.py b/litellm/llms/volcengine/__init__.py new file mode 100644 index 00000000000..0887937bed5 --- /dev/null +++ b/litellm/llms/volcengine/__init__.py @@ -0,0 +1,24 @@ +""" +Volcengine LLM Provider +Support for Volcengine (ByteDance) chat and embedding models +""" + +from .chat.transformation import VolcEngineChatConfig +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 + +__all__ = [ + "VolcEngineChatConfig", + "VolcEngineConfig", # backward compatibility + "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..7b3efc4f961 --- /dev/null +++ b/litellm/llms/volcengine/embedding/__init__.py @@ -0,0 +1,7 @@ +""" +Volcengine Embedding Module +""" + +from .transformation import VolcEngineEmbeddingConfig + +__all__ = ["VolcEngineEmbeddingConfig"] diff --git a/litellm/llms/volcengine/embedding/transformation.py b/litellm/llms/volcengine/embedding/transformation.py new file mode 100644 index 00000000000..20747b76725 --- /dev/null +++ b/litellm/llms/volcengine/embedding/transformation.py @@ -0,0 +1,211 @@ +""" +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 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], + 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_embedding_request( + self, + model: str, + input: AllEmbeddingInputValues, + optional_params: dict, + headers: dict, + ) -> dict: + """Transform embedding request to Volcengine format""" + # 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, + 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""" + 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) + + 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 + 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} + + 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 + # Convert dict to httpx.Headers if needed + if isinstance(headers, dict): + headers = httpx.Headers(headers) + return VolcEngineError( + status_code=status_code, + message=error_message, + headers=headers, + ) diff --git a/litellm/main.py b/litellm/main.py index d0377490942..decbebaf485 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -502,7 +502,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 @@ -3671,7 +3671,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) @@ -3681,7 +3681,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 @@ -4503,6 +4503,36 @@ def embedding( # noqa: PLR0915 client=client, aembedding=aembedding, ) + elif custom_llm_provider == "volcengine": + volcengine_key = ( + api_key + or litellm.api_key + or get_secret_str("ARK_API_KEY") + or get_secret_str("VOLCENGINE_API_KEY") + ) + 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." + ) + if extra_headers is not None and isinstance(extra_headers, dict): + headers = extra_headers + else: + 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/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/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/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/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 68fa80c2b0f..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 @@ -109,7 +111,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 @@ -118,7 +119,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 @@ -132,7 +132,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: @@ -145,7 +144,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}" @@ -168,7 +166,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 @@ -180,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 @@ -430,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] @@ -462,7 +494,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, @@ -480,7 +511,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: @@ -501,7 +531,6 @@ class ProxyBaseLLMRequestProcessing: headers=custom_headers, ) else: - selected_data_generator = select_data_generator( response=response, user_api_key_dict=user_api_key_dict, @@ -740,7 +769,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) ) 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/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 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/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..e15d5401374 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, ) @@ -993,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.""" @@ -3807,13 +3784,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 67de202aa7a..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( @@ -434,15 +413,17 @@ async def cloudzero_dry_run_export( # 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( @@ -494,20 +474,24 @@ async def cloudzero_export( # 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/router.py b/litellm/router.py index 190d19598c3..6255c2fdf92 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: """ @@ -5451,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 @@ -5658,6 +5672,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 +5685,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, ) @@ -6783,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/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/litellm/types/llms/databricks.py b/litellm/types/llms/databricks.py index bb59b692ef7..112427c6b56 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): @@ -35,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[ diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 1b74ee25803..625a76b6789 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: Any class FunctionDeclaration(TypedDict, total=False): 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/litellm/utils.py b/litellm/utils.py index aa582fa751e..14d00cc528b 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7118,6 +7118,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/model_prices_and_context_window.json b/model_prices_and_context_window.json index a7586124509..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", @@ -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/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/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 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 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 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/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 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 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'), \ 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 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..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 @@ -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(): @@ -85,8 +88,101 @@ 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" 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": "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, + } + ], + } + ], + }, + "index": 0, + "finish_reason": "stop", + } + ] + + choices = config._transform_dbrx_choices(choices=databricks_choices) + + assert choices[0].message.provider_specific_fields == { + "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", + } + ] + ] + } + + +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": { + "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, + } + ], + } + + parsed = iterator.chunk_parser(chunk) + assert parsed.choices[0].delta.provider_specific_fields == { + "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, + } + } 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/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/llms/volcengine/test_volcengine_embedding.py b/tests/test_litellm/llms/volcengine/test_volcengine_embedding.py new file mode 100644 index 00000000000..3be7f6ca8d4 --- /dev/null +++ b/tests/test_litellm/llms/volcengine/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 tests.llm_translation.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/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) + 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__]) 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 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) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 6f63b866220..6fb268d942c 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" 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..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 @@ -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" }`} />
@@ -284,7 +285,35 @@ export default function UserInfoView({ Teams
- {userData.teams?.length || 0} teams + {userData.teams?.length && userData.teams?.length > 0 ? ( +
+ {userData.teams?.slice(0, isTeamsExpanded ? userData.teams.length : 20).map((team, index) => ( + + {team.team_alias} + + ))} + {!isTeamsExpanded && userData.teams?.length > 20 && ( + setIsTeamsExpanded(true)} + > + +{userData.teams.length - 20} more + + )} + {isTeamsExpanded && userData.teams?.length > 20 && ( + setIsTeamsExpanded(false)} + > + Show Less + + )} +
+ ) : ( + No teams + )}
@@ -344,9 +373,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" }`} /> @@ -384,11 +413,33 @@ export default function UserInfoView({ Teams
{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 && ( + setIsTeamsExpanded(true)} + > + +{userData.teams.length - 20} more + + )} + {isTeamsExpanded && userData.teams?.length > 20 && ( + setIsTeamsExpanded(false)} + > + Show Less + + )} + ) : ( No teams )}