feat: support multi-image image generation
This commit is contained in:
@@ -82,6 +82,7 @@ class ImageGenerationResult:
|
||||
class MultimodalImage:
|
||||
data: bytes
|
||||
mime_type: str = "image/png"
|
||||
filename: str = "image.png"
|
||||
|
||||
|
||||
class Provider(Protocol):
|
||||
@@ -108,6 +109,7 @@ class Provider(Protocol):
|
||||
image: bytes | None = None,
|
||||
image_mime_type: str = "image/png",
|
||||
image_filename: str = "image.png",
|
||||
images: Sequence[MultimodalImage] | None = None,
|
||||
resolution: str = "1K",
|
||||
aspect_ratio: str = "1:1",
|
||||
parameters: Mapping[str, Any] | None = None,
|
||||
|
||||
@@ -138,6 +138,7 @@ class ChatCompletionsProvider(BaseHttpProvider):
|
||||
image: bytes | None = None,
|
||||
image_mime_type: str = "image/png",
|
||||
image_filename: str = "image.png",
|
||||
images: Sequence[MultimodalImage] | None = None,
|
||||
resolution: str = "1K",
|
||||
aspect_ratio: str = "1:1",
|
||||
parameters: Mapping[str, Any] | None = None,
|
||||
@@ -149,6 +150,7 @@ class ChatCompletionsProvider(BaseHttpProvider):
|
||||
prompt,
|
||||
image=image,
|
||||
image_mime_type=image_mime_type,
|
||||
images=images,
|
||||
parameters=parameters,
|
||||
)
|
||||
response = self.session.post(
|
||||
@@ -243,6 +245,7 @@ class GeminiProvider(ChatCompletionsProvider):
|
||||
image: bytes | None = None,
|
||||
image_mime_type: str = "image/png",
|
||||
image_filename: str = "image.png",
|
||||
images: Sequence[MultimodalImage] | None = None,
|
||||
resolution: str = "1K",
|
||||
aspect_ratio: str = "1:1",
|
||||
parameters: Mapping[str, Any] | None = None,
|
||||
@@ -254,6 +257,7 @@ class GeminiProvider(ChatCompletionsProvider):
|
||||
prompt,
|
||||
image=image,
|
||||
image_mime_type=image_mime_type,
|
||||
images=images,
|
||||
response_modalities=["TEXT", "IMAGE"],
|
||||
parameters=parameters,
|
||||
)
|
||||
@@ -324,6 +328,7 @@ class ImagesGenerationProvider(BaseHttpProvider):
|
||||
image: bytes | None = None,
|
||||
image_mime_type: str = "image/png",
|
||||
image_filename: str = "image.png",
|
||||
images: Sequence[MultimodalImage] | None = None,
|
||||
resolution: str = "1K",
|
||||
aspect_ratio: str = "1:1",
|
||||
parameters: Mapping[str, Any] | None = None,
|
||||
@@ -337,8 +342,17 @@ class ImagesGenerationProvider(BaseHttpProvider):
|
||||
"resolution": resolution,
|
||||
"n": 1,
|
||||
}
|
||||
if image is not None:
|
||||
payload["image_urls"] = [image_bytes_to_data_url(image, image_mime_type)]
|
||||
image_inputs = generation_image_inputs(
|
||||
image=image,
|
||||
image_mime_type=image_mime_type,
|
||||
image_filename=image_filename,
|
||||
images=images,
|
||||
)
|
||||
if image_inputs:
|
||||
payload["image_urls"] = [
|
||||
image_bytes_to_data_url(item.data, item.mime_type)
|
||||
for item in image_inputs
|
||||
]
|
||||
apply_extra_body(payload, model, parameters)
|
||||
response = self.session.post(
|
||||
url,
|
||||
@@ -376,12 +390,19 @@ class ImagesEditsProvider(BaseHttpProvider):
|
||||
image: bytes | None = None,
|
||||
image_mime_type: str = "image/png",
|
||||
image_filename: str = "image.png",
|
||||
images: Sequence[MultimodalImage] | None = None,
|
||||
resolution: str = "1K",
|
||||
aspect_ratio: str = "1:1",
|
||||
parameters: Mapping[str, Any] | None = None,
|
||||
) -> ImageGenerationResult:
|
||||
validate_model_config(model)
|
||||
if image is None:
|
||||
image_inputs = generation_image_inputs(
|
||||
image=image,
|
||||
image_mime_type=image_mime_type,
|
||||
image_filename=image_filename,
|
||||
images=images,
|
||||
)
|
||||
if not image_inputs:
|
||||
raise AiCapabilityError("images edits provider requires an input image")
|
||||
|
||||
url = normalize_api_url(model.url, API_IMAGES_EDITS)
|
||||
@@ -392,7 +413,15 @@ class ImagesEditsProvider(BaseHttpProvider):
|
||||
"size": resolution_to_size(resolution),
|
||||
}
|
||||
apply_extra_body(data, model, parameters)
|
||||
files = {"image": (image_filename, image, image_mime_type)}
|
||||
files: dict[str, tuple[str, bytes, str]] | list[tuple[str, tuple[str, bytes, str]]]
|
||||
if len(image_inputs) == 1:
|
||||
item = image_inputs[0]
|
||||
files = {"image": (item.filename, item.data, item.mime_type)}
|
||||
else:
|
||||
files = [
|
||||
("image", (item.filename, item.data, item.mime_type))
|
||||
for item in image_inputs
|
||||
]
|
||||
response = self.session.post(
|
||||
url,
|
||||
headers=self._headers(model),
|
||||
@@ -443,13 +472,17 @@ def build_chat_image_payload(
|
||||
*,
|
||||
image: bytes | None = None,
|
||||
image_mime_type: str = "image/png",
|
||||
images: Sequence[MultimodalImage] | None = None,
|
||||
parameters: Mapping[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return build_chat_text_payload(
|
||||
return build_chat_vision_payload(
|
||||
model,
|
||||
prompt,
|
||||
image=image,
|
||||
image_mime_type=image_mime_type,
|
||||
images=generation_image_inputs(
|
||||
image=image,
|
||||
image_mime_type=image_mime_type,
|
||||
images=images,
|
||||
),
|
||||
parameters=parameters,
|
||||
)
|
||||
|
||||
@@ -486,12 +519,17 @@ def build_gemini_payload(
|
||||
*,
|
||||
image: bytes | None = None,
|
||||
image_mime_type: str = "image/png",
|
||||
images: Sequence[MultimodalImage] | None = None,
|
||||
response_modalities: list[str],
|
||||
parameters: Mapping[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
parts: list[dict[str, Any]] = [{"text": prompt}]
|
||||
if image is not None:
|
||||
data_url = image_bytes_to_data_url(image, image_mime_type)
|
||||
for image_input in generation_image_inputs(
|
||||
image=image,
|
||||
image_mime_type=image_mime_type,
|
||||
images=images,
|
||||
):
|
||||
data_url = image_bytes_to_data_url(image_input.data, image_input.mime_type)
|
||||
mime_type, data = split_data_url(data_url)
|
||||
parts.append({"inlineData": {"mimeType": mime_type, "data": data}})
|
||||
payload: dict[str, Any] = {
|
||||
@@ -522,6 +560,26 @@ def build_gemini_vision_payload(
|
||||
return payload
|
||||
|
||||
|
||||
def generation_image_inputs(
|
||||
*,
|
||||
image: bytes | None = None,
|
||||
image_mime_type: str = "image/png",
|
||||
image_filename: str = "image.png",
|
||||
images: Sequence[MultimodalImage] | None = None,
|
||||
) -> tuple[MultimodalImage, ...]:
|
||||
if images is not None:
|
||||
return tuple(images)
|
||||
if image is None:
|
||||
return ()
|
||||
return (
|
||||
MultimodalImage(
|
||||
data=image,
|
||||
mime_type=image_mime_type,
|
||||
filename=image_filename,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def apply_extra_body(
|
||||
payload: dict[str, Any],
|
||||
model: ResolvedModel,
|
||||
|
||||
@@ -22,6 +22,7 @@ from apps.ai.providers.openai_compatible import (
|
||||
ChatCompletionsProvider,
|
||||
GeminiProvider,
|
||||
ImagesEditsProvider,
|
||||
ImagesGenerationProvider,
|
||||
)
|
||||
from apps.ai.providers.utils import image_request_timeout, resolution_to_size
|
||||
|
||||
@@ -207,6 +208,34 @@ class ChatCompletionsProviderTests(SimpleTestCase):
|
||||
self.assertEqual(content[0], {"type": "text", "text": "Generate product image"})
|
||||
self.assertTrue(content[1]["image_url"]["url"].startswith("data:image/jpeg;base64,"))
|
||||
|
||||
def test_generate_image_sends_multiple_chat_images_in_order(self):
|
||||
encoded = base64.b64encode(b"generated-image").decode("ascii")
|
||||
session = FakeSession(
|
||||
FakeResponse({"data": [{"b64_json": encoded}]})
|
||||
)
|
||||
provider = ChatCompletionsProvider(session=session)
|
||||
model = ResolvedModel(
|
||||
name="Nano Banana 2",
|
||||
url="https://api.vectorengine.ai/v1/chat/completions",
|
||||
model="gemini-3.1-flash-image-preview",
|
||||
api_key="test-key",
|
||||
api_type="chat",
|
||||
)
|
||||
|
||||
provider.generate_image(
|
||||
"Generate product image",
|
||||
model,
|
||||
images=(
|
||||
MultimodalImage(b"main-image", "image/jpeg", "main.jpg"),
|
||||
MultimodalImage(b"reference-image", "image/png", "reference.png"),
|
||||
),
|
||||
)
|
||||
|
||||
content = session.posts[0]["json"]["messages"][0]["content"]
|
||||
self.assertEqual(content[0]["text"], "Generate product image")
|
||||
self.assertTrue(content[1]["image_url"]["url"].startswith("data:image/jpeg;base64,"))
|
||||
self.assertTrue(content[2]["image_url"]["url"].startswith("data:image/png;base64,"))
|
||||
|
||||
@override_settings(AI_IMAGE_UPSTREAM_DEADLINE_SECONDS=180)
|
||||
def test_generate_image_caps_post_and_download_timeouts(self):
|
||||
session = FakeSession(
|
||||
@@ -345,6 +374,34 @@ class ChatCompletionsProviderTests(SimpleTestCase):
|
||||
)
|
||||
|
||||
|
||||
class ImagesGenerationProviderTests(SimpleTestCase):
|
||||
def test_generate_image_sends_multiple_json_images_in_order(self):
|
||||
encoded = base64.b64encode(b"generated-image").decode("ascii")
|
||||
session = FakeSession(FakeResponse({"data": [{"b64_json": encoded}]}))
|
||||
provider = ImagesGenerationProvider(session=session)
|
||||
model = ResolvedModel(
|
||||
name="JSON image provider",
|
||||
url="https://images.example.test/v1/images/generations",
|
||||
model="image-model",
|
||||
api_key="test-key",
|
||||
api_type="images",
|
||||
)
|
||||
|
||||
provider.generate_image(
|
||||
"Use the first image as the main product",
|
||||
model,
|
||||
images=(
|
||||
MultimodalImage(b"main-image", "image/jpeg", "main.jpg"),
|
||||
MultimodalImage(b"reference-image", "image/png", "reference.png"),
|
||||
),
|
||||
)
|
||||
|
||||
image_urls = session.posts[0]["json"]["image_urls"]
|
||||
self.assertEqual(len(image_urls), 2)
|
||||
self.assertTrue(image_urls[0].startswith("data:image/jpeg;base64,"))
|
||||
self.assertTrue(image_urls[1].startswith("data:image/png;base64,"))
|
||||
|
||||
|
||||
class ImagesEditsProviderTests(SimpleTestCase):
|
||||
def test_generate_image_builds_multipart_request_and_parses_base64(self):
|
||||
generated = b"edited-image"
|
||||
@@ -379,6 +436,35 @@ class ImagesEditsProviderTests(SimpleTestCase):
|
||||
("source.png", b"source-image", "image/png"),
|
||||
)
|
||||
|
||||
def test_generate_image_sends_multiple_multipart_images_in_order(self):
|
||||
encoded = base64.b64encode(b"edited-image").decode("ascii")
|
||||
session = FakeSession(FakeResponse({"data": [{"b64_json": encoded}]}))
|
||||
provider = ImagesEditsProvider(session=session)
|
||||
model = ResolvedModel(
|
||||
name="GPT Image 2",
|
||||
url="https://api.vectorengine.ai/v1/images/edits",
|
||||
model="gpt-image-2",
|
||||
api_key="test-key",
|
||||
api_type="images_edits",
|
||||
)
|
||||
|
||||
provider.generate_image(
|
||||
"Use the first image as the main product",
|
||||
model,
|
||||
images=(
|
||||
MultimodalImage(b"main-image", "image/jpeg", "main.jpg"),
|
||||
MultimodalImage(b"reference-image", "image/png", "reference.png"),
|
||||
),
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
session.posts[0]["files"],
|
||||
[
|
||||
("image", ("main.jpg", b"main-image", "image/jpeg")),
|
||||
("image", ("reference.png", b"reference-image", "image/png")),
|
||||
],
|
||||
)
|
||||
|
||||
def test_parameters_cannot_override_core_images_edits_fields(self):
|
||||
generated = b"edited-image"
|
||||
encoded = base64.b64encode(generated).decode("ascii")
|
||||
|
||||
Reference in New Issue
Block a user