feat: support multi-image image generation

This commit is contained in:
QiuSW
2026-07-17 15:55:38 +08:00
parent bd274e26e3
commit 4a0becbb4d
23 changed files with 690 additions and 80 deletions
+2
View File
@@ -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,
+67 -9
View File
@@ -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,
+86
View File
@@ -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")