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
+45 -1
View File
@@ -1,3 +1,47 @@
from django.contrib import admin
# Register your models here.
from .models import ImageGenerationTask
@admin.register(ImageGenerationTask)
class ImageGenerationTaskAdmin(admin.ModelAdmin):
list_display = (
"task_id",
"user",
"status",
"attempt_count",
"points_balance_after_charge",
"created_at",
"finished_at",
)
list_filter = ("status", "created_at", "finished_at")
search_fields = (
"task_id",
"user__username",
"api_key__key_prefix",
"call_record__id",
"idempotency_key",
"error_code",
)
readonly_fields = (
"task_id",
"user",
"api_key",
"call_record",
"idempotency_key_hash",
"request_hash",
"request_payload",
"input_image",
"result_url",
"points_balance_after_charge",
"started_at",
"finished_at",
"expires_at",
"locked_at",
"lease_expires_at",
"heartbeat_at",
"worker_id",
"attempt_count",
"created_at",
"updated_at",
)
+581
View File
@@ -0,0 +1,581 @@
from __future__ import annotations
import hashlib
import json
import socket
from dataclasses import replace
from datetime import timedelta
from typing import Any, Mapping
from django.conf import settings
from django.core.files.base import ContentFile
from django.db import IntegrityError, connection, transaction
from django.utils import timezone
from rest_framework import status
from apps.billing.models import CallRecord
from apps.billing.services import InvalidCallStateError, refund_call_points
from .generation import (
ApiRequestError,
GenerationInput,
ImageInput,
PrechargedGeneration,
execute_precharged_generation,
prepare_generation,
precharge_generation,
)
from .models import ImageGenerationTask
IDEMPOTENCY_KEY_MAX_LENGTH = 128
TASK_FAILURE_MESSAGE = "图片生成失败,已退回点数"
TASK_TIMEOUT_MESSAGE = "图片生成任务超时,已退回点数"
def create_image_generation_task(
*,
user,
api_key,
request_data: Mapping[str, Any],
idempotency_key: str = "",
) -> tuple[ImageGenerationTask, bool]:
normalized_key = normalize_idempotency_key(idempotency_key)
idempotency_key_hash = hash_text(normalized_key) if normalized_key else None
request_hash = request_hash_for_image_request(request_data)
existing = find_idempotent_task(api_key, idempotency_key_hash)
if existing is not None:
return idempotent_task_or_raise(existing, request_hash), False
prepared = prepare_generation(
GenerationInput(
user=user,
api_key=api_key,
operation_type=CallRecord.OperationType.IMAGE,
prompt=request_data["prompt"],
alias=request_data.get("model") or None,
resolution=request_data.get("resolution") or "1K",
parameters=dict(request_data.get("parameters") or {}),
image_url=str(request_data.get("image_url") or ""),
image_base64=str(request_data.get("image_base64") or ""),
aspect_ratio=request_data.get("aspect_ratio") or "1:1",
)
)
try:
with transaction.atomic():
existing = find_idempotent_task(
api_key,
idempotency_key_hash,
for_update=True,
)
if existing is not None:
return idempotent_task_or_raise(existing, request_hash), False
precharged = precharge_generation(prepared)
now = timezone.now()
task = ImageGenerationTask.objects.create(
user=user,
api_key=api_key,
call_record=precharged.call_record,
idempotency_key=normalized_key,
idempotency_key_hash=idempotency_key_hash,
request_hash=request_hash,
request_payload=task_request_payload(
request_data=request_data,
image_input=prepared.image_input,
),
points_balance_after_charge=precharged.points_balance_after_charge,
expires_at=now + timedelta(hours=image_task_retention_hours()),
)
store_task_input_image(task, prepared.image_input)
return task, True
except IntegrityError:
if idempotency_key_hash:
existing = find_idempotent_task(api_key, idempotency_key_hash)
if existing is not None:
return idempotent_task_or_raise(existing, request_hash), False
raise
def find_idempotent_task(api_key, idempotency_key_hash: str | None, *, for_update=False):
if not idempotency_key_hash:
return None
queryset = ImageGenerationTask.objects.select_related("call_record").filter(
api_key=api_key,
idempotency_key_hash=idempotency_key_hash,
)
if for_update:
queryset = queryset.select_for_update()
return queryset.order_by("id").first()
def idempotent_task_or_raise(
task: ImageGenerationTask,
request_hash: str,
) -> ImageGenerationTask:
if task.request_hash != request_hash:
raise ApiRequestError(
"idempotency_conflict",
"Idempotency-Key 已用于不同请求",
status.HTTP_409_CONFLICT,
)
return task
def normalize_idempotency_key(value: str) -> str:
normalized = str(value or "").strip()
if len(normalized) > IDEMPOTENCY_KEY_MAX_LENGTH:
raise ApiRequestError(
"bad_request",
"Idempotency-Key 过长",
status.HTTP_400_BAD_REQUEST,
)
return normalized
def request_hash_for_image_request(request_data: Mapping[str, Any]) -> str:
image_base64 = str(request_data.get("image_base64") or "")
payload = {
"prompt": str(request_data.get("prompt") or ""),
"model": str(request_data.get("model") or ""),
"resolution": str(request_data.get("resolution") or "1K"),
"aspect_ratio": str(request_data.get("aspect_ratio") or "1:1"),
"parameters": dict(request_data.get("parameters") or {}),
"image_url": str(request_data.get("image_url") or ""),
"image_base64_sha256": hash_text(image_base64) if image_base64 else "",
}
return hash_json(payload)
def task_request_payload(
*,
request_data: Mapping[str, Any],
image_input: ImageInput | None,
) -> dict[str, Any]:
image_source = "none"
if request_data.get("image_base64"):
image_source = "base64_stored"
elif request_data.get("image_url"):
image_source = "url_stored"
payload = {
"prompt": str(request_data.get("prompt") or ""),
"model": str(request_data.get("model") or ""),
"resolution": str(request_data.get("resolution") or "1K"),
"aspect_ratio": str(request_data.get("aspect_ratio") or "1:1"),
"parameters": dict(request_data.get("parameters") or {}),
"image_url": str(request_data.get("image_url") or ""),
"image_input": None,
}
if image_input is not None:
payload["image_input"] = {
"source": image_source,
"mime_type": image_input.mime_type,
"filename": image_input.filename,
"storage_path": "",
}
return payload
def store_task_input_image(
task: ImageGenerationTask,
image_input: ImageInput | None,
) -> None:
if image_input is None:
return
task.input_image.save(
image_input.filename,
ContentFile(image_input.data),
save=True,
)
payload = dict(task.request_payload or {})
image_info = dict(payload.get("image_input") or {})
image_info["storage_path"] = task.input_image.name
payload["image_input"] = image_info
task.request_payload = payload
task.save(update_fields=("request_payload", "updated_at"))
def claim_next_image_task(worker_id: str | None = None) -> ImageGenerationTask | None:
normalized_worker_id = normalize_worker_id(worker_id)
now = timezone.now()
lease_expires_at = now + timedelta(seconds=image_task_lease_seconds())
with transaction.atomic():
queryset = (
ImageGenerationTask.objects.select_related("user", "api_key", "call_record")
.filter(status=ImageGenerationTask.Status.QUEUED)
.order_by("created_at", "id")
)
if connection.features.has_select_for_update_skip_locked:
queryset = queryset.select_for_update(skip_locked=True)
else:
queryset = queryset.select_for_update()
task = queryset.first()
if task is None:
return None
task.status = ImageGenerationTask.Status.RUNNING
task.worker_id = normalized_worker_id
task.locked_at = now
task.lease_expires_at = lease_expires_at
task.heartbeat_at = now
task.started_at = task.started_at or now
task.attempt_count += 1
task.save(
update_fields=(
"status",
"worker_id",
"locked_at",
"lease_expires_at",
"heartbeat_at",
"started_at",
"attempt_count",
"updated_at",
)
)
return task
def run_one_image_task(worker_id: str | None = None) -> ImageGenerationTask | None:
task = claim_next_image_task(worker_id)
if task is None:
return None
return run_image_generation_task(task, worker_id=worker_id)
def run_image_generation_task(
task: ImageGenerationTask,
*,
worker_id: str | None = None,
) -> ImageGenerationTask:
task = refresh_task(task)
if task.status != ImageGenerationTask.Status.RUNNING:
return task
if worker_id:
heartbeat_image_task(task.pk, worker_id)
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)
except Exception as exc:
refund_task_call(task, str(exc))
return mark_task_failed_if_running(
task.pk,
"upstream_error",
TASK_FAILURE_MESSAGE,
)
return mark_task_succeeded_if_running(task.pk, result.image_url)
def precharged_generation_for_task(task: ImageGenerationTask) -> PrechargedGeneration:
payload = dict(task.request_payload or {})
prepared = prepare_generation(
GenerationInput(
user=task.user,
api_key=task.api_key,
operation_type=CallRecord.OperationType.IMAGE,
prompt=str(payload.get("prompt") or ""),
alias=str(payload.get("model") or "") or None,
resolution=str(payload.get("resolution") or "1K"),
parameters=dict(payload.get("parameters") or {}),
aspect_ratio=str(payload.get("aspect_ratio") or "1:1"),
)
)
image_input = stored_image_input_for_task(task)
if image_input is not None:
prepared = replace(prepared, image_input=image_input)
return PrechargedGeneration(
prepared=prepared,
call_record=task.call_record,
points_cost=task.call_record.points_cost,
points_balance_after_charge=task.points_balance_after_charge,
)
def stored_image_input_for_task(task: ImageGenerationTask) -> ImageInput | None:
if not task.input_image:
return None
payload = dict(task.request_payload or {})
image_info = dict(payload.get("image_input") or {})
with task.input_image.open("rb") as image_file:
image = image_file.read()
return ImageInput(
data=image,
mime_type=str(image_info.get("mime_type") or "image/png"),
filename=str(image_info.get("filename") or "input.png"),
)
def heartbeat_image_task(task_pk: int, worker_id: str | None = None) -> None:
now = timezone.now()
ImageGenerationTask.objects.filter(
pk=task_pk,
status=ImageGenerationTask.Status.RUNNING,
).update(
heartbeat_at=now,
lease_expires_at=now + timedelta(seconds=image_task_lease_seconds()),
worker_id=normalize_worker_id(worker_id),
)
def mark_task_succeeded_if_running(
task_pk: int,
result_url: str,
) -> ImageGenerationTask:
now = timezone.now()
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.SUCCEEDED
task.result_url = str(result_url or "")
task.error_code = ""
task.error_message = ""
task.finished_at = now
task.heartbeat_at = now
task.save(
update_fields=(
"status",
"result_url",
"error_code",
"error_message",
"finished_at",
"heartbeat_at",
"updated_at",
)
)
return task
def mark_task_failed_if_running(
task_pk: int,
error_code: str,
error_message: str,
) -> ImageGenerationTask:
now = timezone.now()
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.FAILED
task.error_code = str(error_code or "upstream_error")
task.error_message = str(error_message or TASK_FAILURE_MESSAGE)
task.finished_at = now
task.heartbeat_at = now
task.save(
update_fields=(
"status",
"error_code",
"error_message",
"finished_at",
"heartbeat_at",
"updated_at",
)
)
return task
def reap_stale_image_tasks(
*,
now=None,
limit: int = 100,
) -> int:
current_time = now or timezone.now()
stale_ids = list(
ImageGenerationTask.objects.filter(status=ImageGenerationTask.Status.RUNNING)
.filter(stale_task_filter(current_time))
.order_by("lease_expires_at", "id")
.values_list("id", flat=True)[:limit]
)
reaped = 0
for task_pk in stale_ids:
if reap_stale_image_task(task_pk, current_time):
reaped += 1
return reaped
def reap_stale_image_task(task_pk: int, now) -> bool:
with transaction.atomic():
task = (
ImageGenerationTask.objects.select_for_update()
.select_related("call_record")
.get(pk=task_pk)
)
if task.status != ImageGenerationTask.Status.RUNNING or not task_is_stale(task, now):
return False
if task.call_record.status == CallRecord.Status.SUCCESS:
task.status = ImageGenerationTask.Status.SUCCEEDED
task.result_url = task.call_record.result_ref
task.error_code = ""
task.error_message = ""
task.finished_at = now
task.save(
update_fields=(
"status",
"result_url",
"error_code",
"error_message",
"finished_at",
"updated_at",
)
)
return True
try:
refund_call_points(
task.call_record,
error_message=TASK_TIMEOUT_MESSAGE,
reason="Image generation task lease expired; refund precharged points.",
)
except InvalidCallStateError:
task.call_record.refresh_from_db()
if task.call_record.status == CallRecord.Status.SUCCESS:
task.status = ImageGenerationTask.Status.SUCCEEDED
task.result_url = task.call_record.result_ref
task.finished_at = now
task.save(update_fields=("status", "result_url", "finished_at", "updated_at"))
return True
raise
task.status = ImageGenerationTask.Status.FAILED
task.error_code = "task_timeout"
task.error_message = TASK_TIMEOUT_MESSAGE
task.finished_at = now
task.heartbeat_at = now
task.save(
update_fields=(
"status",
"error_code",
"error_message",
"finished_at",
"heartbeat_at",
"updated_at",
)
)
return True
def stale_task_filter(now):
stale_heartbeat_before = now - timedelta(seconds=image_task_lease_seconds())
return (
models_q("lease_expires_at__lte", now)
| models_q("heartbeat_at__lte", stale_heartbeat_before)
)
def task_is_stale(task: ImageGenerationTask, now) -> bool:
stale_heartbeat_before = now - timedelta(seconds=image_task_lease_seconds())
if task.lease_expires_at and task.lease_expires_at <= now:
return True
return bool(task.heartbeat_at and task.heartbeat_at <= stale_heartbeat_before)
def refund_task_call(task: ImageGenerationTask, error_message: str) -> None:
try:
refund_call_points(
task.call_record,
error_message=error_message,
reason="Image generation task failed; refund precharged points.",
)
except InvalidCallStateError:
return
def refresh_task(task: ImageGenerationTask) -> ImageGenerationTask:
return (
ImageGenerationTask.objects.select_related("user", "api_key", "call_record")
.get(pk=task.pk)
)
def task_submit_response(task: ImageGenerationTask) -> dict[str, Any]:
call_record = task.call_record
return {
"task_id": str(task.task_id),
"status": task.status,
"call_id": call_record.id,
"points_cost": call_record.points_cost,
"points_balance": task.points_balance_after_charge,
"created_at": task.created_at.isoformat(),
"expires_at": task.expires_at.isoformat() if task.expires_at else None,
}
def task_detail_response(task: ImageGenerationTask) -> dict[str, Any]:
call_record = task.call_record
data: dict[str, Any] = {
"task_id": str(task.task_id),
"status": task.status,
"call_id": call_record.id,
"points_cost": call_record.points_cost,
"created_at": task.created_at.isoformat(),
"updated_at": task.updated_at.isoformat(),
"expires_at": task.expires_at.isoformat() if task.expires_at else None,
}
if task.status == ImageGenerationTask.Status.SUCCEEDED:
data["result"] = {"image_url": task.result_url}
elif task.status in {ImageGenerationTask.Status.FAILED, ImageGenerationTask.Status.EXPIRED}:
data["error"] = {
"code": task.error_code or "upstream_error",
"message": task.error_message or TASK_FAILURE_MESSAGE,
}
return data
def media_public_url_builder(url: str) -> str:
if url.startswith(("http://", "https://")):
return url
base_url = str(getattr(settings, "MEDIA_PUBLIC_BASE_URL", "") or "").rstrip("/")
if base_url and url.startswith("/"):
return f"{base_url}{url}"
return url
def image_task_retention_hours() -> int:
return max(1, int(getattr(settings, "IMAGE_TASK_RETENTION_HOURS", 24)))
def image_task_lease_seconds() -> int:
return max(1, int(getattr(settings, "IMAGE_TASK_LEASE_SECONDS", 600)))
def normalize_worker_id(worker_id: str | None) -> str:
normalized = str(worker_id or "").strip()
if normalized:
return normalized[:128]
return f"{socket.gethostname()}:{hash_text(str(timezone.now().timestamp()))[:8]}"
def hash_text(value: str) -> str:
return hashlib.sha256(value.encode("utf-8")).hexdigest()
def hash_json(value: Mapping[str, Any]) -> str:
payload = json.dumps(
value,
sort_keys=True,
separators=(",", ":"),
ensure_ascii=False,
default=str,
)
return hash_text(payload)
def models_q(key: str, value):
from django.db.models import Q
return Q(**{key: value})
+1
View File
@@ -0,0 +1 @@
+1
View File
@@ -0,0 +1 @@
@@ -0,0 +1,66 @@
import time
import uuid
from django.conf import settings
from django.core.management.base import BaseCommand
from apps.api.image_tasks import reap_stale_image_tasks, run_one_image_task
class Command(BaseCommand):
help = "Run asynchronous image generation tasks."
def add_arguments(self, parser):
parser.add_argument(
"--once",
action="store_true",
help="Process at most one queued task and exit.",
)
parser.add_argument(
"--worker-id",
default="",
help="Stable worker identifier recorded on claimed tasks.",
)
parser.add_argument(
"--sleep-seconds",
type=float,
default=1.0,
help="Sleep interval when no queued task is available.",
)
parser.add_argument(
"--reap-only",
action="store_true",
help="Only reap stale running tasks and exit.",
)
def handle(self, *args, **options):
worker_id = options["worker_id"] or f"image-worker-{uuid.uuid4().hex[:8]}"
once = bool(options["once"])
reap_only = bool(options["reap_only"])
sleep_seconds = max(0.1, float(options["sleep_seconds"]))
reaper_interval = max(
1,
int(getattr(settings, "IMAGE_TASK_REAPER_INTERVAL_SECONDS", 60)),
)
last_reap_at = 0.0
while True:
now = time.monotonic()
if reap_only or now - last_reap_at >= reaper_interval:
reaped = reap_stale_image_tasks()
if reaped:
self.stdout.write(f"reaped={reaped}")
last_reap_at = now
if reap_only:
return
task = run_one_image_task(worker_id=worker_id)
if task is not None:
self.stdout.write(f"task={task.task_id} status={task.status}")
if once:
return
continue
if once:
return
time.sleep(sleep_seconds)
+58
View File
@@ -0,0 +1,58 @@
# Generated by Django 5.2.15 on 2026-07-08 13:44
import django.db.models.deletion
import uuid
from django.conf import settings
from django.db import migrations, models
class Migration(migrations.Migration):
initial = True
dependencies = [
('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.CreateModel(
name='ImageGenerationTask',
fields=[
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
('task_id', models.UUIDField(default=uuid.uuid4, editable=False, unique=True, verbose_name='任务 ID')),
('status', models.CharField(choices=[('queued', '排队中'), ('running', '处理中'), ('succeeded', '成功'), ('failed', '失败'), ('expired', '已过期')], default='queued', max_length=20, verbose_name='状态')),
('idempotency_key', models.CharField(blank=True, max_length=128, verbose_name='幂等键')),
('idempotency_key_hash', models.CharField(blank=True, editable=False, max_length=64, null=True, verbose_name='幂等键哈希')),
('request_hash', models.CharField(max_length=64, verbose_name='请求哈希')),
('request_payload', models.JSONField(blank=True, default=dict, verbose_name='请求快照')),
('input_image', models.FileField(blank=True, upload_to='generated/task_inputs/%Y/%m/%d/', verbose_name='输入图片')),
('result_url', models.TextField(blank=True, verbose_name='结果 URL')),
('error_code', models.CharField(blank=True, max_length=64, verbose_name='错误码')),
('error_message', models.TextField(blank=True, verbose_name='错误信息')),
('points_balance_after_charge', models.BigIntegerField(default=0, verbose_name='扣费后余额')),
('started_at', models.DateTimeField(blank=True, null=True, verbose_name='开始时间')),
('finished_at', models.DateTimeField(blank=True, null=True, verbose_name='完成时间')),
('expires_at', models.DateTimeField(blank=True, null=True, verbose_name='任务元数据过期时间')),
('locked_at', models.DateTimeField(blank=True, null=True, verbose_name='锁定时间')),
('lease_expires_at', models.DateTimeField(blank=True, null=True, verbose_name='租约过期时间')),
('heartbeat_at', models.DateTimeField(blank=True, null=True, verbose_name='心跳时间')),
('worker_id', models.CharField(blank=True, max_length=128, verbose_name='Worker ID')),
('attempt_count', models.PositiveIntegerField(default=0, verbose_name='尝试次数')),
('created_at', models.DateTimeField(auto_now_add=True, verbose_name='创建时间')),
('updated_at', models.DateTimeField(auto_now=True, verbose_name='更新时间')),
('api_key', models.ForeignKey(on_delete=django.db.models.deletion.PROTECT, related_name='image_generation_tasks', to='users.apikey', verbose_name='API 密钥')),
('call_record', models.OneToOneField(on_delete=django.db.models.deletion.PROTECT, related_name='image_generation_task', to='billing.callrecord', verbose_name='调用记录')),
('user', models.ForeignKey(on_delete=django.db.models.deletion.PROTECT, related_name='image_generation_tasks', to=settings.AUTH_USER_MODEL, verbose_name='用户')),
],
options={
'verbose_name': '图片生成任务',
'verbose_name_plural': '图片生成任务',
'db_table': 'image_generation_task',
'ordering': ('-created_at', '-id'),
'indexes': [models.Index(fields=['user', 'created_at'], name='image_gener_user_id_1e7164_idx'), models.Index(fields=['api_key', 'created_at'], name='image_gener_api_key_30d9ad_idx'), models.Index(fields=['status', 'created_at'], name='image_gener_status_727db2_idx'), models.Index(fields=['status', 'lease_expires_at'], name='image_gener_status_7a7bd1_idx'), models.Index(fields=['worker_id', 'status'], name='image_gener_worker__f9e534_idx'), models.Index(fields=['expires_at'], name='image_gener_expires_545a78_idx')],
'constraints': [models.UniqueConstraint(fields=('api_key', 'idempotency_key_hash'), name='unique_image_task_idempotency_hash_per_api_key'), models.CheckConstraint(condition=models.Q(('points_balance_after_charge__gte', 0)), name='image_task_balance_after_charge_non_negative')],
},
),
]
+97 -2
View File
@@ -1,3 +1,98 @@
from django.db import models
from __future__ import annotations
# Create your models here.
import uuid
from django.conf import settings
from django.db import models
from django.db.models import Q
class ImageGenerationTask(models.Model):
class Status(models.TextChoices):
QUEUED = "queued", "排队中"
RUNNING = "running", "处理中"
SUCCEEDED = "succeeded", "成功"
FAILED = "failed", "失败"
EXPIRED = "expired", "已过期"
task_id = models.UUIDField("任务 ID", default=uuid.uuid4, unique=True, editable=False)
user = models.ForeignKey(
settings.AUTH_USER_MODEL,
verbose_name="用户",
on_delete=models.PROTECT,
related_name="image_generation_tasks",
)
api_key = models.ForeignKey(
"users.ApiKey",
verbose_name="API 密钥",
on_delete=models.PROTECT,
related_name="image_generation_tasks",
)
call_record = models.OneToOneField(
"billing.CallRecord",
verbose_name="调用记录",
on_delete=models.PROTECT,
related_name="image_generation_task",
)
status = models.CharField(
"状态",
max_length=20,
choices=Status.choices,
default=Status.QUEUED,
)
idempotency_key = models.CharField("幂等键", max_length=128, blank=True)
idempotency_key_hash = models.CharField(
"幂等键哈希",
max_length=64,
null=True,
blank=True,
editable=False,
)
request_hash = models.CharField("请求哈希", max_length=64)
request_payload = models.JSONField("请求快照", default=dict, blank=True)
input_image = models.FileField(
"输入图片",
upload_to="generated/task_inputs/%Y/%m/%d/",
blank=True,
)
result_url = models.TextField("结果 URL", blank=True)
error_code = models.CharField("错误码", max_length=64, blank=True)
error_message = models.TextField("错误信息", blank=True)
points_balance_after_charge = models.BigIntegerField("扣费后余额", default=0)
started_at = models.DateTimeField("开始时间", null=True, blank=True)
finished_at = models.DateTimeField("完成时间", null=True, blank=True)
expires_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)
worker_id = models.CharField("Worker ID", max_length=128, blank=True)
attempt_count = models.PositiveIntegerField("尝试次数", default=0)
created_at = models.DateTimeField("创建时间", auto_now_add=True)
updated_at = models.DateTimeField("更新时间", auto_now=True)
class Meta:
db_table = "image_generation_task"
verbose_name = "图片生成任务"
verbose_name_plural = "图片生成任务"
ordering = ("-created_at", "-id")
constraints = [
models.UniqueConstraint(
fields=("api_key", "idempotency_key_hash"),
name="unique_image_task_idempotency_hash_per_api_key",
),
models.CheckConstraint(
condition=Q(points_balance_after_charge__gte=0),
name="image_task_balance_after_charge_non_negative",
),
]
indexes = [
models.Index(fields=("user", "created_at")),
models.Index(fields=("api_key", "created_at")),
models.Index(fields=("status", "created_at")),
models.Index(fields=("status", "lease_expires_at")),
models.Index(fields=("worker_id", "status")),
models.Index(fields=("expires_at",)),
]
def __str__(self) -> str:
return f"{self.task_id} {self.status}"
+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")
+12
View File
@@ -4,6 +4,8 @@ from .views import (
AlipayRechargeCallbackView,
BalanceView,
ClientLatestReleaseView,
GenerateImageTaskDetailView,
GenerateImageTaskSubmitView,
GenerateImageView,
GenerateTitleView,
ModelsView,
@@ -22,6 +24,16 @@ urlpatterns = [
),
path("v1/generate/title", GenerateTitleView.as_view(), name="api-generate-title"),
path("v1/generate/image", GenerateImageView.as_view(), name="api-generate-image"),
path(
"v1/generate/image/tasks",
GenerateImageTaskSubmitView.as_view(),
name="api-generate-image-task-submit",
),
path(
"v1/generate/image/tasks/<uuid:task_id>",
GenerateImageTaskDetailView.as_view(),
name="api-generate-image-task-detail",
),
path("v1/recharge/create", RechargeCreateView.as_view(), name="api-recharge-create"),
path("v1/recharge/status", RechargeStatusView.as_view(), name="api-recharge-status"),
path(
+43
View File
@@ -18,6 +18,12 @@ from apps.api.generation import (
generate_image_response,
generate_title_response,
)
from apps.api.image_tasks import (
create_image_generation_task,
task_detail_response,
task_submit_response,
)
from apps.api.models import ImageGenerationTask
from apps.api.serializers import (
GenerateImageRequestSerializer,
GenerateTitleRequestSerializer,
@@ -107,6 +113,43 @@ class GenerateImageView(ExternalApiView):
return Response(data, status=status.HTTP_200_OK)
class GenerateImageTaskSubmitView(ExternalApiView):
throttle_classes = (GenerateRateThrottle,)
def post(self, request):
serializer = GenerateImageRequestSerializer(data=request.data)
if not serializer.is_valid():
return Response(
api_error("bad_request", "参数错误"),
status=status.HTTP_400_BAD_REQUEST,
)
try:
task, _created = create_image_generation_task(
user=request.user,
api_key=request.auth,
request_data=serializer.validated_data,
idempotency_key=request.headers.get("Idempotency-Key", ""),
)
except ApiRequestError as exc:
return Response(exc.as_response_data(), status=exc.http_status)
return Response(task_submit_response(task), status=status.HTTP_202_ACCEPTED)
class GenerateImageTaskDetailView(ExternalApiView):
def get(self, request, task_id):
task = (
ImageGenerationTask.objects.select_related("call_record")
.filter(task_id=task_id, user=request.user)
.first()
)
if task is None:
return Response(
api_error("task_not_found", "图片生成任务不存在"),
status=status.HTTP_404_NOT_FOUND,
)
return Response(task_detail_response(task), status=status.HTTP_200_OK)
class BalanceView(ExternalApiView):
def get(self, request):
balance = get_balance_snapshot(request.user)