Merge remote-tracking branches 'origin/litellm_daily_any_cleanup_08_11_2026_3Mn0k_a' and 'origin/litellm_daily_any_cleanup_08_11_2026_3Mn0k_b' into litellm_daily_any_cleanup_08_11_2026_3Mn0k

This commit is contained in:
Devin AI 2026-08-11 13:50:17 +00:00
commit 1e1751ebff
5 changed files with 875 additions and 575 deletions

View file

@ -6,12 +6,14 @@ import random
import time
import uuid
from collections import Counter
from collections.abc import Mapping, Sequence
from collections.abc import Awaitable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, runtime_checkable
import httpx
from pydantic import BaseModel, TypeAdapter, ValidationError
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
@ -48,11 +50,86 @@ _WEBHOOK_PATH_PROMPT_MODERATION: Final = "/v1/before_prompt/openai/v1"
_WEBHOOK_PATH_LOGGING_BATCH: Final = "/v1/litellm/batch"
_MAX_QUEUE_SIZE: Final = 10_000
_DROP_WARNING_INTERVAL_SECONDS: Final = 60.0
_EMPTY_MAPPING: Final[Mapping[str, Any]] = MappingProxyType({})
_EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
_MAPPING_ADAPTER: Final = TypeAdapter(dict[str, object])
_SEQUENCE_ADAPTER: Final = TypeAdapter(tuple[object, ...])
def _as_optional_mapping(value: object) -> Mapping[str, object] | None:
"""``value`` as a string-keyed mapping, or ``None`` when it is not one."""
try:
return _MAPPING_ADAPTER.validate_python(value)
except ValidationError:
return None
def _as_mapping(value: object) -> Mapping[str, object]:
"""``value`` as a string-keyed mapping, empty when it is not one."""
return _as_optional_mapping(value) or _EMPTY_MAPPING
def _as_sequence(value: object) -> tuple[object, ...]:
"""``value`` as a sequence, empty when it is not one."""
try:
return _SEQUENCE_ADAPTER.validate_python(value)
except ValidationError:
return ()
@runtime_checkable
class _ToolCallLike(Protocol):
"""A provider tool-call object exposing ``id`` and ``function`` attributes."""
id: str | None
function: Function
@runtime_checkable
class _HasType(Protocol):
"""Any object carrying a ``type`` discriminator."""
type: str | None
@runtime_checkable
class _HasModel(Protocol):
"""Any object carrying the model name of an LLM call."""
model: str | None
class _ToolCallFunctionPayload(BaseModel):
name: str | None = ""
arguments: str | dict[str, object] | None = ""
class _ToolCallPayload(BaseModel):
id: str | None = ""
type: str | None = "function"
function: _ToolCallFunctionPayload | None = None
class _ModerationToolCall(BaseModel):
id: str | None = None
class _ModerationMessage(BaseModel):
content: str | None = None
tool_calls: tuple[_ModerationToolCall, ...] | None = None
class _ModerationChoice(BaseModel):
message: _ModerationMessage | None = None
class _ModerationResponse(BaseModel):
"""The OpenAI chat-completion shape both Rubrik moderation webhooks return."""
choices: tuple[_ModerationChoice, ...] | None = None
class _MalformedToolBlockingResponseError(Exception):
"""Raised when the response moderation service returns a structurally invalid
"""Raised when a moderation service returns a structurally invalid
response (e.g. empty ``choices``).
Distinct from transient network/HTTP errors so callers can surface a
@ -61,6 +138,14 @@ class _MalformedToolBlockingResponseError(Exception):
"""
def _parse_moderation_response(service_response: Mapping[str, object], service_name: str) -> _ModerationResponse:
"""Validate a webhook response into the chat-completion shape this module reads."""
try:
return _ModerationResponse.model_validate(service_response)
except ValidationError as e:
raise _MalformedToolBlockingResponseError(f"{service_name} returned an unreadable response: {e}") from e
@dataclass
class BlockedResponseResult:
"""Returned by _extract_response_block when the response was blocked
@ -143,7 +228,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
else {"Content-Type": "application/json"}
)
self._periodic_flush_task: asyncio.Task[Any] | None = self._start_periodic_flush_task()
self._periodic_flush_task: asyncio.Task[None] | None = self._start_periodic_flush_task()
@classmethod
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
@ -191,7 +276,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
params={"timeout": httpx.Timeout(5.0, connect=2.0)},
)
def _start_periodic_flush_task(self) -> asyncio.Task[Any] | None:
def _start_periodic_flush_task(self) -> asyncio.Task[None] | None:
"""Start the periodic flush task only when an event loop is already running."""
try:
loop: Final = asyncio.get_running_loop()
@ -212,7 +297,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
Closing them here would close the shared connection pool for every
other logger instance; let LiteLLM manage their lifecycle instead.
"""
task: Final = getattr(self, "_periodic_flush_task", None)
task: Final[asyncio.Task[None] | None] = getattr(self, "_periodic_flush_task", None)
if task is not None:
task.cancel()
@ -253,7 +338,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
@staticmethod
async def _guarded(
coro: Any,
coro: Awaitable[GenericGuardrailAPIInputs],
inputs: GenericGuardrailAPIInputs,
label: str,
) -> GenericGuardrailAPIInputs:
@ -400,34 +485,36 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
request_data["_rubrik_logging_obj"] = logging_obj
@staticmethod
def _normalize_tool_calls(tool_calls: Any) -> tuple[ChatCompletionMessageToolCall, ...]:
def _normalize_tool_calls(tool_calls: Iterable[object]) -> tuple[ChatCompletionMessageToolCall, ...]:
"""Convert tool_calls from inputs to ChatCompletionMessageToolCall objects."""
return tuple(RubrikLogger._normalize_tool_call(tc) for tc in tool_calls)
@staticmethod
def _normalize_tool_call(tc: Any) -> ChatCompletionMessageToolCall:
def _normalize_tool_call(tc: object) -> ChatCompletionMessageToolCall:
if isinstance(tc, ChatCompletionMessageToolCall):
return tc
if isinstance(tc, dict):
func: Final = tc.get("function") or _EMPTY_MAPPING
as_mapping: Final = _as_optional_mapping(tc)
if as_mapping is not None:
payload: Final = _ToolCallPayload.model_validate(as_mapping)
func: Final = payload.function or _ToolCallFunctionPayload()
return ChatCompletionMessageToolCall(
id=tc.get("id", ""),
type=tc.get("type", "function"),
id=payload.id,
type=payload.type,
function=Function(
name=func.get("name", ""),
arguments=func.get("arguments", ""),
name=func.name,
arguments=func.arguments,
),
)
if hasattr(tc, "id") and hasattr(tc, "function"):
if isinstance(tc, _ToolCallLike):
return ChatCompletionMessageToolCall(
id=tc.id or "",
type=getattr(tc, "type", None) or "function",
type=(tc.type if isinstance(tc, _HasType) else None) or "function",
function=tc.function,
)
raise TypeError(f"Cannot normalize tool_call of type {type(tc).__name__}: {tc!r}")
@staticmethod
def _join_texts(texts: Any) -> str:
def _join_texts(texts: Sequence[str] | None) -> str:
"""Join response text segments into the single content string the
webhook evaluates. Empty when there is no assistant text."""
if not texts:
@ -439,14 +526,14 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
tool_calls: Sequence[ChatCompletionMessageToolCall],
content: str,
request_id: str | None,
) -> Mapping[str, Any]:
) -> Mapping[str, object]:
"""Build an OpenAI ChatCompletion-format dict (assistant text + tool
calls) for the after_completion webhook.
``content`` is sent so the webhook can moderate the response text;
``None`` when the assistant produced no text (tool-call-only response).
"""
message: Final[dict[str, Any]] = {
message: Final[dict[str, object]] = {
"role": "assistant",
"content": content or None,
}
@ -467,7 +554,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
}
@staticmethod
def _flatten_messages_for_moderation(messages: Any) -> tuple[Mapping[str, Any], ...]:
def _flatten_messages_for_moderation(messages: Sequence[object] | None) -> tuple[Mapping[str, object], ...]:
"""Collapse each message's content to a plain string for the webhook.
litellm normalizes Anthropic ``/v1/messages`` requests to OpenAI shape,
@ -483,31 +570,30 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
"role": message.get("role"),
"content": "\n".join(p for p in RubrikLogger._moderation_text_parts(message) if p),
}
for message in messages or ()
if isinstance(message, dict)
for message in (_as_optional_mapping(entry) for entry in messages or ())
if message is not None
)
@staticmethod
def _moderation_text_parts(message: Mapping[str, Any]) -> tuple[str, ...]:
def _moderation_text_parts(message: Mapping[str, object]) -> tuple[str, ...]:
"""Every attacker-controlled text segment of a message: its content plus
the arguments of any tool call or deprecated function call."""
fc: Final = message.get("function_call")
fc: Final = _as_mapping(message.get("function_call"))
return (
# Base text content (flattens Anthropic content-part arrays)
convert_content_list_to_str(message), # pyright: ignore[reportArgumentType] # dict[str,Any] is AllMessageValues at runtime
convert_content_list_to_str(message), # pyright: ignore[reportArgumentType] # a str-keyed mapping is AllMessageValues at runtime
*(
str((tc.get("function") or _EMPTY_MAPPING).get("arguments") or "")
for tc in message.get("tool_calls") or ()
if isinstance(tc, dict)
str(_as_mapping(_as_mapping(tc).get("function")).get("arguments") or "")
for tc in _as_sequence(message.get("tool_calls"))
),
str((fc.get("arguments") if isinstance(fc, dict) else None) or ""),
str(fc.get("arguments") or ""),
)
@staticmethod
def _build_prompt_moderation_payload(
inputs: GenericGuardrailAPIInputs,
request_data: Mapping[str, Any],
) -> Mapping[str, Any]:
request_data: Mapping[str, object],
) -> Mapping[str, object]:
"""Build the bare OpenAI request the before_prompt webhook consumes.
Unlike the after_completion envelope, this endpoint takes a raw OpenAI
@ -516,7 +602,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
``/v1/messages`` requests too. Optional fields are sent only when
present so the payload stays clean.
"""
payload: Final[dict[str, Any]] = {
payload: Final[dict[str, object]] = {
"model": inputs.get("model") or request_data.get("model") or "",
"messages": RubrikLogger._flatten_messages_for_moderation(inputs.get("structured_messages")),
}
@ -539,9 +625,9 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
@staticmethod
def _extract_request_data(
call_details: Mapping[str, Any],
request_data: Mapping[str, Any] | None,
) -> Mapping[str, Any]:
call_details: Mapping[str, object],
request_data: Mapping[str, object] | None,
) -> Mapping[str, object]:
"""Extract original request data from model_call_details for the
response moderation service envelope.
@ -553,7 +639,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
return _EMPTY_MAPPING
call_details = call_details or _EMPTY_MAPPING
request_data = request_data or _EMPTY_MAPPING
optional_params: Final = call_details.get("optional_params") or _EMPTY_MAPPING
optional_params: Final = _as_mapping(call_details.get("optional_params"))
# Use ``in`` rather than truthy ``or`` so an explicit empty list
# (caller declared the agent has NO tools) is forwarded as-is.
@ -576,27 +662,32 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
}
@staticmethod
def _sanitize_proxy_server_request(proxy_server_request: Any) -> Any:
def _sanitize_proxy_server_request(proxy_server_request: object) -> object:
"""Allowlist only routing fields (``url``, ``method``) when forwarding
``proxy_server_request`` to an external webhook, dropping inbound
``headers`` (Authorization, Cookie, x-api-key, ...) and the raw
request ``body`` so proxy credentials are not exfiltrated."""
if not isinstance(proxy_server_request, dict):
parsed: Final = _as_optional_mapping(proxy_server_request)
if parsed is None:
return proxy_server_request
return {key: proxy_server_request[key] for key in ("url", "method") if key in proxy_server_request}
return {key: parsed[key] for key in ("url", "method") if key in parsed}
@staticmethod
def _resolve_model(request_data: Mapping[str, Any], call_details: Mapping[str, Any]) -> str:
def _resolve_model(request_data: Mapping[str, object], call_details: Mapping[str, object]) -> str:
"""Get the model name for the ModifyResponseException."""
response: Final = request_data.get("response")
if response and hasattr(response, "model"):
if response and isinstance(response, _HasModel):
return response.model or "unknown"
return call_details.get("model", "unknown")
model: Final = call_details.get("model", "unknown")
return model if isinstance(model, str) else str(model)
# -- Logging hooks ---------------------------------------------------------
@staticmethod
def _correlation_id(call_details: Mapping[str, Any], request_data: Mapping[str, Any] | None = None) -> str | None:
def _correlation_id(
call_details: Mapping[str, object],
request_data: Mapping[str, object] | None = None,
) -> str | None:
"""The id that joins a blocked request's two S3 logs by filename: the
moderation (``_blocking``) log and the failure (response) log.
@ -607,10 +698,13 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
block fires before the response/logging object is populated, so the
two logs correlate for every provider (OpenAI and Anthropic alike).
"""
return call_details.get("litellm_call_id") or (request_data or _EMPTY_MAPPING).get("litellm_call_id")
correlated: Final = call_details.get("litellm_call_id") or (request_data or _EMPTY_MAPPING).get(
"litellm_call_id"
)
return correlated if isinstance(correlated, str) else None
@classmethod
def _apply_correlation_id(cls, payload: dict[str, Any], source: Mapping[str, Any]) -> None:
def _apply_correlation_id(cls, payload: dict[str, object], source: Mapping[str, object]) -> None:
"""Pin ``payload["id"]`` to ``litellm_call_id`` in place so this log
shares its S3 filename id with the moderation (``_blocking``) and
failure logs for the same request -- for every provider.
@ -630,7 +724,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
payload["id"] = correlated
@staticmethod
def _prepend_system_prompt(payload: dict[str, Any], source: Mapping[str, Any]) -> None:
def _prepend_system_prompt(payload: dict[str, object], source: Mapping[str, object]) -> None:
"""Prepend ``source["system"]`` onto ``payload["messages"]``.
Builds a NEW messages list rather than mutating ``payload["messages"]``
@ -658,7 +752,11 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
exc_info=True,
)
async def _prepare_log_payload(self, kwargs: Mapping[str, Any], event_type: str) -> StandardLoggingPayload | None:
async def _prepare_log_payload(
self,
kwargs: Mapping[str, object],
event_type: str,
) -> StandardLoggingPayload | None:
"""Shared logic for success logging (sampled)."""
if random.random() > self.sampling_rate:
verbose_logger.debug("Skipping Rubrik %s logging (sampling_rate=%s)", event_type, self.sampling_rate)
@ -697,7 +795,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
self._dropped_since_warning = 0
self._last_drop_warning_time = now
async def _enqueue_log_event(self, kwargs: Mapping[str, Any], event_type: str):
async def _enqueue_log_event(self, kwargs: Mapping[str, object], event_type: str):
try:
payload: Final = await self._prepare_log_payload(kwargs, event_type)
if payload is None:
@ -857,7 +955,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
- time fields: ``call_details["start_time"]`` reused for all three;
end/completion times are meaningless for a prompt block.
"""
call_details: Final = logging_obj.model_call_details
call_details: Final = _as_mapping(logging_obj.model_call_details)
exception_text: Final = f"{type(exception).__name__}: {exception.message}"
base: Final = call_details.get("standard_logging_object")
@ -906,13 +1004,13 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
@classmethod
def _build_fallback_payload(
cls,
call_details: Mapping[str, Any],
call_details: Mapping[str, object],
user_api_key_dict: "UserAPIKeyAuth",
) -> dict[str, Any]:
) -> dict[str, object]:
# Convert datetime to a Unix float so json.dumps can serialize it.
# httpx's json= parameter uses stdlib json.dumps with no custom encoder.
_raw_start: Final = call_details.get("start_time")
_start: Final = _raw_start.timestamp() if _raw_start is not None else None
_start: Final = _raw_start.timestamp() if isinstance(_raw_start, datetime) else None
return {
"id": call_details.get("litellm_call_id"),
"model": call_details.get("model") or "",
@ -996,7 +1094,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
# -- Webhook services ------------------------------------------------------
async def _post_json(self, endpoint: str, payload: Mapping[str, Any], service_name: str) -> Mapping[str, Any]:
async def _post_json(self, endpoint: str, payload: Mapping[str, object], service_name: str) -> Mapping[str, object]:
"""POST ``payload`` to a Rubrik webhook and return its dict response.
Raises:
@ -1010,20 +1108,21 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
headers=self._headers,
)
http_response.raise_for_status()
result: Final = http_response.json()
if not isinstance(result, dict):
result: Final[object] = http_response.json() # any-ok: httpx exposes the decoded JSON body as Any
parsed: Final = _as_optional_mapping(result)
if parsed is None:
raise TypeError(
f"{service_name} returned non-dict JSON "
f"({type(result).__name__}); expected OpenAI chat completion "
"shape or empty object."
)
return result
return parsed
async def _post_to_response_moderation_endpoint(
self,
response_data: Mapping[str, Any],
request_data: Mapping[str, Any],
) -> Mapping[str, Any]:
response_data: Mapping[str, object],
request_data: Mapping[str, object],
) -> Mapping[str, object]:
"""Post the ``{request, response}`` envelope to the after_completion
webhook and return its (possibly rewritten) response.
@ -1039,7 +1138,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
"Response moderation service",
)
async def _post_to_prompt_moderation_endpoint(self, payload: Mapping[str, Any]) -> Mapping[str, Any]:
async def _post_to_prompt_moderation_endpoint(self, payload: Mapping[str, object]) -> Mapping[str, object]:
"""Post a bare OpenAI request to the before_prompt webhook.
Returns ``{}`` (passthrough) or a synthetic chat.completion (block).
@ -1047,23 +1146,22 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
return await self._post_json(self.prompt_moderation_endpoint, payload, "Prompt moderation service")
@staticmethod
def _extract_prompt_refusal(service_response: Mapping[str, Any]) -> str | None:
def _extract_prompt_refusal(service_response: Mapping[str, object]) -> str | None:
"""Return the refusal text when the prompt was blocked, else None.
The before_prompt webhook returns ``{}`` (passthrough) or a synthetic
chat.completion whose ``choices[0].message.content`` is the refusal
explanation.
"""
choices: Final = service_response.get("choices")
choices: Final = _parse_moderation_response(service_response, "Prompt moderation service").choices or ()
if not choices:
return None
message: Final = choices[0].get("message") or _EMPTY_MAPPING
content: Final = message.get("content")
return content or "Request blocked by policy."
message: Final = choices[0].message
return (message.content if message else None) or "Request blocked by policy."
@staticmethod
def _extract_response_block(
service_response: Mapping[str, Any],
service_response: Mapping[str, object],
all_tool_calls: Sequence[ChatCompletionMessageToolCall],
sent_content: str,
) -> BlockedResponseResult | None:
@ -1086,19 +1184,19 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
Expects service_response in OpenAI chat completion format:
{"choices": [{"message": {"tool_calls": [...], "content": "..."}}]}
"""
choices: Final = service_response.get("choices") or ()
choices: Final = _parse_moderation_response(service_response, "Response moderation service").choices or ()
if not choices:
raise _MalformedToolBlockingResponseError("Response moderation service returned empty response")
message: Final = choices[0].get("message") or _EMPTY_MAPPING
returned_tool_calls: Final = message.get("tool_calls") or ()
returned_content: Final = message.get("content") or ""
message: Final = choices[0].message
returned_tool_calls: Final = (message.tool_calls if message else None) or ()
returned_content: Final = (message.content if message else None) or ""
# Use Counter so duplicate IDs are handled correctly: if the model
# emits two calls with the same ID (one allowed, one prohibited) and
# the service returns only the allowed one, a set-based check would
# miss the block. Counter preserves multiplicity.
returned_id_counts: Final[Counter[str]] = Counter(tc["id"] for tc in returned_tool_calls if tc.get("id"))
returned_id_counts: Final[Counter[str]] = Counter(tc.id for tc in returned_tool_calls if tc.id)
required_id_counts: Final[Counter[str]] = Counter(tc.id for tc in all_tool_calls if tc.id)
# Cardinality check catches ID-less tool calls (not counted in
# required_id_counts because tc.id is falsy); Counter check catches

View file

@ -2,7 +2,8 @@ from datetime import datetime
from typing import TYPE_CHECKING, Any, Final
import httpx
from httpx._types import RequestFiles
from httpx._types import FileContent, RequestFiles
from pydantic import BaseModel, ConfigDict, Field
import litellm
from litellm.constants import RUNWAYML_DEFAULT_API_VERSION
@ -31,6 +32,29 @@ else:
LiteLLMLoggingObj = Any
class RunwayMLTaskResponse(BaseModel):
"""RunwayML task object returned by ``/v1/image_to_video`` and ``/v1/tasks/{id}``."""
model_config = ConfigDict(populate_by_name=True)
id: str = ""
status: str = "pending"
created_at: str | None = Field(default=None, alias="createdAt")
completed_at: str | None = Field(default=None, alias="completedAt")
output: tuple[str, ...] | str | None = None
progress: int | None = None
failure: str | None = None
failure_code: str | None = Field(default=None, alias="failureCode")
class RunwayMLTaskRequest(BaseModel):
"""The RunwayML-shaped request fields this module reads back when building a ``VideoObject``."""
model: str | None = None
ratio: str | None = None
duration: int | None = None
class RunwayMLVideoConfig(BaseVideoConfig):
"""
Configuration class for RunwayML video generation.
@ -78,7 +102,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
- size -> ratio (convert "WIDTHxHEIGHT" to "WIDTH:HEIGHT")
- seconds -> duration (convert to integer)
"""
mapped_params: Final[dict[str, Any]] = {}
mapped_params: Final[dict[str, object]] = {}
# Handle input_reference parameter - map to promptImage
if "input_reference" in video_create_optional_params:
@ -180,7 +204,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
}
"""
# Build the request data
request_data: Final[dict[str, Any]] = {
request_data: Final[dict[str, object]] = {
"model": model,
"promptText": prompt,
}
@ -189,7 +213,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
request_data.update(video_create_optional_request_params)
# RunwayML uses JSON body, no files multipart
files_list: Final[list[tuple[str, Any]]] = []
files_list: Final[RequestFiles] = []
# Append the specific endpoint for video generation
full_api_base: Final = f"{api_base}/image_to_video"
@ -216,60 +240,58 @@ class RunwayMLVideoConfig(BaseVideoConfig):
We map this to OpenAI VideoObject format.
"""
response_data: Final = raw_response.json()
task: Final = RunwayMLTaskResponse.model_validate(raw_response.json())
request_params: Final = RunwayMLTaskRequest.model_validate(request_data or {})
seconds: Final = str(request_params.duration) if request_params.duration is not None else None
# Map RunwayML task response to VideoObject format
video_data: Final[dict[str, Any]] = {
"id": response_data.get("id", ""),
"object": "video",
"status": self._map_runway_status(response_data.get("status", "pending")),
"created_at": self._parse_runway_timestamp(response_data.get("createdAt")),
return VideoObject(
id=self._encoded_video_id(task.id, custom_llm_provider, model),
object="video",
status=self._map_runway_status(task.status),
created_at=self._parse_runway_timestamp(task.created_at),
completed_at=self._completed_at(task),
error=self._build_error(task),
model=request_params.model,
size=self._ratio_to_size(request_params.ratio),
seconds=seconds,
usage=self._usage_from_seconds(seconds),
)
@staticmethod
def _encoded_video_id(video_id: str, custom_llm_provider: str | None, model: str | None) -> str:
if not (custom_llm_provider and video_id):
return video_id
return encode_video_id_with_provider(video_id, custom_llm_provider, model)
@staticmethod
def _ratio_to_size(ratio: str | None) -> str | None:
if not ratio or ":" not in ratio:
return None
return ratio.replace(":", "x")
@staticmethod
def _usage_from_seconds(seconds: str | None) -> dict[str, float]:
if not seconds:
return {}
try:
return {"duration_seconds": float(seconds)}
except ValueError:
return {}
def _completed_at(self, task: RunwayMLTaskResponse) -> int | None:
if "completed_at" not in task.model_fields_set:
return None
return self._parse_runway_timestamp(task.completed_at)
@staticmethod
def _build_error(task: RunwayMLTaskResponse) -> dict[str, str] | None:
if not {"failure", "failure_code"} & task.model_fields_set:
return None
return {
"code": task.failure_code if task.failure_code is not None else "unknown",
"message": task.failure if task.failure is not None else "Video generation failed",
}
# Add optional fields if present
if "output" in response_data and response_data["output"]:
# RunwayML returns output as array of URLs when task succeeds
video_data["output_url"] = (
response_data["output"][0] if isinstance(response_data["output"], list) else response_data["output"]
)
if "completedAt" in response_data:
video_data["completed_at"] = self._parse_runway_timestamp(response_data.get("completedAt"))
if "failureCode" in response_data or "failure" in response_data:
video_data["error"] = {
"code": response_data.get("failureCode", "unknown"),
"message": response_data.get("failure", "Video generation failed"),
}
# Add model and size info if available from request
if request_data:
if "model" in request_data:
video_data["model"] = request_data["model"]
if "ratio" in request_data:
# Convert ratio back to size format
ratio: Final = request_data["ratio"]
if isinstance(ratio, str) and ":" in ratio:
video_data["size"] = ratio.replace(":", "x")
if "duration" in request_data:
video_data["seconds"] = str(request_data["duration"])
video_obj: Final = VideoObject(**video_data)
if custom_llm_provider and video_obj.id:
video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, model)
# Add usage data for cost tracking
usage_data: Final = {}
if video_obj and hasattr(video_obj, "seconds") and video_obj.seconds:
try:
usage_data["duration_seconds"] = float(video_obj.seconds)
except (ValueError, TypeError):
pass
video_obj.usage = usage_data
return video_obj
def _map_runway_status(self, runway_status: str) -> str:
"""
Map RunwayML status to OpenAI status format.
@ -326,33 +348,32 @@ class RunwayMLVideoConfig(BaseVideoConfig):
# Get task status to retrieve video URL
url: Final = f"{api_base}/tasks/{encoded_video_id}"
params: Final[dict[str, Any]] = {}
params: Final[dict[str, str]] = {}
return url, params
def _extract_video_url_from_response(self, response_data: dict[str, Any]) -> str:
def _extract_video_url_from_response(self, task: RunwayMLTaskResponse) -> str:
"""
Helper method to extract video URL from RunwayML response.
Shared between sync and async transforms.
"""
# Extract video URL from the output field
video_url = None
if "output" in response_data and response_data["output"]:
output: Final = response_data["output"]
video_url = output[0] if isinstance(output, list) else output
video_url: Final = self._first_output_url(task.output)
if video_url:
return video_url
if not video_url:
# Check if the video generation failed or is still processing
status: Final = response_data.get("status", "UNKNOWN")
if status in ["PENDING", "RUNNING", "THROTTLED"]:
raise ValueError(f"Video is still processing (status: {status}). Please wait and try again.")
elif status == "FAILED":
failure_reason: Final = response_data.get("failure", "Unknown error")
raise ValueError(f"Video generation failed: {failure_reason}")
else:
raise ValueError("Video URL not found in response. Video may not be ready yet.")
status: Final = task.status
if status in ("PENDING", "RUNNING", "THROTTLED"):
raise ValueError(f"Video is still processing (status: {status}). Please wait and try again.")
if status == "FAILED":
failure_reason: Final = task.failure if task.failure is not None else "Unknown error"
raise ValueError(f"Video generation failed: {failure_reason}")
raise ValueError("Video URL not found in response. Video may not be ready yet.")
return video_url
@staticmethod
def _first_output_url(output: tuple[str, ...] | str | None) -> str | None:
if isinstance(output, tuple):
return output[0] if output else None
return output
def transform_video_content_response(
self,
@ -373,8 +394,8 @@ class RunwayMLVideoConfig(BaseVideoConfig):
"output":["https://dnznrvs05pmza.cloudfront.net/.../video.mp4?_jwt=..."]
}
"""
response_data: Final = raw_response.json()
video_url: Final = self._extract_video_url_from_response(response_data)
task: Final = RunwayMLTaskResponse.model_validate(raw_response.json())
video_url: Final = self._extract_video_url_from_response(task)
# Download the video from the CloudFront URL synchronously
httpx_client: Final[HTTPHandler] = _get_httpx_client()
@ -402,8 +423,8 @@ class RunwayMLVideoConfig(BaseVideoConfig):
"output":["https://dnznrvs05pmza.cloudfront.net/.../video.mp4?_jwt=..."]
}
"""
response_data: Final = raw_response.json()
video_url: Final = self._extract_video_url_from_response(response_data)
task: Final = RunwayMLTaskResponse.model_validate(raw_response.json())
video_url: Final = self._extract_video_url_from_response(task)
# Download the video from the CloudFront URL asynchronously
async_httpx_client: Final[AsyncHTTPHandler] = get_async_httpx_client(
@ -421,7 +442,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
extra_body: dict[str, Any] | None = None,
extra_body: dict[str, object] | None = None,
) -> tuple[str, dict]:
"""
Transform the video remix request for RunwayML API.
@ -448,7 +469,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
after: str | None = None,
limit: int | None = None,
order: str | None = None,
extra_query: dict[str, Any] | None = None,
extra_query: dict[str, object] | None = None,
) -> tuple[str, dict]:
"""
Transform the video list request for RunwayML API.
@ -484,7 +505,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
# Construct the URL for task cancellation
url: Final = f"{api_base}/tasks/{encoded_video_id}/cancel"
data: Final[dict[str, Any]] = {}
data: Final[dict[str, str]] = {}
return url, data
@ -494,17 +515,15 @@ class RunwayMLVideoConfig(BaseVideoConfig):
logging_obj: LiteLLMLoggingObj,
) -> VideoObject:
"""Transform the RunwayML video delete/cancel response."""
response_data: Final = raw_response.json()
task: Final = RunwayMLTaskResponse.model_validate(raw_response.json())
video_obj: Final = VideoObject(
id=response_data.get("id", ""),
return VideoObject(
id=task.id,
object="video",
status="cancelled",
created_at=self._parse_runway_timestamp(response_data.get("createdAt")),
created_at=self._parse_runway_timestamp(task.created_at),
)
return video_obj
def transform_video_status_retrieve_request(
self,
video_id: str,
@ -524,7 +543,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
url: Final = f"{api_base}/tasks/{encoded_video_id}"
# Empty dict for GET request (no body)
data: Final[dict[str, Any]] = {}
data: Final[dict[str, str]] = {}
return url, data
@ -537,42 +556,19 @@ class RunwayMLVideoConfig(BaseVideoConfig):
"""
Transform the RunwayML video status retrieve response.
"""
response_data: Final = raw_response.json()
task: Final = RunwayMLTaskResponse.model_validate(raw_response.json())
# Map RunwayML task response to VideoObject format
video_data: Final[dict[str, Any]] = {
"id": response_data.get("id", ""),
"object": "video",
"status": self._map_runway_status(response_data.get("status", "pending")),
"created_at": self._parse_runway_timestamp(response_data.get("createdAt")),
}
return VideoObject(
id=self._encoded_video_id(task.id, custom_llm_provider, None),
object="video",
status=self._map_runway_status(task.status),
created_at=self._parse_runway_timestamp(task.created_at),
completed_at=self._completed_at(task),
progress=task.progress,
error=self._build_error(task),
)
# Add optional fields if present
if "output" in response_data and response_data["output"]:
video_data["output_url"] = (
response_data["output"][0] if isinstance(response_data["output"], list) else response_data["output"]
)
if "completedAt" in response_data:
video_data["completed_at"] = self._parse_runway_timestamp(response_data.get("completedAt"))
if "progress" in response_data:
video_data["progress"] = response_data["progress"]
if "failureCode" in response_data or "failure" in response_data:
video_data["error"] = {
"code": response_data.get("failureCode", "unknown"),
"message": response_data.get("failure", "Video generation failed"),
}
video_obj: Final = VideoObject(**video_data)
if custom_llm_provider and video_obj.id:
video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, None)
return video_obj
def transform_video_create_character_request(self, name, video, api_base, litellm_params, headers):
def transform_video_create_character_request(self, name, video: FileContent, api_base, litellm_params, headers):
raise NotImplementedError("video create character is not supported for RunwayML")
def transform_video_create_character_response(self, raw_response, logging_obj):

View file

@ -5,12 +5,13 @@ import json
import os
import re
import time
from collections.abc import Callable, Iterable, Iterator
from typing import Any, Final
from collections.abc import Callable, Iterable, Iterator, Mapping
from typing import Final, Protocol
import httpx
from httpx import Headers, Response
from openai.types.file_deleted import FileDeleted
from pydantic import TypeAdapter
import litellm
from litellm._uuid import uuid
@ -46,13 +47,14 @@ from litellm.types.files import StreamingMediaUploadConfig
from litellm.types.llms.openai import (
AllMessageValues,
CreateFileRequest,
FileContentRequest,
FileTypes,
HttpxBinaryResponseContent,
OpenAICreateFileRequestOptionalParams,
OpenAIFileObject,
OpenAIFilesPurpose,
PathLike,
)
from litellm.types.llms.vertex_ai import GcsBucketResponse
from litellm.types.utils import LlmProviders, ModelResponse
from ..common_utils import VertexAIError
@ -61,6 +63,53 @@ from ..vertex_llm_base import VertexBase
_GCP_LABEL_VALUE_MAX_LEN: Final = 63
_CUSTOM_ID_RAW_LABEL_PREFIX: Final = "b32_"
_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object])
_PURPOSE_ADAPTER: Final = TypeAdapter[OpenAIFilesPurpose](OpenAIFilesPurpose)
_JSON_DECODER: Final = json.JSONDecoder()
class _JsonDecoder(Protocol):
def decode(self, s: str, /) -> object: ...
class _JsonResponse(Protocol):
def json(self) -> object: ...
def _json_object(value: object) -> dict[str, object]:
return _JSON_OBJECT_ADAPTER.validate_python(value)
def _json_object_or_empty(value: object) -> dict[str, object]:
return _json_object(value) if value else {}
def _parse_json_object(raw: str, decoder: _JsonDecoder = _JSON_DECODER) -> dict[str, object]:
return _json_object(decoder.decode(raw))
def _response_json_object(response: _JsonResponse) -> dict[str, object]:
return _json_object(response.json())
def _str_field(payload: Mapping[str, object], key: str, default: str = "") -> str:
value: Final = payload.get(key, default)
return value if isinstance(value, str) else default
def _int_field(payload: Mapping[str, object], key: str) -> int:
value: Final = payload.get(key, 0)
return int(value) if isinstance(value, (int, float, str)) else 0
def _purpose_field(payload: Mapping[str, object], default: OpenAIFilesPurpose = "batch") -> OpenAIFilesPurpose:
return _PURPOSE_ADAPTER.validate_python(payload.get("purpose", default))
def _gcs_file_id(payload: Mapping[str, object]) -> str:
gcs_id: Final = _str_field(payload, "id")
return "/".join(gcs_id.split("/")[:-1]) if gcs_id else ""
def _sanitize_gcp_label_value(value: str) -> str:
"""
@ -106,7 +155,7 @@ def _decode_gcp_label_value_chunks(values: list[str]) -> str | None:
return None
def _set_litellm_batch_custom_id_labels(labels: dict[str, str], custom_id: Any) -> None:
def _set_litellm_batch_custom_id_labels(labels: dict[str, str], custom_id: object) -> None:
"""
Store OpenAI batch custom_id for Vertex batch correlation.
@ -122,7 +171,7 @@ def _set_litellm_batch_custom_id_labels(labels: dict[str, str], custom_id: Any)
labels[f"litellm_custom_id_raw_{index}"] = raw_label_chunk
def _get_litellm_batch_custom_id_from_labels(labels: dict[str, Any]) -> str:
def _get_litellm_batch_custom_id_from_labels(labels: Mapping[str, object]) -> str:
"""Prefer encoded custom_id when present (see _set_litellm_batch_custom_id_labels)."""
raw: Final = labels.get("litellm_custom_id_raw")
if raw:
@ -140,10 +189,15 @@ def _get_litellm_batch_custom_id_from_labels(labels: dict[str, Any]) -> str:
return str(labels.get("litellm_custom_id", "unknown"))
# any-ok: batch bodies must reach _transform_request_body unmodified, so no lossy revalidation here
def _batch_messages(openai_request_body) -> list[AllMessageValues]:
return openai_request_body.get("messages", [])
def _openai_batch_jsonl_entry_to_vertex_wrapped_request(
openai_entry: dict[str, Any],
map_openai_to_vertex_params: Callable[[dict[str, Any]], dict[str, Any]],
) -> dict[str, Any]:
openai_entry: Mapping[str, object],
map_openai_to_vertex_params: Callable[[dict[str, object]], dict[str, object]],
) -> dict[str, object]:
"""
Transforms a single OpenAI JSONL batch entry into its Vertex wrapped request.
@ -151,10 +205,10 @@ def _openai_batch_jsonl_entry_to_vertex_wrapped_request(
Example Vertex jsonl
{"request":{"contents": [{"role": "user", "parts": [{"text": "What is the relation between the following video and image samples?"}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/animals.mp4", "mimeType": "video/mp4"}}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/image/cricket.jpeg", "mimeType": "image/jpeg"}}]}]}}
"""
openai_request_body: Final = openai_entry.get("body") or {}
openai_request_body: Final = _json_object_or_empty(openai_entry.get("body"))
vertex_request_body: Final = _transform_request_body(
messages=openai_request_body.get("messages", []),
model=openai_request_body.get("model", ""),
messages=_batch_messages(openai_request_body),
model=_str_field(openai_request_body, "model"),
optional_params=map_openai_to_vertex_params(openai_request_body),
custom_llm_provider="vertex_ai",
litellm_params={},
@ -186,9 +240,7 @@ def _iter_openai_jsonl_lines(openai_file_content: FileTypes) -> Iterator[str]:
``str.splitlines()`` + ``line.strip()`` for ``\\n`` / ``\\r\\n`` delimited
JSONL.
"""
content: Any = openai_file_content
if isinstance(content, tuple):
content = content[1]
content: Final[object] = openai_file_content[1] if isinstance(openai_file_content, tuple) else openai_file_content
if isinstance(content, (bytes, bytearray)):
# Scan for newlines in place so a large in-memory payload is not copied
@ -220,14 +272,13 @@ def _iter_openai_jsonl_lines(openai_file_content: FileTypes) -> Iterator[str]:
# object name, then the body stream), so it must rewind to 0. A
# non-seekable handle would silently resume mid-stream and drop the
# already-consumed first row, so reject it loudly instead.
seek: Final = getattr(content, "seek", None)
if seek is None:
if not hasattr(content, "seek"):
raise ValueError(
"Batch upload file handle must be seekable; got a non-seekable "
"stream. Pass bytes, a path, or a seekable handle."
)
try:
seek(0)
content.seek(0)
except (OSError, ValueError) as e:
raise ValueError(
"Batch upload file handle must be seekable so it can be re-read "
@ -241,9 +292,9 @@ def _iter_openai_jsonl_lines(openai_file_content: FileTypes) -> Iterator[str]:
def _iter_openai_jsonl_entries(
openai_file_content: FileTypes,
) -> Iterator[dict[str, Any]]:
) -> Iterator[dict[str, object]]:
for line in _iter_openai_jsonl_lines(openai_file_content):
yield json.loads(line)
yield _parse_json_object(line)
class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream):
@ -257,7 +308,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, object]], dict[str, object]],
) -> None:
self._openai_file_content = openai_file_content
self._map_openai_to_vertex_params = map_openai_to_vertex_params
@ -308,17 +359,18 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
def _get_gcs_object_name_from_batch_jsonl(
self,
openai_jsonl_content: list[dict[str, Any]],
openai_jsonl_content: list[dict[str, object]],
) -> str:
"""
Gets a unique GCS object name for the VertexAI batch prediction job
named as: litellm-vertex-{model}-{uuid}
"""
_model = openai_jsonl_content[0].get("body", {}).get("model", "")
if "publishers/google/models" not in _model:
_model = f"publishers/google/models/{_model}"
safe_model_path: Final = sanitize_cloud_object_path(_model, fallback="model")
model_name: Final = _str_field(_json_object_or_empty(openai_jsonl_content[0].get("body")), "model")
qualified_model: Final = (
model_name if "publishers/google/models" in model_name else f"publishers/google/models/{model_name}"
)
safe_model_path: Final = sanitize_cloud_object_path(qualified_model, fallback="model")
object_name: Final = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
return object_name
@ -343,9 +395,11 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
fallback_filename="file",
)
def _get_configured_bucket_name(self, litellm_params: dict) -> str:
def _get_configured_bucket_name(self, litellm_params: Mapping[str, object]) -> str:
bucket_name: Final = (
litellm_params.get("gcs_bucket_name") or litellm_params.get("bucket_name") or os.getenv("GCS_BUCKET_NAME")
_str_field(litellm_params, "gcs_bucket_name")
or _str_field(litellm_params, "bucket_name")
or os.getenv("GCS_BUCKET_NAME")
)
if not bucket_name:
raise ValueError("GCS bucket_name is required")
@ -396,8 +450,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
def _map_openai_to_vertex_params(
self,
openai_request_body: dict[str, Any],
) -> dict[str, Any]:
openai_request_body: dict[str, object],
) -> dict[str, object]:
"""
wrapper to call VertexGeminiConfig.map_openai_params
"""
@ -406,9 +460,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
)
config: Final = VertexGeminiConfig()
_model: Final = openai_request_body.get("model", "")
vertex_params: Final = config.map_openai_params(
model=_model,
vertex_params: Final[dict[str, object]] = config.map_openai_params(
model=_str_field(openai_request_body, "model"),
non_default_params=openai_request_body,
optional_params={},
drop_params=False,
@ -463,10 +516,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
"""
Transform VertexAI File upload response into OpenAI-style FileObject
"""
response_json: Final = raw_response.json()
try:
response_object: Final = GcsBucketResponse(**response_json)
response_object: Final = _response_json_object(raw_response)
except Exception as e:
raise VertexAIError(
status_code=raw_response.status_code,
@ -474,19 +525,15 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
headers=raw_response.headers,
)
gcs_id = response_object.get("id", "")
# Remove the last numeric ID from the path
gcs_id = "/".join(gcs_id.split("/")[:-1]) if gcs_id else ""
return OpenAIFileObject(
purpose=response_object.get("purpose", "batch"),
id=f"gs://{gcs_id}",
filename=response_object.get("name", ""),
purpose=_purpose_field(response_object),
id=f"gs://{_gcs_file_id(response_object)}",
filename=_str_field(response_object, "name"),
created_at=_convert_vertex_datetime_to_openai_datetime(
vertex_datetime=response_object.get("timeCreated", "")
vertex_datetime=_str_field(response_object, "timeCreated")
),
status="uploaded",
bytes=int(response_object.get("size", 0)),
bytes=_int_field(response_object, "size"),
object="file",
)
@ -523,18 +570,16 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
) -> OpenAIFileObject:
response_json: Final = raw_response.json()
gcs_id = response_json.get("id", "")
gcs_id = "/".join(gcs_id.split("/")[:-1]) if gcs_id else ""
response_json: Final = _response_json_object(raw_response)
return OpenAIFileObject(
id=f"gs://{gcs_id}",
bytes=int(response_json.get("size", 0)),
id=f"gs://{_gcs_file_id(response_json)}",
bytes=_int_field(response_json, "size"),
created_at=_convert_vertex_datetime_to_openai_datetime(
vertex_datetime=response_json.get("timeCreated", "")
vertex_datetime=_str_field(response_json, "timeCreated")
),
filename=response_json.get("name", ""),
filename=_str_field(response_json, "name"),
object="file",
purpose=response_json.get("metadata", {}).get("purpose", "batch"),
purpose=_purpose_field(_json_object_or_empty(response_json.get("metadata"))),
status="processed",
status_details=None,
)
@ -584,7 +629,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
def transform_file_content_request(
self,
file_content_request,
file_content_request: FileContentRequest,
optional_params: dict,
litellm_params: dict,
) -> tuple[str, dict]:
@ -682,14 +727,15 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
# discriminating fields. Anything else (e.g. a binary file whose
# first line is not valid UTF-8/JSON) raises and falls through to the
# passthrough below, leaving the content untouched.
first_row: Final = json.loads(first_line)
first_row: Final = _parse_json_object(first_line)
first_row_response: Final = _json_object_or_empty(first_row.get("response"))
is_vertex_batch_output: Final = (
"request" in first_row
and "response" in first_row
and "processed_time" in first_row
and (
"candidates" in first_row.get("response", {})
or "promptFeedback" in first_row.get("response", {})
"candidates" in first_row_response
or "promptFeedback" in first_row_response
or bool(first_row.get("status"))
)
)
@ -723,7 +769,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
for line in itertools.chain([first_line], lines):
try:
openai_output = self._transform_single_vertex_batch_output_to_openai(
vertex_output=json.loads(line),
vertex_output=_parse_json_object(line),
vertex_gemini_config=vertex_gemini_config,
logging_obj=batch_transform_logging_obj,
mock_httpx_response=mock_httpx_response,
@ -742,22 +788,22 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
def _transform_single_vertex_batch_output_to_openai(
self,
vertex_output: dict[str, Any],
vertex_output: Mapping[str, object],
vertex_gemini_config: VertexGeminiConfig,
logging_obj: Logging,
mock_httpx_response: httpx.Response,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Transform a single Vertex AI batch output line to OpenAI format.
Uses the existing VertexGeminiConfig transformation for the response.
"""
# Extract custom_id from request labels (prefer raw for OpenAI round-trip)
request_data: Final = vertex_output.get("request", {})
labels: Final = request_data.get("labels", {}) or {}
request_data: Final = _json_object_or_empty(vertex_output.get("request"))
labels: Final = _json_object_or_empty(request_data.get("labels"))
custom_id: Final = _get_litellm_batch_custom_id_from_labels(labels)
# Check if there's an error
status: Final = vertex_output.get("status", "")
status: Final = _str_field(vertex_output, "status")
has_error: Final = bool(status)
if has_error:
@ -772,12 +818,11 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
}
# Transform successful response using existing transformation
vertex_response: Final = vertex_output.get("response", {})
vertex_response: Final = _json_object_or_empty(vertex_output.get("response"))
# Extract model from response
model = vertex_response.get("modelVersion", "gemini-1.5-flash-001")
if "@" in model:
model = model.split("@")[0]
model_version: Final = _str_field(vertex_response, "modelVersion", "gemini-1.5-flash-001")
model: Final = model_version.split("@")[0] if "@" in model_version else model_version
try:
# Use existing VertexGeminiConfig transformation
@ -792,7 +837,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
)
# Convert ModelResponse to dict
response_dict: Final = transformed_response.model_dump()
response_dict: Final = _json_object(transformed_response.model_dump())
# Return in OpenAI batch format
return {
@ -800,7 +845,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
"custom_id": custom_id,
"response": {
"status_code": 200,
"request_id": response_dict.get("id", ""),
"request_id": _str_field(response_dict, "id"),
"body": response_dict,
},
"error": None,

File diff suppressed because it is too large Load diff

View file

@ -7,7 +7,10 @@ from unittest.mock import Mock
import httpx
import pytest
from litellm.llms.runwayml.videos.transformation import RunwayMLVideoConfig
from litellm.llms.runwayml.videos.transformation import (
RunwayMLTaskResponse,
RunwayMLVideoConfig,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.videos.main import VideoObject
@ -126,13 +129,17 @@ class TestRunwayMLVideoTransformation:
"status": "SUCCEEDED",
"output": ["https://dnznrvs05pmza.cloudfront.net/video.mp4"],
}
video_url = self.config._extract_video_url_from_response(response_data)
video_url = self.config._extract_video_url_from_response(
RunwayMLTaskResponse.model_validate(response_data)
)
assert video_url == "https://dnznrvs05pmza.cloudfront.net/video.mp4"
# Test error handling when video is still processing
processing_response = {"id": "test-id", "status": "RUNNING", "output": None}
with pytest.raises(ValueError, match="still processing"):
self.config._extract_video_url_from_response(processing_response)
self.config._extract_video_url_from_response(
RunwayMLTaskResponse.model_validate(processing_response)
)
def test_transform_video_status_encodes_video_id_path_segment(self):
"""Test task IDs are encoded before being appended to Runway URLs."""