Merge remote-tracking branch 'origin/main' into litellm_fal_gpt_image_25_flux_dev_edits

This commit is contained in:
kerry 2026-09-20 17:20:17 +00:00
commit 342bde7a8d
514 changed files with 2244 additions and 2870 deletions

View file

@ -85,7 +85,7 @@ def _filter_reserved_headers(
def _request_scoped_runtime_session_id(
params: Mapping[str, Any],
params: Mapping[str, object],
litellm_params: Mapping[str, Any],
) -> str | None:
context_id: Final = get_session_id_from_a2a_params(params)

View file

@ -20,7 +20,7 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig):
params: dict[str, Any],
api_base: str | None = None,
**kwargs: Any,
) -> dict[str, Any]:
) -> dict[str, object]:
"""Handle a non-streaming A2A request via WXO runs API."""
litellm_params: Final = kwargs.get("litellm_params")
if not litellm_params:
@ -40,7 +40,7 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig):
params: dict[str, Any],
api_base: str | None = None,
**kwargs: Any,
) -> AsyncIterator[dict[str, Any]]:
) -> AsyncIterator[dict[str, object]]:
"""Handle a streaming A2A request via WXO streaming runs API."""
litellm_params: Final = kwargs.get("litellm_params")
if not litellm_params:

View file

@ -17,7 +17,7 @@ class A2ARequestUtils:
"""Utility class for A2A request/response processing."""
@staticmethod
def extract_text_from_message(message: Any) -> str:
def extract_text_from_message(message: object) -> str:
"""
Extract text content from A2A message parts.
@ -142,7 +142,7 @@ class A2ARequestUtils:
return prompt_tokens, completion_tokens, total_tokens
def get_session_id_from_a2a_params(params: Mapping[str, Any]) -> str | None:
def get_session_id_from_a2a_params(params: Mapping[str, object]) -> str | None:
message: Final = params.get("message", {})
if isinstance(message, dict):
return message.get("contextId")
@ -166,7 +166,7 @@ def scope_session_to_principal(session_id: str, principal: str | None) -> str:
# Backwards compatibility aliases
def extract_text_from_a2a_message(message: Any) -> str:
def extract_text_from_a2a_message(message: object) -> str:
return A2ARequestUtils.extract_text_from_message(message)

View file

@ -200,8 +200,8 @@ class GitLabTemplateManager:
metadata=metadata,
)
def _parse_yaml_basic(self, yaml_str: str) -> dict[str, Any]:
result: Final[dict[str, Any]] = {}
def _parse_yaml_basic(self, yaml_str: str) -> dict[str, bool | int | float | str]:
result: Final[dict[str, bool | int | float | str]] = {}
for line in yaml_str.split("\n"):
line = line.strip()
if ":" in line and not line.startswith("#"):

View file

@ -59,7 +59,7 @@ class VantageLogger(FocusLogger):
raw_interval,
)
destination_config: Final[dict[str, Any]] = {}
destination_config: Final[dict[str, str]] = {}
if resolved_api_key:
destination_config["api_key"] = resolved_api_key
if resolved_token:
@ -93,7 +93,7 @@ class VantageLogger(FocusLogger):
pod_lock_manager = None
if proxy_logging_obj is not None:
writer: Final = getattr(proxy_logging_obj, "db_spend_update_writer", None)
writer: Final[object] = getattr(proxy_logging_obj, "db_spend_update_writer", None)
if writer is not None:
pod_lock_manager = getattr(writer, "pod_lock_manager", None)

View file

@ -7,7 +7,7 @@ duplicated. BaseAgentsAPIConfig stays as pure transform code.
"""
from collections.abc import Coroutine, Mapping
from typing import Any, Final
from typing import Final
import httpx
@ -38,7 +38,7 @@ class AgentsHTTPHandler(InteractionsHTTPHandler):
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: dict[str, Any] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: Mapping[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
client: HTTPHandler | None = None,
@ -93,7 +93,7 @@ class AgentsHTTPHandler(InteractionsHTTPHandler):
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: dict[str, Any] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: Mapping[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
client: AsyncHTTPHandler | None = None,
@ -141,7 +141,7 @@ class AgentsHTTPHandler(InteractionsHTTPHandler):
agents_api_config: BaseAgentsAPIConfig,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: dict[str, Any] | None = None,
extra_headers: dict[str, str] | None = None,
timeout: float | httpx.Timeout | None = None,
client: HTTPHandler | None = None,
_is_async: bool = False,
@ -181,7 +181,7 @@ class AgentsHTTPHandler(InteractionsHTTPHandler):
agents_api_config: BaseAgentsAPIConfig,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: dict[str, Any] | None = None,
extra_headers: dict[str, str] | None = None,
timeout: float | httpx.Timeout | None = None,
client: AsyncHTTPHandler | None = None,
) -> AgentListResponse:
@ -216,7 +216,7 @@ class AgentsHTTPHandler(InteractionsHTTPHandler):
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: dict[str, Any] | None = None,
extra_headers: dict[str, str] | None = None,
timeout: float | httpx.Timeout | None = None,
client: HTTPHandler | None = None,
_is_async: bool = False,
@ -259,7 +259,7 @@ class AgentsHTTPHandler(InteractionsHTTPHandler):
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: dict[str, Any] | None = None,
extra_headers: dict[str, str] | None = None,
timeout: float | httpx.Timeout | None = None,
client: AsyncHTTPHandler | None = None,
) -> AgentCreateResponse:
@ -295,7 +295,7 @@ class AgentsHTTPHandler(InteractionsHTTPHandler):
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: dict[str, Any] | None = None,
extra_headers: dict[str, str] | None = None,
timeout: float | httpx.Timeout | None = None,
client: HTTPHandler | None = None,
_is_async: bool = False,
@ -338,7 +338,7 @@ class AgentsHTTPHandler(InteractionsHTTPHandler):
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: dict[str, Any] | None = None,
extra_headers: dict[str, str] | None = None,
timeout: float | httpx.Timeout | None = None,
client: AsyncHTTPHandler | None = None,
) -> AgentDeleteResult:
@ -374,7 +374,7 @@ class AgentsHTTPHandler(InteractionsHTTPHandler):
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: dict[str, Any] | None = None,
extra_headers: dict[str, str] | None = None,
timeout: float | httpx.Timeout | None = None,
client: HTTPHandler | None = None,
_is_async: bool = False,
@ -417,7 +417,7 @@ class AgentsHTTPHandler(InteractionsHTTPHandler):
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: dict[str, Any] | None = None,
extra_headers: dict[str, str] | None = None,
timeout: float | httpx.Timeout | None = None,
client: AsyncHTTPHandler | None = None,
) -> AgentVersionsResponse:

View file

@ -67,7 +67,9 @@ def _truncate_base64_in_string(value: str) -> str:
return _DATA_URI_RE.sub(_base64_data_uri_replacer, value)
def _truncate_base64_in_value(value: Any) -> Any:
def _truncate_base64_in_value(
value: str | dict[str, object] | list[object] | None,
) -> str | dict[str, object] | list[object] | None:
"""Iteratively truncate base64 data URIs in a JSON-like value (str/list/dict).
Uses an explicit stack instead of recursion to satisfy the project's

View file

@ -418,7 +418,7 @@ def _extract_redirect_url(response: httpx.Response, request_url: str) -> str:
return str(httpx.URL(request_url).join(location))
def safe_get(client: Any, url: str, **kwargs: Any) -> httpx.Response:
def safe_get(client: _UrlFetcher, url: str, **kwargs: Any) -> httpx.Response:
"""
Fetch a user-supplied URL with SSRF protection on every redirect hop.
@ -461,7 +461,7 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> httpx.Response:
raise SSRFError("Too many redirects")
async def async_safe_get(client: Any, url: str, **kwargs: Any) -> httpx.Response:
async def async_safe_get(client: _AsyncUrlFetcher, url: str, **kwargs: Any) -> httpx.Response:
"""Async version of safe_get."""
if not getattr(litellm, "user_url_validation", True):
kwargs.setdefault("follow_redirects", True)

View file

@ -596,7 +596,7 @@ class ModelResponseIterator:
self.reasoning_content_chunks: list[str] = []
# Track server tool use inputs and results for code_interpreter_results
self._server_tool_inputs: dict[str, Any] = {}
self._server_tool_inputs: dict[str, object] = {}
self.tool_results: list[dict[str, Any]] = []
self._current_server_tool_id: str | None = None
self._container_id: str | None = None

View file

@ -1,6 +1,7 @@
from collections.abc import Callable
from typing import Any, Final
from typing import Final
import httpx
from openai import AsyncAzureOpenAI, AzureOpenAI
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -191,7 +192,7 @@ class AzureTextCompletion(BaseAzureLLM):
model: str,
api_base: str,
data: dict,
timeout: Any,
timeout: float | httpx.Timeout | None,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
max_retries: int,
@ -253,7 +254,7 @@ class AzureTextCompletion(BaseAzureLLM):
api_version: str,
data: dict,
model: str,
timeout: Any,
timeout: float | httpx.Timeout | None,
azure_ad_token: str | None = None,
client=None,
litellm_params: dict = {},
@ -306,7 +307,7 @@ class AzureTextCompletion(BaseAzureLLM):
api_version: str,
data: dict,
model: str,
timeout: Any,
timeout: float | httpx.Timeout | None,
azure_ad_token: str | None = None,
client=None,
litellm_params: dict = {},

View file

@ -12,7 +12,7 @@ import asyncio
import re
import time
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Final
from urllib.parse import quote
import httpx
@ -127,7 +127,7 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
def map_ocr_params(
self,
non_default_params: dict,
non_default_params: Mapping[str, object],
optional_params: dict,
model: str,
) -> dict:
@ -164,7 +164,7 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
raise UnsupportedParamsError(message=f"{e}", model=model, llm_provider="azure_ai") from e
@staticmethod
def _normalize_pages_param(pages: Any) -> str:
def _normalize_pages_param(pages: object) -> str:
"""
Convert a caller-provided `pages` value to Azure DI's query-string
form. Azure expects 1-based page numbers, grammar: `^(\\d+(-\\d+)?)(,\\s*(\\d+(-\\d+)?))*$`.
@ -412,7 +412,7 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
raise ValueError("Document URL is required")
# Build Azure DI request
data: Final[dict[str, Any]] = {}
data: Final[dict[str, str]] = {}
# Check if it's a data URI (base64)
if document_url.startswith("data:"):

View file

@ -2,7 +2,7 @@ import os
import re
import time
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Literal, cast
from typing import TYPE_CHECKING, Final, Literal, cast
from httpx import Headers, Response
from pydantic import TypeAdapter, ValidationError
@ -170,7 +170,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
create_batch_data: CreateBatchRequest,
optional_params: dict,
litellm_params: dict,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Transform the batch creation request to Bedrock format.
@ -354,7 +354,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
)
@staticmethod
def _get_openai_compatible_batch_metadata(metadata: Any) -> dict[str, str]:
def _get_openai_compatible_batch_metadata(metadata: object) -> dict[str, str]:
"""
OpenAI Batch metadata only accepts string values.
"""
@ -379,7 +379,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
batch_id: str,
optional_params: dict,
litellm_params: dict,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Transform batch retrieval request for Bedrock.
@ -523,7 +523,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
)
# Enrich metadata with useful Bedrock fields
enriched_metadata_raw: Final[dict[str, Any]] = {
enriched_metadata_raw: Final[dict[str, object]] = {
"jobName": response_data.get("jobName"),
"clientRequestToken": response_data.get("clientRequestToken"),
"modelId": response_data.get("modelId"),

View file

@ -110,7 +110,7 @@ class AmazonTwelveLabsPegasusConfig(AmazonInvokeConfig, BaseConfig):
headers: dict,
) -> dict:
input_prompt: Final = self._convert_messages_to_prompt(messages=messages)
request_data: Final[dict[str, Any]] = {"inputPrompt": input_prompt}
request_data: Final[dict[str, object]] = {"inputPrompt": input_prompt}
media_source: Final = self._build_media_source(optional_params)
if media_source is not None:

View file

@ -335,10 +335,10 @@ class BytezChatConfig(BaseConfig):
class BytezCustomStreamWrapper(CustomStreamWrapper):
def chunk_creator(self, chunk: Any):
def chunk_creator(self, chunk: object):
try:
model_response: Final = self.model_response_creator()
response_obj: dict[str, Any] = {}
response_obj: dict[str, object] = {}
response_obj = {
"text": chunk,
@ -346,7 +346,7 @@ class BytezCustomStreamWrapper(CustomStreamWrapper):
"finish_reason": "",
}
completion_obj: Final[dict[str, Any]] = {"content": chunk}
completion_obj: Final[dict[str, object]] = {"content": chunk}
return self.return_processed_chunk_logic(
completion_obj=completion_obj,

View file

@ -1,5 +1,5 @@
import ssl
from collections.abc import Callable
from collections.abc import AsyncIterable, Callable, Iterable
from typing import TYPE_CHECKING, Any, Final, cast
import aiohttp
@ -212,7 +212,7 @@ class BaseLLMAIOHTTPHandler:
litellm_params: dict,
stream: bool = False,
files: dict | None = None,
content: Any = None,
content: str | bytes | Iterable[bytes] | AsyncIterable[bytes] | None = None,
params: dict | None = None,
) -> httpx.Response:
max_retry_on_unprocessable_entity_error: Final = provider_config.max_retry_on_unprocessable_entity_error

View file

@ -146,7 +146,7 @@ class AlephAlphaConfig:
setattr(self.__class__, key, value)
@classmethod
def get_config(cls):
def get_config(cls) -> dict[str, object]:
return {
k: v
for k, v in cls.__dict__.items()

View file

@ -0,0 +1,3 @@
from litellm.llms.fal_ai.videos.transformation import FalAIVideoConfig
__all__ = ("FalAIVideoConfig",)

View file

@ -0,0 +1,516 @@
import math
import time
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final, TypeAlias
import httpx
from httpx._types import FileContent, RequestFiles
from pydantic import TypeAdapter
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
_get_httpx_client, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # shared HTTP factory is private
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # shared HTTP factory lacks typed params
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
from litellm.types.videos.main import (
CharacterObject,
VideoCreateOptionalRequestParams,
VideoObject,
)
from litellm.types.videos.utils import (
decode_video_id_with_provider,
encode_video_id_with_provider,
)
class FalAIVideoError(BaseLLMException):
pass
_ALLOWED_ASPECT_RATIOS: Final[frozenset[str]] = frozenset({"auto", "16:9", "9:16", "1:1", "4:3", "3:4", "21:9"})
_ALLOWED_RESOLUTIONS: Final[frozenset[str]] = frozenset({"480p", "720p", "1080p", "4k"})
_RESOLUTION_TIERS: Final[tuple[tuple[int, str], ...]] = (
(480, "480p"),
(720, "720p"),
(1080, "1080p"),
)
_QUEUE_NAMESPACES: Final[frozenset[str]] = frozenset(("workflows", "comfy"))
_STATUS_MAP: Final[Mapping[str, str]] = MappingProxyType(
{
"IN_QUEUE": "queued",
"IN_PROGRESS": "in_progress",
"COMPLETED": "completed",
}
)
_FAL_AI_PROVIDER: Final[str] = LlmProviders.FAL_AI.value
_SupportedParams: TypeAlias = list[str]
_VideoParams: TypeAlias = dict[str, object]
_VideoHeaders: TypeAlias = dict[str, str]
_VideoStringParams: TypeAlias = dict[str, str]
_VideoFiles: TypeAlias = list[object]
def _queue_request_base_path(model: str) -> str:
segments: Final[tuple[str, ...]] = tuple(model.split("/"))
segment_count: Final[int] = 3 if segments and segments[0] in _QUEUE_NAMESPACES else 2
return "/".join(segments[:segment_count])
def _duration_value(value: object) -> str | None:
if isinstance(value, str) and value == "auto":
return value
if isinstance(value, bool) or not isinstance(value, (int, float, str)):
return None
try:
return str(int(float(value)))
except (TypeError, ValueError):
return None
def _resolution_for_short_side(short_side: int) -> str:
return next((resolution for threshold, resolution in _RESOLUTION_TIERS if short_side <= threshold), "4k")
def _model_path_from_request_url(raw_response: httpx.Response) -> str | None:
segments: Final[tuple[str, ...]] = tuple(segment for segment in raw_response.request.url.path.split("/") if segment)
if "requests" not in segments:
return None
model_segments: Final[tuple[str, ...]] = segments[: segments.index("requests")]
segment_count: Final[int] = 3 if len(model_segments) >= 3 and model_segments[-3] in _QUEUE_NAMESPACES else 2
return "/".join(model_segments[-segment_count:]) if len(model_segments) >= segment_count else None
def _request_id_from_request_url(raw_response: httpx.Response) -> str | None:
segments: Final[tuple[str, ...]] = tuple(segment for segment in raw_response.request.url.path.split("/") if segment)
if "requests" not in segments:
return None
request_index: Final[int] = segments.index("requests")
request_id_index: Final[int] = request_index + 1
return segments[request_id_index] if len(segments) > request_id_index else None
def _size_params(size: object) -> Mapping[str, str]:
if not isinstance(size, str):
return MappingProxyType({})
if size in _ALLOWED_RESOLUTIONS:
return MappingProxyType({"resolution": size})
if size.count("x") != 1:
return MappingProxyType({})
width_text, height_text = size.split("x")
if not (width_text.isdigit() and height_text.isdigit()):
return MappingProxyType({})
width: Final[int] = int(width_text)
height: Final[int] = int(height_text)
if width <= 0 or height <= 0:
return MappingProxyType({})
reduced_gcd: Final[int] = math.gcd(width, height)
aspect_ratio: Final[str] = f"{width // reduced_gcd}:{height // reduced_gcd}"
resolution: Final[str] = _resolution_for_short_side(min(width, height))
if aspect_ratio in _ALLOWED_ASPECT_RATIOS:
return MappingProxyType({"resolution": resolution, "aspect_ratio": aspect_ratio})
return MappingProxyType({"resolution": resolution})
def _numeric_duration(value: object) -> float | None:
duration: Final[str | None] = _duration_value(value)
if duration is None or duration == "auto":
return None
return float(duration)
def _response_data(raw_response: httpx.Response) -> Mapping[str, object]:
return TypeAdapter(Mapping[str, object]).validate_python(raw_response.json())
def _response_string(response_data: Mapping[str, object], key: str, default: str = "") -> str:
value: Final[object] = response_data.get(key)
return value if isinstance(value, str) else default
class FalAIVideoConfig(BaseVideoConfig):
def get_supported_openai_params(self, model: str) -> _SupportedParams:
supported_params: Final[_SupportedParams] = [ # mutable-ok: BaseVideoConfig requires a list
"model",
"prompt",
"input_reference",
"seconds",
"size",
"user",
"extra_headers",
]
return supported_params
def map_openai_params(
self,
video_create_optional_params: VideoCreateOptionalRequestParams,
model: str,
drop_params: bool,
) -> _VideoParams:
supported_params: Final[frozenset[str]] = frozenset(self.get_supported_openai_params(model))
input_reference: Final[object] = video_create_optional_params.get("input_reference")
if "input_reference" in video_create_optional_params and not isinstance(input_reference, str):
raise ValueError("fal.ai needs a public image URL for input_reference")
input_reference_params: Final[Mapping[str, str]] = (
MappingProxyType({})
if not isinstance(input_reference, str)
else MappingProxyType({"image_url": input_reference})
)
duration_params: Final[Mapping[str, str]] = (
MappingProxyType({})
if "seconds" not in video_create_optional_params
else self._duration_params(video_create_optional_params["seconds"])
)
size_params: Final[Mapping[str, str]] = (
_size_params(video_create_optional_params["size"])
if "size" in video_create_optional_params
else MappingProxyType({})
)
user_params: Final[Mapping[str, str]] = (
MappingProxyType({"end_user_id": user})
if isinstance(user := video_create_optional_params.get("user"), str)
else MappingProxyType({})
)
mapped_params: Final[_VideoParams] = {
**input_reference_params,
**duration_params,
**size_params,
**user_params,
**{ # mutable-ok: BaseVideoConfig requires a mutable parameter mapping
key: value for key, value in video_create_optional_params.items() if key not in supported_params
},
}
return mapped_params
@staticmethod
def _duration_params(seconds: object) -> Mapping[str, str]:
duration: Final[str | None] = _duration_value(seconds)
if duration is None:
raise ValueError("fal.ai seconds must be a numeric value")
return MappingProxyType({"duration": duration})
def validate_environment(
self,
headers: _VideoHeaders,
model: str,
api_key: str | None = None,
litellm_params: GenericLiteLLMParams | None = None,
) -> _VideoHeaders:
final_api_key: Final[str | None] = (
api_key
or (litellm_params.api_key if litellm_params is not None else None)
or get_secret_str("FAL_AI_API_KEY")
)
if not final_api_key:
raise ValueError("FAL_AI_API_KEY is not set")
validated_headers: Final[_VideoHeaders] = {
**headers,
"Authorization": f"Key {final_api_key}",
"Content-Type": "application/json",
}
return validated_headers
def get_complete_url(
self,
model: str,
api_base: str | None,
litellm_params: _VideoParams,
) -> str:
return (api_base or "https://queue.fal.run").rstrip("/")
def transform_video_create_request(
self,
model: str,
prompt: str,
api_base: str,
video_create_optional_request_params: _VideoParams,
litellm_params: GenericLiteLLMParams,
headers: _VideoHeaders,
) -> tuple[_VideoParams, RequestFiles, str]:
request_data: Final[_VideoParams] = {
"prompt": prompt,
**{ # mutable-ok: HTTP JSON payload requires a mutable mapping
key: value for key, value in video_create_optional_request_params.items() if key != "model"
},
}
return request_data, [], f"{api_base.rstrip('/')}/{model}" # mutable-ok: HTTP files payload requires a list
def transform_video_create_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: object,
custom_llm_provider: str | None = None,
request_data: Mapping[str, object] | None = None,
) -> VideoObject:
response_data: Final[Mapping[str, object]] = _response_data(raw_response)
request_params: Final[Mapping[str, object]] = request_data or MappingProxyType({})
request_id: Final[str] = _response_string(response_data, "request_id")
provider: Final[str] = custom_llm_provider or _FAL_AI_PROVIDER
duration: Final[float | None] = _numeric_duration(request_params.get("duration"))
resolution: Final[object] = request_params.get("resolution")
seconds: Final[str | None] = _duration_value(request_params["duration"]) if duration is not None else None
size: Final[str | None] = resolution if isinstance(resolution, str) else None
usage: Final[_VideoParams] = { # mutable-ok: VideoObject requires a mutable usage mapping
key: value
for key, value in (
("duration_seconds", duration),
("video_resolution", resolution if isinstance(resolution, str) else "720p"),
)
if value is not None
}
video_object: Final[VideoObject] = VideoObject(
id=encode_video_id_with_provider(request_id, provider, model),
object="video",
status="queued",
created_at=int(time.time()),
model=model,
seconds=seconds,
size=size,
)
video_object.usage = usage
return video_object
def transform_video_status_retrieve_request(
self,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: _VideoHeaders,
) -> tuple[str, _VideoParams]:
request_id, model_id = self._decode_video_id(video_id)
encoded_request_id: Final[str] = encode_url_path_segment(request_id, field_name="video_id")
return (
f"{api_base.rstrip('/')}/{_queue_request_base_path(model_id)}/requests/{encoded_request_id}/status",
{}, # mutable-ok: BaseVideoConfig requires a mutable mapping
)
def transform_video_status_retrieve_response(
self,
raw_response: httpx.Response,
logging_obj: object,
custom_llm_provider: str | None = None,
) -> VideoObject:
response_data: Final[Mapping[str, object]] = _response_data(raw_response)
raw_status: Final[str] = _response_string(response_data, "status", "IN_QUEUE")
status: Final[str] = _STATUS_MAP.get(raw_status, "queued")
error_value: Final[object] = response_data.get("error")
error: Final[str | None] = error_value if isinstance(error_value, str) else None
provider: Final[str] = custom_llm_provider or _FAL_AI_PROVIDER
model_path: Final[str | None] = _model_path_from_request_url(raw_response)
request_id: Final[str] = _response_string(response_data, "request_id") or (
_request_id_from_request_url(raw_response) or ""
)
return VideoObject(
id=encode_video_id_with_provider(request_id, provider, model_path),
object="video",
status="failed" if error else status,
created_at=0,
model=model_path,
error=(
{"code": "fal_error", "message": error} if error else None # mutable-ok: VideoObject requires a dict
),
)
@staticmethod
def _decode_video_id(video_id: str) -> tuple[str, str]:
decoded: Final = decode_video_id_with_provider(video_id)
request_id: Final[str] = decoded.get("video_id", video_id)
model_id: Final[str | None] = decoded.get("model_id")
if not model_id:
raise ValueError("fal.ai video ids must be created through litellm with a model")
return request_id, model_id
def transform_video_content_request(
self,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: _VideoHeaders,
variant: str | None = None,
) -> tuple[str, _VideoStringParams]:
request_id, model_id = self._decode_video_id(video_id)
encoded_request_id: Final[str] = encode_url_path_segment(request_id, field_name="video_id")
return (
f"{api_base.rstrip('/')}/{_queue_request_base_path(model_id)}/requests/{encoded_request_id}",
{}, # mutable-ok: BaseVideoConfig requires a mutable mapping
)
@staticmethod
def _extract_video_url(response_data: Mapping[str, object]) -> str:
raw_video_data: Final[object] = response_data.get("video")
video_data: Final[Mapping[str, object] | None] = (
TypeAdapter(Mapping[str, object]).validate_python(raw_video_data)
if isinstance(raw_video_data, Mapping)
else None
)
if video_data is not None:
video_url: Final[object] = video_data.get("url")
if isinstance(video_url, str) and video_url:
return video_url
error_message: Final[str | None] = next(
(value for key in ("error", "detail") if isinstance(value := response_data.get(key), str)),
None,
)
if error_message:
raise ValueError(f"fal.ai video result did not include a video URL: {error_message}")
raise ValueError("fal.ai video result did not include a video URL")
def transform_video_content_response(self, raw_response: httpx.Response, logging_obj: object) -> bytes:
video_url: Final[str] = self._extract_video_url(_response_data(raw_response))
httpx_client: Final[HTTPHandler] = _get_httpx_client()
video_response: Final[httpx.Response] = httpx_client.get( # pyright: ignore[reportUnknownMemberType] # HTTP handler stubs are untyped
video_url
)
video_response.raise_for_status()
return video_response.content
async def async_transform_video_content_response(self, raw_response: httpx.Response, logging_obj: object) -> bytes:
video_url: Final[str] = self._extract_video_url(_response_data(raw_response))
async_httpx_client: Final[AsyncHTTPHandler] = get_async_httpx_client(llm_provider=LlmProviders.FAL_AI)
video_response: Final[httpx.Response] = await async_httpx_client.get( # pyright: ignore[reportUnknownMemberType] # HTTP handler stubs are untyped
video_url
)
video_response.raise_for_status()
return video_response.content
def transform_video_remix_request(
self,
video_id: str,
prompt: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: _VideoHeaders,
extra_body: Mapping[str, object] | None = None,
) -> tuple[str, _VideoParams]:
raise NotImplementedError("video remix is not supported for fal.ai")
def transform_video_remix_response(
self,
raw_response: httpx.Response,
logging_obj: object,
custom_llm_provider: str | None = None,
) -> VideoObject:
raise NotImplementedError("video remix is not supported for fal.ai")
def transform_video_list_request(
self,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: _VideoHeaders,
after: str | None = None,
limit: int | None = None,
order: str | None = None,
extra_query: Mapping[str, object] | None = None,
) -> tuple[str, _VideoParams]:
raise NotImplementedError("video listing is not supported for fal.ai")
def transform_video_list_response(
self,
raw_response: httpx.Response,
logging_obj: object,
custom_llm_provider: str | None = None,
) -> _VideoStringParams:
raise NotImplementedError("video listing is not supported for fal.ai")
def transform_video_delete_request(
self,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: _VideoHeaders,
) -> tuple[str, _VideoParams]:
raise NotImplementedError("video delete is not supported for fal.ai")
def transform_video_delete_response(self, raw_response: httpx.Response, logging_obj: object) -> VideoObject:
raise NotImplementedError("video delete is not supported for fal.ai")
def transform_video_create_character_request(
self,
name: str,
video: object,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: _VideoHeaders,
) -> tuple[str, _VideoFiles]:
raise NotImplementedError("video character creation is not supported for fal.ai")
def transform_video_create_character_response(
self,
raw_response: httpx.Response,
logging_obj: object,
) -> CharacterObject:
raise NotImplementedError("video character creation is not supported for fal.ai")
def transform_video_get_character_request(
self,
character_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: _VideoHeaders,
) -> tuple[str, _VideoParams]:
raise NotImplementedError("video character retrieval is not supported for fal.ai")
def transform_video_get_character_response(
self,
raw_response: httpx.Response,
logging_obj: object,
) -> CharacterObject:
raise NotImplementedError("video character retrieval is not supported for fal.ai")
def transform_video_edit_request(
self,
prompt: str,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: _VideoHeaders,
video_file: FileContent | None = None,
extra_body: Mapping[str, object] | None = None,
prefetched_source_data: Mapping[str, object] | None = None,
) -> tuple[str, Mapping[str, object], RequestFiles | None]:
raise NotImplementedError("video edit is not supported for fal.ai")
def transform_video_edit_response(
self,
raw_response: httpx.Response,
logging_obj: object,
custom_llm_provider: str | None = None,
request_data: Mapping[str, object] | None = None,
) -> VideoObject:
raise NotImplementedError("video edit is not supported for fal.ai")
def transform_video_extension_request(
self,
prompt: str,
video_id: str,
seconds: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: _VideoHeaders,
extra_body: Mapping[str, object] | None = None,
) -> tuple[str, _VideoParams]:
raise NotImplementedError("video extension is not supported for fal.ai")
def transform_video_extension_response(
self,
raw_response: httpx.Response,
logging_obj: object,
custom_llm_provider: str | None = None,
) -> VideoObject:
raise NotImplementedError("video extension is not supported for fal.ai")
def get_error_class(
self,
error_message: str,
status_code: int,
headers: _VideoHeaders | httpx.Headers,
) -> BaseLLMException:
return FalAIVideoError(status_code=status_code, message=error_message, headers=headers)

View file

@ -170,7 +170,7 @@ class LemonadeChatConfig(OpenAILikeChatConfig):
model: str,
api_base: str | None = None,
api_key: str | None = None,
) -> Any:
) -> dict[str, object]:
if model.startswith("lemonade/"):
model = model.split("/", 1)[1]

View file

@ -66,7 +66,7 @@ class OpenAIImageEditConfig(BaseImageEditConfig):
def _add_image_to_files(
self,
files_list: list[tuple[str, Any]],
image: Any,
image: object,
field_name: str,
) -> None:
"""Add an image to the files list with appropriate content type"""

View file

@ -78,7 +78,7 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
aspeech: bool,
api_base: str | None,
api_key: str | None,
**kwargs: Any,
**kwargs: object,
) -> Union[
"HttpxBinaryResponseContent",
Coroutine[object, object, "HttpxBinaryResponseContent"],

View file

@ -651,7 +651,7 @@ def _openai_batch_jsonl_entry_to_vertex_embeddings_rows(
def _openai_batch_jsonl_entry_to_vertex_rows(
openai_entry: dict[str, Any],
map_openai_to_vertex_params: Callable[[dict[str, Any]], dict[str, Any]],
map_openai_to_vertex_params: Callable[[dict[str, Any]], dict[str, object]],
) -> tuple[Mapping[str, object], ...]:
"""
Transforms a single OpenAI JSONL batch entry into the Vertex rows it maps to.
@ -774,7 +774,7 @@ class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream):
def __init__(
self,
openai_file_content: FileTypes,
map_openai_to_vertex_params: Callable[[dict[str, Any]], dict[str, Any]],
map_openai_to_vertex_params: Callable[[dict[str, Any]], dict[str, object]],
) -> None:
self._openai_file_content = openai_file_content
self._map_openai_to_vertex_params = map_openai_to_vertex_params
@ -948,7 +948,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
def _map_openai_to_vertex_params(
self,
openai_request_body: dict[str, Any],
) -> dict[str, Any]:
) -> dict[str, object]:
"""
wrapper to call VertexGeminiConfig.map_openai_params
"""

View file

@ -3,7 +3,7 @@ Google AI Studio /batchEmbedContents Embeddings Endpoint
"""
import json
from typing import TYPE_CHECKING, Any, Final, Literal
from typing import TYPE_CHECKING, Final, Literal
import httpx
@ -210,7 +210,7 @@ class GoogleBatchEmbeddings(VertexLLM):
)
### TRANSFORMATION (sync path) ###
request_data: Any
request_data: VertexAIBatchEmbeddingsRequestBody | dict[str, object]
if use_embed_content:
resolved_files = {}
if api_key:

View file

@ -64,7 +64,7 @@ def _get_client_from_cache(client_cache_key: str):
return litellm.in_memory_llm_clients_cache.get_cache(client_cache_key)
def _set_client_in_cache(client_cache_key: str, vertex_llm_model: Any):
def _set_client_in_cache(client_cache_key: str, vertex_llm_model: object):
litellm.in_memory_llm_clients_cache.set_cache(
key=client_cache_key,
value=vertex_llm_model,

View file

@ -4,7 +4,7 @@ Translates from OpenAI's `/v1/audio/transcriptions` to IBM WatsonX's `/ml/v1/aud
WatsonX follows the OpenAI spec for audio transcription.
"""
from typing import Any, Final
from typing import Final
from httpx import Response
@ -124,7 +124,7 @@ class IBMWatsonXAudioTranscriptionConfig(IBMWatsonXMixin, OpenAIWhisperAudioTran
}
# Convert TypedDict to regular dict for AudioTranscriptionRequestData
form_data_dict: Final[dict[str, Any]] = dict(form_data)
form_data_dict: Final[dict[str, object]] = dict(form_data)
return AudioTranscriptionRequestData(data=form_data_dict, files=files)

View file

