feat: support multi-image image generation
This commit is contained in:
+146
-1
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user