feat: add async image task API

This commit is contained in:
QiuSW
2026-07-08 22:08:48 +08:00
parent c98f713762
commit 25a4080177
26 changed files with 1531 additions and 48 deletions
+276 -2
View File
@@ -2,6 +2,7 @@ import uuid
import base64
import json
import tempfile
from datetime import timedelta
from decimal import Decimal
from pathlib import Path
from unittest.mock import patch
@@ -28,6 +29,12 @@ from apps.api.generation import (
prepare_generation,
run_synchronous_generation,
)
from apps.api.image_tasks import (
claim_next_image_task,
reap_stale_image_tasks,
run_image_generation_task,
)
from apps.api.models import ImageGenerationTask
from apps.api.throttles import GenerateRateThrottle
from apps.api.views import ClientLatestReleaseView, ExternalApiView, ModelsView
from apps.ai.models import AiModel, ModelAlias
@@ -1157,9 +1164,9 @@ class GenerateApiTests(TestCase):
def auth_header(self) -> dict:
return {"HTTP_AUTHORIZATION": f"Bearer {self.raw_key}"}
def post_with_provider(self, path, payload, provider=None):
def post_with_provider(self, path, payload, provider=None, **extra):
with patch("apps.api.generation.get_provider", return_value=provider or self.provider):
return self.client.post(path, payload, format="json", **self.auth_header())
return self.client.post(path, payload, format="json", **self.auth_header(), **extra)
def assert_generation_not_charged(self):
self.wallet.refresh_from_db()
@@ -1281,6 +1288,273 @@ class GenerateApiTests(TestCase):
self.assertEqual(call.result_summary, "image_bytes=21")
self.assertNotIn("SECRET_RAW", call.result_ref + call.result_summary)
@override_settings(
MODERATION_ENABLED=True,
MODERATION_PROVIDER="keyword",
MODERATION_CACHE_VERSION_KEY="test:api:moderation:sensitive_words:version",
)
def test_async_image_blocked_prompt_creates_no_task_or_charge(self):
SensitiveWord.objects.create(word="敏感词", category="policy")
with (
patch("apps.api.generation.socket.getaddrinfo") as dns_lookup,
patch("apps.api.generation.requests.Session.get") as image_get,
):
response = self.post_with_provider(
"/api/v1/generate/image/tasks",
{
"prompt": "请生成敏-感\u200b 词图片",
"model": self.image_alias,
"image_url": "https://safe.example.com/input.jpg",
"resolution": "1K",
"aspect_ratio": "1:1",
},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "content_blocked")
dns_lookup.assert_not_called()
image_get.assert_not_called()
self.assertFalse(ImageGenerationTask.objects.exists())
self.assert_generation_not_charged()
def test_async_image_insufficient_points_returns_402_without_task(self):
self.wallet.points_balance = 1
self.wallet.save(update_fields=("points_balance", "updated_at"))
response = self.post_with_provider(
"/api/v1/generate/image/tasks",
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
)
self.assertEqual(response.status_code, 402)
self.assertEqual(response.data["error"]["code"], "insufficient_points")
self.assertFalse(ImageGenerationTask.objects.exists())
self.assertEqual(self.provider.image_calls, [])
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 1)
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
def test_async_image_idempotency_reuses_task_and_rejects_conflict(self):
payload = {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}
first = self.post_with_provider(
"/api/v1/generate/image/tasks",
payload,
HTTP_IDEMPOTENCY_KEY="image-job-001",
)
second = self.post_with_provider(
"/api/v1/generate/image/tasks",
payload,
HTTP_IDEMPOTENCY_KEY="image-job-001",
)
conflict = self.post_with_provider(
"/api/v1/generate/image/tasks",
{**payload, "prompt": "生成另一张图片"},
HTTP_IDEMPOTENCY_KEY="image-job-001",
)
self.assertEqual(first.status_code, 202)
self.assertEqual(second.status_code, 202)
self.assertEqual(first.data["task_id"], second.data["task_id"])
self.assertEqual(conflict.status_code, 409)
self.assertEqual(conflict.data["error"]["code"], "idempotency_conflict")
self.assertEqual(ImageGenerationTask.objects.count(), 1)
self.assertEqual(CallRecord.objects.filter(user=self.user).count(), 1)
self.assertEqual(
PointsLedger.objects.filter(
user=self.user,
change_type=PointsLedger.ChangeType.CONSUME,
).count(),
1,
)
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 90)
self.assertEqual(self.provider.image_calls, [])
@override_settings(MEDIA_PUBLIC_BASE_URL="https://cm.example.test")
def test_async_image_worker_success_and_poll_are_idempotent(self):
response = self.post_with_provider(
"/api/v1/generate/image/tasks",
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
)
self.assertEqual(response.status_code, 202)
self.assertEqual(response.data["status"], ImageGenerationTask.Status.QUEUED)
self.assertEqual(response.data["points_balance"], 90)
self.assertEqual(self.provider.image_calls, [])
with patch("apps.api.generation.get_provider", return_value=self.provider):
claimed = claim_next_image_task("worker-a")
self.assertIsNotNone(claimed)
task = run_image_generation_task(claimed, worker_id="worker-a")
self.assertEqual(task.status, ImageGenerationTask.Status.SUCCEEDED)
self.assertTrue(task.result_url.startswith("https://cm.example.test/media/"))
self.assertEqual(len(self.provider.image_calls), 1)
poll = self.client.get(
f"/api/v1/generate/image/tasks/{response.data['task_id']}",
**self.auth_header(),
)
repeat = self.client.get(
f"/api/v1/generate/image/tasks/{response.data['task_id']}",
**self.auth_header(),
)
self.assertEqual(poll.status_code, 200)
self.assertEqual(poll.data["status"], ImageGenerationTask.Status.SUCCEEDED)
self.assertEqual(poll.data["result"]["image_url"], task.result_url)
self.assertEqual(repeat.data["result"]["image_url"], task.result_url)
def test_async_image_poll_rejects_cross_user_access(self):
response = self.post_with_provider(
"/api/v1/generate/image/tasks",
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
)
other_user = get_user_model().objects.create_user(
username=f"other-{uuid.uuid4().hex[:8]}",
email=f"other-{uuid.uuid4().hex[:8]}@example.com",
password="password",
)
_other_key, other_raw_key = ApiKey.create_for_user(other_user, name="other")
denied = self.client.get(
f"/api/v1/generate/image/tasks/{response.data['task_id']}",
HTTP_AUTHORIZATION=f"Bearer {other_raw_key}",
)
self.assertEqual(denied.status_code, 404)
self.assertEqual(denied.data["error"]["code"], "task_not_found")
def test_async_image_worker_failure_refunds_precharged_points(self):
self.provider.image_error = requests.Timeout("image upstream deadline exceeded")
response = self.post_with_provider(
"/api/v1/generate/image/tasks",
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
)
with patch("apps.api.generation.get_provider", return_value=self.provider):
task = run_image_generation_task(
claim_next_image_task("worker-failure"),
worker_id="worker-failure",
)
self.assertEqual(task.status, ImageGenerationTask.Status.FAILED)
self.assertEqual(task.error_code, "upstream_timeout")
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
call = CallRecord.objects.get(pk=response.data["call_id"])
self.assertEqual(call.status, CallRecord.Status.FAILED)
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.REFUND,
).count(),
1,
)
poll = self.client.get(
f"/api/v1/generate/image/tasks/{response.data['task_id']}",
**self.auth_header(),
)
self.assertEqual(poll.data["status"], ImageGenerationTask.Status.FAILED)
self.assertEqual(poll.data["error"]["code"], "upstream_timeout")
def test_async_image_reaper_fails_stale_running_task_and_refunds(self):
response = self.post_with_provider(
"/api/v1/generate/image/tasks",
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
)
claimed = claim_next_image_task("worker-crash")
stale_at = timezone.now() - timedelta(seconds=5)
ImageGenerationTask.objects.filter(pk=claimed.pk).update(
lease_expires_at=stale_at,
heartbeat_at=stale_at,
)
reaped = reap_stale_image_tasks(now=timezone.now())
task = ImageGenerationTask.objects.get(pk=claimed.pk)
self.assertEqual(reaped, 1)
self.assertEqual(task.status, ImageGenerationTask.Status.FAILED)
self.assertEqual(task.error_code, "task_timeout")
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
call = CallRecord.objects.get(pk=response.data["call_id"])
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.REFUND,
).count(),
1,
)
@override_settings(MEDIA_PUBLIC_BASE_URL="https://cm.example.test")
def test_async_image_duplicate_worker_does_not_double_charge_or_refund(self):
response = self.post_with_provider(
"/api/v1/generate/image/tasks",
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
)
with patch("apps.api.generation.get_provider", return_value=self.provider):
task = run_image_generation_task(
claim_next_image_task("worker-a"),
worker_id="worker-a",
)
duplicate = run_image_generation_task(task, worker_id="worker-b")
self.assertEqual(duplicate.status, ImageGenerationTask.Status.SUCCEEDED)
self.assertEqual(duplicate.result_url, task.result_url)
self.assertEqual(len(self.provider.image_calls), 1)
call = CallRecord.objects.get(pk=response.data["call_id"])
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.CONSUME,
).count(),
1,
)
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.REFUND,
).count(),
0,
)
def test_async_image_late_worker_after_reaper_cannot_flip_failed_task(self):
response = self.post_with_provider(
"/api/v1/generate/image/tasks",
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
)
claimed = claim_next_image_task("worker-late")
stale_at = timezone.now() - timedelta(seconds=5)
ImageGenerationTask.objects.filter(pk=claimed.pk).update(
lease_expires_at=stale_at,
heartbeat_at=stale_at,
)
reap_stale_image_tasks(now=timezone.now())
with patch("apps.api.generation.get_provider", return_value=self.provider):
late = run_image_generation_task(claimed, worker_id="worker-late")
self.assertEqual(late.status, ImageGenerationTask.Status.FAILED)
self.assertEqual(late.result_url, "")
self.assertEqual(self.provider.image_calls, [])
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
call = CallRecord.objects.get(pk=response.data["call_id"])
self.assertEqual(call.status, CallRecord.Status.FAILED)
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.REFUND,
).count(),
1,
)
def test_generation_core_saves_image_with_url_builder_without_request(self):
encoded = base64.b64encode(b"input-image").decode("ascii")