@ -22801,6 +22801,127 @@
"/v1/images/generations"
]
},
"fal_ai/bytedance/seedance-2.5/text-to-video": {
"litellm_provider": "fal_ai",
"mode": "video_generation",
"output_cost_per_second": 0.473,
"output_cost_per_second_480p": 0.2205,
"output_cost_per_second_720p": 0.473,
"source": "https://fal.ai/models/bytedance/seedance-2.5/text-to-video",
"supported_endpoints": [
"/v1/videos"
],
"supported_modalities": [
"text"
],
"supported_output_modalities": [
"video"
]
},
"fal_ai/bytedance/seedance-2.5/image-to-video": {
"litellm_provider": "fal_ai",
"mode": "video_generation",
"output_cost_per_second": 0.473,
"output_cost_per_second_480p": 0.2205,
"output_cost_per_second_720p": 0.473,
"source": "https://fal.ai/models/bytedance/seedance-2.5/image-to-video",
"supported_endpoints": [
"/v1/videos"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"video"
]
},
"fal_ai/bytedance/seedance-2.5/reference-to-video": {
"litellm_provider": "fal_ai",
"mode": "video_generation",
"output_cost_per_second": 0.473,
"output_cost_per_second_480p": 0.2205,
"output_cost_per_second_720p": 0.473,
"source": "https://fal.ai/models/bytedance/seedance-2.5/reference-to-video",
"supported_endpoints": [
"/v1/videos"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"video"
]
},
"fal_ai/bytedance/seedance-2.0/text-to-video": {
"litellm_provider": "fal_ai",
"mode": "video_generation",
"output_cost_per_second": 0.3034,
"output_cost_per_second_480p": 0.1346,
"output_cost_per_second_720p": 0.3034,
"output_cost_per_second_1080p": 0.682,
"output_cost_per_second_4k": 1.5552,
"source": "https://fal.ai/models/bytedance/seedance-2.0/text-to-video",
"metadata": {
"comment": "fal bills $0.014 per 1k tokens (480p/720p/1080p) and $0.008 per 1k tokens (4k) with tokens = h*w*seconds*24/1024; 480p and 4k rates derived from that formula at 854x480 and 3840x2160"
},
"supported_endpoints": [
"/v1/videos"
],
"supported_modalities": [
"text"
],
"supported_output_modalities": [
"video"
]
},
"fal_ai/bytedance/seedance-2.0/image-to-video": {
"litellm_provider": "fal_ai",
"mode": "video_generation",
"output_cost_per_second": 0.3034,
"output_cost_per_second_480p": 0.1346,
"output_cost_per_second_720p": 0.3034,
"output_cost_per_second_1080p": 0.682,
"output_cost_per_second_4k": 1.5552,
"source": "https://fal.ai/models/bytedance/seedance-2.0/image-to-video",
"metadata": {
"comment": "fal bills $0.014 per 1k tokens (480p/720p/1080p) and $0.008 per 1k tokens (4k) with tokens = h*w*seconds*24/1024; 480p and 4k rates derived from that formula at 854x480 and 3840x2160"
},
"supported_endpoints": [
"/v1/videos"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"video"
]
},
"fal_ai/bytedance/seedance-2.0/reference-to-video": {
"litellm_provider": "fal_ai",
"mode": "video_generation",
"output_cost_per_second": 0.3034,
"output_cost_per_second_480p": 0.1346,
"output_cost_per_second_720p": 0.3034,
"output_cost_per_second_1080p": 0.682,
"output_cost_per_second_4k": 1.5552,
"source": "https://fal.ai/models/bytedance/seedance-2.0/reference-to-video",
"metadata": {
"comment": "fal bills $0.014 per 1k tokens (480p/720p/1080p) and $0.008 per 1k tokens (4k) with tokens = h*w*seconds*24/1024; 480p and 4k rates derived from that formula at 854x480 and 3840x2160"
},
"supported_endpoints": [
"/v1/videos"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"video"
]
},
"fal_ai/fal-ai/ideogram/v3": {
"litellm_provider": "fal_ai",
"mode": "image_generation",

View file

@ -10,7 +10,7 @@ and uses LiteLLM auth.
import re
from collections.abc import Mapping
from copy import deepcopy
from typing import Any, Final, Literal
from typing import Final, Literal
SupportedA2AVersion = Literal["0.3", "1.0"]
@ -44,7 +44,7 @@ def normalize_protocol_version(version: object) -> SupportedA2AVersion | None:
return next((supported for supported in SUPPORTED_A2A_PROTOCOL_VERSIONS if supported == major_minor), None)
def resolve_served_protocol_version(card: Mapping[str, Any] | None) -> str:
def resolve_served_protocol_version(card: Mapping[str, object] | None) -> str:
"""Return the validated protocol version an agent card pins, else the default."""
normalized: Final = normalize_protocol_version(card.get("protocolVersion") if card else None)
return normalized if normalized is not None else LITELLM_A2A_PROTOCOL_VERSION
@ -53,7 +53,7 @@ def resolve_served_protocol_version(card: Mapping[str, Any] | None) -> str:
# Security scheme exposed by the LiteLLM-fronted agent card. Always replaces
# whatever upstream advertised — the client must authenticate to the proxy,
# not the upstream agent.
LITELLM_SECURITY_SCHEMES: Final[dict[str, dict[str, Any]]] = {
LITELLM_SECURITY_SCHEMES: Final[dict[str, dict[str, str]]] = {
"LiteLLMKey": {
"type": "http",
"scheme": "bearer",
@ -112,7 +112,7 @@ _ALLOWED_TOP_LEVEL_KEYS: Final = {
"url",
}
_DEFAULT_SKILLS: Final[list[dict[str, Any]]] = [
_DEFAULT_SKILLS: Final[list[dict[str, str | list[str]]]] = [
{
"id": "chat",
"name": "Chat",
@ -129,7 +129,7 @@ _DEFAULT_MODES: Final[list[str]] = ["text"]
_DEFAULT_AGENT_VERSION: Final = "1.0.0"
def _filter_capabilities(upstream_capabilities: Any) -> dict[str, Any]:
def _filter_capabilities(upstream_capabilities: object) -> dict[str, object]:
"""Return a capabilities dict containing only allowlisted, truthy keys."""
if not isinstance(upstream_capabilities, dict):
return {}
@ -143,13 +143,13 @@ def _default_litellm_provider(proxy_base_url: str) -> dict[str, str]:
def merge_agent_card(
upstream_card: Mapping[str, Any] | None,
upstream_card: Mapping[str, object] | None,
*,
proxy_url: str,
proxy_base_url: str,
name: str | None = None,
description: str | None = None,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Build the LiteLLM-fronted agent card.
@ -169,7 +169,7 @@ def merge_agent_card(
A dict suitable for serving as the proxy's agent card. Only keys in
the v1.0 AgentCard schema (plus ``supportedInterfaces``) are emitted.
"""
base: Final[dict[str, Any]] = deepcopy(dict(upstream_card)) if upstream_card else {}
base: Final[dict[str, object]] = deepcopy(dict(upstream_card)) if upstream_card else {}
# Keep the upstream ``url`` on the stored card: the runtime A2A
# invocation path reads it from ``agent_card_params`` to know where to

View file

@ -1,3 +1,4 @@
from collections.abc import Mapping
from typing import Any, Final
import requests
@ -69,8 +70,8 @@ class CredentialsManagementClient:
def create(
self,
credential_name: str,
credential_info: dict[str, Any],
credential_values: dict[str, Any],
credential_info: Mapping[str, object],
credential_values: Mapping[str, object],
return_request: bool = False,
) -> dict[str, Any] | requests.Request:
"""

View file

@ -7,7 +7,7 @@
import json
import os
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict
from typing import TYPE_CHECKING, Final, Literal, TypedDict
from fastapi import HTTPException
@ -181,7 +181,7 @@ class GuardrailsAI(CustomGuardrail):
): # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm
return await self.process_input(data=data, call_type=call_type)
async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]:
async def async_logging_hook(self, kwargs: dict, result: object, call_type: str) -> tuple[dict, object]:
if call_type == "acompletion" or call_type == "completion":
kwargs = await self.process_input(data=kwargs, call_type=call_type)

View file

@ -24,7 +24,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk
from litellm.types.proxy.guardrails.guardrail_hooks.base import (
GuardrailConfigModel,
)
@ -36,7 +36,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.singulr import (
ToolCall,
ToolCallFunction,
)
from litellm.types.utils import CallTypes, GenericGuardrailAPIInputs
from litellm.types.utils import CallTypes, ChatCompletionMessageToolCall, GenericGuardrailAPIInputs
_DEFAULT_API_BASE: Final = "http://localhost:8003"
_GUARD_ENDPOINT: Final = "/api/v1/ai-gateway/litellm-v2"
@ -339,7 +339,7 @@ class SingulrGuardrail(CustomGuardrail):
return inputs
@staticmethod
def _build_tool_call(tool_call: Mapping[str, Any]) -> "ToolCall | None":
def _build_tool_call(tool_call: ChatCompletionToolCallChunk | ChatCompletionMessageToolCall) -> "ToolCall | None":
tool_call_id: Final = tool_call.get("id")
fun: Final = tool_call.get("function")
if not tool_call_id or not fun:

View file

@ -319,7 +319,7 @@ async def _authorize_models_this_test_can_call(
its calls through the proxy. Team and member budgets are already enforced on every route.
"""
models: Final = _models_this_test_can_call(config)
if not models:
if not models and config.classifier_type != "jev":
return
from litellm.proxy.proxy_server import proxy_logging_obj
@ -345,6 +345,14 @@ async def _authorize_models_this_test_can_call(
code=status.HTTP_400_BAD_REQUEST,
) from e
if config.classifier_type == "jev" and user_api_key_dict.budget_throttle_pct is not None:
raise ProxyException(
message="Budget has been exceeded! JEV Test Routing requires available budget.",
type=ProxyErrorTypes.budget_exceeded,
param=None,
code=status.HTTP_400_BAD_REQUEST,
)
@router.post(
"/auto_router/validate_complexity_router_config",

View file

@ -8,7 +8,7 @@ pass/fail actions (allow, block, next, modify_response) and data forwarding.
import copy
import time
from collections.abc import Callable, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar
from typing import TYPE_CHECKING, Final, Literal, TypeVar
from pydantic import BaseModel
@ -314,11 +314,11 @@ class PipelineExecutor:
steps: list[PipelineStep],
mode: str,
data: dict,
user_api_key_dict: Any,
user_api_key_dict: "UserAPIKeyAuth",
call_type: str,
policy_name: str,
raw_request_snapshot: dict | None = None, # mutable-ok: same request-payload shape as data
streaming_chunks: list[Any] | None = None, # mutable-ok: shared buffered-stream chunks, read per step
streaming_chunks: list[object] | None = None, # mutable-ok: shared buffered-stream chunks, read per step
endpoint_translation: "BaseTranslation | None" = None,
) -> PipelineExecutionResult:
"""
@ -490,10 +490,10 @@ class PipelineExecutor:
step: PipelineStep,
mode: str,
data: dict,
user_api_key_dict: Any,
user_api_key_dict: "UserAPIKeyAuth",
call_type: str,
raw_request_snapshot: dict | None = None, # mutable-ok: same request-payload shape as data
streaming_chunks: list[Any] | None = None, # mutable-ok: shared buffered-stream chunks, read per step
streaming_chunks: list[object] | None = None, # mutable-ok: shared buffered-stream chunks, read per step
endpoint_translation: "BaseTranslation | None" = None,
) -> tuple[
Literal["pass", "fail", "error"],
@ -722,7 +722,7 @@ def _extract_error_message(e: Exception) -> str:
if isinstance(e, ModifyResponseException):
return str(e)
if HTTPException is not None and isinstance(e, HTTPException):
detail: Final = getattr(e, "detail", None)
detail: Final[object] = getattr(e, "detail", None)
if detail:
return str(detail)
return str(e)

View file

@ -86,7 +86,7 @@ def _resolve_session_key(kwargs: dict[str, Any]) -> str | None:
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
def _last_user_content(messages: list[dict[str, Any]] | None) -> str | None:
def _last_user_content(messages: Sequence[Mapping[str, object]] | None) -> str | None:
if not messages:
return None
for msg in reversed(messages):

View file

@ -3,7 +3,7 @@ Auto-Routing Strategy that works with a Semantic Router Config
"""
import asyncio
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Optional
from pydantic import BaseModel, ConfigDict
@ -158,7 +158,7 @@ class AutoRouter(CustomLogger):
return await asyncio.shield(build_task)
@staticmethod
def _extract_text_from_messages(messages: list[dict[str, Any]]) -> str:
def _extract_text_from_messages(messages: Sequence[Mapping[str, object]]) -> str:
"""
Extract text content from the last user message for routing.

View file

@ -9414,6 +9414,10 @@ class ProviderConfigManager:
from litellm.llms.runwayml.videos.transformation import RunwayMLVideoConfig
return RunwayMLVideoConfig()
elif LlmProviders.FAL_AI == provider:
from litellm.llms.fal_ai.videos.transformation import FalAIVideoConfig
return FalAIVideoConfig()
elif LlmProviders.HOSTED_VLLM == provider:
from litellm.llms.hosted_vllm.videos import get_hosted_vllm_video_config

View file

@ -22801,6 +22801,127 @@
"/v1/images/generations"
]
},
"fal_ai/bytedance/seedance-2.5/text-to-video": {
"litellm_provider": "fal_ai",
"mode": "video_generation",
"output_cost_per_second": 0.473,
"output_cost_per_second_480p": 0.2205,
"output_cost_per_second_720p": 0.473,
"source": "https://fal.ai/models/bytedance/seedance-2.5/text-to-video",
"supported_endpoints": [
"/v1/videos"
],
"supported_modalities": [
"text"
],
"supported_output_modalities": [
"video"
]
},
"fal_ai/bytedance/seedance-2.5/image-to-video": {
"litellm_provider": "fal_ai",
"mode": "video_generation",
"output_cost_per_second": 0.473,
"output_cost_per_second_480p": 0.2205,
"output_cost_per_second_720p": 0.473,
"source": "https://fal.ai/models/bytedance/seedance-2.5/image-to-video",
"supported_endpoints": [
"/v1/videos"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"video"
]
},
"fal_ai/bytedance/seedance-2.5/reference-to-video": {
"litellm_provider": "fal_ai",
"mode": "video_generation",
"output_cost_per_second": 0.473,
"output_cost_per_second_480p": 0.2205,
"output_cost_per_second_720p": 0.473,
"source": "https://fal.ai/models/bytedance/seedance-2.5/reference-to-video",
"supported_endpoints": [
"/v1/videos"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"video"
]
},
"fal_ai/bytedance/seedance-2.0/text-to-video": {
"litellm_provider": "fal_ai",
"mode": "video_generation",
"output_cost_per_second": 0.3034,
"output_cost_per_second_480p": 0.1346,
"output_cost_per_second_720p": 0.3034,
"output_cost_per_second_1080p": 0.682,
"output_cost_per_second_4k": 1.5552,
"source": "https://fal.ai/models/bytedance/seedance-2.0/text-to-video",
"metadata": {
"comment": "fal bills $0.014 per 1k tokens (480p/720p/1080p) and $0.008 per 1k tokens (4k) with tokens = h*w*seconds*24/1024; 480p and 4k rates derived from that formula at 854x480 and 3840x2160"
},
"supported_endpoints": [
"/v1/videos"
],
"supported_modalities": [
"text"
],
"supported_output_modalities": [
"video"
]
},
"fal_ai/bytedance/seedance-2.0/image-to-video": {
"litellm_provider": "fal_ai",
"mode": "video_generation",
"output_cost_per_second": 0.3034,
"output_cost_per_second_480p": 0.1346,
"output_cost_per_second_720p": 0.3034,
"output_cost_per_second_1080p": 0.682,
"output_cost_per_second_4k": 1.5552,
"source": "https://fal.ai/models/bytedance/seedance-2.0/image-to-video",
"metadata": {
"comment": "fal bills $0.014 per 1k tokens (480p/720p/1080p) and $0.008 per 1k tokens (4k) with tokens = h*w*seconds*24/1024; 480p and 4k rates derived from that formula at 854x480 and 3840x2160"
},
"supported_endpoints": [
"/v1/videos"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"video"
]
},
"fal_ai/bytedance/seedance-2.0/reference-to-video": {
"litellm_provider": "fal_ai",
"mode": "video_generation",
"output_cost_per_second": 0.3034,
"output_cost_per_second_480p": 0.1346,
"output_cost_per_second_720p": 0.3034,
"output_cost_per_second_1080p": 0.682,
"output_cost_per_second_4k": 1.5552,
"source": "https://fal.ai/models/bytedance/seedance-2.0/reference-to-video",
"metadata": {
"comment": "fal bills $0.014 per 1k tokens (480p/720p/1080p) and $0.008 per 1k tokens (4k) with tokens = h*w*seconds*24/1024; 480p and 4k rates derived from that formula at 854x480 and 3840x2160"
},
"supported_endpoints": [
"/v1/videos"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"video"
]
},
"fal_ai/fal-ai/ideogram/v3": {
"litellm_provider": "fal_ai",
"mode": "image_generation",

View file

@ -163,6 +163,9 @@
"other.provider_wire.anthropic.tool_history_system_cache_and_internal_fields",
"quota_management.spend_tracking.cache_tokens.disjoint_classes_use_explicit_rates"
],
"tests/integration/providers/test_fal_ai_video_wire.py::test_fal_video_create_status_and_content_follow_queue_wire_contract": [
"other.provider_wire.fal_ai.video_queue_create_status_and_content_download"
],
"tests/integration/mcp/test_mcp_lifecycle.py::test_saved_headers_reach_real_mcp_tool_and_survive_unrelated_edit": [
"mcp.call_tool.saved_headers.reach_actual_transport"
],

View file

@ -0,0 +1,69 @@
import json
import uuid
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
_MODEL: Final = "bytedance/seedance-2.5/text-to-video"
_MP4: Final = b"\x00\x00\x00\x18ftypmp42" + uuid.uuid4().bytes * 4
@pytest.mark.covers("other.provider_wire.fal_ai.video_queue_create_status_and_content_download")
def test_fal_video_create_status_and_content_follow_queue_wire_contract(gateway: Gateway) -> None:
request_id: Final = "fal-req-" + uuid.uuid4().hex
def respond(request: Request) -> Reply:
if request.target == f"/files/{request_id}.mp4":
assert request.method == "GET"
return Reply(body=_MP4, content_type="video/mp4")
assert request.headers["authorization"] == "Key synthetic-fal-key"
if request.method == "POST":
assert request.target == f"/{_MODEL}"
assert json.loads(request.body) == {
"prompt": "a cat playing volleyball on a beach",
"duration": "4",
"resolution": "720p",
"aspect_ratio": "16:9",
}
return Reply(
body=json.dumps({"status": "IN_QUEUE", "request_id": request_id, "queue_position": 0}).encode()
)
assert request.method == "GET"
if request.target == f"/bytedance/seedance-2.5/requests/{request_id}/status":
return Reply(body=json.dumps({"status": "COMPLETED", "request_id": request_id}).encode())
assert request.target == f"/bytedance/seedance-2.5/requests/{request_id}"
return Reply(body=json.dumps({"video": {"url": f"{wire_url}/files/{request_id}.mp4"}}).encode())
with wire_server(respond) as wire, gateway.scenario() as scenario:
wire_url: Final = wire.url
model: Final = scenario.model(
model=f"fal_ai/{_MODEL}",
api_base=wire.url,
api_key="synthetic-fal-key",
)
created: Final = gateway.post(
"/v1/videos",
{
"model": model,
"prompt": "a cat playing volleyball on a beach",
"seconds": "4",
"size": "1280x720",
},
)
assert created["status"] == "queued"
video_id: Final = created["id"]
assert isinstance(video_id, str) and video_id
status: Final = gateway.get(f"/v1/videos/{video_id}")
assert status["status"] == "completed"
content: Final = gateway.request("GET", f"/v1/videos/{video_id}/content")
assert content.status_code == 200, content.text
assert content.headers["content-type"].startswith("video/mp4")
assert content.content == _MP4
assert [(request.method, request.target) for request in wire.drain()] == [
("POST", f"/{_MODEL}"),
("GET", f"/bytedance/seedance-2.5/requests/{request_id}/status"),
("GET", f"/bytedance/seedance-2.5/requests/{request_id}"),
("GET", f"/files/{request_id}.mp4"),
]

View file

@ -1,153 +0,0 @@
"""
Regression test for https://github.com/BerriAI/litellm/issues/28505 -
the Responses API bridge double-strips the provider prefix from the
model name when a Chat Completions request has both `tools` and
`reasoning_effort`.
Root cause: the bridge handler called `litellm.responses()` /
`litellm.aresponses()` without passing the already-resolved
`custom_llm_provider`. The downstream call then re-invoked
`get_llm_provider()` with `custom_llm_provider=None`, which stripped
a second provider prefix from a `provider/provider/model` deployment
string.
This test pins both the sync and async bridge handler call sites:
the resolved `custom_llm_provider` must be forwarded to the underlying
`responses` / `aresponses` call so the provider isn't re-detected.
"""
from unittest.mock import MagicMock, patch
import pytest
from litellm.completion_extras.litellm_responses_transformation.handler import (
ResponsesToCompletionBridgeHandler,
)
def _validated_kwargs():
return {
"model": "openai/openai/openai/gpt-5.5",
"messages": [{"role": "user", "content": "hi"}],
"optional_params": {},
"litellm_params": {},
"headers": {},
"model_response": MagicMock(),
"logging_obj": MagicMock(),
"custom_llm_provider": "openai",
}
def test_sync_completion_forwards_custom_llm_provider():
handler = ResponsesToCompletionBridgeHandler()
handler.transformation_handler = MagicMock()
handler.transformation_handler.transform_request.return_value = {
"model": "openai/openai/openai/gpt-5.5",
"input": [],
# `_build_sanitized_litellm_params` spreads `custom_llm_provider` from
# `litellm_params` into request_data on the real bridge path. Seed
# it here so the test exercises the overwrite (not an explicit kwarg
# that would TypeError against an already-present key).
"custom_llm_provider": "should-be-overwritten",
}
handler.transformation_handler.transform_response.return_value = (
_validated_kwargs()["model_response"]
)
with (
patch.object(
handler, "validate_input_kwargs", return_value=_validated_kwargs()
),
patch(
"litellm.responses",
return_value=MagicMock(spec=[]),
) as mock_responses,
):
# The handler routes ResponsesAPIResponse through transform_response.
# We just want to verify the kwargs going INTO responses().
try:
handler.completion(acompletion=False)
except Exception:
# Downstream handling (transform_response, type checks) is not
# the subject of this test.
pass
assert mock_responses.called
kwargs = mock_responses.call_args.kwargs
assert kwargs.get("custom_llm_provider") == "openai", (
"sync bridge must forward custom_llm_provider to litellm.responses() "
"so the downstream get_llm_provider() call does not re-strip the "
"provider prefix on a provider/provider/model deployment string"
)
@pytest.mark.asyncio
async def test_async_completion_forwards_custom_llm_provider():
handler = ResponsesToCompletionBridgeHandler()
handler.transformation_handler = MagicMock()
handler.transformation_handler.transform_request.return_value = {
"model": "openai/openai/openai/gpt-5.5",
"input": [],
# `_build_sanitized_litellm_params` spreads `custom_llm_provider` from
# `litellm_params` into request_data on the real bridge path. Seed
# it here so the test exercises the overwrite (not an explicit kwarg
# that would TypeError against an already-present key).
"custom_llm_provider": "should-be-overwritten",
}
async def _fake_aresponses(**kwargs):
_fake_aresponses.kwargs = kwargs
return MagicMock(spec=[])
_fake_aresponses.kwargs = {}
with (
patch.object(
handler, "validate_input_kwargs", return_value=_validated_kwargs()
),
patch("litellm.aresponses", _fake_aresponses),
):
try:
await handler.acompletion()
except Exception:
pass
assert _fake_aresponses.kwargs.get("custom_llm_provider") == "openai", (
"async bridge must forward custom_llm_provider to litellm.aresponses() "
"so the downstream get_llm_provider() call does not re-strip the "
"provider prefix on a provider/provider/model deployment string"
)
@pytest.mark.asyncio
async def test_async_completion_forwards_aws_region_name():
handler = ResponsesToCompletionBridgeHandler()
handler.transformation_handler = MagicMock()
handler.transformation_handler.transform_request.return_value = {
"model": "openai.gpt-5.5",
"input": [],
"aws_region_name": "us-east-2",
"api_base": "https://bedrock-mantle.us-east-1.api.aws/v1",
"custom_llm_provider": "bedrock_mantle",
}
async def _fake_aresponses(**kwargs):
_fake_aresponses.kwargs = kwargs
return MagicMock(spec=[])
_fake_aresponses.kwargs = {}
validated = _validated_kwargs()
validated["custom_llm_provider"] = "bedrock_mantle"
validated["litellm_params"] = {
"aws_region_name": "us-east-2",
"api_base": "https://bedrock-mantle.us-east-1.api.aws/v1",
"custom_llm_provider": "bedrock_mantle",
}
with (
patch.object(handler, "validate_input_kwargs", return_value=validated),
patch("litellm.aresponses", _fake_aresponses),
):
try:
await handler.acompletion()
except Exception:
pass
assert _fake_aresponses.kwargs.get("aws_region_name") == "us-east-2"

View file

@ -1,158 +0,0 @@
import json
import os
from unittest.mock import Mock, patch
import pytest
import litellm
from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler
# Mock response for Bedrock image generation
mock_image_response = {"images": ["base64_encoded_image_data"], "error": None}
class TestBedrockImageGeneration:
def test_image_generation_with_api_key_bearer_token(self):
"""Test image generation with bearer token authentication"""
test_api_key = "test-bearer-token-12345"
model = "bedrock/stability.sd3-large-v1:0"
prompt = "A cute baby sea otter"
with patch(
"litellm.llms.bedrock.image_generation.image_handler.BedrockImageGeneration.image_generation"
) as mock_bedrock_image_gen:
# Setup mock response
mock_image_response_obj = litellm.ImageResponse()
mock_image_response_obj.data = [{"url": "https://example.com/image.jpg"}]
mock_bedrock_image_gen.return_value = mock_image_response_obj
response = litellm.image_generation(
model=model,
prompt=prompt,
aws_region_name="us-west-2",
api_key=test_api_key,
)
assert response is not None
assert len(response.data) > 0
mock_bedrock_image_gen.assert_called_once()
for call in mock_bedrock_image_gen.call_args_list:
if "headers" in call.kwargs:
headers = call.kwargs["headers"]
if (
"Authorization" in headers
and headers["Authorization"] == f"Bearer {test_api_key}"
):
break
def test_image_generation_with_env_variable_bearer_token(self, monkeypatch):
"""Test image generation with bearer token from environment variable"""
test_api_key = "env-bearer-token-12345"
model = "bedrock/stability.sd3-large-v1:0"
prompt = "A cute baby sea otter"
# Mock the environment variable
with (
patch.dict(os.environ, {"AWS_BEARER_TOKEN_BEDROCK": test_api_key}),
patch(
"litellm.llms.bedrock.image_generation.image_handler.BedrockImageGeneration.image_generation"
) as mock_bedrock_image_gen,
):
mock_image_response_obj = litellm.ImageResponse()
mock_image_response_obj.data = [{"url": "https://example.com/image.jpg"}]
mock_bedrock_image_gen.return_value = mock_image_response_obj
response = litellm.image_generation(
model=model, prompt=prompt, aws_region_name="us-west-2"
)
assert response is not None
assert len(response.data) > 0
mock_bedrock_image_gen.assert_called_once()
for call in mock_bedrock_image_gen.call_args_list:
if "headers" in call.kwargs:
headers = call.kwargs["headers"]
if (
"Authorization" in headers
and headers["Authorization"] == f"Bearer {test_api_key}"
):
break
@pytest.mark.asyncio
async def test_async_image_generation_with_bearer_token(self):
"""Test async image generation with bearer token authentication"""
test_api_key = "async-bearer-token-12345"
model = "bedrock/stability.sd3-large-v1:0"
prompt = "A cute baby sea otter"
with patch(
"litellm.llms.bedrock.image_generation.image_handler.BedrockImageGeneration.async_image_generation"
) as mock_async_bedrock_image_gen:
mock_image_response_obj = litellm.ImageResponse()
mock_image_response_obj.data = [{"url": "https://example.com/image.jpg"}]
mock_async_bedrock_image_gen.return_value = mock_image_response_obj
# Call async image generation with api_key parameter
response = await litellm.aimage_generation(
model=model,
prompt=prompt,
aws_region_name="us-west-2",
api_key=test_api_key,
)
assert response is not None
assert len(response.data) > 0
mock_async_bedrock_image_gen.assert_called_once()
for call in mock_async_bedrock_image_gen.call_args_list:
if "headers" in call.kwargs:
headers = call.kwargs["headers"]
if (
"Authorization" in headers
and headers["Authorization"] == f"Bearer {test_api_key}"
):
break
def test_image_generation_with_sigv4(self):
"""Test image generation falls back to SigV4 auth when no bearer token is provided"""
model = "bedrock/stability.sd3-large-v1:0"
prompt = "A cute baby sea otter"
with patch(
"litellm.llms.bedrock.image_generation.image_handler.BedrockImageGeneration.image_generation"
) as mock_bedrock_image_gen:
mock_image_response_obj = litellm.ImageResponse()
mock_image_response_obj.data = [{"url": "https://example.com/image.jpg"}]
mock_bedrock_image_gen.return_value = mock_image_response_obj
response = litellm.image_generation(
model=model, prompt=prompt, aws_region_name="us-west-2"
)
assert response is not None
assert len(response.data) > 0
mock_bedrock_image_gen.assert_called_once()
def test_image_generation_bearer_token_never_runs_the_sigv4_credential_chain(monkeypatch):
"""The deployment's AWS profile does not exist, so resolving SigV4 credentials
raises; a bearer-token deployment must still sign the request with the
bearer token alone."""
from litellm.llms.bedrock.image_generation.image_handler import BedrockImageGeneration
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "env-bearer-token-12345")
request = BedrockImageGeneration()._prepare_request(
model="amazon.nova-canvas-v1:0",
prompt="A cute baby sea otter",
optional_params={"aws_region_name": "us-west-2", "aws_profile_name": "litellm-no-such-aws-profile"},
api_base=None,
extra_headers=None,
api_key=None,
logging_obj=Mock(),
)
assert request.prepped.headers["Authorization"] == "Bearer env-bearer-token-12345"

View file

@ -0,0 +1,298 @@
from unittest.mock import Mock
import httpx
import pytest
import litellm
import litellm.llms.fal_ai.videos.transformation as fal_video_module
from litellm.cost_calculator import default_video_cost_calculator
from litellm.llms.fal_ai.videos.transformation import (
FalAIVideoConfig,
FalAIVideoError,
_queue_request_base_path,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
from litellm.types.videos.utils import decode_video_id_with_provider
from litellm.utils import ProviderConfigManager
MODEL = "bytedance/seedance-2.5/text-to-video"
class TestFalAIVideoTransformation:
def setup_method(self):
self.config = FalAIVideoConfig()
self.logging_obj = Mock()
def test_map_openai_params(self):
mapped = self.config.map_openai_params(
{
"seconds": "5",
"size": "1280x720",
"input_reference": "https://example.com/image.png",
"user": "user-123",
"generate_audio": False,
},
MODEL,
False,
)
assert mapped == {
"duration": "5",
"resolution": "720p",
"aspect_ratio": "16:9",
"image_url": "https://example.com/image.png",
"end_user_id": "user-123",
"generate_audio": False,
}
assert self.config.map_openai_params({"size": "1080x1080"}, MODEL, False) == {
"resolution": "1080p",
"aspect_ratio": "1:1",
}
assert self.config.map_openai_params({"size": "720p"}, MODEL, False) == {"resolution": "720p"}
assert self.config.map_openai_params({"size": "720x1280"}, MODEL, False) == {
"resolution": "720p",
"aspect_ratio": "9:16",
}
assert self.config.map_openai_params({"size": "1080x1920"}, MODEL, False) == {
"resolution": "1080p",
"aspect_ratio": "9:16",
}
def test_map_openai_params_rejects_non_url_input_reference(self):
with pytest.raises(ValueError, match="public image URL"):
self.config.map_openai_params({"input_reference": b"image"}, MODEL, False)
def test_transform_video_create_request(self):
body, files, url = self.config.transform_video_create_request(
model=MODEL,
prompt="A quiet ocean at sunrise",
api_base="https://queue.fal.run",
video_create_optional_request_params={
"duration": "5",
"resolution": "480p",
"aspect_ratio": "16:9",
"generate_audio": False,
"model": MODEL,
},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert url == f"https://queue.fal.run/{MODEL}"
assert files == []
assert body == {
"prompt": "A quiet ocean at sunrise",
"duration": "5",
"resolution": "480p",
"aspect_ratio": "16:9",
"generate_audio": False,
}
assert "model" not in body
def test_get_complete_url_respects_api_base_override(self):
url = self.config.get_complete_url(
model=MODEL,
api_base="https://proxy.internal/",
litellm_params={},
)
assert url == "https://proxy.internal"
def test_validate_environment_requires_fal_ai_api_key(self, monkeypatch):
monkeypatch.setattr(fal_video_module, "get_secret_str", lambda _: None)
with pytest.raises(ValueError, match="FAL_AI_API_KEY is not set"):
self.config.validate_environment(
headers={},
model=MODEL,
api_key=None,
litellm_params=GenericLiteLLMParams(),
)
def test_transform_video_create_response_encodes_model_and_usage(self):
response = Mock(spec=httpx.Response)
response.json.return_value = {"request_id": "abc"}
video = self.config.transform_video_create_response(
model=MODEL,
raw_response=response,
logging_obj=self.logging_obj,
custom_llm_provider="fal_ai",
request_data={"duration": "5", "resolution": "480p"},
)
decoded = decode_video_id_with_provider(video.id)
assert decoded["custom_llm_provider"] == "fal_ai"
assert decoded["model_id"] == MODEL
assert decoded["video_id"] == "abc"
assert video.status == "queued"
assert video.usage == {"duration_seconds": 5.0, "video_resolution": "480p"}
auto_video = self.config.transform_video_create_response(
model=MODEL,
raw_response=response,
logging_obj=self.logging_obj,
custom_llm_provider="fal_ai",
request_data={"duration": "auto"},
)
assert auto_video.usage == {"video_resolution": "720p"}
assert auto_video.seconds is None
assert auto_video.size is None
def test_status_request_uses_queue_base_path(self):
response = Mock(spec=httpx.Response)
response.json.return_value = {"request_id": "abc"}
video = self.config.transform_video_create_response(
model=MODEL,
raw_response=response,
logging_obj=self.logging_obj,
custom_llm_provider="fal_ai",
request_data={},
)
url, params = self.config.transform_video_status_retrieve_request(
video_id=video.id,
api_base="https://queue.fal.run",
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert url == "https://queue.fal.run/bytedance/seedance-2.5/requests/abc/status"
assert params == {}
assert _queue_request_base_path("workflows/owner/app/x") == "workflows/owner/app"
assert _queue_request_base_path("comfy/owner/app/x") == "comfy/owner/app"
def test_status_request_rejects_unencoded_video_id(self):
with pytest.raises(ValueError, match="must be created through litellm"):
self.config.transform_video_status_retrieve_request(
video_id="abc",
api_base="https://queue.fal.run",
litellm_params=GenericLiteLLMParams(),
headers={},
)
@pytest.mark.parametrize(
("response_data", "expected_status"),
[
({"request_id": "abc", "status": "IN_QUEUE"}, "queued"),
({"request_id": "abc", "status": "IN_PROGRESS"}, "in_progress"),
({"request_id": "abc", "status": "COMPLETED"}, "completed"),
],
)
def test_status_response_mapping(self, response_data, expected_status):
status_url = "https://queue.fal.run/bytedance/seedance-2.5/requests/abc/status"
response = httpx.Response(200, json=response_data, request=httpx.Request("GET", status_url))
video = self.config.transform_video_status_retrieve_response(
raw_response=response,
logging_obj=self.logging_obj,
custom_llm_provider="fal_ai",
)
assert video.status == expected_status
assert video.created_at == 0
decoded = decode_video_id_with_provider(video.id)
assert decoded["model_id"] == "bytedance/seedance-2.5"
assert decoded["video_id"] == "abc"
poll_url, _ = self.config.transform_video_status_retrieve_request(
video_id=video.id,
api_base="https://queue.fal.run",
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert poll_url == status_url
def test_status_response_error(self):
response_data = {
"request_id": "abc",
"status": "COMPLETED",
"error": "generation failed",
}
response = httpx.Response(
200,
json=response_data,
request=httpx.Request(
"GET",
"https://queue.fal.run/bytedance/seedance-2.5/requests/abc/status",
),
)
video = self.config.transform_video_status_retrieve_response(
raw_response=response,
logging_obj=self.logging_obj,
custom_llm_provider="fal_ai",
)
assert video.status == "failed"
assert video.error == {"code": "fal_error", "message": "generation failed"}
def test_status_response_uses_namespaced_request_url(self):
response = httpx.Response(
200,
json={"status": "IN_PROGRESS"},
request=httpx.Request(
"GET",
"https://example.com/proxy/workflows/owner/app/requests/xyz/status",
),
)
video = self.config.transform_video_status_retrieve_response(
raw_response=response,
logging_obj=self.logging_obj,
custom_llm_provider="fal_ai",
)
decoded = decode_video_id_with_provider(video.id)
assert decoded["model_id"] == "workflows/owner/app"
assert decoded["video_id"] == "xyz"
assert video.model == "workflows/owner/app"
def test_content_response_downloads_video_url(self, monkeypatch):
content_response = httpx.Response(
200,
content=b"video-bytes",
request=httpx.Request("GET", "https://cdn.example.com/video.mp4"),
)
class FakeHTTPClient:
def get(self, url):
assert url == "https://cdn.example.com/video.mp4"
return content_response
monkeypatch.setattr(fal_video_module, "_get_httpx_client", lambda: FakeHTTPClient())
response = Mock(spec=httpx.Response)
response.json.return_value = {"video": {"url": "https://cdn.example.com/video.mp4"}}
assert self.config.transform_video_content_response(response, self.logging_obj) == b"video-bytes"
def test_content_response_rejects_missing_video(self):
response = Mock(spec=httpx.Response)
response.json.return_value = {"error": "generation failed"}
with pytest.raises(ValueError, match="generation failed"):
self.config.transform_video_content_response(response, self.logging_obj)
def test_provider_config_and_error_class(self):
provider_config = ProviderConfigManager.get_provider_video_config(
model=MODEL,
provider=LlmProviders.FAL_AI,
)
assert isinstance(provider_config, FalAIVideoConfig)
assert isinstance(self.config.get_error_class("bad key", 401, {}), FalAIVideoError)
def test_video_cost_uses_tiered_rows(self):
rows = {
model: row
for model, row in litellm.model_cost.items()
if row.get("litellm_provider") == "fal_ai" and row.get("mode") == "video_generation"
}
assert rows
for model, row in rows.items():
assert default_video_cost_calculator(model, 5, "fal_ai", video_resolution="480p") == (
5 * row["output_cost_per_second_480p"]
)
assert default_video_cost_calculator(model, 5, "fal_ai", video_resolution="720p") == (
5 * row["output_cost_per_second"]
)

View file

@ -1,147 +0,0 @@
"""
Test SSL verification for hosted_vllm provider.
This test ensures that the ssl_verify parameter is properly passed through
to the HTTP client when using the hosted_vllm provider.
Issue: ssl_verify parameter was being ignored because hosted_vllm fell through
to the OpenAI catch-all path in main.py, which doesn't pass ssl_verify to the HTTP client.
"""
from unittest.mock import MagicMock, patch
import pytest
import litellm
class TestHostedVLLMSSLVerify:
"""Test suite for SSL verification in hosted_vllm provider."""
@patch("litellm.llms.custom_httpx.llm_http_handler._get_httpx_client")
def test_hosted_vllm_ssl_verify_false_sync(self, mock_get_httpx_client):
"""Test that ssl_verify=False is passed to the HTTP client for sync calls."""
# Setup mock client
mock_client = MagicMock()
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.json.return_value = {
"id": "chatcmpl-test",
"object": "chat.completion",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "Test response",
},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 10,
"completion_tokens": 5,
"total_tokens": 15,
},
}
mock_response.text = '{"id": "chatcmpl-test", "object": "chat.completion", "created": 1234567890, "model": "test-model", "choices": [{"index": 0, "message": {"role": "assistant", "content": "Test response"}, "finish_reason": "stop"}], "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}'
mock_client.post.return_value = mock_response
mock_get_httpx_client.return_value = mock_client
try:
litellm.completion(
model="hosted_vllm/test-model",
messages=[{"role": "user", "content": "Hello"}],
api_base="https://test-vllm.example.com/v1",
ssl_verify=False,
)
except Exception:
# Even if the response parsing fails, we just need to verify
# that the mock was called with the correct ssl_verify parameter
pass
# Verify _get_httpx_client was called with ssl_verify=False
mock_get_httpx_client.assert_called()
call_args = mock_get_httpx_client.call_args
# Check that params contains ssl_verify=False
if call_args[0]:
# Positional argument
params = call_args[0][0]
else:
# Keyword argument
params = call_args[1].get("params", {})
assert (
params.get("ssl_verify") is False
), f"Expected ssl_verify=False in params, got {params}"
@patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client")
@pytest.mark.asyncio
async def test_hosted_vllm_ssl_verify_false_async(
self, mock_get_async_httpx_client
):
"""Test that ssl_verify=False is passed to the HTTP client for async calls."""
# Setup mock async client
mock_client = MagicMock()
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.json.return_value = {
"id": "chatcmpl-test",
"object": "chat.completion",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "Test response",
},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 10,
"completion_tokens": 5,
"total_tokens": 15,
},
}
mock_response.text = '{"id": "chatcmpl-test", "object": "chat.completion", "created": 1234567890, "model": "test-model", "choices": [{"index": 0, "message": {"role": "assistant", "content": "Test response"}, "finish_reason": "stop"}], "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}'
async def mock_post(*args, **kwargs):
return mock_response
mock_client.post = mock_post
mock_get_async_httpx_client.return_value = mock_client
try:
await litellm.acompletion(
model="hosted_vllm/test-model",
messages=[{"role": "user", "content": "Hello"}],
api_base="https://test-vllm.example.com/v1",
ssl_verify=False,
)
except Exception:
# Even if the response parsing fails, we just need to verify
# that the mock was called with the correct ssl_verify parameter
pass
# Verify get_async_httpx_client was called with ssl_verify=False
mock_get_async_httpx_client.assert_called()
call_kwargs = mock_get_async_httpx_client.call_args[1]
# Check that params contains ssl_verify=False
params = call_kwargs.get("params", {})
assert (
params.get("ssl_verify") is False
), f"Expected ssl_verify=False in params, got {params}"
if __name__ == "__main__":
pytest.main([__file__, "-v", "-s"])

View file

@ -1,135 +0,0 @@
"""
Test SSL verification for hosted_vllm provider embeddings.
This test ensures that the ssl_verify parameter is properly passed through
to the HTTP client when using the hosted_vllm provider for embeddings.
Issue: ssl_verify parameter was being ignored because hosted_vllm fell through
to the openai_like catch-all path in main.py, which doesn't pass ssl_verify to the HTTP client.
"""
from unittest.mock import MagicMock, patch
import pytest
import litellm
class TestHostedVLLMEmbeddingSSLVerify:
"""Test suite for SSL verification in hosted_vllm provider embeddings."""
@patch("litellm.llms.custom_httpx.llm_http_handler._get_httpx_client")
def test_hosted_vllm_embedding_ssl_verify_false_sync(self, mock_get_httpx_client):
"""Test that ssl_verify=False is passed to the HTTP client for sync embedding calls."""
# Setup mock client
mock_client = MagicMock()
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.json.return_value = {
"object": "list",
"data": [
{
"object": "embedding",
"index": 0,
"embedding": [0.1, 0.2, 0.3, 0.4, 0.5],
}
],
"model": "text-embedding-model",
"usage": {
"prompt_tokens": 5,
"total_tokens": 5,
},
}
mock_response.text = '{"object": "list", "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3, 0.4, 0.5]}], "model": "text-embedding-model", "usage": {"prompt_tokens": 5, "total_tokens": 5}}'
mock_client.post.return_value = mock_response
mock_get_httpx_client.return_value = mock_client
try:
litellm.embedding(
model="hosted_vllm/text-embedding-model",
input=["hello world"],
api_base="https://test-vllm.example.com/v1",
ssl_verify=False,
)
except Exception:
# Even if the response parsing fails, we just need to verify
# that the mock was called with the correct ssl_verify parameter
pass
# Verify _get_httpx_client was called with ssl_verify=False
mock_get_httpx_client.assert_called()
call_args = mock_get_httpx_client.call_args
# Check that params contains ssl_verify=False
if call_args[0]:
# Positional argument
params = call_args[0][0]
else:
# Keyword argument
params = call_args[1].get("params", {})
assert (
params.get("ssl_verify") is False
), f"Expected ssl_verify=False in params, got {params}"
@patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client")
@pytest.mark.asyncio
async def test_hosted_vllm_embedding_ssl_verify_false_async(
self, mock_get_async_httpx_client
):
"""Test that ssl_verify=False is passed to the HTTP client for async embedding calls."""
# Setup mock async client
mock_client = MagicMock()
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.json.return_value = {
"object": "list",
"data": [
{
"object": "embedding",
"index": 0,
"embedding": [0.1, 0.2, 0.3, 0.4, 0.5],
}
],
"model": "text-embedding-model",
"usage": {
"prompt_tokens": 5,
"total_tokens": 5,
},
}
mock_response.text = '{"object": "list", "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3, 0.4, 0.5]}], "model": "text-embedding-model", "usage": {"prompt_tokens": 5, "total_tokens": 5}}'
async def mock_post(*args, **kwargs):
return mock_response
mock_client.post = mock_post
mock_get_async_httpx_client.return_value = mock_client
try:
await litellm.aembedding(
model="hosted_vllm/text-embedding-model",
input=["hello world"],
api_base="https://test-vllm.example.com/v1",
ssl_verify=False,
)
except Exception:
# Even if the response parsing fails, we just need to verify
# that the mock was called with the correct ssl_verify parameter
pass
# Verify get_async_httpx_client was called with ssl_verify=False
mock_get_async_httpx_client.assert_called()
call_kwargs = mock_get_async_httpx_client.call_args[1]
# Check that params contains ssl_verify=False
params = call_kwargs.get("params", {})
assert (
params.get("ssl_verify") is False
), f"Expected ssl_verify=False in params, got {params}"
if __name__ == "__main__":
pytest.main([__file__, "-v", "-s"])

View file

@ -1 +0,0 @@
"""OpenAI Evals API tests"""

View file

@ -1 +0,0 @@
# Test module for OpenAI-like embedding handler

View file

@ -1,337 +0,0 @@
"""
Tests for IBM WatsonX Audio Transcription.
Validates that litellm.transcription transforms requests correctly for WatsonX.
"""
import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import litellm
from litellm.llms.watsonx.audio_transcription.transformation import (
IBMWatsonXAudioTranscriptionConfig,
)
from litellm.types.utils import TranscriptionResponse
class TestWatsonXAudioTranscription:
"""Tests for WatsonX audio transcription via litellm.transcription."""
@pytest.mark.asyncio
async def test_watsonx_transcription_url_and_headers(self):
"""
Test that litellm.transcription sends request to correct WatsonX URL with proper headers.
"""
captured_request = {}
async def mock_post(*args, **kwargs):
captured_request["url"] = str(kwargs.get("url", args[0] if args else None))
captured_request["headers"] = kwargs.get("headers", {})
captured_request["data"] = kwargs.get("data", {})
captured_request["files"] = kwargs.get("files", {})
mock_response = MagicMock()
mock_response.json.return_value = {
"text": "test transcription",
"duration": 1.0,
}
mock_response.status_code = 200
return mock_response
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=mock_post,
):
try:
await litellm.atranscription(
model="watsonx/whisper-large-v3-turbo",
file=b"fake_audio_data",
api_base="https://us-south.ml.cloud.ibm.com",
api_key="test-api-key",
project_id="test-project-123",
token="test-bearer-token",
)
except Exception:
pass # We just want to capture the request
# Validate URL contains WatsonX audio transcription endpoint
assert "/ml/v1/audio/transcriptions" in captured_request["url"]
assert "version=" in captured_request["url"]
# project_id should NOT be in URL (it should be in form data instead)
assert "project_id=test-project-123" not in captured_request["url"]
# Validate headers contain WatsonX auth
assert "Authorization" in captured_request["headers"]
assert (
"Bearer test-bearer-token" in captured_request["headers"]["Authorization"]
)
# Validate Content-Type is NOT set (httpx sets multipart/form-data automatically)
assert "Content-Type" not in captured_request["headers"]
# Validate project_id is in form data, not URL
assert captured_request["data"].get("project_id") == "test-project-123"
# Validate file is in files dict
assert "file" in captured_request["files"]
@pytest.mark.asyncio
async def test_watsonx_transcription_request_body(self):
"""
Test that litellm.transcription sends correct request body for WatsonX.
Validates that:
- Request uses multipart/form-data (data + files)
- Model name has watsonx/ prefix removed
- project_id is in form data, not URL
- Audio file is in files dict
- OpenAI params are included in form data
"""
captured_request = {}
async def mock_post(*args, **kwargs):
captured_request["data"] = kwargs.get("data", {})
captured_request["files"] = kwargs.get("files", {})
mock_response = MagicMock()
mock_response.json.return_value = {
"text": "test transcription",
"duration": 1.0,
}
mock_response.status_code = 200
return mock_response
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=mock_post,
):
try:
await litellm.atranscription(
model="watsonx/whisper-large-v3-turbo",
file=b"fake_audio_data",
api_base="https://us-south.ml.cloud.ibm.com",
api_key="test-api-key",
project_id="test-project-123",
token="test-bearer-token",
language="en",
temperature=0.5,
)
except Exception:
pass # We just want to capture the request
# Validate form data contains expected fields
data = captured_request.get("data", {})
print("JSON DUMPS captured_request:")
print(json.dumps(captured_request, indent=4, default=str))
# Model name should NOT have watsonx/ prefix
assert data.get("model") == "whisper-large-v3-turbo"
# project_id should be in form data
assert data.get("project_id") == "test-project-123"
# OpenAI params should be in form data
assert data.get("language") == "en"
assert data.get("temperature") == 0.5
# response_format should NOT be set by default - only send what user specifies
assert "response_format" not in data
# Validate file is in files dict (multipart/form-data)
files = captured_request.get("files", {})
assert "file" in files
assert isinstance(
files["file"], tuple
) # Should be (filename, content, content_type)
@pytest.mark.asyncio
async def test_watsonx_transcription_only_user_params_sent_with_project_id(self):
"""
Test that only user-specified params are sent in request body to WatsonX.
LiteLLM should NOT add extra params like response_format if user didn't specify them.
"""
captured_request = {}
async def mock_post(*args, **kwargs):
captured_request["data"] = kwargs.get("data", {})
captured_request["files"] = kwargs.get("files", {})
mock_response = MagicMock()
mock_response.json.return_value = {
"text": "test transcription",
"duration": 1.0,
}
mock_response.status_code = 200
return mock_response
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=mock_post,
):
try:
# Minimal request - only required params
await litellm.atranscription(
model="watsonx/whisper-large-v3-turbo",
file=b"fake_audio_data",
api_base="https://us-south.ml.cloud.ibm.com",
api_key="test-api-key",
project_id="test-project-123",
token="test-bearer-token",
)
except Exception:
pass # We just want to capture the request
data = captured_request.get("data", {})
# These are the ONLY keys that should be in data
expected_keys = {"model", "project_id"}
actual_keys = set(data.keys())
assert actual_keys == expected_keys, (
f"Request body should only contain {expected_keys}, "
f"but got {actual_keys}. "
f"Extra keys: {actual_keys - expected_keys}"
)
# Specifically verify response_format is NOT added
assert (
"response_format" not in data
), "response_format should NOT be added by default"
# Verify file is sent separately
files = captured_request.get("files", {})
assert "file" in files
@pytest.mark.asyncio
async def test_watsonx_transcription_only_user_params_sent_with_space_id(self):
"""
Test that only user-specified params are sent in request body to WatsonX.
LiteLLM should NOT add extra params like response_format if user didn't specify them.
"""
captured_request = {}
async def mock_post(*args, **kwargs):
captured_request["data"] = kwargs.get("data", {})
captured_request["files"] = kwargs.get("files", {})
mock_response = MagicMock()
mock_response.json.return_value = {
"text": "test transcription",
"duration": 1.0,
}
mock_response.status_code = 200
return mock_response
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=mock_post,
):
try:
# Minimal request - only required params
await litellm.atranscription(
model="watsonx/whisper-large-v3-turbo",
file=b"fake_audio_data",
api_base="https://us-south.ml.cloud.ibm.com",
api_key="test-api-key",
space_id="test-space_id-123",
token="test-bearer-token",
)
except Exception:
pass # We just want to capture the request
data = captured_request.get("data", {})
# These are the ONLY keys that should be in data
expected_keys = {"model", "space_id"}
actual_keys = set(data.keys())
assert actual_keys == expected_keys, (
f"Request body should only contain {expected_keys}, "
f"but got {actual_keys}. "
f"Extra keys: {actual_keys - expected_keys}"
)
# Specifically verify response_format is NOT added
assert (
"response_format" not in data
), "response_format should NOT be added by default"
# Verify file is sent separately
files = captured_request.get("files", {})
assert "file" in files
def test_transform_audio_transcription_response_removes_model_field(self):
"""
Test that transform_audio_transcription_response removes the 'model' field
from WatsonX response before creating TranscriptionResponse.
This test ensures that when WatsonX returns a response with a 'model' field,
it is removed before creating the TranscriptionResponse object, since
TranscriptionResponse doesn't accept a 'model' parameter.
"""
handler = IBMWatsonXAudioTranscriptionConfig()
# Mock response with 'model' field (as WatsonX may return)
mock_response = MagicMock()
mock_response.json.return_value = {
"text": "Hello, this is a test transcription.",
"model": "whisper-large-v3-turbo", # This field should be removed
"duration": 5.5,
}
mock_response.text = '{"text": "Hello, this is a test transcription.", "model": "whisper-large-v3-turbo", "duration": 5.5}'
# This should not raise a TypeError - model field should be removed
result = handler.transform_audio_transcription_response(mock_response)
# Verify the result is a TranscriptionResponse
assert isinstance(result, TranscriptionResponse)
# Verify the text is correct
assert result.text == "Hello, this is a test transcription."
# Verify duration is set via dictionary assignment
assert result["duration"] == 5.5
# Verify the model field is NOT in the serialized result
# Check via model_dump() or dict() to ensure it's not in the output
try:
result_dict = result.model_dump()
except AttributeError:
# Fallback for pydantic v1
result_dict = result.dict()
# The 'model' field should not be in the result
assert "model" not in result_dict, "Model field should be removed from response"
def test_transform_audio_transcription_response_without_model_field(self):
"""
Test that transform_audio_transcription_response works correctly
when WatsonX response doesn't include a 'model' field.
"""
handler = IBMWatsonXAudioTranscriptionConfig()
# Mock response without 'model' field
mock_response = MagicMock()
mock_response.json.return_value = {
"text": "Hello, this is a test transcription.",
"duration": 5.5,
}
mock_response.text = (
'{"text": "Hello, this is a test transcription.", "duration": 5.5}'
)
result = handler.transform_audio_transcription_response(mock_response)
# Verify the result is a TranscriptionResponse
assert isinstance(result, TranscriptionResponse)
# Verify the text is correct
assert result.text == "Hello, this is a test transcription."
# Verify duration is set via dictionary assignment
assert result["duration"] == 5.5

View file

@ -1,577 +0,0 @@
import json
from typing import Optional
from unittest.mock import Mock, patch
import pytest
import litellm
from litellm import completion
from litellm.llms.custom_httpx.http_handler import HTTPHandler
@pytest.fixture
def watsonx_chat_completion_call():
def _call(
model="watsonx/my-test-model",
messages=None,
api_key="test_api_key",
space_id: Optional[str] = None,
headers=None,
client=None,
patch_token_call=True,
):
if messages is None:
messages = [{"role": "user", "content": "Hello, how are you?"}]
if client is None:
client = HTTPHandler()
if patch_token_call:
mock_response = Mock()
mock_response.json.return_value = {
"access_token": "mock_access_token",
"expires_in": 3600,
}
mock_response.raise_for_status = Mock() # No-op to simulate no exception
with (
patch.object(client, "post") as mock_post,
patch.object(
litellm.module_level_client, "post", return_value=mock_response
) as mock_get,
):
try:
completion(
model=model,
messages=messages,
api_key=api_key,
headers=headers or {},
client=client,
space_id=space_id,
)
except Exception as e:
print(e)
return mock_post, mock_get
else:
with patch.object(client, "post") as mock_post:
try:
completion(
model=model,
messages=messages,
api_key=api_key,
headers=headers or {},
client=client,
space_id=space_id,
)
except Exception as e:
print(e)
return mock_post, None
return _call
def test_watsonx_deployment_model_id_not_in_payload(
monkeypatch, watsonx_chat_completion_call
):
"""Test that deployment models do not include 'model_id' in the request payload"""
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
model = "watsonx/deployment/test-deployment-id"
messages = [{"role": "user", "content": "Test message"}]
mock_post, _ = watsonx_chat_completion_call(model=model, messages=messages)
assert mock_post.call_count == 1
json_data = json.loads(mock_post.call_args.kwargs["data"])
# Ensure model_id is not in the payload for deployment models
assert "model_id" not in json_data or json_data["model_id"] is None
# Ensure project_id is also not in the payload for deployment models
assert "project_id" not in json_data or json_data["project_id"] is None
def test_watsonx_regular_model_includes_model_id(
monkeypatch, watsonx_chat_completion_call
):
"""Test that regular models include 'model_id' in the request payload"""
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
model = "watsonx/regular-model"
messages = [{"role": "user", "content": "Test message"}]
mock_post, _ = watsonx_chat_completion_call(model=model, messages=messages)
assert mock_post.call_count == 1
json_data = json.loads(mock_post.call_args.kwargs["data"])
# Ensure model_id is included in the payload for regular models
assert "model_id" in json_data
assert json_data["model_id"] == "regular-model" # Provider prefix is stripped
# Ensure project_id is also included for regular models
assert "project_id" in json_data
@pytest.fixture
def watsonx_completion_call():
def _call(
model="watsonx_text/my-test-model",
prompt="Hello, how are you?",
api_key="test_api_key",
space_id: Optional[str] = None,
headers=None,
client=None,
patch_token_call=True,
):
if client is None:
client = HTTPHandler()
if patch_token_call:
mock_response = Mock()
mock_response.json.return_value = {
"access_token": "mock_access_token",
"expires_in": 3600,
}
mock_response.raise_for_status = Mock()
with (
patch.object(client, "post") as mock_post,
patch.object(
litellm.module_level_client, "post", return_value=mock_response
) as mock_get,
):
try:
litellm.text_completion(
model=model,
prompt=prompt,
api_key=api_key,
headers=headers or {},
client=client,
space_id=space_id,
)
except Exception as e:
print(e)
return mock_post, mock_get
else:
with patch.object(client, "post") as mock_post:
try:
litellm.text_completion(
model=model,
prompt=prompt,
api_key=api_key,
headers=headers or {},
client=client,
space_id=space_id,
)
except Exception as e:
print(e)
return mock_post, None
return _call
def test_watsonx_completion_deployment_model_id_not_in_payload(
monkeypatch, watsonx_completion_call
):
"""Test that deployment models do not include 'model_id' in completion request payload"""
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
model = "watsonx_text/deployment/test-deployment-id"
prompt = "Test prompt"
mock_post, _ = watsonx_completion_call(model=model, prompt=prompt)
assert mock_post.call_count == 1
json_data = json.loads(mock_post.call_args.kwargs["data"])
# Ensure model_id is not in the payload for deployment models
assert "model_id" not in json_data
# Ensure project_id is also not in the payload for deployment models
assert "project_id" not in json_data
def test_watsonx_completion_regular_model_includes_model_id(
monkeypatch, watsonx_completion_call
):
"""Test that regular models include 'model_id' in completion request payload"""
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
model = "watsonx_text/regular-model"
prompt = "Test prompt"
mock_post, _ = watsonx_completion_call(model=model, prompt=prompt)
assert mock_post.call_count == 1
json_data = json.loads(mock_post.call_args.kwargs["data"])
# Ensure model_id is included in the payload for regular models
assert "model_id" in json_data
assert json_data["model_id"] == "regular-model" # Provider prefix is stripped
# Ensure project_id is also included for regular models
assert "project_id" in json_data
def test_watsonx_gpt_oss_prompt_transformation(monkeypatch):
"""
Test that gpt-oss-120b model transforms messages to proper format instead of simple concatenation.
This test calls litellm.completion (sync) and verifies what gets sent in the final POST request body.
Input messages should be transformed using the HuggingFace chat template from openai/gpt-oss-120b,
not just concatenated as "You are chatgpt Hi there".
"""
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
# Test with gpt-oss model using watsonx_text provider (text generation endpoint)
model = "watsonx_text/openai/gpt-oss-120b"
# Input messages
messages = [
{"role": "system", "content": "You are chatgpt"},
{"role": "user", "content": "Hi there"},
]
client = HTTPHandler()
# Mock HuggingFace template fetch to make test deterministic and avoid network flakiness.
# The test verifies that prompt transformation occurs (not simple concatenation), not the exact
# HuggingFace template format. Using a mock template that produces the correct format is sufficient.
#
# Mock template that produces gpt-oss-120b-like format.
# Note: This is a simplified version of the actual template. The real template is more complex
# (adds metadata, handles tools, thinking messages, etc.), but this captures the key aspects:
# - Converts system role to developer (matching real template behavior)
# - Uses the same tag structure (<|start|>, <|message|>, <|end|>)
# - Preserves message content
mock_tokenizer_config = {
"status": "success",
"tokenizer": {
"chat_template": "{% for message in messages %}{% if message['role'] == 'system' %}<|start|>developer<|message|>{% else %}<|start|>{{ message['role'] }}<|message|>{% endif %}{{ message['content'] }}<|end|>{% endfor %}",
"bos_token": None,
"eos_token": None,
},
}
# Isolate known_tokenizer_config so parallel tests don't interfere.
# monkeypatch.setitem restores the original value on teardown.
hf_model = "openai/gpt-oss-120b"
monkeypatch.setitem(litellm.known_tokenizer_config, hf_model, mock_tokenizer_config)
# Mock IAM token generation to avoid real HTTP calls.
mock_token_response = Mock()
mock_token_response.json.return_value = {
"access_token": "mock_access_token",
"expires_in": 3600,
}
mock_token_response.raise_for_status = Mock()
with (
patch.object(client, "post") as mock_post,
patch.object(
litellm.module_level_client, "post", return_value=mock_token_response
),
):
try:
completion(
model=model,
messages=messages,
api_key="test_api_key",
client=client,
)
except Exception as e:
print(f"Caught expected exception: {e}")
# Verify the POST was called
assert (
mock_post.call_count == 1
), f"POST should have been called exactly once, got {mock_post.call_count}"
# Get the request body
call_args = mock_post.call_args
assert "data" in call_args.kwargs, "call_args.kwargs should contain 'data'"
json_data = json.loads(call_args.kwargs["data"])
# Verify the transformed input is in the request
assert "input" in json_data, "Request should have 'input' field"
transformed_prompt = json_data["input"]
# Verify it's NOT simple concatenation
simple_concat = "You are chatgpt Hi there"
assert transformed_prompt != simple_concat, (
f"Prompt should not be simple concatenation.\n"
f"Expected: Chat template with <|start|> tags\n"
f"Got: {transformed_prompt}"
)
# Verify it contains proper chat template formatting
assert "<|start|>" in transformed_prompt, "Prompt should contain <|start|> tag"
assert "<|message|>" in transformed_prompt, "Prompt should contain <|message|> tag"
assert "<|end|>" in transformed_prompt, "Prompt should contain <|end|> tag"
assert (
"You are chatgpt" in transformed_prompt
), "Prompt should contain system message content"
assert (
"Hi there" in transformed_prompt
), "Prompt should contain user message content"
@pytest.mark.asyncio
@pytest.mark.xdist_group("watsonx_heavy")
async def test_watsonx_gpt_oss_uses_async_http_handler():
"""
Test that verifies async HTTP client is used when fetching HuggingFace templates.
"""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import (
_aget_chat_template_file,
)
# Mock the async HTTP client
mock_async_client = MagicMock()
mock_get = AsyncMock()
mock_async_client.get = mock_get
# Create mock response for chat template file
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.content = b"test template content"
mock_get.return_value = mock_response
# Test the async function directly
with patch(
"litellm.litellm_core_utils.prompt_templates.huggingface_template_handler.get_async_httpx_client",
return_value=mock_async_client,
):
result = await _aget_chat_template_file(hf_model_name="test/model")
# Verify async HTTP client was called
assert mock_get.called, "Async HTTP client's get method should be called"
assert mock_get.await_count > 0, "Async HTTP client's get should be awaited"
# Verify it was called with HuggingFace URL
call_args = mock_get.call_args
assert call_args is not None, "get should have been called with arguments"
called_url = call_args.kwargs.get("url", "")
assert (
"huggingface.co/test/model" in called_url
), f"Should call HuggingFace API for test/model, got: {called_url}"
assert result["status"] == "success", "Should return success status"
@pytest.mark.parametrize("tokenizer_config_cached", [False, True], ids=["tokenizer_config", "cached_config_jinja"])
async def test_watsonx_text_gpt_oss_async_completion_fetches_hf_template_off_the_event_loop(
monkeypatch, tokenizer_config_cached
):
import httpx
from litellm._uuid import uuid
from litellm.litellm_core_utils.prompt_templates import huggingface_template_handler
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
hf_model = f"openai/gpt-oss-{uuid.uuid4()}"
chat_template = "{% for m in messages %}<|{{ m['role'] }}|>{{ m['content'] }}{% endfor %}"
if tokenizer_config_cached:
cached_config = {"status": "success", "tokenizer": {"bos_token": None, "eos_token": None}}
monkeypatch.setattr(litellm, "known_tokenizer_config", {hf_model: cached_config})
expected_fetch = f"https://huggingface.co/{hf_model}/raw/main/chat_template.jinja"
else:
monkeypatch.setattr(litellm, "known_tokenizer_config", {})
expected_fetch = f"https://huggingface.co/{hf_model}/raw/main/tokenizer_config.json"
hf_fetched = []
captured = {}
def forbid_sync_client():
raise AssertionError("sync HuggingFace fetch ran on the request path")
async def serve_hf_file(url, **kwargs):
hf_fetched.append(url)
if url.endswith(".jinja"):
return httpx.Response(200, content=chat_template.encode())
return httpx.Response(200, json={"chat_template": chat_template, "bos_token": None, "eos_token": None})
monkeypatch.setattr(huggingface_template_handler, "_get_httpx_client", forbid_sync_client)
monkeypatch.setattr(huggingface_template_handler, "get_async_httpx_client", lambda **kwargs: Mock(get=serve_hf_file))
def handle(request):
captured["body"] = json.loads(request.content)
return httpx.Response(
200,
json={
"model_id": hf_model,
"results": [
{
"generated_text": "Hi",
"generated_token_count": 1,
"input_token_count": 1,
"stop_reason": "eos_token",
}
],
},
)
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(handle))
response = await litellm.acompletion(
model=f"watsonx_text/{hf_model}",
messages=[{"role": "user", "content": "Hi there"}],
api_base="https://test-api.watsonx.ai",
project_id="test-project-id",
token="test-token",
client=client,
)
assert response.choices[0].message.content == "Hi"
assert hf_fetched == [expected_fetch]
assert captured["body"]["input"] == "<|user|>Hi there"
def test_watsonx_chat_completion_with_reasoning_effort(monkeypatch):
"""
Test that 'reasoning_effort' is correctly passed through to the WatsonX API payload.
"""
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
model = "watsonx/openai/gpt-oss-120b"
messages = [{"role": "user", "content": "Test message"}]
client = HTTPHandler()
# Mock the token generation call
mock_token_response = Mock()
mock_token_response.json.return_value = {
"access_token": "mock_access_token",
"expires_in": 3600,
}
mock_token_response.raise_for_status = Mock()
# Call litellm.completion with the new parameter
with (
patch.object(client, "post") as mock_post,
patch.object(
litellm.module_level_client, "post", return_value=mock_token_response
),
):
try:
completion(
model=model,
messages=messages,
api_key="test_api_key",
client=client,
reasoning_effort="low",
)
except Exception as e:
print(f"Caught expected exception: {e}")
# Verify the parameter is in the final request payload
assert (
mock_post.call_count == 1
), "The completion endpoint should have been called once."
# Get the JSON data sent in the POST request
request_kwargs = mock_post.call_args.kwargs
json_data = json.loads(request_kwargs["data"])
print("\nRequest payload sent to WatsonX API:")
print(json.dumps(json_data, indent=2))
# Check for the parameter at the top level of the payload
assert (
"reasoning_effort" in json_data
), "'reasoning_effort' should be at the top level of the payload."
assert (
json_data["reasoning_effort"] == "low"
), "The value of 'reasoning_effort' should be 'low'."
def test_watsonx_zen_api_key_from_client(monkeypatch, watsonx_chat_completion_call):
"""
Test that zen_api_key can be passed from client code and is used in Authorization header.
"""
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
model = "watsonx/ibm/granite-3-3-8b-instruct"
messages = [{"role": "user", "content": "What is your favorite color?"}]
client = HTTPHandler()
zen_api_key = "U1ZDLWQo="
# No need to patch token call since zen_api_key should skip token generation
with patch.object(client, "post") as mock_post:
try:
completion(
model=model,
messages=messages,
api_key="test_api_key",
client=client,
zen_api_key=zen_api_key,
)
except Exception as e:
print(f"Caught expected exception: {e}")
# Verify the request was made
assert (
mock_post.call_count == 1
), "The completion endpoint should have been called once."
# Get the headers sent in the POST request
request_kwargs = mock_post.call_args.kwargs
headers = request_kwargs["headers"]
print("\nHeaders sent to WatsonX API:")
print(json.dumps(dict(headers), indent=2))
# Verify Authorization header uses ZenApiKey format
assert "Authorization" in headers, "Authorization header should be present."
assert headers["Authorization"] == f"ZenApiKey {zen_api_key}", (
f"Authorization header should use ZenApiKey format. "
f"Expected: 'ZenApiKey {zen_api_key}', Got: '{headers['Authorization']}'"
)
def test_watsonx_zen_api_key_from_env(monkeypatch, watsonx_chat_completion_call):
"""
Test that zen_api_key from environment variable is used in Authorization header.
"""
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
zen_api_key = "U1ZDLWxpdG--==="
monkeypatch.setenv("WATSONX_ZENAPIKEY", zen_api_key)
model = "watsonx/ibm/granite-3-3-8b-instruct"
messages = [{"role": "user", "content": "What is your favorite color?"}]
client = HTTPHandler()
# No need to patch token call since zen_api_key should skip token generation
with patch.object(client, "post") as mock_post:
try:
completion(
model=model,
messages=messages,
api_key="test_api_key",
client=client,
)
except Exception as e:
print(f"Caught expected exception: {e}")
# Verify the request was made
assert (
mock_post.call_count == 1
), "The completion endpoint should have been called once."
# Get the headers sent in the POST request
request_kwargs = mock_post.call_args.kwargs
headers = request_kwargs["headers"]
print("\nHeaders sent to WatsonX API:")
print(json.dumps(dict(headers), indent=2))
# Verify Authorization header uses ZenApiKey format
assert "Authorization" in headers, "Authorization header should be present."
assert headers["Authorization"] == f"ZenApiKey {zen_api_key}", (
f"Authorization header should use ZenApiKey format. "
f"Expected: 'ZenApiKey {zen_api_key}', Got: '{headers['Authorization']}'"
)

