feat: add async image task API
This commit is contained in:
+276
-2
@@ -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")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user