fix: satisfy rust image edit lint gate

This commit is contained in:
Ishaan Jaff 2026-06-25 12:41:35 -07:00
parent 834e479dc8
commit 620f461739
No known key found for this signature in database
2 changed files with 125 additions and 130 deletions

View file

@ -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

View file

@ -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: