feat: retry async image tasks

This commit is contained in:
QiuSW
2026-07-09 14:46:09 +08:00
parent e3598bfe8b
commit 8878769ec5
16 changed files with 550 additions and 42 deletions
+236
View File
@@ -1507,6 +1507,9 @@ class GenerateApiTests(TestCase):
self.assertEqual(response.status_code, 202)
self.assertEqual(response.data["status"], ImageGenerationTask.Status.QUEUED)
self.assertEqual(response.data["points_balance"], 90)
self.assertEqual(response.data["attempt_count"], 0)
self.assertEqual(response.data["max_attempts"], 3)
self.assertIsNone(response.data["next_attempt_at"])
self.assertEqual(self.provider.image_calls, [])
with patch("apps.api.generation.get_provider", return_value=self.provider):
@@ -1529,6 +1532,9 @@ class GenerateApiTests(TestCase):
self.assertEqual(poll.status_code, 200)
self.assertEqual(poll.data["status"], ImageGenerationTask.Status.SUCCEEDED)
self.assertEqual(poll.data["attempt_count"], 1)
self.assertEqual(poll.data["max_attempts"], 3)
self.assertIsNone(poll.data["next_attempt_at"])
self.assertEqual(poll.data["result"]["image_url"], task.result_url)
self.assertEqual(repeat.data["result"]["image_url"], task.result_url)
@@ -1552,6 +1558,7 @@ class GenerateApiTests(TestCase):
self.assertEqual(denied.status_code, 404)
self.assertEqual(denied.data["error"]["code"], "task_not_found")
@override_settings(IMAGE_TASK_MAX_RETRIES=0)
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(
@@ -1586,7 +1593,199 @@ class GenerateApiTests(TestCase):
)
self.assertEqual(poll.data["status"], ImageGenerationTask.Status.FAILED)
self.assertEqual(poll.data["error"]["code"], "upstream_timeout")
self.assertEqual(poll.data["attempt_count"], 1)
self.assertEqual(poll.data["max_attempts"], 1)
self.assertIsNone(poll.data["next_attempt_at"])
@override_settings(IMAGE_TASK_RETRY_BACKOFF_SECONDS="60,120")
def test_async_image_retryable_timeout_requeues_without_refund_and_respects_backoff(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-retry"),
worker_id="worker-retry",
)
self.assertEqual(task.status, ImageGenerationTask.Status.QUEUED)
self.assertEqual(task.attempt_count, 1)
self.assertEqual(task.error_code, "upstream_timeout")
self.assertEqual(task.error_message, "上游 AI 调用超时,稍后自动重试")
self.assertIsNotNone(task.next_attempt_at)
self.assertGreater(task.next_attempt_at, timezone.now())
self.assertIsNone(claim_next_image_task("worker-too-soon"))
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 90)
call = CallRecord.objects.get(pk=response.data["call_id"])
self.assertEqual(call.status, CallRecord.Status.PENDING)
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.REFUND,
).count(),
0,
)
poll = 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.QUEUED)
self.assertEqual(poll.data["attempt_count"], 1)
self.assertEqual(poll.data["max_attempts"], 3)
self.assertIsNotNone(poll.data["next_attempt_at"])
@override_settings(
MEDIA_PUBLIC_BASE_URL="https://cm.example.test",
IMAGE_TASK_RETRY_BACKOFF_SECONDS="0,0",
)
def test_async_image_retryable_timeouts_then_success_charges_once(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):
self.provider.image_error = requests.Timeout("first timeout")
first = run_image_generation_task(
claim_next_image_task("worker-retry-1"),
worker_id="worker-retry-1",
)
ImageGenerationTask.objects.filter(pk=first.pk).update(
next_attempt_at=timezone.now() - timedelta(seconds=1)
)
self.provider.image_error = requests.Timeout("second timeout")
second = run_image_generation_task(
claim_next_image_task("worker-retry-2"),
worker_id="worker-retry-2",
)
ImageGenerationTask.objects.filter(pk=second.pk).update(
next_attempt_at=timezone.now() - timedelta(seconds=1)
)
self.provider.image_error = None
succeeded = run_image_generation_task(
claim_next_image_task("worker-retry-3"),
worker_id="worker-retry-3",
)
self.assertEqual(first.status, ImageGenerationTask.Status.QUEUED)
self.assertEqual(second.status, ImageGenerationTask.Status.QUEUED)
self.assertEqual(succeeded.status, ImageGenerationTask.Status.SUCCEEDED)
self.assertEqual(succeeded.attempt_count, 3)
self.assertTrue(succeeded.result_url.startswith("https://cm.example.test/media/"))
self.assertEqual(len(self.provider.image_calls), 3)
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 90)
call = CallRecord.objects.get(pk=response.data["call_id"])
self.assertEqual(call.status, CallRecord.Status.SUCCESS)
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,
)
@override_settings(IMAGE_TASK_RETRY_BACKOFF_SECONDS="0,0")
def test_async_image_retryable_timeouts_final_failure_refunds_once(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):
first = run_image_generation_task(
claim_next_image_task("worker-final-1"),
worker_id="worker-final-1",
)
ImageGenerationTask.objects.filter(pk=first.pk).update(
next_attempt_at=timezone.now() - timedelta(seconds=1)
)
second = run_image_generation_task(
claim_next_image_task("worker-final-2"),
worker_id="worker-final-2",
)
ImageGenerationTask.objects.filter(pk=second.pk).update(
next_attempt_at=timezone.now() - timedelta(seconds=1)
)
failed = run_image_generation_task(
claim_next_image_task("worker-final-3"),
worker_id="worker-final-3",
)
self.assertEqual(failed.status, ImageGenerationTask.Status.FAILED)
self.assertEqual(failed.error_code, "upstream_timeout")
self.assertEqual(failed.attempt_count, 3)
self.assertIsNone(failed.next_attempt_at)
self.assertEqual(len(self.provider.image_calls), 3)
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.CONSUME,
).count(),
1,
)
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.REFUND,
).count(),
1,
)
def test_async_image_non_retryable_provider_error_fails_immediately_and_refunds(self):
self.provider.image_error = AiCapabilityError("input image is required")
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):
failed = run_image_generation_task(
claim_next_image_task("worker-no-retry"),
worker_id="worker-no-retry",
)
self.assertEqual(failed.status, ImageGenerationTask.Status.FAILED)
self.assertEqual(failed.error_code, "bad_request")
self.assertEqual(failed.attempt_count, 1)
self.assertIsNone(failed.next_attempt_at)
self.assertEqual(len(self.provider.image_calls), 1)
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,
)
@override_settings(IMAGE_TASK_MAX_RETRIES=0)
def test_run_image_tasks_logs_failed_task_alias_error_and_duration(self):
self.provider.image_error = requests.Timeout("image upstream deadline exceeded")
response = self.post_with_provider(
@@ -1611,10 +1810,47 @@ class GenerateApiTests(TestCase):
self.assertIn(f"task_id={task.task_id}", output)
self.assertIn(f"alias={self.image_alias}", output)
self.assertIn("status=failed", output)
self.assertIn("attempt=1", output)
self.assertIn("max_attempts=1", output)
self.assertIn("retrying=false", output)
self.assertIn("error_code=upstream_timeout", output)
self.assertRegex(output, r"duration_ms=\d+")
self.assertNotIn("生成图片", output)
@override_settings(IMAGE_TASK_RETRY_BACKOFF_SECONDS="60,120")
def test_run_image_tasks_logs_retrying_task_attempt_fields_without_sensitive_data(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"},
)
out = io.StringIO()
with patch("apps.api.generation.get_provider", return_value=self.provider):
call_command(
"run_image_tasks",
"--once",
"--worker-id",
"worker-log-retry",
stdout=out,
)
task = ImageGenerationTask.objects.get(task_id=response.data["task_id"])
output = out.getvalue()
self.assertEqual(task.status, ImageGenerationTask.Status.QUEUED)
self.assertIn("event=image_task_processed", output)
self.assertIn(f"task_id={task.task_id}", output)
self.assertIn(f"alias={self.image_alias}", output)
self.assertIn("status=queued", output)
self.assertIn("attempt=1", output)
self.assertIn("max_attempts=3", output)
self.assertIn("retrying=true", output)
self.assertIn("next_attempt_at=", output)
self.assertIn("error_code=upstream_timeout", output)
self.assertRegex(output, r"duration_ms=\d+")
self.assertNotIn("生成图片", output)
self.assertNotIn(self.raw_key, output)
def test_async_image_reaper_fails_stale_running_task_and_refunds(self):
response = self.post_with_provider(
"/api/v1/generate/image/tasks",