diff --git a/litellm/images/main.py b/litellm/images/main.py index f777201d001..4da8edd25b4 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -1,8 +1,11 @@ +from __future__ import annotations + import asyncio import base64 import contextvars import importlib import os +from dataclasses import dataclass from functools import partial from io import BufferedReader, BytesIO from typing import ( @@ -81,10 +84,25 @@ from litellm.ocr.rust_bridge import ( ) # Cache for ImageEditRequestUtils to avoid repeated __getattr__ calls -_ImageEditRequestUtils_cache: Optional["ImageEditRequestUtils"] = None +_ImageEditRequestUtils_cache: ImageEditRequestUtils | None = None -def _get_ImageEditRequestUtils() -> "ImageEditRequestUtils": +@dataclass(frozen=True) +class _RustImageEditCall: + model: str + images: list[FileTypes] + prompt: str | None + api_key: str | None + api_base: str | None + custom_llm_provider: str + extra_headers: dict[str, Any] | None + optional_params: dict[str, object] + mask: dict[str, object] | None + timeout: float | httpx.Timeout | None + logging_obj: LiteLLMLoggingObj + + +def _get_ImageEditRequestUtils() -> ImageEditRequestUtils: """Get ImageEditRequestUtils, loading it lazily if needed.""" global _ImageEditRequestUtils_cache if _ImageEditRequestUtils_cache is None: @@ -96,8 +114,8 @@ def _get_ImageEditRequestUtils() -> "ImageEditRequestUtils": def _timeout_to_seconds( - timeout: Optional[Union[float, httpx.Timeout]], -) -> Optional[float]: + timeout: float | httpx.Timeout | None, +) -> float | None: if timeout is None: return None if isinstance(timeout, httpx.Timeout): @@ -148,9 +166,9 @@ def _filename_for_file(file_value: Any, default: str) -> str: return default -def _rust_image_file_part(file_value: Any, default_filename: str) -> Dict[str, object]: +def _rust_image_file_part(file_value: Any, default_filename: str) -> dict[str, object]: filename = default_filename - content_type: Optional[str] = None + content_type: str | None = None raw_file = file_value if isinstance(file_value, tuple): @@ -172,7 +190,7 @@ def _rust_image_file_part(file_value: Any, default_filename: str) -> Dict[str, o } -def _rust_image_file_parts(images: List[FileTypes]) -> List[Dict[str, object]]: +def _rust_image_file_parts(images: list[FileTypes]) -> list[dict[str, object]]: return [ _rust_image_file_part(image, f"image-{index}.png") for index, image in enumerate(images) @@ -182,9 +200,9 @@ def _rust_image_file_parts(images: List[FileTypes]) -> List[Dict[str, object]]: def _prepare_rust_image_edit_params( image_edit_optional_params: ImageEditOptionalRequestParams, - non_default_params: Dict[str, Any], -) -> tuple[Dict[str, object], Optional[Dict[str, object]]]: - optional_params: Dict[str, object] = dict(image_edit_optional_params) + non_default_params: dict[str, Any], +) -> tuple[dict[str, object], dict[str, object] | None]: + optional_params: dict[str, object] = dict(image_edit_optional_params) optional_params.update(non_default_params) mask = optional_params.pop("mask", None) mask_part = _rust_image_file_part(mask, "mask.png") if mask is not None else None @@ -193,91 +211,69 @@ def _prepare_rust_image_edit_params( def _run_rust_image_edit( rust_image_edit: RustImageEdit, - *, - model: str, - images: List[FileTypes], - prompt: Optional[str], - api_key: Optional[str], - api_base: Optional[str], - custom_llm_provider: str, - extra_headers: Optional[Dict[str, Any]], - optional_params: Dict[str, object], - mask: Optional[Dict[str, object]], - timeout: Optional[Union[float, httpx.Timeout]], - logging_obj: LiteLLMLoggingObj, + request: _RustImageEditCall, ) -> ImageResponse: - image_parts = _rust_image_file_parts(images) - logging_obj.pre_call( - input=prompt, + image_parts = _rust_image_file_parts(request.images) + request.logging_obj.pre_call( + input=request.prompt, api_key="", additional_args={ "complete_input_dict": { - "model": model, - "prompt": prompt, - **optional_params, + "model": request.model, + "prompt": request.prompt, + **request.optional_params, }, - "api_base": api_base, - "headers": extra_headers or {}, + "api_base": request.api_base, + "headers": request.extra_headers or {}, }, ) return ImageResponse( **rust_image_edit( - model=model, + model=request.model, images=image_parts, - prompt=prompt, - mask=mask, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=cast(Optional[dict[str, object]], extra_headers), - optional_params=optional_params, - timeout_seconds=_timeout_to_seconds(timeout), + prompt=request.prompt, + mask=request.mask, + api_key=request.api_key, + api_base=request.api_base, + custom_llm_provider=request.custom_llm_provider, + extra_headers=cast(dict[str, object] | None, request.extra_headers), + optional_params=request.optional_params, + timeout_seconds=_timeout_to_seconds(request.timeout), ) ) async def _run_rust_aimage_edit( rust_aimage_edit: RustAimageEdit, - *, - model: str, - images: List[FileTypes], - prompt: Optional[str], - api_key: Optional[str], - api_base: Optional[str], - custom_llm_provider: str, - extra_headers: Optional[Dict[str, Any]], - optional_params: Dict[str, object], - mask: Optional[Dict[str, object]], - timeout: Optional[Union[float, httpx.Timeout]], - logging_obj: LiteLLMLoggingObj, + request: _RustImageEditCall, ) -> ImageResponse: - image_parts = _rust_image_file_parts(images) - logging_obj.pre_call( - input=prompt, + image_parts = _rust_image_file_parts(request.images) + request.logging_obj.pre_call( + input=request.prompt, api_key="", additional_args={ "complete_input_dict": { - "model": model, - "prompt": prompt, - **optional_params, + "model": request.model, + "prompt": request.prompt, + **request.optional_params, }, - "api_base": api_base, - "headers": extra_headers or {}, + "api_base": request.api_base, + "headers": request.extra_headers or {}, }, ) return ImageResponse( **( await rust_aimage_edit( - model=model, + model=request.model, images=image_parts, - prompt=prompt, - mask=mask, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=cast(Optional[dict[str, object]], extra_headers), - optional_params=optional_params, - timeout_seconds=_timeout_to_seconds(timeout), + prompt=request.prompt, + mask=request.mask, + api_key=request.api_key, + api_base=request.api_base, + custom_llm_provider=request.custom_llm_provider, + extra_headers=cast(dict[str, object] | None, request.extra_headers), + optional_params=request.optional_params, + timeout_seconds=_timeout_to_seconds(request.timeout), ) ) ) @@ -989,8 +985,8 @@ def image_edit( model_info = kwargs.get("model_info", None) metadata = kwargs.get("metadata", {}) _is_async = kwargs.pop("async_call", False) is True - api_key_param: Optional[str] = kwargs.get("api_key") - api_base_param: Optional[str] = kwargs.get("api_base") + api_key_param: str | None = kwargs.get("api_key") + api_base_param: str | None = kwargs.get("api_base") # add images / or return a single image images = ( @@ -998,7 +994,7 @@ def image_edit( ) headers_from_kwargs = kwargs.get("headers") - merged_extra_headers: Dict[str, Any] = {} + merged_extra_headers: dict[str, Any] = {} if isinstance(headers_from_kwargs, dict): merged_extra_headers.update(headers_from_kwargs) if isinstance(extra_headers, dict): @@ -1101,17 +1097,19 @@ def image_edit( ) return _run_rust_aimage_edit( rust_aimage_edit=rust_aimage_edit, - model=model, - images=images, - prompt=prompt, - api_key=resolved_api_key, - api_base=resolved_api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=rust_optional_params, - mask=rust_mask, - timeout=timeout, - logging_obj=litellm_logging_obj, + request=_RustImageEditCall( + model=model, + images=images, + prompt=prompt, + api_key=resolved_api_key, + api_base=resolved_api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=rust_optional_params, + mask=rust_mask, + timeout=timeout, + logging_obj=litellm_logging_obj, + ), ) else: rust_image_edit = load_rust_image_edit() @@ -1130,17 +1128,19 @@ def image_edit( ) return _run_rust_image_edit( rust_image_edit=rust_image_edit, - model=model, - images=images, - prompt=prompt, - api_key=resolved_api_key, - api_base=resolved_api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=rust_optional_params, - mask=rust_mask, - timeout=timeout, - logging_obj=litellm_logging_obj, + request=_RustImageEditCall( + model=model, + images=images, + prompt=prompt, + api_key=resolved_api_key, + api_base=resolved_api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=rust_optional_params, + mask=rust_mask, + timeout=timeout, + logging_obj=litellm_logging_obj, + ), ) # get provider config diff --git a/litellm/ocr/rust_bridge.py b/litellm/ocr/rust_bridge.py index 553804fa6fc..b7f59402aaf 100644 --- a/litellm/ocr/rust_bridge.py +++ b/litellm/ocr/rust_bridge.py @@ -11,7 +11,8 @@ can import it statically without forming an import cycle. from __future__ import annotations -from typing import Awaitable, Final, Protocol, cast +from collections.abc import Awaitable, Callable +from typing import Final, Protocol, cast class RustOcr(Protocol): @@ -48,42 +49,36 @@ class RustAocr(Protocol): raise NotImplementedError -class RustImageEdit(Protocol): - """Signature of the compiled ``litellm_python_bridge.image_edit`` entrypoint.""" - - def __call__( - self, - model: str, - images: list[dict[str, object]], - prompt: str | None, - mask: dict[str, object] | None, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, - ) -> dict[str, object]: - raise NotImplementedError - - -class RustAimageEdit(Protocol): - """Signature of the compiled ``litellm_python_bridge.aimage_edit`` entrypoint.""" - - def __call__( - self, - model: str, - images: list[dict[str, object]], - prompt: str | None, - mask: dict[str, object] | None, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, - ) -> Awaitable[dict[str, object]]: - raise NotImplementedError +RustImageEdit = Callable[ + [ + str, + list[dict[str, object]], + str | None, + dict[str, object] | None, + str | None, + str | None, + str, + dict[str, object] | None, + dict[str, object], + float | None, + ], + dict[str, object], +] +RustAimageEdit = Callable[ + [ + str, + list[dict[str, object]], + str | None, + dict[str, object] | None, + str | None, + str | None, + str, + dict[str, object] | None, + dict[str, object], + float | None, + ], + Awaitable[dict[str, object]], +] class _Unset: