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
+146 -1
View File
@@ -36,7 +36,7 @@ from apps.api.image_tasks import (
reap_stale_image_tasks,
run_image_generation_task,
)
from apps.api.models import ImageGenerationTask
from apps.api.models import ImageGenerationTask, ImageGenerationTaskInput
from apps.api.throttles import GenerateRateThrottle
from apps.api.views import ClientLatestReleaseView, ExternalApiView, ModelsView
from apps.ai.models import AiModel, ModelAlias
@@ -1672,6 +1672,112 @@ class GenerateApiTests(TestCase):
self.assertEqual(call.result_summary, "image_bytes=21")
self.assertNotIn("SECRET_RAW", call.result_ref + call.result_summary)
def test_generate_image_accepts_ordered_images_and_injects_role_rules(self):
first = base64.b64encode(b"main-image").decode("ascii")
second = base64.b64encode(b"reference-image").decode("ascii")
response = self.post_with_provider(
"/api/v1/generate/image",
{
"prompt": "生成新的商品主图",
"model": self.image_alias,
"images": [
{"image_base64": f"data:image/jpeg;base64,{first}"},
{"image_base64": f"data:image/png;base64,{second}"},
],
},
)
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data["points_cost"], 10)
self.assertEqual(response.data["points_balance"], 90)
provider_call = self.provider.image_calls[0]
self.assertEqual(provider_call["image"], b"main-image")
self.assertEqual(
[image.data for image in provider_call["images"]],
[b"main-image", b"reference-image"],
)
self.assertIn("第 1 张图片是主商品图", provider_call["prompt"])
self.assertIn("第 2 张及之后的图片仅作为", provider_call["prompt"])
self.assertIn("生成新的商品主图", provider_call["prompt"])
def test_generate_image_accepts_mixed_base64_and_url_images_in_order(self):
encoded = base64.b64encode(b"main-image").decode("ascii")
downloaded = FakeImageUrlResponse(
headers={"Content-Type": "image/jpeg"},
chunks=(b"reference-image",),
)
with (
patch(
"apps.api.generation.socket.getaddrinfo",
return_value=dns_result("93.184.216.34"),
),
patch(
"apps.api.generation.requests.Session.get",
return_value=downloaded,
),
):
response = self.post_with_provider(
"/api/v1/generate/image",
{
"prompt": "生成新的商品主图",
"model": self.image_alias,
"images": [
{"image_base64": f"data:image/png;base64,{encoded}"},
{"image_url": "https://images.example.test/reference.jpg"},
],
},
)
self.assertEqual(response.status_code, 200)
provider_call = self.provider.image_calls[0]
self.assertEqual(
[image.data for image in provider_call["images"]],
[b"main-image", b"reference-image"],
)
self.assertEqual(
[image.mime_type for image in provider_call["images"]],
["image/png", "image/jpeg"],
)
def test_generate_image_rejects_mixed_legacy_and_images_inputs_without_charge(self):
encoded = base64.b64encode(b"main-image").decode("ascii")
response = self.post_with_provider(
"/api/v1/generate/image",
{
"prompt": "生成图片",
"model": self.image_alias,
"image_base64": f"data:image/png;base64,{encoded}",
"images": [{"image_base64": f"data:image/png;base64,{encoded}"}],
},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(self.provider.image_calls, [])
self.assert_generation_not_charged()
@override_settings(IMAGE_MAX_INPUT_IMAGES=1)
def test_generate_image_rejects_too_many_input_images_without_charge(self):
encoded = base64.b64encode(b"input-image").decode("ascii")
response = self.post_with_provider(
"/api/v1/generate/image",
{
"prompt": "生成图片",
"model": self.image_alias,
"images": [
{"image_base64": encoded},
{"image_base64": encoded},
],
},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(self.provider.image_calls, [])
self.assert_generation_not_charged()
def test_sync_image_usage_telemetry_logs_safe_client_version_and_key_identity(self):
encoded = base64.b64encode(b"input-image").decode("ascii")
payload = {
@@ -1883,6 +1989,45 @@ class GenerateApiTests(TestCase):
self.assertEqual(poll.data["result"]["image_url"], task.result_url)
self.assertEqual(repeat.data["result"]["image_url"], task.result_url)
@override_settings(MEDIA_PUBLIC_BASE_URL="https://cm.example.test")
def test_async_image_task_stores_and_restores_ordered_inputs(self):
first = base64.b64encode(b"main-image").decode("ascii")
second = base64.b64encode(b"reference-image").decode("ascii")
response = self.post_with_provider(
"/api/v1/generate/image/tasks",
{
"prompt": "生成新的商品主图",
"model": self.image_alias,
"images": [
{"image_base64": f"data:image/jpeg;base64,{first}"},
{"image_base64": f"data:image/png;base64,{second}"},
],
},
)
self.assertEqual(response.status_code, 202)
task = ImageGenerationTask.objects.get(task_id=response.data["task_id"])
stored_inputs = list(task.input_images.order_by("ordinal"))
self.assertEqual(len(stored_inputs), 2)
self.assertEqual([item.ordinal for item in stored_inputs], [0, 1])
self.assertFalse(bool(task.input_image))
self.assertFalse(ImageGenerationTaskInput.objects.filter(task=task, image__isnull=True).exists())
serialized = json.dumps(task.request_payload, ensure_ascii=False)
self.assertNotIn(first, serialized)
self.assertNotIn(second, serialized)
with patch("apps.api.generation.get_provider", return_value=self.provider):
claimed = claim_next_image_task("worker-multi")
completed = run_image_generation_task(claimed, worker_id="worker-multi")
self.assertEqual(completed.status, ImageGenerationTask.Status.SUCCEEDED)
provider_call = self.provider.image_calls[0]
self.assertEqual(
[image.data for image in provider_call["images"]],
[b"main-image", b"reference-image"],
)
self.assertIn("第 1 张图片是主商品图", provider_call["prompt"])
def test_async_image_poll_rejects_cross_user_access(self):
response = self.post_with_provider(
"/api/v1/generate/image/tasks",