Merge branch 'main' into litellm_responses_structured_output

This commit is contained in:
Krish Dholakia 2025-09-06 09:20:00 -07:00 • committed by GitHub
commit 1980960218
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
69 changed files with 3505 additions and 549 deletions

View file

@ -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:

View file

@ -8,10 +8,25 @@ Use web search with litellm
| Feature | Details |
|---------|---------|
| Supported Endpoints | - `/chat/completions` <br/> - `/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:

View file

@ -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<br/>• Manual data export<br/>• Dry run testing<br/>• 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.
<iframe width="840" height="500" src="https://www.loom.com/embed/59b57593183f4cc3b1c05a2dd3277f92" frameborder="0" webkitallowfullscreen mozallowfullscreen allowfullscreen></iframe>
### 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)

View file

@ -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+.

View file

@ -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.

View file

@ -3,7 +3,7 @@ https://www.volcengine.com/docs/82379/1263482
:::tip
**We support ALL Volcengine NIM models, just set `model=volcengine/<any-model-on-volcengine>` as a prefix when sending litellm requests**
**We support ALL Volcengine models including Chat and Embeddings, just set `model=volcengine/<any-model-on-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/<OUR_ENDPOINT_ID>` 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/<OUR_ENDPOINT_ID>` 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/<OUR_ENDPOINT_ID>` as a
```yaml
model_list:
# Chat model
- model_name: volcengine-model
litellm_params:
model: volcengine/<OUR_ENDPOINT_ID>
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"]
}'
```

View file

@ -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**

View file

@ -124,6 +124,8 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \
}'
```
</TabItem>
</Tabs>
### Test - Loadbalancing

View file

@ -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/`)

View file

@ -12,6 +12,12 @@ This tutorial is based on [Anthropic's official LiteLLM configuration documentat
:::
<br />
### Video Walkthrough
<iframe width="840" height="500" src="https://www.loom.com/embed/3c17d683cdb74d36a3698763cc558f56" frameborder="0" webkitallowfullscreen mozallowfullscreen allowfullscreen></iframe>
## 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:

View file

@ -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,

View file

@ -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']}"
)

View file

@ -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))

View file

@ -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")
# 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
)

View file

@ -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'

View file

@ -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)}")

View file

@ -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:<key> 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

View file

@ -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,

View file

@ -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]:
"""

View file

@ -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):

View file

@ -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

View file

@ -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")

View file

@ -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"] = (

View file

@ -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",
]

View file

@ -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

View file

@ -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

View file

@ -0,0 +1,7 @@
"""
Volcengine Embedding Module
"""
from .transformation import VolcEngineEmbeddingConfig
__all__ = ["VolcEngineEmbeddingConfig"]

View file

@ -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,
)

View file

@ -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:

View file

@ -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"
}
}
}

View file

@ -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

View file

@ -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 <your-model-name>`
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:

View file

@ -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)
)

View file

@ -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:

View file

@ -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

View file

@ -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,
}

View file

