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 = "图片生成任务超时,已退回点数" RETRYABLE_TASK_ERROR_CODES = {"upstream_timeout", "upstream_error"} 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) .filter( models_q("next_attempt_at__isnull", True) | models_q("next_attempt_at__lte", now) ) .order_by("created_at", "id") ) if connection.features.has_select_for_update_skip_locked: 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.next_attempt_at = None task.attempt_count += 1 task.save( update_fields=( "status", "worker_id", "locked_at", "lease_expires_at", "heartbeat_at", "started_at", "next_attempt_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) except ApiRequestError as exc: refund_task_call(task, exc.message) return mark_task_failed_if_running(task.pk, exc.code, exc.message) try: result = execute_precharged_generation( precharged, image_url_builder=media_public_url_builder, refund_on_failure=False, ) except ApiRequestError as exc: return handle_task_generation_failure( task, error_code=exc.code, error_message=exc.message, retryable=is_retryable_task_error(exc.code), ) except Exception as exc: return handle_task_generation_failure( task, error_code="upstream_error", error_message=str(exc) or TASK_FAILURE_MESSAGE, retryable=True, ) return mark_task_succeeded_if_running(task.pk, result.image_url) def handle_task_generation_failure( task: ImageGenerationTask, *, error_code: str, error_message: str, retryable: bool, ) -> ImageGenerationTask: normalized_code = str(error_code or "upstream_error") normalized_message = str(error_message or TASK_FAILURE_MESSAGE) if retryable and task.attempt_count < image_task_max_attempts(): return mark_task_retry_if_running( task.pk, normalized_code, retry_error_message(normalized_code, normalized_message), next_attempt_at=timezone.now() + timedelta(seconds=image_task_retry_backoff_seconds(task.attempt_count)), ) refund_task_call(task, normalized_message) return mark_task_failed_if_running(task.pk, normalized_code, normalized_message) def precharged_generation_for_task(task: ImageGenerationTask) -> PrechargedGeneration: payload = dict(task.request_payload or {}) prepared = prepare_generation( 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.next_attempt_at = None task.finished_at = now task.heartbeat_at = now task.save( update_fields=( "status", "result_url", "error_code", "error_message", "next_attempt_at", "finished_at", "heartbeat_at", "updated_at", ) ) return task def mark_task_retry_if_running( task_pk: int, error_code: str, error_message: str, *, next_attempt_at, ) -> ImageGenerationTask: with transaction.atomic(): task = ImageGenerationTask.objects.select_for_update().get(pk=task_pk) if task.status != ImageGenerationTask.Status.RUNNING: return task task.status = ImageGenerationTask.Status.QUEUED task.error_code = str(error_code or "upstream_error") task.error_message = str(error_message or TASK_FAILURE_MESSAGE) task.next_attempt_at = next_attempt_at task.worker_id = "" task.locked_at = None task.lease_expires_at = None task.heartbeat_at = None task.save( update_fields=( "status", "error_code", "error_message", "next_attempt_at", "worker_id", "locked_at", "lease_expires_at", "heartbeat_at", "updated_at", ) ) return task def mark_task_failed_if_running( task_pk: int, error_code: str, 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.next_attempt_at = None task.finished_at = now task.heartbeat_at = now task.save( update_fields=( "status", "error_code", "error_message", "next_attempt_at", "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.next_attempt_at = None task.finished_at = now task.save( update_fields=( "status", "result_url", "error_code", "error_message", "next_attempt_at", "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.next_attempt_at = None task.finished_at = now task.save( update_fields=( "status", "result_url", "next_attempt_at", "finished_at", "updated_at", ) ) return True raise task.status = ImageGenerationTask.Status.FAILED task.error_code = "task_timeout" task.error_message = TASK_TIMEOUT_MESSAGE task.next_attempt_at = None task.finished_at = now task.heartbeat_at = now task.save( update_fields=( "status", "error_code", "error_message", "next_attempt_at", "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, "attempt_count": task.attempt_count, "max_attempts": image_task_max_attempts(), "next_attempt_at": task.next_attempt_at.isoformat() if task.next_attempt_at else None, "created_at": task.created_at.isoformat(), "expires_at": task.expires_at.isoformat() if task.expires_at else None, } 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, "attempt_count": task.attempt_count, "max_attempts": image_task_max_attempts(), "next_attempt_at": task.next_attempt_at.isoformat() if task.next_attempt_at else None, "created_at": task.created_at.isoformat(), "updated_at": task.updated_at.isoformat(), "expires_at": task.expires_at.isoformat() if task.expires_at else None, } 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 image_task_max_retries() -> int: return max(0, int(getattr(settings, "IMAGE_TASK_MAX_RETRIES", 2))) def image_task_max_attempts() -> int: return 1 + image_task_max_retries() def image_task_retry_backoff_seconds(attempt_count: int) -> int: values = image_task_retry_backoff_values() if not values: return 0 index = max(0, int(attempt_count or 1) - 1) return values[min(index, len(values) - 1)] def image_task_retry_backoff_values() -> list[int]: raw = str(getattr(settings, "IMAGE_TASK_RETRY_BACKOFF_SECONDS", "10,30") or "") values: list[int] = [] for part in raw.split(","): item = part.strip() if not item: continue try: values.append(max(0, int(item))) except ValueError: continue return values def is_retryable_task_error(error_code: str) -> bool: return str(error_code or "") in RETRYABLE_TASK_ERROR_CODES def retry_error_message(error_code: str, fallback: str) -> str: if error_code == "upstream_timeout": return "上游 AI 调用超时,稍后自动重试" if error_code == "upstream_error": return "上游 AI 调用失败,稍后自动重试" return fallback def normalize_worker_id(worker_id: str | None) -> str: normalized = str(worker_id or "").strip() if normalized: 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})