mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
chore(typing): clear basedpyright Any errors in rubrik and runwayml videos
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
b0fac57fe4
commit
a10b1fa697
3 changed files with 298 additions and 197 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue