mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
commit
1e1751ebff
5 changed files with 875 additions and 575 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):
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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