celery worker poll 频次机制调整| celery beat 设置poll检测任务| 拆镜复刻切片删除API | 交易流水时区BUG
This commit is contained in:
+83
@@ -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,
|
||||||
|
|||||||
@@ -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_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 注册。
|
# 覆盖 ModuleGenerationStep 提词任务、shot 原视频/片段分析、shot ffmpeg 切割 active 注册。
|
||||||
# 恢复扫描走 CELERY_RECOVERY_QUEUE,真实业务任务回到原始队列。
|
# 恢复扫描走 CELERY_RECOVERY_QUEUE,真实业务任务回到原始队列。
|
||||||
|
|||||||
@@ -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 *
|
||||||
|
|||||||
@@ -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"}
|
||||||
|
|||||||
@@ -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;
|
||||||
@@ -630,4 +658,126 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
|
|||||||
"checked": len(checked_ids),
|
"checked": len(checked_ids),
|
||||||
"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),
|
||||||
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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 投递恢复任务。
|
||||||
"""
|
"""
|
||||||
@@ -182,4 +200,4 @@ def on_worker_process_shutdown(**kwargs):
|
|||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
finally:
|
finally:
|
||||||
close_loop()
|
close_loop()
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -305,4 +333,4 @@ else:
|
|||||||
def apply_async(self, *args, **kwargs):
|
def apply_async(self, *args, **kwargs):
|
||||||
raise RuntimeError("Celery is disabled")
|
raise RuntimeError("Celery is disabled")
|
||||||
|
|
||||||
chatapi_create_generation_task = _DisabledTask()
|
chatapi_create_generation_task = _DisabledTask()
|
||||||
|
|||||||
@@ -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 active,Celery 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()
|
||||||
|
|||||||
Reference in New Issue
Block a user