View file

@ -1 +0,0 @@
# XAI Responses API tests

View file

@ -1,105 +0,0 @@
"""
Tests for XAI Responses API transformation
Tests the XAIResponsesAPIConfig class that handles XAI-specific
transformations for the Responses API.
Source: litellm/llms/xai/responses/transformation.py
"""
import pytest
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams
class TestXAIResponsesAPITransformation:
"""Test XAI Responses API configuration and transformations"""
def test_xai_provider_config_registration(self):
"""Test that XAI provider returns XAIResponsesAPIConfig"""
config = ProviderConfigManager.get_provider_responses_api_config(
model="xai/grok-4-fast",
provider=LlmProviders.XAI,
)
assert config is not None, "Config should not be None for XAI provider"
assert isinstance(
config, XAIResponsesAPIConfig
), f"Expected XAIResponsesAPIConfig, got {type(config)}"
assert (
config.custom_llm_provider == LlmProviders.XAI
), "custom_llm_provider should be XAI"
def test_code_interpreter_container_field_removed(self):
"""Test that container field is removed from code_interpreter tools"""
config = XAIResponsesAPIConfig()
params = ResponsesAPIOptionalRequestParams(
tools=[{"type": "code_interpreter", "container": {"type": "auto"}}]
)
result = config.map_openai_params(
response_api_optional_params=params, model="grok-4-fast", drop_params=False
)
assert "tools" in result
assert len(result["tools"]) == 1
assert result["tools"][0]["type"] == "code_interpreter"
assert (
"container" not in result["tools"][0]
), "Container field should be removed"
def test_instructions_parameter_forwarded(self):
"""xAI supports 'instructions' on /v1/responses, so it must survive param mapping"""
config = XAIResponsesAPIConfig()
params = ResponsesAPIOptionalRequestParams(
instructions="You are a helpful assistant.", temperature=0.7
)
result = config.map_openai_params(
response_api_optional_params=params, model="grok-4-fast", drop_params=False
)
assert result.get("instructions") == "You are a helpful assistant."
assert result.get("temperature") == 0.7, "Other params should be preserved"
def test_supported_params_includes_instructions(self):
"""A system message bridged to 'instructions' must not be rejected for xAI"""
config = XAIResponsesAPIConfig()
supported = config.get_supported_openai_params("grok-4-fast")
assert "instructions" in supported, "instructions should be supported"
assert "tools" in supported, "tools should be supported"
assert "temperature" in supported, "temperature should be supported"
assert "model" in supported, "model should be supported"
def test_xai_responses_endpoint_url(self):
"""Test that get_complete_url returns correct XAI endpoint"""
config = XAIResponsesAPIConfig()
# Test with default XAI API base
url = config.get_complete_url(api_base=None, litellm_params={})
assert (
url == "https://api.x.ai/v1/responses"
), f"Expected XAI responses endpoint, got {url}"
# Test with custom api_base
custom_url = config.get_complete_url(
api_base="https://custom.x.ai/v1", litellm_params={}
)
assert (
custom_url == "https://custom.x.ai/v1/responses"
), f"Expected custom endpoint, got {custom_url}"
# Test with trailing slash
url_with_slash = config.get_complete_url(
api_base="https://api.x.ai/v1/", litellm_params={}
)
assert (
url_with_slash == "https://api.x.ai/v1/responses"
), "Should handle trailing slash"

