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