feat: support multi-image image generation
This commit is contained in:
@@ -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