mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
Merge remote-tracking branch 'origin/main' into litellm_fal_gpt_image_25_flux_dev_edits
This commit is contained in:
commit
342bde7a8d
514 changed files with 2244 additions and 2870 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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("#"):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = {},
|
||||
|
|
|
|||
|
|
@ -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:"):
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
3
litellm/llms/fal_ai/videos/__init__.py
Normal file
3
litellm/llms/fal_ai/videos/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from litellm.llms.fal_ai.videos.transformation import FalAIVideoConfig
|
||||
|
||||
__all__ = ("FalAIVideoConfig",)
|
||||
516
litellm/llms/fal_ai/videos/transformation.py
Normal file
516
litellm/llms/fal_ai/videos/transformation.py
Normal 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)
|
||||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
],
|
||||
|
|
|
|||
69
tests/integration/providers/test_fal_ai_video_wire.py
Normal file
69
tests/integration/providers/test_fal_ai_video_wire.py
Normal 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"),
|
||||
]
|
||||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
|
|
@ -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"]
|
||||
)
|
||||
|
|
@ -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"])
|
||||
|
|
@ -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"])
|
||||
|
|
@ -1 +0,0 @@
|
|||
"""OpenAI Evals API tests"""
|
||||
|
|
@ -1 +0,0 @@
|
|||
# Test module for OpenAI-like embedding handler
|
||||
|
|
@ -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
|
||||
|
|
@ -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']}'"
|
||||
)
|
||||
|
|
@ -1 +0,0 @@
|
|||
# XAI Responses API tests
|
||||
|
|
@ -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"
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
0
tests/unit/a2a_protocol/providers/__init__.py
Normal file
0
tests/unit/a2a_protocol/providers/__init__.py
Normal 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()
|
||||
0
tests/unit/anthropic_interface/__init__.py
Normal file
0
tests/unit/anthropic_interface/__init__.py
Normal file
0
tests/unit/anthropic_interface/exceptions/__init__.py
Normal file
0
tests/unit/anthropic_interface/exceptions/__init__.py
Normal file
0
tests/unit/batches/__init__.py
Normal file
0
tests/unit/batches/__init__.py
Normal 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"])
|
||||
0
tests/unit/chat_completions/__init__.py
Normal file
0
tests/unit/chat_completions/__init__.py
Normal 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)
|
||||
0
tests/unit/completion_extras/__init__.py
Normal file
0
tests/unit/completion_extras/__init__.py
Normal 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"
|
||||
0
tests/unit/compression/__init__.py
Normal file
0
tests/unit/compression/__init__.py
Normal 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()
|
||||
|
|
|
|||
0
tests/unit/endpoints/__init__.py
Normal file
0
tests/unit/endpoints/__init__.py
Normal file
0
tests/unit/endpoints/speech/__init__.py
Normal file
0
tests/unit/endpoints/speech/__init__.py
Normal file
0
tests/unit/enterprise/__init__.py
Normal file
0
tests/unit/enterprise/__init__.py
Normal file
0
tests/unit/enterprise/enterprise_callbacks/__init__.py
Normal file
0
tests/unit/enterprise/enterprise_callbacks/__init__.py
Normal file
0
tests/unit/integrations/__init__.py
Normal file
0
tests/unit/integrations/__init__.py
Normal file
0
tests/unit/integrations/gcs_bucket/__init__.py
Normal file
0
tests/unit/integrations/gcs_bucket/__init__.py
Normal file
0
tests/unit/integrations/gcs_pubsub/__init__.py
Normal file
0
tests/unit/integrations/gcs_pubsub/__init__.py
Normal file
0
tests/unit/integrations/helicone/__init__.py
Normal file
0
tests/unit/integrations/helicone/__init__.py
Normal 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,
|
||||
0
tests/unit/integrations/levo/__init__.py
Normal file
0
tests/unit/integrations/levo/__init__.py
Normal file
0
tests/unit/integrations/litellm_agent/__init__.py
Normal file
0
tests/unit/integrations/litellm_agent/__init__.py
Normal file
0
tests/unit/integrations/mavvrik_focus/__init__.py
Normal file
0
tests/unit/integrations/mavvrik_focus/__init__.py
Normal file
0
tests/unit/integrations/opik/__init__.py
Normal file
0
tests/unit/integrations/opik/__init__.py
Normal file
0
tests/unit/integrations/pointfive/__init__.py
Normal file
0
tests/unit/integrations/pointfive/__init__.py
Normal file
0
tests/unit/litellm_core_utils/__init__.py
Normal file
0
tests/unit/litellm_core_utils/__init__.py
Normal file
0
tests/unit/litellm_core_utils/audio_utils/__init__.py
Normal file
0
tests/unit/litellm_core_utils/audio_utils/__init__.py
Normal file
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue