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
+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})