feat: add async image task API
This commit is contained in:
@@ -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})
|
||||
Reference in New Issue
Block a user