mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(openai/image_edits): Support 'mask' parameter for openai image edits
Closes https://github.com/BerriAI/litellm/issues/13528
This commit is contained in:
parent
1b2ec16eee
commit
b5c4ee60bb
3 changed files with 54 additions and 24 deletions
|
|
@ -80,24 +80,33 @@ class OpenAIImageEditConfig(BaseImageEditConfig):
|
|||
request_dict = cast(Dict, request)
|
||||
|
||||
#########################################################
|
||||
# Separate images as `files` and send other parameters as `data`
|
||||
# Separate images and masks as `files` and send other parameters as `data`
|
||||
#########################################################
|
||||
_images = request_dict.get("image") or []
|
||||
data_without_images = {k: v for k, v in request_dict.items() if k != "image"}
|
||||
_image = request_dict.get("image")
|
||||
_mask = request_dict.get("mask")
|
||||
data_without_files = {
|
||||
k: v for k, v in request_dict.items() if k not in ["image", "mask"]
|
||||
}
|
||||
files_list: List[Tuple[str, Any]] = []
|
||||
for _image in _images:
|
||||
|
||||
# Handle image parameter
|
||||
if _image is not None:
|
||||
image_content_type: str = ImageEditRequestUtils.get_image_content_type(
|
||||
_image
|
||||
)
|
||||
if isinstance(_image, BufferedReader):
|
||||
files_list.append(
|
||||
("image[]", (_image.name, _image, image_content_type))
|
||||
)
|
||||
files_list.append(("image", (_image.name, _image, image_content_type)))
|
||||
else:
|
||||
files_list.append(
|
||||
("image[]", ("image.png", _image, image_content_type))
|
||||
)
|
||||
return data_without_images, files_list
|
||||
files_list.append(("image", ("image.png", _image, image_content_type)))
|
||||
|
||||
# Handle mask parameter if provided
|
||||
if _mask is not None:
|
||||
mask_content_type: str = ImageEditRequestUtils.get_image_content_type(_mask)
|
||||
if isinstance(_mask, BufferedReader):
|
||||
files_list.append(("mask", (_mask.name, _mask, mask_content_type)))
|
||||
else:
|
||||
files_list.append(("mask", ("mask.png", _mask, mask_content_type)))
|
||||
return data_without_files, files_list
|
||||
|
||||
def transform_image_edit_response(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1,13 +1,10 @@
|
|||
model_list:
|
||||
- model_name: fake-openai-endpoint
|
||||
litellm_params:
|
||||
model: openai/fake
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
|
||||
litellm_settings:
|
||||
cache: true
|
||||
cache_params:
|
||||
type: redis
|
||||
ttl: 600
|
||||
supported_call_types: ["acompletion", "completion"]
|
||||
- model_name: fake-openai-endpoint
|
||||
litellm_params:
|
||||
model: openai/fake
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
- model_name: gpt-image-1
|
||||
litellm_params:
|
||||
model: openai/gpt-image-1
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
|
@ -94,4 +94,28 @@ class BaseImageGenTest(ABC):
|
|||
if "Your task failed as a result of our safety system." in str(e):
|
||||
pass
|
||||
else:
|
||||
pytest.fail(f"An exception occurred - {str(e)}")
|
||||
pytest.fail(f"An exception occurred - {str(e)}")
|
||||
|
||||
|
||||
def test_openai_gpt_image_1():
|
||||
from litellm import image_edit
|
||||
from PIL import Image
|
||||
import io
|
||||
|
||||
# Create a simple mask image with alpha channel
|
||||
# Create a 512x512 black image with alpha channel
|
||||
try:
|
||||
response = image_edit(
|
||||
model="openai/gpt-image-1",
|
||||
image=open("test_image_edit.png", "rb"),
|
||||
mask=open("test_image_edit.png", "rb"),
|
||||
prompt="Add a red hat to the person in the image",
|
||||
n=1,
|
||||
size="1024x1024",
|
||||
)
|
||||
print("response: ", response)
|
||||
except Exception as e:
|
||||
if "mask image missing alpha channel" in str(e):
|
||||
pass
|
||||
else:
|
||||
raise e
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue