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:
yuneng-jiang 2026-03-16 15:56:09 -07:00 • committed by GitHub
commit e4c8f95328
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
98 changed files with 5697 additions and 1715 deletions

File diff suppressed because it is too large Load diff

View 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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View 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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View 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

View file

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

View file

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