From f6c5032c7b3d2a6d9674a23ebb339f6a0c354a9f Mon Sep 17 00:00:00 2001 From: GinHa <15201596918@163.com> Date: Thu, 2 Jul 2026 09:36:38 +0800 Subject: [PATCH] =?UTF-8?q?celery=20worker=20poll=20=E9=A2=91=E6=AC=A1?= =?UTF-8?q?=E6=9C=BA=E5=88=B6=E8=B0=83=E6=95=B4|=20celery=20beat=20?= =?UTF-8?q?=E8=AE=BE=E7=BD=AEpoll=E6=A3=80=E6=B5=8B=E4=BB=BB=E5=8A=A1|=20?= =?UTF-8?q?=E6=8B=86=E9=95=9C=E5=A4=8D=E5=88=BB=E5=88=87=E7=89=87=E5=88=A0?= =?UTF-8?q?=E9=99=A4API=20|=20=E4=BA=A4=E6=98=93=E6=B5=81=E6=B0=B4?= =?UTF-8?q?=E6=97=B6=E5=8C=BABUG?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ...dd_chat_generation_poll_schedule_fields.py | 83 ++++++++ video-gen-api/app/api/v1/shot_replicate.py | 30 +++ video-gen-api/app/config.py | 17 +- video-gen-api/app/enums/__init__.py | 1 + video-gen-api/app/enums/celery_queue.py | 21 ++ video-gen-api/app/enums/generation_task.py | 10 + video-gen-api/app/models/__init__.py | 3 +- .../app/models/chat_generation_task.py | 16 ++ video-gen-api/app/schemas/shot_replicate.py | 10 + .../services/admin_credit_record_service.py | 14 +- .../app/services/generation_ai_service.py | 2 +- .../generation_poll_schedule_service.py | 153 ++++++++++++++ .../services/generation_recovery_service.py | 176 ++++++++++++++-- .../generation_task_factory_service.py | 2 +- .../module_generation_flow_base_service.py | 8 +- .../services/resource_accounting_service.py | 1 + .../services/shot_replicate_flow_service.py | 99 ++++++++- .../shot_replicate_taskset_service.py | 101 +++++++++ video-gen-api/app/tasks/celery_app.py | 60 ++++-- .../app/tasks/generation_create_tasks.py | 98 +++++---- .../app/tasks/generation_poll_tasks.py | 197 ++++++++++++++---- .../app/tasks/generation_recovery_tasks.py | 61 +++++- 22 files changed, 1039 insertions(+), 124 deletions(-) create mode 100644 video-gen-api/alembic/versions/ee4b42e57960_add_chat_generation_poll_schedule_fields.py create mode 100644 video-gen-api/app/enums/celery_queue.py create mode 100644 video-gen-api/app/services/generation_poll_schedule_service.py diff --git a/video-gen-api/alembic/versions/ee4b42e57960_add_chat_generation_poll_schedule_fields.py b/video-gen-api/alembic/versions/ee4b42e57960_add_chat_generation_poll_schedule_fields.py new file mode 100644 index 00000000..70bd7e09 --- /dev/null +++ b/video-gen-api/alembic/versions/ee4b42e57960_add_chat_generation_poll_schedule_fields.py @@ -0,0 +1,83 @@ +"""add chat generation poll schedule fields + +Revision ID: ee4b42e57960 +Revises: f7a3b2c1d4e5 +Create Date: 2026-07-01 16:11:28.799227 +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = 'ee4b42e57960' +down_revision: Union[str, None] = 'f7a3b2c1d4e5' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + op.add_column( + 'chat_generation_tasks', + sa.Column('poll_started_at', sa.DateTime(timezone=True), nullable=True), + ) + op.add_column( + 'chat_generation_tasks', + sa.Column('next_poll_at', sa.DateTime(timezone=True), nullable=True), + ) + op.add_column( + 'chat_generation_tasks', + sa.Column( + 'poll_interval_seconds', + sa.Integer(), + nullable=False, + server_default='0', + comment='当前轮询退避间隔秒数', + ), + ) + + # 修复历史“生成中视频任务”: + # 1. deadline_at 改为 created_at + 24 小时,避免沿用旧 30 分钟最终超时。 + # 2. next_poll_at 设置为 now(),让 Beat dispatcher 能尽快接管。 + # 3. poll_started_at 用 last_poll_at / created_at 兜底。 + op.execute( + """ + UPDATE chat_generation_tasks + SET + poll_started_at = COALESCE(poll_started_at, last_poll_at, created_at, NOW()), + next_poll_at = COALESCE(next_poll_at, NOW()), + poll_interval_seconds = COALESCE(poll_interval_seconds, 0), + deadline_at = CASE + WHEN deadline_at IS NULL THEN created_at + INTERVAL '24 hours' + WHEN deadline_at < created_at + INTERVAL '24 hours' THEN created_at + INTERVAL '24 hours' + ELSE deadline_at + END + WHERE deleted_at IS NULL + AND status = 'generating' + AND gen_type = 'video' + """ + ) + + op.create_index( + 'idx_chat_generation_tasks_next_poll_at', + 'chat_generation_tasks', + ['next_poll_at'], + unique=False, + postgresql_where=sa.text( + "deleted_at IS NULL " + "AND status = 'generating' " + "AND gen_type = 'video' " + "AND next_poll_at IS NOT NULL" + ), + ) + + +def downgrade() -> None: + op.drop_index( + 'idx_chat_generation_tasks_next_poll_at', + table_name='chat_generation_tasks', + ) + op.drop_column('chat_generation_tasks', 'poll_interval_seconds') + op.drop_column('chat_generation_tasks', 'next_poll_at') + op.drop_column('chat_generation_tasks', 'poll_started_at') diff --git a/video-gen-api/app/api/v1/shot_replicate.py b/video-gen-api/app/api/v1/shot_replicate.py index d1720205..b64019b2 100644 --- a/video-gen-api/app/api/v1/shot_replicate.py +++ b/video-gen-api/app/api/v1/shot_replicate.py @@ -31,6 +31,7 @@ from app.schemas.shot_replicate import ( ShotReplicateSpecOut, ShotReplicateTaskDetailOut, ShotReplicateVideoPromptSchemaUpdateRequest, + ShotSegmentDeleteOut, ShotSegmentDetailOut, ShotSegmentListOut, ShotSegmentReplicationCreateRequest, @@ -60,6 +61,7 @@ from app.services.shot_replicate_taskset_service import ( create_custom_segment, create_segments_by_ai, create_task_set, + delete_segment, get_segment_for_user, list_segments, list_task_sets, @@ -455,6 +457,34 @@ async def get_segment( return await segment_detail(db, current_user=current_user, segment_id=segment_id) +@router.delete( + "/segments/{segment_id}", + response_model=ShotSegmentDeleteOut, + summary="软删除拆镜片段", + description=( + "软删除单个拆镜片段;不删除 segment_video_path 指向的物理文件。" + "如片段已创建拆镜复刻项目,会联动软删除该项目和已生成资源账本,但用户主动删除不退款。" + "如片段或关联项目仍有处理中任务,会拒绝删除。" + ), +) +async def delete_shot_segment( + segment_id: str = Path(..., description="拆镜片段ID,即 shot_replicate_segments.id"), + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + try: + out = await delete_segment(db, current_user=current_user, segment_id=segment_id) + await db.commit() + return out + except HTTPException: + await db.rollback() + raise + except Exception as exc: + await db.rollback() + _log_api_exception_from_locals(exc, locals(), f"删除拆镜片段失败: {exc}") + raise HTTPException(status_code=500, detail=f"删除拆镜片段失败: {exc}") + + @router.post( "/segments/{segment_id}/replication-projects", response_model=ShotReplicateActionOut, diff --git a/video-gen-api/app/config.py b/video-gen-api/app/config.py index 727c69eb..9ded341f 100644 --- a/video-gen-api/app/config.py +++ b/video-gen-api/app/config.py @@ -121,7 +121,14 @@ class Settings(BaseSettings): CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS: int = 30 CHATAPI_ASYNC_POLL_INTERVAL_SECONDS: int = 30 CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES: int = 10 - CHATAPI_ASYNC_VIDEO_DEADLINE_MINUTES: int = 30 + # 视频异步生成不再使用 30 分钟最终超时;前 10 分钟高频轮询,之后降频,24 小时最后判定失败才退款。 + CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS: int = 24 + CHATAPI_ASYNC_VIDEO_HIGH_FREQ_MINUTES: int = 10 + CHATAPI_ASYNC_VIDEO_HIGH_FREQ_POLL_SECONDS: int = 30 + CHATAPI_ASYNC_VIDEO_BACKOFF_INITIAL_SECONDS: int = 60 + CHATAPI_ASYNC_VIDEO_BACKOFF_MULTIPLIER: int = 2 + CHATAPI_ASYNC_VIDEO_BACKOFF_MAX_SECONDS: int = 3600 + CHATAPI_ASYNC_VIDEO_DIRECT_COUNTDOWN_MAX_SECONDS: int = 300 # Distributed provider concurrency limits. 0 means disabled/no-op. ARK_CHAT_PROMPT_MAX_CONCURRENCY: int = 20 @@ -184,6 +191,14 @@ class Settings(BaseSettings): MODULE_ASYNC_RECOVERY_LOCK_KEY: str = "vg:celery:module_async_recovery_lock" SHOT_SPLIT_RECOVERY_LOCK_KEY: str = "vg:celery:shot_split_recovery_lock" + # 视频到期轮询调度。 + # Celery Beat 每分钟投递轻量 dispatcher 到 gen_recovery;dispatcher 只扫描 next_poll_at 到期的视频任务。 + POLL_DUE_DISPATCH_ENABLED: bool = True + POLL_DUE_DISPATCH_INTERVAL_SECONDS: int = 60 + POLL_DUE_DISPATCH_BATCH_SIZE: int = 100 + POLL_DUE_DISPATCH_LOCK_KEY: str = "vg:celery:poll_due_dispatch_lock" + POLL_DUE_DISPATCH_LOCK_TTL_SECONDS: int = 55 + # 模块异步任务容灾配置。 # 覆盖 ModuleGenerationStep 提词任务、shot 原视频/片段分析、shot ffmpeg 切割 active 注册。 # 恢复扫描走 CELERY_RECOVERY_QUEUE,真实业务任务回到原始队列。 diff --git a/video-gen-api/app/enums/__init__.py b/video-gen-api/app/enums/__init__.py index 0c6efa60..1bec065a 100644 --- a/video-gen-api/app/enums/__init__.py +++ b/video-gen-api/app/enums/__init__.py @@ -14,3 +14,4 @@ from app.enums.notification import * from app.enums.resource_capacity import * from app.enums.team import * from app.enums.home_material import * +from app.enums.celery_queue import * diff --git a/video-gen-api/app/enums/celery_queue.py b/video-gen-api/app/enums/celery_queue.py new file mode 100644 index 00000000..4fea9222 --- /dev/null +++ b/video-gen-api/app/enums/celery_queue.py @@ -0,0 +1,21 @@ +from enum import Enum + + +class CeleryQueue(str, Enum): + GEN_CHATAPI_CREATE = "gen_chatapi_create" + GEN_PROVIDER_POLL = "gen_provider_poll" + GEN_RESULT_DOWNLOAD = "gen_result_download" + GEN_RECOVERY = "gen_recovery" + DEFAULT = "default" + + +class CeleryTaskName(str, Enum): + CHATAPI_CREATE = "generation.chatapi_create_generation_task" + POLL_GENERATION = "generation.poll_generation_task" + DOWNLOAD_GENERATION_RESULT = "generation.download_generation_result_task" + RECOVER_DOWNLOAD = "generation.recover_download_tasks_once" + RECOVER_GENERATION = "generation.recover_generation_tasks_once" + DISPATCH_DUE_POLL = "generation.dispatch_due_poll_tasks" + STARTUP_RECOVERY = "recovery.startup_recovery_once" + MODULE_ASYNC_RECOVERY = "module_async.recover_module_async_tasks_once" + SHOT_SPLIT_RECOVERY = "shot_replicate.recover_split_tasks_once" diff --git a/video-gen-api/app/enums/generation_task.py b/video-gen-api/app/enums/generation_task.py index 9a264566..e7f3a091 100644 --- a/video-gen-api/app/enums/generation_task.py +++ b/video-gen-api/app/enums/generation_task.py @@ -45,12 +45,18 @@ class ChatGenerationTaskEventType(str, Enum): POLL_START = "POLL_START" POLL_PENDING = "POLL_PENDING" + POLL_SCHEDULED = "POLL_SCHEDULED" + POLL_SKIP_NOT_DUE = "POLL_SKIP_NOT_DUE" + POLL_DISPATCH_DUE = "POLL_DISPATCH_DUE" + POLL_DISPATCH_SKIP = "POLL_DISPATCH_SKIP" POLL_SUCCESS = "POLL_SUCCESS" POLL_FAILED = "POLL_FAILED" POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY = "POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY" + FINAL_POLL_BEFORE_TIMEOUT = "FINAL_POLL_BEFORE_TIMEOUT" 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" + GENERATION_RECOVERY_TIMEOUT = "GENERATION_RECOVERY_TIMEOUT" DOWNLOAD_ENQUEUE = "DOWNLOAD_ENQUEUE" DOWNLOAD_ENQUEUE_FAILED = "DOWNLOAD_ENQUEUE_FAILED" @@ -78,6 +84,7 @@ class ChatGenerationTaskEventType(str, Enum): DOWNLOAD_SKIP_DISABLED = "DOWNLOAD_SKIP_DISABLED" TASK_TIMEOUT = "TASK_TIMEOUT" + TASK_FAILED = "TASK_FAILED" ALLOWED_GENERATION_MODES = { @@ -99,3 +106,6 @@ DOWNLOAD_RECOVERABLE_STAGES = { ChatGenerationPipelineStage.DOWNLOADING.value, ChatGenerationPipelineStage.RETRY_WAITING.value, } + +PROVIDER_SUCCESS_STATUSES = {"succeeded", "success", "completed", "done"} +PROVIDER_FAILED_STATUSES = {"failed", "error", "canceled", "cancelled"} diff --git a/video-gen-api/app/models/__init__.py b/video-gen-api/app/models/__init__.py index 1d2f5b9b..763a6f43 100644 --- a/video-gen-api/app/models/__init__.py +++ b/video-gen-api/app/models/__init__.py @@ -31,6 +31,7 @@ from app.models.user_oauth import UserOAuth from app.models.user_oauth_account import UserOAuthAccount from app.models.user_oauth_app import UserOAuthApp from app.models.home_material import HomeMaterialAsset, HomeMaterialCategory, HomeMaterialWatermark +from app.models.contact_request import ContactRequest __all__ = [ "Base", "TimestampMixin", "SoftDeleteMixin", "engine", "async_session", @@ -38,7 +39,7 @@ __all__ = [ "User", "Team", "Project", "GenerationRecord", "CreditRecord", "ModelConfig", "SystemConfig", "Notification", "PaymentOrder", "TokenUsage", "IndustryConfig", "VideoEngine", "CreditRatio", - "MenuConfig", "RechargePackage", "OperationLog", + "MenuConfig", "RechargePackage", "OperationLog", "ContactRequest", "ChatGenerationTask", "ChatGenerationTaskEvent", "ChatProviderCallLog", "GeneratedResource", "UserResourceMonthStat", "UserResourceTotalStat", "UserResourceCapacityConfig", diff --git a/video-gen-api/app/models/chat_generation_task.py b/video-gen-api/app/models/chat_generation_task.py index 7307e1f8..4bb3e1c3 100644 --- a/video-gen-api/app/models/chat_generation_task.py +++ b/video-gen-api/app/models/chat_generation_task.py @@ -26,6 +26,17 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin): unique=True, postgresql_where=text("deleted_at IS NULL AND idempotency_key IS NOT NULL"), ), + # 视频 24 小时降频轮询调度使用。 + Index( + "idx_chat_generation_tasks_next_poll_at", + "next_poll_at", + postgresql_where=text( + "deleted_at IS NULL " + "AND status = 'generating' " + "AND gen_type = 'video' " + "AND next_poll_at IS NOT NULL" + ), + ), ) @@ -72,6 +83,11 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin): retry_count: Mapped[int] = mapped_column(Integer, default=0) poll_count: Mapped[int] = mapped_column(Integer, default=0) last_poll_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) + # 视频降频轮询调度字段。 + # 图片同步生成仍沿用原超时逻辑;这些字段主要给 video + provider poll 使用。 + poll_started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) + next_poll_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) + poll_interval_seconds: Mapped[int] = mapped_column(Integer, default=0) deadline_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) generated_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) error_message: Mapped[str | None] = mapped_column(Text, nullable=True) diff --git a/video-gen-api/app/schemas/shot_replicate.py b/video-gen-api/app/schemas/shot_replicate.py index 89da1ccb..86c809f2 100644 --- a/video-gen-api/app/schemas/shot_replicate.py +++ b/video-gen-api/app/schemas/shot_replicate.py @@ -484,6 +484,16 @@ class ShotReplicateDeleteOut(BaseModel): message: str = Field(..., description="删除结果提示") project_id: str = Field(..., description="被软删除的总任务项目ID") deleted: bool = Field(..., description="是否已软删除") + released_size_bytes: int = Field(0, description="本次软删释放的用户容量占用字节数;不删除物理文件") + + +class ShotSegmentDeleteOut(BaseModel): + message: str = Field(..., description="删除结果提示") + segment_id: str = Field(..., description="被软删除的拆镜片段ID") + task_set_id: str = Field(..., description="所属拆镜总任务集ID") + deleted: bool = Field(..., description="是否已软删除") + deleted_module_project_id: str | None = Field(None, description="联动软删除的拆镜复刻项目ID;没有关联项目时为空") + released_size_bytes: int = Field(0, description="本次释放的用户容量占用字节数;只释放数据账本,不删除物理文件") diff --git a/video-gen-api/app/services/admin_credit_record_service.py b/video-gen-api/app/services/admin_credit_record_service.py index 2f5f3f5b..cbc6b682 100644 --- a/video-gen-api/app/services/admin_credit_record_service.py +++ b/video-gen-api/app/services/admin_credit_record_service.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import datetime +from datetime import datetime, timezone, timedelta from typing import Any from sqlalchemy import and_, case, distinct, func, or_, select @@ -28,15 +28,25 @@ from app.models.shot_replicate_task_set import ShotReplicateTaskSet from app.models.user import User +CST = timezone(timedelta(hours=8)) + + def _iso(dt: Any) -> str | None: if dt is None: return None + + if isinstance(dt, datetime): + if dt.tzinfo is None: + # 数据库已经按东八区业务时间返回但丢了 tzinfo 时,不再额外 +8 + return dt.replace(tzinfo=CST).isoformat() + + return dt.astimezone(CST).isoformat() + try: return dt.isoformat() except Exception: return str(dt) - def _round2(value: Any) -> float: try: return round(float(value or 0), 2) diff --git a/video-gen-api/app/services/generation_ai_service.py b/video-gen-api/app/services/generation_ai_service.py index b1df03df..1fe81b64 100644 --- a/video-gen-api/app/services/generation_ai_service.py +++ b/video-gen-api/app/services/generation_ai_service.py @@ -331,7 +331,7 @@ async def create_async_generation_task(db: AsyncSession, current_user: User, req media_references=_json(refs) if refs else None, credits_cost=round(media_billing.total_charged, 2), idempotency_key=req.idempotency_key, - deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_VIDEO_DEADLINE_MINUTES), + deadline_at=now + timedelta(hours=settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS), ) db.add(task) diff --git a/video-gen-api/app/services/generation_poll_schedule_service.py b/video-gen-api/app/services/generation_poll_schedule_service.py new file mode 100644 index 00000000..4aef5cd9 --- /dev/null +++ b/video-gen-api/app/services/generation_poll_schedule_service.py @@ -0,0 +1,153 @@ +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone + +from app.config import settings +from app.enums.generation_task import GenerationType +from app.models.chat_generation_task import ChatGenerationTask +from app.services.redis_registry_service import ensure_aware_utc + + +@dataclass(slots=True) +class PollScheduleDecision: + delay_seconds: int + next_poll_at: datetime + poll_interval_seconds: int + direct_countdown: bool + final_poll: bool + reason: str + + +def utc_now() -> datetime: + return datetime.now(timezone.utc) + + +def is_video_generation_task(task: ChatGenerationTask) -> bool: + return str(getattr(task, "gen_type", "") or "").lower() == GenerationType.VIDEO.value + + +def video_final_deadline_from(now: datetime | None = None) -> datetime: + current_time = now or utc_now() + hours = max(1, int(settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS or 24)) + return current_time + timedelta(hours=hours) + + +def ensure_video_poll_fields(task: ChatGenerationTask, *, now: datetime | None = None) -> None: + """补齐视频轮询调度字段,兼容历史任务。""" + if not is_video_generation_task(task): + return + + current_time = now or utc_now() + if ensure_aware_utc(getattr(task, "poll_started_at", None)) is None: + task.poll_started_at = current_time + if ensure_aware_utc(getattr(task, "deadline_at", None)) is None: + task.deadline_at = video_final_deadline_from(current_time) + if getattr(task, "poll_interval_seconds", None) is None: + task.poll_interval_seconds = 0 + + +def is_final_poll_due(task: ChatGenerationTask, *, now: datetime | None = None) -> bool: + deadline_at = ensure_aware_utc(getattr(task, "deadline_at", None)) + return bool(deadline_at and deadline_at <= (now or utc_now())) + + +def is_poll_not_due(task: ChatGenerationTask, *, now: datetime | None = None, tolerance_seconds: int = 1) -> bool: + """判断当前 poll 任务是否早于 next_poll_at。只对视频降频轮询生效。""" + if not is_video_generation_task(task): + return False + if is_final_poll_due(task, now=now): + return False + next_poll_at = ensure_aware_utc(getattr(task, "next_poll_at", None)) + if next_poll_at is None: + return False + return next_poll_at > (now or utc_now()) + timedelta(seconds=max(0, int(tolerance_seconds))) + + +def _clamp_positive_seconds(value: int | float | None, default: int) -> int: + try: + parsed = int(value if value is not None else default) + except (TypeError, ValueError): + parsed = default + return max(1, parsed) + + +def build_video_pending_poll_schedule( + task: ChatGenerationTask, + *, + now: datetime | None = None, +) -> PollScheduleDecision: + """计算视频任务 pending/running 后的下一次轮询时间。""" + current_time = now or utc_now() + ensure_video_poll_fields(task, now=current_time) + + deadline_at = ensure_aware_utc(task.deadline_at) + if deadline_at and deadline_at <= current_time: + return PollScheduleDecision( + delay_seconds=0, + next_poll_at=current_time, + poll_interval_seconds=int(task.poll_interval_seconds or 0), + direct_countdown=True, + final_poll=True, + reason="video_final_poll_due", + ) + + poll_started_at = ensure_aware_utc(task.poll_started_at) or current_time + elapsed_seconds = max(0, int((current_time - poll_started_at).total_seconds())) + high_freq_seconds = max(0, int(settings.CHATAPI_ASYNC_VIDEO_HIGH_FREQ_MINUTES or 10)) * 60 + high_freq_poll_seconds = _clamp_positive_seconds(settings.CHATAPI_ASYNC_VIDEO_HIGH_FREQ_POLL_SECONDS, 30) + initial_backoff_seconds = _clamp_positive_seconds(settings.CHATAPI_ASYNC_VIDEO_BACKOFF_INITIAL_SECONDS, 60) + multiplier = max(1, int(settings.CHATAPI_ASYNC_VIDEO_BACKOFF_MULTIPLIER or 2)) + max_backoff_seconds = _clamp_positive_seconds(settings.CHATAPI_ASYNC_VIDEO_BACKOFF_MAX_SECONDS, 3600) + + if elapsed_seconds < high_freq_seconds: + delay_seconds = high_freq_poll_seconds + interval_seconds = int(task.poll_interval_seconds or 0) + reason = "video_high_freq_poll" + else: + previous_interval = int(task.poll_interval_seconds or 0) + if previous_interval < initial_backoff_seconds: + interval_seconds = initial_backoff_seconds + else: + interval_seconds = min(previous_interval * multiplier, max_backoff_seconds) + delay_seconds = interval_seconds + reason = "video_backoff_poll" + + next_poll_at = current_time + timedelta(seconds=delay_seconds) + final_poll = False + if deadline_at and next_poll_at >= deadline_at: + next_poll_at = deadline_at + delay_seconds = max(0, int((deadline_at - current_time).total_seconds())) + final_poll = delay_seconds <= 0 + reason = "video_schedule_to_final_deadline" + + direct_max = max(0, int(settings.CHATAPI_ASYNC_VIDEO_DIRECT_COUNTDOWN_MAX_SECONDS or 300)) + return PollScheduleDecision( + delay_seconds=delay_seconds, + next_poll_at=next_poll_at, + poll_interval_seconds=interval_seconds, + direct_countdown=delay_seconds <= direct_max, + final_poll=final_poll, + reason=reason, + ) + + +def build_default_poll_schedule( + task: ChatGenerationTask, + *, + now: datetime | None = None, + delay_seconds: int | None = None, + reason: str = "default_poll", +) -> PollScheduleDecision: + """图片和旧逻辑兼容用的固定间隔轮询计划。""" + current_time = now or utc_now() + delay = _clamp_positive_seconds(delay_seconds, int(settings.CHATAPI_ASYNC_POLL_INTERVAL_SECONDS or 30)) + next_poll_at = current_time + timedelta(seconds=delay) + return PollScheduleDecision( + delay_seconds=delay, + next_poll_at=next_poll_at, + poll_interval_seconds=int(getattr(task, "poll_interval_seconds", 0) or 0), + direct_countdown=True, + final_poll=False, + reason=reason, + ) diff --git a/video-gen-api/app/services/generation_recovery_service.py b/video-gen-api/app/services/generation_recovery_service.py index 451b26fc..bf455de3 100644 --- a/video-gen-api/app/services/generation_recovery_service.py +++ b/video-gen-api/app/services/generation_recovery_service.py @@ -9,11 +9,13 @@ from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings +from app.enums.celery_queue import CeleryQueue from app.enums.generation_task import ( ALLOWED_GENERATION_MODES, ChatGenerationPipelineStage, ChatGenerationTaskEventType, ChatGenerationTaskStatus, + GenerationType, ) from app.models.chat_generation_task import ChatGenerationTask from app.services.celery_download_recovery_service import ( @@ -25,6 +27,7 @@ from app.services.celery_download_recovery_service import ( ) 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.generation_poll_schedule_service import ensure_video_poll_fields, is_poll_not_due, is_video_generation_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, @@ -35,7 +38,7 @@ from app.services.redis_registry_service import ( logger = logging.getLogger("video_gen") -POLL_QUEUE = "gen_provider_poll" +POLL_QUEUE = CeleryQueue.GEN_PROVIDER_POLL.value def _now() -> datetime: @@ -342,7 +345,7 @@ async def _mark_timeout( await _remove_poll_active(task.id) await log_task_event( task, - event_type="TASK_TIMEOUT", + event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value, to_status="failed", to_stage=ChatGenerationPipelineStage.TIMEOUT.value, ) @@ -414,7 +417,7 @@ async def recover_one_generation_task( await _remove_poll_active(task.id) await log_task_event( task, - event_type="GENERATION_RECOVERY_ENQUEUE", + event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value, message=f"{source} 发现任务已存在 remote_result_url,恢复投递下载队列", detail={ "pipeline_stage": task.pipeline_stage, @@ -439,7 +442,7 @@ async def recover_one_generation_task( await db.commit() await log_task_event( task, - event_type="GENERATION_RECOVERY_ENQUEUE", + event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value, message=f"{source} 发现任务已到 deadline 且存在供应商任务ID,投递 poll 队列做最终查询", detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, ) @@ -453,26 +456,51 @@ async def recover_one_generation_task( await log_task_event( task, - event_type="GENERATION_RECOVERY_TIMEOUT", + event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_TIMEOUT.value, 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 队列。 + # 视频任务如果 next_poll_at 未到期,不提前 poll,只刷新 active 注册表等待 Beat dispatcher 到期投递。 if has_provider_task_id: + if is_video_generation_task(task): + ensure_video_poll_fields(task, now=current_time) + if is_poll_not_due(task, now=current_time): + await db.commit() + await register_poll_active( + task, + check_at=task.next_poll_at, + next_poll_at=task.next_poll_at, + reason=f"{source}_video_poll_not_due", + ) + await log_task_event( + task, + event_type=ChatGenerationTaskEventType.POLL_SKIP_NOT_DUE.value, + message=f"{source} 发现视频任务尚未到下一次轮询时间,启动容灾不提前投递 poll", + detail={ + "pipeline_stage": task.pipeline_stage, + "payload": redis_payload, + "next_poll_at": task.next_poll_at, + }, + ) + return "skip_video_poll_not_due" + task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value + task.next_poll_at = _poll_queue_timeout_at(current_time) await db.commit() await log_task_event( task, - event_type="GENERATION_RECOVERY_ENQUEUE", + event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value, 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(), + check_at=task.next_poll_at, + next_poll_at=task.next_poll_at, reason=f"{source}_has_provider_task_id", ) return "recover_poll_has_provider_id" @@ -499,13 +527,13 @@ async def recover_one_generation_task( await _remove_poll_active(task.id) await log_task_event( task, - event_type="GENERATION_RECOVERY_ENQUEUE", + event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value, message=f"{source} 发现任务未超时且缺少 remote_result_url/供应商任务ID,恢复投递创建队列", detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, ) chatapi_create_generation_task.apply_async( args=[task.id], - queue="gen_chatapi_create", + queue=CeleryQueue.GEN_CHATAPI_CREATE.value, countdown=0, ) return "recover_create_no_remote_no_provider_before_deadline" @@ -517,13 +545,13 @@ async def recover_one_generation_task( await _remove_poll_active(task.id) await log_task_event( task, - event_type="GENERATION_RECOVERY_ENQUEUE", + event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value, 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", + queue=CeleryQueue.GEN_CHATAPI_CREATE.value, countdown=0, ) return "recover_create_result_ready_no_url_before_deadline" @@ -534,7 +562,7 @@ async def recover_one_generation_task( async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]: """启动时生成链路容灾扫描。 - 不新增 Celery beat,不新增 worker 命令;worker 启动时由 Redis 锁保证只投递一次。 + 启动容灾由 worker_ready 触发,只跑一次完整恢复;周期性视频到期轮询由 Celery Beat 调度 dispatch_due_poll_tasks。 恢复顺序: 1. Redis poll active_index 到期任务; 2. DB fallback 扫描 queued/creating/waiting_remote/polling/result_ready; @@ -630,4 +658,126 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]: "checked": len(checked_ids), "db_checked": total_db_checked, "results": results, - } \ No newline at end of file + } + +async def dispatch_due_poll_tasks_once(db: AsyncSession) -> dict[str, Any]: + """周期性轻量到期轮询调度。 + + 只处理视频任务的 next_poll_at 到期记录,不替代启动容灾 recover_generation_tasks_once。 + Beat 每分钟触发本任务,本任务把真实供应商轮询投递到 gen_provider_poll 队列。 + """ + from app.tasks.generation_poll_tasks import poll_generation_task, register_poll_active + + current_time = _now() + batch_size = max(1, int(settings.POLL_DUE_DISPATCH_BATCH_SIZE or 100)) + poll_lease_expired_at = current_time - timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300)) + + query_result = await db.execute( + select(ChatGenerationTask) + .where( + ChatGenerationTask.deleted_at.is_(None), + ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)), + ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value, + ChatGenerationTask.gen_type == GenerationType.VIDEO.value, + ChatGenerationTask.next_poll_at.is_not(None), + ChatGenerationTask.next_poll_at <= current_time, + ChatGenerationTask.pipeline_stage.in_( + [ + ChatGenerationPipelineStage.WAITING_REMOTE.value, + ChatGenerationPipelineStage.POLLING.value, + ] + ), + ) + .order_by(ChatGenerationTask.next_poll_at.asc(), ChatGenerationTask.updated_at.asc()) + .limit(batch_size) + .with_for_update(skip_locked=True) + ) + tasks = query_result.scalars().all() + + results: dict[str, int] = {} + dispatched_task_ids: list[str] = [] + + for task in tasks: + action = "skip_unknown" + try: + if task.pipeline_stage == ChatGenerationPipelineStage.POLLING.value: + last_poll_at = ensure_aware_utc(task.last_poll_at) + if last_poll_at and last_poll_at > poll_lease_expired_at: + await register_poll_active( + task, + check_at=_poll_queue_timeout_at(last_poll_at), + next_poll_at=task.next_poll_at, + reason="due_dispatch_polling_lease_alive", + ) + action = "skip_polling_lease_alive" + continue + + if not (task.seedance_task_id or task.provider_task_id): + # dispatcher 不负责重新 create;没有 provider id 的异常状态交给启动容灾或 create 任务处理。 + await log_task_event( + task, + event_type=ChatGenerationTaskEventType.POLL_DISPATCH_SKIP.value, + message="视频到期轮询调度跳过:缺少外部任务ID", + detail={"pipeline_stage": task.pipeline_stage, "next_poll_at": task.next_poll_at}, + ) + action = "skip_no_provider_task_id" + continue + + task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value + # 设置一个队列消费保护时间,避免 Beat 下一分钟看到旧 next_poll_at 又重复投递。 + task.next_poll_at = _poll_queue_timeout_at(current_time) + dispatched_task_ids.append(task.id) + action = "dispatch_poll" + except Exception as exc: + logger.exception("视频到期轮询调度单条处理失败。task_id=%s", getattr(task, "id", None)) + action = "error" + await log_task_event( + task, + event_type=ChatGenerationTaskEventType.POLL_DISPATCH_SKIP.value, + message=f"视频到期轮询调度单条处理失败:{exc}", + ) + finally: + results[action] = results.get(action, 0) + 1 + + await db.commit() + + fresh_tasks = [] + if dispatched_task_ids: + fresh_result = await db.execute( + select(ChatGenerationTask) + .where( + ChatGenerationTask.id.in_(dispatched_task_ids), + ChatGenerationTask.deleted_at.is_(None), + ) + .execution_options(populate_existing=True) + ) + fresh_tasks = fresh_result.scalars().all() + + enqueued_count = 0 + + for task in fresh_tasks: + await register_poll_active( + task, + check_at=task.next_poll_at or _poll_queue_timeout_at(current_time), + next_poll_at=task.next_poll_at, + reason="due_dispatch_poll_queued", + ) + await log_task_event( + task, + event_type=ChatGenerationTaskEventType.POLL_DISPATCH_DUE.value, + message="视频 next_poll_at 到期,已投递 provider poll 队列", + detail={"next_poll_at": task.next_poll_at, "queue": POLL_QUEUE}, + ) + poll_generation_task.apply_async(args=[task.id], queue=POLL_QUEUE, countdown=0) + enqueued_count += 1 + + # 如果 log_task_event 内部不 commit,这里要提交一次 + if fresh_tasks: + await db.commit() + + return { + "checked": len(tasks), + "dispatched": len(dispatched_task_ids), + "enqueued": enqueued_count, + "results": results, + } diff --git a/video-gen-api/app/services/generation_task_factory_service.py b/video-gen-api/app/services/generation_task_factory_service.py index 0bb51267..a58705ad 100644 --- a/video-gen-api/app/services/generation_task_factory_service.py +++ b/video-gen-api/app/services/generation_task_factory_service.py @@ -195,7 +195,7 @@ async def create_chat_generation_task_for_module( media_references=_json(refs) if refs else None, credits_cost=round(media_billing.total_charged, 2), idempotency_key=backend_idempotency_key, - deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_VIDEO_DEADLINE_MINUTES), + deadline_at=now + timedelta(hours=settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS), ) db.add(task) diff --git a/video-gen-api/app/services/module_generation_flow_base_service.py b/video-gen-api/app/services/module_generation_flow_base_service.py index 1f53cb8f..4be57cb9 100644 --- a/video-gen-api/app/services/module_generation_flow_base_service.py +++ b/video-gen-api/app/services/module_generation_flow_base_service.py @@ -232,6 +232,8 @@ async def soft_delete_steps_from_index( config: ModuleGenerationFlowConfig, log_module_event: LogModuleEventCallable, deleted_at: datetime | None = None, + refund_unfinished: bool = True, + release_stats: dict[str, int] | None = None, ) -> list[ModuleGenerationStep]: deleted_at = deleted_at or utc_now() result = await db.execute( @@ -266,8 +268,10 @@ async def soft_delete_steps_from_index( chat_task = chat_result.scalar_one_or_none() if chat_task: if chat_task.status == "completed": - await soft_delete_chat_task_resources(db, chat_task.id, deleted_at=deleted_at) - elif chat_task.status != "failed": + released_size = await soft_delete_chat_task_resources(db, chat_task.id, deleted_at=deleted_at) + if release_stats is not None: + release_stats["released_size_bytes"] = int(release_stats.get("released_size_bytes", 0)) + int(released_size or 0) + elif refund_unfinished and chat_task.status != "failed": await mark_chat_generation_task_failed_and_refund_once( db, task=chat_task, diff --git a/video-gen-api/app/services/resource_accounting_service.py b/video-gen-api/app/services/resource_accounting_service.py index 723d321b..9a962bf3 100644 --- a/video-gen-api/app/services/resource_accounting_service.py +++ b/video-gen-api/app/services/resource_accounting_service.py @@ -18,6 +18,7 @@ from app.utils.id_gen import generate_id SOURCE_MODEL_CHAT_TASK = "ChatGenerationTask" SOURCE_MODEL_GENERATION_RECORD = "GenerationRecord" +SOURCE_MODEL_SHOT_SEGMENT = "ShotReplicateSegment" @dataclass(slots=True) diff --git a/video-gen-api/app/services/shot_replicate_flow_service.py b/video-gen-api/app/services/shot_replicate_flow_service.py index 3e0d0d68..2fd5d3e2 100644 --- a/video-gen-api/app/services/shot_replicate_flow_service.py +++ b/video-gen-api/app/services/shot_replicate_flow_service.py @@ -12,6 +12,7 @@ from sqlalchemy.orm.attributes import flag_modified from app.config import settings from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum +from app.enums.generation_task import ChatGenerationPipelineStage, ChatGenerationTaskStatus from app.enums.shot_replicate import ShotReplicateGenerationModeEnum, ShotReplicateStepCodeEnum, ModuleCodeEnum from app.models.chat_generation_task import ChatGenerationTask from app.models.module_generation_project import ModuleGenerationProject @@ -335,6 +336,8 @@ async def _soft_delete_steps_from_index( project: ModuleGenerationProject, start_index: int, deleted_at: datetime | None = None, + refund_unfinished: bool = True, + release_stats: dict[str, int] | None = None, ) -> None: await _base_soft_delete_steps_from_index( db, @@ -343,6 +346,8 @@ async def _soft_delete_steps_from_index( config=FLOW_CONFIG, log_module_event=log_module_event, deleted_at=deleted_at, + refund_unfinished=refund_unfinished, + release_stats=release_stats, ) @@ -1406,6 +1411,60 @@ async def handle_chat_generation_task_failed(db: AsyncSession, task: ChatGenerat await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.CHAT_TASK_FAILED.value, message=project.error_message, detail={"chat_task_id": task.id}) + +_ACTIVE_DELETE_BLOCK_STATUSES = { + ChatGenerationTaskStatus.PENDING.value, + ChatGenerationTaskStatus.GENERATING.value, +} +_ACTIVE_DELETE_BLOCK_STAGES = { + ChatGenerationPipelineStage.QUEUED.value, + ChatGenerationPipelineStage.PREPARING.value, + ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, + ChatGenerationPipelineStage.WAITING_REMOTE.value, + ChatGenerationPipelineStage.POLLING.value, + ChatGenerationPipelineStage.RESULT_READY.value, + ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value, + ChatGenerationPipelineStage.DOWNLOADING.value, + ChatGenerationPipelineStage.RETRY_WAITING.value, +} + + +async def _assert_project_has_no_active_chat_tasks_for_delete( + db: AsyncSession, + *, + project: ModuleGenerationProject, +) -> None: + """用户主动删除项目/切片时不退款;如仍有异步生成任务进行中,直接拦截。""" + step_result = await db.execute( + select(ModuleGenerationStep.chat_task_id) + .where( + ModuleGenerationStep.project_id == project.id, + ModuleGenerationStep.module == MODULE, + ModuleGenerationStep.deleted_at.is_(None), + ModuleGenerationStep.chat_task_id.is_not(None), + ) + ) + chat_task_ids = [task_id for task_id in step_result.scalars().all() if task_id] + if not chat_task_ids: + return + + task_result = await db.execute( + select(ChatGenerationTask) + .where( + ChatGenerationTask.id.in_(chat_task_ids), + ChatGenerationTask.deleted_at.is_(None), + ) + .with_for_update() + ) + active_tasks = [] + for task in task_result.scalars().all(): + if task.status in _ACTIVE_DELETE_BLOCK_STATUSES or (task.pipeline_stage in _ACTIVE_DELETE_BLOCK_STAGES): + active_tasks.append(task.id) + + if active_tasks: + raise HTTPException(status_code=400, detail="当前拆镜复刻项目仍有生成中任务,暂不能删除") + + async def mark_shot_replicate_step_dispatch_failed( db: AsyncSession, *, @@ -1462,13 +1521,45 @@ async def mark_shot_replicate_step_dispatch_failed( ) -async def delete_shot_replicate_project(db: AsyncSession, *, current_user: User, project_id: str) -> ShotReplicateDeleteOut: +async def delete_shot_replicate_project( + db: AsyncSession, + *, + current_user: User, + project_id: str, + refund_unfinished: bool = False, +) -> ShotReplicateDeleteOut: project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) deleted_at = _now() + release_stats: dict[str, int] = {"released_size_bytes": 0} + + if not refund_unfinished: + await _assert_project_has_no_active_chat_tasks_for_delete(db, project=project) + project.deleted_at = deleted_at - await _soft_delete_steps_from_index(db, project=project, start_index=1, deleted_at=deleted_at) - await log_module_event(db, project=project, event_type=ModuleEventTypeEnum.PROJECT_DELETED.value, message="软删除拆镜复刻项目") - return ShotReplicateDeleteOut(message="项目已删除", project_id=project.id, deleted=True) + await _soft_delete_steps_from_index( + db, + project=project, + start_index=1, + deleted_at=deleted_at, + refund_unfinished=refund_unfinished, + release_stats=release_stats, + ) + await log_module_event( + db, + project=project, + event_type=ModuleEventTypeEnum.PROJECT_DELETED.value, + message="软删除拆镜复刻项目", + detail={ + "refund_unfinished": refund_unfinished, + "released_size_bytes": int(release_stats.get("released_size_bytes", 0)), + }, + ) + return ShotReplicateDeleteOut( + message="项目已删除", + project_id=project.id, + deleted=True, + released_size_bytes=int(release_stats.get("released_size_bytes", 0)), + ) async def create_shot_replicate_project_from_segment( diff --git a/video-gen-api/app/services/shot_replicate_taskset_service.py b/video-gen-api/app/services/shot_replicate_taskset_service.py index af9c4c91..e94a4a99 100644 --- a/video-gen-api/app/services/shot_replicate_taskset_service.py +++ b/video-gen-api/app/services/shot_replicate_taskset_service.py @@ -25,6 +25,7 @@ from app.models.shot_replicate_task_set import ShotReplicateTaskSet from app.models.user import User from app.schemas.shot_replicate import ( ShotAISuggestionOut, + ShotSegmentDeleteOut, ShotSegmentDetailOut, ShotSegmentListOut, ShotSegmentOut, @@ -38,6 +39,7 @@ from app.schemas.shot_replicate import ( ShotTaskSetOut, ) from app.services.module_generation_log_service import log_module_event_file +from app.services.resource_accounting_service import SOURCE_MODEL_SHOT_SEGMENT, soft_delete_resources_by_source from app.services.upload_video_asset_service import ( build_time_node, validate_split_range, @@ -615,3 +617,102 @@ async def segment_detail(db: AsyncSession, *, current_user: User, segment_id: st project_result = await db.execute(select(ModuleGenerationProject).where(ModuleGenerationProject.id == segment.module_project_id).limit(1)) project = project_result.scalar_one_or_none() return _segment_to_detail_out(segment, project) + +async def delete_segment( + db: AsyncSession, + *, + current_user: User, + segment_id: str, +) -> ShotSegmentDeleteOut: + """软删除拆镜片段。 + + 只释放用户容量账本记录,不删除 segment_video_path 指向的物理文件。 + 如果片段已创建复刻项目,则联动调用项目删除逻辑,但用户主动删除不退款; + 项目仍有生成中任务时会拒绝删除,避免异步任务继续写回软删数据。 + """ + query = select(ShotReplicateSegment).where(ShotReplicateSegment.id == segment_id) + if not current_user.is_admin: + query = query.where(ShotReplicateSegment.user_id == current_user.id) + result = await db.execute(query.with_for_update().limit(1)) + segment = result.scalar_one_or_none() + if not segment: + raise HTTPException(status_code=404, detail="拆镜片段不存在") + + task_set_id = segment.task_set_id + module_project_id = segment.module_project_id + + if segment.deleted_at is not None: + return ShotSegmentDeleteOut( + message="拆镜片段已删除", + segment_id=segment.id, + task_set_id=task_set_id, + deleted=True, + deleted_module_project_id=module_project_id, + released_size_bytes=0, + ) + + if segment.split_status == ShotSplitStatusEnum.PROCESSING.value: + raise HTTPException(status_code=400, detail="当前拆镜片段正在切割处理中,暂不能删除") + if segment.analysis_status == ShotSegmentAnalysisStatusEnum.PROCESSING.value: + raise HTTPException(status_code=400, detail="当前拆镜片段正在分析处理中,暂不能删除") + if segment.replicate_status == ShotSegmentReplicateStatusEnum.PROCESSING.value: + raise HTTPException(status_code=400, detail="当前拆镜片段关联的复刻流程正在处理中,暂不能删除") + + deleted_at = _now() + released_size_bytes = await soft_delete_resources_by_source( + db, + source_model=SOURCE_MODEL_SHOT_SEGMENT, + source_ids=[segment.id], + deleted_at=deleted_at, + ) + + deleted_module_project_id: str | None = None + if module_project_id: + from app.services.shot_replicate_flow_service import delete_shot_replicate_project + + project_delete_out = await delete_shot_replicate_project( + db, + current_user=current_user, + project_id=module_project_id, + refund_unfinished=False, + ) + deleted_module_project_id = project_delete_out.project_id + released_size_bytes += int(project_delete_out.released_size_bytes or 0) + + segment.deleted_at = deleted_at + segment.replicate_status = ( + ShotSegmentReplicateStatusEnum.NOT_STARTED.value + if not deleted_module_project_id + else ShotSegmentReplicateStatusEnum.FAILED.value + ) + + await refresh_task_set_split_summary(db, task_set_id) + await db.flush() + + log_module_event_file( + module=MODULE, + event_type="SHOT_SEGMENT_DELETED", + project_id=task_set_id, + step_id=segment.id, + user_id=segment.user_id, + message="软删除拆镜片段并释放用户容量账本记录", + detail={ + "segment_id": segment.id, + "task_set_id": task_set_id, + "module_project_id": module_project_id, + "deleted_module_project_id": deleted_module_project_id, + "released_size_bytes": released_size_bytes, + "physical_file_deleted": False, + "refund": False, + }, + ) + + return ShotSegmentDeleteOut( + message="拆镜片段已删除", + segment_id=segment.id, + task_set_id=task_set_id, + deleted=True, + deleted_module_project_id=deleted_module_project_id, + released_size_bytes=int(released_size_bytes or 0), + ) + diff --git a/video-gen-api/app/tasks/celery_app.py b/video-gen-api/app/tasks/celery_app.py index 94c777ee..80280283 100644 --- a/video-gen-api/app/tasks/celery_app.py +++ b/video-gen-api/app/tasks/celery_app.py @@ -4,6 +4,7 @@ from celery import Celery from celery.signals import worker_process_init, worker_process_shutdown, worker_ready from app.config import settings +from app.enums.celery_queue import CeleryQueue, CeleryTaskName from app.models.base import engine from app.tasks.async_runner import close_loop, run_async @@ -26,7 +27,7 @@ CELERY_TASK_IMPORTS = ( ) -RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or "gen_recovery" +RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or CeleryQueue.GEN_RECOVERY.value def _derive_redis_db(url: str, db_no: int) -> str: @@ -39,6 +40,21 @@ def _derive_redis_db(url: str, db_no: int) -> str: return url.rstrip("/") + f"/{db_no}" +def _beat_schedule() -> dict: + if not bool(getattr(settings, "POLL_DUE_DISPATCH_ENABLED", True)): + return {} + return { + "dispatch-due-poll-tasks-every-minute": { + "task": CeleryTaskName.DISPATCH_DUE_POLL.value, + "schedule": max(1, int(settings.POLL_DUE_DISPATCH_INTERVAL_SECONDS or 60)), + "options": { + "queue": RECOVERY_QUEUE, + "priority": settings.DOWNLOAD_TASK_PRIORITY_RECOVER, + }, + } + } + + broker_url = settings.CELERY_BROKER_URL or (_derive_redis_db(settings.REDIS_URL, 1) if settings.REDIS_URL else "") backend_url = settings.CELERY_RESULT_BACKEND or (_derive_redis_db(settings.REDIS_URL, 2) if settings.REDIS_URL else "") @@ -58,6 +74,7 @@ if broker_url: task_acks_late=True, task_reject_on_worker_lost=True, task_track_started=True, + beat_schedule=_beat_schedule(), task_annotations={ # 生成链路任务以数据库状态为准,不依赖 Celery result backend。 # 这里忽略结果可避免任务误返回 ORM / 非 JSON 对象时触发结果序列化失败。 @@ -80,24 +97,25 @@ if broker_url: "sep": ":", }, task_routes={ - "generation.chatapi_create_generation_task": {"queue": "gen_chatapi_create"}, - "generation.poll_generation_task": {"queue": "gen_provider_poll"}, - "generation.download_generation_result_task": {"queue": "gen_result_download"}, - "hot_opening.start_image_prompt_optimize": {"queue": "gen_chatapi_create"}, - "hot_opening.start_video_prompt_optimize": {"queue": "gen_chatapi_create"}, - "shot_replicate.analyze_original_video": {"queue": "gen_chatapi_create"}, - "shot_replicate.analyze_custom_segment_video": {"queue": "gen_chatapi_create"}, - "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"}, + CeleryTaskName.CHATAPI_CREATE.value: {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value}, + CeleryTaskName.POLL_GENERATION.value: {"queue": CeleryQueue.GEN_PROVIDER_POLL.value}, + CeleryTaskName.DOWNLOAD_GENERATION_RESULT.value: {"queue": CeleryQueue.GEN_RESULT_DOWNLOAD.value}, + CeleryTaskName.DISPATCH_DUE_POLL.value: {"queue": RECOVERY_QUEUE}, + "hot_opening.start_image_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value}, + "hot_opening.start_video_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value}, + "shot_replicate.analyze_original_video": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value}, + "shot_replicate.analyze_custom_segment_video": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value}, + "shot_replicate.split_one_segment": {"queue": CeleryQueue.GEN_RESULT_DOWNLOAD.value}, + "shot_replicate.start_image_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value}, + "shot_replicate.start_video_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value}, # 恢复扫描统一走独立队列,避免占用下载/轮询/创建业务 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"}, + CeleryTaskName.STARTUP_RECOVERY.value: {"queue": RECOVERY_QUEUE}, + CeleryTaskName.SHOT_SPLIT_RECOVERY.value: {"queue": RECOVERY_QUEUE}, + CeleryTaskName.RECOVER_DOWNLOAD.value: {"queue": RECOVERY_QUEUE}, + CeleryTaskName.RECOVER_GENERATION.value: {"queue": RECOVERY_QUEUE}, + CeleryTaskName.MODULE_ASYNC_RECOVERY.value: {"queue": RECOVERY_QUEUE}, + "user_oauth.update_oauth_accounts": {"queue": CeleryQueue.DEFAULT.value}, + "app.tasks.cleanup.*": {"queue": CeleryQueue.DEFAULT.value}, }, ) else: @@ -121,8 +139,8 @@ def on_worker_ready(sender=None, **kwargs): """Celery worker 启动时做一次容灾恢复。 注意: - - 不启用 Celery beat。 - - 启动容灾保留,但只投递一个 recovery.startup_recovery_once 协调任务。 + - 启动容灾只投递一个 recovery.startup_recovery_once 协调任务。 + - Celery Beat 只用于每分钟触发轻量 generation.dispatch_due_poll_tasks,不跑完整启动容灾。 - 协调任务走独立 gen_recovery 队列,串行扫描并把真实业务任务投回原队列。 - 所有 worker 都尝试抢 Redis 投递锁,只有抢到锁的 worker 投递恢复任务。 """ @@ -182,4 +200,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 94a23c0d..40ac7c47 100644 --- a/video-gen-api/app/tasks/generation_create_tasks.py +++ b/video-gen-api/app/tasks/generation_create_tasks.py @@ -6,21 +6,31 @@ from typing import Any, Optional from sqlalchemy import select from app.config import settings +from app.enums.celery_queue import CeleryQueue +from app.enums.generation_task import ( + ALLOWED_GENERATION_MODES, + ChatGenerationPipelineStage, + ChatGenerationTaskEventType, + ChatGenerationTaskStatus, + GenerationMode, + GenerationType, +) from app.models.base import async_session from app.models.chat_generation_task import ChatGenerationTask from app.services.error_codes import extract_error_message from app.services.generation_log_service import log_task_event +from app.services.generation_poll_schedule_service import ensure_video_poll_fields 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 -ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"} def _now() -> datetime: return datetime.now(timezone.utc) + def _get_first_value(obj: Any, *field_names: str) -> Optional[Any]: """ 兼容不同版本字段名,避免字段调整后 Celery 任务直接报错。 @@ -74,7 +84,7 @@ def _build_optimized_prompt_by_params(task: ChatGenerationTask) -> str: # 爆款开头复刻第5步的视频生成,original_prompt 已经是视频提词 JSON schema。 # 不能再追加“时长/比例/分辨率”中文参数,否则会污染 schema。 - if generation_mode in {"hot_opening_replicate", "shot_replicate"} and gen_type == "video": + if generation_mode in {GenerationMode.HOT_OPENING_REPLICATE.value, GenerationMode.SHOT_REPLICATE.value} and gen_type == GenerationType.VIDEO.value: stripped = base_prompt.strip() if stripped.startswith("{") or stripped.startswith("["): return base_prompt @@ -88,25 +98,25 @@ def _build_optimized_prompt_by_params(task: ChatGenerationTask) -> str: parts = [] - if gen_type == "video": + if gen_type == GenerationType.VIDEO.value: # 时长:4秒,画面比例:16:9,分辨率:480p if duration: parts.append(f"时长:{duration}秒") parts.append(f"画面比例:{aspect_ratio}") parts.append(f"分辨率:{resolution}") else: - parts.append(f"时长:4秒") - parts.append(f"画面比例:16:9") - parts.append(f"分辨率:480p") - elif gen_type == "image": - if image_size : + parts.append("时长:4秒") + parts.append("画面比例:16:9") + parts.append("分辨率:480p") + elif gen_type == GenerationType.IMAGE.value: + if image_size: parts.append(f"分辨率:{image_size}") parts.append(f"画布比例:{image_proportion}") parts.append(f"像素尺寸:{image_px}") else: - parts.append(f"分辨率:2K") - parts.append(f"画布比例:1:1") - parts.append(f"像素尺寸:2048x2048") + parts.append("分辨率:2K") + parts.append("画布比例:1:1") + parts.append("像素尺寸:2048x2048") else: # 未知类型时返回原始字符 return base_prompt @@ -131,7 +141,7 @@ async def _run(task_id: str): if not task or task.generation_mode not in ALLOWED_GENERATION_MODES: return - if task.status != "generating": + if task.status != ChatGenerationTaskStatus.GENERATING.value: return deadline_at = ensure_aware_utc(task.deadline_at) @@ -140,28 +150,37 @@ async def _run(task_id: str): db, task=task, error_message="任务超时", - pipeline_stage="timeout", + pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value, ) await db.commit() - await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout") + await log_task_event( + task, + event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value, + to_status=ChatGenerationTaskStatus.FAILED.value, + to_stage=ChatGenerationPipelineStage.TIMEOUT.value, + ) from app.services.generation_module_hook_service import notify_chat_generation_task_finished await notify_chat_generation_task_finished(db, task) await db.commit() return - if task.pipeline_stage not in ("queued", "preparing", "creating_provider_task"): + if task.pipeline_stage not in ( + ChatGenerationPipelineStage.QUEUED.value, + ChatGenerationPipelineStage.PREPARING.value, + ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, + ): return try: if not task.optimized_prompt: old_stage = task.pipeline_stage - task.pipeline_stage = "preparing" + task.pipeline_stage = ChatGenerationPipelineStage.PREPARING.value await db.commit() await log_task_event( task, - event_type="PROMPT_CONCAT_START", + event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_START.value, from_stage=old_stage, - to_stage="preparing", + to_stage=ChatGenerationPipelineStage.PREPARING.value, message="开始本地拼接提示词,不调用提词优化API", ) @@ -174,8 +193,8 @@ async def _run(task_id: str): await db.commit() await log_task_event( task, - event_type="PROMPT_CONCAT_SUCCESS", - to_stage="preparing", + event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_SUCCESS.value, + to_stage=ChatGenerationPipelineStage.PREPARING.value, detail={ "optimized_prompt": optimized_prompt, "gen_type": task.gen_type, @@ -184,11 +203,14 @@ async def _run(task_id: str): ) if task.seedance_task_id or task.provider_task_id: - task.pipeline_stage = "waiting_remote" + task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value + if task.gen_type == GenerationType.VIDEO.value: + ensure_video_poll_fields(task, now=_now()) + task.next_poll_at = _now() await db.commit() elif task.remote_result_url: - task.pipeline_stage = "result_ready" + task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value await db.commit() from app.tasks.generation_download_tasks import enqueue_download_task @@ -198,13 +220,13 @@ async def _run(task_id: str): else: old_stage = task.pipeline_stage - task.pipeline_stage = "creating_provider_task" + task.pipeline_stage = ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value await db.commit() await log_task_event( task, - event_type="PROVIDER_CREATE_START", + event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_START.value, from_stage=old_stage, - to_stage="creating_provider_task", + to_stage=ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, ) created = await create_provider_task(db, task) @@ -216,7 +238,7 @@ async def _run(task_id: str): task.remote_result_url = created.get("remote_result_url") or task.remote_result_url - if task.gen_type == "image": + if task.gen_type == GenerationType.IMAGE.value: task.image_tokens_used = created.get("image_tokens", task.image_tokens_used or 0) or 0 task.provider_response_json = json.dumps( @@ -228,21 +250,26 @@ async def _run(task_id: str): if task.remote_result_url and not task.seedance_task_id: # 同步图片路径:原 SDK 已经返回最终 URL。 - task.pipeline_stage = "result_ready" + task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value else: # 视频路径:provider 返回 task id,后续轮询。 - task.pipeline_stage = "waiting_remote" + task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value + if task.gen_type == GenerationType.VIDEO.value: + current_time = _now() + ensure_video_poll_fields(task, now=current_time) + task.next_poll_at = current_time + task.poll_interval_seconds = int(task.poll_interval_seconds or 0) - task.status = "generating" + task.status = ChatGenerationTaskStatus.GENERATING.value await db.commit() await log_task_event( task, - event_type="PROVIDER_CREATE_SUCCESS", + event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_SUCCESS.value, to_stage=task.pipeline_stage, detail=created, ) - if task.pipeline_stage == "result_ready": + if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value: from app.tasks.generation_download_tasks import enqueue_download_task await enqueue_download_task(db, task, reason="create_result_ready") @@ -253,10 +280,11 @@ async def _run(task_id: str): task, reason="create_provider_success", check_at=_now() + timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300)), + next_poll_at=getattr(task, "next_poll_at", None), ) poll_generation_task.apply_async( args=[task.id], - queue="gen_provider_poll", + queue=CeleryQueue.GEN_PROVIDER_POLL.value, countdown=0, ) @@ -278,10 +306,10 @@ async def _run(task_id: str): db, task=task, error_message=error_message, - pipeline_stage="failed", + pipeline_stage=ChatGenerationPipelineStage.FAILED.value, ) await db.commit() - await log_task_event(task, event_type="TASK_FAILED", message=task.error_message) + await log_task_event(task, event_type=ChatGenerationTaskEventType.TASK_FAILED.value, message=task.error_message) from app.services.generation_module_hook_service import notify_chat_generation_task_finished await notify_chat_generation_task_finished(db, task) await db.commit() @@ -305,4 +333,4 @@ else: def apply_async(self, *args, **kwargs): raise RuntimeError("Celery is disabled") - chatapi_create_generation_task = _DisabledTask() \ No newline at end of file + chatapi_create_generation_task = _DisabledTask() diff --git a/video-gen-api/app/tasks/generation_poll_tasks.py b/video-gen-api/app/tasks/generation_poll_tasks.py index 575bb3fe..3ea1064c 100644 --- a/video-gen-api/app/tasks/generation_poll_tasks.py +++ b/video-gen-api/app/tasks/generation_poll_tasks.py @@ -6,10 +6,28 @@ from typing import Any from sqlalchemy import select from app.config import settings +from app.enums.celery_queue import CeleryQueue +from app.enums.generation_task import ( + ALLOWED_GENERATION_MODES, + ChatGenerationPipelineStage, + ChatGenerationTaskEventType, + ChatGenerationTaskStatus, + GenerationType, + PROVIDER_FAILED_STATUSES, + PROVIDER_SUCCESS_STATUSES, +) from app.models.base import async_session from app.models.chat_generation_task import ChatGenerationTask 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_poll_schedule_service import ( + build_default_poll_schedule, + build_video_pending_poll_schedule, + ensure_video_poll_fields, + is_final_poll_due, + is_poll_not_due, + is_video_generation_task, +) 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 @@ -22,20 +40,19 @@ from app.services.redis_registry_service import ( ) from app.tasks.celery_app import celery_app -ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"} -POLL_QUEUE = "gen_provider_poll" +POLL_QUEUE = CeleryQueue.GEN_PROVIDER_POLL.value def _now() -> datetime: return datetime.now(timezone.utc) -def _is_success(status: str) -> bool: - return status in ("succeeded", "success", "completed", "done") +def _is_success(status: str | None) -> bool: + return str(status or "").lower() in PROVIDER_SUCCESS_STATUSES -def _is_failed(status: str) -> bool: - return status in ("failed", "error", "canceled", "cancelled") +def _is_failed(status: str | None) -> bool: + return str(status or "").lower() in PROVIDER_FAILED_STATUSES def _engine_snapshot(task: ChatGenerationTask) -> dict: @@ -46,8 +63,7 @@ def _engine_snapshot(task: ChatGenerationTask) -> dict: def _deadline_expired(task: ChatGenerationTask, now: datetime | None = None) -> bool: - deadline_at = ensure_aware_utc(task.deadline_at) - return bool(deadline_at and deadline_at <= (now or _now())) + return is_final_poll_due(task, now=now) def _poll_check_at(*, delay_seconds: int | float | None = None, now: datetime | None = None) -> datetime: @@ -83,6 +99,8 @@ def _build_poll_active_payload( "queue": POLL_QUEUE, "poll_count": int(task.poll_count or 0), "retry_count": int(task.retry_count or 0), + "poll_started_at": datetime_to_epoch(task.poll_started_at) if getattr(task, "poll_started_at", None) else None, + "poll_interval_seconds": int(getattr(task, "poll_interval_seconds", 0) or 0), "last_poll_at": datetime_to_epoch(task.last_poll_at) if task.last_poll_at else None, "next_poll_at": datetime_to_epoch(checked_next_poll_at) if checked_next_poll_at else None, "deadline_at": datetime_to_epoch(task.deadline_at) if task.deadline_at else None, @@ -154,12 +172,18 @@ async def _mark_timeout(db, task: ChatGenerationTask, *, message: str = "任务 db, task=task, error_message=message, - pipeline_stage="timeout", + pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value, ) + task.next_poll_at = None await _notify_finished(db, task) await db.commit() await remove_poll_active(task.id) - await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout") + await log_task_event( + task, + event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value, + to_status=ChatGenerationTaskStatus.FAILED.value, + to_stage=ChatGenerationPipelineStage.TIMEOUT.value, + ) async def _mark_failed(db, task: ChatGenerationTask, *, message: str, detail: Any = None) -> None: @@ -167,12 +191,84 @@ async def _mark_failed(db, task: ChatGenerationTask, *, message: str, detail: An db, task=task, error_message=message, - pipeline_stage="failed", + pipeline_stage=ChatGenerationPipelineStage.FAILED.value, ) + task.next_poll_at = None await _notify_finished(db, task) await db.commit() await remove_poll_active(task.id) - await log_task_event(task, event_type="POLL_FAILED", message=task.error_message, detail=detail) + await log_task_event( + task, + event_type=ChatGenerationTaskEventType.POLL_FAILED.value, + message=task.error_message, + detail=detail, + ) + + +async def _skip_not_due(task: ChatGenerationTask) -> None: + next_poll_at = ensure_aware_utc(task.next_poll_at) + if next_poll_at is None: + return + await register_poll_active( + task, + check_at=next_poll_at, + next_poll_at=next_poll_at, + reason="poll_task_not_due", + ) + await log_task_event( + task, + event_type=ChatGenerationTaskEventType.POLL_SKIP_NOT_DUE.value, + message="视频任务尚未到下一次轮询时间,本次 poll 跳过", + detail={"next_poll_at": next_poll_at, "pipeline_stage": task.pipeline_stage}, + ) + + +async def _schedule_next_poll( + task: ChatGenerationTask, + *, + reason: str, + default_delay_seconds: int | None = None, +) -> None: + current_time = _now() + if is_video_generation_task(task): + schedule = build_video_pending_poll_schedule(task, now=current_time) + else: + schedule = build_default_poll_schedule( + task, + now=current_time, + delay_seconds=default_delay_seconds, + reason=reason, + ) + + task.next_poll_at = schedule.next_poll_at + task.poll_interval_seconds = schedule.poll_interval_seconds + + await log_task_event( + task, + event_type=ChatGenerationTaskEventType.POLL_SCHEDULED.value, + message=f"已登记下一次轮询。reason={schedule.reason}", + detail={ + "delay_seconds": schedule.delay_seconds, + "next_poll_at": schedule.next_poll_at, + "direct_countdown": schedule.direct_countdown, + "poll_interval_seconds": schedule.poll_interval_seconds, + "source_reason": reason, + }, + ) + + await register_poll_active( + task, + check_at=schedule.next_poll_at, + next_poll_at=schedule.next_poll_at, + reason=schedule.reason, + ) + + if schedule.direct_countdown: + poll_generation_task.apply_async( + args=[task.id], + queue=POLL_QUEUE, + countdown=max(0, int(schedule.delay_seconds)), + ) async def _run(task_id: str): @@ -190,11 +286,23 @@ async def _run(task_id: str): return # 只处理正在生成,且处于远程等待/轮询中的任务。 - if task.status != "generating" or task.pipeline_stage not in ("waiting_remote", "polling"): + if task.status != ChatGenerationTaskStatus.GENERATING.value or task.pipeline_stage not in ( + ChatGenerationPipelineStage.WAITING_REMOTE.value, + ChatGenerationPipelineStage.POLLING.value, + ): await remove_poll_active(task.id) return - final_poll_before_timeout = _deadline_expired(task) + current_time = _now() + if is_video_generation_task(task): + ensure_video_poll_fields(task, now=current_time) + if is_poll_not_due(task, now=current_time): + await db.commit() + await _skip_not_due(task) + return + await db.commit() + + final_poll_before_timeout = _deadline_expired(task, current_time) if final_poll_before_timeout and not (task.seedance_task_id or task.provider_task_id): await _mark_timeout(db, task, message="任务轮询超时") return @@ -206,21 +314,23 @@ async def _run(task_id: str): if final_poll_before_timeout: await log_task_event( task, - event_type="FINAL_POLL_BEFORE_TIMEOUT", + event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT.value, message="任务已到 deadline,执行最后一次供应商查询后再判定超时", detail={"deadline_at": task.deadline_at, "stage": task.pipeline_stage}, ) # 标记本次正在轮询,并登记 poll lease。 # 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递。 - task.pipeline_stage = "polling" + task.pipeline_stage = ChatGenerationPipelineStage.POLLING.value task.poll_count = (task.poll_count or 0) + 1 task.last_poll_at = _now() + task.next_poll_at = _poll_lease_until(task.last_poll_at) await db.commit() await register_poll_active( task, check_at=_poll_lease_until(task.last_poll_at), reason="polling_lease", + next_poll_at=task.next_poll_at, ) try: @@ -247,7 +357,7 @@ async def _run(task_id: str): ) if _is_success(status): - if task.gen_type == "image": + if task.gen_type == GenerationType.IMAGE.value: task.remote_result_url = poll_result.get("image_url") task.image_tokens_used = poll_result.get("image_tokens", 0) or 0 else: @@ -261,12 +371,13 @@ async def _run(task_id: str): await _mark_failed(db, task, message="供应商任务成功但未返回结果URL", detail=poll_result) return - task.pipeline_stage = "result_ready" + task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value task.retry_count = 0 + task.next_poll_at = None await db.commit() await remove_poll_active(task.id) - await log_task_event(task, event_type="POLL_SUCCESS", to_stage="result_ready") + await log_task_event(task, event_type=ChatGenerationTaskEventType.POLL_SUCCESS.value, to_stage=ChatGenerationPipelineStage.RESULT_READY.value) from app.tasks.generation_download_tasks import enqueue_download_task @@ -286,7 +397,7 @@ async def _run(task_id: str): if final_poll_before_timeout: await log_task_event( task, - event_type="FINAL_POLL_BEFORE_TIMEOUT_PENDING", + event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT_PENDING.value, message=f"最终查询后供应商仍未完成,按超时处理。status={status}", detail=poll_result, ) @@ -294,26 +405,17 @@ async def _run(task_id: str): return # 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。 - # 同时登记下一次 poll active,Celery countdown 丢失时可由恢复任务拉起。 - task.pipeline_stage = "waiting_remote" + # 视频任务写入 next_poll_at,由 Beat dispatcher 到期投递;短间隔可保留 countdown 兼容。 + task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value task.retry_count = 0 + await _schedule_next_poll(task, reason="poll_pending_next") await db.commit() - await log_task_event(task, event_type="POLL_PENDING", message=f"status={status}") - - delay_seconds = int(settings.CHATAPI_ASYNC_POLL_INTERVAL_SECONDS or 30) - next_poll_at = _now() + timedelta(seconds=max(1, delay_seconds)) - await register_poll_active( + await log_task_event( task, - check_at=_poll_check_at(delay_seconds=delay_seconds), - next_poll_at=next_poll_at, - reason="poll_pending_next", - ) - - poll_generation_task.apply_async( - args=[task.id], - queue=POLL_QUEUE, - countdown=delay_seconds, + event_type=ChatGenerationTaskEventType.POLL_PENDING.value, + message=f"status={status}", + detail={"next_poll_at": task.next_poll_at, "poll_interval_seconds": task.poll_interval_seconds}, ) except Exception as exc: @@ -331,7 +433,7 @@ async def _run(task_id: str): if final_poll_before_timeout: await log_task_event( task, - event_type="FINAL_POLL_BEFORE_TIMEOUT_ERROR", + event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT_ERROR.value, message=str(exc), ) await _mark_timeout(db, task, message="任务轮询超时") @@ -339,21 +441,34 @@ async def _run(task_id: str): task.retry_count = (task.retry_count or 0) + 1 + # 视频轮询的临时异常不再 3 次内直接退款;继续降频到 24 小时最终 deadline。 + if is_video_generation_task(task): + task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value + await _schedule_next_poll(task, reason="poll_exception_retry") + await db.commit() + return + if task.retry_count > settings.CHATAPI_ASYNC_MAX_RETRIES: error_message = extract_error_message(exc, "轮询") if callable(extract_error_message) else str(exc) await _mark_failed(db, task, message=error_message) else: # 临时轮询异常时,不让任务停在 polling。 # 回到 waiting_remote,等待下一次重试轮询。 - task.pipeline_stage = "waiting_remote" + task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value + delay_seconds = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * int(task.retry_count or 1) + default_schedule = build_default_poll_schedule( + task, + now=_now(), + delay_seconds=delay_seconds, + reason="poll_exception_retry", + ) + task.next_poll_at = default_schedule.next_poll_at await db.commit() - delay_seconds = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * int(task.retry_count or 1) - next_poll_at = _now() + timedelta(seconds=max(1, delay_seconds)) await register_poll_active( task, check_at=_poll_check_at(delay_seconds=delay_seconds), - next_poll_at=next_poll_at, + next_poll_at=default_schedule.next_poll_at, reason="poll_exception_retry", ) diff --git a/video-gen-api/app/tasks/generation_recovery_tasks.py b/video-gen-api/app/tasks/generation_recovery_tasks.py index d96d8ea2..b2f3259e 100644 --- a/video-gen-api/app/tasks/generation_recovery_tasks.py +++ b/video-gen-api/app/tasks/generation_recovery_tasks.py @@ -5,6 +5,7 @@ import logging from typing import Any, Awaitable, Callable, Dict from app.config import settings +from app.enums.celery_queue import CeleryQueue from app.models.base import async_session from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock from app.tasks.async_runner import run_async @@ -12,7 +13,7 @@ from app.tasks.celery_app import celery_app logger = logging.getLogger("video_gen") -RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or "gen_recovery" +RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or CeleryQueue.GEN_RECOVERY.value RecoveryRunner = Callable[[], Awaitable[Dict[str, Any]]] @@ -30,6 +31,13 @@ async def _run_generation_once() -> Dict[str, Any]: return await recover_generation_tasks_once(db) +async def _run_due_poll_dispatch_once() -> Dict[str, Any]: + from app.services.generation_recovery_service import dispatch_due_poll_tasks_once + + async with async_session() as db: + return await dispatch_due_poll_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 @@ -49,6 +57,7 @@ async def _run_with_execution_lock( lock_key: str, log_context: str, runner: RecoveryRunner, + ttl_seconds: int | None = None, ) -> Dict[str, Any]: """恢复任务执行锁。 @@ -61,7 +70,7 @@ async def _run_with_execution_lock( 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), + ttl_seconds=int(ttl_seconds or settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600), log_context=log_context, ) if not token: @@ -79,6 +88,36 @@ async def _run_with_execution_lock( await redis_release_lock(lock_key=lock_key, token=token, log_context=log_context) +async def _is_lock_held(lock_key: str) -> bool: + redis = await get_registry_redis() + if redis is None: + return False + try: + return bool(await redis.exists(lock_key)) + except Exception: + return False + + +async def _startup_or_generation_recovery_running() -> str | None: + # Beat 触发 dispatcher 时,如果启动容灾或完整生成容灾还在跑,直接跳过本轮。 + # gen_recovery concurrency=1 已经能串行;这里是多机部署、残留消息、手动触发时的双保险。 + lock_checks = [ + ("startup_recovery", settings.CELERY_RECOVERY_STARTUP_TASK_LOCK_KEY), + ("generation_recovery", settings.GENERATION_RECOVERY_LOCK_KEY), + ] + for name, lock_key in lock_checks: + if await _is_lock_held(lock_key): + return name + return None + + +async def _run_due_poll_dispatch_with_guard() -> Dict[str, Any]: + running = await _startup_or_generation_recovery_running() + if running: + return {"skipped": "recovery_lock_held", "lock": running} + return await _run_due_poll_dispatch_once() + + async def _acquire_download_recovery_loop_lock() -> tuple[bool, str]: """下载恢复循环锁。 @@ -220,6 +259,23 @@ if celery_app: ) ) + + @celery_app.task( + name="generation.dispatch_due_poll_tasks", + bind=True, + soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS, + time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS, + ) + def dispatch_due_poll_tasks(self) -> Dict[str, Any]: + return run_async( + _run_with_execution_lock( + lock_key=settings.POLL_DUE_DISPATCH_LOCK_KEY, + log_context="due_poll_dispatch", + runner=_run_due_poll_dispatch_with_guard, + ttl_seconds=int(settings.POLL_DUE_DISPATCH_LOCK_TTL_SECONDS or 55), + ) + ) + else: class _DisabledTask: @@ -232,3 +288,4 @@ else: startup_recovery_once = _DisabledTask() recover_download_tasks_once = _DisabledTask() recover_generation_tasks_once = _DisabledTask() + dispatch_due_poll_tasks = _DisabledTask()