feat: retry async image tasks
This commit is contained in:
+3
-1
@@ -10,11 +10,12 @@ class ImageGenerationTaskAdmin(admin.ModelAdmin):
|
||||
"user",
|
||||
"status",
|
||||
"attempt_count",
|
||||
"next_attempt_at",
|
||||
"points_balance_after_charge",
|
||||
"created_at",
|
||||
"finished_at",
|
||||
)
|
||||
list_filter = ("status", "created_at", "finished_at")
|
||||
list_filter = ("status", "created_at", "next_attempt_at", "finished_at")
|
||||
search_fields = (
|
||||
"task_id",
|
||||
"user__username",
|
||||
@@ -37,6 +38,7 @@ class ImageGenerationTaskAdmin(admin.ModelAdmin):
|
||||
"started_at",
|
||||
"finished_at",
|
||||
"expires_at",
|
||||
"next_attempt_at",
|
||||
"locked_at",
|
||||
"lease_expires_at",
|
||||
"heartbeat_at",
|
||||
|
||||
+13
-10
@@ -246,6 +246,7 @@ def execute_precharged_generation(
|
||||
precharged: PrechargedGeneration,
|
||||
*,
|
||||
image_url_builder: ImageUrlBuilder | None = None,
|
||||
refund_on_failure: bool = True,
|
||||
) -> GenerationResult:
|
||||
prepared = precharged.prepared
|
||||
try:
|
||||
@@ -259,22 +260,24 @@ def execute_precharged_generation(
|
||||
image_url_builder=image_url_builder,
|
||||
)
|
||||
except AiCapabilityError as exc:
|
||||
refund_call_points(
|
||||
precharged.call_record,
|
||||
error_message=str(exc),
|
||||
reason=f"Provider rejected the {prepared.operation_type} request.",
|
||||
)
|
||||
if refund_on_failure:
|
||||
refund_call_points(
|
||||
precharged.call_record,
|
||||
error_message=str(exc),
|
||||
reason=f"Provider rejected the {prepared.operation_type} request.",
|
||||
)
|
||||
raise ApiRequestError(
|
||||
"bad_request",
|
||||
"请求参数不支持当前模型",
|
||||
status.HTTP_400_BAD_REQUEST,
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
refund_call_points(
|
||||
precharged.call_record,
|
||||
error_message=str(exc),
|
||||
reason=f"Upstream {prepared.operation_type} generation failed.",
|
||||
)
|
||||
if refund_on_failure:
|
||||
refund_call_points(
|
||||
precharged.call_record,
|
||||
error_message=str(exc),
|
||||
reason=f"Upstream {prepared.operation_type} generation failed.",
|
||||
)
|
||||
raise upstream_error(exc) from exc
|
||||
|
||||
return result
|
||||
|
||||
+149
-10
@@ -31,6 +31,7 @@ from .models import ImageGenerationTask
|
||||
IDEMPOTENCY_KEY_MAX_LENGTH = 128
|
||||
TASK_FAILURE_MESSAGE = "图片生成失败,已退回点数"
|
||||
TASK_TIMEOUT_MESSAGE = "图片生成任务超时,已退回点数"
|
||||
RETRYABLE_TASK_ERROR_CODES = {"upstream_timeout", "upstream_error"}
|
||||
|
||||
|
||||
def create_image_generation_task(
|
||||
@@ -207,6 +208,10 @@ def claim_next_image_task(worker_id: str | None = None) -> ImageGenerationTask |
|
||||
queryset = (
|
||||
ImageGenerationTask.objects.select_related("user", "api_key", "call_record")
|
||||
.filter(status=ImageGenerationTask.Status.QUEUED)
|
||||
.filter(
|
||||
models_q("next_attempt_at__isnull", True)
|
||||
| models_q("next_attempt_at__lte", now)
|
||||
)
|
||||
.order_by("created_at", "id")
|
||||
)
|
||||
if connection.features.has_select_for_update_skip_locked:
|
||||
@@ -224,6 +229,7 @@ def claim_next_image_task(worker_id: str | None = None) -> ImageGenerationTask |
|
||||
task.lease_expires_at = lease_expires_at
|
||||
task.heartbeat_at = now
|
||||
task.started_at = task.started_at or now
|
||||
task.next_attempt_at = None
|
||||
task.attempt_count += 1
|
||||
task.save(
|
||||
update_fields=(
|
||||
@@ -233,6 +239,7 @@ def claim_next_image_task(worker_id: str | None = None) -> ImageGenerationTask |
|
||||
"lease_expires_at",
|
||||
"heartbeat_at",
|
||||
"started_at",
|
||||
"next_attempt_at",
|
||||
"attempt_count",
|
||||
"updated_at",
|
||||
)
|
||||
@@ -261,24 +268,56 @@ def run_image_generation_task(
|
||||
|
||||
try:
|
||||
precharged = precharged_generation_for_task(task)
|
||||
result = execute_precharged_generation(
|
||||
precharged,
|
||||
image_url_builder=media_public_url_builder,
|
||||
)
|
||||
except ApiRequestError as exc:
|
||||
refund_task_call(task, exc.message)
|
||||
return mark_task_failed_if_running(task.pk, exc.code, exc.message)
|
||||
|
||||
try:
|
||||
result = execute_precharged_generation(
|
||||
precharged,
|
||||
image_url_builder=media_public_url_builder,
|
||||
refund_on_failure=False,
|
||||
)
|
||||
except ApiRequestError as exc:
|
||||
return handle_task_generation_failure(
|
||||
task,
|
||||
error_code=exc.code,
|
||||
error_message=exc.message,
|
||||
retryable=is_retryable_task_error(exc.code),
|
||||
)
|
||||
except Exception as exc:
|
||||
refund_task_call(task, str(exc))
|
||||
return mark_task_failed_if_running(
|
||||
task.pk,
|
||||
"upstream_error",
|
||||
TASK_FAILURE_MESSAGE,
|
||||
return handle_task_generation_failure(
|
||||
task,
|
||||
error_code="upstream_error",
|
||||
error_message=str(exc) or TASK_FAILURE_MESSAGE,
|
||||
retryable=True,
|
||||
)
|
||||
|
||||
return mark_task_succeeded_if_running(task.pk, result.image_url)
|
||||
|
||||
|
||||
def handle_task_generation_failure(
|
||||
task: ImageGenerationTask,
|
||||
*,
|
||||
error_code: str,
|
||||
error_message: str,
|
||||
retryable: bool,
|
||||
) -> ImageGenerationTask:
|
||||
normalized_code = str(error_code or "upstream_error")
|
||||
normalized_message = str(error_message or TASK_FAILURE_MESSAGE)
|
||||
if retryable and task.attempt_count < image_task_max_attempts():
|
||||
return mark_task_retry_if_running(
|
||||
task.pk,
|
||||
normalized_code,
|
||||
retry_error_message(normalized_code, normalized_message),
|
||||
next_attempt_at=timezone.now()
|
||||
+ timedelta(seconds=image_task_retry_backoff_seconds(task.attempt_count)),
|
||||
)
|
||||
|
||||
refund_task_call(task, normalized_message)
|
||||
return mark_task_failed_if_running(task.pk, normalized_code, normalized_message)
|
||||
|
||||
|
||||
def precharged_generation_for_task(task: ImageGenerationTask) -> PrechargedGeneration:
|
||||
payload = dict(task.request_payload or {})
|
||||
prepared = prepare_generation(
|
||||
@@ -344,6 +383,7 @@ def mark_task_succeeded_if_running(
|
||||
task.result_url = str(result_url or "")
|
||||
task.error_code = ""
|
||||
task.error_message = ""
|
||||
task.next_attempt_at = None
|
||||
task.finished_at = now
|
||||
task.heartbeat_at = now
|
||||
task.save(
|
||||
@@ -352,6 +392,7 @@ def mark_task_succeeded_if_running(
|
||||
"result_url",
|
||||
"error_code",
|
||||
"error_message",
|
||||
"next_attempt_at",
|
||||
"finished_at",
|
||||
"heartbeat_at",
|
||||
"updated_at",
|
||||
@@ -360,6 +401,41 @@ def mark_task_succeeded_if_running(
|
||||
return task
|
||||
|
||||
|
||||
def mark_task_retry_if_running(
|
||||
task_pk: int,
|
||||
error_code: str,
|
||||
error_message: str,
|
||||
*,
|
||||
next_attempt_at,
|
||||
) -> ImageGenerationTask:
|
||||
with transaction.atomic():
|
||||
task = ImageGenerationTask.objects.select_for_update().get(pk=task_pk)
|
||||
if task.status != ImageGenerationTask.Status.RUNNING:
|
||||
return task
|
||||
task.status = ImageGenerationTask.Status.QUEUED
|
||||
task.error_code = str(error_code or "upstream_error")
|
||||
task.error_message = str(error_message or TASK_FAILURE_MESSAGE)
|
||||
task.next_attempt_at = next_attempt_at
|
||||
task.worker_id = ""
|
||||
task.locked_at = None
|
||||
task.lease_expires_at = None
|
||||
task.heartbeat_at = None
|
||||
task.save(
|
||||
update_fields=(
|
||||
"status",
|
||||
"error_code",
|
||||
"error_message",
|
||||
"next_attempt_at",
|
||||
"worker_id",
|
||||
"locked_at",
|
||||
"lease_expires_at",
|
||||
"heartbeat_at",
|
||||
"updated_at",
|
||||
)
|
||||
)
|
||||
return task
|
||||
|
||||
|
||||
def mark_task_failed_if_running(
|
||||
task_pk: int,
|
||||
error_code: str,
|
||||
@@ -373,6 +449,7 @@ def mark_task_failed_if_running(
|
||||
task.status = ImageGenerationTask.Status.FAILED
|
||||
task.error_code = str(error_code or "upstream_error")
|
||||
task.error_message = str(error_message or TASK_FAILURE_MESSAGE)
|
||||
task.next_attempt_at = None
|
||||
task.finished_at = now
|
||||
task.heartbeat_at = now
|
||||
task.save(
|
||||
@@ -380,6 +457,7 @@ def mark_task_failed_if_running(
|
||||
"status",
|
||||
"error_code",
|
||||
"error_message",
|
||||
"next_attempt_at",
|
||||
"finished_at",
|
||||
"heartbeat_at",
|
||||
"updated_at",
|
||||
@@ -422,6 +500,7 @@ def reap_stale_image_task(task_pk: int, now) -> bool:
|
||||
task.result_url = task.call_record.result_ref
|
||||
task.error_code = ""
|
||||
task.error_message = ""
|
||||
task.next_attempt_at = None
|
||||
task.finished_at = now
|
||||
task.save(
|
||||
update_fields=(
|
||||
@@ -429,6 +508,7 @@ def reap_stale_image_task(task_pk: int, now) -> bool:
|
||||
"result_url",
|
||||
"error_code",
|
||||
"error_message",
|
||||
"next_attempt_at",
|
||||
"finished_at",
|
||||
"updated_at",
|
||||
)
|
||||
@@ -446,14 +526,24 @@ def reap_stale_image_task(task_pk: int, now) -> bool:
|
||||
if task.call_record.status == CallRecord.Status.SUCCESS:
|
||||
task.status = ImageGenerationTask.Status.SUCCEEDED
|
||||
task.result_url = task.call_record.result_ref
|
||||
task.next_attempt_at = None
|
||||
task.finished_at = now
|
||||
task.save(update_fields=("status", "result_url", "finished_at", "updated_at"))
|
||||
task.save(
|
||||
update_fields=(
|
||||
"status",
|
||||
"result_url",
|
||||
"next_attempt_at",
|
||||
"finished_at",
|
||||
"updated_at",
|
||||
)
|
||||
)
|
||||
return True
|
||||
raise
|
||||
|
||||
task.status = ImageGenerationTask.Status.FAILED
|
||||
task.error_code = "task_timeout"
|
||||
task.error_message = TASK_TIMEOUT_MESSAGE
|
||||
task.next_attempt_at = None
|
||||
task.finished_at = now
|
||||
task.heartbeat_at = now
|
||||
task.save(
|
||||
@@ -461,6 +551,7 @@ def reap_stale_image_task(task_pk: int, now) -> bool:
|
||||
"status",
|
||||
"error_code",
|
||||
"error_message",
|
||||
"next_attempt_at",
|
||||
"finished_at",
|
||||
"heartbeat_at",
|
||||
"updated_at",
|
||||
@@ -510,6 +601,9 @@ def task_submit_response(task: ImageGenerationTask) -> dict[str, Any]:
|
||||
"call_id": call_record.id,
|
||||
"points_cost": call_record.points_cost,
|
||||
"points_balance": task.points_balance_after_charge,
|
||||
"attempt_count": task.attempt_count,
|
||||
"max_attempts": image_task_max_attempts(),
|
||||
"next_attempt_at": task.next_attempt_at.isoformat() if task.next_attempt_at else None,
|
||||
"created_at": task.created_at.isoformat(),
|
||||
"expires_at": task.expires_at.isoformat() if task.expires_at else None,
|
||||
}
|
||||
@@ -522,6 +616,9 @@ def task_detail_response(task: ImageGenerationTask) -> dict[str, Any]:
|
||||
"status": task.status,
|
||||
"call_id": call_record.id,
|
||||
"points_cost": call_record.points_cost,
|
||||
"attempt_count": task.attempt_count,
|
||||
"max_attempts": image_task_max_attempts(),
|
||||
"next_attempt_at": task.next_attempt_at.isoformat() if task.next_attempt_at else None,
|
||||
"created_at": task.created_at.isoformat(),
|
||||
"updated_at": task.updated_at.isoformat(),
|
||||
"expires_at": task.expires_at.isoformat() if task.expires_at else None,
|
||||
@@ -553,6 +650,48 @@ def image_task_lease_seconds() -> int:
|
||||
return max(1, int(getattr(settings, "IMAGE_TASK_LEASE_SECONDS", 600)))
|
||||
|
||||
|
||||
def image_task_max_retries() -> int:
|
||||
return max(0, int(getattr(settings, "IMAGE_TASK_MAX_RETRIES", 2)))
|
||||
|
||||
|
||||
def image_task_max_attempts() -> int:
|
||||
return 1 + image_task_max_retries()
|
||||
|
||||
|
||||
def image_task_retry_backoff_seconds(attempt_count: int) -> int:
|
||||
values = image_task_retry_backoff_values()
|
||||
if not values:
|
||||
return 0
|
||||
index = max(0, int(attempt_count or 1) - 1)
|
||||
return values[min(index, len(values) - 1)]
|
||||
|
||||
|
||||
def image_task_retry_backoff_values() -> list[int]:
|
||||
raw = str(getattr(settings, "IMAGE_TASK_RETRY_BACKOFF_SECONDS", "10,30") or "")
|
||||
values: list[int] = []
|
||||
for part in raw.split(","):
|
||||
item = part.strip()
|
||||
if not item:
|
||||
continue
|
||||
try:
|
||||
values.append(max(0, int(item)))
|
||||
except ValueError:
|
||||
continue
|
||||
return values
|
||||
|
||||
|
||||
def is_retryable_task_error(error_code: str) -> bool:
|
||||
return str(error_code or "") in RETRYABLE_TASK_ERROR_CODES
|
||||
|
||||
|
||||
def retry_error_message(error_code: str, fallback: str) -> str:
|
||||
if error_code == "upstream_timeout":
|
||||
return "上游 AI 调用超时,稍后自动重试"
|
||||
if error_code == "upstream_error":
|
||||
return "上游 AI 调用失败,稍后自动重试"
|
||||
return fallback
|
||||
|
||||
|
||||
def normalize_worker_id(worker_id: str | None) -> str:
|
||||
normalized = str(worker_id or "").strip()
|
||||
if normalized:
|
||||
|
||||
@@ -4,8 +4,12 @@ import uuid
|
||||
from django.conf import settings
|
||||
from django.core.management.base import BaseCommand
|
||||
|
||||
from apps.api.image_tasks import (
|
||||
image_task_max_attempts,
|
||||
reap_stale_image_tasks,
|
||||
run_one_image_task,
|
||||
)
|
||||
from apps.api.models import ImageGenerationTask
|
||||
from apps.api.image_tasks import reap_stale_image_tasks, run_one_image_task
|
||||
|
||||
|
||||
class Command(BaseCommand):
|
||||
@@ -70,13 +74,25 @@ class Command(BaseCommand):
|
||||
def format_task_log_line(task: ImageGenerationTask, *, started_at: float) -> str:
|
||||
duration_ms = max(0, int((time.monotonic() - started_at) * 1000))
|
||||
alias = str((task.request_payload or {}).get("model") or "")
|
||||
retrying = bool(
|
||||
task.status == ImageGenerationTask.Status.QUEUED
|
||||
and task.error_code
|
||||
and task.next_attempt_at
|
||||
)
|
||||
fields = {
|
||||
"event": "image_task_processed",
|
||||
"task_id": str(task.task_id),
|
||||
"status": task.status,
|
||||
"alias": alias,
|
||||
"attempt": str(task.attempt_count),
|
||||
"max_attempts": str(image_task_max_attempts()),
|
||||
"retrying": str(retrying).lower(),
|
||||
"next_attempt_at": task.next_attempt_at.isoformat() if task.next_attempt_at else "",
|
||||
"duration_ms": str(duration_ms),
|
||||
}
|
||||
if task.status in {ImageGenerationTask.Status.FAILED, ImageGenerationTask.Status.EXPIRED}:
|
||||
if retrying or task.status in {
|
||||
ImageGenerationTask.Status.FAILED,
|
||||
ImageGenerationTask.Status.EXPIRED,
|
||||
}:
|
||||
fields["error_code"] = task.error_code or "upstream_error"
|
||||
return " ".join(f"{key}={value}" for key, value in fields.items())
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
# Generated by Django 5.2.15 on 2026-07-09 06:28
|
||||
|
||||
from django.conf import settings
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('api', '0001_initial'),
|
||||
('billing', '0007_alter_pointsledger_change_type_signupbonusgrant'),
|
||||
('users', '0005_alter_apikey_created_at_alter_apikey_key_hash_and_more'),
|
||||
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='imagegenerationtask',
|
||||
name='next_attempt_at',
|
||||
field=models.DateTimeField(blank=True, null=True, verbose_name='下次重试时间'),
|
||||
),
|
||||
migrations.AddIndex(
|
||||
model_name='imagegenerationtask',
|
||||
index=models.Index(fields=['status', 'next_attempt_at'], name='image_gener_status_e2e2f3_idx'),
|
||||
),
|
||||
]
|
||||
@@ -62,6 +62,7 @@ class ImageGenerationTask(models.Model):
|
||||
started_at = models.DateTimeField("开始时间", null=True, blank=True)
|
||||
finished_at = models.DateTimeField("完成时间", null=True, blank=True)
|
||||
expires_at = models.DateTimeField("任务元数据过期时间", null=True, blank=True)
|
||||
next_attempt_at = models.DateTimeField("下次重试时间", null=True, blank=True)
|
||||
locked_at = models.DateTimeField("锁定时间", null=True, blank=True)
|
||||
lease_expires_at = models.DateTimeField("租约过期时间", null=True, blank=True)
|
||||
heartbeat_at = models.DateTimeField("心跳时间", null=True, blank=True)
|
||||
@@ -89,6 +90,7 @@ class ImageGenerationTask(models.Model):
|
||||
models.Index(fields=("user", "created_at")),
|
||||
models.Index(fields=("api_key", "created_at")),
|
||||
models.Index(fields=("status", "created_at")),
|
||||
models.Index(fields=("status", "next_attempt_at")),
|
||||
models.Index(fields=("status", "lease_expires_at")),
|
||||
models.Index(fields=("worker_id", "status")),
|
||||
models.Index(fields=("expires_at",)),
|
||||
|
||||
@@ -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