diff --git a/litellm/integrations/rubrik.py b/litellm/integrations/rubrik.py index 97e831f5822..fe43019d8ff 100644 --- a/litellm/integrations/rubrik.py +++ b/litellm/integrations/rubrik.py @@ -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 diff --git a/litellm/llms/runwayml/videos/transformation.py b/litellm/llms/runwayml/videos/transformation.py index 2e0ae30a192..a055fec84db 100644 --- a/litellm/llms/runwayml/videos/transformation.py +++ b/litellm/llms/runwayml/videos/transformation.py @@ -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): diff --git a/tests/test_litellm/llms/runwayml/videos/test_runway_video_transformation.py b/tests/test_litellm/llms/runwayml/videos/test_runway_video_transformation.py index 24879ce83f9..e02481bddc2 100644 --- a/tests/test_litellm/llms/runwayml/videos/test_runway_video_transformation.py +++ b/tests/test_litellm/llms/runwayml/videos/test_runway_video_transformation.py @@ -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."""