celery worker poll 频次机制调整| celery beat 设置poll检测任务| 拆镜复刻切片删除API | 交易流水时区BUG

This commit is contained in:
2026-07-02 09:36:38 +08:00
parent c6df04a895
commit f6c5032c7b
22 changed files with 1039 additions and 124 deletions
@@ -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')
@@ -31,6 +31,7 @@ from app.schemas.shot_replicate import (
ShotReplicateSpecOut, ShotReplicateSpecOut,
ShotReplicateTaskDetailOut, ShotReplicateTaskDetailOut,
ShotReplicateVideoPromptSchemaUpdateRequest, ShotReplicateVideoPromptSchemaUpdateRequest,
ShotSegmentDeleteOut,
ShotSegmentDetailOut, ShotSegmentDetailOut,
ShotSegmentListOut, ShotSegmentListOut,
ShotSegmentReplicationCreateRequest, ShotSegmentReplicationCreateRequest,
@@ -60,6 +61,7 @@ from app.services.shot_replicate_taskset_service import (
create_custom_segment, create_custom_segment,
create_segments_by_ai, create_segments_by_ai,
create_task_set, create_task_set,
delete_segment,
get_segment_for_user, get_segment_for_user,
list_segments, list_segments,
list_task_sets, list_task_sets,
@@ -455,6 +457,34 @@ async def get_segment(
return await segment_detail(db, current_user=current_user, segment_id=segment_id) 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( @router.post(
"/segments/{segment_id}/replication-projects", "/segments/{segment_id}/replication-projects",
response_model=ShotReplicateActionOut, response_model=ShotReplicateActionOut,
+16 -1
View File
@@ -121,7 +121,14 @@ class Settings(BaseSettings):
CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS: int = 30 CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS: int = 30
CHATAPI_ASYNC_POLL_INTERVAL_SECONDS: int = 30 CHATAPI_ASYNC_POLL_INTERVAL_SECONDS: int = 30
CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES: int = 10 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. # Distributed provider concurrency limits. 0 means disabled/no-op.
ARK_CHAT_PROMPT_MAX_CONCURRENCY: int = 20 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" MODULE_ASYNC_RECOVERY_LOCK_KEY: str = "vg:celery:module_async_recovery_lock"
SHOT_SPLIT_RECOVERY_LOCK_KEY: str = "vg:celery:shot_split_recovery_lock" SHOT_SPLIT_RECOVERY_LOCK_KEY: str = "vg:celery:shot_split_recovery_lock"
# 视频到期轮询调度。
# Celery Beat 每分钟投递轻量 dispatcher 到 gen_recoverydispatcher 只扫描 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 注册。 # 覆盖 ModuleGenerationStep 提词任务、shot 原视频/片段分析、shot ffmpeg 切割 active 注册。
# 恢复扫描走 CELERY_RECOVERY_QUEUE,真实业务任务回到原始队列。 # 恢复扫描走 CELERY_RECOVERY_QUEUE,真实业务任务回到原始队列。
+1
View File
@@ -14,3 +14,4 @@ from app.enums.notification import *
from app.enums.resource_capacity import * from app.enums.resource_capacity import *
from app.enums.team import * from app.enums.team import *
from app.enums.home_material import * from app.enums.home_material import *
from app.enums.celery_queue import *
+21
View File
@@ -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"
@@ -45,12 +45,18 @@ class ChatGenerationTaskEventType(str, Enum):
POLL_START = "POLL_START" POLL_START = "POLL_START"
POLL_PENDING = "POLL_PENDING" 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_SUCCESS = "POLL_SUCCESS"
POLL_FAILED = "POLL_FAILED" POLL_FAILED = "POLL_FAILED"
POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY = "POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY" 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_ERROR = "FINAL_POLL_BEFORE_TIMEOUT_ERROR"
FINAL_POLL_BEFORE_TIMEOUT_PENDING = "FINAL_POLL_BEFORE_TIMEOUT_PENDING" FINAL_POLL_BEFORE_TIMEOUT_PENDING = "FINAL_POLL_BEFORE_TIMEOUT_PENDING"
GENERATION_RECOVERY_ENQUEUE = "GENERATION_RECOVERY_ENQUEUE" GENERATION_RECOVERY_ENQUEUE = "GENERATION_RECOVERY_ENQUEUE"
GENERATION_RECOVERY_TIMEOUT = "GENERATION_RECOVERY_TIMEOUT"
DOWNLOAD_ENQUEUE = "DOWNLOAD_ENQUEUE" DOWNLOAD_ENQUEUE = "DOWNLOAD_ENQUEUE"
DOWNLOAD_ENQUEUE_FAILED = "DOWNLOAD_ENQUEUE_FAILED" DOWNLOAD_ENQUEUE_FAILED = "DOWNLOAD_ENQUEUE_FAILED"
@@ -78,6 +84,7 @@ class ChatGenerationTaskEventType(str, Enum):
DOWNLOAD_SKIP_DISABLED = "DOWNLOAD_SKIP_DISABLED" DOWNLOAD_SKIP_DISABLED = "DOWNLOAD_SKIP_DISABLED"
TASK_TIMEOUT = "TASK_TIMEOUT" TASK_TIMEOUT = "TASK_TIMEOUT"
TASK_FAILED = "TASK_FAILED"
ALLOWED_GENERATION_MODES = { ALLOWED_GENERATION_MODES = {
@@ -99,3 +106,6 @@ DOWNLOAD_RECOVERABLE_STAGES = {
ChatGenerationPipelineStage.DOWNLOADING.value, ChatGenerationPipelineStage.DOWNLOADING.value,
ChatGenerationPipelineStage.RETRY_WAITING.value, ChatGenerationPipelineStage.RETRY_WAITING.value,
} }
PROVIDER_SUCCESS_STATUSES = {"succeeded", "success", "completed", "done"}
PROVIDER_FAILED_STATUSES = {"failed", "error", "canceled", "cancelled"}
+2 -1
View File
@@ -31,6 +31,7 @@ from app.models.user_oauth import UserOAuth
from app.models.user_oauth_account import UserOAuthAccount from app.models.user_oauth_account import UserOAuthAccount
from app.models.user_oauth_app import UserOAuthApp from app.models.user_oauth_app import UserOAuthApp
from app.models.home_material import HomeMaterialAsset, HomeMaterialCategory, HomeMaterialWatermark from app.models.home_material import HomeMaterialAsset, HomeMaterialCategory, HomeMaterialWatermark
from app.models.contact_request import ContactRequest
__all__ = [ __all__ = [
"Base", "TimestampMixin", "SoftDeleteMixin", "engine", "async_session", "Base", "TimestampMixin", "SoftDeleteMixin", "engine", "async_session",
@@ -38,7 +39,7 @@ __all__ = [
"User", "Team", "Project", "GenerationRecord", "CreditRecord", "User", "Team", "Project", "GenerationRecord", "CreditRecord",
"ModelConfig", "SystemConfig", "Notification", "PaymentOrder", "ModelConfig", "SystemConfig", "Notification", "PaymentOrder",
"TokenUsage", "IndustryConfig", "VideoEngine", "CreditRatio", "TokenUsage", "IndustryConfig", "VideoEngine", "CreditRatio",
"MenuConfig", "RechargePackage", "OperationLog", "MenuConfig", "RechargePackage", "OperationLog", "ContactRequest",
"ChatGenerationTask", "ChatGenerationTaskEvent", "ChatProviderCallLog", "ChatGenerationTask", "ChatGenerationTaskEvent", "ChatProviderCallLog",
"GeneratedResource", "UserResourceMonthStat", "UserResourceTotalStat", "GeneratedResource", "UserResourceMonthStat", "UserResourceTotalStat",
"UserResourceCapacityConfig", "UserResourceCapacityConfig",
@@ -26,6 +26,17 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
unique=True, unique=True,
postgresql_where=text("deleted_at IS NULL AND idempotency_key IS NOT NULL"), 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) retry_count: Mapped[int] = mapped_column(Integer, default=0)
poll_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) 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) deadline_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
generated_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) error_message: Mapped[str | None] = mapped_column(Text, nullable=True)
@@ -484,6 +484,16 @@ class ShotReplicateDeleteOut(BaseModel):
message: str = Field(..., description="删除结果提示") message: str = Field(..., description="删除结果提示")
project_id: str = Field(..., description="被软删除的总任务项目ID") project_id: str = Field(..., description="被软删除的总任务项目ID")
deleted: bool = Field(..., description="是否已软删除") 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="本次释放的用户容量占用字节数;只释放数据账本,不删除物理文件")
@@ -1,6 +1,6 @@
from __future__ import annotations from __future__ import annotations
from datetime import datetime from datetime import datetime, timezone, timedelta
from typing import Any from typing import Any
from sqlalchemy import and_, case, distinct, func, or_, select 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 from app.models.user import User
CST = timezone(timedelta(hours=8))
def _iso(dt: Any) -> str | None: def _iso(dt: Any) -> str | None:
if dt is None: if dt is None:
return 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: try:
return dt.isoformat() return dt.isoformat()
except Exception: except Exception:
return str(dt) return str(dt)
def _round2(value: Any) -> float: def _round2(value: Any) -> float:
try: try:
return round(float(value or 0), 2) return round(float(value or 0), 2)
@@ -331,7 +331,7 @@ async def create_async_generation_task(db: AsyncSession, current_user: User, req
media_references=_json(refs) if refs else None, media_references=_json(refs) if refs else None,
credits_cost=round(media_billing.total_charged, 2), credits_cost=round(media_billing.total_charged, 2),
idempotency_key=req.idempotency_key, 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) db.add(task)
@@ -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,
)
@@ -9,11 +9,13 @@ from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings from app.config import settings
from app.enums.celery_queue import CeleryQueue
from app.enums.generation_task import ( from app.enums.generation_task import (
ALLOWED_GENERATION_MODES, ALLOWED_GENERATION_MODES,
ChatGenerationPipelineStage, ChatGenerationPipelineStage,
ChatGenerationTaskEventType, ChatGenerationTaskEventType,
ChatGenerationTaskStatus, ChatGenerationTaskStatus,
GenerationType,
) )
from app.models.chat_generation_task import ChatGenerationTask from app.models.chat_generation_task import ChatGenerationTask
from app.services.celery_download_recovery_service import ( 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_log_service import log_task_event
from app.services.generation_module_hook_service import notify_chat_generation_task_finished 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.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.redis_registry_service import ( from app.services.redis_registry_service import (
redis_get_due_registry_ids, redis_get_due_registry_ids,
@@ -35,7 +38,7 @@ from app.services.redis_registry_service import (
logger = logging.getLogger("video_gen") logger = logging.getLogger("video_gen")
POLL_QUEUE = "gen_provider_poll" POLL_QUEUE = CeleryQueue.GEN_PROVIDER_POLL.value
def _now() -> datetime: def _now() -> datetime:
@@ -342,7 +345,7 @@ async def _mark_timeout(
await _remove_poll_active(task.id) await _remove_poll_active(task.id)
await log_task_event( await log_task_event(
task, task,
event_type="TASK_TIMEOUT", event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
to_status="failed", to_status="failed",
to_stage=ChatGenerationPipelineStage.TIMEOUT.value, to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
) )
@@ -414,7 +417,7 @@ async def recover_one_generation_task(
await _remove_poll_active(task.id) await _remove_poll_active(task.id)
await log_task_event( await log_task_event(
task, task,
event_type="GENERATION_RECOVERY_ENQUEUE", event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
message=f"{source} 发现任务已存在 remote_result_url,恢复投递下载队列", message=f"{source} 发现任务已存在 remote_result_url,恢复投递下载队列",
detail={ detail={
"pipeline_stage": task.pipeline_stage, "pipeline_stage": task.pipeline_stage,
@@ -439,7 +442,7 @@ async def recover_one_generation_task(
await db.commit() await db.commit()
await log_task_event( await log_task_event(
task, task,
event_type="GENERATION_RECOVERY_ENQUEUE", event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
message=f"{source} 发现任务已到 deadline 且存在供应商任务ID,投递 poll 队列做最终查询", message=f"{source} 发现任务已到 deadline 且存在供应商任务ID,投递 poll 队列做最终查询",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
) )
@@ -453,26 +456,51 @@ async def recover_one_generation_task(
await log_task_event( await log_task_event(
task, task,
event_type="GENERATION_RECOVERY_TIMEOUT", event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_TIMEOUT.value,
message=f"{source} 发现任务已到 deadline,且没有 remote_result_url/供应商任务ID,按超时失败处理", message=f"{source} 发现任务已到 deadline,且没有 remote_result_url/供应商任务ID,按超时失败处理",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
) )
return await _mark_timeout(db, task) return await _mark_timeout(db, task)
# 未过 deadline:有供应商任务 ID 才允许恢复到 poll 队列。 # 未过 deadline:有供应商任务 ID 才允许恢复到 poll 队列。
# 视频任务如果 next_poll_at 未到期,不提前 poll,只刷新 active 注册表等待 Beat dispatcher 到期投递。
if has_provider_task_id: 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.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
task.next_poll_at = _poll_queue_timeout_at(current_time)
await db.commit() await db.commit()
await log_task_event( await log_task_event(
task, task,
event_type="GENERATION_RECOVERY_ENQUEUE", event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
message=f"{source} 发现任务存在供应商任务ID,恢复投递轮询队列", message=f"{source} 发现任务存在供应商任务ID,恢复投递轮询队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
) )
poll_generation_task.apply_async(args=[task.id], queue=POLL_QUEUE, countdown=0) poll_generation_task.apply_async(args=[task.id], queue=POLL_QUEUE, countdown=0)
await register_poll_active( await register_poll_active(
task, 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", reason=f"{source}_has_provider_task_id",
) )
return "recover_poll_has_provider_id" return "recover_poll_has_provider_id"
@@ -499,13 +527,13 @@ async def recover_one_generation_task(
await _remove_poll_active(task.id) await _remove_poll_active(task.id)
await log_task_event( await log_task_event(
task, task,
event_type="GENERATION_RECOVERY_ENQUEUE", event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
message=f"{source} 发现任务未超时且缺少 remote_result_url/供应商任务ID,恢复投递创建队列", message=f"{source} 发现任务未超时且缺少 remote_result_url/供应商任务ID,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
) )
chatapi_create_generation_task.apply_async( chatapi_create_generation_task.apply_async(
args=[task.id], args=[task.id],
queue="gen_chatapi_create", queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0, countdown=0,
) )
return "recover_create_no_remote_no_provider_before_deadline" 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 _remove_poll_active(task.id)
await log_task_event( await log_task_event(
task, task,
event_type="GENERATION_RECOVERY_ENQUEUE", event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 result_ready 但缺少 remote_result_url,未超时,恢复投递创建队列", message=f"{source} 发现 result_ready 但缺少 remote_result_url,未超时,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload}, detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
) )
chatapi_create_generation_task.apply_async( chatapi_create_generation_task.apply_async(
args=[task.id], args=[task.id],
queue="gen_chatapi_create", queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0, countdown=0,
) )
return "recover_create_result_ready_no_url_before_deadline" 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]: 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 到期任务; 1. Redis poll active_index 到期任务;
2. DB fallback 扫描 queued/creating/waiting_remote/polling/result_ready 2. DB fallback 扫描 queued/creating/waiting_remote/polling/result_ready
@@ -631,3 +659,125 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
"db_checked": total_db_checked, "db_checked": total_db_checked,
"results": results, "results": results,
} }
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,
}
@@ -195,7 +195,7 @@ async def create_chat_generation_task_for_module(
media_references=_json(refs) if refs else None, media_references=_json(refs) if refs else None,
credits_cost=round(media_billing.total_charged, 2), credits_cost=round(media_billing.total_charged, 2),
idempotency_key=backend_idempotency_key, 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) db.add(task)
@@ -232,6 +232,8 @@ async def soft_delete_steps_from_index(
config: ModuleGenerationFlowConfig, config: ModuleGenerationFlowConfig,
log_module_event: LogModuleEventCallable, log_module_event: LogModuleEventCallable,
deleted_at: datetime | None = None, deleted_at: datetime | None = None,
refund_unfinished: bool = True,
release_stats: dict[str, int] | None = None,
) -> list[ModuleGenerationStep]: ) -> list[ModuleGenerationStep]:
deleted_at = deleted_at or utc_now() deleted_at = deleted_at or utc_now()
result = await db.execute( result = await db.execute(
@@ -266,8 +268,10 @@ async def soft_delete_steps_from_index(
chat_task = chat_result.scalar_one_or_none() chat_task = chat_result.scalar_one_or_none()
if chat_task: if chat_task:
if chat_task.status == "completed": if chat_task.status == "completed":
await soft_delete_chat_task_resources(db, chat_task.id, deleted_at=deleted_at) released_size = await soft_delete_chat_task_resources(db, chat_task.id, deleted_at=deleted_at)
elif chat_task.status != "failed": 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( await mark_chat_generation_task_failed_and_refund_once(
db, db,
task=chat_task, task=chat_task,
@@ -18,6 +18,7 @@ from app.utils.id_gen import generate_id
SOURCE_MODEL_CHAT_TASK = "ChatGenerationTask" SOURCE_MODEL_CHAT_TASK = "ChatGenerationTask"
SOURCE_MODEL_GENERATION_RECORD = "GenerationRecord" SOURCE_MODEL_GENERATION_RECORD = "GenerationRecord"
SOURCE_MODEL_SHOT_SEGMENT = "ShotReplicateSegment"
@dataclass(slots=True) @dataclass(slots=True)
@@ -12,6 +12,7 @@ from sqlalchemy.orm.attributes import flag_modified
from app.config import settings from app.config import settings
from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum 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.enums.shot_replicate import ShotReplicateGenerationModeEnum, ShotReplicateStepCodeEnum, ModuleCodeEnum
from app.models.chat_generation_task import ChatGenerationTask from app.models.chat_generation_task import ChatGenerationTask
from app.models.module_generation_project import ModuleGenerationProject from app.models.module_generation_project import ModuleGenerationProject
@@ -335,6 +336,8 @@ async def _soft_delete_steps_from_index(
project: ModuleGenerationProject, project: ModuleGenerationProject,
start_index: int, start_index: int,
deleted_at: datetime | None = None, deleted_at: datetime | None = None,
refund_unfinished: bool = True,
release_stats: dict[str, int] | None = None,
) -> None: ) -> None:
await _base_soft_delete_steps_from_index( await _base_soft_delete_steps_from_index(
db, db,
@@ -343,6 +346,8 @@ async def _soft_delete_steps_from_index(
config=FLOW_CONFIG, config=FLOW_CONFIG,
log_module_event=log_module_event, log_module_event=log_module_event,
deleted_at=deleted_at, 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}) 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( async def mark_shot_replicate_step_dispatch_failed(
db: AsyncSession, 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) project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True)
deleted_at = _now() 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 project.deleted_at = deleted_at
await _soft_delete_steps_from_index(db, project=project, start_index=1, deleted_at=deleted_at) await _soft_delete_steps_from_index(
await log_module_event(db, project=project, event_type=ModuleEventTypeEnum.PROJECT_DELETED.value, message="软删除拆镜复刻项目") db,
return ShotReplicateDeleteOut(message="项目已删除", project_id=project.id, deleted=True) 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( async def create_shot_replicate_project_from_segment(
@@ -25,6 +25,7 @@ from app.models.shot_replicate_task_set import ShotReplicateTaskSet
from app.models.user import User from app.models.user import User
from app.schemas.shot_replicate import ( from app.schemas.shot_replicate import (
ShotAISuggestionOut, ShotAISuggestionOut,
ShotSegmentDeleteOut,
ShotSegmentDetailOut, ShotSegmentDetailOut,
ShotSegmentListOut, ShotSegmentListOut,
ShotSegmentOut, ShotSegmentOut,
@@ -38,6 +39,7 @@ from app.schemas.shot_replicate import (
ShotTaskSetOut, ShotTaskSetOut,
) )
from app.services.module_generation_log_service import log_module_event_file 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 ( from app.services.upload_video_asset_service import (
build_time_node, build_time_node,
validate_split_range, 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_result = await db.execute(select(ModuleGenerationProject).where(ModuleGenerationProject.id == segment.module_project_id).limit(1))
project = project_result.scalar_one_or_none() project = project_result.scalar_one_or_none()
return _segment_to_detail_out(segment, project) 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),
)
+38 -20
View File
@@ -4,6 +4,7 @@ from celery import Celery
from celery.signals import worker_process_init, worker_process_shutdown, worker_ready from celery.signals import worker_process_init, worker_process_shutdown, worker_ready
from app.config import settings from app.config import settings
from app.enums.celery_queue import CeleryQueue, CeleryTaskName
from app.models.base import engine from app.models.base import engine
from app.tasks.async_runner import close_loop, run_async 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: 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}" 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 "") 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 "") 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_acks_late=True,
task_reject_on_worker_lost=True, task_reject_on_worker_lost=True,
task_track_started=True, task_track_started=True,
beat_schedule=_beat_schedule(),
task_annotations={ task_annotations={
# 生成链路任务以数据库状态为准,不依赖 Celery result backend。 # 生成链路任务以数据库状态为准,不依赖 Celery result backend。
# 这里忽略结果可避免任务误返回 ORM / 非 JSON 对象时触发结果序列化失败。 # 这里忽略结果可避免任务误返回 ORM / 非 JSON 对象时触发结果序列化失败。
@@ -80,24 +97,25 @@ if broker_url:
"sep": ":", "sep": ":",
}, },
task_routes={ task_routes={
"generation.chatapi_create_generation_task": {"queue": "gen_chatapi_create"}, CeleryTaskName.CHATAPI_CREATE.value: {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"generation.poll_generation_task": {"queue": "gen_provider_poll"}, CeleryTaskName.POLL_GENERATION.value: {"queue": CeleryQueue.GEN_PROVIDER_POLL.value},
"generation.download_generation_result_task": {"queue": "gen_result_download"}, CeleryTaskName.DOWNLOAD_GENERATION_RESULT.value: {"queue": CeleryQueue.GEN_RESULT_DOWNLOAD.value},
"hot_opening.start_image_prompt_optimize": {"queue": "gen_chatapi_create"}, CeleryTaskName.DISPATCH_DUE_POLL.value: {"queue": RECOVERY_QUEUE},
"hot_opening.start_video_prompt_optimize": {"queue": "gen_chatapi_create"}, "hot_opening.start_image_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"shot_replicate.analyze_original_video": {"queue": "gen_chatapi_create"}, "hot_opening.start_video_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"shot_replicate.analyze_custom_segment_video": {"queue": "gen_chatapi_create"}, "shot_replicate.analyze_original_video": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"shot_replicate.split_one_segment": {"queue": "gen_result_download"}, "shot_replicate.analyze_custom_segment_video": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"shot_replicate.start_image_prompt_optimize": {"queue": "gen_chatapi_create"}, "shot_replicate.split_one_segment": {"queue": CeleryQueue.GEN_RESULT_DOWNLOAD.value},
"shot_replicate.start_video_prompt_optimize": {"queue": "gen_chatapi_create"}, "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。 # 恢复扫描统一走独立队列,避免占用下载/轮询/创建业务 worker。
"recovery.startup_recovery_once": {"queue": RECOVERY_QUEUE}, CeleryTaskName.STARTUP_RECOVERY.value: {"queue": RECOVERY_QUEUE},
"shot_replicate.recover_split_tasks_once": {"queue": RECOVERY_QUEUE}, CeleryTaskName.SHOT_SPLIT_RECOVERY.value: {"queue": RECOVERY_QUEUE},
"generation.recover_download_tasks_once": {"queue": RECOVERY_QUEUE}, CeleryTaskName.RECOVER_DOWNLOAD.value: {"queue": RECOVERY_QUEUE},
"generation.recover_generation_tasks_once": {"queue": RECOVERY_QUEUE}, CeleryTaskName.RECOVER_GENERATION.value: {"queue": RECOVERY_QUEUE},
"module_async.recover_module_async_tasks_once": {"queue": RECOVERY_QUEUE}, CeleryTaskName.MODULE_ASYNC_RECOVERY.value: {"queue": RECOVERY_QUEUE},
"user_oauth.update_oauth_accounts": {"queue": "default"}, "user_oauth.update_oauth_accounts": {"queue": CeleryQueue.DEFAULT.value},
"app.tasks.cleanup.*": {"queue": "default"}, "app.tasks.cleanup.*": {"queue": CeleryQueue.DEFAULT.value},
}, },
) )
else: else:
@@ -121,8 +139,8 @@ def on_worker_ready(sender=None, **kwargs):
"""Celery worker 启动时做一次容灾恢复。 """Celery worker 启动时做一次容灾恢复。
注意: 注意:
- 不启用 Celery beat - 启动容灾只投递一个 recovery.startup_recovery_once 协调任务
- 启动容灾保留,但只投递一个 recovery.startup_recovery_once 协调任务 - Celery Beat 只用于每分钟触发轻量 generation.dispatch_due_poll_tasks,不跑完整启动容灾
- 协调任务走独立 gen_recovery 队列,串行扫描并把真实业务任务投回原队列。 - 协调任务走独立 gen_recovery 队列,串行扫描并把真实业务任务投回原队列。
- 所有 worker 都尝试抢 Redis 投递锁,只有抢到锁的 worker 投递恢复任务。 - 所有 worker 都尝试抢 Redis 投递锁,只有抢到锁的 worker 投递恢复任务。
""" """
@@ -6,21 +6,31 @@ from typing import Any, Optional
from sqlalchemy import select from sqlalchemy import select
from app.config import settings 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.base import async_session
from app.models.chat_generation_task import ChatGenerationTask from app.models.chat_generation_task import ChatGenerationTask
from app.services.error_codes import extract_error_message from app.services.error_codes import extract_error_message
from app.services.generation_log_service import log_task_event 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_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_provider_service import create_provider_task 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.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.services.redis_registry_service import ensure_aware_utc
from app.tasks.celery_app import celery_app from app.tasks.celery_app import celery_app
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"}
def _now() -> datetime: def _now() -> datetime:
return datetime.now(timezone.utc) return datetime.now(timezone.utc)
def _get_first_value(obj: Any, *field_names: str) -> Optional[Any]: def _get_first_value(obj: Any, *field_names: str) -> Optional[Any]:
""" """
兼容不同版本字段名,避免字段调整后 Celery 任务直接报错。 兼容不同版本字段名,避免字段调整后 Celery 任务直接报错。
@@ -74,7 +84,7 @@ def _build_optimized_prompt_by_params(task: ChatGenerationTask) -> str:
# 爆款开头复刻第5步的视频生成,original_prompt 已经是视频提词 JSON schema。 # 爆款开头复刻第5步的视频生成,original_prompt 已经是视频提词 JSON schema。
# 不能再追加“时长/比例/分辨率”中文参数,否则会污染 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() stripped = base_prompt.strip()
if stripped.startswith("{") or stripped.startswith("["): if stripped.startswith("{") or stripped.startswith("["):
return base_prompt return base_prompt
@@ -88,25 +98,25 @@ def _build_optimized_prompt_by_params(task: ChatGenerationTask) -> str:
parts = [] parts = []
if gen_type == "video": if gen_type == GenerationType.VIDEO.value:
# 时长:4秒,画面比例:16:9,分辨率:480p # 时长:4秒,画面比例:16:9,分辨率:480p
if duration: if duration:
parts.append(f"时长:{duration}") parts.append(f"时长:{duration}")
parts.append(f"画面比例:{aspect_ratio}") parts.append(f"画面比例:{aspect_ratio}")
parts.append(f"分辨率:{resolution}") parts.append(f"分辨率:{resolution}")
else: else:
parts.append(f"时长:4秒") parts.append("时长:4秒")
parts.append(f"画面比例:16:9") parts.append("画面比例:16:9")
parts.append(f"分辨率:480p") parts.append("分辨率:480p")
elif gen_type == "image": elif gen_type == GenerationType.IMAGE.value:
if image_size : if image_size:
parts.append(f"分辨率:{image_size}") parts.append(f"分辨率:{image_size}")
parts.append(f"画布比例:{image_proportion}") parts.append(f"画布比例:{image_proportion}")
parts.append(f"像素尺寸:{image_px}") parts.append(f"像素尺寸:{image_px}")
else: else:
parts.append(f"分辨率:2K") parts.append("分辨率:2K")
parts.append(f"画布比例:1:1") parts.append("画布比例:1:1")
parts.append(f"像素尺寸:2048x2048") parts.append("像素尺寸:2048x2048")
else: else:
# 未知类型时返回原始字符 # 未知类型时返回原始字符
return base_prompt 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: if not task or task.generation_mode not in ALLOWED_GENERATION_MODES:
return return
if task.status != "generating": if task.status != ChatGenerationTaskStatus.GENERATING.value:
return return
deadline_at = ensure_aware_utc(task.deadline_at) deadline_at = ensure_aware_utc(task.deadline_at)
@@ -140,28 +150,37 @@ async def _run(task_id: str):
db, db,
task=task, task=task,
error_message="任务超时", error_message="任务超时",
pipeline_stage="timeout", pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
) )
await db.commit() 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 from app.services.generation_module_hook_service import notify_chat_generation_task_finished
await notify_chat_generation_task_finished(db, task) await notify_chat_generation_task_finished(db, task)
await db.commit() await db.commit()
return 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 return
try: try:
if not task.optimized_prompt: if not task.optimized_prompt:
old_stage = task.pipeline_stage old_stage = task.pipeline_stage
task.pipeline_stage = "preparing" task.pipeline_stage = ChatGenerationPipelineStage.PREPARING.value
await db.commit() await db.commit()
await log_task_event( await log_task_event(
task, task,
event_type="PROMPT_CONCAT_START", event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_START.value,
from_stage=old_stage, from_stage=old_stage,
to_stage="preparing", to_stage=ChatGenerationPipelineStage.PREPARING.value,
message="开始本地拼接提示词,不调用提词优化API", message="开始本地拼接提示词,不调用提词优化API",
) )
@@ -174,8 +193,8 @@ async def _run(task_id: str):
await db.commit() await db.commit()
await log_task_event( await log_task_event(
task, task,
event_type="PROMPT_CONCAT_SUCCESS", event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_SUCCESS.value,
to_stage="preparing", to_stage=ChatGenerationPipelineStage.PREPARING.value,
detail={ detail={
"optimized_prompt": optimized_prompt, "optimized_prompt": optimized_prompt,
"gen_type": task.gen_type, "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: 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() await db.commit()
elif task.remote_result_url: elif task.remote_result_url:
task.pipeline_stage = "result_ready" task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
await db.commit() await db.commit()
from app.tasks.generation_download_tasks import enqueue_download_task from app.tasks.generation_download_tasks import enqueue_download_task
@@ -198,13 +220,13 @@ async def _run(task_id: str):
else: else:
old_stage = task.pipeline_stage old_stage = task.pipeline_stage
task.pipeline_stage = "creating_provider_task" task.pipeline_stage = ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value
await db.commit() await db.commit()
await log_task_event( await log_task_event(
task, task,
event_type="PROVIDER_CREATE_START", event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_START.value,
from_stage=old_stage, from_stage=old_stage,
to_stage="creating_provider_task", to_stage=ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
) )
created = await create_provider_task(db, task) 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 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.image_tokens_used = created.get("image_tokens", task.image_tokens_used or 0) or 0
task.provider_response_json = json.dumps( 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: if task.remote_result_url and not task.seedance_task_id:
# 同步图片路径:原 SDK 已经返回最终 URL。 # 同步图片路径:原 SDK 已经返回最终 URL。
task.pipeline_stage = "result_ready" task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
else: else:
# 视频路径:provider 返回 task id,后续轮询。 # 视频路径: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 db.commit()
await log_task_event( await log_task_event(
task, task,
event_type="PROVIDER_CREATE_SUCCESS", event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_SUCCESS.value,
to_stage=task.pipeline_stage, to_stage=task.pipeline_stage,
detail=created, 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 from app.tasks.generation_download_tasks import enqueue_download_task
await enqueue_download_task(db, task, reason="create_result_ready") await enqueue_download_task(db, task, reason="create_result_ready")
@@ -253,10 +280,11 @@ async def _run(task_id: str):
task, task,
reason="create_provider_success", reason="create_provider_success",
check_at=_now() + timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300)), 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( poll_generation_task.apply_async(
args=[task.id], args=[task.id],
queue="gen_provider_poll", queue=CeleryQueue.GEN_PROVIDER_POLL.value,
countdown=0, countdown=0,
) )
@@ -278,10 +306,10 @@ async def _run(task_id: str):
db, db,
task=task, task=task,
error_message=error_message, error_message=error_message,
pipeline_stage="failed", pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
) )
await db.commit() 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 from app.services.generation_module_hook_service import notify_chat_generation_task_finished
await notify_chat_generation_task_finished(db, task) await notify_chat_generation_task_finished(db, task)
await db.commit() await db.commit()
+156 -41
View File
@@ -6,10 +6,28 @@ from typing import Any
from sqlalchemy import select from sqlalchemy import select
from app.config import settings 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.base import async_session
from app.models.chat_generation_task import ChatGenerationTask from app.models.chat_generation_task import ChatGenerationTask
from app.services.error_codes import extract_error_message 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_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_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_provider_service import poll_provider_task 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.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 from app.tasks.celery_app import celery_app
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"} POLL_QUEUE = CeleryQueue.GEN_PROVIDER_POLL.value
POLL_QUEUE = "gen_provider_poll"
def _now() -> datetime: def _now() -> datetime:
return datetime.now(timezone.utc) return datetime.now(timezone.utc)
def _is_success(status: str) -> bool: def _is_success(status: str | None) -> bool:
return status in ("succeeded", "success", "completed", "done") return str(status or "").lower() in PROVIDER_SUCCESS_STATUSES
def _is_failed(status: str) -> bool: def _is_failed(status: str | None) -> bool:
return status in ("failed", "error", "canceled", "cancelled") return str(status or "").lower() in PROVIDER_FAILED_STATUSES
def _engine_snapshot(task: ChatGenerationTask) -> dict: 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: def _deadline_expired(task: ChatGenerationTask, now: datetime | None = None) -> bool:
deadline_at = ensure_aware_utc(task.deadline_at) return is_final_poll_due(task, now=now)
return bool(deadline_at and deadline_at <= (now or _now()))
def _poll_check_at(*, delay_seconds: int | float | None = None, now: datetime | None = None) -> datetime: 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, "queue": POLL_QUEUE,
"poll_count": int(task.poll_count or 0), "poll_count": int(task.poll_count or 0),
"retry_count": int(task.retry_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, "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, "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, "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, db,
task=task, task=task,
error_message=message, error_message=message,
pipeline_stage="timeout", pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
) )
task.next_poll_at = None
await _notify_finished(db, task) await _notify_finished(db, task)
await db.commit() await db.commit()
await remove_poll_active(task.id) 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: 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, db,
task=task, task=task,
error_message=message, error_message=message,
pipeline_stage="failed", pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
) )
task.next_poll_at = None
await _notify_finished(db, task) await _notify_finished(db, task)
await db.commit() await db.commit()
await remove_poll_active(task.id) 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): async def _run(task_id: str):
@@ -190,11 +286,23 @@ async def _run(task_id: str):
return 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) await remove_poll_active(task.id)
return 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): if final_poll_before_timeout and not (task.seedance_task_id or task.provider_task_id):
await _mark_timeout(db, task, message="任务轮询超时") await _mark_timeout(db, task, message="任务轮询超时")
return return
@@ -206,21 +314,23 @@ async def _run(task_id: str):
if final_poll_before_timeout: if final_poll_before_timeout:
await log_task_event( await log_task_event(
task, task,
event_type="FINAL_POLL_BEFORE_TIMEOUT", event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT.value,
message="任务已到 deadline,执行最后一次供应商查询后再判定超时", message="任务已到 deadline,执行最后一次供应商查询后再判定超时",
detail={"deadline_at": task.deadline_at, "stage": task.pipeline_stage}, detail={"deadline_at": task.deadline_at, "stage": task.pipeline_stage},
) )
# 标记本次正在轮询,并登记 poll lease。 # 标记本次正在轮询,并登记 poll lease。
# 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递。 # 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递。
task.pipeline_stage = "polling" task.pipeline_stage = ChatGenerationPipelineStage.POLLING.value
task.poll_count = (task.poll_count or 0) + 1 task.poll_count = (task.poll_count or 0) + 1
task.last_poll_at = _now() task.last_poll_at = _now()
task.next_poll_at = _poll_lease_until(task.last_poll_at)
await db.commit() await db.commit()
await register_poll_active( await register_poll_active(
task, task,
check_at=_poll_lease_until(task.last_poll_at), check_at=_poll_lease_until(task.last_poll_at),
reason="polling_lease", reason="polling_lease",
next_poll_at=task.next_poll_at,
) )
try: try:
@@ -247,7 +357,7 @@ async def _run(task_id: str):
) )
if _is_success(status): 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.remote_result_url = poll_result.get("image_url")
task.image_tokens_used = poll_result.get("image_tokens", 0) or 0 task.image_tokens_used = poll_result.get("image_tokens", 0) or 0
else: else:
@@ -261,12 +371,13 @@ async def _run(task_id: str):
await _mark_failed(db, task, message="供应商任务成功但未返回结果URL", detail=poll_result) await _mark_failed(db, task, message="供应商任务成功但未返回结果URL", detail=poll_result)
return return
task.pipeline_stage = "result_ready" task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
task.retry_count = 0 task.retry_count = 0
task.next_poll_at = None
await db.commit() await db.commit()
await remove_poll_active(task.id) 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 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: if final_poll_before_timeout:
await log_task_event( await log_task_event(
task, task,
event_type="FINAL_POLL_BEFORE_TIMEOUT_PENDING", event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT_PENDING.value,
message=f"最终查询后供应商仍未完成,按超时处理。status={status}", message=f"最终查询后供应商仍未完成,按超时处理。status={status}",
detail=poll_result, detail=poll_result,
) )
@@ -294,26 +405,17 @@ async def _run(task_id: str):
return return
# 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。 # 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。
# 同时登记下一次 poll activeCelery countdown 丢失时可由恢复任务拉起 # 视频任务写入 next_poll_at,由 Beat dispatcher 到期投递;短间隔可保留 countdown 兼容
task.pipeline_stage = "waiting_remote" task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
task.retry_count = 0 task.retry_count = 0
await _schedule_next_poll(task, reason="poll_pending_next")
await db.commit() await db.commit()
await log_task_event(task, event_type="POLL_PENDING", message=f"status={status}") await log_task_event(
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(
task, task,
check_at=_poll_check_at(delay_seconds=delay_seconds), event_type=ChatGenerationTaskEventType.POLL_PENDING.value,
next_poll_at=next_poll_at, message=f"status={status}",
reason="poll_pending_next", detail={"next_poll_at": task.next_poll_at, "poll_interval_seconds": task.poll_interval_seconds},
)
poll_generation_task.apply_async(
args=[task.id],
queue=POLL_QUEUE,
countdown=delay_seconds,
) )
except Exception as exc: except Exception as exc:
@@ -331,7 +433,7 @@ async def _run(task_id: str):
if final_poll_before_timeout: if final_poll_before_timeout:
await log_task_event( await log_task_event(
task, task,
event_type="FINAL_POLL_BEFORE_TIMEOUT_ERROR", event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT_ERROR.value,
message=str(exc), message=str(exc),
) )
await _mark_timeout(db, task, message="任务轮询超时") 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 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: if task.retry_count > settings.CHATAPI_ASYNC_MAX_RETRIES:
error_message = extract_error_message(exc, "轮询") if callable(extract_error_message) else str(exc) error_message = extract_error_message(exc, "轮询") if callable(extract_error_message) else str(exc)
await _mark_failed(db, task, message=error_message) await _mark_failed(db, task, message=error_message)
else: else:
# 临时轮询异常时,不让任务停在 polling。 # 临时轮询异常时,不让任务停在 polling。
# 回到 waiting_remote,等待下一次重试轮询。 # 回到 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() 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( await register_poll_active(
task, task,
check_at=_poll_check_at(delay_seconds=delay_seconds), 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", reason="poll_exception_retry",
) )
@@ -5,6 +5,7 @@ import logging
from typing import Any, Awaitable, Callable, Dict from typing import Any, Awaitable, Callable, Dict
from app.config import settings from app.config import settings
from app.enums.celery_queue import CeleryQueue
from app.models.base import async_session 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.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.async_runner import run_async
@@ -12,7 +13,7 @@ from app.tasks.celery_app import celery_app
logger = logging.getLogger("video_gen") 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]]] 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) 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]: async def _run_module_async_once() -> Dict[str, Any]:
from app.services.module_async_recovery_service import recover_module_async_tasks_once 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, lock_key: str,
log_context: str, log_context: str,
runner: RecoveryRunner, runner: RecoveryRunner,
ttl_seconds: int | None = None,
) -> Dict[str, Any]: ) -> Dict[str, Any]:
"""恢复任务执行锁。 """恢复任务执行锁。
@@ -61,7 +70,7 @@ async def _run_with_execution_lock(
if redis is not None: if redis is not None:
token = await redis_acquire_lock( token = await redis_acquire_lock(
lock_key=lock_key, 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, log_context=log_context,
) )
if not token: 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) 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]: 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: else:
class _DisabledTask: class _DisabledTask:
@@ -232,3 +288,4 @@ else:
startup_recovery_once = _DisabledTask() startup_recovery_once = _DisabledTask()
recover_download_tasks_once = _DisabledTask() recover_download_tasks_once = _DisabledTask()
recover_generation_tasks_once = _DisabledTask() recover_generation_tasks_once = _DisabledTask()
dispatch_due_poll_tasks = _DisabledTask()