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 1/3] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E7=94=9F=E8=BE=B0?= =?UTF-8?q?=E4=BB=BB=E5=8A=A1=E6=81=A2=E5=A4=8D=E6=9C=BA=E5=88=B6BUG|?= =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E7=94=9F=E6=88=90=E8=A7=86=E9=A2=91/?= =?UTF-8?q?=E5=9B=BE=E7=89=87token=E5=BF=AB=E7=85=A7=E4=B8=8D=E5=9B=9E?= =?UTF-8?q?=E8=90=BD=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") From ee03242e6c799de18a5e7a4faa440033247c94a8 Mon Sep 17 00:00:00 2001 From: GinHa <15201596918@163.com> Date: Fri, 26 Jun 2026 13:15:25 +0800 Subject: [PATCH 2/3] =?UTF-8?q?celery=E5=BC=82=E6=AD=A5=E6=81=A2=E5=A4=8D?= =?UTF-8?q?=E4=BB=BB=E5=8A=A1=E7=8B=AC=E7=AB=8B=E9=98=9F=E5=88=97|?= =?UTF-8?q?=E6=89=A9=E5=A2=9Ecelery=E5=AD=90=E8=BF=9B=E7=A8=8B=E8=BF=9E?= =?UTF-8?q?=E6=8E=A5=E6=B1=A0=E4=B8=8A=E9=99=90=E9=85=8D=E7=BD=AE?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- video-gen-api/app/config.py | 33 ++- .../services/generation_recovery_service.py | 273 ++++++++---------- video-gen-api/app/tasks/async_runner.py | 110 ++++++- video-gen-api/app/tasks/celery_app.py | 57 ++-- .../app/tasks/generation_create_tasks.py | 8 +- .../app/tasks/generation_download_tasks.py | 8 +- .../app/tasks/generation_poll_tasks.py | 38 ++- .../app/tasks/generation_recovery_tasks.py | 163 ++++++++++- .../app/tasks/module_async_recovery_tasks.py | 36 ++- .../app/tasks/shot_replicate_tasks.py | 36 ++- 10 files changed, 517 insertions(+), 245 deletions(-) diff --git a/video-gen-api/app/config.py b/video-gen-api/app/config.py index 1396d051..29c753ef 100644 --- a/video-gen-api/app/config.py +++ b/video-gen-api/app/config.py @@ -113,8 +113,8 @@ class Settings(BaseSettings): PROVIDER_LIMIT_TOKEN_TTL_SECONDS: int = 600 CELERY_DB_POOL_SIZE: int = 1 - CELERY_DB_MAX_OVERFLOW: int = 1 - CELERY_DB_POOL_TIMEOUT: int = 30 + CELERY_DB_MAX_OVERFLOW: int = 2 + CELERY_DB_POOL_TIMEOUT: int = 60 CELERY_DB_POOL_RECYCLE: int = 1800 # Celery 图片/视频下载容灾配置。 @@ -125,10 +125,10 @@ class Settings(BaseSettings): DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS: int = 30 DOWNLOAD_TASK_LEASE_SECONDS: int = 10 * 60 DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS: int = 5 * 60 - DOWNLOAD_RECOVERY_BATCH_SIZE: int = 100 + DOWNLOAD_RECOVERY_BATCH_SIZE: int = 20 DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS: int = 3 # 下载恢复自循环:不依赖 Celery beat,不新增 worker;由 gen_result_download 队列周期扫描 DB/Redis。 - DOWNLOAD_RECOVERY_LOOP_ENABLED: bool = True + DOWNLOAD_RECOVERY_LOOP_ENABLED: bool = False 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 @@ -141,23 +141,32 @@ class Settings(BaseSettings): # Celery 生成链路 / provider poll 容灾配置。 # 说明: - # - 不新增 Celery worker;恢复任务仍投递到 gen_result_download。 - # - worker_ready 每个 worker 都会尝试抢启动恢复锁,只有抢到锁的 worker 投递恢复任务。 + # - 启动容灾保留,但恢复扫描独立投递到 CELERY_RECOVERY_QUEUE。 + # - worker_ready 每个 worker 都会尝试抢启动恢复锁,只有抢到锁的 worker 投递恢复协调任务。 # - poll active 使用独立 Redis key,避免影响稳定的下载 active 注册表。 - GENERATION_RECOVERY_BATCH_SIZE: int = 100 - GENERATION_RECOVERY_MAX_ROUNDS: int = 5 - POLL_RECOVERY_BATCH_SIZE: int = 100 + GENERATION_RECOVERY_BATCH_SIZE: int = 20 + GENERATION_RECOVERY_MAX_ROUNDS: int = 1 + POLL_RECOVERY_BATCH_SIZE: int = 20 POLL_TASK_LEASE_SECONDS: int = 5 * 60 POLL_TASK_QUEUE_TIMEOUT_SECONDS: int = 2 * 60 POLL_ACTIVE_REDIS_HASH_KEY: str = "vg:celery:poll:active" POLL_ACTIVE_REDIS_ZSET_KEY: str = "vg:celery:poll:active_index" CELERY_STARTUP_RECOVERY_LOCK_KEY: str = "vg:celery:startup_recovery_lock" CELERY_STARTUP_RECOVERY_LOCK_TTL_SECONDS: int = 120 + CELERY_RECOVERY_QUEUE: str = "gen_recovery" + CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS: int = 10 * 60 + CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS: int = 300 + CELERY_RECOVERY_TIME_LIMIT_SECONDS: int = 420 + CELERY_RECOVERY_STARTUP_TASK_LOCK_KEY: str = "vg:celery:startup_recovery_task_lock" + GENERATION_RECOVERY_LOCK_KEY: str = "vg:celery:generation_recovery_lock" + DOWNLOAD_RECOVERY_LOCK_KEY: str = "vg:celery:download_recovery_lock" + MODULE_ASYNC_RECOVERY_LOCK_KEY: str = "vg:celery:module_async_recovery_lock" + SHOT_SPLIT_RECOVERY_LOCK_KEY: str = "vg:celery:shot_split_recovery_lock" # 模块异步任务容灾配置。 # 覆盖 ModuleGenerationStep 提词任务、shot 原视频/片段分析、shot ffmpeg 切割 active 注册。 - # 不新增 worker 队列:恢复扫描仍走 gen_result_download,真实业务任务回到原始队列。 - MODULE_ASYNC_RECOVERY_BATCH_SIZE: int = 100 + # 恢复扫描走 CELERY_RECOVERY_QUEUE,真实业务任务回到原始队列。 + MODULE_ASYNC_RECOVERY_BATCH_SIZE: int = 20 MODULE_ASYNC_QUEUE_TIMEOUT_SECONDS: int = 5 * 60 MODULE_ASYNC_LEASE_SECONDS: int = 10 * 60 MODULE_ASYNC_LOCK_TTL_SECONDS: int = 10 * 60 @@ -201,7 +210,7 @@ class Settings(BaseSettings): SHOT_SPLIT_RETRY_BACKOFF_SECONDS: int = 30 SHOT_SPLIT_LEASE_SECONDS: int = 10 * 60 SHOT_SPLIT_PENDING_TIMEOUT_SECONDS: int = 5 * 60 - SHOT_SPLIT_RECOVERY_BATCH_SIZE: int = 50 + SHOT_SPLIT_RECOVERY_BATCH_SIZE: int = 20 SHOT_SPLIT_LOCK_KEY_PREFIX: str = "vg:shot_replicate:split:lock" SHOT_SPLIT_SEMAPHORE_KEY_PREFIX: str = "vg:shot_replicate:split:semaphore" diff --git a/video-gen-api/app/services/generation_recovery_service.py b/video-gen-api/app/services/generation_recovery_service.py index 0ce6d132..451b26fc 100644 --- a/video-gen-api/app/services/generation_recovery_service.py +++ b/video-gen-api/app/services/generation_recovery_service.py @@ -23,10 +23,8 @@ from app.services.celery_download_recovery_service import ( postpone_download_active_check, remove_download_active, ) -from app.services.generation_log_service import log_provider_call, log_task_event +from app.services.generation_log_service import 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 ( redis_get_due_registry_ids, @@ -145,10 +143,20 @@ async def recover_one_download_task( return "clean_final_state" 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}) + 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}) + 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 @@ -362,94 +370,6 @@ async def _mark_failed( return "mark_failed" -async def _try_final_poll_before_timeout(db: AsyncSession, task: ChatGenerationTask) -> str: - """超时前最后查一次供应商,避免 Celery 中断导致本地假超时。 - - 如果供应商已经成功,继续进入下载;如果仍 running 或查询失败,再按超时处理。 - """ - from app.tasks.generation_download_tasks import enqueue_download_task - - if not (task.provider_task_id or task.seedance_task_id): - return await _mark_timeout(db, task) - - try: - poll_result = await poll_provider_task(db, task) - status = poll_result.get("status") - response_data = poll_result.get("response_data") - except Exception as exc: - await log_task_event( - task, - event_type="FINAL_POLL_BEFORE_TIMEOUT_ERROR", - message=str(exc), - ) - return await _mark_timeout(db, task) - - try: - provider_response = json.loads(response_data or "{}") - except Exception: - provider_response = {"raw": response_data} - - snapshot = _engine_snapshot(task) - await log_provider_call( - task, - provider=snapshot.get("provider") or "ark", - api_type=f"{task.gen_type}_final_poll_before_timeout", - model=snapshot.get("model_name"), - engine_id=task.engine_id, - status="success", - provider_task_id=task.seedance_task_id or task.provider_task_id, - response_data=provider_response, - ) - - if _is_success(status): - if task.gen_type == "image": - task.remote_result_url = poll_result.get("image_url") - task.image_tokens_used = poll_result.get("image_tokens", 0) or 0 - else: - task.remote_result_url = poll_result.get("video_url") - 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, - task, - error_message="供应商任务成功但未返回结果URL", - detail=poll_result, - ) - - 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=ChatGenerationPipelineStage.RESULT_READY.value, - detail=poll_result, - ) - await enqueue_download_task(db, task, recover=True, reason="final_poll_before_timeout_success") - return "recover_timeout_success_to_download" - - if _is_failed(status): - task.provider_response_json = response_data - return await _mark_failed( - db, - task, - error_message=poll_result.get("error") or f"供应商任务失败: {status}", - detail=poll_result, - ) - - await log_task_event( - task, - event_type="FINAL_POLL_BEFORE_TIMEOUT_PENDING", - message=f"status={status}", - detail=poll_result, - ) - return await _mark_timeout(db, task) - - async def recover_one_generation_task( db: AsyncSession, task: ChatGenerationTask, @@ -457,6 +377,14 @@ async def recover_one_generation_task( payload: dict[str, Any] | None = None, source: str = "startup_db", ) -> str: + """恢复单个生成任务。 + + 分流原则: + 1. 已有 remote_result_url:只恢复下载,不 poll,不重新 create。 + 2. 已有 provider_task_id/seedance_task_id:恢复 poll。 + 3. 无结果 URL、无供应商任务 ID:deadline 未过才恢复 create。 + 4. 无结果 URL、无供应商任务 ID:deadline 已过直接超时失败,不再补救生成。 + """ from app.tasks.generation_create_tasks import chatapi_create_generation_task from app.tasks.generation_download_tasks import enqueue_download_task from app.tasks.generation_poll_tasks import poll_generation_task, register_poll_active @@ -472,37 +400,107 @@ async def recover_one_generation_task( if _is_final_task_state(task): await _remove_poll_active(task.id) return "clean_final_state" - if task.status != "generating": + if task.status != ChatGenerationTaskStatus.GENERATING.value: await _remove_poll_active(task.id) return "clean_not_generating" - if task.deadline_at and _is_expired(task.deadline_at, current_time): - 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) + has_remote_result = bool(str(task.remote_result_url or "").strip()) + has_provider_task_id = bool(str(task.provider_task_id or "").strip() or str(task.seedance_task_id or "").strip()) + is_deadline_expired = bool(task.deadline_at and _is_expired(task.deadline_at, current_time)) - 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: + # 最高优先级:只要远程结果 URL 已经落库,说明生成侧已经成功。 + # 不管当前 pipeline_stage 是 queued/creating/waiting/result_ready/download_*,恢复时都不能重复 create 或 poll。 + if has_remote_result: + await _remove_poll_active(task.id) + await log_task_event( + task, + event_type="GENERATION_RECOVERY_ENQUEUE", + message=f"{source} 发现任务已存在 remote_result_url,恢复投递下载队列", + detail={ + "pipeline_stage": task.pipeline_stage, + "payload": redis_payload, + "deadline_expired": is_deadline_expired, + }, + ) + await enqueue_download_task( + db, + task, + recover=True, + reason=f"{source}_has_remote_result_url", + ) + return "recover_download_has_remote_result" + + # 已经过 deadline 且没有结果 URL: + # - 有供应商任务 ID:交给 poll worker 做最后一次状态确认; + # - 没有供应商任务 ID:说明没有可查询的远程任务,直接按超时失败处理,不再重新 create。 + if is_deadline_expired: + if has_provider_task_id: task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value await db.commit() await log_task_event( task, event_type="GENERATION_RECOVERY_ENQUEUE", - message=f"{source} 发现创建阶段已存在供应商任务ID,恢复投递轮询队列", + message=f"{source} 发现任务已到 deadline 且存在供应商任务ID,投递 poll 队列做最终查询", detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, ) poll_generation_task.apply_async(args=[task.id], queue=POLL_QUEUE, countdown=0) await register_poll_active( task, check_at=_poll_queue_timeout_at(), - reason=f"{source}_create_stage_has_provider_id", + reason=f"{source}_deadline_final_poll", ) - return "recover_poll_from_create_stage" + return "recover_deadline_final_poll" + await log_task_event( + task, + event_type="GENERATION_RECOVERY_TIMEOUT", + message=f"{source} 发现任务已到 deadline,且没有 remote_result_url/供应商任务ID,按超时失败处理", + detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, + ) + return await _mark_timeout(db, task) + + # 未过 deadline:有供应商任务 ID 才允许恢复到 poll 队列。 + if has_provider_task_id: + task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value + await db.commit() await log_task_event( task, event_type="GENERATION_RECOVERY_ENQUEUE", - message=f"{source} 发现创建阶段任务未完成,恢复投递创建队列", + message=f"{source} 发现任务存在供应商任务ID,恢复投递轮询队列", + detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, + ) + poll_generation_task.apply_async(args=[task.id], queue=POLL_QUEUE, countdown=0) + await register_poll_active( + task, + check_at=_poll_queue_timeout_at(), + reason=f"{source}_has_provider_task_id", + ) + return "recover_poll_has_provider_id" + + # 未过 deadline,且没有结果 URL / 供应商任务 ID: + # 图片同步任务会重新进入 submit_image_task;视频/其它任务会重新创建供应商任务。 + # 这里不能投 poll,因为没有 provider_task_id/seedance_task_id 可查询。 + recoverable_create_stages = { + ChatGenerationPipelineStage.QUEUED.value, + ChatGenerationPipelineStage.PREPARING.value, + ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, + ChatGenerationPipelineStage.WAITING_REMOTE.value, + ChatGenerationPipelineStage.POLLING.value, + } + if task.pipeline_stage in recoverable_create_stages: + if task.pipeline_stage not in ( + ChatGenerationPipelineStage.QUEUED.value, + ChatGenerationPipelineStage.PREPARING.value, + ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, + ): + task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value + await db.commit() + + await _remove_poll_active(task.id) + await log_task_event( + task, + event_type="GENERATION_RECOVERY_ENQUEUE", + message=f"{source} 发现任务未超时且缺少 remote_result_url/供应商任务ID,恢复投递创建队列", detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, ) chatapi_create_generation_task.apply_async( @@ -510,67 +508,25 @@ async def recover_one_generation_task( queue="gen_chatapi_create", countdown=0, ) - return "recover_create" + return "recover_create_no_remote_no_provider_before_deadline" - 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( - db, - task, - recover=True, - reason=f"{source}_waiting_remote_has_result", - ) - return "recover_waiting_has_result" - - if task.provider_task_id or task.seedance_task_id: - await log_task_event( - task, - event_type="GENERATION_RECOVERY_ENQUEUE", - message=f"{source} 发现远程等待/轮询阶段任务未完成,恢复投递轮询队列", - detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, - ) - task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value - await db.commit() - poll_generation_task.apply_async( - args=[task.id], - queue=POLL_QUEUE, - countdown=0, - ) - await register_poll_active( - task, - check_at=_poll_queue_timeout_at(), - reason=f"{source}_recover_poll", - ) - return "recover_poll" - - await log_task_event( - task, - event_type="GENERATION_RECOVERY_ENQUEUE", - message=f"{source} 发现任务缺少供应商任务ID,恢复投递创建队列", - detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, - ) + # result_ready 但没有 URL 是脏状态;未过 deadline 时回创建队列重新处理,过期上面已标记超时。 + if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value: task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value await db.commit() await _remove_poll_active(task.id) + await log_task_event( + task, + event_type="GENERATION_RECOVERY_ENQUEUE", + message=f"{source} 发现 result_ready 但缺少 remote_result_url,未超时,恢复投递创建队列", + detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, + ) chatapi_create_generation_task.apply_async( args=[task.id], queue="gen_chatapi_create", countdown=0, ) - return "recover_create_missing_provider_id" - - if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value: - await _remove_poll_active(task.id) - if task.remote_result_url: - await enqueue_download_task( - db, - task, - recover=True, - reason=f"{source}_generation_result_ready", - ) - return "recover_result_ready" - return "skip_result_ready_no_url" + return "recover_create_result_ready_no_url_before_deadline" return f"skip_stage_{task.pipeline_stage}" @@ -670,11 +626,8 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]: if len(tasks) < batch_size or progressed_this_round <= 0: break - # 下载阶段单独跑 DB fallback。 - download_result = await recover_download_tasks_once(db) return { "checked": len(checked_ids), "db_checked": total_db_checked, "results": results, - "download_recovery": download_result, - } + } \ No newline at end of file diff --git a/video-gen-api/app/tasks/async_runner.py b/video-gen-api/app/tasks/async_runner.py index 24964dad..e362fc4d 100644 --- a/video-gen-api/app/tasks/async_runner.py +++ b/video-gen-api/app/tasks/async_runner.py @@ -1,13 +1,16 @@ from __future__ import annotations import asyncio +import logging import os import threading -from concurrent.futures import Future +from concurrent.futures import Future, TimeoutError as FutureTimeoutError from typing import Awaitable, TypeVar from app.config import settings +logger = logging.getLogger("video_gen") + T = TypeVar("T") _thread_local = threading.local() @@ -18,6 +21,28 @@ _single_loop_pid: int | None = None _single_loop_ready: threading.Event | None = None +async def _dispose_async_resources() -> None: + """释放当前 async loop 内缓存的异步资源。 + + Celery soft time limit 会打断同步等待 future.result() 的线程;如果不主动 + cancel coroutine 并释放 engine/redis,后台 loop 里残留的协程可能继续占用 + SQLAlchemy QueuePool 连接,后续任务就会出现 QueuePool timeout。 + """ + try: + from app.services.redis_registry_service import close_registry_redis + + await close_registry_redis() + except Exception: + logger.debug("关闭 Celery Redis registry 连接失败", exc_info=True) + + try: + from app.models.base import engine + + await engine.dispose() + except Exception: + logger.debug("dispose Celery SQLAlchemy engine 失败", exc_info=True) + + def _runner_mode() -> str: mode = str(getattr(settings, "CELERY_ASYNC_RUNNER_MODE", "single_loop") or "single_loop").strip().lower() if mode not in {"single_loop", "direct"}: @@ -46,16 +71,18 @@ def _get_or_create_thread_local_loop() -> asyncio.AbstractEventLoop: def _single_loop_worker(loop: asyncio.AbstractEventLoop, ready: threading.Event) -> None: asyncio.set_event_loop(loop) ready.set() - loop.run_forever() + try: + loop.run_forever() + finally: + pending = [task for task in asyncio.all_tasks(loop) if not task.done()] + if pending: + for task in pending: + task.cancel() + loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True)) - pending = [task for task in asyncio.all_tasks(loop) if not task.done()] - if pending: - for task in pending: - task.cancel() - loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True)) - - loop.run_until_complete(loop.shutdown_asyncgens()) - loop.close() + loop.run_until_complete(_dispose_async_resources()) + loop.run_until_complete(loop.shutdown_asyncgens()) + loop.close() def _get_or_create_single_loop() -> asyncio.AbstractEventLoop: @@ -92,13 +119,28 @@ def _get_or_create_single_loop() -> asyncio.AbstractEventLoop: return _single_loop +def _cancel_future_and_reset_loop(future: Future[T] | None, *, reason: str) -> None: + """取消当前协程并重置当前进程内 event loop。""" + if future is not None and not future.done(): + future.cancel() + try: + future.result(timeout=2) + except Exception: + pass + + logger.warning("Celery async_runner 正在重置 event loop。reason=%s", reason) + close_loop() + + def run_async(coro: Awaitable[T]) -> T: """Celery 同步 task 调用异步协程的统一入口。 默认 single_loop 模式: - 一个 Celery 子进程只有一个专用 event loop; - 所有 asyncpg / redis.asyncio 操作都在这个 loop 内创建和使用; - - 避免 got Future attached to a different loop。 + - 避免 got Future attached to a different loop; + - 当 Celery soft time limit 打断 future.result() 时,主动 cancel 后台协程并 + 释放连接池,避免 QueuePool 被残留任务长期占用。 降级 direct 模式: - 兼容旧的线程本地 loop 方案; @@ -106,7 +148,15 @@ def run_async(coro: Awaitable[T]) -> T: """ if _runner_mode() == "direct": loop = _get_or_create_thread_local_loop() - return loop.run_until_complete(coro) + try: + return loop.run_until_complete(coro) + except BaseException: + try: + if not loop.is_closed(): + loop.run_until_complete(_dispose_async_resources()) + finally: + close_loop() + raise loop = _get_or_create_single_loop() try: @@ -118,7 +168,16 @@ def run_async(coro: Awaitable[T]) -> T: raise RuntimeError("run_async() 不能在 Celery async_runner 的事件循环内部被同步调用") future: Future[T] = asyncio.run_coroutine_threadsafe(coro, loop) - return future.result() + try: + return future.result() + except FutureTimeoutError: + _cancel_future_and_reset_loop(future, reason="future_result_timeout") + raise + except BaseException: + # Celery SoftTimeLimitExceeded/worker shutdown 等异常会从这里抛出。 + # 必须重置 loop,否则后台协程继续运行会拖住 DB 连接池。 + _cancel_future_and_reset_loop(future, reason="base_exception") + raise def close_loop() -> None: @@ -130,8 +189,16 @@ def close_loop() -> None: loop = _single_loop thread = _single_loop_thread if loop is not None and not loop.is_closed() and thread is not None and thread.is_alive(): - loop.call_soon_threadsafe(loop.stop) - thread.join(timeout=5) + try: + cleanup_future = asyncio.run_coroutine_threadsafe(_dispose_async_resources(), loop) + cleanup_future.result(timeout=5) + except Exception: + logger.debug("关闭 loop 前清理 async 资源失败", exc_info=True) + try: + loop.call_soon_threadsafe(loop.stop) + thread.join(timeout=5) + except Exception: + logger.debug("关闭 Celery async_runner loop 失败", exc_info=True) _single_loop = None _single_loop_thread = None @@ -141,6 +208,17 @@ def close_loop() -> None: # 关闭 direct 降级模式的线程本地 loop。 loop = getattr(_thread_local, "loop", None) if loop is not None and not loop.is_closed(): - loop.close() + try: + pending = [task for task in asyncio.all_tasks(loop) if not task.done()] + for task in pending: + task.cancel() + if pending: + loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True)) + loop.run_until_complete(_dispose_async_resources()) + loop.run_until_complete(loop.shutdown_asyncgens()) + except Exception: + logger.debug("关闭 direct loop 前清理失败", exc_info=True) + finally: + loop.close() _thread_local.loop = None _thread_local.pid = None diff --git a/video-gen-api/app/tasks/celery_app.py b/video-gen-api/app/tasks/celery_app.py index 9ef946f7..5a61f251 100644 --- a/video-gen-api/app/tasks/celery_app.py +++ b/video-gen-api/app/tasks/celery_app.py @@ -26,6 +26,9 @@ CELERY_TASK_IMPORTS = ( ) +RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or "gen_recovery" + + def _derive_redis_db(url: str, db_no: int) -> str: if not url: return url @@ -73,10 +76,12 @@ if broker_url: "shot_replicate.split_one_segment": {"queue": "gen_result_download"}, "shot_replicate.start_image_prompt_optimize": {"queue": "gen_chatapi_create"}, "shot_replicate.start_video_prompt_optimize": {"queue": "gen_chatapi_create"}, - "shot_replicate.recover_split_tasks_once": {"queue": "gen_result_download"}, - "generation.recover_download_tasks_once": {"queue": "gen_result_download"}, - "generation.recover_generation_tasks_once": {"queue": "gen_result_download"}, - "module_async.recover_module_async_tasks_once": {"queue": "gen_result_download"}, + # 恢复扫描统一走独立队列,避免占用下载/轮询/创建业务 worker。 + "recovery.startup_recovery_once": {"queue": RECOVERY_QUEUE}, + "shot_replicate.recover_split_tasks_once": {"queue": RECOVERY_QUEUE}, + "generation.recover_download_tasks_once": {"queue": RECOVERY_QUEUE}, + "generation.recover_generation_tasks_once": {"queue": RECOVERY_QUEUE}, + "module_async.recover_module_async_tasks_once": {"queue": RECOVERY_QUEUE}, "user_oauth.update_oauth_accounts": {"queue": "default"}, "app.tasks.cleanup.*": {"queue": "default"}, }, @@ -86,7 +91,7 @@ else: async def _try_acquire_startup_recovery_lock() -> bool: - """任意 worker 启动时都可尝试抢恢复锁,避免依赖 hostname 命名。""" + """任意 worker 启动时都可尝试抢恢复投递锁,避免依赖 hostname 命名。""" from app.services.redis_registry_service import redis_acquire_lock token = await redis_acquire_lock( @@ -103,9 +108,9 @@ def on_worker_ready(sender=None, **kwargs): 注意: - 不启用 Celery beat。 - - 不要求新增第四条启动命令。 - - 不再依赖 worker hostname 是否包含 gen_result_download。 - - 所有 worker 都尝试抢 Redis 锁,只有抢到锁的 worker 投递恢复任务。 + - 启动容灾保留,但只投递一个 recovery.startup_recovery_once 协调任务。 + - 协调任务走独立 gen_recovery 队列,串行扫描并把真实业务任务投回原队列。 + - 所有 worker 都尝试抢 Redis 投递锁,只有抢到锁的 worker 投递恢复任务。 """ if celery_app is None: return @@ -122,39 +127,21 @@ def on_worker_ready(sender=None, **kwargs): return try: - from app.tasks.generation_recovery_tasks import ( - recover_download_tasks_once, - recover_generation_tasks_once, - ) - from app.tasks.shot_replicate_tasks import recover_split_tasks_once - from app.tasks.module_async_recovery_tasks import recover_module_async_tasks_once_task + from app.tasks.generation_recovery_tasks import startup_recovery_once countdown = max(0, int(settings.DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS or 0)) - - recover_generation_tasks_once.apply_async( + startup_recovery_once.apply_async( countdown=countdown, - queue="gen_result_download", + queue=RECOVERY_QUEUE, priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, ) - recover_download_tasks_once.apply_async( - countdown=countdown + 5, - queue="gen_result_download", - priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, + logger.info( + "启动容灾恢复协调任务已投递。queue=%s countdown=%s", + RECOVERY_QUEUE, + countdown, ) - recover_split_tasks_once.apply_async( - countdown=countdown + 10, - queue="gen_result_download", - priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, - ) - recover_module_async_tasks_once_task.apply_async( - countdown=countdown + 15, - queue="gen_result_download", - priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, - ) - - logger.info("启动容灾恢复任务已投递。countdown=%s", countdown) except Exception: - logger.exception("启动容灾恢复任务投递失败") + logger.exception("启动容灾恢复协调任务投递失败") @worker_process_init.connect @@ -181,4 +168,4 @@ def on_worker_process_shutdown(**kwargs): except Exception: pass finally: - close_loop() \ No newline at end of file + close_loop() diff --git a/video-gen-api/app/tasks/generation_create_tasks.py b/video-gen-api/app/tasks/generation_create_tasks.py index 6a25c28e..94a23c0d 100644 --- a/video-gen-api/app/tasks/generation_create_tasks.py +++ b/video-gen-api/app/tasks/generation_create_tasks.py @@ -290,7 +290,13 @@ async def _run(task_id: str): if celery_app: @celery_app.task(name="generation.chatapi_create_generation_task", bind=True, max_retries=3, default_retry_delay=30) def chatapi_create_generation_task(self, task_id: str): - return run_async(_run(task_id)) + try: + return run_async(_run(task_id)) + except Exception as exc: + # 只处理 run_async/连接池/worker 中断等基础设施异常;业务异常已在 _run 内落库并退款。 + retries = int(getattr(self.request, "retries", 0) or 0) + 1 + countdown = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * max(1, retries) + raise self.retry(exc=exc, countdown=countdown) else: class _DisabledTask: def delay(self, *args, **kwargs): diff --git a/video-gen-api/app/tasks/generation_download_tasks.py b/video-gen-api/app/tasks/generation_download_tasks.py index 78968a79..118170af 100644 --- a/video-gen-api/app/tasks/generation_download_tasks.py +++ b/video-gen-api/app/tasks/generation_download_tasks.py @@ -598,7 +598,13 @@ async def _run(task_id: str): if celery_app: @celery_app.task(name="generation.download_generation_result_task", bind=True, max_retries=3, default_retry_delay=30) def download_generation_result_task(self, task_id: str): - return run_async(_run(task_id)) + try: + return run_async(_run(task_id)) + except Exception as exc: + # 只重试 run_async/连接池/worker 中断等基础设施异常;下载业务异常已在 _run 内写入 retry_waiting。 + retries = int(getattr(self.request, "retries", 0) or 0) + 1 + countdown = int(settings.DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS or 30) * max(1, retries) + raise self.retry(exc=exc, countdown=countdown) else: class _DisabledTask: def delay(self, *args, **kwargs): diff --git a/video-gen-api/app/tasks/generation_poll_tasks.py b/video-gen-api/app/tasks/generation_poll_tasks.py index ac4aab17..575bb3fe 100644 --- a/video-gen-api/app/tasks/generation_poll_tasks.py +++ b/video-gen-api/app/tasks/generation_poll_tasks.py @@ -194,7 +194,8 @@ async def _run(task_id: str): await remove_poll_active(task.id) return - if _deadline_expired(task): + final_poll_before_timeout = _deadline_expired(task) + if final_poll_before_timeout and not (task.seedance_task_id or task.provider_task_id): await _mark_timeout(db, task, message="任务轮询超时") return @@ -202,6 +203,14 @@ async def _run(task_id: str): await _mark_failed(db, task, message="缺少外部任务ID") return + if final_poll_before_timeout: + await log_task_event( + task, + event_type="FINAL_POLL_BEFORE_TIMEOUT", + message="任务已到 deadline,执行最后一次供应商查询后再判定超时", + detail={"deadline_at": task.deadline_at, "stage": task.pipeline_stage}, + ) + # 标记本次正在轮询,并登记 poll lease。 # 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递。 task.pipeline_stage = "polling" @@ -274,6 +283,16 @@ async def _run(task_id: str): ) return + if final_poll_before_timeout: + await log_task_event( + task, + event_type="FINAL_POLL_BEFORE_TIMEOUT_PENDING", + message=f"最终查询后供应商仍未完成,按超时处理。status={status}", + detail=poll_result, + ) + await _mark_timeout(db, task, message="任务轮询超时") + return + # 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。 # 同时登记下一次 poll active,Celery countdown 丢失时可由恢复任务拉起。 task.pipeline_stage = "waiting_remote" @@ -309,6 +328,15 @@ async def _run(task_id: str): await remove_poll_active(task_id) return + if final_poll_before_timeout: + await log_task_event( + task, + event_type="FINAL_POLL_BEFORE_TIMEOUT_ERROR", + message=str(exc), + ) + await _mark_timeout(db, task, message="任务轮询超时") + return + task.retry_count = (task.retry_count or 0) + 1 if task.retry_count > settings.CHATAPI_ASYNC_MAX_RETRIES: @@ -339,7 +367,13 @@ async def _run(task_id: str): if celery_app: @celery_app.task(name="generation.poll_generation_task", bind=True, max_retries=3, default_retry_delay=30) def poll_generation_task(self, task_id: str): - return run_async(_run(task_id)) + try: + return run_async(_run(task_id)) + except Exception as exc: + # 只重试基础设施异常;供应商失败/业务失败已在 _run 内处理。 + retries = int(getattr(self.request, "retries", 0) or 0) + 1 + countdown = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * max(1, retries) + raise self.retry(exc=exc, countdown=countdown) else: class _DisabledTask: def delay(self, *args, **kwargs): diff --git a/video-gen-api/app/tasks/generation_recovery_tasks.py b/video-gen-api/app/tasks/generation_recovery_tasks.py index dee9039c..d96d8ea2 100644 --- a/video-gen-api/app/tasks/generation_recovery_tasks.py +++ b/video-gen-api/app/tasks/generation_recovery_tasks.py @@ -2,16 +2,19 @@ from __future__ import annotations import logging -from typing import Any, Dict +from typing import Any, Awaitable, Callable, 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.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock from app.tasks.async_runner import run_async from app.tasks.celery_app import celery_app logger = logging.getLogger("video_gen") +RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or "gen_recovery" +RecoveryRunner = Callable[[], Awaitable[Dict[str, Any]]] + async def _run_download_once() -> Dict[str, Any]: from app.services.generation_recovery_service import recover_download_tasks_once @@ -27,6 +30,55 @@ async def _run_generation_once() -> Dict[str, Any]: return await recover_generation_tasks_once(db) +async def _run_module_async_once() -> Dict[str, Any]: + from app.services.module_async_recovery_service import recover_module_async_tasks_once + + async with async_session() as db: + return await recover_module_async_tasks_once(db) + + +async def _run_shot_split_once() -> Dict[str, Any]: + from app.services.shot_replicate_recovery_service import recover_shot_split_tasks_once + + async with async_session() as db: + return await recover_shot_split_tasks_once(db) + + +async def _run_with_execution_lock( + *, + lock_key: str, + log_context: str, + runner: RecoveryRunner, +) -> Dict[str, Any]: + """恢复任务执行锁。 + + worker_ready 的启动锁只保证“只投递一次”;如果 broker 中残留旧消息, + 或者人工手动触发恢复任务,仍可能并发执行。这里再加执行锁,避免多个 + 恢复扫描同时扫库、抢行锁、抢连接池。 + """ + redis = await get_registry_redis() + token: str | None = None + if redis is not None: + token = await redis_acquire_lock( + lock_key=lock_key, + ttl_seconds=int(settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600), + log_context=log_context, + ) + if not token: + return {"skipped": "lock_held", "lock_key": lock_key} + else: + # Redis 不可用时仍允许 DB fallback 执行一次,避免恢复能力彻底失效。 + logger.warning("恢复任务执行锁不可用,降级直接执行。context=%s", log_context) + + try: + result = await runner() + result["execution_lock"] = "lock_acquired" if token else "redis_unavailable_run_db_fallback" + return result + finally: + if token: + await redis_release_lock(lock_key=lock_key, token=token, log_context=log_context) + + async def _acquire_download_recovery_loop_lock() -> tuple[bool, str]: """下载恢复循环锁。 @@ -45,36 +97,128 @@ async def _acquire_download_recovery_loop_lock() -> tuple[bool, str]: def _schedule_next_download_recovery_loop() -> None: - if not celery_app or not bool(getattr(settings, "DOWNLOAD_RECOVERY_LOOP_ENABLED", True)): + if not celery_app or not bool(getattr(settings, "DOWNLOAD_RECOVERY_LOOP_ENABLED", False)): return try: recover_download_tasks_once.apply_async( countdown=max(1, int(settings.DOWNLOAD_RECOVERY_INTERVAL_SECONDS or 60)), - queue="gen_result_download", + queue=RECOVERY_QUEUE, priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, ) except Exception: logger.exception("下载恢复循环下一轮投递失败") +async def _run_startup_recovery_once() -> Dict[str, Any]: + """启动容灾协调器:串行跑恢复扫描。 + + 真实业务任务仍投递回原队列: + - 创建/提词/视频分析 -> gen_chatapi_create + - provider poll -> gen_provider_poll + - 下载/ffmpeg 切片 -> gen_result_download + 恢复扫描本身只走 gen_recovery,避免堵住业务 worker。 + """ + return await _run_with_execution_lock( + lock_key=settings.CELERY_RECOVERY_STARTUP_TASK_LOCK_KEY, + log_context="startup_recovery_once", + runner=_run_startup_recovery_steps, + ) + + +async def _run_startup_recovery_steps() -> Dict[str, Any]: + results: Dict[str, Any] = {} + + steps: list[tuple[str, str, str, RecoveryRunner]] = [ + ( + "module_async", + settings.MODULE_ASYNC_RECOVERY_LOCK_KEY, + "module_async_recovery", + _run_module_async_once, + ), + ( + "shot_split", + settings.SHOT_SPLIT_RECOVERY_LOCK_KEY, + "shot_split_recovery", + _run_shot_split_once, + ), + ( + "generation", + settings.GENERATION_RECOVERY_LOCK_KEY, + "generation_recovery", + _run_generation_once, + ), + ( + "download", + settings.DOWNLOAD_RECOVERY_LOCK_KEY, + "download_recovery", + _run_download_once, + ), + ] + + for name, lock_key, log_context, runner in steps: + try: + results[name] = await _run_with_execution_lock( + lock_key=lock_key, + log_context=log_context, + runner=runner, + ) + except Exception as exc: + logger.exception("启动容灾步骤执行失败。step=%s", name) + results[name] = {"error": str(exc)} + + return {"steps": results} + + if celery_app: - @celery_app.task(name="generation.recover_download_tasks_once", bind=True) + @celery_app.task( + name="recovery.startup_recovery_once", + bind=True, + soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS, + time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS, + ) + def startup_recovery_once(self) -> Dict[str, Any]: + return run_async(_run_startup_recovery_once()) + + + @celery_app.task( + name="generation.recover_download_tasks_once", + bind=True, + soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS, + time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS, + ) 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 = run_async( + _run_with_execution_lock( + lock_key=settings.DOWNLOAD_RECOVERY_LOCK_KEY, + log_context="download_recovery", + runner=_run_download_once, + ) + ) result["loop_lock"] = reason return result finally: _schedule_next_download_recovery_loop() - @celery_app.task(name="generation.recover_generation_tasks_once") - def recover_generation_tasks_once() -> Dict[str, Any]: - return run_async(_run_generation_once()) + @celery_app.task( + name="generation.recover_generation_tasks_once", + bind=True, + soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS, + time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS, + ) + def recover_generation_tasks_once(self) -> Dict[str, Any]: + return run_async( + _run_with_execution_lock( + lock_key=settings.GENERATION_RECOVERY_LOCK_KEY, + log_context="generation_recovery", + runner=_run_generation_once, + ) + ) else: @@ -85,5 +229,6 @@ else: def apply_async(self, *args: Any, **kwargs: Any) -> None: raise RuntimeError("Celery is disabled") + startup_recovery_once = _DisabledTask() recover_download_tasks_once = _DisabledTask() recover_generation_tasks_once = _DisabledTask() diff --git a/video-gen-api/app/tasks/module_async_recovery_tasks.py b/video-gen-api/app/tasks/module_async_recovery_tasks.py index 6025eaec..9dd37677 100644 --- a/video-gen-api/app/tasks/module_async_recovery_tasks.py +++ b/video-gen-api/app/tasks/module_async_recovery_tasks.py @@ -2,21 +2,49 @@ from __future__ import annotations from typing import Any +from app.config import settings from app.models.base import async_session from app.services.module_async_recovery_service import recover_module_async_tasks_once +from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock from app.tasks.async_runner import run_async from app.tasks.celery_app import celery_app async def _run_recover_module_async_tasks_once() -> dict[str, Any]: - async with async_session() as db: - return await recover_module_async_tasks_once(db) + redis = await get_registry_redis() + token: str | None = None + if redis is not None: + token = await redis_acquire_lock( + lock_key=settings.MODULE_ASYNC_RECOVERY_LOCK_KEY, + ttl_seconds=int(settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600), + log_context="module_async_recovery", + ) + if not token: + return {"skipped": "lock_held", "lock_key": settings.MODULE_ASYNC_RECOVERY_LOCK_KEY} + + try: + async with async_session() as db: + result = await recover_module_async_tasks_once(db) + result["execution_lock"] = "lock_acquired" if token else "redis_unavailable_run_db_fallback" + return result + finally: + if token: + await redis_release_lock( + lock_key=settings.MODULE_ASYNC_RECOVERY_LOCK_KEY, + token=token, + log_context="module_async_recovery", + ) if celery_app: - @celery_app.task(name="module_async.recover_module_async_tasks_once") - def recover_module_async_tasks_once_task() -> dict[str, Any]: + @celery_app.task( + name="module_async.recover_module_async_tasks_once", + bind=True, + soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS, + time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS, + ) + def recover_module_async_tasks_once_task(self) -> dict[str, Any]: return run_async(_run_recover_module_async_tasks_once()) else: diff --git a/video-gen-api/app/tasks/shot_replicate_tasks.py b/video-gen-api/app/tasks/shot_replicate_tasks.py index 8c5a5c73..db149962 100644 --- a/video-gen-api/app/tasks/shot_replicate_tasks.py +++ b/video-gen-api/app/tasks/shot_replicate_tasks.py @@ -20,7 +20,7 @@ from app.models.base import async_session from app.models.shot_replicate_segment import ShotReplicateSegment from app.models.shot_replicate_task_set import ShotReplicateTaskSet from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event -from app.services.redis_registry_service import redis_acquire_lock, redis_release_lock +from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock from app.services.module_async_recovery_service import ( OBJECT_SHOT_SEGMENT_ANALYSIS, OBJECT_SHOT_SPLIT_SEGMENT, @@ -561,8 +561,29 @@ async def _run_split_one_segment(segment_id: str) -> None: async def _run_recover_split_tasks_once() -> dict[str, Any]: from app.services.shot_replicate_recovery_service import recover_shot_split_tasks_once - async with async_session() as db: - return await recover_shot_split_tasks_once(db) + redis = await get_registry_redis() + token: str | None = None + if redis is not None: + token = await redis_acquire_lock( + lock_key=settings.SHOT_SPLIT_RECOVERY_LOCK_KEY, + ttl_seconds=int(settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600), + log_context="shot_split_recovery", + ) + if not token: + return {"skipped": "lock_held", "lock_key": settings.SHOT_SPLIT_RECOVERY_LOCK_KEY} + + try: + async with async_session() as db: + result = await recover_shot_split_tasks_once(db) + result["execution_lock"] = "lock_acquired" if token else "redis_unavailable_run_db_fallback" + return result + finally: + if token: + await redis_release_lock( + lock_key=settings.SHOT_SPLIT_RECOVERY_LOCK_KEY, + token=token, + log_context="shot_split_recovery", + ) if celery_app: @@ -582,8 +603,13 @@ if celery_app: return run_async(_run_analyze_custom_segment_video(segment_id)) - @celery_app.task(name="shot_replicate.recover_split_tasks_once") - def recover_split_tasks_once() -> dict[str, Any]: + @celery_app.task( + name="shot_replicate.recover_split_tasks_once", + bind=True, + soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS, + time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS, + ) + def recover_split_tasks_once(self) -> dict[str, Any]: return run_async(_run_recover_split_tasks_once()) else: From 9bf181e14cd3bbce91f2f8048d36cbe1a72a551b Mon Sep 17 00:00:00 2001 From: GinHa <15201596918@163.com> Date: Fri, 26 Jun 2026 13:59:58 +0800 Subject: [PATCH 3/3] =?UTF-8?q?=E4=BF=AE=E5=A4=8Dcelery=E5=BC=82=E6=AD=A5?= =?UTF-8?q?=E8=BE=93=E5=87=BAjson=E5=BA=8F=E5=88=97=E5=8C=96=E5=BC=82?= =?UTF-8?q?=E5=B8=B8BUG?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- video-gen-api/app/tasks/celery_app.py | 16 ++++++++++++- .../app/tasks/hot_opening_replicate_tasks.py | 24 ++++++++++--------- .../app/tasks/shot_replicate_flow_tasks.py | 24 ++++++++++--------- 3 files changed, 41 insertions(+), 23 deletions(-) diff --git a/video-gen-api/app/tasks/celery_app.py b/video-gen-api/app/tasks/celery_app.py index 5a61f251..94c777ee 100644 --- a/video-gen-api/app/tasks/celery_app.py +++ b/video-gen-api/app/tasks/celery_app.py @@ -58,6 +58,20 @@ if broker_url: task_acks_late=True, task_reject_on_worker_lost=True, task_track_started=True, + task_annotations={ + # 生成链路任务以数据库状态为准,不依赖 Celery result backend。 + # 这里忽略结果可避免任务误返回 ORM / 非 JSON 对象时触发结果序列化失败。 + # "generation.chatapi_create_generation_task": {"ignore_result": True}, + # "generation.poll_generation_task": {"ignore_result": True}, + # "generation.download_generation_result_task": {"ignore_result": True}, + "hot_opening.start_image_prompt_optimize": {"ignore_result": True}, + "hot_opening.start_video_prompt_optimize": {"ignore_result": True}, + "shot_replicate.analyze_original_video": {"ignore_result": True}, + "shot_replicate.analyze_custom_segment_video": {"ignore_result": True}, + "shot_replicate.split_one_segment": {"ignore_result": True}, + "shot_replicate.start_image_prompt_optimize": {"ignore_result": True}, + "shot_replicate.start_video_prompt_optimize": {"ignore_result": True}, + }, worker_prefetch_multiplier=1, broker_transport_options={ "visibility_timeout": 3600, @@ -168,4 +182,4 @@ def on_worker_process_shutdown(**kwargs): except Exception: pass finally: - close_loop() + close_loop() \ No newline at end of file diff --git a/video-gen-api/app/tasks/hot_opening_replicate_tasks.py b/video-gen-api/app/tasks/hot_opening_replicate_tasks.py index 9b172f97..f94db7c7 100644 --- a/video-gen-api/app/tasks/hot_opening_replicate_tasks.py +++ b/video-gen-api/app/tasks/hot_opening_replicate_tasks.py @@ -21,7 +21,7 @@ from app.tasks.celery_app import celery_app MODULE = ModuleCodeEnum.HOT_OPENING_REPLICATE.value -async def _run_image_prompt(project_id: str, step_id: str | None = None): +async def _run_image_prompt(project_id: str, step_id: str | None = None) -> None: lock_token: str | None = None if step_id: lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id) @@ -38,17 +38,17 @@ async def _run_image_prompt(project_id: str, step_id: str | None = None): await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id) try: async with async_session() as db: - result = await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id) + await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id) await db.commit() if step_id: await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id) - return result + return None finally: if step_id: await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token) -async def _run_video_prompt(project_id: str, step_id: str | None = None): +async def _run_video_prompt(project_id: str, step_id: str | None = None) -> None: lock_token: str | None = None if step_id: lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id) @@ -64,18 +64,18 @@ async def _run_video_prompt(project_id: str, step_id: str | None = None): await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id) try: async with async_session() as db: - result = await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id) + await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id) await db.commit() if step_id: await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id) - return result + return None finally: if step_id: await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token) if celery_app: - @celery_app.task(name="hot_opening.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30) + @celery_app.task(name="hot_opening.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30, ignore_result=True) def start_image_prompt_optimize(self, project_id: str, step_id: str | None = None): """手动触发后的图片 AI 提词任务。 @@ -84,11 +84,12 @@ if celery_app: service 内部已经落库为业务失败的情况不会抛出异常,也不会重复 retry。 """ try: - return run_async(_run_image_prompt(project_id, step_id)) + run_async(_run_image_prompt(project_id, step_id)) + return None except Exception as exc: raise self.retry(exc=exc) from exc - @celery_app.task(name="hot_opening.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30) + @celery_app.task(name="hot_opening.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30, ignore_result=True) def start_video_prompt_optimize(self, project_id: str, step_id: str | None = None): """手动触发后的视频 AI 提词任务。 @@ -97,7 +98,8 @@ if celery_app: service 内部已经落库为业务失败的情况不会抛出异常,也不会重复 retry。 """ try: - return run_async(_run_video_prompt(project_id, step_id)) + run_async(_run_video_prompt(project_id, step_id)) + return None except Exception as exc: raise self.retry(exc=exc) from exc else: @@ -109,4 +111,4 @@ else: raise RuntimeError("Celery is disabled") start_image_prompt_optimize = _DisabledTask() - start_video_prompt_optimize = _DisabledTask() + start_video_prompt_optimize = _DisabledTask() \ No newline at end of file diff --git a/video-gen-api/app/tasks/shot_replicate_flow_tasks.py b/video-gen-api/app/tasks/shot_replicate_flow_tasks.py index 6c5105ef..a60d6a10 100644 --- a/video-gen-api/app/tasks/shot_replicate_flow_tasks.py +++ b/video-gen-api/app/tasks/shot_replicate_flow_tasks.py @@ -21,7 +21,7 @@ from app.tasks.celery_app import celery_app MODULE = ModuleCodeEnum.SHOT_REPLICATE.value -async def _run_image_prompt(project_id: str, step_id: str | None = None): +async def _run_image_prompt(project_id: str, step_id: str | None = None) -> None: lock_token: str | None = None if step_id: lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id) @@ -37,17 +37,17 @@ async def _run_image_prompt(project_id: str, step_id: str | None = None): await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id) try: async with async_session() as db: - result = await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id) + await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id) await db.commit() if step_id: await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id) - return result + return None finally: if step_id: await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token) -async def _run_video_prompt(project_id: str, step_id: str | None = None): +async def _run_video_prompt(project_id: str, step_id: str | None = None) -> None: lock_token: str | None = None if step_id: lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id) @@ -63,11 +63,11 @@ async def _run_video_prompt(project_id: str, step_id: str | None = None): await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id) try: async with async_session() as db: - result = await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id) + await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id) await db.commit() if step_id: await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id) - return result + return None finally: if step_id: await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token) @@ -75,7 +75,7 @@ async def _run_video_prompt(project_id: str, step_id: str | None = None): if celery_app: - @celery_app.task(name="shot_replicate.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30) + @celery_app.task(name="shot_replicate.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30, ignore_result=True) def start_image_prompt_optimize(self, project_id: str, step_id: str | None = None): """手动触发后的图片 AI 提词任务。 @@ -84,12 +84,13 @@ if celery_app: service 内部已经落库为业务失败的情况不会抛出异常,也不会重复 retry。 """ try: - return run_async(_run_image_prompt(project_id, step_id)) + run_async(_run_image_prompt(project_id, step_id)) + return None except Exception as exc: raise self.retry(exc=exc) from exc - @celery_app.task(name="shot_replicate.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30) + @celery_app.task(name="shot_replicate.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30, ignore_result=True) def start_video_prompt_optimize(self, project_id: str, step_id: str | None = None): """手动触发后的视频 AI 提词任务。 @@ -98,7 +99,8 @@ if celery_app: service 内部已经落库为业务失败的情况不会抛出异常,也不会重复 retry。 """ try: - return run_async(_run_video_prompt(project_id, step_id)) + run_async(_run_video_prompt(project_id, step_id)) + return None except Exception as exc: raise self.retry(exc=exc) from exc @@ -112,4 +114,4 @@ else: raise RuntimeError("Celery is disabled") start_image_prompt_optimize = _DisabledTask() - start_video_prompt_optimize = _DisabledTask() + start_video_prompt_optimize = _DisabledTask() \ No newline at end of file