@ -3,3 +3,5 @@ model_list:
litellm_params:
model: openai/*
api_base: https://exampleopenaiendpoint-production-0ee2.up.railway.app/
litellm_settings:
callbacks: ["cloudzero"]

View file

@ -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

View file

@ -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:

View file

@ -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:

View file

@ -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 []

View file

@ -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

View file

@ -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[

View file

@ -113,6 +113,7 @@ class Schema(TypedDict, total=False):
pattern: str
example: Any
anyOf: List["Schema"]
additionalProperties: Any
class FunctionDeclaration(TypedDict, total=False):

View file

@ -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):

View file

@ -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

View file

@ -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"
}
}
}

12
poetry.lock generated
View file

@ -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"

View file

@ -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}

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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()

View file

@ -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"

View file

@ -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

View file

@ -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

View file

@ -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'), \

View file

@ -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

View file

@ -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,
}
}

View file

@ -0,0 +1 @@
# Volcengine tests

View file

@ -0,0 +1 @@
# Volcengine embedding tests

View file

@ -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

View file

@ -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__])

View file

@ -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)

View file

@ -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__])

View file

@ -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"

View file

@ -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)

View file

@ -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"

View file

@ -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<UserInfo | null>(null);
const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false);
const [isLoading, setIsLoading] = useState(true);
const [isEditing, setIsEditing] = useState(startInEditMode);
const [userModels, setUserModels] = useState<string[]>([]);
const [isInvitationLinkModalVisible, setIsInvitationLinkModalVisible] = useState(false);
const [invitationLinkData, setInvitationLinkData] = useState<InvitationLink | null>(null);
const [baseUrl, setBaseUrl] = useState<string | null>(null);
const [activeTab, setActiveTab] = useState(initialTab);
const [copiedStates, setCopiedStates] = useState<Record<string, boolean>>({});
const [userData, setUserData] = useState<UserInfo | null>(null)
const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false)
const [isLoading, setIsLoading] = useState(true)
const [isEditing, setIsEditing] = useState(startInEditMode)
const [userModels, setUserModels] = useState<string[]>([])
const [isInvitationLinkModalVisible, setIsInvitationLinkModalVisible] = useState(false)
const [invitationLinkData, setInvitationLinkData] = useState<InvitationLink | null>(null)
const [baseUrl, setBaseUrl] = useState<string | null>(null)
const [activeTab, setActiveTab] = useState(initialTab)
const [copiedStates, setCopiedStates] = useState<Record<string, boolean>>({})
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 (
<div className="p-4">
@ -200,9 +201,9 @@ export default function UserInfoView({
icon={copiedStates["user-id"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
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"
}`}
/>
</div>
@ -284,7 +285,35 @@ export default function UserInfoView({
<Card>
<Text>Teams</Text>
<div className="mt-2">
<Text>{userData.teams?.length || 0} teams</Text>
{userData.teams?.length && userData.teams?.length > 0 ? (
<div className="flex flex-wrap gap-2">
{userData.teams?.slice(0, isTeamsExpanded ? userData.teams.length : 20).map((team, index) => (
<Badge key={index} color="blue" title={team.team_alias}>
{team.team_alias}
</Badge>
))}
{!isTeamsExpanded && userData.teams?.length > 20 && (
<Badge
color="gray"
className="cursor-pointer hover:bg-gray-200 transition-colors"
onClick={() => setIsTeamsExpanded(true)}
>
+{userData.teams.length - 20} more
</Badge>
)}
{isTeamsExpanded && userData.teams?.length > 20 && (
<Badge
color="gray"
className="cursor-pointer hover:bg-gray-200 transition-colors"
onClick={() => setIsTeamsExpanded(false)}
>
Show Less
</Badge>
)}
</div>
) : (
<Text>No teams</Text>
)}
</div>
</Card>
@ -344,9 +373,9 @@ export default function UserInfoView({
icon={copiedStates["user-id"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
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"
}`}
/>
</div>
@ -384,11 +413,33 @@ export default function UserInfoView({
<Text className="font-medium">Teams</Text>
<div className="flex flex-wrap gap-2 mt-1">
{userData.teams?.length && userData.teams?.length > 0 ? (
userData.teams?.map((team, index) => (
<span key={index} className="px-2 py-1 bg-blue-100 rounded text-xs">
{team.team_alias || team.team_id}
</span>
))
<>
{userData.teams?.slice(0, isTeamsExpanded ? userData.teams.length : 20).map((team, index) => (
<span
key={index}
className="px-2 py-1 bg-blue-100 rounded text-xs"
title={team.team_alias || team.team_id}
>
{team.team_alias || team.team_id}
</span>
))}
{!isTeamsExpanded && userData.teams?.length > 20 && (
<span
className="px-2 py-1 bg-gray-100 rounded text-xs cursor-pointer hover:bg-gray-200 transition-colors"
onClick={() => setIsTeamsExpanded(true)}
>
+{userData.teams.length - 20} more
</span>
)}
{isTeamsExpanded && userData.teams?.length > 20 && (
<span
className="px-2 py-1 bg-gray-100 rounded text-xs cursor-pointer hover:bg-gray-200 transition-colors"
onClick={() => setIsTeamsExpanded(false)}
>
Show Less
</span>
)}
</>
) : (
<Text>No teams</Text>
)}