mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
fix: satisfy rust image edit lint gate
This commit is contained in:
parent
834e479dc8
commit
620f461739
2 changed files with 125 additions and 130 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue