From 06cdab9cbc104d3965c9088b6a816ae91e0829c0 Mon Sep 17 00:00:00 2001 From: GinHa <15201596918@163.com> Date: Fri, 26 Jun 2026 11:51:02 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E7=94=9F=E8=BE=B0=E4=BB=BB?= =?UTF-8?q?=E5=8A=A1=E6=81=A2=E5=A4=8D=E6=9C=BA=E5=88=B6BUG|=E4=BF=AE?= =?UTF-8?q?=E5=A4=8D=E7=94=9F=E6=88=90=E8=A7=86=E9=A2=91/=E5=9B=BE?= =?UTF-8?q?=E7=89=87token=E5=BF=AB=E7=85=A7=E4=B8=8D=E5=9B=9E=E8=90=BD?= =?UTF-8?q?=E4=BA=A4=E6=98=93=E6=B5=81=E6=B0=B4BUG?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- video-gen-api/app/api/v1/generation.py | 2 + video-gen-api/app/config.py | 9 + video-gen-api/app/enums/__init__.py | 1 + video-gen-api/app/enums/generation_task.py | 101 +++++ .../generation_module_hook_service.py | 22 +- .../services/generation_recovery_service.py | 72 ++-- .../media_token_usage_snapshot_service.py | 253 ++++++++++++ video-gen-api/app/services/video_queue.py | 3 + .../app/tasks/generation_create_tasks.py | 2 + .../app/tasks/generation_download_tasks.py | 382 ++++++++++++++---- .../app/tasks/generation_poll_tasks.py | 2 + .../app/tasks/generation_recovery_tasks.py | 49 ++- 12 files changed, 773 insertions(+), 125 deletions(-) create mode 100644 video-gen-api/app/enums/generation_task.py create mode 100644 video-gen-api/app/services/media_token_usage_snapshot_service.py diff --git a/video-gen-api/app/api/v1/generation.py b/video-gen-api/app/api/v1/generation.py index a4c40c9f..a2410d30 100644 --- a/video-gen-api/app/api/v1/generation.py +++ b/video-gen-api/app/api/v1/generation.py @@ -43,6 +43,7 @@ from app.services.generation_billing_service import ( get_next_credit_attempt_no, ) from app.services.generation_refund_service import mark_generation_record_failed_and_refund_once +from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot from app.services.credit_record_meta_service import build_generation_record_prompt_meta from app.services.video_cover_service import async_create_video_cover_for_local_video from app.utils.id_gen import generate_id @@ -724,6 +725,7 @@ async def seedance_callback(request: Request, db: AsyncSession = Depends(get_db) usage = data.get("usage", {}) if usage: record.video_tokens_used = usage.get("total_tokens", 0) + await sync_generation_record_media_token_snapshot(db, record, provider_response=data) # Log callback response from app.services.video_gen import _log_video_response _log_video_response(record.id, data) diff --git a/video-gen-api/app/config.py b/video-gen-api/app/config.py index f126cfcf..1396d051 100644 --- a/video-gen-api/app/config.py +++ b/video-gen-api/app/config.py @@ -127,6 +127,15 @@ class Settings(BaseSettings): DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS: int = 5 * 60 DOWNLOAD_RECOVERY_BATCH_SIZE: int = 100 DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS: int = 3 + # 下载恢复自循环:不依赖 Celery beat,不新增 worker;由 gen_result_download 队列周期扫描 DB/Redis。 + DOWNLOAD_RECOVERY_LOOP_ENABLED: bool = True + DOWNLOAD_RECOVERY_INTERVAL_SECONDS: int = 60 + DOWNLOAD_RECOVERY_LOOP_LOCK_KEY: str = "vg:celery:download_recovery_loop_lock" + DOWNLOAD_RECOVERY_LOOP_LOCK_TTL_SECONDS: int = 55 + DOWNLOAD_RETRY_COUNTDOWN_EXTRA_SECONDS: int = 1 + DOWNLOAD_NON_RETRYABLE_LOCAL_ERRORS: bool = True + DOWNLOAD_EVENT_VERBOSE_ENABLED: bool = True + MEDIA_TOKEN_SNAPSHOT_ENABLED: bool = True DOWNLOAD_ACTIVE_REDIS_HASH_KEY: str = "vg:celery:download:active" DOWNLOAD_ACTIVE_REDIS_ZSET_KEY: str = "vg:celery:download:active_index" diff --git a/video-gen-api/app/enums/__init__.py b/video-gen-api/app/enums/__init__.py index a116920b..730424e5 100644 --- a/video-gen-api/app/enums/__init__.py +++ b/video-gen-api/app/enums/__init__.py @@ -6,3 +6,4 @@ from app.enums.video_prompt_schema import * from app.enums.user import * from app.enums.credit_record import * from app.enums.token_usage import * +from app.enums.generation_task import * diff --git a/video-gen-api/app/enums/generation_task.py b/video-gen-api/app/enums/generation_task.py new file mode 100644 index 00000000..9a264566 --- /dev/null +++ b/video-gen-api/app/enums/generation_task.py @@ -0,0 +1,101 @@ +from enum import Enum + + +class GenerationMode(str, Enum): + CHATAPI_ASYNC = "chatapi_async" + HOT_OPENING_REPLICATE = "hot_opening_replicate" + SHOT_REPLICATE = "shot_replicate" + + +class GenerationType(str, Enum): + IMAGE = "image" + VIDEO = "video" + + +class ChatGenerationTaskStatus(str, Enum): + PENDING = "pending" + GENERATING = "generating" + COMPLETED = "completed" + FAILED = "failed" + + +class ChatGenerationPipelineStage(str, Enum): + QUEUED = "queued" + PREPARING = "preparing" + CREATING_PROVIDER_TASK = "creating_provider_task" + WAITING_REMOTE = "waiting_remote" + POLLING = "polling" + RESULT_READY = "result_ready" + DOWNLOAD_QUEUED = "download_queued" + DOWNLOADING = "downloading" + RETRY_WAITING = "retry_waiting" + DONE = "done" + FAILED = "failed" + TIMEOUT = "timeout" + DOWNLOAD_FAILED = "download_failed" + + +class ChatGenerationTaskEventType(str, Enum): + PROMPT_CONCAT_START = "PROMPT_CONCAT_START" + PROMPT_CONCAT_SUCCESS = "PROMPT_CONCAT_SUCCESS" + + PROVIDER_CREATE_START = "PROVIDER_CREATE_START" + PROVIDER_CREATE_SUCCESS = "PROVIDER_CREATE_SUCCESS" + PROVIDER_CREATE_FAILED = "PROVIDER_CREATE_FAILED" + + POLL_START = "POLL_START" + POLL_PENDING = "POLL_PENDING" + POLL_SUCCESS = "POLL_SUCCESS" + POLL_FAILED = "POLL_FAILED" + POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY = "POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY" + FINAL_POLL_BEFORE_TIMEOUT_ERROR = "FINAL_POLL_BEFORE_TIMEOUT_ERROR" + FINAL_POLL_BEFORE_TIMEOUT_PENDING = "FINAL_POLL_BEFORE_TIMEOUT_PENDING" + GENERATION_RECOVERY_ENQUEUE = "GENERATION_RECOVERY_ENQUEUE" + + DOWNLOAD_ENQUEUE = "DOWNLOAD_ENQUEUE" + DOWNLOAD_ENQUEUE_FAILED = "DOWNLOAD_ENQUEUE_FAILED" + DOWNLOAD_START = "DOWNLOAD_START" + DOWNLOAD_SUCCESS = "DOWNLOAD_SUCCESS" + DOWNLOAD_RETRY_WAITING = "DOWNLOAD_RETRY_WAITING" + DOWNLOAD_RETRY_ENQUEUE = "DOWNLOAD_RETRY_ENQUEUE" + DOWNLOAD_RETRY_ENQUEUE_FAILED = "DOWNLOAD_RETRY_ENQUEUE_FAILED" + DOWNLOAD_RETRY_NOT_DUE = "DOWNLOAD_RETRY_NOT_DUE" + DOWNLOAD_STUCK_RECOVER = "DOWNLOAD_STUCK_RECOVER" + DOWNLOAD_RECOVERY_ENQUEUE = "DOWNLOAD_RECOVERY_ENQUEUE" + DOWNLOAD_RECOVERY_ENQUEUE_FAILED = "DOWNLOAD_RECOVERY_ENQUEUE_FAILED" + DOWNLOAD_FAILED = "DOWNLOAD_FAILED" + DOWNLOAD_FAILED_NON_RETRYABLE = "DOWNLOAD_FAILED_NON_RETRYABLE" + + DOWNLOAD_SKIP_TASK_MISSING = "DOWNLOAD_SKIP_TASK_MISSING" + DOWNLOAD_SKIP_INVALID_MODE = "DOWNLOAD_SKIP_INVALID_MODE" + DOWNLOAD_SKIP_NOT_GENERATING = "DOWNLOAD_SKIP_NOT_GENERATING" + DOWNLOAD_SKIP_ALREADY_COMPLETED = "DOWNLOAD_SKIP_ALREADY_COMPLETED" + DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL = "DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL" + DOWNLOAD_SKIP_STAGE_NOT_ALLOWED = "DOWNLOAD_SKIP_STAGE_NOT_ALLOWED" + DOWNLOAD_SKIP_DOWNLOADING_LEASE_ALIVE = "DOWNLOAD_SKIP_DOWNLOADING_LEASE_ALIVE" + DOWNLOAD_SKIP_RETRY_NOT_DUE = "DOWNLOAD_SKIP_RETRY_NOT_DUE" + DOWNLOAD_SKIP_FINAL_STATE = "DOWNLOAD_SKIP_FINAL_STATE" + DOWNLOAD_SKIP_DISABLED = "DOWNLOAD_SKIP_DISABLED" + + TASK_TIMEOUT = "TASK_TIMEOUT" + + +ALLOWED_GENERATION_MODES = { + GenerationMode.CHATAPI_ASYNC.value, + GenerationMode.HOT_OPENING_REPLICATE.value, + GenerationMode.SHOT_REPLICATE.value, +} + +FINAL_CHAT_GENERATION_STAGES = { + ChatGenerationPipelineStage.DONE.value, + ChatGenerationPipelineStage.FAILED.value, + ChatGenerationPipelineStage.TIMEOUT.value, + ChatGenerationPipelineStage.DOWNLOAD_FAILED.value, +} + +DOWNLOAD_RECOVERABLE_STAGES = { + ChatGenerationPipelineStage.RESULT_READY.value, + ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value, + ChatGenerationPipelineStage.DOWNLOADING.value, + ChatGenerationPipelineStage.RETRY_WAITING.value, +} diff --git a/video-gen-api/app/services/generation_module_hook_service.py b/video-gen-api/app/services/generation_module_hook_service.py index 9a518725..c7f5b2bd 100644 --- a/video-gen-api/app/services/generation_module_hook_service.py +++ b/video-gen-api/app/services/generation_module_hook_service.py @@ -2,36 +2,40 @@ from __future__ import annotations from sqlalchemy.ext.asyncio import AsyncSession +from app.enums.generation_task import ChatGenerationTaskStatus, GenerationMode from app.models.chat_generation_task import ChatGenerationTask async def notify_chat_generation_task_finished(db: AsyncSession, task: ChatGenerationTask) -> None: """通知业务模块 ChatGenerationTask 已进入终态。 - 当前用于爆款开头复刻: - - image_generate 完成后自动进入 video_prompt_optimize - - video_generate 完成后总任务完成 + 该方法必须幂等:下载恢复任务、重试任务、服务重启补偿都可能重复调用。 + 具体模块服务需要自行判断 step/project 是否已经完成或失败,避免重复推进。 """ if not task: return - if task.generation_mode == "hot_opening_replicate": + + status = getattr(task, "status", None) + generation_mode = getattr(task, "generation_mode", None) + + if generation_mode == GenerationMode.HOT_OPENING_REPLICATE.value: from app.services.hot_opening_replicate_service import ( handle_chat_generation_task_completed, handle_chat_generation_task_failed, ) - if task.status == "completed": + if status == ChatGenerationTaskStatus.COMPLETED.value: await handle_chat_generation_task_completed(db, task) - elif task.status == "failed": + elif status == ChatGenerationTaskStatus.FAILED.value: await handle_chat_generation_task_failed(db, task) return - if task.generation_mode == "shot_replicate": + if generation_mode == GenerationMode.SHOT_REPLICATE.value: from app.services.shot_replicate_flow_service import ( handle_chat_generation_task_completed, handle_chat_generation_task_failed, ) - if task.status == "completed": + if status == ChatGenerationTaskStatus.COMPLETED.value: await handle_chat_generation_task_completed(db, task) - elif task.status == "failed": + elif status == ChatGenerationTaskStatus.FAILED.value: await handle_chat_generation_task_failed(db, task) return diff --git a/video-gen-api/app/services/generation_recovery_service.py b/video-gen-api/app/services/generation_recovery_service.py index ea9ebe18..0ce6d132 100644 --- a/video-gen-api/app/services/generation_recovery_service.py +++ b/video-gen-api/app/services/generation_recovery_service.py @@ -9,6 +9,12 @@ from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings +from app.enums.generation_task import ( + ALLOWED_GENERATION_MODES, + ChatGenerationPipelineStage, + ChatGenerationTaskEventType, + ChatGenerationTaskStatus, +) from app.models.chat_generation_task import ChatGenerationTask from app.services.celery_download_recovery_service import ( ensure_aware_utc, @@ -19,6 +25,7 @@ from app.services.celery_download_recovery_service import ( ) from app.services.generation_log_service import log_provider_call, log_task_event from app.services.generation_module_hook_service import notify_chat_generation_task_finished +from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot from app.services.generation_provider_service import poll_provider_task from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once from app.services.redis_registry_service import ( @@ -30,7 +37,6 @@ from app.services.redis_registry_service import ( logger = logging.getLogger("video_gen") -ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"} POLL_QUEUE = "gen_provider_poll" @@ -59,11 +65,11 @@ def _is_queue_timeout(task: ChatGenerationTask, now: datetime | None = None) -> def _is_final_task_state(task: ChatGenerationTask) -> bool: - return task.status in ("completed", "failed") or task.pipeline_stage in ( - "done", - "failed", - "timeout", - "download_failed", + return task.status in (ChatGenerationTaskStatus.COMPLETED.value, ChatGenerationTaskStatus.FAILED.value) or task.pipeline_stage in ( + ChatGenerationPipelineStage.DONE.value, + ChatGenerationPipelineStage.FAILED.value, + ChatGenerationPipelineStage.TIMEOUT.value, + ChatGenerationPipelineStage.DOWNLOAD_FAILED.value, ) @@ -137,19 +143,21 @@ async def recover_one_download_task( if _is_final_task_state(task): await remove_download_active(task.id) return "clean_final_state" - if task.status != "generating": + if task.status != ChatGenerationTaskStatus.GENERATING.value: await remove_download_active(task.id) + await log_task_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING.value, message=f"{source} 下载恢复跳过:任务不是 generating", detail={"status": task.status, "stage": task.pipeline_stage}) return "clean_not_generating" if not task.remote_result_url: + await log_task_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL.value, message=f"{source} 下载恢复跳过:缺少 remote_result_url", detail={"status": task.status, "stage": task.pipeline_stage}) return "skip_no_remote_result_url" stage = task.pipeline_stage redis_payload = payload or {} - if stage == "result_ready": + if stage == ChatGenerationPipelineStage.RESULT_READY.value: await log_task_event( task, - event_type="DOWNLOAD_RECOVERY_ENQUEUE", + event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value, message=f"{source} 发现 result_ready 未完成下载,启动时恢复投递下载任务", detail={"payload": redis_payload}, ) @@ -165,7 +173,7 @@ async def recover_one_download_task( if _is_queue_timeout(task, current_time): await log_task_event( task, - event_type="DOWNLOAD_RECOVERY_ENQUEUE", + event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value, message=f"{source} 发现 download_queued 长时间未消费,启动时恢复投递下载任务", detail={"payload": redis_payload}, ) @@ -188,7 +196,7 @@ async def recover_one_download_task( if _is_expired(task.download_lease_until, current_time): await log_task_event( task, - event_type="DOWNLOAD_RECOVERY_ENQUEUE", + event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value, message=f"{source} 发现 downloading lease 过期,启动时恢复投递下载任务", detail={"payload": redis_payload}, ) @@ -211,7 +219,7 @@ async def recover_one_download_task( if _is_expired(task.download_next_retry_at, current_time): await log_task_event( task, - event_type="DOWNLOAD_RECOVERY_ENQUEUE", + event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value, message=f"{source} 发现 retry_waiting 到期,启动时恢复投递下载任务", detail={"payload": redis_payload}, ) @@ -276,11 +284,16 @@ async def recover_download_tasks_once(db: AsyncSession) -> dict[str, Any]: select(ChatGenerationTask) .where( ChatGenerationTask.deleted_at.is_(None), - ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate", "shot_replicate"]), - ChatGenerationTask.status == "generating", + ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)), + ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value, ChatGenerationTask.remote_result_url.is_not(None), ChatGenerationTask.pipeline_stage.in_( - ["result_ready", "download_queued", "downloading", "retry_waiting"] + [ + ChatGenerationPipelineStage.RESULT_READY.value, + ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value, + ChatGenerationPipelineStage.DOWNLOADING.value, + ChatGenerationPipelineStage.RETRY_WAITING.value, + ] ), ) .order_by(ChatGenerationTask.updated_at.asc()) @@ -314,7 +327,7 @@ async def _mark_timeout( db, task=task, error_message=error_message, - pipeline_stage="timeout", + pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value, ) await notify_chat_generation_task_finished(db, task) await db.commit() @@ -323,7 +336,7 @@ async def _mark_timeout( task, event_type="TASK_TIMEOUT", to_status="failed", - to_stage="timeout", + to_stage=ChatGenerationPipelineStage.TIMEOUT.value, ) return "mark_timeout" @@ -340,7 +353,7 @@ async def _mark_failed( db, task=task, error_message=error_message, - pipeline_stage="failed", + pipeline_stage=ChatGenerationPipelineStage.FAILED.value, ) await notify_chat_generation_task_finished(db, task) await db.commit() @@ -397,6 +410,7 @@ async def _try_final_poll_before_timeout(db: AsyncSession, task: ChatGenerationT task.video_tokens_used = poll_result.get("video_tokens", 0) or 0 task.provider_response_json = response_data + await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=response_data) if not task.remote_result_url: return await _mark_failed( db, @@ -405,14 +419,14 @@ async def _try_final_poll_before_timeout(db: AsyncSession, task: ChatGenerationT detail=poll_result, ) - task.pipeline_stage = "result_ready" + task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value task.retry_count = 0 await db.commit() await _remove_poll_active(task.id) await log_task_event( task, event_type="POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY", - to_stage="result_ready", + to_stage=ChatGenerationPipelineStage.RESULT_READY.value, detail=poll_result, ) await enqueue_download_task(db, task, recover=True, reason="final_poll_before_timeout_success") @@ -463,13 +477,13 @@ async def recover_one_generation_task( return "clean_not_generating" if task.deadline_at and _is_expired(task.deadline_at, current_time): - if task.pipeline_stage in ("waiting_remote", "polling"): + if task.pipeline_stage in (ChatGenerationPipelineStage.WAITING_REMOTE.value, ChatGenerationPipelineStage.POLLING.value): return await _try_final_poll_before_timeout(db, task) return await _mark_timeout(db, task) - if task.pipeline_stage in ("queued", "preparing", "creating_provider_task"): + if task.pipeline_stage in (ChatGenerationPipelineStage.QUEUED.value, ChatGenerationPipelineStage.PREPARING.value, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value): if task.provider_task_id or task.seedance_task_id: - task.pipeline_stage = "waiting_remote" + task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value await db.commit() await log_task_event( task, @@ -498,7 +512,7 @@ async def recover_one_generation_task( ) return "recover_create" - if task.pipeline_stage in ("waiting_remote", "polling"): + if task.pipeline_stage in (ChatGenerationPipelineStage.WAITING_REMOTE.value, ChatGenerationPipelineStage.POLLING.value): if task.remote_result_url: await _remove_poll_active(task.id) await enqueue_download_task( @@ -516,7 +530,7 @@ async def recover_one_generation_task( message=f"{source} 发现远程等待/轮询阶段任务未完成,恢复投递轮询队列", detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, ) - task.pipeline_stage = "waiting_remote" + task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value await db.commit() poll_generation_task.apply_async( args=[task.id], @@ -536,7 +550,7 @@ async def recover_one_generation_task( message=f"{source} 发现任务缺少供应商任务ID,恢复投递创建队列", detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, ) - task.pipeline_stage = "queued" + task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value await db.commit() await _remove_poll_active(task.id) chatapi_create_generation_task.apply_async( @@ -546,7 +560,7 @@ async def recover_one_generation_task( ) return "recover_create_missing_provider_id" - if task.pipeline_stage == "result_ready": + if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value: await _remove_poll_active(task.id) if task.remote_result_url: await enqueue_download_task( @@ -617,8 +631,8 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]: select(ChatGenerationTask) .where( ChatGenerationTask.deleted_at.is_(None), - ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate", "shot_replicate"]), - ChatGenerationTask.status == "generating", + ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)), + ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value, ChatGenerationTask.pipeline_stage.in_( [ "queued", diff --git a/video-gen-api/app/services/media_token_usage_snapshot_service.py b/video-gen-api/app/services/media_token_usage_snapshot_service.py new file mode 100644 index 00000000..9bfef50a --- /dev/null +++ b/video-gen-api/app/services/media_token_usage_snapshot_service.py @@ -0,0 +1,253 @@ +from __future__ import annotations + +import json +from typing import Any, Mapping + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from app.config import settings +from app.enums.credit_record import ( + CreditRecordAction, + CreditRecordChargeKind, + CreditRecordOwnerType, + CreditRecordSourceModule, +) +from app.models.chat_generation_task import ChatGenerationTask +from app.models.credit_record import CreditRecord +from app.models.generation_record import GenerationRecord +from app.models.token_usage import TokenUsage +from app.utils.id_gen import generate_id + + +def _safe_int(value: Any, default: int = 0) -> int: + try: + if value is None or value == "": + return default + return int(value) + except Exception: + return default + + +def _safe_json_dict(value: Any) -> dict[str, Any]: + if not value: + return {} + if isinstance(value, dict): + return value + try: + parsed = json.loads(value) + return parsed if isinstance(parsed, dict) else {} + except Exception: + return {} + + +def _extract_usage(provider_response: Any) -> dict[str, Any]: + data = _safe_json_dict(provider_response) + usage = data.get("usage") + return usage if isinstance(usage, dict) else {} + + +def _normalize_media_tokens( + *, + gen_type: str | None, + provider_response: Any = None, + fallback_total: int | None = None, +) -> tuple[int, int, int]: + usage = _extract_usage(provider_response) + input_tokens = _safe_int(usage.get("input_tokens"), 0) + output_tokens = _safe_int( + usage.get("output_tokens"), + _safe_int(usage.get("generated_tokens"), 0), + ) + total_tokens = _safe_int(usage.get("total_tokens"), 0) + + if total_tokens <= 0: + total_tokens = _safe_int(fallback_total, 0) + if output_tokens <= 0: + output_tokens = max(0, total_tokens - input_tokens) + if total_tokens <= 0: + total_tokens = input_tokens + output_tokens + + # 图片生成多数供应商只返回 output/total,没有 input;保持 input=0。视频同理兼容缺字段。 + return input_tokens, output_tokens, total_tokens + + +def _engine_model_from_provider_response(provider_response: Any) -> str | None: + data = _safe_json_dict(provider_response) + model = data.get("model") + return str(model) if model else None + + +async def _find_latest_media_charge( + db: AsyncSession, + *, + user_id: str, + owner_type: str, + owner_id: str, + media_type: str | None, +) -> CreditRecord | None: + query = ( + select(CreditRecord) + .where(CreditRecord.user_id == user_id) + .where(CreditRecord.type == "consume") + .where(CreditRecord.owner_type == owner_type) + .where(CreditRecord.owner_id == owner_id) + .where(CreditRecord.charge_kind == CreditRecordChargeKind.MEDIA.value) + .where(CreditRecord.charge_action == CreditRecordAction.CHARGE.value) + ) + if media_type: + query = query.where(CreditRecord.media_type == media_type) + query = query.order_by(CreditRecord.attempt_no.desc().nullslast(), CreditRecord.created_at.desc()).limit(1) + result = await db.execute(query) + return result.scalar_one_or_none() + + +async def _get_or_create_token_usage( + db: AsyncSession, + *, + charge: CreditRecord, + input_tokens: int, + output_tokens: int, + total_tokens: int, + model_config_id: str | None = None, +) -> TokenUsage: + token_usage: TokenUsage | None = None + if charge.token_usage_id: + result = await db.execute(select(TokenUsage).where(TokenUsage.id == charge.token_usage_id).limit(1)) + token_usage = result.scalar_one_or_none() + if token_usage is None and charge.biz_key: + result = await db.execute(select(TokenUsage).where(TokenUsage.biz_key == charge.biz_key).limit(1)) + token_usage = result.scalar_one_or_none() + if token_usage is None: + token_usage = TokenUsage( + id=generate_id(), + user_id=charge.user_id, + model_config_id=model_config_id, + owner_type=charge.owner_type, + owner_id=charge.owner_id, + biz_key=charge.biz_key, + source_module=charge.source_module, + source_step_code=charge.source_step_code, + input_tokens=input_tokens, + output_tokens=output_tokens, + total_tokens=total_tokens, + ) + db.add(token_usage) + await db.flush() + else: + token_usage.user_id = token_usage.user_id or charge.user_id + token_usage.model_config_id = token_usage.model_config_id or model_config_id + token_usage.owner_type = token_usage.owner_type or charge.owner_type + token_usage.owner_id = token_usage.owner_id or charge.owner_id + token_usage.biz_key = token_usage.biz_key or charge.biz_key + token_usage.source_module = token_usage.source_module or charge.source_module + token_usage.source_step_code = token_usage.source_step_code or charge.source_step_code + token_usage.input_tokens = input_tokens + token_usage.output_tokens = output_tokens + token_usage.total_tokens = total_tokens + return token_usage + + +async def _sync_charge_snapshot( + db: AsyncSession, + *, + charge: CreditRecord | None, + gen_type: str | None, + provider_response: Any = None, + fallback_total: int | None = None, +) -> CreditRecord | None: + if not charge: + return None + + input_tokens, output_tokens, total_tokens = _normalize_media_tokens( + gen_type=gen_type, + provider_response=provider_response, + fallback_total=fallback_total, + ) + if total_tokens <= 0: + return charge + + token_usage = await _get_or_create_token_usage( + db, + charge=charge, + input_tokens=input_tokens, + output_tokens=output_tokens, + total_tokens=total_tokens, + model_config_id=None, + ) + + charge.token_usage_id = token_usage.id + charge.input_tokens = input_tokens + charge.output_tokens = output_tokens + charge.total_tokens = total_tokens + + # 兼容旧流水扣费时未冷备 engine_model_name 的场景,能从 provider response 推出来就补充。 + provider_model = _engine_model_from_provider_response(provider_response) + if provider_model and not charge.engine_model_name: + charge.engine_model_name = provider_model + return charge + + +async def sync_chat_generation_task_media_token_snapshot( + db: AsyncSession, + task: ChatGenerationTask, + *, + provider_response: Any = None, +) -> CreditRecord | None: + """把 ChatGenerationTask 图片/视频媒体生成 token 后置快照回填到积分流水。 + + 媒体扣费发生在创建任务前,供应商 usage 只能在创建/轮询成功后拿到, + 所以这里按 owner_type + owner_id + media_type 找到对应 media charge 流水并回填。 + """ + if not bool(getattr(settings, "MEDIA_TOKEN_SNAPSHOT_ENABLED", True)): + return None + if not task: + return None + + gen_type = (getattr(task, "gen_type", None) or "").lower().strip() + fallback_total = task.image_tokens_used if gen_type == "image" else task.video_tokens_used + response = provider_response if provider_response is not None else getattr(task, "provider_response_json", None) + charge = await _find_latest_media_charge( + db, + user_id=task.user_id, + owner_type=CreditRecordOwnerType.CHAT_GENERATION_TASK.value, + owner_id=task.id, + media_type=gen_type or None, + ) + return await _sync_charge_snapshot( + db, + charge=charge, + gen_type=gen_type, + provider_response=response, + fallback_total=fallback_total, + ) + + +async def sync_generation_record_media_token_snapshot( + db: AsyncSession, + record: GenerationRecord, + *, + provider_response: Any = None, +) -> CreditRecord | None: + """把旧 GenerationRecord 图片/视频媒体生成 token 后置快照回填到积分流水。""" + if not bool(getattr(settings, "MEDIA_TOKEN_SNAPSHOT_ENABLED", True)): + return None + if not record: + return None + + gen_type = (getattr(record, "gen_type", None) or "").lower().strip() + fallback_total = record.image_tokens_used if gen_type == "image" else record.video_tokens_used + charge = await _find_latest_media_charge( + db, + user_id=record.user_id, + owner_type=CreditRecordOwnerType.GENERATION_RECORD.value, + owner_id=record.id, + media_type=gen_type or None, + ) + return await _sync_charge_snapshot( + db, + charge=charge, + gen_type=gen_type, + provider_response=provider_response, + fallback_total=fallback_total, + ) diff --git a/video-gen-api/app/services/video_queue.py b/video-gen-api/app/services/video_queue.py index 7129972b..95687748 100644 --- a/video-gen-api/app/services/video_queue.py +++ b/video-gen-api/app/services/video_queue.py @@ -10,6 +10,7 @@ from app.models.base import async_session from app.models.generation_record import GenerationRecord from app.services.video_gen import get_active_engine, poll_task_status, download_video, _log_video_response from app.services.image_gen import get_active_image_engine, download_image +from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot from app.services.resource_accounting_service import ( record_generation_record_generated_resource, safe_file_size, @@ -157,6 +158,7 @@ class TaskQueue: else: record.video_url = file_url record.video_tokens_used = poll_result.get("video_tokens", 0) + await sync_generation_record_media_token_snapshot(db, record, provider_response=resp_data) record.status = "completed" record.generated_at = datetime.now() if record.video_url: @@ -235,6 +237,7 @@ class TaskQueue: else: record.image_url = remote_url record.image_tokens_used = poll_result.get("image_tokens", 0) + await sync_generation_record_media_token_snapshot(db, record, provider_response=poll_result) record.status = "completed" record.generated_at = datetime.now() if record.image_url: diff --git a/video-gen-api/app/tasks/generation_create_tasks.py b/video-gen-api/app/tasks/generation_create_tasks.py index 460e50f3..6a25c28e 100644 --- a/video-gen-api/app/tasks/generation_create_tasks.py +++ b/video-gen-api/app/tasks/generation_create_tasks.py @@ -12,6 +12,7 @@ from app.services.error_codes import extract_error_message from app.services.generation_log_service import log_task_event from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once from app.services.generation_provider_service import create_provider_task +from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot from app.services.redis_registry_service import ensure_aware_utc from app.tasks.celery_app import celery_app @@ -223,6 +224,7 @@ async def _run(task_id: str): ensure_ascii=False, default=str, ) + await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=task.provider_response_json) if task.remote_result_url and not task.seedance_task_id: # 同步图片路径:原 SDK 已经返回最终 URL。 diff --git a/video-gen-api/app/tasks/generation_download_tasks.py b/video-gen-api/app/tasks/generation_download_tasks.py index df9a9c54..78968a79 100644 --- a/video-gen-api/app/tasks/generation_download_tasks.py +++ b/video-gen-api/app/tasks/generation_download_tasks.py @@ -1,11 +1,24 @@ +from __future__ import annotations + from app.tasks.async_runner import run_async -from datetime import datetime, timezone, timedelta + +import errno +import math import uuid +from datetime import datetime, timedelta, timezone +from typing import Any from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings +from app.enums.generation_task import ( + ALLOWED_GENERATION_MODES, + ChatGenerationPipelineStage, + ChatGenerationTaskEventType, + ChatGenerationTaskStatus, + GenerationType, +) from app.models.base import async_session from app.models.chat_generation_task import ChatGenerationTask from app.services.celery_download_recovery_service import ( @@ -18,17 +31,17 @@ from app.services.error_codes import extract_error_message from app.services.generation_download_service import download_generation_result from app.services.generation_log_service import log_task_event from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once +from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot from app.services.resource_accounting_service import record_chat_task_generated_resource from app.tasks.celery_app import celery_app -ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"} - DOWNLOAD_QUEUE = "gen_result_download" -DOWNLOAD_STAGE_QUEUED = "download_queued" -DOWNLOAD_STAGE_DOWNLOADING = "downloading" -DOWNLOAD_STAGE_RETRY_WAITING = "retry_waiting" -DOWNLOAD_STAGE_DONE = "done" -DOWNLOAD_STAGE_FAILED = "download_failed" +DOWNLOAD_STAGE_QUEUED = ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value +DOWNLOAD_STAGE_DOWNLOADING = ChatGenerationPipelineStage.DOWNLOADING.value +DOWNLOAD_STAGE_RETRY_WAITING = ChatGenerationPipelineStage.RETRY_WAITING.value +DOWNLOAD_STAGE_DONE = ChatGenerationPipelineStage.DONE.value +DOWNLOAD_STAGE_FAILED = ChatGenerationPipelineStage.DOWNLOAD_FAILED.value +RESULT_READY_STAGE = ChatGenerationPipelineStage.RESULT_READY.value def _now() -> datetime: @@ -56,6 +69,15 @@ def _retry_at(attempt: int, now: datetime | None = None) -> datetime: return now + timedelta(seconds=max(1, base * max(1, attempt))) +def _countdown_until(value: datetime | None, *, minimum: int = 1) -> int: + target = ensure_aware_utc(value) + if target is None: + return minimum + extra = int(getattr(settings, "DOWNLOAD_RETRY_COUNTDOWN_EXTRA_SECONDS", 1) or 0) + seconds = (target - _now()).total_seconds() + return max(minimum, math.ceil(seconds) + extra) + + def _is_expired(value: datetime | None, now: datetime | None = None) -> bool: value = ensure_aware_utc(value) if value is None: @@ -64,10 +86,10 @@ def _is_expired(value: datetime | None, now: datetime | None = None) -> bool: def _is_already_completed(task: ChatGenerationTask) -> bool: - if task.status == "completed" or task.pipeline_stage == DOWNLOAD_STAGE_DONE: - if task.gen_type == "image" and task.image_url: + if task.status == ChatGenerationTaskStatus.COMPLETED.value or task.pipeline_stage == DOWNLOAD_STAGE_DONE: + if task.gen_type == GenerationType.IMAGE.value and task.image_url: return True - if task.gen_type == "video" and task.video_url: + if task.gen_type == GenerationType.VIDEO.value and task.video_url: return True return False @@ -77,6 +99,61 @@ def _build_celery_task_id(task_id: str, attempt: int | None = None, reason: str return f"download:{task_id}:{int(attempt or 0)}:{safe_reason}:{uuid.uuid4().hex[:12]}" +def _is_non_retryable_download_error(exc: Exception) -> bool: + if not bool(getattr(settings, "DOWNLOAD_NON_RETRYABLE_LOCAL_ERRORS", True)): + return False + if isinstance(exc, PermissionError): + return True + if isinstance(exc, OSError) and getattr(exc, "errno", None) in {errno.EACCES, errno.EPERM, errno.ENOSPC, errno.EROFS, errno.ENAMETOOLONG}: + return True + message = str(exc).lower() + non_retryable_fragments = ( + "permission denied", + "no space left on device", + "read-only file system", + "file name too long", + "invalid argument", + ) + return any(fragment in message for fragment in non_retryable_fragments) + + +async def _log_download_event( + task: ChatGenerationTask | None = None, + *, + task_id: str | None = None, + event_type: ChatGenerationTaskEventType | str, + from_status: str | None = None, + to_status: str | None = None, + from_stage: str | None = None, + to_stage: str | None = None, + message: str | None = None, + detail: Any = None, +) -> None: + if not bool(getattr(settings, "DOWNLOAD_EVENT_VERBOSE_ENABLED", True)): + # 成功/失败关键事件仍保留;只关闭 verbose skip 事件。 + critical = { + ChatGenerationTaskEventType.DOWNLOAD_START.value, + ChatGenerationTaskEventType.DOWNLOAD_SUCCESS.value, + ChatGenerationTaskEventType.DOWNLOAD_FAILED.value, + ChatGenerationTaskEventType.DOWNLOAD_FAILED_NON_RETRYABLE.value, + ChatGenerationTaskEventType.DOWNLOAD_RETRY_WAITING.value, + } + event_value = event_type.value if hasattr(event_type, "value") else str(event_type) + if event_value not in critical: + return + await log_task_event( + task, + task_id=task_id, + event_type=event_type.value if hasattr(event_type, "value") else str(event_type), + from_status=from_status, + to_status=to_status, + from_stage=from_stage, + to_stage=to_stage, + message=message, + detail=detail, + ) + + async def _register_active_from_task( task: ChatGenerationTask, *, @@ -101,6 +178,62 @@ async def _register_active_from_task( await upsert_download_active(record_id=task.id, payload=payload, check_at=check_at) +async def _apply_download_async( + task: ChatGenerationTask, + *, + priority: int, + countdown: int | None, + reason: str, + event_type: ChatGenerationTaskEventType, + failed_event_type: ChatGenerationTaskEventType, +) -> bool: + if not celery_app: + await _log_download_event( + task, + event_type=failed_event_type, + message="Celery 未启用,下载任务无法投递", + detail={"queue": DOWNLOAD_QUEUE, "priority": priority, "countdown": countdown, "reason": reason}, + ) + return False + try: + download_generation_result_task.apply_async( + args=[task.id], + queue=DOWNLOAD_QUEUE, + priority=priority, + countdown=countdown, + task_id=task.download_celery_task_id, + ) + await _log_download_event( + task, + event_type=event_type, + to_stage=task.pipeline_stage, + detail={ + "queue": DOWNLOAD_QUEUE, + "priority": priority, + "countdown": countdown, + "download_celery_task_id": task.download_celery_task_id, + "reason": reason, + "attempt": task.download_attempt_count, + }, + ) + return True + except Exception as exc: + await _log_download_event( + task, + event_type=failed_event_type, + message=str(exc), + detail={ + "queue": DOWNLOAD_QUEUE, + "priority": priority, + "countdown": countdown, + "download_celery_task_id": task.download_celery_task_id, + "reason": reason, + "attempt": task.download_attempt_count, + }, + ) + return False + + async def enqueue_download_task( db: AsyncSession, task: ChatGenerationTask, @@ -110,11 +243,17 @@ async def enqueue_download_task( countdown: int | None = None, ) -> str | None: """统一投递图片/视频下载任务,并同步 DB + Redis active 注册表。""" - if not task or task.generation_mode not in ALLOWED_GENERATION_MODES: + reason = reason or ("recover" if recover else "normal") + if not task: return None - if task.status != "generating": + if task.generation_mode not in ALLOWED_GENERATION_MODES: + await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_INVALID_MODE, message="不支持的 generation_mode", detail={"reason": reason}) + return None + if task.status != ChatGenerationTaskStatus.GENERATING.value: + await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING, message="任务不是 generating 状态", detail={"reason": reason, "status": task.status}) return None if not task.remote_result_url: + await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL, message="缺少 remote_result_url", detail={"reason": reason}) return None now = _now() @@ -122,9 +261,10 @@ async def enqueue_download_task( celery_task_id = _build_celery_task_id( task.id, attempt=task.download_attempt_count or task.retry_count or 0, - reason=reason or ("recover" if recover else "normal"), + reason=reason, ) + old_stage = task.pipeline_stage task.pipeline_stage = DOWNLOAD_STAGE_QUEUED task.download_celery_task_id = celery_task_id task.download_enqueued_at = now @@ -137,13 +277,22 @@ async def enqueue_download_task( check_at = _queue_timeout_at(now) await _register_active_from_task(task, check_at=check_at, priority=priority, reason=reason) - if celery_app: - download_generation_result_task.apply_async( - args=[task.id], - queue=DOWNLOAD_QUEUE, - priority=priority, - countdown=countdown, - task_id=celery_task_id, + await _apply_download_async( + task, + priority=priority, + countdown=countdown, + reason=reason, + event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE, + failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE_FAILED if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE_FAILED, + ) + if old_stage != DOWNLOAD_STAGE_QUEUED: + # 独立记录阶段变化的上下文,便于和真正投递事件对照。 + await _log_download_event( + task, + event_type=ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE, + from_stage=old_stage, + to_stage=DOWNLOAD_STAGE_QUEUED, + detail={"reason": reason, "recover": recover, "check_at": check_at}, ) return celery_task_id @@ -158,34 +307,81 @@ async def _reload_task(db: AsyncSession, task_id: str) -> ChatGenerationTask | N return result.scalar_one_or_none() +async def _reschedule_not_due_retry(task: ChatGenerationTask, *, now: datetime) -> None: + countdown = _countdown_until(task.download_next_retry_at) + await _register_active_from_task( + task, + check_at=ensure_aware_utc(task.download_next_retry_at) or (now + timedelta(seconds=countdown)), + priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, + reason="retry_waiting_not_due", + ) + await _apply_download_async( + task, + priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, + countdown=countdown, + reason="retry_waiting_not_due", + event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE, + failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE_FAILED, + ) + + async def _claim_download_lease(db: AsyncSession, task: ChatGenerationTask) -> bool: now = _now() - if not task or task.generation_mode not in ALLOWED_GENERATION_MODES: + if not task: return False - if task.status != "generating": + if task.generation_mode not in ALLOWED_GENERATION_MODES: + await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_INVALID_MODE, message="下载任务跳过:不支持的 generation_mode") + return False + if task.status != ChatGenerationTaskStatus.GENERATING.value: + await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING, message="下载任务跳过:任务不是 generating 状态", detail={"status": task.status, "stage": task.pipeline_stage}) return False if _is_already_completed(task): + await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_ALREADY_COMPLETED, message="下载任务跳过:任务已完成", detail={"status": task.status, "stage": task.pipeline_stage}) return False if not task.remote_result_url: + await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL, message="下载任务跳过:缺少 remote_result_url") return False stage = task.pipeline_stage if stage == DOWNLOAD_STAGE_DOWNLOADING: if not _is_expired(task.download_lease_until, now): + await _log_download_event( + task, + event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_DOWNLOADING_LEASE_ALIVE, + message="下载任务跳过:已有 downloading lease 且未过期", + detail={"lease_until": task.download_lease_until, "download_celery_task_id": task.download_celery_task_id}, + ) return False - await log_task_event( + await _log_download_event( task, - event_type="DOWNLOAD_STUCK_RECOVER", + event_type=ChatGenerationTaskEventType.DOWNLOAD_STUCK_RECOVER, message=f"downloading lease 已过期,重新抢占下载。lease_until={task.download_lease_until}", ) elif stage == DOWNLOAD_STAGE_RETRY_WAITING: if not _is_expired(task.download_next_retry_at, now): + await _log_download_event( + task, + event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_NOT_DUE, + message="下载重试提前触发,尚未到 next_retry_at,已重新投递延后重试", + detail={ + "now": now, + "download_next_retry_at": task.download_next_retry_at, + "download_celery_task_id": task.download_celery_task_id, + }, + ) + await _reschedule_not_due_retry(task, now=now) return False - elif stage in (DOWNLOAD_STAGE_QUEUED, "result_ready"): + elif stage in (DOWNLOAD_STAGE_QUEUED, RESULT_READY_STAGE): pass else: + await _log_download_event( + task, + event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_STAGE_NOT_ALLOWED, + message="下载任务跳过:当前阶段不允许下载", + detail={"stage": stage, "status": task.status, "download_celery_task_id": task.download_celery_task_id}, + ) return False old_stage = stage @@ -207,9 +403,9 @@ async def _claim_download_lease(db: AsyncSession, task: ChatGenerationTask) -> b reason="claim_download_lease", ) - await log_task_event( + await _log_download_event( task, - event_type="DOWNLOAD_START", + event_type=ChatGenerationTaskEventType.DOWNLOAD_START, from_stage=old_stage, to_stage=DOWNLOAD_STAGE_DOWNLOADING, detail={ @@ -244,9 +440,9 @@ async def _mark_retry_waiting(db: AsyncSession, task: ChatGenerationTask, exc: E reason="download_retry_waiting", ) - await log_task_event( + await _log_download_event( task, - event_type="DOWNLOAD_RETRY_WAITING", + event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_WAITING, message=error_message, to_stage=DOWNLOAD_STAGE_RETRY_WAITING, detail={ @@ -262,16 +458,53 @@ def _should_final_fail(task: ChatGenerationTask) -> bool: return int(task.download_attempt_count or task.retry_count or 0) >= int(settings.DOWNLOAD_TASK_MAX_ATTEMPTS or 3) +async def _mark_download_failed( + db: AsyncSession, + task: ChatGenerationTask, + *, + exc: Exception, + non_retryable: bool = False, +) -> None: + error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc) + await mark_chat_generation_task_failed_and_refund_once( + db, + task=task, + error_message=error_message, + pipeline_stage=DOWNLOAD_STAGE_FAILED, + ) + task.download_last_error = error_message + task.download_lease_until = None + task.download_next_retry_at = None + + from app.services.generation_module_hook_service import notify_chat_generation_task_finished + + await notify_chat_generation_task_finished(db, task) + await db.commit() + await remove_download_active(task.id) + + await _log_download_event( + task, + event_type=ChatGenerationTaskEventType.DOWNLOAD_FAILED_NON_RETRYABLE if non_retryable else ChatGenerationTaskEventType.DOWNLOAD_FAILED, + message=task.error_message, + detail={ + "download_attempt_count": task.download_attempt_count, + "max_attempts": settings.DOWNLOAD_TASK_MAX_ATTEMPTS, + "non_retryable": non_retryable, + "download_last_error": task.download_last_error, + }, + ) + + async def _run(task_id: str): async with async_session() as db: - result = await db.execute( - select(ChatGenerationTask).where( - ChatGenerationTask.id == task_id, - ChatGenerationTask.deleted_at.is_(None), - ).with_for_update().limit(1) - ) - task = result.scalar_one_or_none() + task = await _reload_task(db, task_id) if not task: + await _log_download_event( + task_id=task_id, + event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_TASK_MISSING, + message="下载任务跳过:ChatGenerationTask 不存在或已软删", + ) + await remove_download_active(task_id) return claimed = await _claim_download_lease(db, task) @@ -283,15 +516,21 @@ async def _run(task_id: str): task = await _reload_task(db, task_id) if not task: + await _log_download_event( + task_id=task_id, + event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_TASK_MISSING, + message="下载完成后任务不存在或已软删", + ) + await remove_download_active(task_id) return - if task.gen_type == "image": + if task.gen_type == GenerationType.IMAGE.value: task.image_url = downloaded.url else: task.video_url = downloaded.url task.video_cover_url = downloaded.cover_url - task.status = "completed" + task.status = ChatGenerationTaskStatus.COMPLETED.value task.pipeline_stage = DOWNLOAD_STAGE_DONE task.generated_at = _now() task.retry_count = 0 @@ -308,18 +547,18 @@ async def _run(task_id: str): remote_url=task.remote_result_url, generated_at=task.generated_at, ) + await sync_chat_generation_task_media_token_snapshot(db, task) + from app.services.generation_module_hook_service import notify_chat_generation_task_finished + + await notify_chat_generation_task_finished(db, task) await db.commit() await remove_download_active(task.id) - from app.services.generation_module_hook_service import notify_chat_generation_task_finished - await notify_chat_generation_task_finished(db, task) - await db.commit() - - await log_task_event( + await _log_download_event( task, - event_type="DOWNLOAD_SUCCESS", - to_status="completed", + event_type=ChatGenerationTaskEventType.DOWNLOAD_SUCCESS, + to_status=ChatGenerationTaskStatus.COMPLETED.value, to_stage=DOWNLOAD_STAGE_DONE, detail={ "resource_url": downloaded.url, @@ -337,48 +576,23 @@ async def _run(task_id: str): task = await _reload_task(db, task_id) if not task: + await remove_download_active(task_id) return - if _should_final_fail(task): - error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc) - await mark_chat_generation_task_failed_and_refund_once( - db, - task=task, - error_message=error_message, - pipeline_stage=DOWNLOAD_STAGE_FAILED, - ) - task.download_last_error = error_message - task.download_lease_until = None - task.download_next_retry_at = None - await db.commit() - - await remove_download_active(task.id) - - from app.services.generation_module_hook_service import notify_chat_generation_task_finished - await notify_chat_generation_task_finished(db, task) - await db.commit() - - await log_task_event( - task, - event_type="DOWNLOAD_FAILED", - message=task.error_message, - detail={ - "download_attempt_count": task.download_attempt_count, - "max_attempts": settings.DOWNLOAD_TASK_MAX_ATTEMPTS, - }, - ) + non_retryable = _is_non_retryable_download_error(exc) + if non_retryable or _should_final_fail(task): + await _mark_download_failed(db, task, exc=exc, non_retryable=non_retryable) else: next_retry_at = await _mark_retry_waiting(db, task, exc) - - if celery_app: - delay_seconds = max(1, int((next_retry_at - _now()).total_seconds())) - download_generation_result_task.apply_async( - args=[task.id], - queue=DOWNLOAD_QUEUE, - priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, - countdown=delay_seconds, - task_id=task.download_celery_task_id, - ) + delay_seconds = _countdown_until(next_retry_at) + await _apply_download_async( + task, + priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, + countdown=delay_seconds, + reason="download_exception_retry", + event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE, + failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE_FAILED, + ) if celery_app: diff --git a/video-gen-api/app/tasks/generation_poll_tasks.py b/video-gen-api/app/tasks/generation_poll_tasks.py index ffc585f2..ac4aab17 100644 --- a/video-gen-api/app/tasks/generation_poll_tasks.py +++ b/video-gen-api/app/tasks/generation_poll_tasks.py @@ -12,6 +12,7 @@ from app.services.error_codes import extract_error_message from app.services.generation_log_service import log_task_event, log_provider_call from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once from app.services.generation_provider_service import poll_provider_task +from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot from app.services.redis_registry_service import ( datetime_to_epoch, ensure_aware_utc, @@ -245,6 +246,7 @@ async def _run(task_id: str): task.video_tokens_used = poll_result.get("video_tokens", 0) or 0 task.provider_response_json = response_data + await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=response_data) if not task.remote_result_url: await _mark_failed(db, task, message="供应商任务成功但未返回结果URL", detail=poll_result) diff --git a/video-gen-api/app/tasks/generation_recovery_tasks.py b/video-gen-api/app/tasks/generation_recovery_tasks.py index f27aa64d..dee9039c 100644 --- a/video-gen-api/app/tasks/generation_recovery_tasks.py +++ b/video-gen-api/app/tasks/generation_recovery_tasks.py @@ -1,12 +1,17 @@ # app/tasks/generation_recovery_tasks.py from __future__ import annotations +import logging from typing import Any, Dict +from app.config import settings from app.models.base import async_session +from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock from app.tasks.async_runner import run_async from app.tasks.celery_app import celery_app +logger = logging.getLogger("video_gen") + async def _run_download_once() -> Dict[str, Any]: from app.services.generation_recovery_service import recover_download_tasks_once @@ -22,11 +27,49 @@ async def _run_generation_once() -> Dict[str, Any]: return await recover_generation_tasks_once(db) +async def _acquire_download_recovery_loop_lock() -> tuple[bool, str]: + """下载恢复循环锁。 + + Redis 不可用时降级为直接执行 DB fallback,避免恢复能力彻底失效; + Redis 可用但锁被其他 worker 持有时,本轮跳过,不再重复投递下一轮。 + """ + redis = await get_registry_redis() + if redis is None: + return True, "redis_unavailable_run_db_fallback" + token = await redis_acquire_lock( + lock_key=settings.DOWNLOAD_RECOVERY_LOOP_LOCK_KEY, + ttl_seconds=int(settings.DOWNLOAD_RECOVERY_LOOP_LOCK_TTL_SECONDS or 55), + log_context="download_recovery_loop", + ) + return (bool(token), "lock_acquired" if token else "lock_held") + + +def _schedule_next_download_recovery_loop() -> None: + if not celery_app or not bool(getattr(settings, "DOWNLOAD_RECOVERY_LOOP_ENABLED", True)): + return + try: + recover_download_tasks_once.apply_async( + countdown=max(1, int(settings.DOWNLOAD_RECOVERY_INTERVAL_SECONDS or 60)), + queue="gen_result_download", + priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, + ) + except Exception: + logger.exception("下载恢复循环下一轮投递失败") + + if celery_app: - @celery_app.task(name="generation.recover_download_tasks_once") - def recover_download_tasks_once() -> Dict[str, Any]: - return run_async(_run_download_once()) + @celery_app.task(name="generation.recover_download_tasks_once", bind=True) + def recover_download_tasks_once(self) -> Dict[str, Any]: + acquired, reason = run_async(_acquire_download_recovery_loop_lock()) + if not acquired: + return {"skipped": reason} + try: + result = run_async(_run_download_once()) + result["loop_lock"] = reason + return result + finally: + _schedule_next_download_recovery_loop() @celery_app.task(name="generation.recover_generation_tasks_once")