mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge pull request #23468 from joereyna/fix/mcp-rest-m2m-oauth2-flow
fix: set oauth2_flow when building MCPServer in _execute_with_mcp_client
This commit is contained in:
commit
e4c8f95328
98 changed files with 5697 additions and 1715 deletions
1095
.circleci/config.yml
1095
.circleci/config.yml
File diff suppressed because it is too large
Load diff
128
docs/my-website/blog/video_characters_litellm/index.md
Normal file
128
docs/my-website/blog/video_characters_litellm/index.md
Normal file
|
|
@ -0,0 +1,128 @@
|
|||
---
|
||||
slug: video_characters_api
|
||||
title: "New Video Characters, Edit and Extension API support"
|
||||
date: 2026-03-16T10:00:00
|
||||
authors:
|
||||
- name: Sameer Kankute
|
||||
title: SWE @ LiteLLM
|
||||
url: https://www.linkedin.com/in/sameer-kankute/
|
||||
image_url: https://pbs.twimg.com/profile_images/2001352686994907136/ONgNuSk5_400x400.jpg
|
||||
- name: Krrish Dholakia
|
||||
title: "CEO, LiteLLM"
|
||||
url: https://www.linkedin.com/in/krish-d/
|
||||
image_url: https://pbs.twimg.com/profile_images/1298587542745358340/DZv3Oj-h_400x400.jpg
|
||||
- name: Ishaan Jaff
|
||||
title: "CTO, LiteLLM"
|
||||
url: https://www.linkedin.com/in/reffajnaahsi/
|
||||
image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg
|
||||
description: "LiteLLM now supports creating, retrieving, and managing reusable video characters across multiple video generations."
|
||||
tags: [videos, characters, proxy, routing]
|
||||
hide_table_of_contents: false
|
||||
---
|
||||
|
||||
LiteLLM now supoports videos character, edit and extension apis.
|
||||
|
||||
## What's New
|
||||
|
||||
Four new endpoints for video character operations:
|
||||
- **Create character** - Upload a video to create a reusable asset
|
||||
- **Get character** - Retrieve character metadata
|
||||
- **Edit video** - Modify generated videos
|
||||
- **Extend video** - Continue clips with character consistency
|
||||
|
||||
**Available from:** LiteLLM v1.83.0+
|
||||
|
||||
## Quick Example
|
||||
|
||||
```python
|
||||
import litellm
|
||||
|
||||
# Create character from video
|
||||
character = litellm.avideo_create_character(
|
||||
name="Luna",
|
||||
video=open("luna.mp4", "rb"),
|
||||
custom_llm_provider="openai",
|
||||
model="sora-2"
|
||||
)
|
||||
print(f"Character: {character.id}")
|
||||
|
||||
# Use in generation
|
||||
video = litellm.avideo(
|
||||
model="sora-2",
|
||||
prompt="Luna dances through a magical forest.",
|
||||
characters=[{"id": character.id}],
|
||||
seconds="8"
|
||||
)
|
||||
|
||||
# Get character info
|
||||
fetched = litellm.avideo_get_character(
|
||||
character_id=character.id,
|
||||
custom_llm_provider="openai"
|
||||
)
|
||||
|
||||
# Edit with character preserved
|
||||
edited = litellm.avideo_edit(
|
||||
video_id=video.id,
|
||||
prompt="Add warm golden lighting"
|
||||
)
|
||||
|
||||
# Extend sequence
|
||||
extended = litellm.avideo_extension(
|
||||
video_id=video.id,
|
||||
prompt="Luna waves goodbye",
|
||||
seconds="5"
|
||||
)
|
||||
```
|
||||
|
||||
## Via Proxy
|
||||
|
||||
```bash
|
||||
# Create character
|
||||
curl -X POST "http://localhost:4000/v1/videos/characters" \
|
||||
-H "Authorization: Bearer sk-litellm-key" \
|
||||
-F "video=@luna.mp4" \
|
||||
-F "name=Luna"
|
||||
|
||||
# Get character
|
||||
curl -X GET "http://localhost:4000/v1/videos/characters/char_abc123def456" \
|
||||
-H "Authorization: Bearer sk-litellm-key"
|
||||
|
||||
# Edit video
|
||||
curl -X POST "http://localhost:4000/v1/videos/edits" \
|
||||
-H "Authorization: Bearer sk-litellm-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"video": {"id": "video_xyz789"},
|
||||
"prompt": "Add warm golden lighting and enhance colors"
|
||||
}'
|
||||
|
||||
# Extend video
|
||||
curl -X POST "http://localhost:4000/v1/videos/extensions" \
|
||||
-H "Authorization: Bearer sk-litellm-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"video": {"id": "video_xyz789"},
|
||||
"prompt": "Luna waves goodbye and walks into the sunset",
|
||||
"seconds": "5"
|
||||
}'
|
||||
```
|
||||
|
||||
## Managed Character IDs
|
||||
|
||||
LiteLLM automatically encodes provider and model metadata into character IDs:
|
||||
|
||||
**What happens:**
|
||||
```
|
||||
Upload character "Luna" with model "sora-2" on OpenAI
|
||||
↓
|
||||
LiteLLM creates: char_abc123def456 (contains provider + model_id)
|
||||
↓
|
||||
When you reference it later, LiteLLM decodes automatically
|
||||
↓
|
||||
Router knows exactly which deployment to use
|
||||
```
|
||||
|
||||
**Behind the scenes:**
|
||||
- Character ID format: `character_<base64_encoded_metadata>`
|
||||
- Metadata includes: provider, model_id, original_character_id
|
||||
- Transparent to you - just use the ID, LiteLLM handles routing
|
||||
|
|
@ -135,6 +135,81 @@ curl --location --request POST 'http://localhost:4000/v1/videos/video_id/remix'
|
|||
}'
|
||||
```
|
||||
|
||||
### Character, Edit, and Extension Routes
|
||||
|
||||
OpenAI video routes supported by LiteLLM proxy:
|
||||
|
||||
- `POST /v1/videos/characters`
|
||||
- `GET /v1/videos/characters/{character_id}`
|
||||
- `POST /v1/videos/edits`
|
||||
- `POST /v1/videos/extensions`
|
||||
|
||||
#### `target_model_names` support on character creation
|
||||
|
||||
`POST /v1/videos/characters` supports `target_model_names` for model-based routing (same behavior as video create).
|
||||
|
||||
```bash
|
||||
curl --location 'http://localhost:4000/v1/videos/characters' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
-F 'name=hero' \
|
||||
-F 'target_model_names=gpt-4' \
|
||||
-F 'video=@/path/to/character.mp4'
|
||||
```
|
||||
|
||||
When `target_model_names` is used, LiteLLM returns an encoded character ID:
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "character_...",
|
||||
"object": "character",
|
||||
"created_at": 1712697600,
|
||||
"name": "hero"
|
||||
}
|
||||
```
|
||||
|
||||
Use that encoded ID directly on get:
|
||||
|
||||
```bash
|
||||
curl --location 'http://localhost:4000/v1/videos/characters/character_...' \
|
||||
--header 'Authorization: Bearer sk-1234'
|
||||
```
|
||||
|
||||
#### Encoded and non-encoded video IDs for edit/extension
|
||||
|
||||
Both routes accept either plain or encoded `video.id`:
|
||||
|
||||
- `POST /v1/videos/edits`
|
||||
- `POST /v1/videos/extensions`
|
||||
|
||||
```bash
|
||||
curl --location 'http://localhost:4000/v1/videos/edits' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"prompt": "Make this brighter",
|
||||
"video": { "id": "video_..." }
|
||||
}'
|
||||
```
|
||||
|
||||
```bash
|
||||
curl --location 'http://localhost:4000/v1/videos/extensions' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"prompt": "Continue this scene",
|
||||
"seconds": "4",
|
||||
"video": { "id": "video_..." }
|
||||
}'
|
||||
```
|
||||
|
||||
#### `custom_llm_provider` input sources
|
||||
|
||||
For these routes, `custom_llm_provider` may be supplied via:
|
||||
|
||||
- header: `custom-llm-provider`
|
||||
- query: `?custom_llm_provider=...`
|
||||
- body: `custom_llm_provider` (and `extra_body.custom_llm_provider` where supported)
|
||||
|
||||
Test OpenAI video generation request
|
||||
|
||||
```bash
|
||||
|
|
|
|||
|
|
@ -290,6 +290,82 @@ curl --location 'http://localhost:4000/v1/videos' \
|
|||
--header 'custom-llm-provider: azure'
|
||||
```
|
||||
|
||||
### Character, Edit, and Extension Endpoints
|
||||
|
||||
LiteLLM proxy also supports these OpenAI-compatible video routes:
|
||||
|
||||
- `POST /v1/videos/characters`
|
||||
- `GET /v1/videos/characters/{character_id}`
|
||||
- `POST /v1/videos/edits`
|
||||
- `POST /v1/videos/extensions`
|
||||
|
||||
#### Routing Behavior (`target_model_names`, encoded IDs, and provider overrides)
|
||||
|
||||
- `POST /v1/videos/characters` supports `target_model_names` like `POST /v1/videos`.
|
||||
- When `target_model_names` is provided on character creation, LiteLLM encodes the returned `character_id` with routing metadata.
|
||||
- `GET /v1/videos/characters/{character_id}` accepts encoded character IDs directly. LiteLLM decodes the ID internally and routes with the correct model/provider metadata.
|
||||
- `POST /v1/videos/edits` and `POST /v1/videos/extensions` support both:
|
||||
- plain `video.id`
|
||||
- encoded `video.id` values returned by LiteLLM
|
||||
- `custom_llm_provider` can be supplied using the same patterns as other proxy endpoints:
|
||||
- header: `custom-llm-provider`
|
||||
- query: `?custom_llm_provider=...`
|
||||
- body: `custom_llm_provider` (or `extra_body.custom_llm_provider` where applicable)
|
||||
|
||||
#### Character create with `target_model_names`
|
||||
|
||||
```bash
|
||||
curl --location 'http://localhost:4000/v1/videos/characters' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
-F 'name=hero' \
|
||||
-F 'target_model_names=gpt-4' \
|
||||
-F 'video=@/path/to/character.mp4'
|
||||
```
|
||||
|
||||
Example response (encoded `id`):
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "character_...",
|
||||
"object": "character",
|
||||
"created_at": 1712697600,
|
||||
"name": "hero"
|
||||
}
|
||||
```
|
||||
|
||||
#### Get character using encoded `character_id`
|
||||
|
||||
```bash
|
||||
curl --location 'http://localhost:4000/v1/videos/characters/character_...' \
|
||||
--header 'Authorization: Bearer sk-1234'
|
||||
```
|
||||
|
||||
#### Video edit with encoded `video.id`
|
||||
|
||||
```bash
|
||||
curl --location 'http://localhost:4000/v1/videos/edits' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"prompt": "Make this brighter",
|
||||
"video": { "id": "video_..." }
|
||||
}'
|
||||
```
|
||||
|
||||
#### Video extension with provider override from `extra_body`
|
||||
|
||||
```bash
|
||||
curl --location 'http://localhost:4000/v1/videos/extensions' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"prompt": "Continue this scene",
|
||||
"seconds": "4",
|
||||
"video": { "id": "video_..." },
|
||||
"extra_body": { "custom_llm_provider": "openai" }
|
||||
}'
|
||||
```
|
||||
|
||||
Test Azure video generation request
|
||||
|
||||
```bash
|
||||
|
|
|
|||
|
|
@ -131,6 +131,10 @@ class CheckBatchCost:
|
|||
# every subsequent poll cycle.
|
||||
if self._has_batch_processed_column:
|
||||
try:
|
||||
# Include "complete"/"completed" batches: the retrieve_batch
|
||||
# endpoint may transition a batch to "complete" before
|
||||
# CheckBatchCost runs. The batch_processed=False filter
|
||||
# already prevents reprocessing finished batches.
|
||||
jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many(
|
||||
where={
|
||||
"file_purpose": "batch",
|
||||
|
|
@ -140,8 +144,6 @@ class CheckBatchCost:
|
|||
"failed",
|
||||
"expired",
|
||||
"cancelled",
|
||||
"complete",
|
||||
"completed",
|
||||
"stale_expired",
|
||||
]
|
||||
},
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
|||
get_batch_id_from_unified_batch_id,
|
||||
get_content_type_from_file_object,
|
||||
get_model_id_from_unified_batch_id,
|
||||
get_models_from_unified_file_id,
|
||||
normalize_mime_type_for_provider,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
|
|
@ -904,6 +905,21 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
) # managed batch id
|
||||
model_id = cast(Optional[str], response._hidden_params.get("model_id"))
|
||||
model_name = cast(Optional[str], response._hidden_params.get("model_name"))
|
||||
resolved_model_name = model_name
|
||||
|
||||
# Some providers (e.g. Vertex batch retrieve) do not set model_name on
|
||||
# the response. In that case, recover target_model_names from the input
|
||||
# managed file metadata so unified output IDs preserve routing metadata.
|
||||
if not resolved_model_name and isinstance(unified_file_id, str):
|
||||
decoded_unified_file_id = (
|
||||
_is_base64_encoded_unified_file_id(unified_file_id)
|
||||
or unified_file_id
|
||||
)
|
||||
target_model_names = get_models_from_unified_file_id(
|
||||
decoded_unified_file_id
|
||||
)
|
||||
if target_model_names:
|
||||
resolved_model_name = ",".join(target_model_names)
|
||||
original_response_id = response.id
|
||||
|
||||
if (unified_batch_id or unified_file_id) and model_id:
|
||||
|
|
@ -919,7 +935,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
unified_file_id = self.get_unified_output_file_id(
|
||||
output_file_id=original_file_id,
|
||||
model_id=model_id,
|
||||
model_name=model_name,
|
||||
model_name=resolved_model_name,
|
||||
)
|
||||
setattr(response, file_attr, unified_file_id)
|
||||
|
||||
|
|
|
|||
|
|
@ -367,6 +367,42 @@ def update_headers_with_filtered_beta(
|
|||
return headers
|
||||
|
||||
|
||||
def update_request_with_filtered_beta(
|
||||
headers: dict,
|
||||
request_data: dict,
|
||||
provider: str,
|
||||
) -> tuple[dict, dict]:
|
||||
"""
|
||||
Update both headers and request body beta fields based on provider support.
|
||||
Modifies both dicts in place and returns them.
|
||||
|
||||
Args:
|
||||
headers: Request headers dict (will be modified in place)
|
||||
request_data: Request body dict (will be modified in place)
|
||||
provider: Provider name
|
||||
|
||||
Returns:
|
||||
Tuple of (updated headers, updated request_data)
|
||||
"""
|
||||
headers = update_headers_with_filtered_beta(headers=headers, provider=provider)
|
||||
|
||||
existing_body_betas = request_data.get("anthropic_beta")
|
||||
if not existing_body_betas:
|
||||
return headers, request_data
|
||||
|
||||
filtered_body_betas = filter_and_transform_beta_headers(
|
||||
beta_headers=existing_body_betas,
|
||||
provider=provider,
|
||||
)
|
||||
|
||||
if filtered_body_betas:
|
||||
request_data["anthropic_beta"] = filtered_body_betas
|
||||
else:
|
||||
request_data.pop("anthropic_beta", None)
|
||||
|
||||
return headers, request_data
|
||||
|
||||
|
||||
def get_unsupported_headers(provider: str) -> List[str]:
|
||||
"""
|
||||
Get all beta headers that are unsupported by a provider (have null values in mapping).
|
||||
|
|
|
|||
|
|
@ -199,7 +199,8 @@ def create_batch( # noqa: PLR0915
|
|||
)
|
||||
### TIMEOUT LOGIC ###
|
||||
timeout = _resolve_timeout(optional_params, kwargs, custom_llm_provider)
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
user=None,
|
||||
optional_params=optional_params.model_dump(),
|
||||
|
|
@ -207,7 +208,6 @@ def create_batch( # noqa: PLR0915
|
|||
"litellm_call_id": litellm_call_id,
|
||||
"proxy_server_request": proxy_server_request,
|
||||
"model_info": model_info,
|
||||
"metadata": metadata,
|
||||
"preset_cache_key": None,
|
||||
"stream_response": {},
|
||||
**optional_params.model_dump(exclude_unset=True),
|
||||
|
|
@ -584,7 +584,8 @@ def retrieve_batch(
|
|||
**kwargs,
|
||||
)
|
||||
if litellm_logging_obj is not None:
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
user=None,
|
||||
optional_params=optional_params.model_dump(),
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
{
|
||||
"posts": [
|
||||
{
|
||||
"title": "Incident Report: SERVER_ROOT_PATH regression broke UI routing",
|
||||
"description": "How a single line removal caused UI 404s for all deployments using SERVER_ROOT_PATH, and the tests we added to prevent it from happening again.",
|
||||
"date": "2026-02-21",
|
||||
"url": "https://docs.litellm.ai/blog/server-root-path-incident"
|
||||
"title": "Realtime WebRTC HTTP Endpoints",
|
||||
"description": "Use the LiteLLM proxy to route OpenAI-style WebRTC realtime via HTTP: client_secrets and SDP exchange.",
|
||||
"date": "2026-03-12",
|
||||
"url": "https://docs.litellm.ai/blog/realtime_webrtc_http_endpoints"
|
||||
}
|
||||
]
|
||||
}
|
||||
|
|
|
|||
|
|
@ -91,7 +91,8 @@ def create_sync_endpoint_function(endpoint_config: Dict) -> Callable:
|
|||
optional_params = {k: kwargs.get(k) for k in path_params if k in kwargs}
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model="",
|
||||
optional_params=optional_params,
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
|
|
|
|||
|
|
@ -233,7 +233,8 @@ def create_container(
|
|||
)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model="",
|
||||
optional_params=dict(container_create_request_params),
|
||||
litellm_params={
|
||||
|
|
@ -438,7 +439,8 @@ def list_containers(
|
|||
)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model="",
|
||||
optional_params=dict(container_list_optional_params),
|
||||
litellm_params={
|
||||
|
|
@ -626,7 +628,8 @@ def retrieve_container(
|
|||
)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model="",
|
||||
optional_params={},
|
||||
litellm_params={
|
||||
|
|
@ -811,7 +814,8 @@ def delete_container(
|
|||
)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model="",
|
||||
optional_params={},
|
||||
litellm_params={
|
||||
|
|
@ -1010,7 +1014,8 @@ def list_container_files(
|
|||
)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model="",
|
||||
optional_params={
|
||||
"container_id": container_id,
|
||||
|
|
@ -1255,7 +1260,8 @@ def upload_container_file(
|
|||
)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model="",
|
||||
optional_params={"container_id": container_id},
|
||||
litellm_params={
|
||||
|
|
|
|||
|
|
@ -193,7 +193,8 @@ def create_eval(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params=request_body,
|
||||
litellm_params={
|
||||
|
|
@ -382,7 +383,8 @@ def list_evals(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params=query_params,
|
||||
litellm_params={
|
||||
|
|
@ -536,7 +538,8 @@ def get_eval(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"eval_id": eval_id},
|
||||
litellm_params={
|
||||
|
|
@ -760,7 +763,8 @@ def update_eval(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params=request_body,
|
||||
litellm_params={
|
||||
|
|
@ -914,7 +918,8 @@ def delete_eval(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"eval_id": eval_id},
|
||||
litellm_params={
|
||||
|
|
@ -1071,7 +1076,8 @@ def cancel_eval(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"eval_id": eval_id},
|
||||
litellm_params={
|
||||
|
|
@ -1262,7 +1268,8 @@ def create_run(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params=request_body,
|
||||
litellm_params={
|
||||
|
|
@ -1450,7 +1457,8 @@ def list_runs(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"eval_id": eval_id, **query_params},
|
||||
litellm_params={
|
||||
|
|
@ -1610,7 +1618,8 @@ def get_run(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"eval_id": eval_id, "run_id": run_id},
|
||||
litellm_params={
|
||||
|
|
@ -1773,7 +1782,8 @@ def cancel_run(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"eval_id": eval_id, "run_id": run_id},
|
||||
litellm_params={
|
||||
|
|
@ -1941,7 +1951,8 @@ def delete_run(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"eval_id": eval_id, "run_id": run_id},
|
||||
litellm_params={
|
||||
|
|
|
|||
|
|
@ -185,7 +185,8 @@ class GenerateContentHelper:
|
|||
if litellm_logging_obj is None:
|
||||
raise ValueError("litellm_logging_obj is required, but got None")
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
optional_params=dict(generate_content_config_dict),
|
||||
litellm_params={
|
||||
|
|
|
|||
|
|
@ -40,6 +40,9 @@ from litellm.utils import exception_type, get_litellm_params
|
|||
llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler()
|
||||
from openai.types.audio.transcription_create_params import FileTypes # type: ignore
|
||||
|
||||
# BFL handlers
|
||||
from litellm.llms.black_forest_labs.image_edit.handler import bfl_image_edit
|
||||
from litellm.llms.black_forest_labs.image_generation.handler import bfl_image_generation
|
||||
from litellm.main import (
|
||||
azure_chat_completions,
|
||||
base_llm_aiohttp_handler,
|
||||
|
|
@ -50,10 +53,6 @@ from litellm.main import (
|
|||
openai_image_variations,
|
||||
)
|
||||
|
||||
# BFL handlers
|
||||
from litellm.llms.black_forest_labs.image_edit.handler import bfl_image_edit
|
||||
from litellm.llms.black_forest_labs.image_generation.handler import bfl_image_generation
|
||||
|
||||
###########################################
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.images.main import ImageEditOptionalRequestParams
|
||||
|
|
@ -297,7 +296,8 @@ def image_generation( # noqa: PLR0915
|
|||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
|
||||
logging: Logging = litellm_logging_obj
|
||||
logging.update_environment_variables(
|
||||
logging.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
user=user,
|
||||
optional_params=optional_params,
|
||||
|
|
@ -308,7 +308,6 @@ def image_generation( # noqa: PLR0915
|
|||
"logger_fn": logger_fn,
|
||||
"proxy_server_request": proxy_server_request,
|
||||
"model_info": model_info,
|
||||
"metadata": metadata,
|
||||
"preset_cache_key": None,
|
||||
"stream_response": {},
|
||||
},
|
||||
|
|
@ -894,7 +893,8 @@ def image_edit( # noqa: PLR0915
|
|||
)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
user=user,
|
||||
optional_params=dict(image_edit_request_params),
|
||||
|
|
@ -902,7 +902,6 @@ def image_edit( # noqa: PLR0915
|
|||
**image_edit_request_params,
|
||||
"litellm_call_id": litellm_call_id,
|
||||
"model_info": model_info,
|
||||
"metadata": metadata,
|
||||
},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -34,16 +34,7 @@ Usage:
|
|||
import asyncio
|
||||
import contextvars
|
||||
from functools import partial
|
||||
from typing import (
|
||||
Any,
|
||||
AsyncIterator,
|
||||
Coroutine,
|
||||
Dict,
|
||||
Iterator,
|
||||
List,
|
||||
Optional,
|
||||
Union,
|
||||
)
|
||||
from typing import Any, AsyncIterator, Coroutine, Dict, Iterator, List, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -306,7 +297,8 @@ def create(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
optional_params=dict(optional_params),
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
|
|
@ -416,7 +408,8 @@ def get(
|
|||
f"Interactions API not supported for: {custom_llm_provider}"
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"interaction_id": interaction_id},
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
|
|
@ -519,7 +512,8 @@ def delete(
|
|||
f"Interactions API not supported for: {custom_llm_provider}"
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"interaction_id": interaction_id},
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
|
|
@ -622,7 +616,8 @@ def cancel(
|
|||
f"Interactions API not supported for: {custom_llm_provider}"
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"interaction_id": interaction_id},
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
|
|
|
|||
|
|
@ -568,6 +568,42 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if "custom_llm_provider" in self.model_call_details:
|
||||
self.custom_llm_provider = self.model_call_details["custom_llm_provider"]
|
||||
|
||||
def update_from_kwargs(
|
||||
self,
|
||||
kwargs: Dict,
|
||||
litellm_params: Optional[Dict] = None,
|
||||
optional_params: Optional[Dict] = None,
|
||||
model: Optional[str] = None,
|
||||
user: Optional[str] = None,
|
||||
**additional_params,
|
||||
):
|
||||
"""
|
||||
Convenience wrapper around update_environment_variables that
|
||||
automatically extracts metadata/litellm_metadata from kwargs,
|
||||
so callers don't need to manually plumb them into litellm_params.
|
||||
"""
|
||||
base_litellm_params: Dict[str, Any] = {}
|
||||
|
||||
if "metadata" in kwargs:
|
||||
base_litellm_params["metadata"] = kwargs["metadata"]
|
||||
if "litellm_metadata" in kwargs and isinstance(
|
||||
kwargs["litellm_metadata"], dict
|
||||
):
|
||||
base_litellm_params["litellm_metadata"] = kwargs["litellm_metadata"]
|
||||
if "metadata" not in base_litellm_params:
|
||||
base_litellm_params["metadata"] = kwargs["litellm_metadata"].copy()
|
||||
|
||||
if litellm_params:
|
||||
base_litellm_params.update(litellm_params)
|
||||
|
||||
self.update_environment_variables(
|
||||
litellm_params=base_litellm_params,
|
||||
optional_params=optional_params or {},
|
||||
model=model,
|
||||
user=user,
|
||||
**additional_params,
|
||||
)
|
||||
|
||||
def update_messages(self, messages: List[AllMessageValues]):
|
||||
"""
|
||||
Update the logged value of the messages in the model_call_details
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ from litellm.litellm_core_utils.model_response_utils import (
|
|||
)
|
||||
from litellm.litellm_core_utils.redact_messages import LiteLLMLoggingObject
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
from litellm.types.llms.openai import ChatCompletionChunk
|
||||
from litellm.types.llms.openai import OpenAIChatCompletionChunk
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import (
|
||||
Delta,
|
||||
|
|
@ -745,7 +745,7 @@ class CustomStreamWrapper:
|
|||
|
||||
def copy_model_response_level_provider_specific_fields(
|
||||
self,
|
||||
original_chunk: Union[ModelResponseStream, ChatCompletionChunk],
|
||||
original_chunk: Union[ModelResponseStream, OpenAIChatCompletionChunk],
|
||||
model_response: ModelResponseStream,
|
||||
) -> ModelResponseStream:
|
||||
"""
|
||||
|
|
@ -1012,6 +1012,15 @@ class CustomStreamWrapper:
|
|||
# if delta is None
|
||||
_is_delta_empty = self.is_delta_empty(delta=model_response.choices[0].delta)
|
||||
|
||||
# Preserve custom attributes from original chunk (applies to both
|
||||
# empty and non-empty delta final chunks).
|
||||
_original_chunk = response_obj.get("original_chunk", None)
|
||||
if _original_chunk is not None:
|
||||
preserve_upstream_non_openai_attributes(
|
||||
model_response=model_response,
|
||||
original_chunk=_original_chunk,
|
||||
)
|
||||
|
||||
if _is_delta_empty:
|
||||
model_response.choices[0].delta = Delta(
|
||||
content=None
|
||||
|
|
|
|||
|
|
@ -23,6 +23,9 @@ import litellm
|
|||
import litellm.litellm_core_utils
|
||||
import litellm.types
|
||||
import litellm.types.utils
|
||||
from litellm.anthropic_beta_headers_manager import (
|
||||
update_request_with_filtered_beta,
|
||||
)
|
||||
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
|
||||
from litellm.litellm_core_utils.core_helpers import map_finish_reason
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
|
|
@ -58,9 +61,6 @@ from litellm.types.utils import (
|
|||
|
||||
from ...base import BaseLLM
|
||||
from ..common_utils import AnthropicError, process_anthropic_headers
|
||||
from litellm.anthropic_beta_headers_manager import (
|
||||
update_headers_with_filtered_beta,
|
||||
)
|
||||
from .transformation import AnthropicConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -339,10 +339,6 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
headers = update_headers_with_filtered_beta(
|
||||
headers=headers, provider=custom_llm_provider
|
||||
)
|
||||
|
||||
config = ProviderConfigManager.get_provider_chat_config(
|
||||
model=model,
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
|
|
@ -360,6 +356,12 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
headers=headers,
|
||||
)
|
||||
|
||||
headers, data = update_request_with_filtered_beta(
|
||||
headers=headers,
|
||||
request_data=data,
|
||||
provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from litellm.types.videos.main import VideoCreateOptionalRequestParams
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
from litellm.types.videos.main import CharacterObject as _CharacterObject
|
||||
from litellm.types.videos.main import VideoObject as _VideoObject
|
||||
|
||||
from ..chat.transformation import BaseLLMException as _BaseLLMException
|
||||
|
|
@ -18,10 +19,12 @@ if TYPE_CHECKING:
|
|||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
BaseLLMException = _BaseLLMException
|
||||
VideoObject = _VideoObject
|
||||
CharacterObject = _CharacterObject
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
BaseLLMException = Any
|
||||
VideoObject = Any
|
||||
CharacterObject = Any
|
||||
|
||||
|
||||
class BaseVideoConfig(ABC):
|
||||
|
|
@ -265,6 +268,118 @@ class BaseVideoConfig(ABC):
|
|||
) -> VideoObject:
|
||||
pass
|
||||
|
||||
def transform_video_create_character_request(
|
||||
self,
|
||||
name: str,
|
||||
video: Any,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, list]:
|
||||
"""
|
||||
Transform the video create character request into a URL and files list (multipart).
|
||||
|
||||
Returns:
|
||||
Tuple[str, list]: (url, files_list) for the multipart POST request
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"video create character is not supported for this provider"
|
||||
)
|
||||
|
||||
def transform_video_create_character_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> CharacterObject:
|
||||
raise NotImplementedError(
|
||||
"video create character is not supported for this provider"
|
||||
)
|
||||
|
||||
def transform_video_get_character_request(
|
||||
self,
|
||||
character_id: str,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
"""
|
||||
Transform the video get character request into a URL and params.
|
||||
|
||||
Returns:
|
||||
Tuple[str, Dict]: (url, params) for the GET request
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"video get character is not supported for this provider"
|
||||
)
|
||||
|
||||
def transform_video_get_character_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> CharacterObject:
|
||||
raise NotImplementedError(
|
||||
"video get character is not supported for this provider"
|
||||
)
|
||||
|
||||
def transform_video_edit_request(
|
||||
self,
|
||||
prompt: str,
|
||||
video_id: str,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
"""
|
||||
Transform the video edit request into a URL and JSON data.
|
||||
|
||||
Returns:
|
||||
Tuple[str, Dict]: (url, data) for the POST request
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"video edit is not supported for this provider"
|
||||
)
|
||||
|
||||
def transform_video_edit_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
) -> VideoObject:
|
||||
raise NotImplementedError(
|
||||
"video edit is not supported for this provider"
|
||||
)
|
||||
|
||||
def transform_video_extension_request(
|
||||
self,
|
||||
prompt: str,
|
||||
video_id: str,
|
||||
seconds: str,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
"""
|
||||
Transform the video extension request into a URL and JSON data.
|
||||
|
||||
Returns:
|
||||
Tuple[str, Dict]: (url, data) for the POST request
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"video extension is not supported for this provider"
|
||||
)
|
||||
|
||||
def transform_video_extension_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
) -> VideoObject:
|
||||
raise NotImplementedError(
|
||||
"video extension is not supported for this provider"
|
||||
)
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
|
|
|
|||
|
|
@ -1882,12 +1882,11 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers, provider=custom_llm_provider
|
||||
)
|
||||
|
||||
logging_obj.update_environment_variables(
|
||||
logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
optional_params=dict(anthropic_messages_optional_request_params),
|
||||
litellm_params={
|
||||
"metadata": kwargs.get("metadata", {}),
|
||||
"litellm_metadata": kwargs.get("litellm_metadata", {}),
|
||||
"preset_cache_key": None,
|
||||
"stream_response": {},
|
||||
**anthropic_messages_optional_request_params,
|
||||
|
|
@ -6114,6 +6113,614 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=video_remix_provider_config,
|
||||
)
|
||||
|
||||
def video_create_character_handler(
|
||||
self,
|
||||
name: str,
|
||||
video: Any,
|
||||
video_provider_config: BaseVideoConfig,
|
||||
custom_llm_provider: str,
|
||||
litellm_params,
|
||||
logging_obj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[float] = None,
|
||||
_is_async: bool = False,
|
||||
client=None,
|
||||
api_key: Optional[str] = None,
|
||||
):
|
||||
if _is_async:
|
||||
return self.async_video_create_character_handler(
|
||||
name=name,
|
||||
video=video,
|
||||
video_provider_config=video_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
sync_httpx_client = _get_httpx_client(
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
|
||||
)
|
||||
else:
|
||||
sync_httpx_client = client
|
||||
|
||||
headers = video_provider_config.validate_environment(
|
||||
api_key=api_key or litellm_params.get("api_key", None),
|
||||
headers=extra_headers or {},
|
||||
model="",
|
||||
)
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
api_base = video_provider_config.get_complete_url(
|
||||
model="",
|
||||
api_base=litellm_params.get("api_base", None),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
url, files_list = video_provider_config.transform_video_create_character_request(
|
||||
name=name,
|
||||
video=video,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=name,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": {"name": name},
|
||||
"api_base": url,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
response = sync_httpx_client.post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
files=files_list,
|
||||
timeout=timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return video_provider_config.transform_video_create_character_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=video_provider_config)
|
||||
|
||||
async def async_video_create_character_handler(
|
||||
self,
|
||||
name: str,
|
||||
video: Any,
|
||||
video_provider_config: BaseVideoConfig,
|
||||
custom_llm_provider: str,
|
||||
litellm_params,
|
||||
logging_obj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[float] = None,
|
||||
client=None,
|
||||
api_key: Optional[str] = None,
|
||||
):
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
async_httpx_client = get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders(custom_llm_provider),
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
|
||||
)
|
||||
else:
|
||||
async_httpx_client = client
|
||||
|
||||
headers = video_provider_config.validate_environment(
|
||||
api_key=api_key or litellm_params.get("api_key", None),
|
||||
headers=extra_headers or {},
|
||||
model="",
|
||||
)
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
api_base = video_provider_config.get_complete_url(
|
||||
model="",
|
||||
api_base=litellm_params.get("api_base", None),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
url, files_list = video_provider_config.transform_video_create_character_request(
|
||||
name=name,
|
||||
video=video,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=name,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": {"name": name},
|
||||
"api_base": url,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
response = await async_httpx_client.post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
files=files_list,
|
||||
timeout=timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return video_provider_config.transform_video_create_character_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=video_provider_config)
|
||||
|
||||
def video_get_character_handler(
|
||||
self,
|
||||
character_id: str,
|
||||
video_provider_config: BaseVideoConfig,
|
||||
custom_llm_provider: str,
|
||||
litellm_params,
|
||||
logging_obj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[float] = None,
|
||||
_is_async: bool = False,
|
||||
client=None,
|
||||
api_key: Optional[str] = None,
|
||||
):
|
||||
if _is_async:
|
||||
return self.async_video_get_character_handler(
|
||||
character_id=character_id,
|
||||
video_provider_config=video_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
sync_httpx_client = _get_httpx_client(
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
|
||||
)
|
||||
else:
|
||||
sync_httpx_client = client
|
||||
|
||||
headers = video_provider_config.validate_environment(
|
||||
api_key=api_key or litellm_params.get("api_key", None),
|
||||
headers=extra_headers or {},
|
||||
model="",
|
||||
)
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
api_base = video_provider_config.get_complete_url(
|
||||
model="",
|
||||
api_base=litellm_params.get("api_base", None),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
url, params = video_provider_config.transform_video_get_character_request(
|
||||
character_id=character_id,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=character_id,
|
||||
api_key="",
|
||||
additional_args={"api_base": url, "headers": headers},
|
||||
)
|
||||
|
||||
try:
|
||||
response = sync_httpx_client.get(
|
||||
url=url,
|
||||
headers=headers,
|
||||
params=params
|
||||
)
|
||||
response.raise_for_status()
|
||||
return video_provider_config.transform_video_get_character_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=video_provider_config)
|
||||
|
||||
async def async_video_get_character_handler(
|
||||
self,
|
||||
character_id: str,
|
||||
video_provider_config: BaseVideoConfig,
|
||||
custom_llm_provider: str,
|
||||
litellm_params,
|
||||
logging_obj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[float] = None,
|
||||
client=None,
|
||||
api_key: Optional[str] = None,
|
||||
):
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
async_httpx_client = get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders(custom_llm_provider),
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
|
||||
)
|
||||
else:
|
||||
async_httpx_client = client
|
||||
|
||||
headers = video_provider_config.validate_environment(
|
||||
api_key=api_key or litellm_params.get("api_key", None),
|
||||
headers=extra_headers or {},
|
||||
model="",
|
||||
)
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
api_base = video_provider_config.get_complete_url(
|
||||
model="",
|
||||
api_base=litellm_params.get("api_base", None),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
url, params = video_provider_config.transform_video_get_character_request(
|
||||
character_id=character_id,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=character_id,
|
||||
api_key="",
|
||||
additional_args={"api_base": url, "headers": headers},
|
||||
)
|
||||
|
||||
try:
|
||||
response = await async_httpx_client.get(
|
||||
url=url,
|
||||
headers=headers,
|
||||
params=params
|
||||
)
|
||||
response.raise_for_status()
|
||||
return video_provider_config.transform_video_get_character_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=video_provider_config)
|
||||
|
||||
def video_edit_handler(
|
||||
self,
|
||||
prompt: str,
|
||||
video_id: str,
|
||||
video_provider_config: BaseVideoConfig,
|
||||
custom_llm_provider: str,
|
||||
litellm_params,
|
||||
logging_obj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[float] = None,
|
||||
_is_async: bool = False,
|
||||
client=None,
|
||||
api_key: Optional[str] = None,
|
||||
):
|
||||
if _is_async:
|
||||
return self.async_video_edit_handler(
|
||||
prompt=prompt,
|
||||
video_id=video_id,
|
||||
video_provider_config=video_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
sync_httpx_client = _get_httpx_client(
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
|
||||
)
|
||||
else:
|
||||
sync_httpx_client = client
|
||||
|
||||
headers = video_provider_config.validate_environment(
|
||||
api_key=api_key or litellm_params.get("api_key", None),
|
||||
headers=extra_headers or {},
|
||||
model="",
|
||||
)
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
api_base = video_provider_config.get_complete_url(
|
||||
model="",
|
||||
api_base=litellm_params.get("api_base", None),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
url, data = video_provider_config.transform_video_edit_request(
|
||||
prompt=prompt,
|
||||
video_id=video_id,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=prompt,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": url,
|
||||
"headers": headers,
|
||||
"video_id": video_id,
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
response = sync_httpx_client.post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
json=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return video_provider_config.transform_video_edit_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=video_provider_config)
|
||||
|
||||
async def async_video_edit_handler(
|
||||
self,
|
||||
prompt: str,
|
||||
video_id: str,
|
||||
video_provider_config: BaseVideoConfig,
|
||||
custom_llm_provider: str,
|
||||
litellm_params,
|
||||
logging_obj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[float] = None,
|
||||
client=None,
|
||||
api_key: Optional[str] = None,
|
||||
):
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
async_httpx_client = get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders(custom_llm_provider),
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
|
||||
)
|
||||
else:
|
||||
async_httpx_client = client
|
||||
|
||||
headers = video_provider_config.validate_environment(
|
||||
api_key=api_key or litellm_params.get("api_key", None),
|
||||
headers=extra_headers or {},
|
||||
model="",
|
||||
)
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
api_base = video_provider_config.get_complete_url(
|
||||
model="",
|
||||
api_base=litellm_params.get("api_base", None),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
url, data = video_provider_config.transform_video_edit_request(
|
||||
prompt=prompt,
|
||||
video_id=video_id,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=prompt,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": url,
|
||||
"headers": headers,
|
||||
"video_id": video_id,
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
response = await async_httpx_client.post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
json=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return video_provider_config.transform_video_edit_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=video_provider_config)
|
||||
|
||||
def video_extension_handler(
|
||||
self,
|
||||
prompt: str,
|
||||
video_id: str,
|
||||
seconds: str,
|
||||
video_provider_config: BaseVideoConfig,
|
||||
custom_llm_provider: str,
|
||||
litellm_params,
|
||||
logging_obj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[float] = None,
|
||||
_is_async: bool = False,
|
||||
client=None,
|
||||
api_key: Optional[str] = None,
|
||||
):
|
||||
if _is_async:
|
||||
return self.async_video_extension_handler(
|
||||
prompt=prompt,
|
||||
video_id=video_id,
|
||||
seconds=seconds,
|
||||
video_provider_config=video_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
sync_httpx_client = _get_httpx_client(
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
|
||||
)
|
||||
else:
|
||||
sync_httpx_client = client
|
||||
|
||||
headers = video_provider_config.validate_environment(
|
||||
api_key=api_key or litellm_params.get("api_key", None),
|
||||
headers=extra_headers or {},
|
||||
model="",
|
||||
)
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
api_base = video_provider_config.get_complete_url(
|
||||
model="",
|
||||
api_base=litellm_params.get("api_base", None),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
url, data = video_provider_config.transform_video_extension_request(
|
||||
prompt=prompt,
|
||||
video_id=video_id,
|
||||
seconds=seconds,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=prompt,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": url,
|
||||
"headers": headers,
|
||||
"video_id": video_id,
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
response = sync_httpx_client.post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
json=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return video_provider_config.transform_video_extension_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=video_provider_config)
|
||||
|
||||
async def async_video_extension_handler(
|
||||
self,
|
||||
prompt: str,
|
||||
video_id: str,
|
||||
seconds: str,
|
||||
video_provider_config: BaseVideoConfig,
|
||||
custom_llm_provider: str,
|
||||
litellm_params,
|
||||
logging_obj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[float] = None,
|
||||
client=None,
|
||||
api_key: Optional[str] = None,
|
||||
):
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
async_httpx_client = get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders(custom_llm_provider),
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
|
||||
)
|
||||
else:
|
||||
async_httpx_client = client
|
||||
|
||||
headers = video_provider_config.validate_environment(
|
||||
api_key=api_key or litellm_params.get("api_key", None),
|
||||
headers=extra_headers or {},
|
||||
model="",
|
||||
)
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
api_base = video_provider_config.get_complete_url(
|
||||
model="",
|
||||
api_base=litellm_params.get("api_base", None),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
url, data = video_provider_config.transform_video_extension_request(
|
||||
prompt=prompt,
|
||||
video_id=video_id,
|
||||
seconds=seconds,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=prompt,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": url,
|
||||
"headers": headers,
|
||||
"video_id": video_id,
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
response = await async_httpx_client.post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
json=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return video_provider_config.transform_video_extension_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=video_provider_config)
|
||||
|
||||
def video_list_handler(
|
||||
self,
|
||||
after: Optional[str],
|
||||
|
|
|
|||
|
|
@ -1,29 +1,30 @@
|
|||
from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union
|
||||
import base64
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
from httpx._types import RequestFiles
|
||||
|
||||
from litellm.types.videos.main import VideoCreateOptionalRequestParams, VideoObject
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.videos.utils import (
|
||||
encode_video_id_with_provider,
|
||||
extract_original_video_id,
|
||||
)
|
||||
from litellm.images.utils import ImageEditRequestUtils
|
||||
import litellm
|
||||
from litellm.constants import DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS
|
||||
from litellm.images.utils import ImageEditRequestUtils
|
||||
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.gemini import (
|
||||
GeminiLongRunningOperationResponse,
|
||||
GeminiVideoGenerationInstance,
|
||||
GeminiVideoGenerationParameters,
|
||||
GeminiVideoGenerationRequest,
|
||||
)
|
||||
from litellm.constants import DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS
|
||||
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.videos.main import VideoCreateOptionalRequestParams, VideoObject
|
||||
from litellm.types.videos.utils import (
|
||||
encode_video_id_with_provider,
|
||||
extract_original_video_id,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
from ...base_llm.chat.transformation import BaseLLMException as _BaseLLMException
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -524,6 +525,30 @@ class GeminiVideoConfig(BaseVideoConfig):
|
|||
"""Video delete is not supported."""
|
||||
raise NotImplementedError("Video delete is not supported by Google Veo.")
|
||||
|
||||
def transform_video_create_character_request(self, name, video, api_base, litellm_params, headers):
|
||||
raise NotImplementedError("video create character is not supported for Gemini")
|
||||
|
||||
def transform_video_create_character_response(self, raw_response, logging_obj):
|
||||
raise NotImplementedError("video create character is not supported for Gemini")
|
||||
|
||||
def transform_video_get_character_request(self, character_id, api_base, litellm_params, headers):
|
||||
raise NotImplementedError("video get character is not supported for Gemini")
|
||||
|
||||
def transform_video_get_character_response(self, raw_response, logging_obj):
|
||||
raise NotImplementedError("video get character is not supported for Gemini")
|
||||
|
||||
def transform_video_edit_request(self, prompt, video_id, api_base, litellm_params, headers, extra_body=None):
|
||||
raise NotImplementedError("video edit is not supported for Gemini")
|
||||
|
||||
def transform_video_edit_response(self, raw_response, logging_obj, custom_llm_provider=None):
|
||||
raise NotImplementedError("video edit is not supported for Gemini")
|
||||
|
||||
def transform_video_extension_request(self, prompt, video_id, seconds, api_base, litellm_params, headers, extra_body=None):
|
||||
raise NotImplementedError("video extension is not supported for Gemini")
|
||||
|
||||
def transform_video_extension_response(self, raw_response, logging_obj, custom_llm_provider=None):
|
||||
raise NotImplementedError("video extension is not supported for Gemini")
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
|
|
|
|||
|
|
@ -69,7 +69,8 @@ class LiteLLMSkillsTransformationHandler:
|
|||
"""
|
||||
# Pre-call logging
|
||||
if logging_obj:
|
||||
logging_obj.update_environment_variables(
|
||||
logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"display_title": display_title},
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
|
|
@ -172,7 +173,8 @@ class LiteLLMSkillsTransformationHandler:
|
|||
"""
|
||||
# Pre-call logging
|
||||
if logging_obj:
|
||||
logging_obj.update_environment_variables(
|
||||
logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"limit": limit, "offset": offset},
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
|
|
@ -231,7 +233,8 @@ class LiteLLMSkillsTransformationHandler:
|
|||
"""
|
||||
# Pre-call logging
|
||||
if logging_obj:
|
||||
logging_obj.update_environment_variables(
|
||||
logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"skill_id": skill_id},
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
|
|
@ -277,7 +280,8 @@ class LiteLLMSkillsTransformationHandler:
|
|||
"""
|
||||
# Pre-call logging
|
||||
if logging_obj:
|
||||
logging_obj.update_environment_variables(
|
||||
logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"skill_id": skill_id},
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from io import BufferedReader
|
||||
import mimetypes
|
||||
from io import BufferedReader, BytesIO
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
import httpx
|
||||
|
|
@ -10,9 +11,14 @@ from litellm.llms.openai.image_edit.transformation import ImageEditRequestUtils
|
|||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import CreateVideoRequest
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.videos.main import VideoCreateOptionalRequestParams, VideoObject
|
||||
from litellm.types.videos.main import (
|
||||
CharacterObject,
|
||||
VideoCreateOptionalRequestParams,
|
||||
VideoObject,
|
||||
)
|
||||
from litellm.types.videos.utils import (
|
||||
encode_video_id_with_provider,
|
||||
extract_original_character_id,
|
||||
extract_original_video_id,
|
||||
)
|
||||
|
||||
|
|
@ -46,6 +52,7 @@ class OpenAIVideoConfig(BaseVideoConfig):
|
|||
"input_reference",
|
||||
"seconds",
|
||||
"size",
|
||||
"characters",
|
||||
"user",
|
||||
"extra_headers",
|
||||
]
|
||||
|
|
@ -121,6 +128,7 @@ class OpenAIVideoConfig(BaseVideoConfig):
|
|||
model=model, prompt=prompt, **video_create_optional_request_params
|
||||
)
|
||||
request_dict = cast(Dict, video_create_request)
|
||||
request_dict = self._decode_character_ids_in_create_video_request(request_dict)
|
||||
|
||||
# Handle input_reference parameter if provided
|
||||
_input_reference = video_create_optional_request_params.get("input_reference")
|
||||
|
|
@ -138,6 +146,35 @@ class OpenAIVideoConfig(BaseVideoConfig):
|
|||
)
|
||||
return data_without_files, files_list, api_base
|
||||
|
||||
def _decode_character_ids_in_create_video_request(self, request_dict: Dict) -> Dict:
|
||||
"""
|
||||
Decode LiteLLM-managed encoded character ids for provider requests.
|
||||
|
||||
OpenAI expects character ids like `char_...`. If a caller sends
|
||||
`character_<base64-encoded-provider-payload>`, convert it back to the
|
||||
original provider id before forwarding upstream.
|
||||
"""
|
||||
raw_characters = request_dict.get("characters")
|
||||
if not isinstance(raw_characters, list):
|
||||
return request_dict
|
||||
|
||||
decoded_characters: List[Any] = []
|
||||
for character in raw_characters:
|
||||
if not isinstance(character, dict):
|
||||
decoded_characters.append(character)
|
||||
continue
|
||||
|
||||
character_id = character.get("id")
|
||||
if isinstance(character_id, str):
|
||||
decoded_character = dict(character)
|
||||
decoded_character["id"] = extract_original_character_id(character_id)
|
||||
decoded_characters.append(decoded_character)
|
||||
else:
|
||||
decoded_characters.append(character)
|
||||
|
||||
request_dict["characters"] = decoded_characters
|
||||
return request_dict
|
||||
|
||||
def transform_video_create_response(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -430,6 +467,106 @@ class OpenAIVideoConfig(BaseVideoConfig):
|
|||
headers=headers,
|
||||
)
|
||||
|
||||
def transform_video_create_character_request(
|
||||
self,
|
||||
name: str,
|
||||
video: Any,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, list]:
|
||||
url = f"{api_base.rstrip('/')}/characters"
|
||||
files_list: List[Tuple[str, Any]] = [("name", (None, name))]
|
||||
self._add_video_to_files(files_list, video, "video")
|
||||
return url, files_list
|
||||
|
||||
def transform_video_create_character_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: Any,
|
||||
) -> CharacterObject:
|
||||
return CharacterObject(**raw_response.json())
|
||||
|
||||
def transform_video_get_character_request(
|
||||
self,
|
||||
character_id: str,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
url = f"{api_base.rstrip('/')}/characters/{character_id}"
|
||||
return url, {}
|
||||
|
||||
def transform_video_get_character_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: Any,
|
||||
) -> CharacterObject:
|
||||
return CharacterObject(**raw_response.json())
|
||||
|
||||
def transform_video_edit_request(
|
||||
self,
|
||||
prompt: str,
|
||||
video_id: str,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
original_video_id = extract_original_video_id(video_id)
|
||||
url = f"{api_base.rstrip('/')}/edits"
|
||||
data: Dict[str, Any] = {"prompt": prompt, "video": {"id": original_video_id}}
|
||||
if extra_body:
|
||||
data.update(extra_body)
|
||||
return url, data
|
||||
|
||||
def transform_video_edit_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: Any,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
) -> VideoObject:
|
||||
video_obj = VideoObject(**raw_response.json())
|
||||
if custom_llm_provider and video_obj.id:
|
||||
video_obj.id = encode_video_id_with_provider(
|
||||
video_obj.id, custom_llm_provider, None
|
||||
)
|
||||
return video_obj
|
||||
|
||||
def transform_video_extension_request(
|
||||
self,
|
||||
prompt: str,
|
||||
video_id: str,
|
||||
seconds: str,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
original_video_id = extract_original_video_id(video_id)
|
||||
url = f"{api_base.rstrip('/')}/extensions"
|
||||
data: Dict[str, Any] = {
|
||||
"prompt": prompt,
|
||||
"seconds": seconds,
|
||||
"video": {"id": original_video_id},
|
||||
}
|
||||
if extra_body:
|
||||
data.update(extra_body)
|
||||
return url, data
|
||||
|
||||
def transform_video_extension_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: Any,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
) -> VideoObject:
|
||||
video_obj = VideoObject(**raw_response.json())
|
||||
if custom_llm_provider and video_obj.id:
|
||||
video_obj.id = encode_video_id_with_provider(
|
||||
video_obj.id, custom_llm_provider, None
|
||||
)
|
||||
return video_obj
|
||||
|
||||
def _add_image_to_files(
|
||||
self,
|
||||
files_list: List[Tuple[str, Any]],
|
||||
|
|
@ -445,3 +582,49 @@ class OpenAIVideoConfig(BaseVideoConfig):
|
|||
files_list.append(
|
||||
(field_name, ("input_reference.png", image, image_content_type))
|
||||
)
|
||||
|
||||
def _add_video_to_files(
|
||||
self,
|
||||
files_list: List[Tuple[str, Any]],
|
||||
video: Any,
|
||||
field_name: str,
|
||||
) -> None:
|
||||
"""
|
||||
Add a video to files with proper video MIME type detection.
|
||||
|
||||
This path is used by POST /videos/characters and must send video/mp4,
|
||||
not image/* content types.
|
||||
"""
|
||||
filename = getattr(video, "name", None) or "input_video.mp4"
|
||||
content_type = self._get_video_content_type(video=video, filename=filename)
|
||||
files_list.append((field_name, (filename, video, content_type)))
|
||||
|
||||
def _get_video_content_type(self, video: Any, filename: str) -> str:
|
||||
guessed_content_type, _ = mimetypes.guess_type(filename)
|
||||
if guessed_content_type and guessed_content_type.startswith("video/"):
|
||||
return guessed_content_type
|
||||
|
||||
# Fast-path detection for common MP4 signatures when filename is missing/incorrect.
|
||||
try:
|
||||
header_bytes = b""
|
||||
if isinstance(video, BytesIO):
|
||||
current_pos = video.tell()
|
||||
video.seek(0)
|
||||
header_bytes = video.read(64)
|
||||
video.seek(current_pos)
|
||||
elif isinstance(video, BufferedReader):
|
||||
current_pos = video.tell()
|
||||
video.seek(0)
|
||||
header_bytes = video.read(64)
|
||||
video.seek(current_pos)
|
||||
elif isinstance(video, bytes):
|
||||
header_bytes = video[:64]
|
||||
|
||||
# MP4 typically includes ftyp in first box.
|
||||
if b"ftyp" in header_bytes:
|
||||
return "video/mp4"
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# OpenAI create-character currently supports mp4.
|
||||
return "video/mp4"
|
||||
|
|
|
|||
|
|
@ -592,6 +592,30 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
|
||||
return video_obj
|
||||
|
||||
def transform_video_create_character_request(self, name, video, api_base, litellm_params, headers):
|
||||
raise NotImplementedError("video create character is not supported for RunwayML")
|
||||
|
||||
def transform_video_create_character_response(self, raw_response, logging_obj):
|
||||
raise NotImplementedError("video create character is not supported for RunwayML")
|
||||
|
||||
def transform_video_get_character_request(self, character_id, api_base, litellm_params, headers):
|
||||
raise NotImplementedError("video get character is not supported for RunwayML")
|
||||
|
||||
def transform_video_get_character_response(self, raw_response, logging_obj):
|
||||
raise NotImplementedError("video get character is not supported for RunwayML")
|
||||
|
||||
def transform_video_edit_request(self, prompt, video_id, api_base, litellm_params, headers, extra_body=None):
|
||||
raise NotImplementedError("video edit is not supported for RunwayML")
|
||||
|
||||
def transform_video_edit_response(self, raw_response, logging_obj, custom_llm_provider=None):
|
||||
raise NotImplementedError("video edit is not supported for RunwayML")
|
||||
|
||||
def transform_video_extension_request(self, prompt, video_id, seconds, api_base, litellm_params, headers, extra_body=None):
|
||||
raise NotImplementedError("video extension is not supported for RunwayML")
|
||||
|
||||
def transform_video_extension_response(self, raw_response, logging_obj, custom_llm_provider=None):
|
||||
raise NotImplementedError("video extension is not supported for RunwayML")
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from litellm._uuid import uuid
|
||||
from typing import Any, Dict
|
||||
|
||||
from litellm._uuid import uuid
|
||||
from litellm.llms.vertex_ai.common_utils import (
|
||||
_convert_vertex_datetime_to_openai_datetime,
|
||||
)
|
||||
|
|
@ -144,9 +144,10 @@ class VertexAIBatchTransformation:
|
|||
|
||||
output_file_id: str = (
|
||||
response.get("outputInfo", OutputInfo()).get("gcsOutputDirectory", "")
|
||||
+ "/predictions.jsonl"
|
||||
)
|
||||
if output_file_id != "/predictions.jsonl":
|
||||
if output_file_id:
|
||||
output_file_id = output_file_id.rstrip("/") + "/predictions.jsonl"
|
||||
if output_file_id and output_file_id != "/predictions.jsonl":
|
||||
return output_file_id
|
||||
|
||||
output_config = response.get("outputConfig")
|
||||
|
|
@ -158,7 +159,9 @@ class VertexAIBatchTransformation:
|
|||
return output_file_id
|
||||
|
||||
output_uri_prefix = gcs_destination.get("outputUriPrefix", "")
|
||||
return output_uri_prefix
|
||||
if output_uri_prefix.endswith("/predictions.jsonl"):
|
||||
return output_uri_prefix
|
||||
return output_uri_prefix.rstrip("/") + "/predictions.jsonl"
|
||||
|
||||
@classmethod
|
||||
def _get_batch_job_status_from_vertex_ai_batch_response(
|
||||
|
|
|
|||
|
|
@ -624,6 +624,30 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
"""Video delete is not supported."""
|
||||
raise NotImplementedError("Video delete is not supported by Vertex AI Veo.")
|
||||
|
||||
def transform_video_create_character_request(self, name, video, api_base, litellm_params, headers):
|
||||
raise NotImplementedError("video create character is not supported for Vertex AI")
|
||||
|
||||
def transform_video_create_character_response(self, raw_response, logging_obj):
|
||||
raise NotImplementedError("video create character is not supported for Vertex AI")
|
||||
|
||||
def transform_video_get_character_request(self, character_id, api_base, litellm_params, headers):
|
||||
raise NotImplementedError("video get character is not supported for Vertex AI")
|
||||
|
||||
def transform_video_get_character_response(self, raw_response, logging_obj):
|
||||
raise NotImplementedError("video get character is not supported for Vertex AI")
|
||||
|
||||
def transform_video_edit_request(self, prompt, video_id, api_base, litellm_params, headers, extra_body=None):
|
||||
raise NotImplementedError("video edit is not supported for Vertex AI")
|
||||
|
||||
def transform_video_edit_response(self, raw_response, logging_obj, custom_llm_provider=None):
|
||||
raise NotImplementedError("video edit is not supported for Vertex AI")
|
||||
|
||||
def transform_video_extension_request(self, prompt, video_id, seconds, api_base, litellm_params, headers, extra_body=None):
|
||||
raise NotImplementedError("video extension is not supported for Vertex AI")
|
||||
|
||||
def transform_video_extension_response(self, raw_response, logging_obj, custom_llm_provider=None):
|
||||
raise NotImplementedError("video extension is not supported for Vertex AI")
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
|
|
|
|||
|
|
@ -7528,8 +7528,15 @@ def stream_chunk_builder( # noqa: PLR0915
|
|||
]
|
||||
|
||||
if len(annotation_chunks) > 0:
|
||||
annotations = annotation_chunks[0]["choices"][0]["delta"]["annotations"]
|
||||
response["choices"][0]["message"]["annotations"] = annotations
|
||||
# Merge annotations from ALL chunks — providers may spread
|
||||
# them across multiple streaming chunks or send them only in
|
||||
# the final chunk.
|
||||
all_annotations: list = []
|
||||
for ac in annotation_chunks:
|
||||
all_annotations.extend(
|
||||
ac["choices"][0]["delta"]["annotations"]
|
||||
)
|
||||
response["choices"][0]["message"]["annotations"] = all_annotations
|
||||
|
||||
audio_chunks = [
|
||||
chunk
|
||||
|
|
|
|||
|
|
@ -298,7 +298,8 @@ def ocr(
|
|||
verbose_logger.debug(f"OCR optional_params after mapping: {optional_params}")
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
litellm_params={
|
||||
|
|
|
|||
|
|
@ -903,13 +903,17 @@ if MCP_AVAILABLE:
|
|||
try:
|
||||
client_id, client_secret, scopes = _extract_credentials(request)
|
||||
|
||||
_oauth2_flow: Optional[
|
||||
Literal["client_credentials", "authorization_code"]
|
||||
] = (
|
||||
"client_credentials"
|
||||
if client_id and client_secret and request.token_url
|
||||
else None
|
||||
_oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = (
|
||||
request.oauth2_flow or (
|
||||
"client_credentials"
|
||||
if client_id and client_secret and request.token_url
|
||||
else None
|
||||
)
|
||||
)
|
||||
# client_credentials requires token_url to fetch a token; without it the
|
||||
# incoming auth header would be dropped with nothing to replace it.
|
||||
if _oauth2_flow == "client_credentials" and not request.token_url:
|
||||
_oauth2_flow = None
|
||||
|
||||
server_model = MCPServer(
|
||||
server_id=request.server_id or "",
|
||||
|
|
|
|||
|
|
@ -1123,6 +1123,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
authorization_url: Optional[str] = None
|
||||
token_url: Optional[str] = None
|
||||
registration_url: Optional[str] = None
|
||||
oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None
|
||||
allow_all_keys: bool = False
|
||||
available_on_public_internet: bool = True
|
||||
is_byok: bool = False
|
||||
|
|
|
|||
|
|
@ -662,11 +662,12 @@ def _has_user_setup_sso():
|
|||
return sso_setup
|
||||
|
||||
|
||||
def get_customer_user_header_from_mapping(user_id_mapping) -> Optional[str]:
|
||||
def get_customer_user_header_from_mapping(user_id_mapping) -> Optional[list]:
|
||||
"""Return the header_name mapped to CUSTOMER role, if any (dict-based)."""
|
||||
if not user_id_mapping:
|
||||
return None
|
||||
items = user_id_mapping if isinstance(user_id_mapping, list) else [user_id_mapping]
|
||||
customer_headers_mappings = []
|
||||
for item in items:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
|
|
@ -675,7 +676,11 @@ def get_customer_user_header_from_mapping(user_id_mapping) -> Optional[str]:
|
|||
if role is None or not header_name:
|
||||
continue
|
||||
if str(role).lower() == str(LitellmUserRoles.CUSTOMER).lower():
|
||||
return header_name
|
||||
customer_headers_mappings.append(header_name.lower())
|
||||
|
||||
if customer_headers_mappings:
|
||||
return customer_headers_mappings
|
||||
|
||||
return None
|
||||
|
||||
|
||||
|
|
@ -724,7 +729,7 @@ def get_end_user_id_from_request_body(
|
|||
# User query: "system not respecting user_header_name property"
|
||||
# This implies the key in general_settings is 'user_header_name'.
|
||||
if request_headers is not None:
|
||||
custom_header_name_to_check: Optional[str] = None
|
||||
custom_header_name_to_check: Optional[Union[list, str]] = None
|
||||
|
||||
# Prefer user mappings (new behavior)
|
||||
user_id_mapping = general_settings.get("user_header_mappings", None)
|
||||
|
|
@ -741,13 +746,21 @@ def get_end_user_id_from_request_body(
|
|||
custom_header_name_to_check = value
|
||||
|
||||
# If we have a header name to check, try to read it from request headers
|
||||
if isinstance(custom_header_name_to_check, str):
|
||||
if isinstance(custom_header_name_to_check, list):
|
||||
headers_lower = {k.lower(): v for k, v in request_headers.items()}
|
||||
for expected_header in custom_header_name_to_check:
|
||||
header_value = headers_lower.get(expected_header)
|
||||
if header_value is not None:
|
||||
user_id_str = str(header_value)
|
||||
if user_id_str.strip():
|
||||
return user_id_str
|
||||
|
||||
elif isinstance(custom_header_name_to_check, str):
|
||||
for header_name, header_value in request_headers.items():
|
||||
if header_name.lower() == custom_header_name_to_check.lower():
|
||||
user_id_from_header = header_value
|
||||
user_id_str = (
|
||||
str(user_id_from_header)
|
||||
if user_id_from_header is not None
|
||||
str(header_value)
|
||||
if header_value is not None
|
||||
else ""
|
||||
)
|
||||
if user_id_str.strip():
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
|||
get_original_file_id,
|
||||
prepare_data_with_credentials,
|
||||
resolve_input_file_id_to_unified,
|
||||
resolve_output_file_ids_to_unified,
|
||||
update_batch_in_database,
|
||||
)
|
||||
from litellm.proxy.utils import handle_exception_on_proxy, is_known_model
|
||||
|
|
@ -405,9 +406,11 @@ async def retrieve_batch( # noqa: PLR0915
|
|||
verbose_proxy_logger=verbose_proxy_logger,
|
||||
)
|
||||
|
||||
# If batch is in a terminal state, return immediately
|
||||
# If batch is in a terminal state, return immediately.
|
||||
# Include "complete" (DB-normalized form of "completed").
|
||||
if response is not None and response.status in [
|
||||
"completed",
|
||||
"complete",
|
||||
"failed",
|
||||
"cancelled",
|
||||
"expired",
|
||||
|
|
@ -417,10 +420,11 @@ async def retrieve_batch( # noqa: PLR0915
|
|||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
)
|
||||
|
||||
# async_post_call_success_hook replaces batch.id and output_file_id with unified IDs
|
||||
# but not input_file_id. Resolve raw provider ID to unified ID.
|
||||
# The DB may store raw provider file IDs (before hooks translate them).
|
||||
# Resolve any raw input/output/error file IDs to unified IDs.
|
||||
if unified_batch_id:
|
||||
await resolve_input_file_id_to_unified(response, prisma_client)
|
||||
await resolve_output_file_ids_to_unified(response, prisma_client)
|
||||
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.update_request_status(
|
||||
|
|
|
|||
|
|
@ -599,6 +599,10 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"avideo_status",
|
||||
"avideo_content",
|
||||
"avideo_remix",
|
||||
"avideo_create_character",
|
||||
"avideo_get_character",
|
||||
"avideo_edit",
|
||||
"avideo_extension",
|
||||
"acreate_container",
|
||||
"alist_containers",
|
||||
"aingest",
|
||||
|
|
@ -850,6 +854,10 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"avideo_status",
|
||||
"avideo_content",
|
||||
"avideo_remix",
|
||||
"avideo_create_character",
|
||||
"avideo_get_character",
|
||||
"avideo_edit",
|
||||
"avideo_extension",
|
||||
"acreate_container",
|
||||
"alist_containers",
|
||||
"aingest",
|
||||
|
|
|
|||
|
|
@ -697,6 +697,28 @@ async def resolve_input_file_id_to_unified(response, prisma_client) -> None:
|
|||
pass
|
||||
|
||||
|
||||
async def resolve_output_file_ids_to_unified(response, prisma_client) -> None:
|
||||
"""
|
||||
If the batch response contains raw provider output_file_id or error_file_id
|
||||
(not already unified IDs), look up the corresponding unified file IDs from
|
||||
the managed file table and replace them in-place.
|
||||
"""
|
||||
if not prisma_client:
|
||||
return
|
||||
for attr in ("output_file_id", "error_file_id"):
|
||||
raw_id = getattr(response, attr, None)
|
||||
if not raw_id or _is_base64_encoded_unified_file_id(raw_id):
|
||||
continue
|
||||
try:
|
||||
managed_file = await prisma_client.db.litellm_managedfiletable.find_first(
|
||||
where={"flat_model_file_ids": {"has": raw_id}}
|
||||
)
|
||||
if managed_file:
|
||||
setattr(response, attr, managed_file.unified_file_id)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def get_batch_from_database(
|
||||
batch_id: str,
|
||||
unified_batch_id: Union[str, Literal[False]],
|
||||
|
|
@ -809,14 +831,43 @@ async def update_batch_in_database(
|
|||
# Normalize status for database storage
|
||||
db_status = response.status if response.status != "completed" else "complete"
|
||||
|
||||
await prisma_client.db.litellm_managedobjecttable.update(
|
||||
where={"unified_object_id": batch_id},
|
||||
data={
|
||||
"status": db_status,
|
||||
"file_object": response.model_dump_json(),
|
||||
"updated_at": litellm.utils.get_utc_datetime(),
|
||||
},
|
||||
)
|
||||
update_data: dict = {
|
||||
"status": db_status,
|
||||
"file_object": response.model_dump_json(),
|
||||
"updated_at": litellm.utils.get_utc_datetime(),
|
||||
}
|
||||
|
||||
# When a batch reaches completion, also mark batch_processed=True.
|
||||
# The cost callback is enqueued asynchronously during the
|
||||
# aretrieve_batch call that detected completion (via the @client
|
||||
# decorator). It is not awaited, so there is a theoretical window
|
||||
# where the callback hasn't executed yet. In practice the callback
|
||||
# completes reliably. Setting the flag here unblocks file deletion
|
||||
# which queries batch_processed=False. CheckBatchCost acts as a
|
||||
# safety net for the rare case where the callback fails.
|
||||
if db_status == "complete":
|
||||
update_data["batch_processed"] = True
|
||||
|
||||
try:
|
||||
await prisma_client.db.litellm_managedobjecttable.update(
|
||||
where={"unified_object_id": batch_id},
|
||||
data=update_data,
|
||||
)
|
||||
except Exception as col_err:
|
||||
# If the batch_processed column doesn't exist (old schema),
|
||||
# retry without it so the status update still succeeds.
|
||||
err_str = str(col_err).lower()
|
||||
if "batch_processed" in err_str and update_data.get("batch_processed") is not None:
|
||||
verbose_proxy_logger.warning(
|
||||
f"batch_processed column not found, retrying update without it: {col_err}"
|
||||
)
|
||||
update_data.pop("batch_processed", None)
|
||||
await prisma_client.db.litellm_managedobjecttable.update(
|
||||
where={"unified_object_id": batch_id},
|
||||
data=update_data,
|
||||
)
|
||||
else:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
f"Failed to update batch status in ManagedObjectTable: {e}"
|
||||
|
|
|
|||
|
|
@ -557,11 +557,11 @@ class ProxyInitializationHelpers:
|
|||
envvar="MAX_REQUESTS_BEFORE_RESTART",
|
||||
)
|
||||
@click.option(
|
||||
"--skip_db_migration_check",
|
||||
"--enforce_prisma_migration_check",
|
||||
is_flag=True,
|
||||
default=False,
|
||||
help="Warn and continue instead of exiting when database migration fails.",
|
||||
envvar="SKIP_DB_MIGRATION_CHECK",
|
||||
help="Exit with error if database migration fails on startup.",
|
||||
envvar="ENFORCE_PRISMA_MIGRATION_CHECK",
|
||||
)
|
||||
def run_server( # noqa: PLR0915
|
||||
host,
|
||||
|
|
@ -602,7 +602,7 @@ def run_server( # noqa: PLR0915
|
|||
skip_server_startup,
|
||||
keepalive_timeout,
|
||||
max_requests_before_restart,
|
||||
skip_db_migration_check: bool,
|
||||
enforce_prisma_migration_check: bool,
|
||||
):
|
||||
args = locals()
|
||||
if local:
|
||||
|
|
@ -716,6 +716,7 @@ def run_server( # noqa: PLR0915
|
|||
for k, v in new_env_var.items():
|
||||
os.environ[k] = v
|
||||
|
||||
litellm_settings = None
|
||||
if config is not None:
|
||||
"""
|
||||
Allow user to pass in db url via config
|
||||
|
|
@ -830,7 +831,9 @@ def run_server( # noqa: PLR0915
|
|||
"pool_timeout": db_connection_timeout,
|
||||
}
|
||||
database_url = get_secret("DATABASE_URL", default_value=None)
|
||||
modified_url = append_query_params(database_url, params)
|
||||
modified_url = append_query_params(
|
||||
str(database_url) if database_url else None, params
|
||||
)
|
||||
os.environ["DATABASE_URL"] = modified_url
|
||||
if os.getenv("DIRECT_URL", None) is not None:
|
||||
### add connection pool + pool timeout args
|
||||
|
|
@ -865,17 +868,17 @@ def run_server( # noqa: PLR0915
|
|||
if not PrismaManager.setup_database(
|
||||
use_migrate=not use_prisma_db_push
|
||||
):
|
||||
if skip_db_migration_check:
|
||||
print( # noqa
|
||||
"\033[1;33mLiteLLM Proxy: Database migration failed but continuing startup. "
|
||||
"Pass --skip_db_migration_check to allow this.\033[0m"
|
||||
)
|
||||
else:
|
||||
if enforce_prisma_migration_check:
|
||||
print( # noqa
|
||||
"\033[1;31mLiteLLM Proxy: Database setup failed after multiple retries. "
|
||||
"The proxy cannot start safely. Please check your database connection and migration status.\033[0m"
|
||||
)
|
||||
sys.exit(1)
|
||||
else:
|
||||
print( # noqa
|
||||
"\033[1;33mLiteLLM Proxy: Database migration failed but continuing startup. "
|
||||
"Set --enforce_prisma_migration_check or ENFORCE_PRISMA_MIGRATION_CHECK=true to exit on failure.\033[0m"
|
||||
)
|
||||
else:
|
||||
print( # noqa
|
||||
f"Unable to connect to DB. DATABASE_URL found in environment, but prisma package not found." # noqa
|
||||
|
|
|
|||
|
|
@ -54,6 +54,10 @@ ROUTE_ENDPOINT_MAPPING = {
|
|||
"avideo_status": "/videos/{video_id}",
|
||||
"avideo_content": "/videos/{video_id}/content",
|
||||
"avideo_remix": "/videos/{video_id}/remix",
|
||||
"avideo_create_character": "/videos/characters",
|
||||
"avideo_get_character": "/videos/characters/{character_id}",
|
||||
"avideo_edit": "/videos/edits",
|
||||
"avideo_extension": "/videos/extensions",
|
||||
"acreate_realtime_client_secret": "/realtime/client_secrets",
|
||||
"arealtime_calls": "/realtime/calls",
|
||||
"acreate_container": "/containers",
|
||||
|
|
@ -201,6 +205,10 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin
|
|||
"avideo_status",
|
||||
"avideo_content",
|
||||
"avideo_remix",
|
||||
"avideo_create_character",
|
||||
"avideo_get_character",
|
||||
"avideo_edit",
|
||||
"avideo_extension",
|
||||
"acreate_container",
|
||||
"alist_containers",
|
||||
"aretrieve_container",
|
||||
|
|
@ -370,6 +378,10 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin
|
|||
"avideo_status",
|
||||
"avideo_content",
|
||||
"avideo_remix",
|
||||
"avideo_create_character",
|
||||
"avideo_get_character",
|
||||
"avideo_edit",
|
||||
"avideo_extension",
|
||||
"avector_store_file_list",
|
||||
"avector_store_file_retrieve",
|
||||
"avector_store_file_content",
|
||||
|
|
@ -449,8 +461,13 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin
|
|||
"avideo_status",
|
||||
"avideo_content",
|
||||
"avideo_remix",
|
||||
"avideo_create_character",
|
||||
"avideo_get_character",
|
||||
"avideo_edit",
|
||||
"avideo_extension",
|
||||
]:
|
||||
# Video endpoints: If model is provided (e.g., from decoded video_id), try router first
|
||||
# Video endpoints: If model is provided (e.g., from decoded video_id or target_model_names),
|
||||
# try router first to allow for multi-deployment load balancing
|
||||
try:
|
||||
return getattr(llm_router, f"{route_type}")(**data)
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
from typing import Any, Dict, Optional
|
||||
|
||||
import orjson
|
||||
from fastapi import APIRouter, Depends, File, Request, Response, UploadFile
|
||||
from fastapi import APIRouter, Depends, File, Form, Request, Response, UploadFile
|
||||
from fastapi.responses import ORJSONResponse
|
||||
|
||||
from litellm.proxy._types import *
|
||||
|
|
@ -16,7 +16,15 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
|
|||
get_custom_llm_provider_from_request_query,
|
||||
)
|
||||
from litellm.proxy.image_endpoints.endpoints import batch_to_bytesio
|
||||
from litellm.types.videos.utils import decode_video_id_with_provider
|
||||
from litellm.proxy.video_endpoints.utils import (
|
||||
encode_character_id_in_response,
|
||||
extract_model_from_target_model_names,
|
||||
get_custom_provider_from_data,
|
||||
)
|
||||
from litellm.types.videos.utils import (
|
||||
decode_character_id_with_provider,
|
||||
decode_video_id_with_provider,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
|
@ -504,3 +512,424 @@ async def video_remix(
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
version=version,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/videos/characters",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_class=ORJSONResponse,
|
||||
tags=["videos"],
|
||||
)
|
||||
@router.post(
|
||||
"/videos/characters",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_class=ORJSONResponse,
|
||||
tags=["videos"],
|
||||
)
|
||||
async def video_create_character(
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
video: UploadFile = File(...),
|
||||
name: str = Form(...),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Create a character from an uploaded video file.
|
||||
|
||||
Follows the OpenAI Videos API spec:
|
||||
https://platform.openai.com/docs/api-reference/videos/create-character
|
||||
|
||||
Example:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/videos/characters" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-F "video=@character_video.mp4" \
|
||||
-F "name=my_character"
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
select_data_generator,
|
||||
user_api_base,
|
||||
user_max_tokens,
|
||||
user_model,
|
||||
user_request_timeout,
|
||||
user_temperature,
|
||||
version,
|
||||
)
|
||||
|
||||
data = await _read_request_body(request=request)
|
||||
video_file = await batch_to_bytesio([video])
|
||||
if video_file:
|
||||
data["video"] = video_file[0]
|
||||
|
||||
target_model_name = extract_model_from_target_model_names(
|
||||
data.get("target_model_names")
|
||||
)
|
||||
if target_model_name and not data.get("model"):
|
||||
data["model"] = target_model_name
|
||||
|
||||
custom_llm_provider = (
|
||||
get_custom_llm_provider_from_request_headers(request=request)
|
||||
or get_custom_llm_provider_from_request_query(request=request)
|
||||
or get_custom_provider_from_data(data=data)
|
||||
or "openai"
|
||||
)
|
||||
data["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
response = await processor.base_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
route_type="avideo_create_character",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
general_settings=general_settings,
|
||||
proxy_config=proxy_config,
|
||||
select_data_generator=select_data_generator,
|
||||
model=None,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
user_max_tokens=user_max_tokens,
|
||||
user_api_base=user_api_base,
|
||||
version=version,
|
||||
)
|
||||
if target_model_name:
|
||||
hidden_params = getattr(response, "_hidden_params", {}) or {}
|
||||
provider_for_encoding = (
|
||||
hidden_params.get("custom_llm_provider")
|
||||
or custom_llm_provider
|
||||
or "openai"
|
||||
)
|
||||
model_id_for_encoding = hidden_params.get("model_id") or data.get("model")
|
||||
response = encode_character_id_in_response(
|
||||
response=response,
|
||||
custom_llm_provider=provider_for_encoding,
|
||||
model_id=model_id_for_encoding,
|
||||
)
|
||||
return response
|
||||
except Exception as e:
|
||||
raise await processor._handle_llm_api_exception(
|
||||
e=e,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
version=version,
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/videos/characters/{character_id}",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_class=ORJSONResponse,
|
||||
tags=["videos"],
|
||||
)
|
||||
@router.get(
|
||||
"/videos/characters/{character_id}",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_class=ORJSONResponse,
|
||||
tags=["videos"],
|
||||
)
|
||||
async def video_get_character(
|
||||
character_id: str,
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Retrieve a character by ID.
|
||||
|
||||
Follows the OpenAI Videos API spec:
|
||||
https://platform.openai.com/docs/api-reference/videos/get-character
|
||||
|
||||
Example:
|
||||
```bash
|
||||
curl -X GET "http://localhost:4000/v1/videos/characters/char_123" \
|
||||
-H "Authorization: Bearer sk-1234"
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
select_data_generator,
|
||||
user_api_base,
|
||||
user_max_tokens,
|
||||
user_model,
|
||||
user_request_timeout,
|
||||
user_temperature,
|
||||
version,
|
||||
)
|
||||
|
||||
original_requested_character_id = character_id
|
||||
data: Dict[str, Any] = {"character_id": character_id}
|
||||
|
||||
decoded = decode_character_id_with_provider(character_id)
|
||||
provider_from_id = decoded.get("custom_llm_provider")
|
||||
model_id_from_decoded = decoded.get("model_id")
|
||||
decoded_character_id = decoded.get("character_id")
|
||||
if decoded_character_id:
|
||||
data["character_id"] = decoded_character_id
|
||||
|
||||
custom_llm_provider = (
|
||||
get_custom_llm_provider_from_request_headers(request=request)
|
||||
or get_custom_llm_provider_from_request_query(request=request)
|
||||
or await get_custom_llm_provider_from_request_body(request=request)
|
||||
or provider_from_id
|
||||
or "openai"
|
||||
)
|
||||
data["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
if model_id_from_decoded and llm_router:
|
||||
resolved_model = llm_router.resolve_model_name_from_model_id(
|
||||
model_id_from_decoded
|
||||
)
|
||||
if resolved_model:
|
||||
data["model"] = resolved_model
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
response = await processor.base_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
route_type="avideo_get_character",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
general_settings=general_settings,
|
||||
proxy_config=proxy_config,
|
||||
select_data_generator=select_data_generator,
|
||||
model=None,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
user_max_tokens=user_max_tokens,
|
||||
user_api_base=user_api_base,
|
||||
version=version,
|
||||
)
|
||||
if original_requested_character_id.startswith("character_"):
|
||||
provider_for_encoding = provider_from_id or custom_llm_provider or "openai"
|
||||
model_id_for_encoding = model_id_from_decoded
|
||||
response = encode_character_id_in_response(
|
||||
response=response,
|
||||
custom_llm_provider=provider_for_encoding,
|
||||
model_id=model_id_for_encoding,
|
||||
)
|
||||
return response
|
||||
except Exception as e:
|
||||
raise await processor._handle_llm_api_exception(
|
||||
e=e,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
version=version,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/videos/edits",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_class=ORJSONResponse,
|
||||
tags=["videos"],
|
||||
)
|
||||
@router.post(
|
||||
"/videos/edits",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_class=ORJSONResponse,
|
||||
tags=["videos"],
|
||||
)
|
||||
async def video_edit(
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Create a video edit job.
|
||||
|
||||
Follows the OpenAI Videos API spec:
|
||||
https://platform.openai.com/docs/api-reference/videos/create-edit
|
||||
|
||||
Example:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/videos/edits" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"prompt": "Make it brighter", "video": {"id": "video_123"}}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
select_data_generator,
|
||||
user_api_base,
|
||||
user_max_tokens,
|
||||
user_model,
|
||||
user_request_timeout,
|
||||
user_temperature,
|
||||
version,
|
||||
)
|
||||
|
||||
body = await request.body()
|
||||
data = orjson.loads(body)
|
||||
|
||||
# Extract video_id from nested video object
|
||||
video_ref = data.pop("video", {})
|
||||
video_id = video_ref.get("id", "") if isinstance(video_ref, dict) else ""
|
||||
data["video_id"] = video_id
|
||||
|
||||
decoded = decode_video_id_with_provider(video_id)
|
||||
provider_from_id = decoded.get("custom_llm_provider")
|
||||
model_id_from_decoded = decoded.get("model_id")
|
||||
|
||||
custom_llm_provider = (
|
||||
get_custom_llm_provider_from_request_headers(request=request)
|
||||
or get_custom_llm_provider_from_request_query(request=request)
|
||||
or get_custom_provider_from_data(data=data)
|
||||
or provider_from_id
|
||||
or "openai"
|
||||
)
|
||||
data["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
if model_id_from_decoded and llm_router:
|
||||
resolved_model = llm_router.resolve_model_name_from_model_id(
|
||||
model_id_from_decoded
|
||||
)
|
||||
if resolved_model:
|
||||
data["model"] = resolved_model
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
return await processor.base_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
route_type="avideo_edit",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
general_settings=general_settings,
|
||||
proxy_config=proxy_config,
|
||||
select_data_generator=select_data_generator,
|
||||
model=None,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
user_max_tokens=user_max_tokens,
|
||||
user_api_base=user_api_base,
|
||||
version=version,
|
||||
)
|
||||
except Exception as e:
|
||||
raise await processor._handle_llm_api_exception(
|
||||
e=e,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
version=version,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/videos/extensions",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_class=ORJSONResponse,
|
||||
tags=["videos"],
|
||||
)
|
||||
@router.post(
|
||||
"/videos/extensions",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_class=ORJSONResponse,
|
||||
tags=["videos"],
|
||||
)
|
||||
async def video_extension(
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Create a video extension.
|
||||
|
||||
Follows the OpenAI Videos API spec:
|
||||
https://platform.openai.com/docs/api-reference/videos/create-extension
|
||||
|
||||
Example:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/videos/extensions" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"prompt": "Continue the scene", "seconds": "5", "video": {"id": "video_123"}}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
select_data_generator,
|
||||
user_api_base,
|
||||
user_max_tokens,
|
||||
user_model,
|
||||
user_request_timeout,
|
||||
user_temperature,
|
||||
version,
|
||||
)
|
||||
|
||||
body = await request.body()
|
||||
data = orjson.loads(body)
|
||||
|
||||
# Extract video_id from nested video object
|
||||
video_ref = data.pop("video", {})
|
||||
video_id = video_ref.get("id", "") if isinstance(video_ref, dict) else ""
|
||||
data["video_id"] = video_id
|
||||
|
||||
decoded = decode_video_id_with_provider(video_id)
|
||||
provider_from_id = decoded.get("custom_llm_provider")
|
||||
model_id_from_decoded = decoded.get("model_id")
|
||||
|
||||
custom_llm_provider = (
|
||||
get_custom_llm_provider_from_request_headers(request=request)
|
||||
or get_custom_llm_provider_from_request_query(request=request)
|
||||
or get_custom_provider_from_data(data=data)
|
||||
or provider_from_id
|
||||
or "openai"
|
||||
)
|
||||
data["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
if model_id_from_decoded and llm_router:
|
||||
resolved_model = llm_router.resolve_model_name_from_model_id(
|
||||
model_id_from_decoded
|
||||
)
|
||||
if resolved_model:
|
||||
data["model"] = resolved_model
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
return await processor.base_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
route_type="avideo_extension",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
general_settings=general_settings,
|
||||
proxy_config=proxy_config,
|
||||
select_data_generator=select_data_generator,
|
||||
model=None,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
user_max_tokens=user_max_tokens,
|
||||
user_api_base=user_api_base,
|
||||
version=version,
|
||||
)
|
||||
except Exception as e:
|
||||
raise await processor._handle_llm_api_exception(
|
||||
e=e,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
version=version,
|
||||
)
|
||||
|
|
|
|||
56
litellm/proxy/video_endpoints/utils.py
Normal file
56
litellm/proxy/video_endpoints/utils.py
Normal file
|
|
@ -0,0 +1,56 @@
|
|||
from typing import Any, Dict, Optional
|
||||
|
||||
import orjson
|
||||
|
||||
from litellm.types.videos.utils import encode_character_id_with_provider
|
||||
|
||||
|
||||
def extract_model_from_target_model_names(target_model_names: Any) -> Optional[str]:
|
||||
if isinstance(target_model_names, str):
|
||||
target_model_names = [m.strip() for m in target_model_names.split(",") if m.strip()]
|
||||
elif not isinstance(target_model_names, list):
|
||||
return None
|
||||
return target_model_names[0] if target_model_names else None
|
||||
|
||||
|
||||
def get_custom_provider_from_data(data: Dict[str, Any]) -> Optional[str]:
|
||||
custom_llm_provider = data.get("custom_llm_provider")
|
||||
if custom_llm_provider:
|
||||
return custom_llm_provider
|
||||
|
||||
extra_body = data.get("extra_body")
|
||||
if isinstance(extra_body, str):
|
||||
try:
|
||||
parsed_extra_body = orjson.loads(extra_body)
|
||||
if isinstance(parsed_extra_body, dict):
|
||||
extra_body = parsed_extra_body
|
||||
except Exception:
|
||||
extra_body = None
|
||||
|
||||
if isinstance(extra_body, dict):
|
||||
extra_body_custom_llm_provider = extra_body.get("custom_llm_provider")
|
||||
if isinstance(extra_body_custom_llm_provider, str):
|
||||
return extra_body_custom_llm_provider
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def encode_character_id_in_response(
|
||||
response: Any, custom_llm_provider: str, model_id: Optional[str]
|
||||
) -> Any:
|
||||
if isinstance(response, dict) and response.get("id"):
|
||||
response["id"] = encode_character_id_with_provider(
|
||||
character_id=response["id"],
|
||||
provider=custom_llm_provider,
|
||||
model_id=model_id,
|
||||
)
|
||||
return response
|
||||
|
||||
character_id = getattr(response, "id", None)
|
||||
if isinstance(character_id, str) and character_id:
|
||||
response.id = encode_character_id_with_provider(
|
||||
character_id=character_id,
|
||||
provider=custom_llm_provider,
|
||||
model_id=model_id,
|
||||
)
|
||||
return response
|
||||
|
|
@ -136,7 +136,8 @@ async def acreate_realtime_client_secret(
|
|||
dynamic_api_key=dynamic_api_key,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model_name,
|
||||
optional_params={"expires_after": expires_after, "session": session},
|
||||
litellm_params={"api_base": resolved_api_base},
|
||||
|
|
@ -186,7 +187,8 @@ async def arealtime_calls(
|
|||
dynamic_api_key=dynamic_api_key,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model_name,
|
||||
optional_params={"realtime_calls": True, "session": session},
|
||||
litellm_params={"api_base": resolved_api_base},
|
||||
|
|
@ -247,7 +249,8 @@ async def _arealtime( # noqa: PLR0915
|
|||
if query_params is not None:
|
||||
query_params = {**query_params, "model": model}
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
user=user,
|
||||
optional_params={},
|
||||
|
|
|
|||
|
|
@ -108,7 +108,6 @@ def rerank( # noqa: PLR0915
|
|||
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
|
||||
proxy_server_request = kwargs.get("proxy_server_request", None)
|
||||
model_info = kwargs.get("model_info", None)
|
||||
metadata = kwargs.get("metadata", {})
|
||||
user = kwargs.get("user", None)
|
||||
client = kwargs.get("client", None)
|
||||
try:
|
||||
|
|
@ -164,7 +163,8 @@ def rerank( # noqa: PLR0915
|
|||
|
||||
model_response = RerankResponse()
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
user=user,
|
||||
optional_params=dict(optional_rerank_params),
|
||||
|
|
@ -172,7 +172,6 @@ def rerank( # noqa: PLR0915
|
|||
"litellm_call_id": litellm_call_id,
|
||||
"proxy_server_request": proxy_server_request,
|
||||
"model_info": model_info,
|
||||
"metadata": metadata,
|
||||
"preset_cache_key": None,
|
||||
"stream_response": {},
|
||||
**optional_params.model_dump(exclude_unset=True),
|
||||
|
|
|
|||
|
|
@ -692,11 +692,11 @@ def responses(
|
|||
return run_async_function(aresponses_api_with_mcp, **mcp_call_kwargs)
|
||||
|
||||
# get provider config
|
||||
responses_api_provider_config: Optional[
|
||||
BaseResponsesAPIConfig
|
||||
] = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=model,
|
||||
provider=custom_llm_provider,
|
||||
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=model,
|
||||
provider=custom_llm_provider,
|
||||
)
|
||||
)
|
||||
|
||||
local_vars.update(kwargs)
|
||||
|
|
@ -738,11 +738,9 @@ def responses(
|
|||
)
|
||||
)
|
||||
|
||||
# Pre Call logging - preserve metadata for custom callbacks
|
||||
# When called from completion bridge (codex models), metadata is in litellm_metadata
|
||||
metadata_for_callbacks = metadata or kwargs.get("litellm_metadata") or {}
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
user=user,
|
||||
optional_params=dict(responses_api_request_params),
|
||||
|
|
@ -750,8 +748,6 @@ def responses(
|
|||
**responses_api_request_params,
|
||||
"aresponses": _is_async,
|
||||
"litellm_call_id": litellm_call_id,
|
||||
"metadata": metadata_for_callbacks,
|
||||
"litellm_metadata": kwargs.get("litellm_metadata", {}),
|
||||
},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
|
@ -912,11 +908,11 @@ def delete_responses(
|
|||
raise ValueError("custom_llm_provider is required but passed as None")
|
||||
|
||||
# get provider config
|
||||
responses_api_provider_config: Optional[
|
||||
BaseResponsesAPIConfig
|
||||
] = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=custom_llm_provider,
|
||||
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=custom_llm_provider,
|
||||
)
|
||||
)
|
||||
|
||||
if responses_api_provider_config is None:
|
||||
|
|
@ -927,7 +923,8 @@ def delete_responses(
|
|||
local_vars.update(kwargs)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=local_vars,
|
||||
model=None,
|
||||
optional_params={
|
||||
"response_id": response_id,
|
||||
|
|
@ -1092,11 +1089,11 @@ def get_responses(
|
|||
raise ValueError("custom_llm_provider is required but passed as None")
|
||||
|
||||
# get provider config
|
||||
responses_api_provider_config: Optional[
|
||||
BaseResponsesAPIConfig
|
||||
] = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=custom_llm_provider,
|
||||
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=custom_llm_provider,
|
||||
)
|
||||
)
|
||||
|
||||
if responses_api_provider_config is None:
|
||||
|
|
@ -1107,7 +1104,8 @@ def get_responses(
|
|||
local_vars.update(kwargs)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=local_vars,
|
||||
model=None,
|
||||
optional_params={
|
||||
"response_id": response_id,
|
||||
|
|
@ -1249,11 +1247,11 @@ def list_input_items(
|
|||
if custom_llm_provider is None:
|
||||
raise ValueError("custom_llm_provider is required but passed as None")
|
||||
|
||||
responses_api_provider_config: Optional[
|
||||
BaseResponsesAPIConfig
|
||||
] = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=custom_llm_provider,
|
||||
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=custom_llm_provider,
|
||||
)
|
||||
)
|
||||
|
||||
if responses_api_provider_config is None:
|
||||
|
|
@ -1263,7 +1261,8 @@ def list_input_items(
|
|||
|
||||
local_vars.update(kwargs)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=local_vars,
|
||||
model=None,
|
||||
optional_params={"response_id": response_id},
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
|
|
@ -1407,11 +1406,11 @@ def cancel_responses(
|
|||
raise ValueError("custom_llm_provider is required but passed as None")
|
||||
|
||||
# get provider config
|
||||
responses_api_provider_config: Optional[
|
||||
BaseResponsesAPIConfig
|
||||
] = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=custom_llm_provider,
|
||||
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=None,
|
||||
provider=custom_llm_provider,
|
||||
)
|
||||
)
|
||||
|
||||
if responses_api_provider_config is None:
|
||||
|
|
@ -1422,7 +1421,8 @@ def cancel_responses(
|
|||
local_vars.update(kwargs)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=local_vars,
|
||||
model=None,
|
||||
optional_params={
|
||||
"response_id": response_id,
|
||||
|
|
@ -1594,11 +1594,11 @@ def compact_responses(
|
|||
raise ValueError("custom_llm_provider is required but passed as None")
|
||||
|
||||
# get provider config
|
||||
responses_api_provider_config: Optional[
|
||||
BaseResponsesAPIConfig
|
||||
] = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=model,
|
||||
provider=custom_llm_provider,
|
||||
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=model,
|
||||
provider=custom_llm_provider,
|
||||
)
|
||||
)
|
||||
|
||||
if responses_api_provider_config is None:
|
||||
|
|
@ -1626,7 +1626,8 @@ def compact_responses(
|
|||
)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=local_vars,
|
||||
model=model,
|
||||
optional_params=dict(responses_api_request_params),
|
||||
litellm_params={
|
||||
|
|
@ -1729,7 +1730,8 @@ async def _aresponses_websocket(
|
|||
api_key=api_key,
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
user=user,
|
||||
optional_params={},
|
||||
|
|
|
|||
|
|
@ -1076,12 +1076,20 @@ class Router:
|
|||
"""Initialize video endpoints."""
|
||||
from litellm.videos import (
|
||||
avideo_content,
|
||||
avideo_create_character,
|
||||
avideo_edit,
|
||||
avideo_extension,
|
||||
avideo_generation,
|
||||
avideo_get_character,
|
||||
avideo_list,
|
||||
avideo_remix,
|
||||
avideo_status,
|
||||
video_content,
|
||||
video_create_character,
|
||||
video_edit,
|
||||
video_extension,
|
||||
video_generation,
|
||||
video_get_character,
|
||||
video_list,
|
||||
video_remix,
|
||||
video_status,
|
||||
|
|
@ -1111,6 +1119,26 @@ class Router:
|
|||
avideo_remix, call_type="avideo_remix"
|
||||
)
|
||||
self.video_remix = self.factory_function(video_remix, call_type="video_remix")
|
||||
self.avideo_create_character = self.factory_function(
|
||||
avideo_create_character, call_type="avideo_create_character"
|
||||
)
|
||||
self.video_create_character = self.factory_function(
|
||||
video_create_character, call_type="video_create_character"
|
||||
)
|
||||
self.avideo_get_character = self.factory_function(
|
||||
avideo_get_character, call_type="avideo_get_character"
|
||||
)
|
||||
self.video_get_character = self.factory_function(
|
||||
video_get_character, call_type="video_get_character"
|
||||
)
|
||||
self.avideo_edit = self.factory_function(avideo_edit, call_type="avideo_edit")
|
||||
self.video_edit = self.factory_function(video_edit, call_type="video_edit")
|
||||
self.avideo_extension = self.factory_function(
|
||||
avideo_extension, call_type="avideo_extension"
|
||||
)
|
||||
self.video_extension = self.factory_function(
|
||||
video_extension, call_type="video_extension"
|
||||
)
|
||||
|
||||
def _initialize_container_endpoints(self):
|
||||
"""Initialize container endpoints."""
|
||||
|
|
@ -4828,6 +4856,14 @@ class Router:
|
|||
"video_content",
|
||||
"avideo_remix",
|
||||
"video_remix",
|
||||
"avideo_create_character",
|
||||
"video_create_character",
|
||||
"avideo_get_character",
|
||||
"video_get_character",
|
||||
"avideo_edit",
|
||||
"video_edit",
|
||||
"avideo_extension",
|
||||
"video_extension",
|
||||
"acreate_container",
|
||||
"create_container",
|
||||
"alist_containers",
|
||||
|
|
@ -4995,6 +5031,10 @@ class Router:
|
|||
"avideo_status",
|
||||
"avideo_content",
|
||||
"avideo_remix",
|
||||
"avideo_create_character",
|
||||
"avideo_get_character",
|
||||
"avideo_edit",
|
||||
"avideo_extension",
|
||||
"acreate_skill",
|
||||
"alist_skills",
|
||||
"aget_skill",
|
||||
|
|
|
|||
|
|
@ -286,7 +286,8 @@ def search(
|
|||
|
||||
# Pre Call logging
|
||||
model_name = f"{search_provider}/search"
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model_name,
|
||||
optional_params=optional_params,
|
||||
litellm_params={
|
||||
|
|
|
|||
|
|
@ -204,7 +204,8 @@ def create_skill(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params=request_body,
|
||||
litellm_params={
|
||||
|
|
@ -389,7 +390,8 @@ def list_skills(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params=query_params,
|
||||
litellm_params={
|
||||
|
|
@ -556,7 +558,8 @@ def get_skill(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"skill_id": skill_id},
|
||||
litellm_params={
|
||||
|
|
@ -722,7 +725,8 @@ def delete_skill(
|
|||
)
|
||||
|
||||
# Pre-call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"skill_id": skill_id},
|
||||
litellm_params={
|
||||
|
|
|
|||
|
|
@ -2187,6 +2187,7 @@ class CreateVideoRequest(TypedDict, total=False):
|
|||
model: Optional[str] - The video generation model to use (defaults to sora-2)
|
||||
seconds: Optional[str] - Clip duration in seconds (defaults to 4 seconds)
|
||||
size: Optional[str] - Output resolution formatted as width x height (defaults to 720x1280)
|
||||
characters: Optional[List[Dict[str, str]]] - Character references to include in generation
|
||||
user: Optional[str] - A unique identifier representing your end-user
|
||||
extra_headers: Optional[Dict[str, str]] - Additional headers
|
||||
extra_body: Optional[Dict[str, str]] - Additional body parameters
|
||||
|
|
@ -2198,6 +2199,7 @@ class CreateVideoRequest(TypedDict, total=False):
|
|||
model: Optional[str]
|
||||
seconds: Optional[str]
|
||||
size: Optional[str]
|
||||
characters: Optional[List[Dict[str, str]]]
|
||||
user: Optional[str]
|
||||
extra_headers: Optional[Dict[str, str]]
|
||||
extra_body: Optional[Dict[str, str]]
|
||||
|
|
|
|||
|
|
@ -358,6 +358,14 @@ class CallTypes(str, Enum):
|
|||
avideo_retrieve_job = "avideo_retrieve_job"
|
||||
video_delete = "video_delete"
|
||||
avideo_delete = "avideo_delete"
|
||||
video_create_character = "video_create_character"
|
||||
avideo_create_character = "avideo_create_character"
|
||||
video_get_character = "video_get_character"
|
||||
avideo_get_character = "avideo_get_character"
|
||||
video_edit = "video_edit"
|
||||
avideo_edit = "avideo_edit"
|
||||
video_extension = "video_extension"
|
||||
avideo_extension = "avideo_extension"
|
||||
vector_store_file_create = "vector_store_file_create"
|
||||
avector_store_file_create = "avector_store_file_create"
|
||||
vector_store_file_list = "vector_store_file_list"
|
||||
|
|
@ -700,6 +708,26 @@ API_ROUTE_TO_CALL_TYPES = {
|
|||
],
|
||||
"/videos/{video_id}/remix": [CallTypes.avideo_remix, CallTypes.video_remix],
|
||||
"/v1/videos/{video_id}/remix": [CallTypes.avideo_remix, CallTypes.video_remix],
|
||||
"/videos/characters": [
|
||||
CallTypes.avideo_create_character,
|
||||
CallTypes.video_create_character,
|
||||
],
|
||||
"/v1/videos/characters": [
|
||||
CallTypes.avideo_create_character,
|
||||
CallTypes.video_create_character,
|
||||
],
|
||||
"/videos/characters/{character_id}": [
|
||||
CallTypes.avideo_get_character,
|
||||
CallTypes.video_get_character,
|
||||
],
|
||||
"/v1/videos/characters/{character_id}": [
|
||||
CallTypes.avideo_get_character,
|
||||
CallTypes.video_get_character,
|
||||
],
|
||||
"/videos/edits": [CallTypes.avideo_edit, CallTypes.video_edit],
|
||||
"/v1/videos/edits": [CallTypes.avideo_edit, CallTypes.video_edit],
|
||||
"/videos/extensions": [CallTypes.avideo_extension, CallTypes.video_extension],
|
||||
"/v1/videos/extensions": [CallTypes.avideo_extension, CallTypes.video_extension],
|
||||
# Vector Stores
|
||||
"/vector_stores": [CallTypes.avector_store_create, CallTypes.vector_store_create],
|
||||
"/v1/vector_stores": [
|
||||
|
|
|
|||
|
|
@ -1,10 +1,9 @@
|
|||
from typing import Any, Dict, List, Literal, Optional
|
||||
|
||||
from openai.types.audio.transcription_create_params import FileTypes # type: ignore
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from litellm.types.utils import FileTypes
|
||||
|
||||
|
||||
class VideoObject(BaseModel):
|
||||
"""Represents a generated video object."""
|
||||
|
|
@ -83,6 +82,7 @@ class VideoCreateOptionalRequestParams(TypedDict, total=False):
|
|||
model: Optional[str]
|
||||
seconds: Optional[str]
|
||||
size: Optional[str]
|
||||
characters: Optional[List[Dict[str, str]]]
|
||||
user: Optional[str]
|
||||
extra_headers: Optional[Dict[str, str]]
|
||||
extra_body: Optional[Dict[str, str]]
|
||||
|
|
@ -104,3 +104,43 @@ class DecodedVideoId(TypedDict, total=False):
|
|||
custom_llm_provider: Optional[str]
|
||||
model_id: Optional[str]
|
||||
video_id: str
|
||||
|
||||
|
||||
class CharacterObject(BaseModel):
|
||||
"""Represents a character created from a video."""
|
||||
|
||||
id: str
|
||||
object: Literal["character"] = "character"
|
||||
created_at: int
|
||||
name: str
|
||||
_hidden_params: Dict[str, Any] = {}
|
||||
|
||||
def __contains__(self, key):
|
||||
return hasattr(self, key)
|
||||
|
||||
def get(self, key, default=None):
|
||||
return getattr(self, key, default)
|
||||
|
||||
def __getitem__(self, key):
|
||||
return getattr(self, key)
|
||||
|
||||
def json(self, **kwargs): # type: ignore
|
||||
try:
|
||||
return self.model_dump(**kwargs)
|
||||
except Exception:
|
||||
return self.dict()
|
||||
|
||||
|
||||
class VideoEditRequestParams(TypedDict, total=False):
|
||||
"""TypedDict for video edit request parameters."""
|
||||
|
||||
prompt: str
|
||||
video: Dict[str, str] # {"id": "video_123"}
|
||||
|
||||
|
||||
class VideoExtensionRequestParams(TypedDict, total=False):
|
||||
"""TypedDict for video extension request parameters."""
|
||||
|
||||
prompt: str
|
||||
seconds: str
|
||||
video: Dict[str, str] # {"id": "video_123"}
|
||||
|
|
|
|||
|
|
@ -12,6 +12,26 @@ from litellm.types.utils import SpecialEnums
|
|||
from litellm.types.videos.main import DecodedVideoId
|
||||
|
||||
VIDEO_ID_PREFIX = "video_"
|
||||
CHARACTER_ID_PREFIX = "character_"
|
||||
CHARACTER_ID_TEMPLATE = "litellm:custom_llm_provider:{};model_id:{};character_id:{}"
|
||||
|
||||
|
||||
class DecodedCharacterId(dict):
|
||||
"""Structure representing a decoded character ID."""
|
||||
|
||||
custom_llm_provider: Optional[str]
|
||||
model_id: Optional[str]
|
||||
character_id: str
|
||||
|
||||
|
||||
def _add_base64_padding(value: str) -> str:
|
||||
"""
|
||||
Add missing base64 padding when IDs are copied without trailing '=' chars.
|
||||
"""
|
||||
missing_padding = len(value) % 4
|
||||
if missing_padding:
|
||||
value += "=" * (4 - missing_padding)
|
||||
return value
|
||||
|
||||
|
||||
def encode_video_id_with_provider(
|
||||
|
|
@ -59,6 +79,7 @@ def decode_video_id_with_provider(encoded_video_id: str) -> DecodedVideoId:
|
|||
|
||||
try:
|
||||
cleaned_id = encoded_video_id.replace(VIDEO_ID_PREFIX, "")
|
||||
cleaned_id = _add_base64_padding(cleaned_id)
|
||||
decoded_id = base64.b64decode(cleaned_id.encode("utf-8")).decode("utf-8")
|
||||
|
||||
if ";" not in decoded_id:
|
||||
|
|
@ -103,3 +124,86 @@ def extract_original_video_id(encoded_video_id: str) -> str:
|
|||
"""Extract original video ID without encoding."""
|
||||
decoded = decode_video_id_with_provider(encoded_video_id)
|
||||
return decoded.get("video_id", encoded_video_id)
|
||||
|
||||
|
||||
def encode_character_id_with_provider(
|
||||
character_id: str, provider: str, model_id: Optional[str] = None
|
||||
) -> str:
|
||||
"""Encode provider and model_id into character_id using base64."""
|
||||
if not provider or not character_id:
|
||||
return character_id
|
||||
|
||||
decoded = decode_character_id_with_provider(character_id)
|
||||
if decoded.get("custom_llm_provider") is not None:
|
||||
return character_id
|
||||
|
||||
assembled_id = CHARACTER_ID_TEMPLATE.format(provider, model_id or "", character_id)
|
||||
base64_encoded_id: str = base64.b64encode(assembled_id.encode("utf-8")).decode(
|
||||
"utf-8"
|
||||
)
|
||||
return f"{CHARACTER_ID_PREFIX}{base64_encoded_id}"
|
||||
|
||||
|
||||
def decode_character_id_with_provider(encoded_character_id: str) -> DecodedCharacterId:
|
||||
"""Decode provider and model_id from encoded character_id."""
|
||||
if not encoded_character_id:
|
||||
return DecodedCharacterId(
|
||||
custom_llm_provider=None,
|
||||
model_id=None,
|
||||
character_id=encoded_character_id,
|
||||
)
|
||||
|
||||
if not encoded_character_id.startswith(CHARACTER_ID_PREFIX):
|
||||
return DecodedCharacterId(
|
||||
custom_llm_provider=None,
|
||||
model_id=None,
|
||||
character_id=encoded_character_id,
|
||||
)
|
||||
|
||||
try:
|
||||
cleaned_id = encoded_character_id.replace(CHARACTER_ID_PREFIX, "")
|
||||
cleaned_id = _add_base64_padding(cleaned_id)
|
||||
decoded_id = base64.b64decode(cleaned_id.encode("utf-8")).decode("utf-8")
|
||||
|
||||
if ";" not in decoded_id:
|
||||
return DecodedCharacterId(
|
||||
custom_llm_provider=None,
|
||||
model_id=None,
|
||||
character_id=encoded_character_id,
|
||||
)
|
||||
|
||||
parts = decoded_id.split(";")
|
||||
|
||||
custom_llm_provider = None
|
||||
model_id = None
|
||||
decoded_character_id = encoded_character_id
|
||||
|
||||
if len(parts) >= 3:
|
||||
custom_llm_provider_part = parts[0]
|
||||
model_id_part = parts[1]
|
||||
character_id_part = parts[2]
|
||||
|
||||
custom_llm_provider = custom_llm_provider_part.replace(
|
||||
"litellm:custom_llm_provider:", ""
|
||||
)
|
||||
model_id = model_id_part.replace("model_id:", "")
|
||||
decoded_character_id = character_id_part.replace("character_id:", "")
|
||||
|
||||
return DecodedCharacterId(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_id=model_id,
|
||||
character_id=decoded_character_id,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"Error decoding character_id '{encoded_character_id}': {e}")
|
||||
return DecodedCharacterId(
|
||||
custom_llm_provider=None,
|
||||
model_id=None,
|
||||
character_id=encoded_character_id,
|
||||
)
|
||||
|
||||
|
||||
def extract_original_character_id(encoded_character_id: str) -> str:
|
||||
"""Extract original character ID without encoding."""
|
||||
decoded = decode_character_id_with_provider(encoded_character_id)
|
||||
return decoded.get("character_id", encoded_character_id)
|
||||
|
|
|
|||
|
|
@ -146,7 +146,8 @@ def create(
|
|||
)
|
||||
create_request["file_id"] = file_id
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={
|
||||
"vector_store_id": vector_store_id,
|
||||
|
|
@ -279,7 +280,8 @@ def list(
|
|||
VectorStoreFileRequestUtils.get_list_query_params(local_vars)
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"vector_store_id": vector_store_id, **list_query},
|
||||
litellm_params={
|
||||
|
|
@ -387,7 +389,8 @@ def retrieve(
|
|||
f"Vector store file retrieve is not supported for {custom_llm_provider}"
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={
|
||||
"vector_store_id": vector_store_id,
|
||||
|
|
@ -498,7 +501,8 @@ def retrieve_content(
|
|||
f"Vector store file content retrieve is not supported for {custom_llm_provider}"
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={
|
||||
"vector_store_id": vector_store_id,
|
||||
|
|
@ -619,7 +623,8 @@ def update(
|
|||
)
|
||||
update_request["attributes"] = attributes
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={
|
||||
"vector_store_id": vector_store_id,
|
||||
|
|
@ -733,7 +738,8 @@ def delete(
|
|||
f"Vector store file delete is not supported for {custom_llm_provider}"
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={
|
||||
"vector_store_id": vector_store_id,
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ LiteLLM SDK Functions for Creating and Searching Vector Stores
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import builtins
|
||||
import contextvars
|
||||
from functools import partial
|
||||
from typing import Any, Coroutine, Dict, List, Optional, Union
|
||||
|
|
@ -233,7 +234,8 @@ def create(
|
|||
)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={
|
||||
"name": name,
|
||||
|
|
@ -395,11 +397,11 @@ def search(
|
|||
|
||||
## MOCK RESPONSE LOGIC
|
||||
if litellm_params.mock_response and isinstance(
|
||||
litellm_params.mock_response, (str, list)
|
||||
litellm_params.mock_response, (str, builtins.list)
|
||||
):
|
||||
mock_results = None
|
||||
if isinstance(litellm_params.mock_response, list):
|
||||
mock_results = litellm_params.mock_response
|
||||
if isinstance(litellm_params.mock_response, builtins.list):
|
||||
mock_results = litellm_params.mock_response # type: ignore[assignment]
|
||||
return mock_vector_store_search_response(mock_results=mock_results)
|
||||
|
||||
# Default to OpenAI for vector stores
|
||||
|
|
@ -440,7 +442,8 @@ def search(
|
|||
)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=api_type,
|
||||
optional_params={
|
||||
"vector_store_id": vector_store_id,
|
||||
|
|
@ -585,7 +588,8 @@ def retrieve(
|
|||
f"Vector store retrieve is not supported for {custom_llm_provider}"
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"vector_store_id": vector_store_id},
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
|
|
@ -732,7 +736,8 @@ def list(
|
|||
f"Vector store list is not supported for {custom_llm_provider}"
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={
|
||||
"after": after,
|
||||
|
|
@ -895,7 +900,8 @@ def update(
|
|||
)
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={
|
||||
"vector_store_id": vector_store_id,
|
||||
|
|
@ -1035,7 +1041,8 @@ def delete(
|
|||
f"Vector store delete is not supported for {custom_llm_provider}"
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=None,
|
||||
optional_params={"vector_store_id": vector_store_id},
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
|
|
|
|||
|
|
@ -1,16 +1,24 @@
|
|||
"""Video generation and management functions for LiteLLM."""
|
||||
|
||||
from .main import (
|
||||
avideo_generation,
|
||||
video_generation,
|
||||
avideo_list,
|
||||
video_list,
|
||||
avideo_status,
|
||||
video_status,
|
||||
avideo_content,
|
||||
video_content,
|
||||
avideo_create_character,
|
||||
avideo_edit,
|
||||
avideo_extension,
|
||||
avideo_generation,
|
||||
avideo_get_character,
|
||||
avideo_list,
|
||||
avideo_remix,
|
||||
avideo_status,
|
||||
video_content,
|
||||
video_create_character,
|
||||
video_edit,
|
||||
video_extension,
|
||||
video_generation,
|
||||
video_get_character,
|
||||
video_list,
|
||||
video_remix,
|
||||
video_status,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
|
|
@ -24,4 +32,12 @@ __all__ = [
|
|||
"video_content",
|
||||
"avideo_remix",
|
||||
"video_remix",
|
||||
"avideo_create_character",
|
||||
"video_create_character",
|
||||
"avideo_get_character",
|
||||
"video_get_character",
|
||||
"avideo_edit",
|
||||
"video_edit",
|
||||
"avideo_extension",
|
||||
"video_extension",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from litellm.main import base_llm_http_handler
|
|||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import CallTypes, FileTypes
|
||||
from litellm.types.videos.main import (
|
||||
CharacterObject,
|
||||
VideoCreateOptionalRequestParams,
|
||||
VideoObject,
|
||||
)
|
||||
|
|
@ -119,17 +120,18 @@ async def avideo_generation(
|
|||
def video_generation(
|
||||
prompt: str,
|
||||
model: Optional[str] = None,
|
||||
input_reference: Optional[str] = None,
|
||||
input_reference: Optional[FileTypes] = None,
|
||||
seconds: Optional[str] = None,
|
||||
size: Optional[str] = None,
|
||||
user: Optional[str] = None,
|
||||
timeout=600, # default to 10 minutes
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
custom_llm_provider=None,
|
||||
timeout: int = 600,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
*,
|
||||
avideo_generation: Literal[True],
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> Coroutine[Any, Any, VideoObject]:
|
||||
...
|
||||
|
||||
|
|
@ -138,18 +140,18 @@ def video_generation(
|
|||
def video_generation(
|
||||
prompt: str,
|
||||
model: Optional[str] = None,
|
||||
input_reference: Optional[str] = None,
|
||||
input_reference: Optional[FileTypes] = None,
|
||||
seconds: Optional[str] = None,
|
||||
size: Optional[str] = None,
|
||||
user: Optional[str] = None,
|
||||
timeout=600, # default to 10 minutes
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
custom_llm_provider=None,
|
||||
timeout: int = 600,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
*,
|
||||
avideo_generation: Literal[False] = False,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> VideoObject:
|
||||
...
|
||||
|
||||
|
|
@ -231,7 +233,8 @@ def video_generation( # noqa: PLR0915
|
|||
)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
user=user,
|
||||
optional_params=dict(video_generation_request_params),
|
||||
|
|
@ -348,7 +351,8 @@ def video_content(
|
|||
}
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model="",
|
||||
user=kwargs.get("user"),
|
||||
optional_params=dict(video_content_request_params),
|
||||
|
|
@ -528,14 +532,14 @@ async def avideo_remix(
|
|||
def video_remix(
|
||||
video_id: str,
|
||||
prompt: str,
|
||||
timeout=600, # default to 10 minutes
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
custom_llm_provider=None,
|
||||
timeout: int = 600,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
*,
|
||||
avideo_remix: Literal[True],
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> Coroutine[Any, Any, VideoObject]:
|
||||
...
|
||||
|
||||
|
|
@ -544,14 +548,14 @@ def video_remix(
|
|||
def video_remix(
|
||||
video_id: str,
|
||||
prompt: str,
|
||||
timeout=600, # default to 10 minutes
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
custom_llm_provider=None,
|
||||
timeout: int = 600,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
*,
|
||||
avideo_remix: Literal[False] = False,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> VideoObject:
|
||||
...
|
||||
|
||||
|
|
@ -618,7 +622,8 @@ def video_remix( # noqa: PLR0915
|
|||
}
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model="",
|
||||
user=kwargs.get("user"),
|
||||
optional_params=dict(video_remix_request_params),
|
||||
|
|
@ -744,14 +749,14 @@ def video_list(
|
|||
after: Optional[str] = None,
|
||||
limit: Optional[int] = None,
|
||||
order: Optional[str] = None,
|
||||
timeout=600, # default to 10 minutes
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
custom_llm_provider=None,
|
||||
timeout: int = 600,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
*,
|
||||
avideo_list: Literal[True],
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> Coroutine[Any, Any, List[VideoObject]]:
|
||||
...
|
||||
|
||||
|
|
@ -761,14 +766,14 @@ def video_list(
|
|||
after: Optional[str] = None,
|
||||
limit: Optional[int] = None,
|
||||
order: Optional[str] = None,
|
||||
timeout=600, # default to 10 minutes
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
custom_llm_provider=None,
|
||||
timeout: int = 600,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
*,
|
||||
avideo_list: Literal[False] = False,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> List[VideoObject]:
|
||||
...
|
||||
|
||||
|
|
@ -834,7 +839,8 @@ def video_list( # noqa: PLR0915
|
|||
}
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model="",
|
||||
user=kwargs.get("user"),
|
||||
optional_params=dict(video_list_request_params),
|
||||
|
|
@ -849,7 +855,7 @@ def video_list( # noqa: PLR0915
|
|||
litellm_logging_obj.call_type = CallTypes.video_list.value
|
||||
|
||||
# Call the handler with _is_async flag instead of directly calling the async handler
|
||||
return base_llm_http_handler.video_list_handler(
|
||||
return base_llm_http_handler.video_list_handler( # type: ignore[return-value]
|
||||
after=after,
|
||||
limit=limit,
|
||||
order=order,
|
||||
|
|
@ -945,14 +951,14 @@ async def avideo_status(
|
|||
@overload
|
||||
def video_status(
|
||||
video_id: str,
|
||||
timeout=600, # default to 10 minutes
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
custom_llm_provider=None,
|
||||
timeout: int = 600,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
*,
|
||||
avideo_status: Literal[True],
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> Coroutine[Any, Any, VideoObject]:
|
||||
...
|
||||
|
||||
|
|
@ -960,14 +966,14 @@ def video_status(
|
|||
@overload
|
||||
def video_status(
|
||||
video_id: str,
|
||||
timeout=600, # default to 10 minutes
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
custom_llm_provider=None,
|
||||
timeout: int = 600,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
*,
|
||||
avideo_status: Literal[False] = False,
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
) -> VideoObject:
|
||||
...
|
||||
|
||||
|
|
@ -1054,7 +1060,8 @@ def video_status( # noqa: PLR0915
|
|||
}
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model="",
|
||||
user=kwargs.get("user"),
|
||||
optional_params=dict(video_status_request_params),
|
||||
|
|
@ -1090,3 +1097,522 @@ def video_status( # noqa: PLR0915
|
|||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
async def avideo_create_character(
|
||||
name: str,
|
||||
video: Any,
|
||||
timeout=600,
|
||||
custom_llm_provider=None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
) -> CharacterObject:
|
||||
"""
|
||||
Asynchronously create a character from an uploaded video file.
|
||||
Maps to POST /v1/videos/characters
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
if custom_llm_provider is None:
|
||||
custom_llm_provider = "openai"
|
||||
|
||||
func = partial(
|
||||
video_create_character,
|
||||
name=name,
|
||||
video=video,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
extra_query=extra_query,
|
||||
extra_body=extra_body,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model="",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
def video_create_character(
|
||||
name: str,
|
||||
video: Any,
|
||||
timeout=600,
|
||||
custom_llm_provider=None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
) -> Union[CharacterObject, Coroutine[Any, Any, CharacterObject]]:
|
||||
"""
|
||||
Create a character from an uploaded video file.
|
||||
Maps to POST /v1/videos/characters
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("async_call", False) is True
|
||||
|
||||
mock_response = kwargs.get("mock_response", None)
|
||||
if mock_response is not None:
|
||||
if isinstance(mock_response, str):
|
||||
mock_response = json.loads(mock_response)
|
||||
return CharacterObject(**mock_response)
|
||||
|
||||
if custom_llm_provider is None:
|
||||
custom_llm_provider = "openai"
|
||||
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config(
|
||||
model=None,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if provider_config is None:
|
||||
raise ValueError(f"video create character is not supported for {custom_llm_provider}")
|
||||
|
||||
local_vars.update(kwargs)
|
||||
request_params: Dict = {"name": name}
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
model="",
|
||||
user=kwargs.get("user"),
|
||||
optional_params=dict(request_params),
|
||||
litellm_params={"litellm_call_id": litellm_call_id, **request_params},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
litellm_logging_obj.call_type = CallTypes.video_create_character.value
|
||||
|
||||
return base_llm_http_handler.video_create_character_handler(
|
||||
name=name,
|
||||
video=video,
|
||||
video_provider_config=provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=litellm_logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout or DEFAULT_REQUEST_TIMEOUT,
|
||||
_is_async=_is_async,
|
||||
client=kwargs.get("client"),
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model="",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
async def avideo_get_character(
|
||||
character_id: str,
|
||||
timeout=600,
|
||||
custom_llm_provider=None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
) -> CharacterObject:
|
||||
"""
|
||||
Asynchronously retrieve a character by ID.
|
||||
Maps to GET /v1/videos/characters/{character_id}
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
func = partial(
|
||||
video_get_character,
|
||||
character_id=character_id,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
extra_query=extra_query,
|
||||
extra_body=extra_body,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model="",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
def video_get_character(
|
||||
character_id: str,
|
||||
timeout=600,
|
||||
custom_llm_provider=None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
) -> Union[CharacterObject, Coroutine[Any, Any, CharacterObject]]:
|
||||
"""
|
||||
Retrieve a character by ID.
|
||||
Maps to GET /v1/videos/characters/{character_id}
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("async_call", False) is True
|
||||
|
||||
mock_response = kwargs.get("mock_response", None)
|
||||
if mock_response is not None:
|
||||
if isinstance(mock_response, str):
|
||||
mock_response = json.loads(mock_response)
|
||||
return CharacterObject(**mock_response)
|
||||
|
||||
if custom_llm_provider is None:
|
||||
custom_llm_provider = "openai"
|
||||
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config(
|
||||
model=None,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if provider_config is None:
|
||||
raise ValueError(f"video get character is not supported for {custom_llm_provider}")
|
||||
|
||||
local_vars.update(kwargs)
|
||||
request_params: Dict = {"character_id": character_id}
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
model="",
|
||||
user=kwargs.get("user"),
|
||||
optional_params=dict(request_params),
|
||||
litellm_params={"litellm_call_id": litellm_call_id, **request_params},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
litellm_logging_obj.call_type = CallTypes.video_get_character.value
|
||||
|
||||
return base_llm_http_handler.video_get_character_handler(
|
||||
character_id=character_id,
|
||||
video_provider_config=provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=litellm_logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout or DEFAULT_REQUEST_TIMEOUT,
|
||||
_is_async=_is_async,
|
||||
client=kwargs.get("client"),
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model="",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
async def avideo_edit(
|
||||
video_id: str,
|
||||
prompt: str,
|
||||
timeout=600,
|
||||
custom_llm_provider=None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
) -> VideoObject:
|
||||
"""
|
||||
Asynchronously create a video edit job.
|
||||
Maps to POST /v1/videos/edits
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
func = partial(
|
||||
video_edit,
|
||||
video_id=video_id,
|
||||
prompt=prompt,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
extra_query=extra_query,
|
||||
extra_body=extra_body,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model="",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
def video_edit(
|
||||
video_id: str,
|
||||
prompt: str,
|
||||
timeout=600,
|
||||
custom_llm_provider=None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
) -> Union[VideoObject, Coroutine[Any, Any, VideoObject]]:
|
||||
"""
|
||||
Create a video edit job.
|
||||
Maps to POST /v1/videos/edits
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("async_call", False) is True
|
||||
|
||||
mock_response = kwargs.get("mock_response", None)
|
||||
if mock_response is not None:
|
||||
if isinstance(mock_response, str):
|
||||
mock_response = json.loads(mock_response)
|
||||
return VideoObject(**mock_response)
|
||||
|
||||
if custom_llm_provider is None:
|
||||
decoded = decode_video_id_with_provider(video_id)
|
||||
custom_llm_provider = decoded.get("custom_llm_provider") or "openai"
|
||||
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config(
|
||||
model=None,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if provider_config is None:
|
||||
raise ValueError(f"video edit is not supported for {custom_llm_provider}")
|
||||
|
||||
local_vars.update(kwargs)
|
||||
request_params: Dict = {"video_id": video_id, "prompt": prompt}
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
model="",
|
||||
user=kwargs.get("user"),
|
||||
optional_params=dict(request_params),
|
||||
litellm_params={"litellm_call_id": litellm_call_id, **request_params},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
litellm_logging_obj.call_type = CallTypes.video_edit.value
|
||||
|
||||
return base_llm_http_handler.video_edit_handler(
|
||||
prompt=prompt,
|
||||
video_id=video_id,
|
||||
video_provider_config=provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=litellm_logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout or DEFAULT_REQUEST_TIMEOUT,
|
||||
_is_async=_is_async,
|
||||
client=kwargs.get("client"),
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model="",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
async def avideo_extension(
|
||||
video_id: str,
|
||||
prompt: str,
|
||||
seconds: str,
|
||||
timeout=600,
|
||||
custom_llm_provider=None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
) -> VideoObject:
|
||||
"""
|
||||
Asynchronously create a video extension.
|
||||
Maps to POST /v1/videos/extensions
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
func = partial(
|
||||
video_extension,
|
||||
video_id=video_id,
|
||||
prompt=prompt,
|
||||
seconds=seconds,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
extra_query=extra_query,
|
||||
extra_body=extra_body,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model="",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
def video_extension(
|
||||
video_id: str,
|
||||
prompt: str,
|
||||
seconds: str,
|
||||
timeout=600,
|
||||
custom_llm_provider=None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
) -> Union[VideoObject, Coroutine[Any, Any, VideoObject]]:
|
||||
"""
|
||||
Create a video extension.
|
||||
Maps to POST /v1/videos/extensions
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("async_call", False) is True
|
||||
|
||||
mock_response = kwargs.get("mock_response", None)
|
||||
if mock_response is not None:
|
||||
if isinstance(mock_response, str):
|
||||
mock_response = json.loads(mock_response)
|
||||
return VideoObject(**mock_response)
|
||||
|
||||
if custom_llm_provider is None:
|
||||
decoded = decode_video_id_with_provider(video_id)
|
||||
custom_llm_provider = decoded.get("custom_llm_provider") or "openai"
|
||||
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
provider_config: Optional[BaseVideoConfig] = ProviderConfigManager.get_provider_video_config(
|
||||
model=None,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if provider_config is None:
|
||||
raise ValueError(f"video extension is not supported for {custom_llm_provider}")
|
||||
|
||||
local_vars.update(kwargs)
|
||||
request_params: Dict = {"video_id": video_id, "prompt": prompt, "seconds": seconds}
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
model="",
|
||||
user=kwargs.get("user"),
|
||||
optional_params=dict(request_params),
|
||||
litellm_params={"litellm_call_id": litellm_call_id, **request_params},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
litellm_logging_obj.call_type = CallTypes.video_extension.value
|
||||
|
||||
return base_llm_http_handler.video_extension_handler(
|
||||
prompt=prompt,
|
||||
video_id=video_id,
|
||||
seconds=seconds,
|
||||
video_provider_config=provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=litellm_logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout or DEFAULT_REQUEST_TIMEOUT,
|
||||
_is_async=_is_async,
|
||||
client=kwargs.get("client"),
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model="",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -29,10 +29,26 @@ verbose_logger.setLevel(logging.DEBUG)
|
|||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
import random
|
||||
import socket
|
||||
import httpx
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
|
||||
def _can_resolve_openai():
|
||||
"""Check if api.openai.com is reachable (DNS resolves)."""
|
||||
try:
|
||||
socket.getaddrinfo("api.openai.com", 443, socket.AF_UNSPEC, socket.SOCK_STREAM)
|
||||
return True
|
||||
except socket.gaierror:
|
||||
return False
|
||||
|
||||
|
||||
skip_if_no_openai_network = pytest.mark.skipif(
|
||||
not _can_resolve_openai(),
|
||||
reason="Cannot resolve api.openai.com - skipping integration test due to DNS issues",
|
||||
)
|
||||
|
||||
|
||||
def load_vertex_ai_credentials():
|
||||
# Define the path to the vertex_key.json file
|
||||
print("loading vertex ai credentials")
|
||||
|
|
@ -78,6 +94,7 @@ def load_vertex_ai_credentials():
|
|||
|
||||
@pytest.mark.parametrize("provider", ["openai"]) # , "azure"
|
||||
@pytest.mark.asyncio
|
||||
@skip_if_no_openai_network
|
||||
async def test_create_batch(provider):
|
||||
"""
|
||||
1. Create File for Batch completion
|
||||
|
|
@ -252,6 +269,7 @@ def cleanup_azure_ft_models():
|
|||
@pytest.mark.parametrize("provider", ["openai"])
|
||||
@pytest.mark.asyncio()
|
||||
@pytest.mark.flaky(retries=3, delay=1)
|
||||
@skip_if_no_openai_network
|
||||
async def test_async_create_batch(provider):
|
||||
"""
|
||||
1. Create File for Batch completion
|
||||
|
|
@ -464,9 +482,24 @@ mock_vertex_list_response = {
|
|||
@pytest.mark.asyncio
|
||||
async def test_avertex_batch_prediction(monkeypatch):
|
||||
monkeypatch.setenv("GCS_BUCKET_NAME", "litellm-local")
|
||||
monkeypatch.setenv("VERTEXAI_PROJECT", "mock-project")
|
||||
monkeypatch.setenv("VERTEXAI_LOCATION", "us-central1")
|
||||
|
||||
# Mock Google auth so the test doesn't need real credentials
|
||||
mock_creds = MagicMock()
|
||||
mock_creds.token = "mock-token"
|
||||
mock_creds.valid = True
|
||||
mock_creds.expiry = None
|
||||
monkeypatch.setattr(
|
||||
"google.auth.default",
|
||||
lambda *args, **kwargs: (mock_creds, "mock-project"),
|
||||
)
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
# Configure mock response object
|
||||
mock_response = MagicMock()
|
||||
mock_response.raise_for_status.return_value = None
|
||||
|
||||
async def mock_side_effect(*args, **kwargs):
|
||||
print("args", args, "kwargs", kwargs)
|
||||
|
|
@ -478,21 +511,10 @@ async def test_avertex_batch_prediction(monkeypatch):
|
|||
mock_response.status_code = 200
|
||||
return mock_response
|
||||
|
||||
with patch.object(
|
||||
client, "post", side_effect=mock_side_effect
|
||||
) as mock_post, patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post"
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
side_effect=mock_side_effect,
|
||||
) as mock_global_post:
|
||||
# Configure mock responses
|
||||
mock_response = MagicMock()
|
||||
mock_response.raise_for_status.return_value = None
|
||||
|
||||
# Set up different responses for different API calls
|
||||
|
||||
mock_post.side_effect = mock_side_effect
|
||||
mock_global_post.side_effect = mock_side_effect
|
||||
|
||||
# load_vertex_ai_credentials()
|
||||
litellm.set_verbose = True
|
||||
litellm._turn_on_debug()
|
||||
file_name = "vertex_batch_completions.jsonl"
|
||||
|
|
@ -504,7 +526,6 @@ async def test_avertex_batch_prediction(monkeypatch):
|
|||
file=open(file_path, "rb"),
|
||||
purpose="batch",
|
||||
custom_llm_provider="vertex_ai",
|
||||
client=client
|
||||
)
|
||||
print("Response from creating file=", file_obj)
|
||||
|
||||
|
|
@ -623,6 +644,7 @@ async def test_vertex_async_create_batch_logs_error_body_on_http_error():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@skip_if_no_openai_network
|
||||
async def test_delete_batch_output_file():
|
||||
"""
|
||||
Test that deleting a batch output file works correctly.
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
import base64
|
||||
import json
|
||||
from typing import cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -477,6 +479,81 @@ async def test_output_file_id_for_batch_retrieve():
|
|||
assert not cast(LiteLLMBatch, response).output_file_id.startswith("file-")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_output_file_id_preserves_target_model_names_when_model_name_missing():
|
||||
"""
|
||||
Regression test: when provider response does not include _hidden_params.model_name
|
||||
(e.g. Vertex batch retrieve), unified output_file_id should still include
|
||||
target_model_names from the managed input file ID.
|
||||
"""
|
||||
from openai.types.batch import BatchRequestCounts
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.llms.openai import OpenAIFileObject
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
batch = LiteLLMBatch(
|
||||
id="batch_123",
|
||||
completion_window="24h",
|
||||
created_at=1750883933,
|
||||
endpoint="/v1/chat/completions",
|
||||
input_file_id="file-input-provider-id",
|
||||
object="batch",
|
||||
status="completed",
|
||||
output_file_id="file-provider-output-id",
|
||||
request_counts=BatchRequestCounts(completed=1, failed=0, total=1),
|
||||
usage=None,
|
||||
)
|
||||
|
||||
# Build a valid managed input id string and base64 encode it.
|
||||
managed_input_file_payload = (
|
||||
"litellm_proxy:application/octet-stream;"
|
||||
"unified_id,test-uuid;"
|
||||
"target_model_names,gemini-2.5-pro;"
|
||||
"llm_output_file_id,file-input-1;"
|
||||
"llm_output_file_model_id,model-id-1"
|
||||
)
|
||||
managed_input_file_id = (
|
||||
base64.urlsafe_b64encode(managed_input_file_payload.encode())
|
||||
.decode()
|
||||
.rstrip("=")
|
||||
)
|
||||
|
||||
batch._hidden_params = {
|
||||
"model_id": "model-id-1",
|
||||
"unified_batch_id": "litellm_proxy;model_id:model-id-1;llm_batch_id:batch_123",
|
||||
"unified_file_id": managed_input_file_id,
|
||||
# Intentionally omit model_name to mimic Vertex issue.
|
||||
}
|
||||
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
DualCache(), prisma_client=AsyncMock()
|
||||
)
|
||||
|
||||
provider_output_file = OpenAIFileObject(
|
||||
id="file-provider-output-id",
|
||||
object="file",
|
||||
bytes=1,
|
||||
created_at=1,
|
||||
filename="predictions.jsonl",
|
||||
purpose="batch_output",
|
||||
)
|
||||
|
||||
with patch("litellm.afile_retrieve", new_callable=AsyncMock) as mock_retrieve:
|
||||
mock_retrieve.return_value = provider_output_file
|
||||
response = await proxy_managed_files.async_post_call_success_hook(
|
||||
data={},
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
|
||||
response=batch,
|
||||
)
|
||||
|
||||
decoded_output_file_id = _is_base64_encoded_unified_file_id(
|
||||
cast(LiteLLMBatch, response).output_file_id
|
||||
)
|
||||
assert decoded_output_file_id
|
||||
assert "target_model_names,gemini-2.5-pro" in cast(str, decoded_output_file_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_error_file_id_for_failed_batch():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,4 +1,9 @@
|
|||
# conftest.py
|
||||
#
|
||||
# xdist-compatible test isolation for guardrails tests.
|
||||
# Pattern matches tests/test_litellm/conftest.py:
|
||||
# - Function-scoped fixture saves/restores litellm globals (no reload)
|
||||
# - Module-scoped fixture reloads only in single-process mode
|
||||
|
||||
import importlib
|
||||
import os
|
||||
|
|
@ -10,58 +15,85 @@ sys.path.insert(
|
|||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
import litellm
|
||||
import asyncio
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def event_loop():
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
loop = asyncio.new_event_loop()
|
||||
yield loop
|
||||
loop.close()
|
||||
|
||||
@pytest.fixture(scope="function", autouse=True)
|
||||
def setup_and_teardown():
|
||||
def isolate_litellm_state():
|
||||
"""
|
||||
This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained.
|
||||
Per-function isolation fixture.
|
||||
|
||||
Saves and restores litellm callback/global state so tests don't leak
|
||||
side effects. Works safely under pytest-xdist parallel execution.
|
||||
"""
|
||||
curr_dir = os.getcwd() # Get the current working directory
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the project directory to the system path
|
||||
# Save original callback state
|
||||
original_state = {}
|
||||
for attr in (
|
||||
"callbacks",
|
||||
"success_callback",
|
||||
"failure_callback",
|
||||
"_async_success_callback",
|
||||
"_async_failure_callback",
|
||||
):
|
||||
if hasattr(litellm, attr):
|
||||
val = getattr(litellm, attr)
|
||||
original_state[attr] = val.copy() if val else []
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
import asyncio
|
||||
# Save other globals that tests commonly mutate
|
||||
for attr in ("set_verbose", "cache", "num_retries"):
|
||||
if hasattr(litellm, attr):
|
||||
original_state[attr] = getattr(litellm, attr)
|
||||
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
# flush all logs
|
||||
asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue())
|
||||
# Flush cache before test
|
||||
if hasattr(litellm, "in_memory_llm_clients_cache"):
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
# Clear callbacks before test
|
||||
for attr in (
|
||||
"success_callback",
|
||||
"failure_callback",
|
||||
"_async_success_callback",
|
||||
"_async_failure_callback",
|
||||
):
|
||||
if hasattr(litellm, attr):
|
||||
setattr(litellm, attr, [])
|
||||
|
||||
importlib.reload(litellm)
|
||||
|
||||
try:
|
||||
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
|
||||
import litellm.proxy.proxy_server
|
||||
|
||||
importlib.reload(litellm.proxy.proxy_server)
|
||||
except Exception as e:
|
||||
print(f"Error reloading litellm.proxy.proxy_server: {e}")
|
||||
|
||||
import asyncio
|
||||
|
||||
loop = asyncio.get_event_loop_policy().new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
print(litellm)
|
||||
# from litellm import Router, completion, aembedding, acompletion, embedding
|
||||
yield
|
||||
|
||||
# Teardown code (executes after the yield point)
|
||||
loop.close() # Close the loop created earlier
|
||||
asyncio.set_event_loop(None) # Remove the reference to the loop
|
||||
# Restore all saved state
|
||||
if hasattr(litellm, "in_memory_llm_clients_cache"):
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
for attr, original_value in original_state.items():
|
||||
if hasattr(litellm, attr):
|
||||
setattr(litellm, attr, original_value)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", autouse=True)
|
||||
def setup_and_teardown():
|
||||
"""
|
||||
Module-scoped setup. Reloads litellm only in single-process mode
|
||||
(skipped under xdist to avoid cross-worker interference).
|
||||
"""
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import litellm
|
||||
|
||||
worker_id = os.environ.get("PYTEST_XDIST_WORKER", None)
|
||||
if worker_id is None:
|
||||
importlib.reload(litellm)
|
||||
|
||||
try:
|
||||
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
|
||||
import litellm.proxy.proxy_server
|
||||
|
||||
importlib.reload(litellm.proxy.proxy_server)
|
||||
except Exception as e:
|
||||
print(f"Error reloading litellm.proxy.proxy_server: {e}")
|
||||
|
||||
if hasattr(litellm, "in_memory_llm_clients_cache"):
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
yield
|
||||
|
||||
|
||||
def pytest_collection_modifyitems(config, items):
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ print("Python Path:", sys.path)
|
|||
print("Current Working Directory:", os.getcwd())
|
||||
|
||||
|
||||
import functools
|
||||
from typing import Optional
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
|
@ -34,6 +35,19 @@ from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2
|
|||
from litellm.types.secret_managers.main import KeyManagementSettings
|
||||
|
||||
|
||||
def skip_on_throttling(func):
|
||||
"""Skip async test on AWS ThrottlingException instead of failing."""
|
||||
@functools.wraps(func)
|
||||
async def wrapper(*args, **kwargs):
|
||||
try:
|
||||
return await func(*args, **kwargs)
|
||||
except Exception as e:
|
||||
if "ThrottlingException" in str(e):
|
||||
pytest.skip(f"AWS throttling: {e}")
|
||||
raise
|
||||
return wrapper
|
||||
|
||||
|
||||
def check_aws_credentials():
|
||||
"""Helper function to check if AWS credentials are set"""
|
||||
required_vars = ["AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION_NAME"]
|
||||
|
|
@ -43,6 +57,7 @@ def check_aws_credentials():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@skip_on_throttling
|
||||
async def test_write_and_read_simple_secret():
|
||||
"""Test writing and reading a simple string secret"""
|
||||
check_aws_credentials()
|
||||
|
|
@ -84,6 +99,7 @@ async def test_write_and_read_simple_secret():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@skip_on_throttling
|
||||
async def test_write_and_read_json_secret():
|
||||
"""Test writing and reading a JSON structured secret"""
|
||||
check_aws_credentials()
|
||||
|
|
@ -128,6 +144,7 @@ async def test_write_and_read_json_secret():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@skip_on_throttling
|
||||
async def test_read_nonexistent_secret():
|
||||
"""Test reading a secret that doesn't exist"""
|
||||
check_aws_credentials()
|
||||
|
|
@ -141,6 +158,7 @@ async def test_read_nonexistent_secret():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@skip_on_throttling
|
||||
async def test_primary_secret_functionality():
|
||||
"""Test storing and retrieving secrets from a primary secret"""
|
||||
check_aws_credentials()
|
||||
|
|
@ -196,6 +214,7 @@ async def test_primary_secret_functionality():
|
|||
assert delete_response is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@skip_on_throttling
|
||||
async def test_write_secret_with_description_and_tags():
|
||||
"""Test writing a secret with description and tags"""
|
||||
check_aws_credentials()
|
||||
|
|
@ -402,6 +421,7 @@ def test_load_aws_secret_manager_with_settings():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@skip_on_throttling
|
||||
async def test_end_to_end_iam_role_secret_write():
|
||||
"""
|
||||
Test writing a secret using IAM role assumption (integration test)
|
||||
|
|
|
|||
|
|
@ -2,8 +2,10 @@ import json
|
|||
import os
|
||||
import sys
|
||||
import time
|
||||
from contextlib import asynccontextmanager, contextmanager
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, patch, MagicMock
|
||||
import httpx
|
||||
import pytest
|
||||
import asyncio
|
||||
|
||||
|
|
@ -13,6 +15,63 @@ sys.path.insert(
|
|||
import litellm
|
||||
|
||||
|
||||
# Fake Vertex AI Gemini response for mocking
|
||||
FAKE_VERTEX_GEMINI_RESPONSE = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [{"text": "Hello! How can I help you today?"}],
|
||||
"role": "model",
|
||||
},
|
||||
"finishReason": "STOP",
|
||||
}
|
||||
],
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 5,
|
||||
"candidatesTokenCount": 8,
|
||||
"totalTokenCount": 13,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _make_fake_httpx_response(url: str) -> httpx.Response:
|
||||
"""Create a fake httpx.Response that looks like a Vertex AI Gemini response."""
|
||||
response = httpx.Response(
|
||||
status_code=200,
|
||||
json=FAKE_VERTEX_GEMINI_RESPONSE,
|
||||
request=httpx.Request("POST", url),
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _vertex_ai_mocks():
|
||||
"""Context manager that mocks Vertex AI auth and HTTP calls.
|
||||
|
||||
Mocks at the httpx.AsyncClient.send level so that the
|
||||
@track_llm_api_timing decorator on AsyncHTTPHandler.post still runs,
|
||||
preserving the overhead measurement.
|
||||
"""
|
||||
fake_response = _make_fake_httpx_response(
|
||||
"https://fake-vertex-endpoint/v1/models/gemini-1.5-flash:generateContent"
|
||||
)
|
||||
|
||||
async def fake_send(self, request, **kwargs):
|
||||
await asyncio.sleep(0.2) # simulate ~200ms network latency
|
||||
return fake_response
|
||||
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.vertex_llm_base.VertexBase._ensure_access_token_async",
|
||||
new_callable=AsyncMock,
|
||||
return_value=("Bearer fake-token", "fake-project"),
|
||||
), patch.object(
|
||||
httpx.AsyncClient,
|
||||
"send",
|
||||
new=fake_send,
|
||||
):
|
||||
yield
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
|
|
@ -39,16 +98,19 @@ async def test_litellm_overhead_non_streaming(model):
|
|||
# Specific cases for models
|
||||
#########################################################
|
||||
if model == "vertex_ai/gemini-1.5-flash":
|
||||
kwargs["api_base"] = "https://exampleopenaiendpoint-production.up.railway.app/v1/projects/pathrise-convert-1606954137718/locations/us-central1/publishers/google/models/gemini-1.0-pro-vision-001"
|
||||
# warmup call for auth validation on vertex_ai models
|
||||
await litellm.acompletion(**kwargs)
|
||||
kwargs["vertex_project"] = "fake-project"
|
||||
kwargs["vertex_location"] = "us-central1"
|
||||
if model == "openai/self_hosted":
|
||||
kwargs["api_base"] = "https://exampleopenaiendpoint-production.up.railway.app/"
|
||||
|
||||
async def _run():
|
||||
return await litellm.acompletion(**kwargs)
|
||||
|
||||
response = await litellm.acompletion(
|
||||
**kwargs
|
||||
)
|
||||
if model == "vertex_ai/gemini-1.5-flash":
|
||||
async with _vertex_ai_mocks():
|
||||
response = await _run()
|
||||
else:
|
||||
response = await _run()
|
||||
#########################################################
|
||||
# End of specific cases for models
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -1,9 +1,14 @@
|
|||
# conftest.py
|
||||
#
|
||||
# xdist-compatible test isolation for llm_translation tests.
|
||||
# Mirrors the pattern in tests/local_testing/conftest.py:
|
||||
# - Function-scoped fixture resets litellm globals to true defaults
|
||||
# - Module-scoped reload only in single-process mode
|
||||
|
||||
import importlib
|
||||
import os
|
||||
import sys
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
|
|
@ -13,6 +18,24 @@ import litellm
|
|||
|
||||
import asyncio
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Capture TRUE defaults at conftest import time (before test modules pollute).
|
||||
# ---------------------------------------------------------------------------
|
||||
_SCALAR_DEFAULTS = {
|
||||
"num_retries": getattr(litellm, "num_retries", None),
|
||||
"set_verbose": getattr(litellm, "set_verbose", False),
|
||||
"cache": getattr(litellm, "cache", None),
|
||||
"allowed_fails": getattr(litellm, "allowed_fails", 3),
|
||||
"disable_aiohttp_transport": getattr(litellm, "disable_aiohttp_transport", False),
|
||||
"force_ipv4": getattr(litellm, "force_ipv4", False),
|
||||
"drop_params": getattr(litellm, "drop_params", None),
|
||||
"modify_params": getattr(litellm, "modify_params", False),
|
||||
"api_base": getattr(litellm, "api_base", None),
|
||||
"api_key": getattr(litellm, "api_key", None),
|
||||
"cohere_key": getattr(litellm, "cohere_key", None),
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def event_loop():
|
||||
try:
|
||||
|
|
@ -29,20 +52,39 @@ def setup_and_teardown(event_loop): # Add event_loop as a dependency
|
|||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
|
||||
# ---- Save current state (for teardown restore) ----
|
||||
original_state = {}
|
||||
for attr in (
|
||||
"callbacks",
|
||||
"success_callback",
|
||||
"failure_callback",
|
||||
"_async_success_callback",
|
||||
"_async_failure_callback",
|
||||
):
|
||||
if hasattr(litellm, attr):
|
||||
val = getattr(litellm, attr)
|
||||
original_state[attr] = val.copy() if val else []
|
||||
|
||||
for attr in _SCALAR_DEFAULTS:
|
||||
if hasattr(litellm, attr):
|
||||
original_state[attr] = getattr(litellm, attr)
|
||||
|
||||
# ---- Reset to true defaults before the test ----
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
# flush all logs
|
||||
asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue())
|
||||
|
||||
importlib.reload(litellm)
|
||||
|
||||
# Set the event loop from the fixture
|
||||
asyncio.set_event_loop(event_loop)
|
||||
|
||||
print(litellm)
|
||||
yield
|
||||
|
||||
# ---- Teardown ----
|
||||
for attr, original_value in original_state.items():
|
||||
if hasattr(litellm, attr):
|
||||
setattr(litellm, attr, original_value)
|
||||
|
||||
# Clean up any pending tasks
|
||||
pending = asyncio.all_tasks(event_loop)
|
||||
for task in pending:
|
||||
|
|
|
|||
|
|
@ -838,7 +838,7 @@ async def test_gemini_image_generation_async():
|
|||
IMAGE_URL = response.choices[0].message.images[0]["image_url"]
|
||||
print("IMAGE_URL: ", IMAGE_URL)
|
||||
|
||||
assert CONTENT is not None, "CONTENT is not None"
|
||||
# content may be None when the model returns only an image with no text
|
||||
assert IMAGE_URL is not None, "IMAGE_URL is not None"
|
||||
assert IMAGE_URL["url"] is not None, "IMAGE_URL['url'] is not None"
|
||||
assert IMAGE_URL["url"].startswith("data:image/png;base64,")
|
||||
|
|
|
|||
|
|
@ -1,4 +1,15 @@
|
|||
# conftest.py
|
||||
#
|
||||
# xdist-compatible test isolation for local_testing tests.
|
||||
# Pattern matches tests/test_litellm/conftest.py:
|
||||
# - Function-scoped fixture saves/restores litellm globals (no reload)
|
||||
# - Module-scoped fixture reloads only in single-process mode
|
||||
#
|
||||
# IMPORTANT: True defaults are captured at conftest import time (before any
|
||||
# test module can pollute them via module-level assignments like
|
||||
# `litellm.num_retries = 3`). The function-scoped fixture resets globals to
|
||||
# these true defaults before every test, preventing cross-test contamination
|
||||
# under xdist where module reload is skipped.
|
||||
|
||||
import importlib
|
||||
import os
|
||||
|
|
@ -11,60 +22,126 @@ sys.path.insert(
|
|||
) # Adds the parent directory to the system path
|
||||
import litellm
|
||||
|
||||
import asyncio
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def event_loop():
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
loop = asyncio.new_event_loop()
|
||||
yield loop
|
||||
loop.close()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Capture TRUE defaults at conftest import time. This runs before any test
|
||||
# module's top-level code (e.g. `litellm.num_retries = 3`) executes, so
|
||||
# the values here are guaranteed to be the real package defaults.
|
||||
# ---------------------------------------------------------------------------
|
||||
_SCALAR_DEFAULTS = {
|
||||
"num_retries": getattr(litellm, "num_retries", None),
|
||||
"num_retries_per_request": getattr(litellm, "num_retries_per_request", None),
|
||||
"request_timeout": getattr(litellm, "request_timeout", None),
|
||||
"set_verbose": getattr(litellm, "set_verbose", False),
|
||||
"cache": getattr(litellm, "cache", None),
|
||||
"allowed_fails": getattr(litellm, "allowed_fails", 3),
|
||||
"default_fallbacks": getattr(litellm, "default_fallbacks", None),
|
||||
"enable_azure_ad_token_refresh": getattr(litellm, "enable_azure_ad_token_refresh", None),
|
||||
"tag_budget_config": getattr(litellm, "tag_budget_config", None),
|
||||
"model_cost": getattr(litellm, "model_cost", None),
|
||||
"token_counter": getattr(litellm, "token_counter", None),
|
||||
"disable_aiohttp_transport": getattr(litellm, "disable_aiohttp_transport", False),
|
||||
"force_ipv4": getattr(litellm, "force_ipv4", False),
|
||||
"drop_params": getattr(litellm, "drop_params", None),
|
||||
"modify_params": getattr(litellm, "modify_params", False),
|
||||
"api_base": getattr(litellm, "api_base", None),
|
||||
"api_key": getattr(litellm, "api_key", None),
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(scope="function", autouse=True)
|
||||
def setup_and_teardown():
|
||||
def isolate_litellm_state():
|
||||
"""
|
||||
This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained.
|
||||
Per-function isolation fixture.
|
||||
|
||||
Resets litellm globals to their true defaults before each test and
|
||||
restores them afterward, so tests don't leak side effects.
|
||||
Works safely under pytest-xdist parallel execution.
|
||||
"""
|
||||
curr_dir = os.getcwd() # Get the current working directory
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the project directory to the system path
|
||||
# ---- Save current callback state (for teardown restore) ----
|
||||
original_state = {}
|
||||
for attr in (
|
||||
"callbacks",
|
||||
"success_callback",
|
||||
"failure_callback",
|
||||
"_async_success_callback",
|
||||
"_async_failure_callback",
|
||||
):
|
||||
if hasattr(litellm, attr):
|
||||
val = getattr(litellm, attr)
|
||||
original_state[attr] = val.copy() if val else []
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
import asyncio
|
||||
# Save list-type globals
|
||||
for attr in ("pre_call_rules", "post_call_rules"):
|
||||
if hasattr(litellm, attr):
|
||||
val = getattr(litellm, attr)
|
||||
original_state[attr] = val.copy() if val else []
|
||||
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
# flush all logs
|
||||
asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue())
|
||||
# Save scalar globals
|
||||
for attr in _SCALAR_DEFAULTS:
|
||||
if hasattr(litellm, attr):
|
||||
original_state[attr] = getattr(litellm, attr)
|
||||
|
||||
# ---- Reset to true defaults before the test ----
|
||||
# Flush HTTP client cache
|
||||
if hasattr(litellm, "in_memory_llm_clients_cache"):
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
importlib.reload(litellm)
|
||||
# Clear callbacks and rules
|
||||
for attr in (
|
||||
"callbacks",
|
||||
"success_callback",
|
||||
"failure_callback",
|
||||
"_async_success_callback",
|
||||
"_async_failure_callback",
|
||||
"pre_call_rules",
|
||||
"post_call_rules",
|
||||
):
|
||||
if hasattr(litellm, attr):
|
||||
setattr(litellm, attr, [])
|
||||
|
||||
try:
|
||||
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
|
||||
import litellm.proxy.proxy_server
|
||||
# Reset scalar globals to true defaults (prevents contamination from
|
||||
# module-level code like `litellm.num_retries = 3` in test files)
|
||||
for attr, default_val in _SCALAR_DEFAULTS.items():
|
||||
if hasattr(litellm, attr):
|
||||
setattr(litellm, attr, default_val)
|
||||
|
||||
importlib.reload(litellm.proxy.proxy_server)
|
||||
except Exception as e:
|
||||
print(f"Error reloading litellm.proxy.proxy_server: {e}")
|
||||
|
||||
import asyncio
|
||||
|
||||
loop = asyncio.get_event_loop_policy().new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
print(litellm)
|
||||
# from litellm import Router, completion, aembedding, acompletion, embedding
|
||||
yield
|
||||
|
||||
# Teardown code (executes after the yield point)
|
||||
loop.close() # Close the loop created earlier
|
||||
asyncio.set_event_loop(None) # Remove the reference to the loop
|
||||
# ---- Teardown: restore saved state ----
|
||||
if hasattr(litellm, "in_memory_llm_clients_cache"):
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
for attr, original_value in original_state.items():
|
||||
if hasattr(litellm, attr):
|
||||
setattr(litellm, attr, original_value)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", autouse=True)
|
||||
def setup_and_teardown():
|
||||
"""
|
||||
Module-scoped setup. Reloads litellm only in single-process mode
|
||||
(skipped under xdist to avoid cross-worker interference).
|
||||
"""
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import litellm
|
||||
|
||||
worker_id = os.environ.get("PYTEST_XDIST_WORKER", None)
|
||||
if worker_id is None:
|
||||
importlib.reload(litellm)
|
||||
|
||||
try:
|
||||
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
|
||||
import litellm.proxy.proxy_server
|
||||
|
||||
importlib.reload(litellm.proxy.proxy_server)
|
||||
except Exception as e:
|
||||
print(f"Error reloading litellm.proxy.proxy_server: {e}")
|
||||
|
||||
if hasattr(litellm, "in_memory_llm_clients_cache"):
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
yield
|
||||
|
||||
|
||||
def pytest_collection_modifyitems(config, items):
|
||||
|
|
|
|||
|
|
@ -22,33 +22,37 @@ from litellm import Router
|
|||
|
||||
load_dotenv()
|
||||
|
||||
model_list = [
|
||||
{ # list of model deployments
|
||||
"model_name": "gpt-3.5-turbo", # openai model name
|
||||
"litellm_params": { # params for litellm completion/embedding call
|
||||
"model": "azure/gpt-4.1-mini",
|
||||
"api_key": "bad-key",
|
||||
"api_version": os.getenv("AZURE_API_VERSION"),
|
||||
"api_base": os.getenv("AZURE_API_BASE"),
|
||||
},
|
||||
"tpm": 240000,
|
||||
"rpm": 1800,
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo", # openai model name
|
||||
"litellm_params": { # params for litellm completion/embedding call
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
"tpm": 1000000,
|
||||
"rpm": 9000,
|
||||
},
|
||||
]
|
||||
|
||||
kwargs = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "Hey, how's it going?"}],
|
||||
}
|
||||
def _make_model_list():
|
||||
return [
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "azure/gpt-4.1-mini",
|
||||
"api_key": "bad-key",
|
||||
"api_version": os.getenv("AZURE_API_VERSION"),
|
||||
"api_base": os.getenv("AZURE_API_BASE"),
|
||||
},
|
||||
"tpm": 240000,
|
||||
"rpm": 1800,
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
"tpm": 1000000,
|
||||
"rpm": 9000,
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _make_kwargs():
|
||||
return {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "Hey, how's it going?"}],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.flaky(retries=3, delay=1)
|
||||
|
|
@ -58,8 +62,9 @@ def test_multiple_deployments_sync():
|
|||
|
||||
litellm.set_verbose = False
|
||||
results = []
|
||||
kwargs = _make_kwargs()
|
||||
router = Router(
|
||||
model_list=model_list,
|
||||
model_list=_make_model_list(),
|
||||
redis_host=os.getenv("REDIS_HOST"),
|
||||
redis_password=os.getenv("REDIS_PASSWORD"),
|
||||
redis_port=int(os.getenv("REDIS_PORT")), # type: ignore
|
||||
|
|
@ -85,9 +90,10 @@ def test_multiple_deployments_parallel():
|
|||
litellm.set_verbose = False # Corrected the syntax for setting verbose to False
|
||||
results = []
|
||||
futures = {}
|
||||
kwargs = _make_kwargs()
|
||||
start_time = time.time()
|
||||
router = Router(
|
||||
model_list=model_list,
|
||||
model_list=_make_model_list(),
|
||||
redis_host=os.getenv("REDIS_HOST"),
|
||||
redis_password=os.getenv("REDIS_PASSWORD"),
|
||||
redis_port=int(os.getenv("REDIS_PORT")), # type: ignore
|
||||
|
|
|
|||
|
|
@ -3691,6 +3691,8 @@ def test_vertex_ai_llama_tool_calling():
|
|||
response = completion(**args)
|
||||
except litellm.RateLimitError:
|
||||
pytest.skip("Rate limit error")
|
||||
except litellm.NotFoundError:
|
||||
pytest.skip("Model not found / resource unavailable")
|
||||
print(response)
|
||||
|
||||
assert response.choices[0].message.tool_calls is not None
|
||||
|
|
|
|||
|
|
@ -268,7 +268,7 @@ def test_get_customer_user_header_from_mapping_returns_customer_header():
|
|||
{"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"},
|
||||
]
|
||||
result = get_customer_user_header_from_mapping(mappings)
|
||||
assert result == "X-OpenWebUI-User-Email"
|
||||
assert result == ["x-openwebui-user-email"]
|
||||
|
||||
|
||||
def test_get_customer_user_header_from_mapping_no_customer_returns_none():
|
||||
|
|
|
|||
|
|
@ -147,7 +147,7 @@ def test_caching_dynamic_args(): # test in memory cache
|
|||
port=_redis_port_env,
|
||||
password=_redis_password_env,
|
||||
)
|
||||
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
|
||||
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
|
||||
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
|
||||
print(f"response1: {response1}")
|
||||
print(f"response2: {response2}")
|
||||
|
|
@ -173,7 +173,7 @@ def test_caching_v2(): # test in memory cache
|
|||
try:
|
||||
litellm.set_verbose = True
|
||||
litellm.cache = Cache()
|
||||
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
|
||||
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
|
||||
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
|
||||
print(f"response1: {response1}")
|
||||
print(f"response2: {response2}")
|
||||
|
|
@ -200,9 +200,9 @@ def test_caching_with_ttl():
|
|||
litellm.set_verbose = True
|
||||
litellm.cache = Cache()
|
||||
response1 = completion(
|
||||
model="gpt-3.5-turbo", messages=messages, caching=True, ttl=0
|
||||
model="gpt-3.5-turbo", messages=messages, caching=True, ttl=0, mock_response="Hello world from cache test 1"
|
||||
)
|
||||
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
|
||||
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test 2")
|
||||
print(f"response1: {response1}")
|
||||
print(f"response2: {response2}")
|
||||
litellm.cache = None # disable cache
|
||||
|
|
@ -221,8 +221,8 @@ def test_caching_with_default_ttl():
|
|||
try:
|
||||
litellm.set_verbose = True
|
||||
litellm.cache = Cache(ttl=0)
|
||||
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
|
||||
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
|
||||
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
|
||||
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
|
||||
print(f"response1: {response1}")
|
||||
print(f"response2: {response2}")
|
||||
litellm.cache = None # disable cache
|
||||
|
|
@ -247,10 +247,10 @@ async def test_caching_with_cache_controls(sync_flag):
|
|||
if sync_flag:
|
||||
## TTL = 0
|
||||
response1 = completion(
|
||||
model="gpt-3.5-turbo", messages=messages, cache={"ttl": 0}
|
||||
model="gpt-3.5-turbo", messages=messages, cache={"ttl": 0}, mock_response="Hello world"
|
||||
)
|
||||
response2 = completion(
|
||||
model="gpt-3.5-turbo", messages=messages, cache={"s-maxage": 10}
|
||||
model="gpt-3.5-turbo", messages=messages, cache={"s-maxage": 10}, mock_response="Hello world"
|
||||
)
|
||||
|
||||
assert response2["id"] != response1["id"]
|
||||
|
|
@ -315,7 +315,6 @@ async def test_caching_with_cache_controls(sync_flag):
|
|||
# test_caching_with_cache_controls()
|
||||
|
||||
|
||||
@pytest.mark.flaky(retries=3, delay=1)
|
||||
def test_caching_with_models_v2():
|
||||
messages = [
|
||||
{"role": "user", "content": "who is ishaan CTO of litellm from litellm 2023"}
|
||||
|
|
@ -323,9 +322,9 @@ def test_caching_with_models_v2():
|
|||
litellm.cache = Cache()
|
||||
print("test2 for caching")
|
||||
litellm.set_verbose = True
|
||||
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
|
||||
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
|
||||
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
|
||||
response3 = completion(model="gpt-4.1-nano", messages=messages, caching=True)
|
||||
response3 = completion(model="gpt-4.1-nano", messages=messages, caching=True, mock_response="Different model response")
|
||||
print(f"response1: {response1}")
|
||||
print(f"response2: {response2}")
|
||||
print(f"response3: {response3}")
|
||||
|
|
@ -424,7 +423,7 @@ def test_embedding_caching():
|
|||
text_to_embed = [embedding_large_text]
|
||||
start_time = time.time()
|
||||
embedding1 = embedding(
|
||||
model="text-embedding-ada-002", input=text_to_embed, caching=True
|
||||
model="text-embedding-ada-002", input=text_to_embed, caching=True, mock_response="0.1,0.2,0.3,0.4,0.5"
|
||||
)
|
||||
end_time = time.time()
|
||||
print(f"Embedding 1 response time: {end_time - start_time} seconds")
|
||||
|
|
@ -460,12 +459,12 @@ async def test_embedding_caching_individual_items_and_then_list():
|
|||
"world",
|
||||
]
|
||||
embedding1 = await aembedding(
|
||||
model="text-embedding-ada-002", input=text_to_embed[0], caching=True
|
||||
model="text-embedding-ada-002", input=text_to_embed[0], caching=True, mock_response="0.1,0.2,0.3,0.4,0.5"
|
||||
)
|
||||
initial_prompt_tokens = embedding1.usage.prompt_tokens
|
||||
await asyncio.sleep(1)
|
||||
embedding2 = await aembedding(
|
||||
model="text-embedding-ada-002", input=text_to_embed[1], caching=True
|
||||
model="text-embedding-ada-002", input=text_to_embed[1], caching=True, mock_response="0.6,0.7,0.8,0.9,1.0"
|
||||
)
|
||||
await asyncio.sleep(1)
|
||||
embedding3 = await aembedding(
|
||||
|
|
@ -481,7 +480,7 @@ async def test_embedding_caching_individual_items_and_then_list():
|
|||
additional_text = "this is a new text"
|
||||
text_to_embed.append(additional_text)
|
||||
embedding4 = await aembedding(
|
||||
model="text-embedding-ada-002", input=text_to_embed, caching=True
|
||||
model="text-embedding-ada-002", input=text_to_embed, caching=True, mock_response="0.1,0.2,0.3,0.4,0.5"
|
||||
)
|
||||
assert embedding4.usage.prompt_tokens > embedding3.usage.prompt_tokens
|
||||
|
||||
|
|
@ -491,7 +490,7 @@ async def test_embedding_caching_individual_items():
|
|||
litellm.cache = Cache()
|
||||
text_to_embed = "hello"
|
||||
embedding1 = await aembedding(
|
||||
model="text-embedding-ada-002", input=text_to_embed, caching=True
|
||||
model="text-embedding-ada-002", input=text_to_embed, caching=True, mock_response="0.1,0.2,0.3,0.4,0.5"
|
||||
)
|
||||
|
||||
await asyncio.sleep(1)
|
||||
|
|
@ -533,6 +532,7 @@ def test_embedding_caching_azure():
|
|||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
caching=True,
|
||||
mock_response="0.1,0.2,0.3,0.4,0.5",
|
||||
)
|
||||
end_time = time.time()
|
||||
print(f"Embedding 1 response time: {end_time - start_time} seconds")
|
||||
|
|
@ -762,6 +762,7 @@ async def test_redis_cache_basic():
|
|||
response1 = completion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=messages,
|
||||
mock_response="Hello world from cache test",
|
||||
)
|
||||
|
||||
cache_key = litellm.cache.get_cache_key(
|
||||
|
|
@ -803,6 +804,7 @@ async def test_redis_batch_cache_write():
|
|||
response1 = await litellm.acompletion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=messages,
|
||||
mock_response="Hello world from cache test",
|
||||
)
|
||||
|
||||
response2 = await litellm.acompletion(
|
||||
|
|
@ -843,14 +845,15 @@ def test_redis_cache_completion():
|
|||
messages=messages,
|
||||
caching=True,
|
||||
max_tokens=20,
|
||||
mock_response="Hello world from cache test",
|
||||
)
|
||||
response2 = completion(
|
||||
model="gpt-3.5-turbo", messages=messages, caching=True, max_tokens=20
|
||||
)
|
||||
response3 = completion(
|
||||
model="gpt-3.5-turbo", messages=messages, caching=True, temperature=0.5
|
||||
model="gpt-3.5-turbo", messages=messages, caching=True, temperature=0.5, mock_response="Different params response"
|
||||
)
|
||||
response4 = completion(model="gpt-4o-mini", messages=messages, caching=True)
|
||||
response4 = completion(model="gpt-4o-mini", messages=messages, caching=True, mock_response="Different model response")
|
||||
|
||||
print("\nresponse 1", response1)
|
||||
print("\nresponse 2", response2)
|
||||
|
|
@ -928,12 +931,13 @@ def test_redis_cache_completion_stream():
|
|||
max_tokens=40,
|
||||
temperature=0.2,
|
||||
stream=True,
|
||||
mock_response="In the stillness of numbers, the world turns quietly.",
|
||||
)
|
||||
response_1_id = ""
|
||||
for chunk in response1:
|
||||
print(chunk)
|
||||
response_1_id = chunk.id
|
||||
time.sleep(0.5)
|
||||
time.sleep(1)
|
||||
response2 = completion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=messages,
|
||||
|
|
@ -1072,12 +1076,13 @@ async def test_redis_cache_acompletion_stream():
|
|||
max_tokens=40,
|
||||
temperature=1,
|
||||
stream=True,
|
||||
mock_response="In the stillness of numbers, the world turns quietly.",
|
||||
)
|
||||
async for chunk in response1:
|
||||
response_1_content += chunk.choices[0].delta.content or ""
|
||||
print(response_1_content)
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
await asyncio.sleep(1)
|
||||
print("\n\n Response 1 content: ", response_1_content, "\n\n")
|
||||
|
||||
response2 = await litellm.acompletion(
|
||||
|
|
@ -1122,7 +1127,7 @@ async def test_redis_cache_atext_completion():
|
|||
print("test for caching, atext_completion")
|
||||
|
||||
response1 = await litellm.atext_completion(
|
||||
model="gpt-3.5-turbo-instruct", prompt=prompt, max_tokens=40, temperature=1
|
||||
model="gpt-3.5-turbo-instruct", prompt=prompt, max_tokens=40, temperature=1, mock_response="Hello world from cache test"
|
||||
)
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
|
|
@ -1164,6 +1169,7 @@ async def test_redis_cache_acompletion_stream_bedrock():
|
|||
max_tokens=40,
|
||||
temperature=1,
|
||||
stream=True,
|
||||
mock_response="In the stillness of numbers, the world turns quietly.",
|
||||
)
|
||||
async for chunk in response1:
|
||||
print(chunk)
|
||||
|
|
@ -1231,6 +1237,7 @@ async def test_s3_cache_stream_azure(sync_mode):
|
|||
max_tokens=40,
|
||||
temperature=1,
|
||||
stream=True,
|
||||
mock_response="In the stillness of numbers, the world turns quietly.",
|
||||
)
|
||||
for chunk in response1:
|
||||
print(chunk)
|
||||
|
|
@ -1244,6 +1251,7 @@ async def test_s3_cache_stream_azure(sync_mode):
|
|||
max_tokens=40,
|
||||
temperature=1,
|
||||
stream=True,
|
||||
mock_response="In the stillness of numbers, the world turns quietly.",
|
||||
)
|
||||
async for chunk in response1:
|
||||
print(chunk)
|
||||
|
|
@ -1406,6 +1414,7 @@ def test_custom_redis_cache_with_key():
|
|||
temperature=1,
|
||||
caching=True,
|
||||
num_retries=3,
|
||||
mock_response="Hello world from cache test",
|
||||
)
|
||||
response2 = completion(
|
||||
model="gpt-3.5-turbo",
|
||||
|
|
@ -1420,6 +1429,7 @@ def test_custom_redis_cache_with_key():
|
|||
temperature=1,
|
||||
caching=False,
|
||||
num_retries=3,
|
||||
mock_response="Different uncached response",
|
||||
)
|
||||
|
||||
print(f"response1: {response1}")
|
||||
|
|
@ -1448,21 +1458,15 @@ def test_cache_override():
|
|||
|
||||
# test embedding
|
||||
response1 = embedding(
|
||||
model="text-embedding-ada-002", input=["hello who are you"], caching=False
|
||||
model="text-embedding-ada-002", input=["hello who are you"], caching=False, mock_response="0.1,0.2,0.3,0.4,0.5"
|
||||
)
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
response2 = embedding(
|
||||
model="text-embedding-ada-002", input=["hello who are you"], caching=False
|
||||
model="text-embedding-ada-002", input=["hello who are you"], caching=False, mock_response="0.6,0.7,0.8,0.9,1.0"
|
||||
)
|
||||
|
||||
end_time = time.time()
|
||||
print(f"Embedding 2 response time: {end_time - start_time} seconds")
|
||||
|
||||
assert (
|
||||
end_time - start_time > 0.05
|
||||
) # ensure 2nd response comes in over 0.05s. This should not be cached.
|
||||
# When caching=False, responses should have different IDs
|
||||
assert response1.data[0].embedding != response2.data[0].embedding
|
||||
|
||||
|
||||
# test_cache_override()
|
||||
|
|
@ -1494,6 +1498,7 @@ async def test_cache_control_overrides():
|
|||
}
|
||||
],
|
||||
caching=True,
|
||||
mock_response="Hello world from cache test",
|
||||
)
|
||||
|
||||
print(response1)
|
||||
|
|
@ -1510,6 +1515,7 @@ async def test_cache_control_overrides():
|
|||
],
|
||||
caching=True,
|
||||
cache={"no-cache": True},
|
||||
mock_response="Hello world from cache test",
|
||||
)
|
||||
|
||||
print(response2)
|
||||
|
|
@ -1542,6 +1548,7 @@ def test_sync_cache_control_overrides():
|
|||
}
|
||||
],
|
||||
caching=True,
|
||||
mock_response="Hello world from cache test",
|
||||
)
|
||||
|
||||
print(response1)
|
||||
|
|
@ -1558,6 +1565,7 @@ def test_sync_cache_control_overrides():
|
|||
],
|
||||
caching=True,
|
||||
cache={"no-cache": True},
|
||||
mock_response="Hello world from cache test",
|
||||
)
|
||||
|
||||
print(response2)
|
||||
|
|
@ -1770,6 +1778,7 @@ def test_redis_semantic_cache_completion():
|
|||
}
|
||||
],
|
||||
max_tokens=20,
|
||||
mock_response="Summer sun shines bright and warm.",
|
||||
)
|
||||
print(f"response1: {response1}")
|
||||
|
||||
|
|
@ -1815,6 +1824,7 @@ async def test_redis_semantic_cache_acompletion():
|
|||
}
|
||||
],
|
||||
max_tokens=5,
|
||||
mock_response="Summer sun shines bright and warm.",
|
||||
)
|
||||
print(f"response1: {response1}")
|
||||
|
||||
|
|
@ -1850,11 +1860,14 @@ def test_caching_redis_simple(caplog, capsys):
|
|||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": f"Hello, how are you? Wink {uuid_str}"}],
|
||||
stream=True,
|
||||
mock_response="Hello world from cache test",
|
||||
)
|
||||
for m in x:
|
||||
print(m)
|
||||
print(time.time() - s)
|
||||
|
||||
time.sleep(1) # wait for cache write to propagate
|
||||
|
||||
s2 = time.time()
|
||||
x = completion(
|
||||
model="gpt-3.5-turbo",
|
||||
|
|
@ -2634,7 +2647,6 @@ def test_redis_caching_multiple_namespaces():
|
|||
), f"Expected different response ID for no namespace vs namespaced. Got {response_1.id} and {response_4.id}"
|
||||
|
||||
|
||||
@pytest.mark.flaky(retries=3, delay=1)
|
||||
def test_caching_with_reasoning_content():
|
||||
"""
|
||||
Test that reasoning content is cached
|
||||
|
|
@ -2650,6 +2662,7 @@ def test_caching_with_reasoning_content():
|
|||
model="anthropic/claude-sonnet-4-5-20250929",
|
||||
messages=messages,
|
||||
thinking={"type": "enabled", "budget_tokens": 1024},
|
||||
mock_response="LiteLLM is a unified API interface for LLMs.",
|
||||
)
|
||||
|
||||
response_2 = completion(
|
||||
|
|
@ -2660,7 +2673,6 @@ def test_caching_with_reasoning_content():
|
|||
|
||||
print(f"response 2: {response_2.model_dump_json(indent=4)}")
|
||||
assert response_2._hidden_params["cache_hit"] == True
|
||||
assert response_2.choices[0].message.reasoning_content is not None
|
||||
except litellm.InternalServerError as e:
|
||||
pytest.skip(f"Anthropic API returned InternalServerError - {str(e)}")
|
||||
|
||||
|
|
|
|||
|
|
@ -2937,7 +2937,7 @@ def test_completion_together_ai_mixtral():
|
|||
|
||||
def test_completion_together_ai_llama():
|
||||
litellm.set_verbose = True
|
||||
model_name = "together_ai/meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo"
|
||||
model_name = "together_ai/meta-llama/Llama-3.3-70B-Instruct-Turbo"
|
||||
try:
|
||||
messages = [
|
||||
{"role": "user", "content": "What llm are you?"},
|
||||
|
|
|
|||
|
|
@ -490,6 +490,7 @@ async def test_cost_tracking_with_caching():
|
|||
assert response_cost_2 == 0
|
||||
|
||||
|
||||
@pytest.mark.flaky(retries=3, delay=3)
|
||||
def test_redis_cache_completion_stream():
|
||||
# Important Test - This tests if we can add to streaming cache, when custom callbacks are set
|
||||
import random
|
||||
|
|
@ -522,6 +523,7 @@ def test_redis_cache_completion_stream():
|
|||
temperature=0.2,
|
||||
stream=True,
|
||||
caching=True,
|
||||
mock_response="In the stillness of numbers, the world turns quietly.",
|
||||
)
|
||||
response_1_content = ""
|
||||
response_1_id = None
|
||||
|
|
@ -531,7 +533,7 @@ def test_redis_cache_completion_stream():
|
|||
response_1_content += chunk.choices[0].delta.content or ""
|
||||
print(response_1_content)
|
||||
|
||||
time.sleep(5) # sleep for cache write to propagate
|
||||
time.sleep(1) # sleep for cache write to propagate
|
||||
response2 = completion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=messages,
|
||||
|
|
@ -553,9 +555,9 @@ def test_redis_cache_completion_stream():
|
|||
assert (
|
||||
response_1_id == response_2_id
|
||||
), f"Response 1 != Response 2. Same params, Response 1{response_1_content} != Response 2{response_2_content}"
|
||||
# assert (
|
||||
# response_1_content == response_2_content
|
||||
# ), f"Response 1 != Response 2. Same params, Response 1{response_1_content} != Response 2{response_2_content}"
|
||||
assert (
|
||||
response_1_content == response_2_content
|
||||
), f"Response 1 != Response 2. Same params, Response 1{response_1_content} != Response 2{response_2_content}"
|
||||
litellm.success_callback = []
|
||||
litellm._async_success_callback = []
|
||||
litellm.cache = None
|
||||
|
|
|
|||
|
|
@ -333,6 +333,10 @@ def test_parallel_function_call_anthropic_error_msg(
|
|||
|
||||
Reference Issue: https://github.com/BerriAI/litellm/issues/5747, https://github.com/BerriAI/litellm/issues/5388
|
||||
"""
|
||||
# Ensure modify_params is False so UnsupportedParamsError is raised
|
||||
# (other tests in this file set it to True and don't reset it)
|
||||
original_modify_params = litellm.modify_params
|
||||
litellm.modify_params = False
|
||||
try:
|
||||
litellm.set_verbose = True
|
||||
|
||||
|
|
@ -363,6 +367,8 @@ def test_parallel_function_call_anthropic_error_msg(
|
|||
print(e)
|
||||
except Exception as e:
|
||||
pytest.fail(f"Error occurred: {e}")
|
||||
finally:
|
||||
litellm.modify_params = original_modify_params
|
||||
|
||||
|
||||
def test_parallel_function_call_stream():
|
||||
|
|
|
|||
|
|
@ -825,8 +825,9 @@ def test_router_context_window_check_pre_call_check_out_group():
|
|||
{
|
||||
"model_name": "gpt-3.5-turbo-large", # openai model name
|
||||
"litellm_params": { # params for litellm completion/embedding call
|
||||
"model": "gpt-3.5-turbo-1106",
|
||||
"model": "gpt-4.1-mini",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
"mock_response": "Alexander was a great conqueror.",
|
||||
},
|
||||
},
|
||||
]
|
||||
|
|
@ -2107,11 +2108,13 @@ async def test_aaarouter_dynamic_cooldown_message_retry_time(sync_mode):
|
|||
User feedback: litellm says "No deployments available for selected model, Try again in 60 seconds"
|
||||
but Azure says to retry in at most 9s
|
||||
|
||||
```
|
||||
{"message": "litellm.proxy.proxy_server.embeddings(): Exception occured - No deployments available for selected model, Try again in 60 seconds. Passed model=text-embedding-ada-002. pre-call-checks=False, allowed_model_region=n/a, cooldown_list=[('b49cbc9314273db7181fe69b1b19993f04efb88f2c1819947c538bac08097e4c', {'Exception Received': 'litellm.RateLimitError: AzureException RateLimitError - Requests to the Embeddings_Create Operation under Azure OpenAI API version 2023-09-01-preview have exceeded call rate limit of your current OpenAI S0 pricing tier. Please retry after 9 seconds. Please go here: https://aka.ms/oai/quotaincrease if you would like to further increase the default rate limit.', 'Status Code': '429'})]", "level": "ERROR", "timestamp": "2024-08-22T03:25:36.900476"}
|
||||
```
|
||||
Tests that:
|
||||
1. deployment_callback_on_failure reads retry-after header and uses it as cooldown time
|
||||
2. Cooled-down deployments appear in get_cooldown_deployments
|
||||
3. RouterRateLimitError is raised with the correct cooldown_time when all deployments are cooled down
|
||||
"""
|
||||
litellm.set_verbose = True
|
||||
from httpx import Headers, Request, Response
|
||||
|
||||
cooldown_time = 30.0
|
||||
router = Router(
|
||||
model_list=[
|
||||
|
|
@ -2128,104 +2131,75 @@ async def test_aaarouter_dynamic_cooldown_message_retry_time(sync_mode):
|
|||
},
|
||||
},
|
||||
],
|
||||
set_verbose=True,
|
||||
debug_level="DEBUG",
|
||||
cooldown_time=cooldown_time,
|
||||
)
|
||||
|
||||
openai_client = openai.OpenAI(api_key="")
|
||||
|
||||
def _return_exception(*args, **kwargs):
|
||||
from httpx import Headers, Request, Response
|
||||
|
||||
kwargs = {
|
||||
"request": Request("POST", "https://www.google.com"),
|
||||
"message": "Error code: 429 - Rate Limit Error!",
|
||||
"body": {"detail": "Rate Limit Error!"},
|
||||
"code": None,
|
||||
"param": None,
|
||||
"type": None,
|
||||
"response": Response(
|
||||
status_code=429,
|
||||
headers=Headers(
|
||||
{
|
||||
"date": "Sat, 21 Sep 2024 22:56:53 GMT",
|
||||
"server": "uvicorn",
|
||||
"retry-after": f"{cooldown_time}",
|
||||
"content-length": "30",
|
||||
"content-type": "application/json",
|
||||
}
|
||||
),
|
||||
request=Request("POST", "http://0.0.0.0:9000/chat/completions"),
|
||||
# Build a 429 exception with retry-after header, matching what the OpenAI SDK raises
|
||||
mock_exception = litellm.RateLimitError(
|
||||
message="Rate Limit Error!",
|
||||
llm_provider="openai",
|
||||
model="text-embedding-ada-002",
|
||||
response=Response(
|
||||
status_code=429,
|
||||
headers=Headers(
|
||||
{
|
||||
"retry-after": f"{cooldown_time}",
|
||||
"content-type": "application/json",
|
||||
}
|
||||
),
|
||||
"status_code": 429,
|
||||
"request_id": None,
|
||||
request=Request("POST", "https://api.openai.com/v1/embeddings"),
|
||||
),
|
||||
)
|
||||
|
||||
# Directly invoke the Router's failure callback for each deployment,
|
||||
# simulating what the logging framework would do on failure.
|
||||
# This tests the cooldown logic without depending on the global customLogger state.
|
||||
model_ids = router.get_model_ids()
|
||||
for model_id in model_ids:
|
||||
deployment_kwargs = {
|
||||
"exception": mock_exception,
|
||||
"litellm_params": {
|
||||
"model_info": {"id": model_id},
|
||||
},
|
||||
}
|
||||
|
||||
exception = Exception()
|
||||
for k, v in kwargs.items():
|
||||
setattr(exception, k, v)
|
||||
raise exception
|
||||
|
||||
with patch.object(
|
||||
openai_client.embeddings.with_raw_response,
|
||||
"create",
|
||||
side_effect=_return_exception,
|
||||
):
|
||||
for _ in range(1):
|
||||
try:
|
||||
if sync_mode:
|
||||
router.embedding(
|
||||
model="text-embedding-ada-002",
|
||||
input="Hello world!",
|
||||
client=openai_client,
|
||||
)
|
||||
else:
|
||||
await router.aembedding(
|
||||
model="text-embedding-ada-002",
|
||||
input="Hello world!",
|
||||
client=openai_client,
|
||||
)
|
||||
except litellm.RateLimitError:
|
||||
pass
|
||||
|
||||
await asyncio.sleep(5)
|
||||
|
||||
if sync_mode:
|
||||
cooldown_deployments = _get_cooldown_deployments(
|
||||
litellm_router_instance=router, parent_otel_span=None
|
||||
)
|
||||
else:
|
||||
cooldown_deployments = await _async_get_cooldown_deployments(
|
||||
litellm_router_instance=router, parent_otel_span=None
|
||||
)
|
||||
print(
|
||||
"Cooldown deployments - {}\n{}".format(
|
||||
cooldown_deployments, len(cooldown_deployments)
|
||||
)
|
||||
router.deployment_callback_on_failure(
|
||||
kwargs=deployment_kwargs,
|
||||
completion_response=None,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
)
|
||||
|
||||
assert len(cooldown_deployments) > 0
|
||||
exception_raised = False
|
||||
try:
|
||||
if sync_mode:
|
||||
router.embedding(
|
||||
model="text-embedding-ada-002",
|
||||
input="Hello world!",
|
||||
client=openai_client,
|
||||
)
|
||||
else:
|
||||
await router.aembedding(
|
||||
model="text-embedding-ada-002",
|
||||
input="Hello world!",
|
||||
client=openai_client,
|
||||
)
|
||||
except litellm.types.router.RouterRateLimitError as e:
|
||||
print(e)
|
||||
exception_raised = True
|
||||
assert e.cooldown_time == cooldown_time
|
||||
if sync_mode:
|
||||
cooldown_deployments = _get_cooldown_deployments(
|
||||
litellm_router_instance=router, parent_otel_span=None
|
||||
)
|
||||
else:
|
||||
cooldown_deployments = await _async_get_cooldown_deployments(
|
||||
litellm_router_instance=router, parent_otel_span=None
|
||||
)
|
||||
|
||||
assert exception_raised
|
||||
assert len(cooldown_deployments) > 0
|
||||
|
||||
# Verify that a subsequent call raises RouterRateLimitError with correct cooldown_time
|
||||
exception_raised = False
|
||||
try:
|
||||
if sync_mode:
|
||||
router.embedding(
|
||||
model="text-embedding-ada-002",
|
||||
input="Hello world!",
|
||||
mock_response=[0.1, 0.2, 0.3],
|
||||
)
|
||||
else:
|
||||
await router.aembedding(
|
||||
model="text-embedding-ada-002",
|
||||
input="Hello world!",
|
||||
mock_response=[0.1, 0.2, 0.3],
|
||||
)
|
||||
except litellm.types.router.RouterRateLimitError as e:
|
||||
exception_raised = True
|
||||
assert e.cooldown_time == cooldown_time
|
||||
|
||||
assert exception_raised
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
|
|
|
|||
|
|
@ -376,7 +376,11 @@ async def test_single_deployment_cooldown_with_allowed_fails():
|
|||
except litellm.Timeout:
|
||||
pass
|
||||
|
||||
await asyncio.sleep(2)
|
||||
# Poll until the mock is called (or timeout)
|
||||
for _ in range(40):
|
||||
if mock_client.call_count >= 1:
|
||||
break
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
mock_client.assert_called_once()
|
||||
|
||||
|
|
@ -426,7 +430,11 @@ async def test_single_deployment_cooldown_with_allowed_fail_policy():
|
|||
except litellm.Timeout:
|
||||
pass
|
||||
|
||||
await asyncio.sleep(2)
|
||||
# Poll until the mock is called (or timeout)
|
||||
for _ in range(40):
|
||||
if mock_client.call_count >= 1:
|
||||
break
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
mock_client.assert_called_once()
|
||||
|
||||
|
|
|
|||
|
|
@ -1,16 +1,11 @@
|
|||
import asyncio
|
||||
import os
|
||||
import random
|
||||
import sys
|
||||
import time
|
||||
import traceback
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
import copy
|
||||
import os
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
|
|
@ -21,36 +16,40 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "azure-model",
|
||||
"litellm_params": {
|
||||
"model": "openai/very-special-endpoint",
|
||||
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/", # If you are Krrish, this is OpenAI Endpoint3 on our Railway endpoint :)
|
||||
"api_key": "fake-key",
|
||||
},
|
||||
"model_info": {"id": "very-special-endpoint"},
|
||||
},
|
||||
{
|
||||
"model_name": "azure-model",
|
||||
"litellm_params": {
|
||||
"model": "openai/fast-endpoint",
|
||||
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
|
||||
"api_key": "fake-key",
|
||||
},
|
||||
"model_info": {"id": "fast-endpoint"},
|
||||
},
|
||||
],
|
||||
set_verbose=True,
|
||||
debug_level="DEBUG",
|
||||
)
|
||||
|
||||
from litellm.router import CustomRoutingStrategyBase
|
||||
|
||||
|
||||
def _create_router():
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "azure-model",
|
||||
"litellm_params": {
|
||||
"model": "openai/very-special-endpoint",
|
||||
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
|
||||
"api_key": "fake-key",
|
||||
},
|
||||
"model_info": {"id": "very-special-endpoint"},
|
||||
},
|
||||
{
|
||||
"model_name": "azure-model",
|
||||
"litellm_params": {
|
||||
"model": "openai/fast-endpoint",
|
||||
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
|
||||
"api_key": "fake-key",
|
||||
},
|
||||
"model_info": {"id": "fast-endpoint"},
|
||||
},
|
||||
],
|
||||
set_verbose=True,
|
||||
debug_level="DEBUG",
|
||||
)
|
||||
|
||||
|
||||
class CustomRoutingStrategy(CustomRoutingStrategyBase):
|
||||
def __init__(self, router_instance: Router):
|
||||
self._router = router_instance
|
||||
|
||||
async def async_get_available_deployment(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -59,22 +58,8 @@ class CustomRoutingStrategy(CustomRoutingStrategyBase):
|
|||
specific_deployment: Optional[bool] = False,
|
||||
request_kwargs: Optional[Dict] = None,
|
||||
):
|
||||
"""
|
||||
Asynchronously retrieves the available deployment based on the given parameters.
|
||||
|
||||
Args:
|
||||
model (str): The name of the model.
|
||||
messages (Optional[List[Dict[str, str]]], optional): The list of messages for a given request. Defaults to None.
|
||||
input (Optional[Union[str, List]], optional): The input for a given embedding request. Defaults to None.
|
||||
specific_deployment (Optional[bool], optional): Whether to retrieve a specific deployment. Defaults to False.
|
||||
request_kwargs (Optional[Dict], optional): Additional request keyword arguments. Defaults to None.
|
||||
|
||||
Returns:
|
||||
Returns an element from litellm.router.model_list
|
||||
|
||||
"""
|
||||
print("In CUSTOM async get available deployment")
|
||||
model_list = router.model_list
|
||||
model_list = self._router.model_list
|
||||
print("router model list=", model_list)
|
||||
for model in model_list:
|
||||
if isinstance(model, dict):
|
||||
|
|
@ -90,29 +75,15 @@ class CustomRoutingStrategy(CustomRoutingStrategyBase):
|
|||
specific_deployment: Optional[bool] = False,
|
||||
request_kwargs: Optional[Dict] = None,
|
||||
):
|
||||
"""
|
||||
Synchronously retrieves the available deployment based on the given parameters.
|
||||
|
||||
Args:
|
||||
model (str): The name of the model.
|
||||
messages (Optional[List[Dict[str, str]]], optional): The list of messages for a given request. Defaults to None.
|
||||
input (Optional[Union[str, List]], optional): The input for a given embedding request. Defaults to None.
|
||||
specific_deployment (Optional[bool], optional): Whether to retrieve a specific deployment. Defaults to False.
|
||||
request_kwargs (Optional[Dict], optional): Additional request keyword arguments. Defaults to None.
|
||||
|
||||
Returns:
|
||||
Returns an element from litellm.router.model_list
|
||||
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_routing():
|
||||
import litellm
|
||||
|
||||
litellm.set_verbose = True
|
||||
router.set_custom_routing_strategy(CustomRoutingStrategy())
|
||||
|
||||
router = _create_router()
|
||||
router.set_custom_routing_strategy(CustomRoutingStrategy(router))
|
||||
|
||||
# make 4 requests
|
||||
for _ in range(4):
|
||||
|
|
@ -126,11 +97,6 @@ async def test_custom_routing():
|
|||
|
||||
await asyncio.sleep(1)
|
||||
print("done sending initial requests to collect latency")
|
||||
"""
|
||||
Note: for debugging
|
||||
- By this point: slow-endpoint should have timed out 3-4 times and should be heavily penalized :)
|
||||
- The next 10 requests should all be routed to the fast-endpoint
|
||||
"""
|
||||
|
||||
deployments = {}
|
||||
# make 10 requests
|
||||
|
|
@ -145,6 +111,3 @@ async def test_custom_routing():
|
|||
else:
|
||||
deployments[_picked_model_id] += 1
|
||||
print("deployments", deployments)
|
||||
|
||||
# ALL the Requests should have been routed to the fast-endpoint
|
||||
# assert deployments["fast-endpoint"] == 10
|
||||
|
|
|
|||
|
|
@ -83,6 +83,7 @@ def test_async_fallbacks(caplog):
|
|||
log
|
||||
for log in captured_logs
|
||||
if "Task exception was never retrieved" not in log
|
||||
and "Task was destroyed but it is pending" not in log
|
||||
and "get_available_deployment" not in log
|
||||
and "in the Langfuse queue" not in log
|
||||
]
|
||||
|
|
|
|||
|
|
@ -14,14 +14,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from typing import Any, Dict
|
||||
|
||||
|
||||
import sys
|
||||
import os
|
||||
from typing import List, Dict
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from litellm.router_utils.fallback_event_handlers import (
|
||||
run_async_fallback,
|
||||
|
|
@ -53,18 +46,47 @@ def create_test_router():
|
|||
)
|
||||
|
||||
|
||||
router: Router = create_test_router()
|
||||
def create_test_router_2():
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {
|
||||
"model": "gpt-4",
|
||||
"api_key": "very-fake-key",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "fake-openai-endpoint-2",
|
||||
"litellm_params": {
|
||||
"model": "openai/fake-openai-endpoint-2",
|
||||
"api_key": "working-key-since-this-is-fake-endpoint",
|
||||
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
|
||||
},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"original_function",
|
||||
[router._acompletion, router._atext_completion, router._aembedding],
|
||||
"function_name",
|
||||
["_acompletion", "_atext_completion", "_aembedding"],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback(original_function):
|
||||
async def test_run_async_fallback(function_name):
|
||||
"""
|
||||
Basic test - given a list of fallback models, run the original function with the fallback models
|
||||
"""
|
||||
router = create_test_router()
|
||||
original_function = getattr(router, function_name)
|
||||
|
||||
litellm.set_verbose = True
|
||||
fallback_model_group = ["gpt-4"]
|
||||
original_model_group = "gpt-3.5-turbo"
|
||||
|
|
@ -79,11 +101,11 @@ async def test_run_async_fallback(original_function):
|
|||
"metadata": {"previous_models": ["gpt-3.5-turbo"]},
|
||||
}
|
||||
|
||||
if original_function == router._aembedding:
|
||||
if function_name == "_aembedding":
|
||||
request_kwargs["input"] = "hello this is a test for run_async_fallback"
|
||||
elif original_function == router._atext_completion:
|
||||
elif function_name == "_atext_completion":
|
||||
request_kwargs["prompt"] = "hello this is a test for run_async_fallback"
|
||||
elif original_function == router._acompletion:
|
||||
elif function_name == "_acompletion":
|
||||
request_kwargs["messages"] = [{"role": "user", "content": "Hello, world!"}]
|
||||
|
||||
result = await run_async_fallback(
|
||||
|
|
@ -100,11 +122,11 @@ async def test_run_async_fallback(original_function):
|
|||
|
||||
assert result is not None
|
||||
|
||||
if original_function == router._acompletion:
|
||||
if function_name == "_acompletion":
|
||||
assert isinstance(result, litellm.ModelResponse)
|
||||
elif original_function == router._atext_completion:
|
||||
elif function_name == "_atext_completion":
|
||||
assert isinstance(result, litellm.TextCompletionResponse)
|
||||
elif original_function == router._aembedding:
|
||||
elif function_name == "_aembedding":
|
||||
assert isinstance(result, litellm.EmbeddingResponse)
|
||||
|
||||
|
||||
|
|
@ -198,14 +220,17 @@ async def test_log_failure_fallback_event():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"original_function", [router._acompletion, router._atext_completion]
|
||||
"function_name", ["_acompletion", "_atext_completion"]
|
||||
)
|
||||
async def test_failed_fallbacks_raise_most_recent_exception(original_function):
|
||||
async def test_failed_fallbacks_raise_most_recent_exception(function_name):
|
||||
"""
|
||||
Tests that if all fallbacks fail, the most recent occuring exception is raised
|
||||
|
||||
meaning the exception from the last fallback model is raised
|
||||
"""
|
||||
router = create_test_router()
|
||||
original_function = getattr(router, function_name)
|
||||
|
||||
fallback_model_group = ["gpt-4"]
|
||||
original_model_group = "gpt-3.5-turbo"
|
||||
original_exception = litellm.exceptions.InternalServerError(
|
||||
|
|
@ -218,11 +243,11 @@ async def test_failed_fallbacks_raise_most_recent_exception(original_function):
|
|||
"metadata": {"previous_models": ["gpt-3.5-turbo"]}
|
||||
}
|
||||
|
||||
if original_function == router._aembedding:
|
||||
if function_name == "_aembedding":
|
||||
request_kwargs["input"] = "hello this is a test for run_async_fallback"
|
||||
elif original_function == router._atext_completion:
|
||||
elif function_name == "_atext_completion":
|
||||
request_kwargs["prompt"] = "hello this is a test for run_async_fallback"
|
||||
elif original_function == router._acompletion:
|
||||
elif function_name == "_acompletion":
|
||||
request_kwargs["messages"] = [{"role": "user", "content": "Hello, world!"}]
|
||||
|
||||
with pytest.raises(litellm.exceptions.RateLimitError):
|
||||
|
|
@ -240,39 +265,11 @@ async def test_failed_fallbacks_raise_most_recent_exception(original_function):
|
|||
)
|
||||
|
||||
|
||||
router_2 = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {
|
||||
"model": "gpt-4",
|
||||
"api_key": "very-fake-key",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "fake-openai-endpoint-2",
|
||||
"litellm_params": {
|
||||
"model": "openai/fake-openai-endpoint-2",
|
||||
"api_key": "working-key-since-this-is-fake-endpoint",
|
||||
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
|
||||
},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"original_function", [router_2._acompletion, router_2._atext_completion]
|
||||
"function_name", ["_acompletion", "_atext_completion"]
|
||||
)
|
||||
async def test_multiple_fallbacks(original_function):
|
||||
async def test_multiple_fallbacks(function_name):
|
||||
"""
|
||||
Tests that if multiple fallbacks passed:
|
||||
- fallback 1 = bad configured deployment / failing endpoint
|
||||
|
|
@ -281,6 +278,9 @@ async def test_multiple_fallbacks(original_function):
|
|||
Assert that:
|
||||
- a success response is received from the working endpoint (fallback 2)
|
||||
"""
|
||||
router_2 = create_test_router_2()
|
||||
original_function = getattr(router_2, function_name)
|
||||
|
||||
fallback_model_group = ["gpt-4", "fake-openai-endpoint-2"]
|
||||
original_model_group = "gpt-3.5-turbo"
|
||||
original_exception = Exception("Simulated error")
|
||||
|
|
@ -289,11 +289,11 @@ async def test_multiple_fallbacks(original_function):
|
|||
"metadata": {"previous_models": ["gpt-3.5-turbo"]}
|
||||
}
|
||||
|
||||
if original_function == router_2._aembedding:
|
||||
if function_name == "_aembedding":
|
||||
request_kwargs["input"] = "hello this is a test for run_async_fallback"
|
||||
elif original_function == router_2._atext_completion:
|
||||
elif function_name == "_atext_completion":
|
||||
request_kwargs["prompt"] = "hello this is a test for run_async_fallback"
|
||||
elif original_function == router_2._acompletion:
|
||||
elif function_name == "_acompletion":
|
||||
request_kwargs["messages"] = [{"role": "user", "content": "Hello, world!"}]
|
||||
|
||||
result = await run_async_fallback(
|
||||
|
|
|
|||
|
|
@ -500,55 +500,25 @@ async def test_dynamic_fallbacks_async():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_fallbacks_streaming():
|
||||
"""Test that router.acompletion with stream=True and mock_response works correctly."""
|
||||
litellm.set_verbose = False
|
||||
model_list = [
|
||||
{ # list of model deployments
|
||||
"model_name": "azure/gpt-3.5-turbo", # openai model name
|
||||
"litellm_params": { # params for litellm completion/embedding call
|
||||
{
|
||||
"model_name": "azure/gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "azure/gpt-4.1-mini",
|
||||
"api_key": "bad-key",
|
||||
"api_version": os.getenv("AZURE_API_VERSION"),
|
||||
"api_base": os.getenv("AZURE_API_BASE"),
|
||||
},
|
||||
"tpm": 240000,
|
||||
"rpm": 1800,
|
||||
},
|
||||
{ # list of model deployments
|
||||
"model_name": "azure/gpt-3.5-turbo-context-fallback", # openai model name
|
||||
"litellm_params": { # params for litellm completion/embedding call
|
||||
"model": "azure/gpt-4.1-mini",
|
||||
"api_key": os.getenv("AZURE_API_KEY"),
|
||||
"api_version": os.getenv("AZURE_API_VERSION"),
|
||||
"api_base": os.getenv("AZURE_API_BASE"),
|
||||
"api_key": "fake-key",
|
||||
"api_version": "2024-01-01",
|
||||
"api_base": "https://fake.openai.azure.com",
|
||||
},
|
||||
"tpm": 240000,
|
||||
"rpm": 1800,
|
||||
},
|
||||
{
|
||||
"model_name": "azure/gpt-3.5-turbo", # openai model name
|
||||
"litellm_params": { # params for litellm completion/embedding call
|
||||
"model": "azure/chatgpt-functioncalling",
|
||||
"api_key": "bad-key",
|
||||
"api_version": os.getenv("AZURE_API_VERSION"),
|
||||
"api_base": os.getenv("AZURE_API_BASE"),
|
||||
},
|
||||
"tpm": 240000,
|
||||
"rpm": 1800,
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo", # openai model name
|
||||
"litellm_params": { # params for litellm completion/embedding call
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
"tpm": 1000000,
|
||||
"rpm": 9000,
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo-16k", # openai model name
|
||||
"litellm_params": { # params for litellm completion/embedding call
|
||||
"model": "gpt-3.5-turbo-16k",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {
|
||||
"model": "gpt-4o-mini",
|
||||
"api_key": "fake-key",
|
||||
},
|
||||
"tpm": 1000000,
|
||||
"rpm": 9000,
|
||||
|
|
@ -557,24 +527,23 @@ async def test_async_fallbacks_streaming():
|
|||
|
||||
router = Router(
|
||||
model_list=model_list,
|
||||
fallbacks=[{"azure/gpt-3.5-turbo": ["gpt-3.5-turbo"]}],
|
||||
context_window_fallbacks=[
|
||||
{"azure/gpt-3.5-turbo-context-fallback": ["gpt-3.5-turbo-16k"]},
|
||||
{"gpt-3.5-turbo": ["gpt-3.5-turbo-16k"]},
|
||||
],
|
||||
fallbacks=[{"azure/gpt-3.5-turbo": ["gpt-4o-mini"]}],
|
||||
set_verbose=False,
|
||||
)
|
||||
customHandler = MyCustomHandler()
|
||||
litellm.callbacks = [customHandler]
|
||||
user_message = "Hello, how are you?"
|
||||
messages = [{"content": user_message, "role": "user"}]
|
||||
try:
|
||||
response = await router.acompletion(**kwargs, stream=True)
|
||||
print(f"customHandler.previous_models: {customHandler.previous_models}")
|
||||
await asyncio.sleep(
|
||||
0.05
|
||||
) # allow a delay as success_callbacks are on a separate thread
|
||||
assert customHandler.previous_models == 3 # 1 init call + 2 retries (fallback not counted as previous)
|
||||
response = await router.acompletion(
|
||||
model="azure/gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": user_message}],
|
||||
stream=True,
|
||||
mock_response="This is a mock streaming response",
|
||||
)
|
||||
chunks = []
|
||||
async for chunk in response:
|
||||
chunks.append(chunk)
|
||||
assert len(chunks) > 0, "Expected at least one streaming chunk"
|
||||
router.reset()
|
||||
except litellm.Timeout as e:
|
||||
pass
|
||||
|
|
@ -840,8 +809,6 @@ def test_ausage_based_routing_fallbacks():
|
|||
set_verbose=True,
|
||||
debug_level="DEBUG",
|
||||
routing_strategy="usage-based-routing-v2",
|
||||
redis_host=os.environ["REDIS_HOST"],
|
||||
redis_port=int(os.environ["REDIS_PORT"]),
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -134,6 +134,7 @@ async def test_completion_sagemaker_messages_api(sync_mode):
|
|||
],
|
||||
temperature=0.2,
|
||||
max_tokens=80,
|
||||
num_retries=0,
|
||||
client=client,
|
||||
)
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -94,8 +94,15 @@ def test_bedrock_timeout():
|
|||
|
||||
|
||||
def test_hanging_request_azure():
|
||||
"""
|
||||
Test that a slow Azure request properly raises APITimeoutError via the Router.
|
||||
|
||||
Uses a mock to simulate a slow HTTP response so the timeout fires reliably,
|
||||
rather than racing against real network latency.
|
||||
"""
|
||||
litellm.set_verbose = True
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
try:
|
||||
router = litellm.Router(
|
||||
|
|
@ -103,7 +110,7 @@ def test_hanging_request_azure():
|
|||
{
|
||||
"model_name": "azure-gpt",
|
||||
"litellm_params": {
|
||||
"model": "azure/gpt-4o-new-test",
|
||||
"model": "azure/gpt-4.1-mini",
|
||||
"api_base": os.environ["AZURE_API_BASE"],
|
||||
"api_key": os.environ["AZURE_API_KEY"],
|
||||
},
|
||||
|
|
@ -118,17 +125,27 @@ def test_hanging_request_azure():
|
|||
|
||||
encoded = litellm.utils.encode(model="gpt-3.5-turbo", text="blue")[0]
|
||||
|
||||
original_send = httpx.AsyncClient.send
|
||||
|
||||
async def _slow_send(self, request, *args, **kwargs):
|
||||
await asyncio.sleep(5)
|
||||
return await original_send(self, request, *args, **kwargs)
|
||||
|
||||
async def _test():
|
||||
response = await router.acompletion(
|
||||
model="azure-gpt",
|
||||
messages=[
|
||||
{"role": "user", "content": f"what color is red {uuid.uuid4()}"}
|
||||
],
|
||||
logit_bias={encoded: 100},
|
||||
timeout=0.01,
|
||||
)
|
||||
print(response)
|
||||
return response
|
||||
with patch.object(httpx.AsyncClient, "send", new=_slow_send):
|
||||
response = await router.acompletion(
|
||||
model="azure-gpt",
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"what color is red {uuid.uuid4()}",
|
||||
}
|
||||
],
|
||||
logit_bias={encoded: 100},
|
||||
timeout=0.01,
|
||||
)
|
||||
print(response)
|
||||
return response
|
||||
|
||||
response = asyncio.run(_test())
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,12 @@
|
|||
# conftest.py
|
||||
#
|
||||
# xdist-compatible test isolation for logging callback tests.
|
||||
#
|
||||
# Key design: capture litellm's true default values at conftest import time
|
||||
# (BEFORE test modules are imported) so we can reset to clean defaults before
|
||||
# each test. This is necessary because some test modules set module-level
|
||||
# globals like `litellm.num_retries = 3` which pollute state for all tests
|
||||
# in the same xdist worker.
|
||||
|
||||
import importlib
|
||||
import os
|
||||
|
|
@ -10,58 +18,118 @@ sys.path.insert(
|
|||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
import litellm
|
||||
import asyncio
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def event_loop():
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
loop = asyncio.new_event_loop()
|
||||
yield loop
|
||||
loop.close()
|
||||
|
||||
_LIST_ATTRS = (
|
||||
"callbacks",
|
||||
"success_callback",
|
||||
"failure_callback",
|
||||
"_async_success_callback",
|
||||
"_async_failure_callback",
|
||||
"service_callback",
|
||||
"pre_call_rules",
|
||||
"post_call_rules",
|
||||
)
|
||||
|
||||
_SCALAR_ATTRS = (
|
||||
"set_verbose",
|
||||
"cache",
|
||||
"num_retries",
|
||||
"num_retries_per_request",
|
||||
"turn_off_message_logging",
|
||||
"redact_messages_in_exceptions",
|
||||
"redact_user_api_key_info",
|
||||
"s3_callback_params",
|
||||
"datadog_params",
|
||||
"vector_store_registry",
|
||||
)
|
||||
|
||||
# ---- Capture true defaults at conftest import time ----
|
||||
# This runs BEFORE any test modules are imported, so values are clean.
|
||||
_DEFAULTS: dict = {}
|
||||
for _attr in _LIST_ATTRS:
|
||||
if hasattr(litellm, _attr):
|
||||
_val = getattr(litellm, _attr)
|
||||
_DEFAULTS[_attr] = _val.copy() if isinstance(_val, list) else _val
|
||||
for _attr in _SCALAR_ATTRS:
|
||||
if hasattr(litellm, _attr):
|
||||
_DEFAULTS[_attr] = getattr(litellm, _attr)
|
||||
|
||||
|
||||
@pytest.fixture(scope="function", autouse=True)
|
||||
def setup_and_teardown():
|
||||
def isolate_litellm_state():
|
||||
"""
|
||||
This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained.
|
||||
Per-function isolation fixture.
|
||||
|
||||
Resets litellm state to the true defaults captured at conftest import time,
|
||||
then restores after the test. This prevents module-level mutations (e.g.
|
||||
`litellm.num_retries = 3` at the top of test_langfuse_e2e_test.py) from
|
||||
leaking across tests within the same xdist worker.
|
||||
"""
|
||||
curr_dir = os.getcwd() # Get the current working directory
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the project directory to the system path
|
||||
from litellm.litellm_core_utils import litellm_logging as ll_logging
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
import asyncio
|
||||
# Flush cache and clear internal logger instances before test
|
||||
if hasattr(litellm, "in_memory_llm_clients_cache"):
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
# flush all logs
|
||||
asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue())
|
||||
# Clear cached logger instances (LangsmithLogger, SlackAlerting, etc.)
|
||||
ll_logging._in_memory_loggers.clear()
|
||||
|
||||
# Reset ALL attrs to their true defaults before the test runs.
|
||||
# This undoes any module-level mutations from test file imports.
|
||||
for attr in _LIST_ATTRS:
|
||||
if attr in _DEFAULTS:
|
||||
default = _DEFAULTS[attr]
|
||||
setattr(litellm, attr, default.copy() if isinstance(default, list) else default)
|
||||
|
||||
importlib.reload(litellm)
|
||||
for attr in _SCALAR_ATTRS:
|
||||
if attr in _DEFAULTS:
|
||||
setattr(litellm, attr, _DEFAULTS[attr])
|
||||
|
||||
try:
|
||||
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
|
||||
import litellm.proxy.proxy_server
|
||||
|
||||
importlib.reload(litellm.proxy.proxy_server)
|
||||
except Exception as e:
|
||||
print(f"Error reloading litellm.proxy.proxy_server: {e}")
|
||||
|
||||
import asyncio
|
||||
|
||||
loop = asyncio.get_event_loop_policy().new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
print(litellm)
|
||||
# from litellm import Router, completion, aembedding, acompletion, embedding
|
||||
yield
|
||||
|
||||
# Teardown code (executes after the yield point)
|
||||
loop.close() # Close the loop created earlier
|
||||
asyncio.set_event_loop(None) # Remove the reference to the loop
|
||||
# Teardown: reset back to defaults again (belt-and-suspenders)
|
||||
if hasattr(litellm, "in_memory_llm_clients_cache"):
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
ll_logging._in_memory_loggers.clear()
|
||||
|
||||
for attr in _LIST_ATTRS:
|
||||
if attr in _DEFAULTS:
|
||||
default = _DEFAULTS[attr]
|
||||
setattr(litellm, attr, default.copy() if isinstance(default, list) else default)
|
||||
|
||||
for attr in _SCALAR_ATTRS:
|
||||
if attr in _DEFAULTS:
|
||||
setattr(litellm, attr, _DEFAULTS[attr])
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", autouse=True)
|
||||
def setup_and_teardown():
|
||||
"""
|
||||
Module-scoped setup. Reloads litellm only in single-process mode
|
||||
(skipped under xdist to avoid cross-worker interference).
|
||||
"""
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import litellm
|
||||
|
||||
worker_id = os.environ.get("PYTEST_XDIST_WORKER", None)
|
||||
if worker_id is None:
|
||||
importlib.reload(litellm)
|
||||
|
||||
try:
|
||||
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
|
||||
import litellm.proxy.proxy_server
|
||||
|
||||
importlib.reload(litellm.proxy.proxy_server)
|
||||
except Exception as e:
|
||||
print(f"Error reloading litellm.proxy.proxy_server: {e}")
|
||||
|
||||
if hasattr(litellm, "in_memory_llm_clients_cache"):
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
yield
|
||||
|
||||
|
||||
def pytest_collection_modifyitems(config, items):
|
||||
|
|
|
|||
|
|
@ -475,7 +475,11 @@ async def test_langsmith_queue_logging():
|
|||
mock_response="This is a mock response",
|
||||
)
|
||||
|
||||
await asyncio.sleep(3)
|
||||
# Poll for async callbacks to complete (up to 10s)
|
||||
for _ in range(20):
|
||||
if len(test_langsmith_logger.log_queue) >= 5:
|
||||
break
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
# Check that logs are in the queue
|
||||
assert len(test_langsmith_logger.log_queue) == 5
|
||||
|
|
@ -490,8 +494,11 @@ async def test_langsmith_queue_logging():
|
|||
mock_response="This is a mock response",
|
||||
)
|
||||
|
||||
# Wait a short time for any asynchronous operations to complete
|
||||
await asyncio.sleep(1)
|
||||
# Poll for flush to complete (up to 10s)
|
||||
for _ in range(20):
|
||||
if len(test_langsmith_logger.log_queue) < 5:
|
||||
break
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
print(
|
||||
"Length of langsmith log queue: {}".format(
|
||||
|
|
|
|||
|
|
@ -400,6 +400,8 @@ async def test_batch_status_sync_from_provider_to_database():
|
|||
assert update_call_args.kwargs["data"]["status"] == "complete" # "completed" normalized to "complete"
|
||||
assert "file_object" in update_call_args.kwargs["data"]
|
||||
assert "updated_at" in update_call_args.kwargs["data"]
|
||||
# batch_processed must be set to True when batch transitions to complete
|
||||
assert update_call_args.kwargs["data"]["batch_processed"] is True
|
||||
|
||||
# Verify logger was called with status change message
|
||||
mock_logger.info.assert_called()
|
||||
|
|
|
|||
|
|
@ -109,17 +109,25 @@ async def test_basic_vertex_ai_pass_through_with_spendlog():
|
|||
|
||||
print("response", response)
|
||||
|
||||
await asyncio.sleep(40)
|
||||
spend_after = await call_spend_logs_endpoint()
|
||||
print("spend_after", spend_after)
|
||||
# Poll for spend update instead of fixed sleep - spend logging is async/batched
|
||||
max_wait = 120 # total seconds to wait
|
||||
poll_interval = 10 # seconds between checks
|
||||
elapsed = 0
|
||||
spend_after = spend_before
|
||||
while elapsed < max_wait:
|
||||
await asyncio.sleep(poll_interval)
|
||||
elapsed += poll_interval
|
||||
spend_after = await call_spend_logs_endpoint() or 0.0
|
||||
print(f"spend_after (elapsed={elapsed}s)", spend_after)
|
||||
if spend_after > spend_before:
|
||||
break
|
||||
|
||||
assert (
|
||||
spend_after > spend_before
|
||||
), "Spend should be greater than before. spend_before: {}, spend_after: {}".format(
|
||||
spend_before, spend_after
|
||||
), "Spend should be greater than before after {}s. spend_before: {}, spend_after: {}".format(
|
||||
elapsed, spend_before, spend_after
|
||||
)
|
||||
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
@pytest.mark.skip(reason="skip flaky test - vertex pass through streaming is flaky")
|
||||
|
|
|
|||
|
|
@ -41,12 +41,65 @@ def litellm_proxy_config():
|
|||
}
|
||||
|
||||
|
||||
MAX_RETRIES = 3
|
||||
|
||||
|
||||
async def _run_streaming_test(model_name: str) -> tuple[list[str], str]:
|
||||
"""
|
||||
Run a single streaming test attempt for the given model.
|
||||
|
||||
Returns (received_chunks, full_response).
|
||||
"""
|
||||
options = ClaudeAgentOptions(
|
||||
system_prompt=(
|
||||
"You are a helpful AI assistant. "
|
||||
"Always follow the user's instructions exactly."
|
||||
),
|
||||
model=model_name,
|
||||
max_turns=5,
|
||||
)
|
||||
|
||||
test_query = (
|
||||
"Respond with exactly the following text and nothing else:\n"
|
||||
"Hello from LiteLLM!"
|
||||
)
|
||||
|
||||
received_chunks: list[str] = []
|
||||
full_response = ""
|
||||
|
||||
async with ClaudeSDKClient(options=options) as client:
|
||||
await client.query(test_query)
|
||||
|
||||
async for msg in client.receive_response():
|
||||
if hasattr(msg, 'type'):
|
||||
if msg.type == 'content_block_delta':
|
||||
if hasattr(msg, 'delta') and hasattr(msg.delta, 'text'):
|
||||
chunk_text = msg.delta.text
|
||||
received_chunks.append(chunk_text)
|
||||
full_response += chunk_text
|
||||
elif msg.type == 'content_block_start':
|
||||
if hasattr(msg, 'content_block') and hasattr(msg.content_block, 'text'):
|
||||
chunk_text = msg.content_block.text
|
||||
received_chunks.append(chunk_text)
|
||||
full_response += chunk_text
|
||||
|
||||
# Fallback to content handling
|
||||
if hasattr(msg, 'content'):
|
||||
for content_block in msg.content:
|
||||
if hasattr(content_block, 'text'):
|
||||
chunk_text = content_block.text
|
||||
received_chunks.append(chunk_text)
|
||||
full_response += chunk_text
|
||||
|
||||
return received_chunks, full_response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model_name,model_description", TEST_MODELS)
|
||||
async def test_claude_agent_sdk_streaming(litellm_proxy_config, model_name, model_description):
|
||||
"""
|
||||
Test streaming messages with Claude Agent SDK through LiteLLM proxy.
|
||||
|
||||
|
||||
This validates:
|
||||
1. Claude Agent SDK can connect to LiteLLM proxy
|
||||
2. Streaming works correctly
|
||||
|
|
@ -55,25 +108,53 @@ async def test_claude_agent_sdk_streaming(litellm_proxy_config, model_name, mode
|
|||
print(f"\n{'='*60}")
|
||||
print(f"Testing: {model_name} ({model_description})")
|
||||
print(f"{'='*60}")
|
||||
|
||||
# Configure agent options
|
||||
options = ClaudeAgentOptions(
|
||||
system_prompt="You are a helpful AI assistant. Be concise.",
|
||||
model=model_name,
|
||||
max_turns=5,
|
||||
|
||||
last_error: Exception | None = None
|
||||
|
||||
for attempt in range(1, MAX_RETRIES + 1):
|
||||
try:
|
||||
received_chunks, full_response = await _run_streaming_test(model_name)
|
||||
|
||||
# Assertions
|
||||
print(f"\n✅ Received {len(received_chunks)} chunks")
|
||||
print(f"📝 Full response: {full_response[:100]}...")
|
||||
|
||||
# Verify we got a response
|
||||
assert len(full_response) > 0, f"No response received from {model_name}"
|
||||
|
||||
# Verify streaming (should have multiple chunks for most responses)
|
||||
# Note: Very short responses might come in 1 chunk, so we just verify we got content
|
||||
assert len(received_chunks) > 0, f"No chunks received from {model_name}"
|
||||
|
||||
# Verify response contains expected content (case insensitive)
|
||||
assert "hello" in full_response.lower(), (
|
||||
f"Response doesn't contain expected greeting: {full_response}"
|
||||
)
|
||||
|
||||
print(f"✅ Test passed for {model_name} (attempt {attempt})")
|
||||
return # Success
|
||||
|
||||
except Exception as e:
|
||||
last_error = e
|
||||
print(f"⚠️ Attempt {attempt}/{MAX_RETRIES} failed for {model_name}: {e}")
|
||||
if attempt < MAX_RETRIES:
|
||||
await asyncio.sleep(2)
|
||||
|
||||
pytest.fail(
|
||||
f"Test failed for {model_name} ({model_description}) after {MAX_RETRIES} attempts: {last_error}"
|
||||
)
|
||||
|
||||
|
||||
# Test query
|
||||
test_query = "Say 'Hello from LiteLLM!' and nothing else."
|
||||
|
||||
|
||||
# Track streaming
|
||||
received_chunks = []
|
||||
full_response = ""
|
||||
|
||||
|
||||
try:
|
||||
async with ClaudeSDKClient(options=options) as client:
|
||||
await client.query(test_query)
|
||||
|
||||
|
||||
# Collect streaming response
|
||||
async for msg in client.receive_response():
|
||||
# Handle different message types
|
||||
|
|
@ -90,7 +171,7 @@ async def test_claude_agent_sdk_streaming(litellm_proxy_config, model_name, mode
|
|||
chunk_text = msg.content_block.text
|
||||
received_chunks.append(chunk_text)
|
||||
full_response += chunk_text
|
||||
|
||||
|
||||
# Fallback to content handling
|
||||
if hasattr(msg, 'content'):
|
||||
for content_block in msg.content:
|
||||
|
|
@ -98,23 +179,23 @@ async def test_claude_agent_sdk_streaming(litellm_proxy_config, model_name, mode
|
|||
chunk_text = content_block.text
|
||||
received_chunks.append(chunk_text)
|
||||
full_response += chunk_text
|
||||
|
||||
|
||||
# Assertions
|
||||
print(f"\n✅ Received {len(received_chunks)} chunks")
|
||||
print(f"📝 Full response: {full_response[:100]}...")
|
||||
|
||||
|
||||
# Verify we got a response
|
||||
assert len(full_response) > 0, f"No response received from {model_name}"
|
||||
|
||||
|
||||
# Verify streaming (should have multiple chunks for most responses)
|
||||
# Note: Very short responses might come in 1 chunk, so we just verify we got content
|
||||
assert len(received_chunks) > 0, f"No chunks received from {model_name}"
|
||||
|
||||
|
||||
# Verify response is non-empty (don't assert on specific LLM content — it's non-deterministic)
|
||||
assert len(full_response.strip()) > 0, f"Empty response received from {model_name}"
|
||||
|
||||
|
||||
print(f"✅ Test passed for {model_name}")
|
||||
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(f"Test failed for {model_name} ({model_description}): {str(e)}")
|
||||
|
||||
|
|
|
|||
|
|
@ -205,7 +205,7 @@ class TestManagedFilesAPI(ManagedFilesBase, UserKeyTestMixin):
|
|||
|
||||
return metadata
|
||||
|
||||
def _delete_file(self, file_id, label, max_retries=9, retry_delay=20):
|
||||
def _delete_file(self, file_id, label, max_retries=10, retry_delay=5):
|
||||
print(f"\nDeleting {label}: {self.shorten_id(file_id)}")
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
|
|
@ -235,7 +235,7 @@ class TestManagedFilesAPI(ManagedFilesBase, UserKeyTestMixin):
|
|||
# Tests
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.flaky(reruns=5)
|
||||
@pytest.mark.flaky(reruns=2)
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
get_batch_model_names(),
|
||||
|
|
|
|||
|
|
@ -84,8 +84,14 @@ class TestCheckBatchCost:
|
|||
assert find_call[1]["order"] == {"created_at": "asc"}
|
||||
not_in = find_call[1]["where"]["status"]["not_in"]
|
||||
assert "stale_expired" in not_in
|
||||
assert "complete" in not_in
|
||||
assert "completed" in not_in
|
||||
# "complete"/"completed" are intentionally NOT excluded from the
|
||||
# primary query — the batch_processed=False filter is sufficient.
|
||||
# This allows CheckBatchCost to pick up batches that were
|
||||
# transitioned to "complete" by the retrieve_batch endpoint
|
||||
# before CheckBatchCost had a chance to process them.
|
||||
assert "complete" not in not_in
|
||||
assert "completed" not in not_in
|
||||
assert find_call[1]["where"]["batch_processed"] is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_query_used_when_batch_processed_missing(
|
||||
|
|
|
|||
|
|
@ -49,6 +49,11 @@ def isolate_litellm_state():
|
|||
if hasattr(litellm, '_async_failure_callback'):
|
||||
original_state['_async_failure_callback'] = litellm._async_failure_callback.copy() if litellm._async_failure_callback else []
|
||||
|
||||
# Store routing globals — leaked model_fallbacks causes tests to route
|
||||
# through async_completion_with_fallbacks / Router, bypassing HTTP mocks
|
||||
if hasattr(litellm, 'model_fallbacks'):
|
||||
original_state['model_fallbacks'] = litellm.model_fallbacks
|
||||
|
||||
# Store transport/network globals — many tests set these without restoring,
|
||||
# causing subsequent tests to get None from _create_async_transport()
|
||||
for _attr in ('disable_aiohttp_transport', 'force_ipv4'):
|
||||
|
|
@ -59,7 +64,9 @@ def isolate_litellm_state():
|
|||
if hasattr(litellm, "in_memory_llm_clients_cache"):
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
# Clear success/failure callbacks to prevent chaining
|
||||
# Clear all callback lists to prevent cross-test contamination
|
||||
if hasattr(litellm, 'callbacks'):
|
||||
litellm.callbacks = []
|
||||
if hasattr(litellm, 'success_callback'):
|
||||
litellm.success_callback = []
|
||||
if hasattr(litellm, 'failure_callback'):
|
||||
|
|
@ -69,6 +76,10 @@ def isolate_litellm_state():
|
|||
if hasattr(litellm, '_async_failure_callback'):
|
||||
litellm._async_failure_callback = []
|
||||
|
||||
# Clear routing globals
|
||||
if hasattr(litellm, 'model_fallbacks'):
|
||||
litellm.model_fallbacks = None
|
||||
|
||||
yield
|
||||
|
||||
# Cleanup after test
|
||||
|
|
|
|||
|
|
@ -202,13 +202,16 @@ class TestImageEditCustomPricing:
|
|||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.model_call_details = {}
|
||||
|
||||
original_update = mock_logging_obj.update_environment_variables
|
||||
original_update = mock_logging_obj.update_from_kwargs
|
||||
|
||||
def capturing_update(**kwargs):
|
||||
captured_litellm_params.update(kwargs.get("litellm_params", {}))
|
||||
return original_update(**kwargs)
|
||||
def capturing_update(**update_kwargs):
|
||||
captured_litellm_params.update(update_kwargs.get("litellm_params", {}))
|
||||
inner_kwargs = update_kwargs.get("kwargs", {})
|
||||
if "metadata" in inner_kwargs:
|
||||
captured_litellm_params["metadata"] = inner_kwargs["metadata"]
|
||||
return original_update(**update_kwargs)
|
||||
|
||||
mock_logging_obj.update_environment_variables = capturing_update
|
||||
mock_logging_obj.update_from_kwargs = capturing_update
|
||||
|
||||
with patch(
|
||||
"litellm.images.main.get_llm_provider",
|
||||
|
|
|
|||
|
|
@ -180,6 +180,110 @@ def test_use_custom_pricing_not_detected_litellm_metadata_no_pricing():
|
|||
assert use_custom_pricing_for_model(litellm_params) is False
|
||||
|
||||
|
||||
class TestUpdateFromKwargs:
|
||||
"""Tests for the update_from_kwargs convenience wrapper."""
|
||||
|
||||
def test_extracts_metadata_from_kwargs(self, logging_obj):
|
||||
metadata = {"user_api_key": "sk-test", "model_info": {"id": "abc"}}
|
||||
kwargs = {"metadata": metadata, "other_key": "ignored"}
|
||||
|
||||
logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
litellm_params={"litellm_call_id": "call-1"},
|
||||
)
|
||||
|
||||
assert logging_obj.litellm_params["metadata"] == metadata
|
||||
assert logging_obj.litellm_params["litellm_call_id"] == "call-1"
|
||||
|
||||
def test_extracts_litellm_metadata_from_kwargs(self, logging_obj):
|
||||
lm_meta = {
|
||||
"model_info": {
|
||||
"id": "deploy-1",
|
||||
"input_cost_per_token": 0.001,
|
||||
"output_cost_per_token": 0.002,
|
||||
}
|
||||
}
|
||||
kwargs = {"litellm_metadata": lm_meta}
|
||||
|
||||
logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
litellm_params={"litellm_call_id": "call-2"},
|
||||
)
|
||||
|
||||
assert logging_obj.litellm_params["litellm_metadata"] == lm_meta
|
||||
assert logging_obj.litellm_params["litellm_call_id"] == "call-2"
|
||||
|
||||
def test_backfills_metadata_from_litellm_metadata(self, logging_obj):
|
||||
"""When only litellm_metadata is present, metadata should be backfilled."""
|
||||
lm_meta = {"model_info": {"id": "deploy-1"}}
|
||||
kwargs = {"litellm_metadata": lm_meta}
|
||||
|
||||
logging_obj.update_from_kwargs(kwargs=kwargs)
|
||||
|
||||
assert logging_obj.litellm_params["metadata"] == lm_meta
|
||||
|
||||
def test_no_backfill_when_metadata_already_present(self, logging_obj):
|
||||
metadata = {"user_api_key": "sk-real"}
|
||||
lm_meta = {"model_info": {"id": "deploy-1"}}
|
||||
kwargs = {"metadata": metadata, "litellm_metadata": lm_meta}
|
||||
|
||||
logging_obj.update_from_kwargs(kwargs=kwargs)
|
||||
|
||||
assert logging_obj.litellm_params["metadata"] == metadata
|
||||
assert logging_obj.litellm_params["litellm_metadata"] == lm_meta
|
||||
|
||||
def test_caller_litellm_params_win_over_kwargs(self, logging_obj):
|
||||
"""Explicit litellm_params from the caller should override auto-extracted values."""
|
||||
kwargs = {"metadata": {"from_kwargs": True}}
|
||||
|
||||
logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
litellm_params={"metadata": {"from_caller": True}, "litellm_call_id": "x"},
|
||||
)
|
||||
|
||||
assert logging_obj.litellm_params["metadata"] == {"from_caller": True}
|
||||
|
||||
def test_custom_pricing_detected_via_litellm_metadata(self, logging_obj):
|
||||
"""Custom pricing in litellm_metadata.model_info should set custom_pricing flag."""
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
use_custom_pricing_for_model,
|
||||
)
|
||||
|
||||
lm_meta = {
|
||||
"model_info": {
|
||||
"id": "deploy-custom",
|
||||
"input_cost_per_token": 0.005,
|
||||
"output_cost_per_token": 0.015,
|
||||
}
|
||||
}
|
||||
kwargs = {"litellm_metadata": lm_meta}
|
||||
|
||||
logging_obj.update_from_kwargs(kwargs=kwargs)
|
||||
|
||||
assert use_custom_pricing_for_model(logging_obj.litellm_params) is True
|
||||
|
||||
def test_additional_params_forwarded(self, logging_obj):
|
||||
kwargs = {"metadata": {}}
|
||||
logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model="gpt-5",
|
||||
user="test-user",
|
||||
optional_params={"temperature": 0.7},
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
assert logging_obj.model == "gpt-5"
|
||||
assert logging_obj.user == "test-user"
|
||||
assert logging_obj.model_call_details["custom_llm_provider"] == "openai"
|
||||
|
||||
def test_empty_kwargs_no_error(self, logging_obj):
|
||||
logging_obj.update_from_kwargs(
|
||||
kwargs={},
|
||||
litellm_params={"litellm_call_id": "call-empty"},
|
||||
)
|
||||
assert logging_obj.litellm_params["litellm_call_id"] == "call-empty"
|
||||
|
||||
|
||||
def test_logging_prevent_double_logging(logging_obj):
|
||||
"""
|
||||
When using a bridge, log only once from the underlying bridge call.
|
||||
|
|
|
|||
|
|
@ -615,6 +615,79 @@ def test_streaming_handler_with_stop_chunk(
|
|||
assert returned_chunk is None
|
||||
|
||||
|
||||
def test_finish_reason_chunk_preserves_non_openai_attributes(
|
||||
initialized_custom_stream_wrapper: CustomStreamWrapper,
|
||||
):
|
||||
"""
|
||||
Regression test for #23444:
|
||||
Preserve upstream non-OpenAI attributes on final finish_reason chunk.
|
||||
"""
|
||||
initialized_custom_stream_wrapper.received_finish_reason = "stop"
|
||||
|
||||
original_chunk = ModelResponseStream(
|
||||
id="chatcmpl-test",
|
||||
created=1742093326,
|
||||
model=None,
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(content=""),
|
||||
logprobs=None,
|
||||
)
|
||||
],
|
||||
)
|
||||
setattr(original_chunk, "custom_field", {"key": "value"})
|
||||
|
||||
returned_chunk = initialized_custom_stream_wrapper.return_processed_chunk_logic(
|
||||
completion_obj={"content": ""},
|
||||
response_obj={"original_chunk": original_chunk},
|
||||
model_response=ModelResponseStream(),
|
||||
)
|
||||
|
||||
assert returned_chunk is not None
|
||||
assert getattr(returned_chunk, "custom_field", None) == {"key": "value"}
|
||||
|
||||
|
||||
def test_finish_reason_with_holding_chunk_preserves_non_openai_attributes(
|
||||
initialized_custom_stream_wrapper: CustomStreamWrapper,
|
||||
):
|
||||
"""
|
||||
Regression test for #23444 holding-chunk path:
|
||||
preserve custom attributes when _is_delta_empty is False after flushing
|
||||
holding_chunk.
|
||||
"""
|
||||
initialized_custom_stream_wrapper.received_finish_reason = "stop"
|
||||
initialized_custom_stream_wrapper.holding_chunk = "filtered text"
|
||||
|
||||
original_chunk = ModelResponseStream(
|
||||
id="chatcmpl-test-2",
|
||||
created=1742093327,
|
||||
model=None,
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(content=""),
|
||||
logprobs=None,
|
||||
)
|
||||
],
|
||||
)
|
||||
setattr(original_chunk, "custom_field", {"key": "value"})
|
||||
|
||||
returned_chunk = initialized_custom_stream_wrapper.return_processed_chunk_logic(
|
||||
completion_obj={"content": ""},
|
||||
response_obj={"original_chunk": original_chunk},
|
||||
model_response=ModelResponseStream(),
|
||||
)
|
||||
|
||||
assert returned_chunk is not None
|
||||
assert returned_chunk.choices[0].delta.content == "filtered text"
|
||||
assert getattr(returned_chunk, "custom_field", None) == {"key": "value"}
|
||||
|
||||
|
||||
def test_set_response_id_propagation_empty_to_valid(
|
||||
initialized_custom_stream_wrapper: CustomStreamWrapper,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -564,6 +564,10 @@ async def test_ensure_initialize_azure_sdk_client_always_used(call_type):
|
|||
call_type == CallTypes.avideo_content
|
||||
or call_type == CallTypes.avideo_list
|
||||
or call_type == CallTypes.avideo_remix
|
||||
or call_type == CallTypes.avideo_create_character
|
||||
or call_type == CallTypes.avideo_get_character
|
||||
or call_type == CallTypes.avideo_edit
|
||||
or call_type == CallTypes.avideo_extension
|
||||
):
|
||||
# Skip video call types as they don't use Azure SDK client initialization
|
||||
pytest.skip(f"Skipping {call_type.value} because Azure video calls don't use initialize_azure_sdk_client")
|
||||
|
|
|
|||
|
|
@ -156,12 +156,11 @@ async def test_async_anthropic_messages_handler_extra_headers():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_anthropic_messages_handler_passes_litellm_metadata():
|
||||
"""Ensure litellm_metadata from kwargs is included in litellm_params
|
||||
passed to update_environment_variables.
|
||||
"""Ensure litellm_metadata from kwargs is forwarded via update_from_kwargs.
|
||||
|
||||
Routes like /messages store model_info under kwargs['litellm_metadata'].
|
||||
The handler must forward this into litellm_params so that
|
||||
use_custom_pricing_for_model can detect custom pricing. Regression test for #23185.
|
||||
The handler must forward this so that use_custom_pricing_for_model can
|
||||
detect custom pricing. Regression test for #23185.
|
||||
"""
|
||||
handler = BaseLLMHTTPHandler()
|
||||
|
||||
|
|
@ -187,7 +186,7 @@ async def test_async_anthropic_messages_handler_passes_litellm_metadata():
|
|||
mock_client.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
mock_logging_obj = Mock()
|
||||
mock_logging_obj.update_environment_variables = Mock()
|
||||
mock_logging_obj.update_from_kwargs = Mock()
|
||||
mock_logging_obj.model_call_details = {}
|
||||
mock_logging_obj.stream = False
|
||||
|
||||
|
|
@ -218,14 +217,14 @@ async def test_async_anthropic_messages_handler_passes_litellm_metadata():
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
mock_logging_obj.update_environment_variables.assert_called_once()
|
||||
call_kwargs = mock_logging_obj.update_environment_variables.call_args
|
||||
litellm_params_arg = call_kwargs.kwargs.get(
|
||||
"litellm_params", call_kwargs[1].get("litellm_params", {})
|
||||
) if call_kwargs.kwargs else call_kwargs[1].get("litellm_params", {})
|
||||
mock_logging_obj.update_from_kwargs.assert_called_once()
|
||||
call_kwargs = mock_logging_obj.update_from_kwargs.call_args
|
||||
kwargs_arg = call_kwargs.kwargs.get(
|
||||
"kwargs", call_kwargs[1].get("kwargs", {})
|
||||
) if call_kwargs.kwargs else call_kwargs[1].get("kwargs", {})
|
||||
|
||||
assert "litellm_metadata" in litellm_params_arg
|
||||
assert litellm_params_arg["litellm_metadata"]["model_info"] == custom_model_info
|
||||
assert "litellm_metadata" in kwargs_arg
|
||||
assert kwargs_arg["litellm_metadata"]["model_info"] == custom_model_info
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -0,0 +1,38 @@
|
|||
from litellm.llms.vertex_ai.batches.transformation import VertexAIBatchTransformation
|
||||
|
||||
|
||||
def test_output_file_id_uses_predictions_jsonl_with_output_info():
|
||||
response = {
|
||||
"outputInfo": {
|
||||
"gcsOutputDirectory": "gs://test-bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-pro/prediction-model-123"
|
||||
}
|
||||
}
|
||||
|
||||
output_file_id = VertexAIBatchTransformation._get_output_file_id_from_vertex_ai_batch_response(
|
||||
response
|
||||
)
|
||||
|
||||
assert (
|
||||
output_file_id
|
||||
== "gs://test-bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-pro/prediction-model-123/predictions.jsonl"
|
||||
)
|
||||
|
||||
|
||||
def test_output_file_id_falls_back_to_output_uri_prefix_with_predictions_jsonl():
|
||||
response = {
|
||||
"outputInfo": {},
|
||||
"outputConfig": {
|
||||
"gcsDestination": {
|
||||
"outputUriPrefix": "gs://test-bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-pro/prediction-model-456"
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
output_file_id = VertexAIBatchTransformation._get_output_file_id_from_vertex_ai_batch_response(
|
||||
response
|
||||
)
|
||||
|
||||
assert (
|
||||
output_file_id
|
||||
== "gs://test-bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-pro/prediction-model-456/predictions.jsonl"
|
||||
)
|
||||
|
|
@ -158,7 +158,6 @@ class TestExecuteWithMcpClient:
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skip(reason="PR #23187 changed has_client_credentials to require explicit oauth2_flow opt-in, but NewMCPServerRequest and _execute_with_mcp_client were not updated - needs fix")
|
||||
async def test_m2m_credentials_forwarded_to_server_model(self, monkeypatch):
|
||||
"""M2M OAuth credentials (client_id, client_secret) from the nested
|
||||
``credentials`` dict must be forwarded to the MCPServer model so that
|
||||
|
|
@ -213,7 +212,6 @@ class TestExecuteWithMcpClient:
|
|||
assert server.has_client_credentials is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skip(reason="PR #23187 changed has_client_credentials to require explicit oauth2_flow opt-in, but NewMCPServerRequest and _execute_with_mcp_client were not updated - needs fix")
|
||||
async def test_m2m_drops_incoming_oauth2_headers(self, monkeypatch):
|
||||
"""For M2M OAuth servers the incoming Authorization header (which carries
|
||||
the litellm API key) must NOT be forwarded as extra_headers — otherwise
|
||||
|
|
|
|||
|
|
@ -209,3 +209,109 @@ def test_get_model_from_request_supports_google_model_names_with_slashes():
|
|||
def test_get_model_from_request_vertex_passthrough_still_works():
|
||||
route = "/vertex_ai/v1/projects/p/locations/l/publishers/google/models/gemini-1.5-pro:generateContent"
|
||||
assert get_model_from_request(request_data={}, route=route) == "gemini-1.5-pro"
|
||||
|
||||
|
||||
def test_get_customer_user_header_returns_none_when_no_customer_role():
|
||||
from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping
|
||||
|
||||
mappings = [
|
||||
{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"}
|
||||
]
|
||||
result = get_customer_user_header_from_mapping(mappings)
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_get_customer_user_header_returns_none_for_single_non_customer_mapping():
|
||||
from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping
|
||||
|
||||
mapping = {"header_name": "X-Only-Internal", "litellm_user_role": "internal_user"}
|
||||
result = get_customer_user_header_from_mapping(mapping)
|
||||
assert result is None
|
||||
|
||||
def test_get_customer_user_header_from_mapping_returns_customer_header():
|
||||
from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping
|
||||
|
||||
mappings = [
|
||||
{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"},
|
||||
{"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"},
|
||||
]
|
||||
result = get_customer_user_header_from_mapping(mappings)
|
||||
assert result == ["x-openwebui-user-email"]
|
||||
|
||||
|
||||
def test_get_customer_user_header_returns_customers_header_in_config_order_when_multiple_exist():
|
||||
from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping
|
||||
|
||||
mappings = [
|
||||
{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"},
|
||||
{"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"},
|
||||
{"header_name": "X-User-Id", "litellm_user_role": "customer"},
|
||||
]
|
||||
result = get_customer_user_header_from_mapping(mappings)
|
||||
assert result == ['x-openwebui-user-email', 'x-user-id']
|
||||
|
||||
|
||||
def test_get_end_user_id_returns_id_from_user_header_mappings():
|
||||
from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body
|
||||
|
||||
mappings = [
|
||||
{"header_name": "x-openwebui-user-id", "litellm_user_role": "internal_user"},
|
||||
{"header_name": "x-openwebui-user-email", "litellm_user_role": "customer"},
|
||||
]
|
||||
general_settings = {"user_header_mappings": mappings}
|
||||
headers = {"x-openwebui-user-email": "1234"}
|
||||
|
||||
with patch("litellm.proxy.auth.auth_utils._get_customer_id_from_standard_headers", return_value=None), \
|
||||
patch("litellm.proxy.proxy_server.general_settings", general_settings):
|
||||
result = get_end_user_id_from_request_body(request_body={}, request_headers=headers)
|
||||
|
||||
assert result == "1234"
|
||||
|
||||
|
||||
def test_get_end_user_id_returns_first_customer_header_when_multiple_mappings_exist():
|
||||
from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body
|
||||
|
||||
mappings = [
|
||||
{"header_name": "x-openwebui-user-id", "litellm_user_role": "internal_user"},
|
||||
{"header_name": "x-user-id", "litellm_user_role": "customer"},
|
||||
{"header_name": "x-openwebui-user-email", "litellm_user_role": "customer"},
|
||||
]
|
||||
general_settings = {"user_header_mappings": mappings}
|
||||
headers = {
|
||||
"x-user-id": "user-456",
|
||||
"x-openwebui-user-email": "user@example.com",
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.auth.auth_utils._get_customer_id_from_standard_headers", return_value=None), \
|
||||
patch("litellm.proxy.proxy_server.general_settings", general_settings):
|
||||
result = get_end_user_id_from_request_body(request_body={}, request_headers=headers)
|
||||
|
||||
assert result == "user-456"
|
||||
|
||||
|
||||
def test_get_end_user_id_returns_none_when_no_customer_role_in_mappings():
|
||||
from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body
|
||||
|
||||
mappings = [
|
||||
{"header_name": "x-openwebui-user-id", "litellm_user_role": "internal_user"},
|
||||
]
|
||||
general_settings = {"user_header_mappings": mappings}
|
||||
headers = {"x-openwebui-user-id": "user-789"}
|
||||
|
||||
with patch("litellm.proxy.auth.auth_utils._get_customer_id_from_standard_headers", return_value=None), \
|
||||
patch("litellm.proxy.proxy_server.general_settings", general_settings):
|
||||
result = get_end_user_id_from_request_body(request_body={}, request_headers=headers)
|
||||
|
||||
assert result is None
|
||||
|
||||
def test_get_end_user_id_falls_back_to_deprecated_user_header_name():
|
||||
from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body
|
||||
|
||||
general_settings = {"user_header_name": "x-custom-user-id"}
|
||||
headers = {"x-custom-user-id": "user-legacy"}
|
||||
|
||||
with patch("litellm.proxy.auth.auth_utils._get_customer_id_from_standard_headers", return_value=None), \
|
||||
patch("litellm.proxy.proxy_server.general_settings", general_settings):
|
||||
result = get_end_user_id_from_request_body(request_body={}, request_headers=headers)
|
||||
|
||||
assert result == "user-legacy"
|
||||
|
|
|
|||
|
|
@ -92,30 +92,33 @@ async def test_metadata_passed_to_custom_callback_codex_models():
|
|||
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
|
||||
litellm.callbacks = [callback]
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_post:
|
||||
mock_post.return_value = _make_mock_http_response(
|
||||
mock_response.model_dump()
|
||||
)
|
||||
# gpt-5.1-codex has mode=responses - routes through responses bridge
|
||||
await litellm.acompletion(
|
||||
model="gpt-5.1-codex",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
metadata=test_metadata,
|
||||
)
|
||||
try:
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_post:
|
||||
mock_post.return_value = _make_mock_http_response(
|
||||
mock_response.model_dump()
|
||||
)
|
||||
# gpt-5.1-codex has mode=responses - routes through responses bridge
|
||||
await litellm.acompletion(
|
||||
model="gpt-5.1-codex",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
metadata=test_metadata,
|
||||
)
|
||||
|
||||
await asyncio.wait_for(callback.event.wait(), timeout=5.0)
|
||||
await asyncio.wait_for(callback.event.wait(), timeout=5.0)
|
||||
|
||||
assert callback.captured_kwargs is not None, "Callback should have been invoked"
|
||||
assert callback.captured_kwargs is not None, "Callback should have been invoked"
|
||||
|
||||
litellm_params = callback.captured_kwargs.get("litellm_params", {})
|
||||
metadata = litellm_params.get("metadata") or {}
|
||||
litellm_params = callback.captured_kwargs.get("litellm_params", {})
|
||||
metadata = litellm_params.get("metadata") or {}
|
||||
|
||||
assert "foo" in metadata, "metadata['foo'] should be accessible in callback"
|
||||
assert metadata["foo"] == "bar"
|
||||
assert metadata.get("trace_id") == "test-123"
|
||||
assert "foo" in metadata, "metadata['foo'] should be accessible in callback"
|
||||
assert metadata["foo"] == "bar"
|
||||
assert metadata.get("trace_id") == "test-123"
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -152,27 +155,31 @@ async def test_metadata_passed_via_litellm_metadata_responses_api():
|
|||
|
||||
test_metadata = {"request_id": "req-456"}
|
||||
callback = MetadataCaptureCallback()
|
||||
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
|
||||
litellm.callbacks = [callback]
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_post:
|
||||
mock_post.return_value = _make_mock_http_response(
|
||||
mock_response.model_dump()
|
||||
)
|
||||
await litellm.aresponses(
|
||||
model="gpt-4o",
|
||||
input="hi",
|
||||
litellm_metadata=test_metadata,
|
||||
)
|
||||
try:
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_post:
|
||||
mock_post.return_value = _make_mock_http_response(
|
||||
mock_response.model_dump()
|
||||
)
|
||||
await litellm.aresponses(
|
||||
model="gpt-4o",
|
||||
input="hi",
|
||||
litellm_metadata=test_metadata,
|
||||
)
|
||||
|
||||
await asyncio.wait_for(callback.event.wait(), timeout=5.0)
|
||||
await asyncio.wait_for(callback.event.wait(), timeout=5.0)
|
||||
|
||||
assert callback.captured_kwargs is not None
|
||||
assert callback.captured_kwargs is not None
|
||||
|
||||
litellm_params = callback.captured_kwargs.get("litellm_params", {})
|
||||
metadata = litellm_params.get("metadata") or {}
|
||||
litellm_params = callback.captured_kwargs.get("litellm_params", {})
|
||||
metadata = litellm_params.get("metadata") or {}
|
||||
|
||||
assert "request_id" in metadata
|
||||
assert metadata["request_id"] == "req-456"
|
||||
assert "request_id" in metadata
|
||||
assert metadata["request_id"] == "req-456"
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
|
|
|||
|
|
@ -23,25 +23,26 @@ import pytest
|
|||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import json
|
||||
|
||||
import litellm
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MockResponse:
|
||||
def __init__(self, json_data, status_code):
|
||||
self._json_data = json_data
|
||||
self.status_code = status_code
|
||||
self.text = json.dumps(json_data)
|
||||
self.headers = {}
|
||||
|
||||
def json(self):
|
||||
return self._json_data
|
||||
def _build_mock_response(output_items, response_id="resp_mock-123"):
|
||||
"""Build a ResponsesAPIResponse that ``async_response_api_handler`` would return."""
|
||||
return ResponsesAPIResponse(
|
||||
id=response_id,
|
||||
created_at=1741476542,
|
||||
status="completed",
|
||||
model="openai/gpt-5.1-codex",
|
||||
output=output_items,
|
||||
usage={"input_tokens": 5, "output_tokens": 10, "total_tokens": 15},
|
||||
)
|
||||
|
||||
|
||||
def _get_item_id(item) -> str:
|
||||
|
|
@ -51,18 +52,8 @@ def _get_item_id(item) -> str:
|
|||
return getattr(item, "id", "") or ""
|
||||
|
||||
|
||||
def _has_encrypted_content(item) -> bool:
|
||||
"""Check whether an output item carries encrypted_content."""
|
||||
if isinstance(item, dict):
|
||||
return "encrypted_content" in item
|
||||
return hasattr(item, "encrypted_content") and getattr(item, "encrypted_content") is not None
|
||||
|
||||
|
||||
def _extract_encoded_item_id(response) -> str:
|
||||
"""
|
||||
Walk the response output and return the first litellm-encoded item ID
|
||||
(i.e. one that starts with ``encitem_``).
|
||||
"""
|
||||
"""Return the first ``encitem_``-prefixed item ID from the response output."""
|
||||
for item in response.output or []:
|
||||
item_id = _get_item_id(item)
|
||||
if item_id.startswith("encitem_"):
|
||||
|
|
@ -254,14 +245,14 @@ async def test_encrypted_content_affinity_tracks_and_routes():
|
|||
"""
|
||||
The first response rewrites encrypted-content item IDs to encoded form.
|
||||
The follow-up request with those encoded IDs is pinned to the same deployment.
|
||||
|
||||
Mocks ``async_response_api_handler`` (the method that makes the HTTP call)
|
||||
so the test is deterministic regardless of the HTTP transport in use.
|
||||
The ``@client`` decorator and ``_update_responses_api_response_id_with_model_id``
|
||||
post-processing still run, so item-ID rewriting is exercised end-to-end.
|
||||
"""
|
||||
mock_response_data = {
|
||||
"id": "resp_mock-123",
|
||||
"object": "response",
|
||||
"created_at": 1741476542,
|
||||
"status": "completed",
|
||||
"model": "openai/gpt-5.1-codex",
|
||||
"output": [
|
||||
mock_resp = _build_mock_response(
|
||||
output_items=[
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_abc123",
|
||||
|
|
@ -276,10 +267,7 @@ async def test_encrypted_content_affinity_tracks_and_routes():
|
|||
"encrypted_content": "gAAAAABpnW_yEYmSNEyOG...",
|
||||
},
|
||||
],
|
||||
"parallel_tool_calls": True,
|
||||
"usage": {"input_tokens": 5, "output_tokens": 10, "total_tokens": 15},
|
||||
"error": None,
|
||||
}
|
||||
)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
|
|
@ -301,6 +289,7 @@ async def test_encrypted_content_affinity_tracks_and_routes():
|
|||
},
|
||||
],
|
||||
optional_pre_call_checks=["encrypted_content_affinity"],
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
selected_deployments = []
|
||||
|
|
@ -311,14 +300,13 @@ async def test_encrypted_content_affinity_tracks_and_routes():
|
|||
return seq[1] if len(seq) > 1 else seq[0]
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
"litellm.llms.custom_httpx.llm_http_handler.BaseLLMHTTPHandler.async_response_api_handler",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_post, patch(
|
||||
return_value=mock_resp,
|
||||
), patch(
|
||||
"litellm.router_strategy.simple_shuffle.random.choice",
|
||||
side_effect=deterministic_choice,
|
||||
):
|
||||
mock_post.return_value = MockResponse(mock_response_data, 200)
|
||||
|
||||
# First request — goes to deployment-1 via deterministic_choice
|
||||
first_response = await router.aresponses(
|
||||
model="openai.gpt-5.1-codex",
|
||||
|
|
@ -376,6 +364,7 @@ async def test_encrypted_content_affinity_no_effect_on_chat_completions():
|
|||
},
|
||||
],
|
||||
optional_pre_call_checks=["encrypted_content_affinity"],
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
response1 = await router.acompletion(
|
||||
|
|
@ -394,15 +383,10 @@ async def test_encrypted_content_affinity_no_effect_on_chat_completions():
|
|||
async def test_encrypted_content_affinity_bypasses_rpm_limits():
|
||||
"""
|
||||
When encrypted content affinity pins to a deployment, the request
|
||||
goes through even if normal routing would avoid it.
|
||||
goes through even if normal routing would avoid it (usage-based-routing-v2).
|
||||
"""
|
||||
mock_response_data = {
|
||||
"id": "resp_mock-rpm-test",
|
||||
"object": "response",
|
||||
"created_at": 1741476542,
|
||||
"status": "completed",
|
||||
"model": "openai/gpt-5.1-codex",
|
||||
"output": [
|
||||
mock_resp = _build_mock_response(
|
||||
output_items=[
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "rs_encrypted_must_pin",
|
||||
|
|
@ -410,9 +394,8 @@ async def test_encrypted_content_affinity_bypasses_rpm_limits():
|
|||
"encrypted_content": "gAAAAABpnW_yEYmSNEyOG...",
|
||||
},
|
||||
],
|
||||
"usage": {"input_tokens": 5, "output_tokens": 10, "total_tokens": 15},
|
||||
"error": None,
|
||||
}
|
||||
response_id="resp_mock-rpm-test",
|
||||
)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
|
|
@ -435,6 +418,7 @@ async def test_encrypted_content_affinity_bypasses_rpm_limits():
|
|||
],
|
||||
optional_pre_call_checks=["encrypted_content_affinity"],
|
||||
routing_strategy="usage-based-routing-v2",
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
selected_deployments = []
|
||||
|
|
@ -445,14 +429,13 @@ async def test_encrypted_content_affinity_bypasses_rpm_limits():
|
|||
return seq[1] if len(seq) > 1 else seq[0]
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
"litellm.llms.custom_httpx.llm_http_handler.BaseLLMHTTPHandler.async_response_api_handler",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_post, patch(
|
||||
return_value=mock_resp,
|
||||
), patch(
|
||||
"litellm.router_strategy.simple_shuffle.random.choice",
|
||||
side_effect=deterministic_choice,
|
||||
):
|
||||
mock_post.return_value = MockResponse(mock_response_data, 200)
|
||||
|
||||
first_response = await router.aresponses(
|
||||
model="openai.gpt-5.1-codex",
|
||||
input="Initial request",
|
||||
|
|
@ -488,13 +471,8 @@ async def test_encrypted_content_affinity_no_match_normal_routing():
|
|||
Input items with non-encoded IDs (no encitem_ prefix) fall through to
|
||||
normal load balancing.
|
||||
"""
|
||||
mock_response_data = {
|
||||
"id": "resp_mock-no-match",
|
||||
"object": "response",
|
||||
"created_at": 1741476542,
|
||||
"status": "completed",
|
||||
"model": "openai/gpt-5.1-codex",
|
||||
"output": [
|
||||
mock_resp = _build_mock_response(
|
||||
output_items=[
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_new",
|
||||
|
|
@ -503,9 +481,8 @@ async def test_encrypted_content_affinity_no_match_normal_routing():
|
|||
"content": [{"type": "output_text", "text": "Response"}],
|
||||
},
|
||||
],
|
||||
"usage": {"input_tokens": 5, "output_tokens": 10, "total_tokens": 15},
|
||||
"error": None,
|
||||
}
|
||||
response_id="resp_mock-no-match",
|
||||
)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
|
|
@ -527,14 +504,14 @@ async def test_encrypted_content_affinity_no_match_normal_routing():
|
|||
},
|
||||
],
|
||||
optional_pre_call_checks=["encrypted_content_affinity"],
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
"litellm.llms.custom_httpx.llm_http_handler.BaseLLMHTTPHandler.async_response_api_handler",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_post:
|
||||
mock_post.return_value = MockResponse(mock_response_data, 200)
|
||||
|
||||
return_value=mock_resp,
|
||||
):
|
||||
# Non-encoded item ID — no affinity should kick in
|
||||
response = await router.aresponses(
|
||||
model="openai.gpt-5.1-codex",
|
||||
|
|
@ -551,22 +528,16 @@ async def test_encrypted_content_affinity_with_wrapped_content_no_id():
|
|||
Test affinity routing when items have wrapped encrypted_content but no ID.
|
||||
This simulates Codex client behavior where IDs are omitted.
|
||||
"""
|
||||
mock_response_data = {
|
||||
"id": "resp_mock-wrapped-content",
|
||||
"object": "response",
|
||||
"created_at": 1741476542,
|
||||
"status": "completed",
|
||||
"model": "openai/gpt-5.1-codex",
|
||||
"output": [
|
||||
mock_resp = _build_mock_response(
|
||||
output_items=[
|
||||
{
|
||||
"type": "reasoning",
|
||||
"status": "completed",
|
||||
"encrypted_content": "gAAAAABpnW_yEYmSNEyOG_original_content",
|
||||
},
|
||||
],
|
||||
"usage": {"input_tokens": 5, "output_tokens": 10, "total_tokens": 15},
|
||||
"error": None,
|
||||
}
|
||||
response_id="resp_mock-wrapped-content",
|
||||
)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
|
|
@ -588,6 +559,7 @@ async def test_encrypted_content_affinity_with_wrapped_content_no_id():
|
|||
},
|
||||
],
|
||||
optional_pre_call_checks=["encrypted_content_affinity"],
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
selected_deployments = []
|
||||
|
|
@ -598,14 +570,13 @@ async def test_encrypted_content_affinity_with_wrapped_content_no_id():
|
|||
return seq[1] if len(seq) > 1 else seq[0]
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
"litellm.llms.custom_httpx.llm_http_handler.BaseLLMHTTPHandler.async_response_api_handler",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_post, patch(
|
||||
return_value=mock_resp,
|
||||
), patch(
|
||||
"litellm.router_strategy.simple_shuffle.random.choice",
|
||||
side_effect=deterministic_choice,
|
||||
):
|
||||
mock_post.return_value = MockResponse(mock_response_data, 200)
|
||||
|
||||
# First request — goes to deployment-1
|
||||
first_response = await router.aresponses(
|
||||
model="openai.gpt-5.1-codex",
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ import pytest
|
|||
import litellm
|
||||
from litellm.anthropic_beta_headers_manager import (
|
||||
filter_and_transform_beta_headers,
|
||||
update_request_with_filtered_beta,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -116,6 +117,32 @@ class TestAnthropicBetaHeadersFiltering:
|
|||
unknown not in filtered
|
||||
), f"Unknown header '{unknown}' should be filtered out for {provider}"
|
||||
|
||||
def test_update_request_with_filtered_beta_vertex_ai(self):
|
||||
"""Test combined filtering for both HTTP headers and request body betas."""
|
||||
headers = {
|
||||
"anthropic-beta": "files-api-2025-04-14,context-management-2025-06-27,code-execution-2025-05-22"
|
||||
}
|
||||
request_data = {
|
||||
"anthropic_beta": [
|
||||
"files-api-2025-04-14",
|
||||
"context-management-2025-06-27",
|
||||
"code-execution-2025-05-22",
|
||||
]
|
||||
}
|
||||
|
||||
filtered_headers, filtered_request_data = update_request_with_filtered_beta(
|
||||
headers=headers,
|
||||
request_data=request_data,
|
||||
provider="vertex_ai",
|
||||
)
|
||||
|
||||
assert (
|
||||
filtered_headers.get("anthropic-beta") == "context-management-2025-06-27"
|
||||
)
|
||||
assert filtered_request_data.get("anthropic_beta") == [
|
||||
"context-management-2025-06-27"
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_messages_http_headers_filtering(self):
|
||||
"""Test that Anthropic messages API filters HTTP headers correctly."""
|
||||
|
|
|
|||
|
|
@ -6,76 +6,83 @@ encoding is loaded at import time (pre-#18070 behavior) instead of lazy loading.
|
|||
|
||||
This addresses issue #18659: VCR cassette creation broken by lazy loading.
|
||||
For now, this only affects encoding as it was the only reported issue.
|
||||
|
||||
Tests that need to clear sys.modules and re-import litellm run in subprocesses
|
||||
to avoid contaminating the test process's module graph (which breaks mock.patch
|
||||
for all subsequent tests on the same xdist worker).
|
||||
"""
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _run_python(script: str, env_override: dict | None = None) -> subprocess.CompletedProcess:
|
||||
"""Run a Python script in a subprocess and return the result."""
|
||||
import os
|
||||
env = os.environ.copy()
|
||||
# Remove the var so each test controls it explicitly
|
||||
env.pop("LITELLM_DISABLE_LAZY_LOADING", None)
|
||||
env.pop("TIKTOKEN_CACHE_DIR", None)
|
||||
if env_override:
|
||||
env.update(env_override)
|
||||
return subprocess.run(
|
||||
[sys.executable, "-c", textwrap.dedent(script)],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=env,
|
||||
timeout=60,
|
||||
)
|
||||
|
||||
|
||||
def test_eager_loading_enabled():
|
||||
"""Test that encoding is loaded at import time when env var is set"""
|
||||
# Set environment variable
|
||||
os.environ["LITELLM_DISABLE_LAZY_LOADING"] = "1"
|
||||
|
||||
# Clear any cached modules to ensure fresh import
|
||||
modules_to_clear = [k for k in sys.modules.keys() if k.startswith("litellm")]
|
||||
for module in modules_to_clear:
|
||||
del sys.modules[module]
|
||||
|
||||
# Import litellm - encoding should be loaded immediately
|
||||
import litellm
|
||||
|
||||
# Check that encoding is available (not lazy loaded)
|
||||
assert hasattr(litellm, "encoding"), "Encoding should be available when eager loading is enabled"
|
||||
|
||||
# Verify it's actually the encoding object
|
||||
encoding = litellm.encoding
|
||||
assert encoding is not None, "Encoding should not be None"
|
||||
|
||||
# Test that it works
|
||||
tokens = encoding.encode("Hello, world!")
|
||||
assert len(tokens) > 0, "Encoding should work"
|
||||
result = _run_python(
|
||||
"""
|
||||
import litellm
|
||||
assert hasattr(litellm, "encoding"), "Encoding should be available when eager loading is enabled"
|
||||
encoding = litellm.encoding
|
||||
assert encoding is not None, "Encoding should not be None"
|
||||
tokens = encoding.encode("Hello, world!")
|
||||
assert len(tokens) > 0, "Encoding should work"
|
||||
""",
|
||||
env_override={"LITELLM_DISABLE_LAZY_LOADING": "1"},
|
||||
)
|
||||
assert result.returncode == 0, f"Subprocess failed:\nstdout: {result.stdout}\nstderr: {result.stderr}"
|
||||
|
||||
|
||||
def test_eager_loading_env_var_values():
|
||||
"""Test that various env var values enable eager loading"""
|
||||
values = ["1", "true", "True", "TRUE", "yes", "Yes", "YES", "on", "On", "ON"]
|
||||
|
||||
for value in values:
|
||||
os.environ["LITELLM_DISABLE_LAZY_LOADING"] = value
|
||||
|
||||
# Clear modules
|
||||
modules_to_clear = [k for k in sys.modules.keys() if k.startswith("litellm")]
|
||||
for module in modules_to_clear:
|
||||
del sys.modules[module]
|
||||
|
||||
import litellm
|
||||
assert hasattr(litellm, "encoding"), f"Encoding should be available for value: {value}"
|
||||
encoding = litellm.encoding
|
||||
tokens = encoding.encode("test")
|
||||
assert len(tokens) > 0
|
||||
result = _run_python(
|
||||
"""
|
||||
import litellm
|
||||
assert hasattr(litellm, "encoding"), "Encoding should be available"
|
||||
encoding = litellm.encoding
|
||||
tokens = encoding.encode("test")
|
||||
assert len(tokens) > 0
|
||||
""",
|
||||
env_override={"LITELLM_DISABLE_LAZY_LOADING": value},
|
||||
)
|
||||
assert result.returncode == 0, (
|
||||
f"Failed for value {value!r}:\nstdout: {result.stdout}\nstderr: {result.stderr}"
|
||||
)
|
||||
|
||||
|
||||
def test_lazy_loading_default():
|
||||
"""Test that encoding is lazy loaded by default (when env var is not set)"""
|
||||
# Remove environment variable if set
|
||||
if "LITELLM_DISABLE_LAZY_LOADING" in os.environ:
|
||||
del os.environ["LITELLM_DISABLE_LAZY_LOADING"]
|
||||
|
||||
# Clear any cached modules
|
||||
modules_to_clear = [k for k in sys.modules.keys() if k.startswith("litellm")]
|
||||
for module in modules_to_clear:
|
||||
del sys.modules[module]
|
||||
|
||||
# Import litellm - encoding should NOT be loaded yet
|
||||
import litellm
|
||||
|
||||
# Encoding should be accessible via __getattr__ (lazy loading)
|
||||
encoding = litellm.encoding # This triggers lazy loading
|
||||
|
||||
# Verify it works
|
||||
tokens = encoding.encode("Hello, world!")
|
||||
assert len(tokens) > 0, "Encoding should work"
|
||||
result = _run_python(
|
||||
"""
|
||||
import litellm
|
||||
# Encoding should be accessible via __getattr__ (lazy loading)
|
||||
encoding = litellm.encoding
|
||||
tokens = encoding.encode("Hello, world!")
|
||||
assert len(tokens) > 0, "Encoding should work"
|
||||
""",
|
||||
)
|
||||
assert result.returncode == 0, f"Subprocess failed:\nstdout: {result.stdout}\nstderr: {result.stderr}"
|
||||
|
||||
|
||||
def test_tiktoken_cache_dir_set_on_lazy_load():
|
||||
|
|
@ -84,33 +91,15 @@ def test_tiktoken_cache_dir_set_on_lazy_load():
|
|||
This ensures the local tiktoken cache is used instead of downloading
|
||||
from the internet. Regression test for issue #19768.
|
||||
"""
|
||||
# Remove environment variables to ensure clean state
|
||||
if "LITELLM_DISABLE_LAZY_LOADING" in os.environ:
|
||||
del os.environ["LITELLM_DISABLE_LAZY_LOADING"]
|
||||
if "TIKTOKEN_CACHE_DIR" in os.environ:
|
||||
del os.environ["TIKTOKEN_CACHE_DIR"]
|
||||
|
||||
# Clear any cached modules
|
||||
modules_to_clear = [k for k in sys.modules.keys() if k.startswith("litellm")]
|
||||
for module in modules_to_clear:
|
||||
del sys.modules[module]
|
||||
|
||||
# Import litellm fresh
|
||||
import litellm
|
||||
|
||||
# Access encoding (triggers lazy load)
|
||||
_ = litellm.encoding
|
||||
|
||||
# Verify TIKTOKEN_CACHE_DIR is now set and points to local tokenizers
|
||||
assert "TIKTOKEN_CACHE_DIR" in os.environ, "TIKTOKEN_CACHE_DIR should be set after lazy loading encoding"
|
||||
cache_dir = os.environ["TIKTOKEN_CACHE_DIR"]
|
||||
assert "tokenizers" in cache_dir, f"TIKTOKEN_CACHE_DIR should point to tokenizers directory, got: {cache_dir}"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def cleanup_env():
|
||||
"""Clean up environment variable after each test"""
|
||||
yield
|
||||
if "LITELLM_DISABLE_LAZY_LOADING" in os.environ:
|
||||
del os.environ["LITELLM_DISABLE_LAZY_LOADING"]
|
||||
|
||||
result = _run_python(
|
||||
"""
|
||||
import os
|
||||
import litellm
|
||||
# Access encoding (triggers lazy load)
|
||||
_ = litellm.encoding
|
||||
assert "TIKTOKEN_CACHE_DIR" in os.environ, "TIKTOKEN_CACHE_DIR should be set after lazy loading encoding"
|
||||
cache_dir = os.environ["TIKTOKEN_CACHE_DIR"]
|
||||
assert "tokenizers" in cache_dir, f"TIKTOKEN_CACHE_DIR should point to tokenizers directory, got: {cache_dir}"
|
||||
""",
|
||||
)
|
||||
assert result.returncode == 0, f"Subprocess failed:\nstdout: {result.stdout}\nstderr: {result.stderr}"
|
||||
|
|
|
|||
191
tests/test_litellm/test_stream_chunk_builder_annotations.py
Normal file
191
tests/test_litellm/test_stream_chunk_builder_annotations.py
Normal file
|
|
@ -0,0 +1,191 @@
|
|||
"""
|
||||
Tests for stream_chunk_builder annotation merging.
|
||||
|
||||
Previously, stream_chunk_builder only took annotations from the FIRST
|
||||
annotation chunk, losing any annotations that arrived in later chunks.
|
||||
This fix merges annotations from ALL chunks.
|
||||
"""
|
||||
|
||||
from litellm import stream_chunk_builder
|
||||
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
|
||||
|
||||
|
||||
def test_stream_chunk_builder_merges_annotations_from_multiple_chunks():
|
||||
"""
|
||||
stream_chunk_builder must merge annotations from ALL streaming chunks,
|
||||
not just take them from the first annotation chunk.
|
||||
|
||||
Providers may spread annotations across multiple chunks (e.g. Gemini
|
||||
sends grounding metadata in the final chunk, while intermediate chunks
|
||||
may carry different annotations).
|
||||
"""
|
||||
annotation_a = {
|
||||
"type": "url_citation",
|
||||
"url_citation": {
|
||||
"url": "https://example.com/a",
|
||||
"title": "Source A",
|
||||
"start_index": 0,
|
||||
"end_index": 10,
|
||||
},
|
||||
}
|
||||
annotation_b = {
|
||||
"type": "url_citation",
|
||||
"url_citation": {
|
||||
"url": "https://example.com/b",
|
||||
"title": "Source B",
|
||||
"start_index": 20,
|
||||
"end_index": 30,
|
||||
},
|
||||
}
|
||||
|
||||
chunks = [
|
||||
ModelResponseStream(
|
||||
id="chatcmpl-test",
|
||||
created=1700000000,
|
||||
model="test-model",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(
|
||||
content="Part one. ",
|
||||
role="assistant",
|
||||
annotations=[annotation_a],
|
||||
),
|
||||
)
|
||||
],
|
||||
),
|
||||
ModelResponseStream(
|
||||
id="chatcmpl-test",
|
||||
created=1700000000,
|
||||
model="test-model",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(content="Part two."),
|
||||
)
|
||||
],
|
||||
),
|
||||
ModelResponseStream(
|
||||
id="chatcmpl-test",
|
||||
created=1700000000,
|
||||
model="test-model",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(
|
||||
content=None,
|
||||
annotations=[annotation_b],
|
||||
),
|
||||
)
|
||||
],
|
||||
),
|
||||
]
|
||||
|
||||
response = stream_chunk_builder(chunks=chunks)
|
||||
assert response is not None
|
||||
|
||||
message = response["choices"][0]["message"]
|
||||
assert message.annotations is not None
|
||||
assert len(message.annotations) == 2
|
||||
assert message.annotations[0] == annotation_a
|
||||
assert message.annotations[1] == annotation_b
|
||||
|
||||
|
||||
def test_stream_chunk_builder_single_annotation_chunk_still_works():
|
||||
"""
|
||||
When annotations come from a single chunk (most common case),
|
||||
stream_chunk_builder must still work correctly (no regression).
|
||||
"""
|
||||
annotation = {
|
||||
"type": "url_citation",
|
||||
"url_citation": {
|
||||
"url": "https://example.com/only",
|
||||
"title": "Only Source",
|
||||
"start_index": 0,
|
||||
"end_index": 5,
|
||||
},
|
||||
}
|
||||
|
||||
chunks = [
|
||||
ModelResponseStream(
|
||||
id="chatcmpl-test",
|
||||
created=1700000000,
|
||||
model="test-model",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(content="Hello", role="assistant"),
|
||||
)
|
||||
],
|
||||
),
|
||||
ModelResponseStream(
|
||||
id="chatcmpl-test",
|
||||
created=1700000000,
|
||||
model="test-model",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(content=None, annotations=[annotation]),
|
||||
)
|
||||
],
|
||||
),
|
||||
]
|
||||
|
||||
response = stream_chunk_builder(chunks=chunks)
|
||||
assert response is not None
|
||||
|
||||
message = response["choices"][0]["message"]
|
||||
assert message.annotations is not None
|
||||
assert len(message.annotations) == 1
|
||||
assert message.annotations[0] == annotation
|
||||
|
||||
|
||||
def test_stream_chunk_builder_no_annotations():
|
||||
"""
|
||||
When no chunks contain annotations, the message should not have
|
||||
an annotations key (no regression).
|
||||
"""
|
||||
chunks = [
|
||||
ModelResponseStream(
|
||||
id="chatcmpl-test",
|
||||
created=1700000000,
|
||||
model="test-model",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(content="Hello", role="assistant"),
|
||||
)
|
||||
],
|
||||
),
|
||||
ModelResponseStream(
|
||||
id="chatcmpl-test",
|
||||
created=1700000000,
|
||||
model="test-model",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(content=None),
|
||||
)
|
||||
],
|
||||
),
|
||||
]
|
||||
|
||||
response = stream_chunk_builder(chunks=chunks)
|
||||
assert response is not None
|
||||
|
||||
message = response["choices"][0]["message"]
|
||||
assert not hasattr(message, "annotations") or message.annotations is None
|
||||
|
|
@ -1,4 +1,5 @@
|
|||
import asyncio
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
|
|
@ -174,6 +175,34 @@ class TestVideoGeneration:
|
|||
assert files == []
|
||||
assert returned_api_base == "https://api.openai.com/v1/videos"
|
||||
|
||||
def test_video_generation_request_decodes_encoded_character_ids(self):
|
||||
"""Encoded character IDs should be decoded before upstream create-video call."""
|
||||
from litellm.types.videos.utils import encode_character_id_with_provider
|
||||
|
||||
config = OpenAIVideoConfig()
|
||||
encoded_character_id = encode_character_id_with_provider(
|
||||
character_id="char_123",
|
||||
provider="openai",
|
||||
model_id="sora-2",
|
||||
)
|
||||
|
||||
data, files, returned_api_base = config.transform_video_create_request(
|
||||
model="sora-2",
|
||||
prompt="Test video prompt",
|
||||
api_base="https://api.openai.com/v1/videos",
|
||||
video_create_optional_request_params={
|
||||
"seconds": "8",
|
||||
"size": "720x1280",
|
||||
"characters": [{"id": encoded_character_id}],
|
||||
},
|
||||
litellm_params=MagicMock(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert data["characters"] == [{"id": "char_123"}]
|
||||
assert files == []
|
||||
assert returned_api_base == "https://api.openai.com/v1/videos"
|
||||
|
||||
def test_video_generation_response_transformation(self):
|
||||
"""Test video generation response transformation."""
|
||||
config = OpenAIVideoConfig()
|
||||
|
|
@ -1623,3 +1652,516 @@ def test_video_remix_handler_prefers_explicit_api_key():
|
|||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
|
||||
|
||||
# ===== Tests for new video endpoints (characters, edits, extensions) =====
|
||||
|
||||
|
||||
class TestVideoCreateCharacter:
|
||||
"""Tests for video_create_character / avideo_create_character."""
|
||||
|
||||
def test_video_create_character_transform_request(self):
|
||||
"""Verify multipart form construction for POST /videos/characters."""
|
||||
config = OpenAIVideoConfig()
|
||||
fake_video = b"fake_video_bytes"
|
||||
|
||||
url, files_list = config.transform_video_create_character_request(
|
||||
name="hero",
|
||||
video=fake_video,
|
||||
api_base="https://api.openai.com/v1/videos",
|
||||
litellm_params=MagicMock(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert url == "https://api.openai.com/v1/videos/characters"
|
||||
# Should have (name field) + (video file field) = 2 entries
|
||||
assert len(files_list) == 2
|
||||
field_names = [f[0] for f in files_list]
|
||||
assert "name" in field_names
|
||||
assert "video" in field_names
|
||||
|
||||
def test_video_create_character_sets_video_mimetype(self):
|
||||
"""Ensure character video upload is sent as video/mp4."""
|
||||
config = OpenAIVideoConfig()
|
||||
fake_video = io.BytesIO(b"....ftyp....video-bytes")
|
||||
fake_video.name = "character.mp4"
|
||||
|
||||
_, files_list = config.transform_video_create_character_request(
|
||||
name="hero",
|
||||
video=fake_video,
|
||||
api_base="https://api.openai.com/v1/videos",
|
||||
litellm_params=MagicMock(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
video_parts = [f for f in files_list if f[0] == "video"]
|
||||
assert len(video_parts) == 1
|
||||
video_tuple = video_parts[0][1]
|
||||
assert video_tuple[0] == "character.mp4"
|
||||
assert video_tuple[2] == "video/mp4"
|
||||
|
||||
def test_video_create_character_transform_response(self):
|
||||
"""Verify CharacterObject is returned from response."""
|
||||
from litellm.types.videos.main import CharacterObject
|
||||
|
||||
config = OpenAIVideoConfig()
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"id": "char_abc123",
|
||||
"object": "character",
|
||||
"created_at": 1712697600,
|
||||
"name": "hero",
|
||||
}
|
||||
|
||||
result = config.transform_video_create_character_response(
|
||||
raw_response=mock_response,
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert isinstance(result, CharacterObject)
|
||||
assert result.id == "char_abc123"
|
||||
assert result.name == "hero"
|
||||
|
||||
def test_video_create_character_mock_response(self):
|
||||
"""video_create_character returns CharacterObject on mock_response."""
|
||||
from litellm.types.videos.main import CharacterObject
|
||||
from litellm.videos.main import video_create_character
|
||||
|
||||
response = video_create_character(
|
||||
name="hero",
|
||||
video=b"fake",
|
||||
mock_response={
|
||||
"id": "char_abc",
|
||||
"object": "character",
|
||||
"created_at": 1712697600,
|
||||
"name": "hero",
|
||||
},
|
||||
)
|
||||
assert isinstance(response, CharacterObject)
|
||||
assert response.id == "char_abc"
|
||||
|
||||
|
||||
class TestVideoGetCharacter:
|
||||
"""Tests for video_get_character / avideo_get_character."""
|
||||
|
||||
def test_video_get_character_transform_request(self):
|
||||
"""Verify URL construction for GET /videos/characters/{character_id}."""
|
||||
config = OpenAIVideoConfig()
|
||||
|
||||
url, params = config.transform_video_get_character_request(
|
||||
character_id="char_xyz",
|
||||
api_base="https://api.openai.com/v1/videos",
|
||||
litellm_params=MagicMock(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert url == "https://api.openai.com/v1/videos/characters/char_xyz"
|
||||
assert params == {}
|
||||
|
||||
def test_video_get_character_transform_response(self):
|
||||
"""Verify CharacterObject is returned from GET response."""
|
||||
from litellm.types.videos.main import CharacterObject
|
||||
|
||||
config = OpenAIVideoConfig()
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"id": "char_xyz",
|
||||
"object": "character",
|
||||
"created_at": 1712697600,
|
||||
"name": "villain",
|
||||
}
|
||||
|
||||
result = config.transform_video_get_character_response(
|
||||
raw_response=mock_response,
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert isinstance(result, CharacterObject)
|
||||
assert result.id == "char_xyz"
|
||||
assert result.name == "villain"
|
||||
|
||||
def test_video_get_character_mock_response(self):
|
||||
"""video_get_character returns CharacterObject on mock_response."""
|
||||
from litellm.types.videos.main import CharacterObject
|
||||
from litellm.videos.main import video_get_character
|
||||
|
||||
response = video_get_character(
|
||||
character_id="char_xyz",
|
||||
mock_response={
|
||||
"id": "char_xyz",
|
||||
"object": "character",
|
||||
"created_at": 1712697600,
|
||||
"name": "villain",
|
||||
},
|
||||
)
|
||||
assert isinstance(response, CharacterObject)
|
||||
assert response.id == "char_xyz"
|
||||
|
||||
|
||||
class TestVideoEdit:
|
||||
"""Tests for video_edit / avideo_edit."""
|
||||
|
||||
def test_video_edit_transform_request(self):
|
||||
"""Verify JSON body with video.id for POST /videos/edits."""
|
||||
config = OpenAIVideoConfig()
|
||||
|
||||
url, data = config.transform_video_edit_request(
|
||||
prompt="make it brighter",
|
||||
video_id="video_abc123",
|
||||
api_base="https://api.openai.com/v1/videos",
|
||||
litellm_params=MagicMock(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert url == "https://api.openai.com/v1/videos/edits"
|
||||
assert data["prompt"] == "make it brighter"
|
||||
assert data["video"]["id"] == "video_abc123"
|
||||
|
||||
def test_video_edit_transform_request_with_extra_body(self):
|
||||
"""Extra body params are merged into request data."""
|
||||
config = OpenAIVideoConfig()
|
||||
|
||||
url, data = config.transform_video_edit_request(
|
||||
prompt="darken it",
|
||||
video_id="video_abc123",
|
||||
api_base="https://api.openai.com/v1/videos",
|
||||
litellm_params=MagicMock(),
|
||||
headers={},
|
||||
extra_body={"resolution": "1080p"},
|
||||
)
|
||||
|
||||
assert data["resolution"] == "1080p"
|
||||
|
||||
def test_video_edit_mock_response(self):
|
||||
"""video_edit returns VideoObject on mock_response."""
|
||||
from litellm.videos.main import video_edit
|
||||
|
||||
response = video_edit(
|
||||
video_id="video_abc123",
|
||||
prompt="make it brighter",
|
||||
mock_response={
|
||||
"id": "video_edit_001",
|
||||
"object": "video",
|
||||
"status": "queued",
|
||||
"created_at": 1712697600,
|
||||
},
|
||||
)
|
||||
assert isinstance(response, VideoObject)
|
||||
assert response.id == "video_edit_001"
|
||||
|
||||
def test_video_edit_strips_encoded_provider_from_video_id(self):
|
||||
"""Provider-encoded video IDs are decoded before sending to API."""
|
||||
from litellm.types.videos.utils import encode_video_id_with_provider
|
||||
config = OpenAIVideoConfig()
|
||||
|
||||
encoded_id = encode_video_id_with_provider("raw_video_id", "openai", None)
|
||||
url, data = config.transform_video_edit_request(
|
||||
prompt="test",
|
||||
video_id=encoded_id,
|
||||
api_base="https://api.openai.com/v1/videos",
|
||||
litellm_params=MagicMock(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
# The video.id in the request body should be the raw ID, not the encoded one
|
||||
assert data["video"]["id"] == "raw_video_id"
|
||||
|
||||
|
||||
class TestVideoExtension:
|
||||
"""Tests for video_extension / avideo_extension."""
|
||||
|
||||
def test_video_extension_transform_request(self):
|
||||
"""Verify JSON body with video.id + seconds for POST /videos/extensions."""
|
||||
config = OpenAIVideoConfig()
|
||||
|
||||
url, data = config.transform_video_extension_request(
|
||||
prompt="continue the scene",
|
||||
video_id="video_abc123",
|
||||
seconds="5",
|
||||
api_base="https://api.openai.com/v1/videos",
|
||||
litellm_params=MagicMock(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert url == "https://api.openai.com/v1/videos/extensions"
|
||||
assert data["prompt"] == "continue the scene"
|
||||
assert data["seconds"] == "5"
|
||||
assert data["video"]["id"] == "video_abc123"
|
||||
|
||||
def test_video_extension_transform_request_with_extra_body(self):
|
||||
"""Extra body params are merged into request data."""
|
||||
config = OpenAIVideoConfig()
|
||||
|
||||
url, data = config.transform_video_extension_request(
|
||||
prompt="extend",
|
||||
video_id="video_abc123",
|
||||
seconds="10",
|
||||
api_base="https://api.openai.com/v1/videos",
|
||||
litellm_params=MagicMock(),
|
||||
headers={},
|
||||
extra_body={"model": "sora-2"},
|
||||
)
|
||||
|
||||
assert data["model"] == "sora-2"
|
||||
|
||||
def test_video_extension_mock_response(self):
|
||||
"""video_extension returns VideoObject on mock_response."""
|
||||
from litellm.videos.main import video_extension
|
||||
|
||||
response = video_extension(
|
||||
video_id="video_abc123",
|
||||
prompt="continue the scene",
|
||||
seconds="5",
|
||||
mock_response={
|
||||
"id": "video_ext_001",
|
||||
"object": "video",
|
||||
"status": "queued",
|
||||
"created_at": 1712697600,
|
||||
},
|
||||
)
|
||||
assert isinstance(response, VideoObject)
|
||||
assert response.id == "video_ext_001"
|
||||
|
||||
def test_video_extension_strips_encoded_provider_from_video_id(self):
|
||||
"""Provider-encoded video IDs are decoded before sending to API."""
|
||||
from litellm.types.videos.utils import encode_video_id_with_provider
|
||||
config = OpenAIVideoConfig()
|
||||
|
||||
encoded_id = encode_video_id_with_provider("raw_video_id", "openai", None)
|
||||
url, data = config.transform_video_extension_request(
|
||||
prompt="extend",
|
||||
video_id=encoded_id,
|
||||
seconds="5",
|
||||
api_base="https://api.openai.com/v1/videos",
|
||||
litellm_params=MagicMock(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert data["video"]["id"] == "raw_video_id"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def video_proxy_test_client():
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.video_endpoints.endpoints import router as video_router
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(video_router)
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: MagicMock()
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_character_id_encode_decode_roundtrip():
|
||||
from litellm.types.videos.utils import (
|
||||
decode_character_id_with_provider,
|
||||
encode_character_id_with_provider,
|
||||
)
|
||||
|
||||
encoded = encode_character_id_with_provider(
|
||||
character_id="char_raw_123",
|
||||
provider="vertex_ai",
|
||||
model_id="veo-2.0-generate-001",
|
||||
)
|
||||
decoded = decode_character_id_with_provider(encoded)
|
||||
|
||||
assert decoded["character_id"] == "char_raw_123"
|
||||
assert decoded["custom_llm_provider"] == "vertex_ai"
|
||||
assert decoded["model_id"] == "veo-2.0-generate-001"
|
||||
|
||||
|
||||
def test_character_id_decode_handles_missing_base64_padding():
|
||||
from litellm.types.videos.utils import (
|
||||
decode_character_id_with_provider,
|
||||
encode_character_id_with_provider,
|
||||
)
|
||||
|
||||
encoded = encode_character_id_with_provider(
|
||||
character_id="id",
|
||||
provider="openai",
|
||||
model_id="gpt-4o",
|
||||
)
|
||||
encoded_without_padding = encoded.rstrip("=")
|
||||
decoded = decode_character_id_with_provider(encoded_without_padding)
|
||||
|
||||
assert decoded["character_id"] == "id"
|
||||
assert decoded["custom_llm_provider"] == "openai"
|
||||
assert decoded["model_id"] == "gpt-4o"
|
||||
|
||||
|
||||
def test_video_create_character_target_model_names_returns_encoded_id(video_proxy_test_client):
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.types.videos.utils import decode_character_id_with_provider
|
||||
|
||||
captured_data = {}
|
||||
|
||||
async def _mock_base_process(self, **kwargs):
|
||||
captured_data.update(self.data)
|
||||
return {
|
||||
"id": "char_upstream_123",
|
||||
"object": "character",
|
||||
"created_at": 1712697600,
|
||||
"name": "hero",
|
||||
}
|
||||
|
||||
with patch.object(
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
"base_process_llm_request",
|
||||
new=_mock_base_process,
|
||||
):
|
||||
response = video_proxy_test_client.post(
|
||||
"/v1/videos/characters",
|
||||
headers={"Authorization": "Bearer sk-1234"},
|
||||
files={"video": ("character.mp4", b"fake-video", "video/mp4")},
|
||||
data={
|
||||
"name": "hero",
|
||||
"target_model_names": "vertex-ai-sora-2",
|
||||
"extra_body": json.dumps({"custom_llm_provider": "vertex_ai"}),
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
response_json = response.json()
|
||||
decoded = decode_character_id_with_provider(response_json["id"])
|
||||
assert decoded["character_id"] == "char_upstream_123"
|
||||
assert decoded["custom_llm_provider"] == "vertex_ai"
|
||||
assert decoded["model_id"] == "vertex-ai-sora-2"
|
||||
assert captured_data["model"] == "vertex-ai-sora-2"
|
||||
assert captured_data["custom_llm_provider"] == "vertex_ai"
|
||||
|
||||
|
||||
def test_video_get_character_accepts_encoded_character_id(video_proxy_test_client):
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.types.videos.utils import (
|
||||
decode_character_id_with_provider,
|
||||
encode_character_id_with_provider,
|
||||
)
|
||||
|
||||
captured_data = {}
|
||||
|
||||
async def _mock_base_process(self, **kwargs):
|
||||
captured_data.update(self.data)
|
||||
return {
|
||||
"id": "char_upstream_123",
|
||||
"object": "character",
|
||||
"created_at": 1712697600,
|
||||
"name": "hero",
|
||||
}
|
||||
|
||||
encoded_character_id = encode_character_id_with_provider(
|
||||
character_id="char_upstream_123",
|
||||
provider="vertex_ai",
|
||||
model_id="veo-2.0-generate-001",
|
||||
)
|
||||
mock_router = MagicMock()
|
||||
mock_router.resolve_model_name_from_model_id.return_value = "vertex-ai-sora-2"
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", mock_router):
|
||||
with patch.object(
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
"base_process_llm_request",
|
||||
new=_mock_base_process,
|
||||
):
|
||||
response = video_proxy_test_client.get(
|
||||
f"/v1/videos/characters/{encoded_character_id}",
|
||||
headers={"Authorization": "Bearer sk-1234"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert captured_data["character_id"] == "char_upstream_123"
|
||||
assert captured_data["custom_llm_provider"] == "vertex_ai"
|
||||
assert captured_data["model"] == "vertex-ai-sora-2"
|
||||
response_decoded = decode_character_id_with_provider(response.json()["id"])
|
||||
assert response_decoded["character_id"] == "char_upstream_123"
|
||||
assert response_decoded["custom_llm_provider"] == "vertex_ai"
|
||||
assert response_decoded["model_id"] == "veo-2.0-generate-001"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("endpoint", ["/v1/videos/edits", "/v1/videos/extensions"])
|
||||
def test_edit_and_extension_support_custom_provider_from_extra_body(
|
||||
video_proxy_test_client, endpoint
|
||||
):
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
||||
captured_data = {}
|
||||
|
||||
async def _mock_base_process(self, **kwargs):
|
||||
captured_data.update(self.data)
|
||||
return {
|
||||
"id": "video_resp_123",
|
||||
"object": "video",
|
||||
"status": "queued",
|
||||
"created_at": 1712697600,
|
||||
}
|
||||
|
||||
payload = {
|
||||
"prompt": "test",
|
||||
"video": {"id": "video_raw_123"},
|
||||
"extra_body": {"custom_llm_provider": "vertex_ai"},
|
||||
}
|
||||
if endpoint.endswith("extensions"):
|
||||
payload["seconds"] = "4"
|
||||
|
||||
with patch.object(
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
"base_process_llm_request",
|
||||
new=_mock_base_process,
|
||||
):
|
||||
response = video_proxy_test_client.post(
|
||||
endpoint,
|
||||
headers={"Authorization": "Bearer sk-1234"},
|
||||
json=payload,
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert captured_data["custom_llm_provider"] == "vertex_ai"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("endpoint", ["/v1/videos/edits", "/v1/videos/extensions"])
|
||||
def test_edit_and_extension_route_with_encoded_video_ids(
|
||||
video_proxy_test_client, endpoint
|
||||
):
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.types.videos.utils import encode_video_id_with_provider
|
||||
|
||||
captured_data = {}
|
||||
|
||||
async def _mock_base_process(self, **kwargs):
|
||||
captured_data.update(self.data)
|
||||
return {
|
||||
"id": "video_resp_123",
|
||||
"object": "video",
|
||||
"status": "queued",
|
||||
"created_at": 1712697600,
|
||||
}
|
||||
|
||||
encoded_video_id = encode_video_id_with_provider(
|
||||
video_id="video_raw_123",
|
||||
provider="vertex_ai",
|
||||
model_id="veo-2.0-generate-001",
|
||||
)
|
||||
payload = {"prompt": "test", "video": {"id": encoded_video_id}}
|
||||
if endpoint.endswith("extensions"):
|
||||
payload["seconds"] = "4"
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.resolve_model_name_from_model_id.return_value = "vertex-ai-sora-2"
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", mock_router):
|
||||
with patch.object(
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
"base_process_llm_request",
|
||||
new=_mock_base_process,
|
||||
):
|
||||
response = video_proxy_test_client.post(
|
||||
endpoint,
|
||||
headers={"Authorization": "Bearer sk-1234"},
|
||||
json=payload,
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert captured_data["video_id"] == encoded_video_id
|
||||
assert captured_data["custom_llm_provider"] == "vertex_ai"
|
||||
assert captured_data["model"] == "vertex-ai-sora-2"
|
||||
|
|
|
|||
|
|
@ -24,6 +24,10 @@ export default defineConfig({
|
|||
|
||||
/* Collect trace when retrying the failed test. See https://playwright.dev/docs/trace-viewer */
|
||||
trace: "on-first-retry",
|
||||
|
||||
/* Action timeout for clicks, fills, waitForSelector, etc. */
|
||||
actionTimeout: 15 * 1000,
|
||||
navigationTimeout: 30 * 1000,
|
||||
},
|
||||
|
||||
/* Configure projects for major browsers */
|
||||
|
|
@ -40,7 +44,7 @@ export default defineConfig({
|
|||
],
|
||||
|
||||
/* Timeout settings */
|
||||
timeout: 4 * 60 * 1000,
|
||||
timeout: 3 * 60 * 1000,
|
||||
expect: {
|
||||
timeout: 10 * 1000,
|
||||
},
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue