mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge branch 'main' into litellm_responses_structured_output
This commit is contained in:
commit
1980960218
69 changed files with 3505 additions and 549 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
209
docs/my-website/docs/observability/cloudzero.md
Normal file
209
docs/my-website/docs/observability/cloudzero.md
Normal 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)
|
||||
|
|
@ -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+.
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
}'
|
||||
```
|
||||
|
|
@ -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**
|
||||
|
|
|
|||
|
|
@ -124,6 +124,8 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \
|
|||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Test - Loadbalancing
|
||||
|
||||
|
|
|
|||
|
|
@ -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/`)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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']}"
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
@ -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'
|
||||
|
||||
|
|
|
|||
|
|
@ -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)}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"] = (
|
||||
|
|
|
|||
24
litellm/llms/volcengine/__init__.py
Normal file
24
litellm/llms/volcengine/__init__.py
Normal 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",
|
||||
]
|
||||
|
|
@ -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
|
||||
62
litellm/llms/volcengine/common_utils.py
Normal file
62
litellm/llms/volcengine/common_utils.py
Normal 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
|
||||
7
litellm/llms/volcengine/embedding/__init__.py
Normal file
7
litellm/llms/volcengine/embedding/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""
|
||||
Volcengine Embedding Module
|
||||
"""
|
||||
|
||||
from .transformation import VolcEngineEmbeddingConfig
|
||||
|
||||
__all__ = ["VolcEngineEmbeddingConfig"]
|
||||
211
litellm/llms/volcengine/embedding/transformation.py
Normal file
211
litellm/llms/volcengine/embedding/transformation.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -3,3 +3,5 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/*
|
||||
api_base: https://exampleopenaiendpoint-production-0ee2.up.railway.app/
|
||||
litellm_settings:
|
||||
callbacks: ["cloudzero"]
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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[
|
||||
|
|
|
|||
|
|
@ -113,6 +113,7 @@ class Schema(TypedDict, total=False):
|
|||
pattern: str
|
||||
example: Any
|
||||
anyOf: List["Schema"]
|
||||
additionalProperties: Any
|
||||
|
||||
|
||||
class FunctionDeclaration(TypedDict, total=False):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
12
poetry.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
47
tests/proxy_unit_tests/test_client_disconnection.py
Normal file
47
tests/proxy_unit_tests/test_client_disconnection.py
Normal 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()
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
183
tests/test_litellm/integrations/cloudzero/test_transform.py
Normal file
183
tests/test_litellm/integrations/cloudzero/test_transform.py
Normal 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
|
||||
|
|
@ -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'), \
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
}
|
||||
|
|
|
|||
1
tests/test_litellm/llms/volcengine/__init__.py
Normal file
1
tests/test_litellm/llms/volcengine/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
# Volcengine tests
|
||||
1
tests/test_litellm/llms/volcengine/embedding/__init__.py
Normal file
1
tests/test_litellm/llms/volcengine/embedding/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
# Volcengine embedding tests
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
262
tests/test_litellm/llms/volcengine/test_volcengine_embedding.py
Normal file
262
tests/test_litellm/llms/volcengine/test_volcengine_embedding.py
Normal 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__])
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
109
tests/test_litellm/test_redis.py
Normal file
109
tests/test_litellm/test_redis.py
Normal 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)
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
)}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue