feat: retry async image tasks
This commit is contained in:
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user