View file

@ -3,23 +3,34 @@ Unit tests for auto router management endpoints
"""
from collections.abc import Mapping, Sequence
from functools import partial
from pathlib import Path
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException, Request
from pydantic import ValidationError
import litellm
from litellm.proxy import proxy_server
from litellm.proxy._types import (
LitellmUserRoles,
ProxyErrorTypes,
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.management_endpoints import auto_router_endpoints
from litellm.proxy.management_endpoints.auto_router_endpoints import (
preview_auto_router_routing,
)
from litellm.router import Router
from litellm.router_strategy.complexity_router import ComplexityRouter
from litellm.router_strategy.complexity_router.jev_classifier import (
JevChoiceAnswer,
JevClassifierClient,
JevSystemOneResponse,
)
from litellm.types.management_endpoints.auto_router_endpoints import (
AutoRouterBenchmarksResponse,
AutoRouterRoutingTestRequest,
@ -422,8 +433,115 @@ async def test_a_key_over_its_budget_cannot_run_a_classifier_config(monkeypatch:
assert calls == []
@pytest.mark.parametrize(
"max_budget, spend, denied",
(
pytest.param(0.0, 0.0, True, id="zero-budget"),
pytest.param(1.0, 1.0, True, id="budget-reached"),
pytest.param(1.0, 2.0, True, id="budget-exceeded"),
pytest.param(1.0, 0.5, False, id="budget-remaining"),
pytest.param(None, 2.0, False, id="unlimited"),
),
)
@pytest.mark.asyncio
async def test_a_heuristic_config_does_not_need_a_budget(monkeypatch: pytest.MonkeyPatch):
async def test_jev_test_routing_enforces_key_budget_before_provider_invocation(
monkeypatch: pytest.MonkeyPatch, max_budget: float | None, spend: float, denied: bool
) -> None:
client: Final = AsyncMock(spec=JevClassifierClient)
client.evaluate.return_value = JevSystemOneResponse(
model="jev-test",
answers={
"tier": JevChoiceAnswer(type="choice", choice="SIMPLE", probabilities={"SIMPLE": 1.0}, confidence=1.0)
},
)
monkeypatch.setattr(proxy_server, "llm_router", _router())
monkeypatch.setattr(auto_router_endpoints, "ComplexityRouter", partial(ComplexityRouter, jev_client=client))
actor: Final = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-jev-budget-test",
user_id="admin",
models=["cheap-model", "typesafe/jev-test"],
max_budget=max_budget,
spend=spend,
)
request: Final = _request(
"what is 2+2",
classifier_type="jev",
jev_classifier_config={"model": "jev-test"},
)
if denied:
with pytest.raises(ProxyException) as exc_info:
await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST, data=request, user_api_key_dict=actor)
assert exc_info.value.type == ProxyErrorTypes.budget_exceeded
assert exc_info.value.code == "400"
assert exc_info.value.param is None
assert "Budget has been exceeded!" in exc_info.value.message
client.evaluate.assert_not_called()
return
response: Final = await preview_auto_router_routing(
http_request=ROUTING_HTTP_REQUEST, data=request, user_api_key_dict=actor
)
assert response.routed_model == "cheap-model"
assert response.routing_decision["cause"] == "jev_classifier"
assert response.routing_decision["classifier_model"] == "typesafe/jev-test"
client.evaluate.assert_awaited_once()
@pytest.mark.parametrize(
"max_budget, spend, denied",
((0.0, 0.0, True), (1.0, 2.0, True), (1.0, 0.5, False), (None, 2.0, False)),
)
@pytest.mark.asyncio
async def test_jev_test_routing_hard_blocks_exhausted_throttle_enabled_keys(
monkeypatch: pytest.MonkeyPatch, max_budget: float | None, spend: float, denied: bool
) -> None:
client: Final = AsyncMock(spec=JevClassifierClient)
client.evaluate.return_value = JevSystemOneResponse(
model="jev-test",
answers={
"tier": JevChoiceAnswer(type="choice", choice="SIMPLE", probabilities={"SIMPLE": 1.0}, confidence=1.0)
},
)
monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", 0.1)
monkeypatch.setattr(proxy_server, "llm_router", _router())
monkeypatch.setattr(auto_router_endpoints, "ComplexityRouter", partial(ComplexityRouter, jev_client=client))
actor: Final = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-jev-throttle-test",
user_id="admin",
models=["cheap-model", "typesafe/jev-test"],
max_budget=max_budget,
spend=spend,
rpm_limit=100,
metadata={"throttle_on_budget_exceeded": True},
)
request: Final = _request(
"what is 2+2",
classifier_type="jev",
jev_classifier_config={"model": "jev-test"},
)
if denied:
with pytest.raises(ProxyException) as exc_info:
await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST, data=request, user_api_key_dict=actor)
assert exc_info.value.type == ProxyErrorTypes.budget_exceeded
assert exc_info.value.code == "400"
client.evaluate.assert_not_called()
return
response: Final = await preview_auto_router_routing(
http_request=ROUTING_HTTP_REQUEST, data=request, user_api_key_dict=actor
)
assert response.routing_decision["cause"] == "jev_classifier"
client.evaluate.assert_awaited_once()
@pytest.mark.parametrize("max_budget, spend", ((0.0, 0.0), (1.0, 2.0)))
@pytest.mark.asyncio
async def test_a_heuristic_config_does_not_need_a_budget(
monkeypatch: pytest.MonkeyPatch, max_budget: float, spend: float
):
import litellm.proxy.proxy_server as proxy_server
monkeypatch.setattr(proxy_server, "llm_router", _router())
@ -435,8 +553,8 @@ async def test_a_heuristic_config_does_not_need_a_budget(monkeypatch: pytest.Mon
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-broke",
user_id="admin",
max_budget=1.0,
spend=2.0,
max_budget=max_budget,
spend=spend,
models=["cheap-model"],
),
)
@ -877,7 +995,6 @@ class TestAutoRouterBenchmarks:
# ---------------------------------------------------------------------------
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy.management_endpoints.auto_router_endpoints import (
get_shadow_eval_job,

View file

@ -941,6 +941,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"/v1/audio/transcriptions",
"/v1/audio/speech",
"/v1/ocr",
"/v1/videos",
"/vertex_ai/live",
"/v1/listen",
"/v1beta/interactions",
@ -1070,6 +1071,9 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
# Add any model IDs that should be exempt from the cost validation
# Example: "expensive-model-id",
"runwayml/seedance2", # 4K output is 150 credits/second = $1.50/second
"fal_ai/bytedance/seedance-2.0/text-to-video",
"fal_ai/bytedance/seedance-2.0/image-to-video",
"fal_ai/bytedance/seedance-2.0/reference-to-video",
]
is_valid, violations = validate_model_cost_values(actual_json, exceptions)

View file

@ -1,7 +1,6 @@
import asyncio
import json
import time
from pathlib import Path
import httpx
import pytest
@ -571,20 +570,3 @@ def test_config_manager_returns_wxo_provider():
)
assert config is not None
assert config.__class__.__name__ == "WatsonxOrchestrateA2AConfig"
def test_wxo_dashboard_auth_fields():
fields_path = (
Path(__file__).resolve().parents[5]
/ "litellm/proxy/public_endpoints/agent_create_fields.json"
)
agent_fields = json.loads(fields_path.read_text())
wxo_agent = next(
agent for agent in agent_fields if agent["agent_type"] == "watsonx_orchestrate"
)
fields_by_key = {field["key"]: field for field in wxo_agent["credential_fields"]}
assert fields_by_key["auth_mode"]["default_value"] == "cp4d"
# Username is CP4D-only; UI does not require it so ibm_cloud users are not blocked.
assert fields_by_key["username"]["required"] is False
assert "cp4d" in fields_by_key["username"]["tooltip"].lower()

View file

View file

@ -12,17 +12,28 @@ Line shape decides the parse, not the batch's declared endpoint, so an output
file mixing Responses-shaped and chat-shaped lines sums across both.
"""
from typing import Literal, get_args, get_type_hints
import pytest
import litellm
import litellm.batches.batch_utils as bu
from litellm.types.llms.openai import CreateBatchRequest
MODEL = "gpt-5.6"
@pytest.fixture
def local_model_cost_map(monkeypatch):
original_model_cost = litellm.model_cost
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
litellm.get_model_info.cache_clear()
try:
yield
finally:
litellm.model_cost = original_model_cost
litellm.get_model_info.cache_clear()
def _responses_line(input_tokens: int, output_tokens: int) -> dict:
return {
"response": {
@ -107,13 +118,3 @@ async def test_mixed_shape_batch_output_sums_across_both_line_shapes(local_model
assert result.cost == pytest.approx(
133 * model_info["input_cost_per_token_batches"] + 107 * model_info["output_cost_per_token_batches"]
)
def test_create_batch_endpoint_accepts_v1_responses():
"""A type-checked caller can pass endpoint="/v1/responses", which the runtime
already forwarded correctly."""
endpoint_annotation = get_type_hints(CreateBatchRequest)["endpoint"]
assert "/v1/responses" in get_args(endpoint_annotation)
for create_fn in (litellm.create_batch, litellm.acreate_batch):
assert "/v1/responses" in get_args(get_type_hints(create_fn)["endpoint"])

View file

View file

@ -1,11 +1,9 @@
import inspect
from collections.abc import Awaitable, Callable, Mapping
from typing import Final, cast # noqa: TID251 # narrows legacy callable signatures for inspect
from typing import Final, cast # noqa: TID251 # narrows legacy callable signatures
import pytest
import litellm
from litellm import main as python_chat
from litellm.chat_completions.dispatch import (
_ADISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
_DISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
@ -40,15 +38,6 @@ def acompletion_binding(native: NativeAcompletion | None) -> NativeBinding[Nativ
return binding
def test_public_signature_is_the_legacy_signature() -> None:
public_completion: Final = cast(Callable[..., object], litellm.completion)
legacy_completion: Final = cast(Callable[..., object], python_chat.completion)
public_acompletion: Final = cast(Callable[..., object], litellm.acompletion)
legacy_acompletion: Final = cast(Callable[..., object], python_chat.acompletion)
assert inspect.signature(public_completion) == inspect.signature(legacy_completion)
assert inspect.signature(public_acompletion) == inspect.signature(legacy_acompletion)
def test_python_route_forwards_original_call_shape() -> None:
metadata: Final = {"user_id": "u"}
args: Final[tuple[object, ...]] = ("gpt-4o", MESSAGES)

View file

View file

@ -0,0 +1,111 @@
from datetime import datetime
from unittest.mock import patch
import pytest
from litellm.completion_extras.litellm_responses_transformation.handler import (
ResponsesToCompletionBridgeHandler,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.types.utils import ModelResponse
MODEL = "openai.gpt-5.5"
REGION = "us-east-2"
def _bedrock_mantle_kwargs() -> dict:
messages = [{"role": "user", "content": "hi"}]
logging_obj = LiteLLMLogging(
litellm_call_id="test-call",
call_type="acompletion",
model=MODEL,
messages=messages,
function_id="fn-id",
stream=False,
start_time=datetime.now(),
)
return {
"model": MODEL,
"custom_llm_provider": "bedrock_mantle",
"messages": messages,
"optional_params": {},
"litellm_params": {
"aws_region_name": REGION,
"api_base": "https://bedrock-mantle.us-east-1.api.aws/v1",
"custom_llm_provider": "bedrock_mantle",
},
"headers": {},
"model_response": ModelResponse(),
"logging_obj": logging_obj,
}
def _openai_kwargs() -> dict:
messages = [{"role": "user", "content": "hi"}]
logging_obj = LiteLLMLogging(
litellm_call_id="test-call",
call_type="completion",
model="gpt-5.5",
messages=messages,
function_id="fn-id",
stream=False,
start_time=datetime.now(),
)
return {
"model": "gpt-5.5",
"custom_llm_provider": "openai",
"messages": messages,
"optional_params": {},
"litellm_params": {},
"headers": {},
"model_response": ModelResponse(),
"logging_obj": logging_obj,
}
def test_completion_forwards_custom_llm_provider_to_responses():
bridge = ResponsesToCompletionBridgeHandler()
cached = ModelResponse(id="chatcmpl-cached", model="gpt-5.5")
with patch("litellm.responses", return_value=cached) as fake_responses:
result = bridge.completion(**_openai_kwargs())
assert result is cached
assert fake_responses.call_args.kwargs["custom_llm_provider"] == "openai"
@pytest.mark.asyncio
async def test_acompletion_forwards_custom_llm_provider_to_aresponses():
bridge = ResponsesToCompletionBridgeHandler()
cached = ModelResponse(id="chatcmpl-cached", model="gpt-5.5")
async def _fake_aresponses(**kwargs):
_fake_aresponses.kwargs = kwargs
return cached
_fake_aresponses.kwargs = {}
with patch("litellm.aresponses", _fake_aresponses):
result = await bridge.acompletion(**_openai_kwargs())
assert result is cached
assert _fake_aresponses.kwargs["custom_llm_provider"] == "openai"
@pytest.mark.asyncio
async def test_acompletion_forwards_aws_region_name_to_aresponses():
bridge = ResponsesToCompletionBridgeHandler()
cached = ModelResponse(id="chatcmpl-cached", model=MODEL)
async def _fake_aresponses(**kwargs):
_fake_aresponses.kwargs = kwargs
return cached
_fake_aresponses.kwargs = {}
with patch("litellm.aresponses", _fake_aresponses):
result = await bridge.acompletion(**_bedrock_mantle_kwargs())
assert result is cached
assert _fake_aresponses.kwargs["aws_region_name"] == REGION
assert _fake_aresponses.kwargs["custom_llm_provider"] == "bedrock_mantle"

View file

View file

@ -1,4 +1,5 @@
import os
from collections.abc import Iterator
from typing import Final
import pytest
@ -6,7 +7,19 @@ from pytest_socket import enable_socket, socket_allow_hosts
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
import litellm # noqa: E402 # litellm reads LITELLM_LOCAL_MODEL_COST_MAP at import
import litellm.router as litellm_router_module # noqa: E402 # same import-time dependency
import litellm.utils as litellm_utils_module # noqa: E402 # same import-time dependency
LOOPBACK_HOSTS: Final = ["127.0.0.1", "::1"]
AMBIENT_AZURE_CREDENTIAL_ENV_VARS: Final = (
"AZURE_AD_TOKEN",
"AZURE_TENANT_ID",
"AZURE_CLIENT_ID",
"AZURE_CLIENT_SECRET",
"AZURE_USERNAME",
"AZURE_PASSWORD",
)
def _allow_loopback_only() -> None:
@ -21,5 +34,38 @@ def pytest_runtest_setup() -> None:
_allow_loopback_only()
@pytest.fixture(autouse=True)
def isolate_router_model_cost_state() -> Iterator[None]:
original_live_routers: Final = frozenset(litellm_router_module._live_routers)
original_runtime_registered_model_cost: Final = {
model_key: dict(model_value)
for model_key, model_value in litellm_utils_module._runtime_registered_model_cost.items()
}
yield
for router in tuple(litellm_router_module._live_routers):
litellm_router_module._live_routers.discard(router)
for router in original_live_routers:
litellm_router_module._live_routers.add(router)
litellm_utils_module._runtime_registered_model_cost.clear()
litellm_utils_module._runtime_registered_model_cost.update(original_runtime_registered_model_cost)
litellm_utils_module._invalidate_model_cost_lowercase_map()
litellm.get_model_info.cache_clear()
@pytest.fixture
def local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
litellm.get_model_info.cache_clear()
yield
litellm.get_model_info.cache_clear()
@pytest.fixture
def no_ambient_azure_credentials(monkeypatch: pytest.MonkeyPatch) -> None:
for name in AMBIENT_AZURE_CREDENTIAL_ENV_VARS:
monkeypatch.delenv(name, raising=False)
def pytest_sessionfinish() -> None:
enable_socket()

View file

View file

View file

View file

View file

@ -3,7 +3,6 @@ Test HeliconeLogger Gemini/Vertex AI support.
Fixes: https://github.com/BerriAI/litellm/issues/19093
"""
import pytest
def test_helicone_gemini_model_in_list():
@ -36,39 +35,6 @@ def test_helicone_gemini_models_recognized():
assert is_recognized, f"{model} should be recognized by helicone_model_list"
def test_helicone_vertex_ai_models_recognized():
"""
Test that Vertex AI models (GLM, DeepSeek, etc.) are recognized via custom_llm_provider.
"""
# Test models that don't contain "gemini" but are vertex_ai
test_models = [
"vertex_ai/zai-org/glm-4.7-maas",
"vertex_ai/deepseek-ai/deepseek-v3",
"vertex_ai/meta/llama-3.1-405b",
]
for model in test_models:
is_vertex_ai = model.startswith("vertex_ai/")
assert is_vertex_ai, f"{model} should be recognized as vertex_ai model"
def test_helicone_vertex_ai_via_custom_llm_provider():
"""
Test that vertex_ai models are recognized when custom_llm_provider is set.
"""
# Models without vertex_ai/ prefix but with custom_llm_provider="vertex_ai"
test_cases = [
("zai-org/glm-4.7-maas", "vertex_ai"),
("deepseek-ai/deepseek-v3", "vertex_ai"),
]
for model, custom_llm_provider in test_cases:
is_vertex_ai = custom_llm_provider == "vertex_ai" or model.startswith(
"vertex_ai/"
)
assert (
is_vertex_ai
), f"{model} with custom_llm_provider={custom_llm_provider} should be recognized as vertex_ai"
def test_helicone_vertex_gemini_gets_vertex_provider_url():
"""
Test that vertex_ai/gemini-* models route to aiplatform.googleapis.com,

View file

View file

Some files were not shown because too many files have changed in this diff Show more