Merge branch 'main' of gitee.com:wg123/video-gen into main

This commit is contained in:
18610128193
2026-06-26 14:00:27 +08:00
18 changed files with 1312 additions and 374 deletions
+2
View File
@@ -43,6 +43,7 @@ from app.services.generation_billing_service import (
get_next_credit_attempt_no,
)
from app.services.generation_refund_service import mark_generation_record_failed_and_refund_once
from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot
from app.services.credit_record_meta_service import build_generation_record_prompt_meta
from app.services.video_cover_service import async_create_video_cover_for_local_video
from app.utils.id_gen import generate_id
@@ -724,6 +725,7 @@ async def seedance_callback(request: Request, db: AsyncSession = Depends(get_db)
usage = data.get("usage", {})
if usage:
record.video_tokens_used = usage.get("total_tokens", 0)
await sync_generation_record_media_token_snapshot(db, record, provider_response=data)
# Log callback response
from app.services.video_gen import _log_video_response
_log_video_response(record.id, data)
+29 -11
View File
@@ -113,8 +113,8 @@ class Settings(BaseSettings):
PROVIDER_LIMIT_TOKEN_TTL_SECONDS: int = 600
CELERY_DB_POOL_SIZE: int = 1
CELERY_DB_MAX_OVERFLOW: int = 1
CELERY_DB_POOL_TIMEOUT: int = 30
CELERY_DB_MAX_OVERFLOW: int = 2
CELERY_DB_POOL_TIMEOUT: int = 60
CELERY_DB_POOL_RECYCLE: int = 1800
# Celery 图片/视频下载容灾配置。
@@ -125,30 +125,48 @@ class Settings(BaseSettings):
DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS: int = 30
DOWNLOAD_TASK_LEASE_SECONDS: int = 10 * 60
DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS: int = 5 * 60
DOWNLOAD_RECOVERY_BATCH_SIZE: int = 100
DOWNLOAD_RECOVERY_BATCH_SIZE: int = 20
DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS: int = 3
# 下载恢复自循环:不依赖 Celery beat,不新增 worker;由 gen_result_download 队列周期扫描 DB/Redis。
DOWNLOAD_RECOVERY_LOOP_ENABLED: bool = False
DOWNLOAD_RECOVERY_INTERVAL_SECONDS: int = 60
DOWNLOAD_RECOVERY_LOOP_LOCK_KEY: str = "vg:celery:download_recovery_loop_lock"
DOWNLOAD_RECOVERY_LOOP_LOCK_TTL_SECONDS: int = 55
DOWNLOAD_RETRY_COUNTDOWN_EXTRA_SECONDS: int = 1
DOWNLOAD_NON_RETRYABLE_LOCAL_ERRORS: bool = True
DOWNLOAD_EVENT_VERBOSE_ENABLED: bool = True
MEDIA_TOKEN_SNAPSHOT_ENABLED: bool = True
DOWNLOAD_ACTIVE_REDIS_HASH_KEY: str = "vg:celery:download:active"
DOWNLOAD_ACTIVE_REDIS_ZSET_KEY: str = "vg:celery:download:active_index"
# Celery 生成链路 / provider poll 容灾配置。
# 说明:
# - 不新增 Celery worker;恢复任务仍投递到 gen_result_download
# - worker_ready 每个 worker 都会尝试抢启动恢复锁,只有抢到锁的 worker 投递恢复任务。
# - 启动容灾保留,但恢复扫描独立投递到 CELERY_RECOVERY_QUEUE
# - worker_ready 每个 worker 都会尝试抢启动恢复锁,只有抢到锁的 worker 投递恢复协调任务。
# - poll active 使用独立 Redis key,避免影响稳定的下载 active 注册表。
GENERATION_RECOVERY_BATCH_SIZE: int = 100
GENERATION_RECOVERY_MAX_ROUNDS: int = 5
POLL_RECOVERY_BATCH_SIZE: int = 100
GENERATION_RECOVERY_BATCH_SIZE: int = 20
GENERATION_RECOVERY_MAX_ROUNDS: int = 1
POLL_RECOVERY_BATCH_SIZE: int = 20
POLL_TASK_LEASE_SECONDS: int = 5 * 60
POLL_TASK_QUEUE_TIMEOUT_SECONDS: int = 2 * 60
POLL_ACTIVE_REDIS_HASH_KEY: str = "vg:celery:poll:active"
POLL_ACTIVE_REDIS_ZSET_KEY: str = "vg:celery:poll:active_index"
CELERY_STARTUP_RECOVERY_LOCK_KEY: str = "vg:celery:startup_recovery_lock"
CELERY_STARTUP_RECOVERY_LOCK_TTL_SECONDS: int = 120
CELERY_RECOVERY_QUEUE: str = "gen_recovery"
CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS: int = 10 * 60
CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS: int = 300
CELERY_RECOVERY_TIME_LIMIT_SECONDS: int = 420
CELERY_RECOVERY_STARTUP_TASK_LOCK_KEY: str = "vg:celery:startup_recovery_task_lock"
GENERATION_RECOVERY_LOCK_KEY: str = "vg:celery:generation_recovery_lock"
DOWNLOAD_RECOVERY_LOCK_KEY: str = "vg:celery:download_recovery_lock"
MODULE_ASYNC_RECOVERY_LOCK_KEY: str = "vg:celery:module_async_recovery_lock"
SHOT_SPLIT_RECOVERY_LOCK_KEY: str = "vg:celery:shot_split_recovery_lock"
# 模块异步任务容灾配置。
# 覆盖 ModuleGenerationStep 提词任务、shot 原视频/片段分析、shot ffmpeg 切割 active 注册。
# 不新增 worker 队列:恢复扫描gen_result_download,真实业务任务回到原始队列。
MODULE_ASYNC_RECOVERY_BATCH_SIZE: int = 100
# 恢复扫描走 CELERY_RECOVERY_QUEUE,真实业务任务回到原始队列。
MODULE_ASYNC_RECOVERY_BATCH_SIZE: int = 20
MODULE_ASYNC_QUEUE_TIMEOUT_SECONDS: int = 5 * 60
MODULE_ASYNC_LEASE_SECONDS: int = 10 * 60
MODULE_ASYNC_LOCK_TTL_SECONDS: int = 10 * 60
@@ -192,7 +210,7 @@ class Settings(BaseSettings):
SHOT_SPLIT_RETRY_BACKOFF_SECONDS: int = 30
SHOT_SPLIT_LEASE_SECONDS: int = 10 * 60
SHOT_SPLIT_PENDING_TIMEOUT_SECONDS: int = 5 * 60
SHOT_SPLIT_RECOVERY_BATCH_SIZE: int = 50
SHOT_SPLIT_RECOVERY_BATCH_SIZE: int = 20
SHOT_SPLIT_LOCK_KEY_PREFIX: str = "vg:shot_replicate:split:lock"
SHOT_SPLIT_SEMAPHORE_KEY_PREFIX: str = "vg:shot_replicate:split:semaphore"
+1
View File
@@ -6,3 +6,4 @@ from app.enums.video_prompt_schema import *
from app.enums.user import *
from app.enums.credit_record import *
from app.enums.token_usage import *
from app.enums.generation_task import *
+101
View File
@@ -0,0 +1,101 @@
from enum import Enum
class GenerationMode(str, Enum):
CHATAPI_ASYNC = "chatapi_async"
HOT_OPENING_REPLICATE = "hot_opening_replicate"
SHOT_REPLICATE = "shot_replicate"
class GenerationType(str, Enum):
IMAGE = "image"
VIDEO = "video"
class ChatGenerationTaskStatus(str, Enum):
PENDING = "pending"
GENERATING = "generating"
COMPLETED = "completed"
FAILED = "failed"
class ChatGenerationPipelineStage(str, Enum):
QUEUED = "queued"
PREPARING = "preparing"
CREATING_PROVIDER_TASK = "creating_provider_task"
WAITING_REMOTE = "waiting_remote"
POLLING = "polling"
RESULT_READY = "result_ready"
DOWNLOAD_QUEUED = "download_queued"
DOWNLOADING = "downloading"
RETRY_WAITING = "retry_waiting"
DONE = "done"
FAILED = "failed"
TIMEOUT = "timeout"
DOWNLOAD_FAILED = "download_failed"
class ChatGenerationTaskEventType(str, Enum):
PROMPT_CONCAT_START = "PROMPT_CONCAT_START"
PROMPT_CONCAT_SUCCESS = "PROMPT_CONCAT_SUCCESS"
PROVIDER_CREATE_START = "PROVIDER_CREATE_START"
PROVIDER_CREATE_SUCCESS = "PROVIDER_CREATE_SUCCESS"
PROVIDER_CREATE_FAILED = "PROVIDER_CREATE_FAILED"
POLL_START = "POLL_START"
POLL_PENDING = "POLL_PENDING"
POLL_SUCCESS = "POLL_SUCCESS"
POLL_FAILED = "POLL_FAILED"
POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY = "POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY"
FINAL_POLL_BEFORE_TIMEOUT_ERROR = "FINAL_POLL_BEFORE_TIMEOUT_ERROR"
FINAL_POLL_BEFORE_TIMEOUT_PENDING = "FINAL_POLL_BEFORE_TIMEOUT_PENDING"
GENERATION_RECOVERY_ENQUEUE = "GENERATION_RECOVERY_ENQUEUE"
DOWNLOAD_ENQUEUE = "DOWNLOAD_ENQUEUE"
DOWNLOAD_ENQUEUE_FAILED = "DOWNLOAD_ENQUEUE_FAILED"
DOWNLOAD_START = "DOWNLOAD_START"
DOWNLOAD_SUCCESS = "DOWNLOAD_SUCCESS"
DOWNLOAD_RETRY_WAITING = "DOWNLOAD_RETRY_WAITING"
DOWNLOAD_RETRY_ENQUEUE = "DOWNLOAD_RETRY_ENQUEUE"
DOWNLOAD_RETRY_ENQUEUE_FAILED = "DOWNLOAD_RETRY_ENQUEUE_FAILED"
DOWNLOAD_RETRY_NOT_DUE = "DOWNLOAD_RETRY_NOT_DUE"
DOWNLOAD_STUCK_RECOVER = "DOWNLOAD_STUCK_RECOVER"
DOWNLOAD_RECOVERY_ENQUEUE = "DOWNLOAD_RECOVERY_ENQUEUE"
DOWNLOAD_RECOVERY_ENQUEUE_FAILED = "DOWNLOAD_RECOVERY_ENQUEUE_FAILED"
DOWNLOAD_FAILED = "DOWNLOAD_FAILED"
DOWNLOAD_FAILED_NON_RETRYABLE = "DOWNLOAD_FAILED_NON_RETRYABLE"
DOWNLOAD_SKIP_TASK_MISSING = "DOWNLOAD_SKIP_TASK_MISSING"
DOWNLOAD_SKIP_INVALID_MODE = "DOWNLOAD_SKIP_INVALID_MODE"
DOWNLOAD_SKIP_NOT_GENERATING = "DOWNLOAD_SKIP_NOT_GENERATING"
DOWNLOAD_SKIP_ALREADY_COMPLETED = "DOWNLOAD_SKIP_ALREADY_COMPLETED"
DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL = "DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL"
DOWNLOAD_SKIP_STAGE_NOT_ALLOWED = "DOWNLOAD_SKIP_STAGE_NOT_ALLOWED"
DOWNLOAD_SKIP_DOWNLOADING_LEASE_ALIVE = "DOWNLOAD_SKIP_DOWNLOADING_LEASE_ALIVE"
DOWNLOAD_SKIP_RETRY_NOT_DUE = "DOWNLOAD_SKIP_RETRY_NOT_DUE"
DOWNLOAD_SKIP_FINAL_STATE = "DOWNLOAD_SKIP_FINAL_STATE"
DOWNLOAD_SKIP_DISABLED = "DOWNLOAD_SKIP_DISABLED"
TASK_TIMEOUT = "TASK_TIMEOUT"
ALLOWED_GENERATION_MODES = {
GenerationMode.CHATAPI_ASYNC.value,
GenerationMode.HOT_OPENING_REPLICATE.value,
GenerationMode.SHOT_REPLICATE.value,
}
FINAL_CHAT_GENERATION_STAGES = {
ChatGenerationPipelineStage.DONE.value,
ChatGenerationPipelineStage.FAILED.value,
ChatGenerationPipelineStage.TIMEOUT.value,
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value,
}
DOWNLOAD_RECOVERABLE_STAGES = {
ChatGenerationPipelineStage.RESULT_READY.value,
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
ChatGenerationPipelineStage.DOWNLOADING.value,
ChatGenerationPipelineStage.RETRY_WAITING.value,
}
@@ -2,36 +2,40 @@ from __future__ import annotations
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.generation_task import ChatGenerationTaskStatus, GenerationMode
from app.models.chat_generation_task import ChatGenerationTask
async def notify_chat_generation_task_finished(db: AsyncSession, task: ChatGenerationTask) -> None:
"""通知业务模块 ChatGenerationTask 已进入终态。
当前用于爆款开头复刻:
- image_generate 完成后自动进入 video_prompt_optimize
- video_generate 完成后总任务完成
该方法必须幂等:下载恢复任务、重试任务、服务重启补偿都可能重复调用。
具体模块服务需要自行判断 step/project 是否已经完成或失败,避免重复推进。
"""
if not task:
return
if task.generation_mode == "hot_opening_replicate":
status = getattr(task, "status", None)
generation_mode = getattr(task, "generation_mode", None)
if generation_mode == GenerationMode.HOT_OPENING_REPLICATE.value:
from app.services.hot_opening_replicate_service import (
handle_chat_generation_task_completed,
handle_chat_generation_task_failed,
)
if task.status == "completed":
if status == ChatGenerationTaskStatus.COMPLETED.value:
await handle_chat_generation_task_completed(db, task)
elif task.status == "failed":
elif status == ChatGenerationTaskStatus.FAILED.value:
await handle_chat_generation_task_failed(db, task)
return
if task.generation_mode == "shot_replicate":
if generation_mode == GenerationMode.SHOT_REPLICATE.value:
from app.services.shot_replicate_flow_service import (
handle_chat_generation_task_completed,
handle_chat_generation_task_failed,
)
if task.status == "completed":
if status == ChatGenerationTaskStatus.COMPLETED.value:
await handle_chat_generation_task_completed(db, task)
elif task.status == "failed":
elif status == ChatGenerationTaskStatus.FAILED.value:
await handle_chat_generation_task_failed(db, task)
return
@@ -9,6 +9,12 @@ from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.enums.generation_task import (
ALLOWED_GENERATION_MODES,
ChatGenerationPipelineStage,
ChatGenerationTaskEventType,
ChatGenerationTaskStatus,
)
from app.models.chat_generation_task import ChatGenerationTask
from app.services.celery_download_recovery_service import (
ensure_aware_utc,
@@ -17,9 +23,8 @@ from app.services.celery_download_recovery_service import (
postpone_download_active_check,
remove_download_active,
)
from app.services.generation_log_service import log_provider_call, log_task_event
from app.services.generation_log_service import log_task_event
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
from app.services.generation_provider_service import poll_provider_task
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.redis_registry_service import (
redis_get_due_registry_ids,
@@ -30,7 +35,6 @@ from app.services.redis_registry_service import (
logger = logging.getLogger("video_gen")
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"}
POLL_QUEUE = "gen_provider_poll"
@@ -59,11 +63,11 @@ def _is_queue_timeout(task: ChatGenerationTask, now: datetime | None = None) ->
def _is_final_task_state(task: ChatGenerationTask) -> bool:
return task.status in ("completed", "failed") or task.pipeline_stage in (
"done",
"failed",
"timeout",
"download_failed",
return task.status in (ChatGenerationTaskStatus.COMPLETED.value, ChatGenerationTaskStatus.FAILED.value) or task.pipeline_stage in (
ChatGenerationPipelineStage.DONE.value,
ChatGenerationPipelineStage.FAILED.value,
ChatGenerationPipelineStage.TIMEOUT.value,
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value,
)
@@ -137,19 +141,31 @@ async def recover_one_download_task(
if _is_final_task_state(task):
await remove_download_active(task.id)
return "clean_final_state"
if task.status != "generating":
if task.status != ChatGenerationTaskStatus.GENERATING.value:
await remove_download_active(task.id)
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING.value,
message=f"{source} 下载恢复跳过:任务不是 generating",
detail={"status": task.status, "stage": task.pipeline_stage},
)
return "clean_not_generating"
if not task.remote_result_url:
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL.value,
message=f"{source} 下载恢复跳过:缺少 remote_result_url",
detail={"status": task.status, "stage": task.pipeline_stage},
)
return "skip_no_remote_result_url"
stage = task.pipeline_stage
redis_payload = payload or {}
if stage == "result_ready":
if stage == ChatGenerationPipelineStage.RESULT_READY.value:
await log_task_event(
task,
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 result_ready 未完成下载,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
@@ -165,7 +181,7 @@ async def recover_one_download_task(
if _is_queue_timeout(task, current_time):
await log_task_event(
task,
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 download_queued 长时间未消费,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
@@ -188,7 +204,7 @@ async def recover_one_download_task(
if _is_expired(task.download_lease_until, current_time):
await log_task_event(
task,
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 downloading lease 过期,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
@@ -211,7 +227,7 @@ async def recover_one_download_task(
if _is_expired(task.download_next_retry_at, current_time):
await log_task_event(
task,
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 retry_waiting 到期,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
@@ -276,11 +292,16 @@ async def recover_download_tasks_once(db: AsyncSession) -> dict[str, Any]:
select(ChatGenerationTask)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate", "shot_replicate"]),
ChatGenerationTask.status == "generating",
ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)),
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
ChatGenerationTask.remote_result_url.is_not(None),
ChatGenerationTask.pipeline_stage.in_(
["result_ready", "download_queued", "downloading", "retry_waiting"]
[
ChatGenerationPipelineStage.RESULT_READY.value,
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
ChatGenerationPipelineStage.DOWNLOADING.value,
ChatGenerationPipelineStage.RETRY_WAITING.value,
]
),
)
.order_by(ChatGenerationTask.updated_at.asc())
@@ -314,7 +335,7 @@ async def _mark_timeout(
db,
task=task,
error_message=error_message,
pipeline_stage="timeout",
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
await notify_chat_generation_task_finished(db, task)
await db.commit()
@@ -323,7 +344,7 @@ async def _mark_timeout(
task,
event_type="TASK_TIMEOUT",
to_status="failed",
to_stage="timeout",
to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
return "mark_timeout"
@@ -340,7 +361,7 @@ async def _mark_failed(
db,
task=task,
error_message=error_message,
pipeline_stage="failed",
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
)
await notify_chat_generation_task_finished(db, task)
await db.commit()
@@ -349,93 +370,6 @@ async def _mark_failed(
return "mark_failed"
async def _try_final_poll_before_timeout(db: AsyncSession, task: ChatGenerationTask) -> str:
"""超时前最后查一次供应商,避免 Celery 中断导致本地假超时。
如果供应商已经成功,继续进入下载;如果仍 running 或查询失败,再按超时处理。
"""
from app.tasks.generation_download_tasks import enqueue_download_task
if not (task.provider_task_id or task.seedance_task_id):
return await _mark_timeout(db, task)
try:
poll_result = await poll_provider_task(db, task)
status = poll_result.get("status")
response_data = poll_result.get("response_data")
except Exception as exc:
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT_ERROR",
message=str(exc),
)
return await _mark_timeout(db, task)
try:
provider_response = json.loads(response_data or "{}")
except Exception:
provider_response = {"raw": response_data}
snapshot = _engine_snapshot(task)
await log_provider_call(
task,
provider=snapshot.get("provider") or "ark",
api_type=f"{task.gen_type}_final_poll_before_timeout",
model=snapshot.get("model_name"),
engine_id=task.engine_id,
status="success",
provider_task_id=task.seedance_task_id or task.provider_task_id,
response_data=provider_response,
)
if _is_success(status):
if task.gen_type == "image":
task.remote_result_url = poll_result.get("image_url")
task.image_tokens_used = poll_result.get("image_tokens", 0) or 0
else:
task.remote_result_url = poll_result.get("video_url")
task.video_tokens_used = poll_result.get("video_tokens", 0) or 0
task.provider_response_json = response_data
if not task.remote_result_url:
return await _mark_failed(
db,
task,
error_message="供应商任务成功但未返回结果URL",
detail=poll_result,
)
task.pipeline_stage = "result_ready"
task.retry_count = 0
await db.commit()
await _remove_poll_active(task.id)
await log_task_event(
task,
event_type="POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY",
to_stage="result_ready",
detail=poll_result,
)
await enqueue_download_task(db, task, recover=True, reason="final_poll_before_timeout_success")
return "recover_timeout_success_to_download"
if _is_failed(status):
task.provider_response_json = response_data
return await _mark_failed(
db,
task,
error_message=poll_result.get("error") or f"供应商任务失败: {status}",
detail=poll_result,
)
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT_PENDING",
message=f"status={status}",
detail=poll_result,
)
return await _mark_timeout(db, task)
async def recover_one_generation_task(
db: AsyncSession,
task: ChatGenerationTask,
@@ -443,6 +377,14 @@ async def recover_one_generation_task(
payload: dict[str, Any] | None = None,
source: str = "startup_db",
) -> str:
"""恢复单个生成任务。
分流原则:
1. 已有 remote_result_url:只恢复下载,不 poll,不重新 create。
2. 已有 provider_task_id/seedance_task_id:恢复 poll。
3. 无结果 URL、无供应商任务 IDdeadline 未过才恢复 create。
4. 无结果 URL、无供应商任务 ID:deadline 已过直接超时失败,不再补救生成。
"""
from app.tasks.generation_create_tasks import chatapi_create_generation_task
from app.tasks.generation_download_tasks import enqueue_download_task
from app.tasks.generation_poll_tasks import poll_generation_task, register_poll_active
@@ -458,37 +400,107 @@ async def recover_one_generation_task(
if _is_final_task_state(task):
await _remove_poll_active(task.id)
return "clean_final_state"
if task.status != "generating":
if task.status != ChatGenerationTaskStatus.GENERATING.value:
await _remove_poll_active(task.id)
return "clean_not_generating"
if task.deadline_at and _is_expired(task.deadline_at, current_time):
if task.pipeline_stage in ("waiting_remote", "polling"):
return await _try_final_poll_before_timeout(db, task)
return await _mark_timeout(db, task)
has_remote_result = bool(str(task.remote_result_url or "").strip())
has_provider_task_id = bool(str(task.provider_task_id or "").strip() or str(task.seedance_task_id or "").strip())
is_deadline_expired = bool(task.deadline_at and _is_expired(task.deadline_at, current_time))
if task.pipeline_stage in ("queued", "preparing", "creating_provider_task"):
if task.provider_task_id or task.seedance_task_id:
task.pipeline_stage = "waiting_remote"
# 最高优先级:只要远程结果 URL 已经落库,说明生成侧已经成功。
# 不管当前 pipeline_stage 是 queued/creating/waiting/result_ready/download_*,恢复时都不能重复 create 或 poll。
if has_remote_result:
await _remove_poll_active(task.id)
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message=f"{source} 发现任务已存在 remote_result_url,恢复投递下载队列",
detail={
"pipeline_stage": task.pipeline_stage,
"payload": redis_payload,
"deadline_expired": is_deadline_expired,
},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_has_remote_result_url",
)
return "recover_download_has_remote_result"
# 已经过 deadline 且没有结果 URL
# - 有供应商任务 ID:交给 poll worker 做最后一次状态确认;
# - 没有供应商任务 ID:说明没有可查询的远程任务,直接按超时失败处理,不再重新 create。
if is_deadline_expired:
if has_provider_task_id:
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
await db.commit()
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message=f"{source} 发现创建阶段已存在供应商任务ID,恢复投递轮询队列",
message=f"{source} 发现任务已到 deadline 且存在供应商任务ID,投递 poll 队列做最终查询",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
poll_generation_task.apply_async(args=[task.id], queue=POLL_QUEUE, countdown=0)
await register_poll_active(
task,
check_at=_poll_queue_timeout_at(),
reason=f"{source}_create_stage_has_provider_id",
reason=f"{source}_deadline_final_poll",
)
return "recover_poll_from_create_stage"
return "recover_deadline_final_poll"
await log_task_event(
task,
event_type="GENERATION_RECOVERY_TIMEOUT",
message=f"{source} 发现任务已到 deadline,且没有 remote_result_url/供应商任务ID,按超时失败处理",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
return await _mark_timeout(db, task)
# 未过 deadline:有供应商任务 ID 才允许恢复到 poll 队列。
if has_provider_task_id:
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
await db.commit()
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message=f"{source} 发现创建阶段任务未完成,恢复投递创建队列",
message=f"{source} 发现任务存在供应商任务ID,恢复投递轮询队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
poll_generation_task.apply_async(args=[task.id], queue=POLL_QUEUE, countdown=0)
await register_poll_active(
task,
check_at=_poll_queue_timeout_at(),
reason=f"{source}_has_provider_task_id",
)
return "recover_poll_has_provider_id"
# 未过 deadline,且没有结果 URL / 供应商任务 ID:
# 图片同步任务会重新进入 submit_image_task;视频/其它任务会重新创建供应商任务。
# 这里不能投 poll,因为没有 provider_task_id/seedance_task_id 可查询。
recoverable_create_stages = {
ChatGenerationPipelineStage.QUEUED.value,
ChatGenerationPipelineStage.PREPARING.value,
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
ChatGenerationPipelineStage.WAITING_REMOTE.value,
ChatGenerationPipelineStage.POLLING.value,
}
if task.pipeline_stage in recoverable_create_stages:
if task.pipeline_stage not in (
ChatGenerationPipelineStage.QUEUED.value,
ChatGenerationPipelineStage.PREPARING.value,
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
):
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
await db.commit()
await _remove_poll_active(task.id)
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message=f"{source} 发现任务未超时且缺少 remote_result_url/供应商任务ID,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
chatapi_create_generation_task.apply_async(
@@ -496,67 +508,25 @@ async def recover_one_generation_task(
queue="gen_chatapi_create",
countdown=0,
)
return "recover_create"
return "recover_create_no_remote_no_provider_before_deadline"
if task.pipeline_stage in ("waiting_remote", "polling"):
if task.remote_result_url:
await _remove_poll_active(task.id)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_waiting_remote_has_result",
)
return "recover_waiting_has_result"
if task.provider_task_id or task.seedance_task_id:
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message=f"{source} 发现远程等待/轮询阶段任务未完成,恢复投递轮询队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
task.pipeline_stage = "waiting_remote"
await db.commit()
poll_generation_task.apply_async(
args=[task.id],
queue=POLL_QUEUE,
countdown=0,
)
await register_poll_active(
task,
check_at=_poll_queue_timeout_at(),
reason=f"{source}_recover_poll",
)
return "recover_poll"
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message=f"{source} 发现任务缺少供应商任务ID,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
task.pipeline_stage = "queued"
# result_ready 但没有 URL 是脏状态;未过 deadline 时回创建队列重新处理,过期上面已标记超时。
if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value:
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
await db.commit()
await _remove_poll_active(task.id)
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message=f"{source} 发现 result_ready 但缺少 remote_result_url,未超时,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
chatapi_create_generation_task.apply_async(
args=[task.id],
queue="gen_chatapi_create",
countdown=0,
)
return "recover_create_missing_provider_id"
if task.pipeline_stage == "result_ready":
await _remove_poll_active(task.id)
if task.remote_result_url:
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_generation_result_ready",
)
return "recover_result_ready"
return "skip_result_ready_no_url"
return "recover_create_result_ready_no_url_before_deadline"
return f"skip_stage_{task.pipeline_stage}"
@@ -617,8 +587,8 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
select(ChatGenerationTask)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate", "shot_replicate"]),
ChatGenerationTask.status == "generating",
ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)),
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
ChatGenerationTask.pipeline_stage.in_(
[
"queued",
@@ -656,11 +626,8 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
if len(tasks) < batch_size or progressed_this_round <= 0:
break
# 下载阶段单独跑 DB fallback。
download_result = await recover_download_tasks_once(db)
return {
"checked": len(checked_ids),
"db_checked": total_db_checked,
"results": results,
"download_recovery": download_result,
}
}
@@ -0,0 +1,253 @@
from __future__ import annotations
import json
from typing import Any, Mapping
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.enums.credit_record import (
CreditRecordAction,
CreditRecordChargeKind,
CreditRecordOwnerType,
CreditRecordSourceModule,
)
from app.models.chat_generation_task import ChatGenerationTask
from app.models.credit_record import CreditRecord
from app.models.generation_record import GenerationRecord
from app.models.token_usage import TokenUsage
from app.utils.id_gen import generate_id
def _safe_int(value: Any, default: int = 0) -> int:
try:
if value is None or value == "":
return default
return int(value)
except Exception:
return default
def _safe_json_dict(value: Any) -> dict[str, Any]:
if not value:
return {}
if isinstance(value, dict):
return value
try:
parsed = json.loads(value)
return parsed if isinstance(parsed, dict) else {}
except Exception:
return {}
def _extract_usage(provider_response: Any) -> dict[str, Any]:
data = _safe_json_dict(provider_response)
usage = data.get("usage")
return usage if isinstance(usage, dict) else {}
def _normalize_media_tokens(
*,
gen_type: str | None,
provider_response: Any = None,
fallback_total: int | None = None,
) -> tuple[int, int, int]:
usage = _extract_usage(provider_response)
input_tokens = _safe_int(usage.get("input_tokens"), 0)
output_tokens = _safe_int(
usage.get("output_tokens"),
_safe_int(usage.get("generated_tokens"), 0),
)
total_tokens = _safe_int(usage.get("total_tokens"), 0)
if total_tokens <= 0:
total_tokens = _safe_int(fallback_total, 0)
if output_tokens <= 0:
output_tokens = max(0, total_tokens - input_tokens)
if total_tokens <= 0:
total_tokens = input_tokens + output_tokens
# 图片生成多数供应商只返回 output/total,没有 input;保持 input=0。视频同理兼容缺字段。
return input_tokens, output_tokens, total_tokens
def _engine_model_from_provider_response(provider_response: Any) -> str | None:
data = _safe_json_dict(provider_response)
model = data.get("model")
return str(model) if model else None
async def _find_latest_media_charge(
db: AsyncSession,
*,
user_id: str,
owner_type: str,
owner_id: str,
media_type: str | None,
) -> CreditRecord | None:
query = (
select(CreditRecord)
.where(CreditRecord.user_id == user_id)
.where(CreditRecord.type == "consume")
.where(CreditRecord.owner_type == owner_type)
.where(CreditRecord.owner_id == owner_id)
.where(CreditRecord.charge_kind == CreditRecordChargeKind.MEDIA.value)
.where(CreditRecord.charge_action == CreditRecordAction.CHARGE.value)
)
if media_type:
query = query.where(CreditRecord.media_type == media_type)
query = query.order_by(CreditRecord.attempt_no.desc().nullslast(), CreditRecord.created_at.desc()).limit(1)
result = await db.execute(query)
return result.scalar_one_or_none()
async def _get_or_create_token_usage(
db: AsyncSession,
*,
charge: CreditRecord,
input_tokens: int,
output_tokens: int,
total_tokens: int,
model_config_id: str | None = None,
) -> TokenUsage:
token_usage: TokenUsage | None = None
if charge.token_usage_id:
result = await db.execute(select(TokenUsage).where(TokenUsage.id == charge.token_usage_id).limit(1))
token_usage = result.scalar_one_or_none()
if token_usage is None and charge.biz_key:
result = await db.execute(select(TokenUsage).where(TokenUsage.biz_key == charge.biz_key).limit(1))
token_usage = result.scalar_one_or_none()
if token_usage is None:
token_usage = TokenUsage(
id=generate_id(),
user_id=charge.user_id,
model_config_id=model_config_id,
owner_type=charge.owner_type,
owner_id=charge.owner_id,
biz_key=charge.biz_key,
source_module=charge.source_module,
source_step_code=charge.source_step_code,
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
)
db.add(token_usage)
await db.flush()
else:
token_usage.user_id = token_usage.user_id or charge.user_id
token_usage.model_config_id = token_usage.model_config_id or model_config_id
token_usage.owner_type = token_usage.owner_type or charge.owner_type
token_usage.owner_id = token_usage.owner_id or charge.owner_id
token_usage.biz_key = token_usage.biz_key or charge.biz_key
token_usage.source_module = token_usage.source_module or charge.source_module
token_usage.source_step_code = token_usage.source_step_code or charge.source_step_code
token_usage.input_tokens = input_tokens
token_usage.output_tokens = output_tokens
token_usage.total_tokens = total_tokens
return token_usage
async def _sync_charge_snapshot(
db: AsyncSession,
*,
charge: CreditRecord | None,
gen_type: str | None,
provider_response: Any = None,
fallback_total: int | None = None,
) -> CreditRecord | None:
if not charge:
return None
input_tokens, output_tokens, total_tokens = _normalize_media_tokens(
gen_type=gen_type,
provider_response=provider_response,
fallback_total=fallback_total,
)
if total_tokens <= 0:
return charge
token_usage = await _get_or_create_token_usage(
db,
charge=charge,
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
model_config_id=None,
)
charge.token_usage_id = token_usage.id
charge.input_tokens = input_tokens
charge.output_tokens = output_tokens
charge.total_tokens = total_tokens
# 兼容旧流水扣费时未冷备 engine_model_name 的场景,能从 provider response 推出来就补充。
provider_model = _engine_model_from_provider_response(provider_response)
if provider_model and not charge.engine_model_name:
charge.engine_model_name = provider_model
return charge
async def sync_chat_generation_task_media_token_snapshot(
db: AsyncSession,
task: ChatGenerationTask,
*,
provider_response: Any = None,
) -> CreditRecord | None:
"""把 ChatGenerationTask 图片/视频媒体生成 token 后置快照回填到积分流水。
媒体扣费发生在创建任务前,供应商 usage 只能在创建/轮询成功后拿到,
所以这里按 owner_type + owner_id + media_type 找到对应 media charge 流水并回填。
"""
if not bool(getattr(settings, "MEDIA_TOKEN_SNAPSHOT_ENABLED", True)):
return None
if not task:
return None
gen_type = (getattr(task, "gen_type", None) or "").lower().strip()
fallback_total = task.image_tokens_used if gen_type == "image" else task.video_tokens_used
response = provider_response if provider_response is not None else getattr(task, "provider_response_json", None)
charge = await _find_latest_media_charge(
db,
user_id=task.user_id,
owner_type=CreditRecordOwnerType.CHAT_GENERATION_TASK.value,
owner_id=task.id,
media_type=gen_type or None,
)
return await _sync_charge_snapshot(
db,
charge=charge,
gen_type=gen_type,
provider_response=response,
fallback_total=fallback_total,
)
async def sync_generation_record_media_token_snapshot(
db: AsyncSession,
record: GenerationRecord,
*,
provider_response: Any = None,
) -> CreditRecord | None:
"""把旧 GenerationRecord 图片/视频媒体生成 token 后置快照回填到积分流水。"""
if not bool(getattr(settings, "MEDIA_TOKEN_SNAPSHOT_ENABLED", True)):
return None
if not record:
return None
gen_type = (getattr(record, "gen_type", None) or "").lower().strip()
fallback_total = record.image_tokens_used if gen_type == "image" else record.video_tokens_used
charge = await _find_latest_media_charge(
db,
user_id=record.user_id,
owner_type=CreditRecordOwnerType.GENERATION_RECORD.value,
owner_id=record.id,
media_type=gen_type or None,
)
return await _sync_charge_snapshot(
db,
charge=charge,
gen_type=gen_type,
provider_response=provider_response,
fallback_total=fallback_total,
)
@@ -10,6 +10,7 @@ from app.models.base import async_session
from app.models.generation_record import GenerationRecord
from app.services.video_gen import get_active_engine, poll_task_status, download_video, _log_video_response
from app.services.image_gen import get_active_image_engine, download_image
from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot
from app.services.resource_accounting_service import (
record_generation_record_generated_resource,
safe_file_size,
@@ -157,6 +158,7 @@ class TaskQueue:
else:
record.video_url = file_url
record.video_tokens_used = poll_result.get("video_tokens", 0)
await sync_generation_record_media_token_snapshot(db, record, provider_response=resp_data)
record.status = "completed"
record.generated_at = datetime.now()
if record.video_url:
@@ -235,6 +237,7 @@ class TaskQueue:
else:
record.image_url = remote_url
record.image_tokens_used = poll_result.get("image_tokens", 0)
await sync_generation_record_media_token_snapshot(db, record, provider_response=poll_result)
record.status = "completed"
record.generated_at = datetime.now()
if record.image_url:
+94 -16
View File
@@ -1,13 +1,16 @@
from __future__ import annotations
import asyncio
import logging
import os
import threading
from concurrent.futures import Future
from concurrent.futures import Future, TimeoutError as FutureTimeoutError
from typing import Awaitable, TypeVar
from app.config import settings
logger = logging.getLogger("video_gen")
T = TypeVar("T")
_thread_local = threading.local()
@@ -18,6 +21,28 @@ _single_loop_pid: int | None = None
_single_loop_ready: threading.Event | None = None
async def _dispose_async_resources() -> None:
"""释放当前 async loop 内缓存的异步资源。
Celery soft time limit 会打断同步等待 future.result() 的线程;如果不主动
cancel coroutine 并释放 engine/redis,后台 loop 里残留的协程可能继续占用
SQLAlchemy QueuePool 连接,后续任务就会出现 QueuePool timeout。
"""
try:
from app.services.redis_registry_service import close_registry_redis
await close_registry_redis()
except Exception:
logger.debug("关闭 Celery Redis registry 连接失败", exc_info=True)
try:
from app.models.base import engine
await engine.dispose()
except Exception:
logger.debug("dispose Celery SQLAlchemy engine 失败", exc_info=True)
def _runner_mode() -> str:
mode = str(getattr(settings, "CELERY_ASYNC_RUNNER_MODE", "single_loop") or "single_loop").strip().lower()
if mode not in {"single_loop", "direct"}:
@@ -46,16 +71,18 @@ def _get_or_create_thread_local_loop() -> asyncio.AbstractEventLoop:
def _single_loop_worker(loop: asyncio.AbstractEventLoop, ready: threading.Event) -> None:
asyncio.set_event_loop(loop)
ready.set()
loop.run_forever()
try:
loop.run_forever()
finally:
pending = [task for task in asyncio.all_tasks(loop) if not task.done()]
if pending:
for task in pending:
task.cancel()
loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True))
pending = [task for task in asyncio.all_tasks(loop) if not task.done()]
if pending:
for task in pending:
task.cancel()
loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True))
loop.run_until_complete(loop.shutdown_asyncgens())
loop.close()
loop.run_until_complete(_dispose_async_resources())
loop.run_until_complete(loop.shutdown_asyncgens())
loop.close()
def _get_or_create_single_loop() -> asyncio.AbstractEventLoop:
@@ -92,13 +119,28 @@ def _get_or_create_single_loop() -> asyncio.AbstractEventLoop:
return _single_loop
def _cancel_future_and_reset_loop(future: Future[T] | None, *, reason: str) -> None:
"""取消当前协程并重置当前进程内 event loop。"""
if future is not None and not future.done():
future.cancel()
try:
future.result(timeout=2)
except Exception:
pass
logger.warning("Celery async_runner 正在重置 event loop。reason=%s", reason)
close_loop()
def run_async(coro: Awaitable[T]) -> T:
"""Celery 同步 task 调用异步协程的统一入口。
默认 single_loop 模式:
- 一个 Celery 子进程只有一个专用 event loop;
- 所有 asyncpg / redis.asyncio 操作都在这个 loop 内创建和使用;
- 避免 got Future attached to a different loop
- 避免 got Future attached to a different loop
- 当 Celery soft time limit 打断 future.result() 时,主动 cancel 后台协程并
释放连接池,避免 QueuePool 被残留任务长期占用。
降级 direct 模式:
- 兼容旧的线程本地 loop 方案;
@@ -106,7 +148,15 @@ def run_async(coro: Awaitable[T]) -> T:
"""
if _runner_mode() == "direct":
loop = _get_or_create_thread_local_loop()
return loop.run_until_complete(coro)
try:
return loop.run_until_complete(coro)
except BaseException:
try:
if not loop.is_closed():
loop.run_until_complete(_dispose_async_resources())
finally:
close_loop()
raise
loop = _get_or_create_single_loop()
try:
@@ -118,7 +168,16 @@ def run_async(coro: Awaitable[T]) -> T:
raise RuntimeError("run_async() 不能在 Celery async_runner 的事件循环内部被同步调用")
future: Future[T] = asyncio.run_coroutine_threadsafe(coro, loop)
return future.result()
try:
return future.result()
except FutureTimeoutError:
_cancel_future_and_reset_loop(future, reason="future_result_timeout")
raise
except BaseException:
# Celery SoftTimeLimitExceeded/worker shutdown 等异常会从这里抛出。
# 必须重置 loop,否则后台协程继续运行会拖住 DB 连接池。
_cancel_future_and_reset_loop(future, reason="base_exception")
raise
def close_loop() -> None:
@@ -130,8 +189,16 @@ def close_loop() -> None:
loop = _single_loop
thread = _single_loop_thread
if loop is not None and not loop.is_closed() and thread is not None and thread.is_alive():
loop.call_soon_threadsafe(loop.stop)
thread.join(timeout=5)
try:
cleanup_future = asyncio.run_coroutine_threadsafe(_dispose_async_resources(), loop)
cleanup_future.result(timeout=5)
except Exception:
logger.debug("关闭 loop 前清理 async 资源失败", exc_info=True)
try:
loop.call_soon_threadsafe(loop.stop)
thread.join(timeout=5)
except Exception:
logger.debug("关闭 Celery async_runner loop 失败", exc_info=True)
_single_loop = None
_single_loop_thread = None
@@ -141,6 +208,17 @@ def close_loop() -> None:
# 关闭 direct 降级模式的线程本地 loop。
loop = getattr(_thread_local, "loop", None)
if loop is not None and not loop.is_closed():
loop.close()
try:
pending = [task for task in asyncio.all_tasks(loop) if not task.done()]
for task in pending:
task.cancel()
if pending:
loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True))
loop.run_until_complete(_dispose_async_resources())
loop.run_until_complete(loop.shutdown_asyncgens())
except Exception:
logger.debug("关闭 direct loop 前清理失败", exc_info=True)
finally:
loop.close()
_thread_local.loop = None
_thread_local.pid = None
+35 -34
View File
@@ -26,6 +26,9 @@ CELERY_TASK_IMPORTS = (
)
RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or "gen_recovery"
def _derive_redis_db(url: str, db_no: int) -> str:
if not url:
return url
@@ -55,6 +58,20 @@ if broker_url:
task_acks_late=True,
task_reject_on_worker_lost=True,
task_track_started=True,
task_annotations={
# 生成链路任务以数据库状态为准,不依赖 Celery result backend。
# 这里忽略结果可避免任务误返回 ORM / 非 JSON 对象时触发结果序列化失败。
# "generation.chatapi_create_generation_task": {"ignore_result": True},
# "generation.poll_generation_task": {"ignore_result": True},
# "generation.download_generation_result_task": {"ignore_result": True},
"hot_opening.start_image_prompt_optimize": {"ignore_result": True},
"hot_opening.start_video_prompt_optimize": {"ignore_result": True},
"shot_replicate.analyze_original_video": {"ignore_result": True},
"shot_replicate.analyze_custom_segment_video": {"ignore_result": True},
"shot_replicate.split_one_segment": {"ignore_result": True},
"shot_replicate.start_image_prompt_optimize": {"ignore_result": True},
"shot_replicate.start_video_prompt_optimize": {"ignore_result": True},
},
worker_prefetch_multiplier=1,
broker_transport_options={
"visibility_timeout": 3600,
@@ -73,10 +90,12 @@ if broker_url:
"shot_replicate.split_one_segment": {"queue": "gen_result_download"},
"shot_replicate.start_image_prompt_optimize": {"queue": "gen_chatapi_create"},
"shot_replicate.start_video_prompt_optimize": {"queue": "gen_chatapi_create"},
"shot_replicate.recover_split_tasks_once": {"queue": "gen_result_download"},
"generation.recover_download_tasks_once": {"queue": "gen_result_download"},
"generation.recover_generation_tasks_once": {"queue": "gen_result_download"},
"module_async.recover_module_async_tasks_once": {"queue": "gen_result_download"},
# 恢复扫描统一走独立队列,避免占用下载/轮询/创建业务 worker。
"recovery.startup_recovery_once": {"queue": RECOVERY_QUEUE},
"shot_replicate.recover_split_tasks_once": {"queue": RECOVERY_QUEUE},
"generation.recover_download_tasks_once": {"queue": RECOVERY_QUEUE},
"generation.recover_generation_tasks_once": {"queue": RECOVERY_QUEUE},
"module_async.recover_module_async_tasks_once": {"queue": RECOVERY_QUEUE},
"user_oauth.update_oauth_accounts": {"queue": "default"},
"app.tasks.cleanup.*": {"queue": "default"},
},
@@ -86,7 +105,7 @@ else:
async def _try_acquire_startup_recovery_lock() -> bool:
"""任意 worker 启动时都可尝试抢恢复锁,避免依赖 hostname 命名。"""
"""任意 worker 启动时都可尝试抢恢复投递锁,避免依赖 hostname 命名。"""
from app.services.redis_registry_service import redis_acquire_lock
token = await redis_acquire_lock(
@@ -103,9 +122,9 @@ def on_worker_ready(sender=None, **kwargs):
注意:
- 不启用 Celery beat。
- 不要求新增第四条启动命令
- 不再依赖 worker hostname 是否包含 gen_result_download
- 所有 worker 都尝试抢 Redis 锁,只有抢到锁的 worker 投递恢复任务。
- 启动容灾保留,但只投递一个 recovery.startup_recovery_once 协调任务
- 协调任务走独立 gen_recovery 队列,串行扫描并把真实业务任务投回原队列
- 所有 worker 都尝试抢 Redis 投递锁,只有抢到锁的 worker 投递恢复任务。
"""
if celery_app is None:
return
@@ -122,39 +141,21 @@ def on_worker_ready(sender=None, **kwargs):
return
try:
from app.tasks.generation_recovery_tasks import (
recover_download_tasks_once,
recover_generation_tasks_once,
)
from app.tasks.shot_replicate_tasks import recover_split_tasks_once
from app.tasks.module_async_recovery_tasks import recover_module_async_tasks_once_task
from app.tasks.generation_recovery_tasks import startup_recovery_once
countdown = max(0, int(settings.DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS or 0))
recover_generation_tasks_once.apply_async(
startup_recovery_once.apply_async(
countdown=countdown,
queue="gen_result_download",
queue=RECOVERY_QUEUE,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
recover_download_tasks_once.apply_async(
countdown=countdown + 5,
queue="gen_result_download",
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
logger.info(
"启动容灾恢复协调任务已投递。queue=%s countdown=%s",
RECOVERY_QUEUE,
countdown,
)
recover_split_tasks_once.apply_async(
countdown=countdown + 10,
queue="gen_result_download",
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
recover_module_async_tasks_once_task.apply_async(
countdown=countdown + 15,
queue="gen_result_download",
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
logger.info("启动容灾恢复任务已投递。countdown=%s", countdown)
except Exception:
logger.exception("启动容灾恢复任务投递失败")
logger.exception("启动容灾恢复协调任务投递失败")
@worker_process_init.connect
@@ -12,6 +12,7 @@ from app.services.error_codes import extract_error_message
from app.services.generation_log_service import log_task_event
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_provider_service import create_provider_task
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
from app.services.redis_registry_service import ensure_aware_utc
from app.tasks.celery_app import celery_app
@@ -223,6 +224,7 @@ async def _run(task_id: str):
ensure_ascii=False,
default=str,
)
await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=task.provider_response_json)
if task.remote_result_url and not task.seedance_task_id:
# 同步图片路径:原 SDK 已经返回最终 URL。
@@ -288,7 +290,13 @@ async def _run(task_id: str):
if celery_app:
@celery_app.task(name="generation.chatapi_create_generation_task", bind=True, max_retries=3, default_retry_delay=30)
def chatapi_create_generation_task(self, task_id: str):
return run_async(_run(task_id))
try:
return run_async(_run(task_id))
except Exception as exc:
# 只处理 run_async/连接池/worker 中断等基础设施异常;业务异常已在 _run 内落库并退款。
retries = int(getattr(self.request, "retries", 0) or 0) + 1
countdown = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * max(1, retries)
raise self.retry(exc=exc, countdown=countdown)
else:
class _DisabledTask:
def delay(self, *args, **kwargs):
@@ -1,11 +1,24 @@
from __future__ import annotations
from app.tasks.async_runner import run_async
from datetime import datetime, timezone, timedelta
import errno
import math
import uuid
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.enums.generation_task import (
ALLOWED_GENERATION_MODES,
ChatGenerationPipelineStage,
ChatGenerationTaskEventType,
ChatGenerationTaskStatus,
GenerationType,
)
from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask
from app.services.celery_download_recovery_service import (
@@ -18,17 +31,17 @@ from app.services.error_codes import extract_error_message
from app.services.generation_download_service import download_generation_result
from app.services.generation_log_service import log_task_event
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
from app.services.resource_accounting_service import record_chat_task_generated_resource
from app.tasks.celery_app import celery_app
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"}
DOWNLOAD_QUEUE = "gen_result_download"
DOWNLOAD_STAGE_QUEUED = "download_queued"
DOWNLOAD_STAGE_DOWNLOADING = "downloading"
DOWNLOAD_STAGE_RETRY_WAITING = "retry_waiting"
DOWNLOAD_STAGE_DONE = "done"
DOWNLOAD_STAGE_FAILED = "download_failed"
DOWNLOAD_STAGE_QUEUED = ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value
DOWNLOAD_STAGE_DOWNLOADING = ChatGenerationPipelineStage.DOWNLOADING.value
DOWNLOAD_STAGE_RETRY_WAITING = ChatGenerationPipelineStage.RETRY_WAITING.value
DOWNLOAD_STAGE_DONE = ChatGenerationPipelineStage.DONE.value
DOWNLOAD_STAGE_FAILED = ChatGenerationPipelineStage.DOWNLOAD_FAILED.value
RESULT_READY_STAGE = ChatGenerationPipelineStage.RESULT_READY.value
def _now() -> datetime:
@@ -56,6 +69,15 @@ def _retry_at(attempt: int, now: datetime | None = None) -> datetime:
return now + timedelta(seconds=max(1, base * max(1, attempt)))
def _countdown_until(value: datetime | None, *, minimum: int = 1) -> int:
target = ensure_aware_utc(value)
if target is None:
return minimum
extra = int(getattr(settings, "DOWNLOAD_RETRY_COUNTDOWN_EXTRA_SECONDS", 1) or 0)
seconds = (target - _now()).total_seconds()
return max(minimum, math.ceil(seconds) + extra)
def _is_expired(value: datetime | None, now: datetime | None = None) -> bool:
value = ensure_aware_utc(value)
if value is None:
@@ -64,10 +86,10 @@ def _is_expired(value: datetime | None, now: datetime | None = None) -> bool:
def _is_already_completed(task: ChatGenerationTask) -> bool:
if task.status == "completed" or task.pipeline_stage == DOWNLOAD_STAGE_DONE:
if task.gen_type == "image" and task.image_url:
if task.status == ChatGenerationTaskStatus.COMPLETED.value or task.pipeline_stage == DOWNLOAD_STAGE_DONE:
if task.gen_type == GenerationType.IMAGE.value and task.image_url:
return True
if task.gen_type == "video" and task.video_url:
if task.gen_type == GenerationType.VIDEO.value and task.video_url:
return True
return False
@@ -77,6 +99,61 @@ def _build_celery_task_id(task_id: str, attempt: int | None = None, reason: str
return f"download:{task_id}:{int(attempt or 0)}:{safe_reason}:{uuid.uuid4().hex[:12]}"
def _is_non_retryable_download_error(exc: Exception) -> bool:
if not bool(getattr(settings, "DOWNLOAD_NON_RETRYABLE_LOCAL_ERRORS", True)):
return False
if isinstance(exc, PermissionError):
return True
if isinstance(exc, OSError) and getattr(exc, "errno", None) in {errno.EACCES, errno.EPERM, errno.ENOSPC, errno.EROFS, errno.ENAMETOOLONG}:
return True
message = str(exc).lower()
non_retryable_fragments = (
"permission denied",
"no space left on device",
"read-only file system",
"file name too long",
"invalid argument",
)
return any(fragment in message for fragment in non_retryable_fragments)
async def _log_download_event(
task: ChatGenerationTask | None = None,
*,
task_id: str | None = None,
event_type: ChatGenerationTaskEventType | str,
from_status: str | None = None,
to_status: str | None = None,
from_stage: str | None = None,
to_stage: str | None = None,
message: str | None = None,
detail: Any = None,
) -> None:
if not bool(getattr(settings, "DOWNLOAD_EVENT_VERBOSE_ENABLED", True)):
# 成功/失败关键事件仍保留;只关闭 verbose skip 事件。
critical = {
ChatGenerationTaskEventType.DOWNLOAD_START.value,
ChatGenerationTaskEventType.DOWNLOAD_SUCCESS.value,
ChatGenerationTaskEventType.DOWNLOAD_FAILED.value,
ChatGenerationTaskEventType.DOWNLOAD_FAILED_NON_RETRYABLE.value,
ChatGenerationTaskEventType.DOWNLOAD_RETRY_WAITING.value,
}
event_value = event_type.value if hasattr(event_type, "value") else str(event_type)
if event_value not in critical:
return
await log_task_event(
task,
task_id=task_id,
event_type=event_type.value if hasattr(event_type, "value") else str(event_type),
from_status=from_status,
to_status=to_status,
from_stage=from_stage,
to_stage=to_stage,
message=message,
detail=detail,
)
async def _register_active_from_task(
task: ChatGenerationTask,
*,
@@ -101,6 +178,62 @@ async def _register_active_from_task(
await upsert_download_active(record_id=task.id, payload=payload, check_at=check_at)
async def _apply_download_async(
task: ChatGenerationTask,
*,
priority: int,
countdown: int | None,
reason: str,
event_type: ChatGenerationTaskEventType,
failed_event_type: ChatGenerationTaskEventType,
) -> bool:
if not celery_app:
await _log_download_event(
task,
event_type=failed_event_type,
message="Celery 未启用,下载任务无法投递",
detail={"queue": DOWNLOAD_QUEUE, "priority": priority, "countdown": countdown, "reason": reason},
)
return False
try:
download_generation_result_task.apply_async(
args=[task.id],
queue=DOWNLOAD_QUEUE,
priority=priority,
countdown=countdown,
task_id=task.download_celery_task_id,
)
await _log_download_event(
task,
event_type=event_type,
to_stage=task.pipeline_stage,
detail={
"queue": DOWNLOAD_QUEUE,
"priority": priority,
"countdown": countdown,
"download_celery_task_id": task.download_celery_task_id,
"reason": reason,
"attempt": task.download_attempt_count,
},
)
return True
except Exception as exc:
await _log_download_event(
task,
event_type=failed_event_type,
message=str(exc),
detail={
"queue": DOWNLOAD_QUEUE,
"priority": priority,
"countdown": countdown,
"download_celery_task_id": task.download_celery_task_id,
"reason": reason,
"attempt": task.download_attempt_count,
},
)
return False
async def enqueue_download_task(
db: AsyncSession,
task: ChatGenerationTask,
@@ -110,11 +243,17 @@ async def enqueue_download_task(
countdown: int | None = None,
) -> str | None:
"""统一投递图片/视频下载任务,并同步 DB + Redis active 注册表。"""
if not task or task.generation_mode not in ALLOWED_GENERATION_MODES:
reason = reason or ("recover" if recover else "normal")
if not task:
return None
if task.status != "generating":
if task.generation_mode not in ALLOWED_GENERATION_MODES:
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_INVALID_MODE, message="不支持的 generation_mode", detail={"reason": reason})
return None
if task.status != ChatGenerationTaskStatus.GENERATING.value:
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING, message="任务不是 generating 状态", detail={"reason": reason, "status": task.status})
return None
if not task.remote_result_url:
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL, message="缺少 remote_result_url", detail={"reason": reason})
return None
now = _now()
@@ -122,9 +261,10 @@ async def enqueue_download_task(
celery_task_id = _build_celery_task_id(
task.id,
attempt=task.download_attempt_count or task.retry_count or 0,
reason=reason or ("recover" if recover else "normal"),
reason=reason,
)
old_stage = task.pipeline_stage
task.pipeline_stage = DOWNLOAD_STAGE_QUEUED
task.download_celery_task_id = celery_task_id
task.download_enqueued_at = now
@@ -137,13 +277,22 @@ async def enqueue_download_task(
check_at = _queue_timeout_at(now)
await _register_active_from_task(task, check_at=check_at, priority=priority, reason=reason)
if celery_app:
download_generation_result_task.apply_async(
args=[task.id],
queue=DOWNLOAD_QUEUE,
priority=priority,
countdown=countdown,
task_id=celery_task_id,
await _apply_download_async(
task,
priority=priority,
countdown=countdown,
reason=reason,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE,
failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE_FAILED if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE_FAILED,
)
if old_stage != DOWNLOAD_STAGE_QUEUED:
# 独立记录阶段变化的上下文,便于和真正投递事件对照。
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE,
from_stage=old_stage,
to_stage=DOWNLOAD_STAGE_QUEUED,
detail={"reason": reason, "recover": recover, "check_at": check_at},
)
return celery_task_id
@@ -158,34 +307,81 @@ async def _reload_task(db: AsyncSession, task_id: str) -> ChatGenerationTask | N
return result.scalar_one_or_none()
async def _reschedule_not_due_retry(task: ChatGenerationTask, *, now: datetime) -> None:
countdown = _countdown_until(task.download_next_retry_at)
await _register_active_from_task(
task,
check_at=ensure_aware_utc(task.download_next_retry_at) or (now + timedelta(seconds=countdown)),
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
reason="retry_waiting_not_due",
)
await _apply_download_async(
task,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
countdown=countdown,
reason="retry_waiting_not_due",
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE,
failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE_FAILED,
)
async def _claim_download_lease(db: AsyncSession, task: ChatGenerationTask) -> bool:
now = _now()
if not task or task.generation_mode not in ALLOWED_GENERATION_MODES:
if not task:
return False
if task.status != "generating":
if task.generation_mode not in ALLOWED_GENERATION_MODES:
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_INVALID_MODE, message="下载任务跳过:不支持的 generation_mode")
return False
if task.status != ChatGenerationTaskStatus.GENERATING.value:
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING, message="下载任务跳过:任务不是 generating 状态", detail={"status": task.status, "stage": task.pipeline_stage})
return False
if _is_already_completed(task):
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_ALREADY_COMPLETED, message="下载任务跳过:任务已完成", detail={"status": task.status, "stage": task.pipeline_stage})
return False
if not task.remote_result_url:
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL, message="下载任务跳过:缺少 remote_result_url")
return False
stage = task.pipeline_stage
if stage == DOWNLOAD_STAGE_DOWNLOADING:
if not _is_expired(task.download_lease_until, now):
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_DOWNLOADING_LEASE_ALIVE,
message="下载任务跳过:已有 downloading lease 且未过期",
detail={"lease_until": task.download_lease_until, "download_celery_task_id": task.download_celery_task_id},
)
return False
await log_task_event(
await _log_download_event(
task,
event_type="DOWNLOAD_STUCK_RECOVER",
event_type=ChatGenerationTaskEventType.DOWNLOAD_STUCK_RECOVER,
message=f"downloading lease 已过期,重新抢占下载。lease_until={task.download_lease_until}",
)
elif stage == DOWNLOAD_STAGE_RETRY_WAITING:
if not _is_expired(task.download_next_retry_at, now):
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_NOT_DUE,
message="下载重试提前触发,尚未到 next_retry_at,已重新投递延后重试",
detail={
"now": now,
"download_next_retry_at": task.download_next_retry_at,
"download_celery_task_id": task.download_celery_task_id,
},
)
await _reschedule_not_due_retry(task, now=now)
return False
elif stage in (DOWNLOAD_STAGE_QUEUED, "result_ready"):
elif stage in (DOWNLOAD_STAGE_QUEUED, RESULT_READY_STAGE):
pass
else:
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_STAGE_NOT_ALLOWED,
message="下载任务跳过:当前阶段不允许下载",
detail={"stage": stage, "status": task.status, "download_celery_task_id": task.download_celery_task_id},
)
return False
old_stage = stage
@@ -207,9 +403,9 @@ async def _claim_download_lease(db: AsyncSession, task: ChatGenerationTask) -> b
reason="claim_download_lease",
)
await log_task_event(
await _log_download_event(
task,
event_type="DOWNLOAD_START",
event_type=ChatGenerationTaskEventType.DOWNLOAD_START,
from_stage=old_stage,
to_stage=DOWNLOAD_STAGE_DOWNLOADING,
detail={
@@ -244,9 +440,9 @@ async def _mark_retry_waiting(db: AsyncSession, task: ChatGenerationTask, exc: E
reason="download_retry_waiting",
)
await log_task_event(
await _log_download_event(
task,
event_type="DOWNLOAD_RETRY_WAITING",
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_WAITING,
message=error_message,
to_stage=DOWNLOAD_STAGE_RETRY_WAITING,
detail={
@@ -262,16 +458,53 @@ def _should_final_fail(task: ChatGenerationTask) -> bool:
return int(task.download_attempt_count or task.retry_count or 0) >= int(settings.DOWNLOAD_TASK_MAX_ATTEMPTS or 3)
async def _mark_download_failed(
db: AsyncSession,
task: ChatGenerationTask,
*,
exc: Exception,
non_retryable: bool = False,
) -> None:
error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc)
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage=DOWNLOAD_STAGE_FAILED,
)
task.download_last_error = error_message
task.download_lease_until = None
task.download_next_retry_at = None
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
await notify_chat_generation_task_finished(db, task)
await db.commit()
await remove_download_active(task.id)
await _log_download_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_FAILED_NON_RETRYABLE if non_retryable else ChatGenerationTaskEventType.DOWNLOAD_FAILED,
message=task.error_message,
detail={
"download_attempt_count": task.download_attempt_count,
"max_attempts": settings.DOWNLOAD_TASK_MAX_ATTEMPTS,
"non_retryable": non_retryable,
"download_last_error": task.download_last_error,
},
)
async def _run(task_id: str):
async with async_session() as db:
result = await db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).with_for_update().limit(1)
)
task = result.scalar_one_or_none()
task = await _reload_task(db, task_id)
if not task:
await _log_download_event(
task_id=task_id,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_TASK_MISSING,
message="下载任务跳过:ChatGenerationTask 不存在或已软删",
)
await remove_download_active(task_id)
return
claimed = await _claim_download_lease(db, task)
@@ -283,15 +516,21 @@ async def _run(task_id: str):
task = await _reload_task(db, task_id)
if not task:
await _log_download_event(
task_id=task_id,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_TASK_MISSING,
message="下载完成后任务不存在或已软删",
)
await remove_download_active(task_id)
return
if task.gen_type == "image":
if task.gen_type == GenerationType.IMAGE.value:
task.image_url = downloaded.url
else:
task.video_url = downloaded.url
task.video_cover_url = downloaded.cover_url
task.status = "completed"
task.status = ChatGenerationTaskStatus.COMPLETED.value
task.pipeline_stage = DOWNLOAD_STAGE_DONE
task.generated_at = _now()
task.retry_count = 0
@@ -308,18 +547,18 @@ async def _run(task_id: str):
remote_url=task.remote_result_url,
generated_at=task.generated_at,
)
await sync_chat_generation_task_media_token_snapshot(db, task)
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
await notify_chat_generation_task_finished(db, task)
await db.commit()
await remove_download_active(task.id)
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
await notify_chat_generation_task_finished(db, task)
await db.commit()
await log_task_event(
await _log_download_event(
task,
event_type="DOWNLOAD_SUCCESS",
to_status="completed",
event_type=ChatGenerationTaskEventType.DOWNLOAD_SUCCESS,
to_status=ChatGenerationTaskStatus.COMPLETED.value,
to_stage=DOWNLOAD_STAGE_DONE,
detail={
"resource_url": downloaded.url,
@@ -337,54 +576,35 @@ async def _run(task_id: str):
task = await _reload_task(db, task_id)
if not task:
await remove_download_active(task_id)
return
if _should_final_fail(task):
error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc)
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage=DOWNLOAD_STAGE_FAILED,
)
task.download_last_error = error_message
task.download_lease_until = None
task.download_next_retry_at = None
await db.commit()
await remove_download_active(task.id)
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
await notify_chat_generation_task_finished(db, task)
await db.commit()
await log_task_event(
task,
event_type="DOWNLOAD_FAILED",
message=task.error_message,
detail={
"download_attempt_count": task.download_attempt_count,
"max_attempts": settings.DOWNLOAD_TASK_MAX_ATTEMPTS,
},
)
non_retryable = _is_non_retryable_download_error(exc)
if non_retryable or _should_final_fail(task):
await _mark_download_failed(db, task, exc=exc, non_retryable=non_retryable)
else:
next_retry_at = await _mark_retry_waiting(db, task, exc)
if celery_app:
delay_seconds = max(1, int((next_retry_at - _now()).total_seconds()))
download_generation_result_task.apply_async(
args=[task.id],
queue=DOWNLOAD_QUEUE,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
countdown=delay_seconds,
task_id=task.download_celery_task_id,
)
delay_seconds = _countdown_until(next_retry_at)
await _apply_download_async(
task,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
countdown=delay_seconds,
reason="download_exception_retry",
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE,
failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE_FAILED,
)
if celery_app:
@celery_app.task(name="generation.download_generation_result_task", bind=True, max_retries=3, default_retry_delay=30)
def download_generation_result_task(self, task_id: str):
return run_async(_run(task_id))
try:
return run_async(_run(task_id))
except Exception as exc:
# 只重试 run_async/连接池/worker 中断等基础设施异常;下载业务异常已在 _run 内写入 retry_waiting。
retries = int(getattr(self.request, "retries", 0) or 0) + 1
countdown = int(settings.DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS or 30) * max(1, retries)
raise self.retry(exc=exc, countdown=countdown)
else:
class _DisabledTask:
def delay(self, *args, **kwargs):
@@ -12,6 +12,7 @@ from app.services.error_codes import extract_error_message
from app.services.generation_log_service import log_task_event, log_provider_call
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_provider_service import poll_provider_task
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
from app.services.redis_registry_service import (
datetime_to_epoch,
ensure_aware_utc,
@@ -193,7 +194,8 @@ async def _run(task_id: str):
await remove_poll_active(task.id)
return
if _deadline_expired(task):
final_poll_before_timeout = _deadline_expired(task)
if final_poll_before_timeout and not (task.seedance_task_id or task.provider_task_id):
await _mark_timeout(db, task, message="任务轮询超时")
return
@@ -201,6 +203,14 @@ async def _run(task_id: str):
await _mark_failed(db, task, message="缺少外部任务ID")
return
if final_poll_before_timeout:
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT",
message="任务已到 deadline,执行最后一次供应商查询后再判定超时",
detail={"deadline_at": task.deadline_at, "stage": task.pipeline_stage},
)
# 标记本次正在轮询,并登记 poll lease。
# 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递。
task.pipeline_stage = "polling"
@@ -245,6 +255,7 @@ async def _run(task_id: str):
task.video_tokens_used = poll_result.get("video_tokens", 0) or 0
task.provider_response_json = response_data
await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=response_data)
if not task.remote_result_url:
await _mark_failed(db, task, message="供应商任务成功但未返回结果URL", detail=poll_result)
@@ -272,6 +283,16 @@ async def _run(task_id: str):
)
return
if final_poll_before_timeout:
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT_PENDING",
message=f"最终查询后供应商仍未完成,按超时处理。status={status}",
detail=poll_result,
)
await _mark_timeout(db, task, message="任务轮询超时")
return
# 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。
# 同时登记下一次 poll activeCelery countdown 丢失时可由恢复任务拉起。
task.pipeline_stage = "waiting_remote"
@@ -307,6 +328,15 @@ async def _run(task_id: str):
await remove_poll_active(task_id)
return
if final_poll_before_timeout:
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT_ERROR",
message=str(exc),
)
await _mark_timeout(db, task, message="任务轮询超时")
return
task.retry_count = (task.retry_count or 0) + 1
if task.retry_count > settings.CHATAPI_ASYNC_MAX_RETRIES:
@@ -337,7 +367,13 @@ async def _run(task_id: str):
if celery_app:
@celery_app.task(name="generation.poll_generation_task", bind=True, max_retries=3, default_retry_delay=30)
def poll_generation_task(self, task_id: str):
return run_async(_run(task_id))
try:
return run_async(_run(task_id))
except Exception as exc:
# 只重试基础设施异常;供应商失败/业务失败已在 _run 内处理。
retries = int(getattr(self.request, "retries", 0) or 0) + 1
countdown = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * max(1, retries)
raise self.retry(exc=exc, countdown=countdown)
else:
class _DisabledTask:
def delay(self, *args, **kwargs):
@@ -1,12 +1,20 @@
# app/tasks/generation_recovery_tasks.py
from __future__ import annotations
from typing import Any, Dict
import logging
from typing import Any, Awaitable, Callable, Dict
from app.config import settings
from app.models.base import async_session
from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock
from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
logger = logging.getLogger("video_gen")
RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or "gen_recovery"
RecoveryRunner = Callable[[], Awaitable[Dict[str, Any]]]
async def _run_download_once() -> Dict[str, Any]:
from app.services.generation_recovery_service import recover_download_tasks_once
@@ -22,16 +30,195 @@ async def _run_generation_once() -> Dict[str, Any]:
return await recover_generation_tasks_once(db)
async def _run_module_async_once() -> Dict[str, Any]:
from app.services.module_async_recovery_service import recover_module_async_tasks_once
async with async_session() as db:
return await recover_module_async_tasks_once(db)
async def _run_shot_split_once() -> Dict[str, Any]:
from app.services.shot_replicate_recovery_service import recover_shot_split_tasks_once
async with async_session() as db:
return await recover_shot_split_tasks_once(db)
async def _run_with_execution_lock(
*,
lock_key: str,
log_context: str,
runner: RecoveryRunner,
) -> Dict[str, Any]:
"""恢复任务执行锁。
worker_ready 的启动锁只保证“只投递一次”;如果 broker 中残留旧消息,
或者人工手动触发恢复任务,仍可能并发执行。这里再加执行锁,避免多个
恢复扫描同时扫库、抢行锁、抢连接池。
"""
redis = await get_registry_redis()
token: str | None = None
if redis is not None:
token = await redis_acquire_lock(
lock_key=lock_key,
ttl_seconds=int(settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
log_context=log_context,
)
if not token:
return {"skipped": "lock_held", "lock_key": lock_key}
else:
# Redis 不可用时仍允许 DB fallback 执行一次,避免恢复能力彻底失效。
logger.warning("恢复任务执行锁不可用,降级直接执行。context=%s", log_context)
try:
result = await runner()
result["execution_lock"] = "lock_acquired" if token else "redis_unavailable_run_db_fallback"
return result
finally:
if token:
await redis_release_lock(lock_key=lock_key, token=token, log_context=log_context)
async def _acquire_download_recovery_loop_lock() -> tuple[bool, str]:
"""下载恢复循环锁。
Redis 不可用时降级为直接执行 DB fallback,避免恢复能力彻底失效;
Redis 可用但锁被其他 worker 持有时,本轮跳过,不再重复投递下一轮。
"""
redis = await get_registry_redis()
if redis is None:
return True, "redis_unavailable_run_db_fallback"
token = await redis_acquire_lock(
lock_key=settings.DOWNLOAD_RECOVERY_LOOP_LOCK_KEY,
ttl_seconds=int(settings.DOWNLOAD_RECOVERY_LOOP_LOCK_TTL_SECONDS or 55),
log_context="download_recovery_loop",
)
return (bool(token), "lock_acquired" if token else "lock_held")
def _schedule_next_download_recovery_loop() -> None:
if not celery_app or not bool(getattr(settings, "DOWNLOAD_RECOVERY_LOOP_ENABLED", False)):
return
try:
recover_download_tasks_once.apply_async(
countdown=max(1, int(settings.DOWNLOAD_RECOVERY_INTERVAL_SECONDS or 60)),
queue=RECOVERY_QUEUE,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
except Exception:
logger.exception("下载恢复循环下一轮投递失败")
async def _run_startup_recovery_once() -> Dict[str, Any]:
"""启动容灾协调器:串行跑恢复扫描。
真实业务任务仍投递回原队列:
- 创建/提词/视频分析 -> gen_chatapi_create
- provider poll -> gen_provider_poll
- 下载/ffmpeg 切片 -> gen_result_download
恢复扫描本身只走 gen_recovery,避免堵住业务 worker。
"""
return await _run_with_execution_lock(
lock_key=settings.CELERY_RECOVERY_STARTUP_TASK_LOCK_KEY,
log_context="startup_recovery_once",
runner=_run_startup_recovery_steps,
)
async def _run_startup_recovery_steps() -> Dict[str, Any]:
results: Dict[str, Any] = {}
steps: list[tuple[str, str, str, RecoveryRunner]] = [
(
"module_async",
settings.MODULE_ASYNC_RECOVERY_LOCK_KEY,
"module_async_recovery",
_run_module_async_once,
),
(
"shot_split",
settings.SHOT_SPLIT_RECOVERY_LOCK_KEY,
"shot_split_recovery",
_run_shot_split_once,
),
(
"generation",
settings.GENERATION_RECOVERY_LOCK_KEY,
"generation_recovery",
_run_generation_once,
),
(
"download",
settings.DOWNLOAD_RECOVERY_LOCK_KEY,
"download_recovery",
_run_download_once,
),
]
for name, lock_key, log_context, runner in steps:
try:
results[name] = await _run_with_execution_lock(
lock_key=lock_key,
log_context=log_context,
runner=runner,
)
except Exception as exc:
logger.exception("启动容灾步骤执行失败。step=%s", name)
results[name] = {"error": str(exc)}
return {"steps": results}
if celery_app:
@celery_app.task(name="generation.recover_download_tasks_once")
def recover_download_tasks_once() -> Dict[str, Any]:
return run_async(_run_download_once())
@celery_app.task(
name="recovery.startup_recovery_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def startup_recovery_once(self) -> Dict[str, Any]:
return run_async(_run_startup_recovery_once())
@celery_app.task(name="generation.recover_generation_tasks_once")
def recover_generation_tasks_once() -> Dict[str, Any]:
return run_async(_run_generation_once())
@celery_app.task(
name="generation.recover_download_tasks_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def recover_download_tasks_once(self) -> Dict[str, Any]:
acquired, reason = run_async(_acquire_download_recovery_loop_lock())
if not acquired:
return {"skipped": reason}
try:
result = run_async(
_run_with_execution_lock(
lock_key=settings.DOWNLOAD_RECOVERY_LOCK_KEY,
log_context="download_recovery",
runner=_run_download_once,
)
)
result["loop_lock"] = reason
return result
finally:
_schedule_next_download_recovery_loop()
@celery_app.task(
name="generation.recover_generation_tasks_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def recover_generation_tasks_once(self) -> Dict[str, Any]:
return run_async(
_run_with_execution_lock(
lock_key=settings.GENERATION_RECOVERY_LOCK_KEY,
log_context="generation_recovery",
runner=_run_generation_once,
)
)
else:
@@ -42,5 +229,6 @@ else:
def apply_async(self, *args: Any, **kwargs: Any) -> None:
raise RuntimeError("Celery is disabled")
startup_recovery_once = _DisabledTask()
recover_download_tasks_once = _DisabledTask()
recover_generation_tasks_once = _DisabledTask()
@@ -21,7 +21,7 @@ from app.tasks.celery_app import celery_app
MODULE = ModuleCodeEnum.HOT_OPENING_REPLICATE.value
async def _run_image_prompt(project_id: str, step_id: str | None = None):
async def _run_image_prompt(project_id: str, step_id: str | None = None) -> None:
lock_token: str | None = None
if step_id:
lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id)
@@ -38,17 +38,17 @@ async def _run_image_prompt(project_id: str, step_id: str | None = None):
await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id)
try:
async with async_session() as db:
result = await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id)
await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id)
await db.commit()
if step_id:
await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id)
return result
return None
finally:
if step_id:
await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token)
async def _run_video_prompt(project_id: str, step_id: str | None = None):
async def _run_video_prompt(project_id: str, step_id: str | None = None) -> None:
lock_token: str | None = None
if step_id:
lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id)
@@ -64,18 +64,18 @@ async def _run_video_prompt(project_id: str, step_id: str | None = None):
await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id)
try:
async with async_session() as db:
result = await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id)
await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id)
await db.commit()
if step_id:
await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id)
return result
return None
finally:
if step_id:
await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token)
if celery_app:
@celery_app.task(name="hot_opening.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30)
@celery_app.task(name="hot_opening.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30, ignore_result=True)
def start_image_prompt_optimize(self, project_id: str, step_id: str | None = None):
"""手动触发后的图片 AI 提词任务。
@@ -84,11 +84,12 @@ if celery_app:
service 内部已经落库为业务失败的情况不会抛出异常,也不会重复 retry。
"""
try:
return run_async(_run_image_prompt(project_id, step_id))
run_async(_run_image_prompt(project_id, step_id))
return None
except Exception as exc:
raise self.retry(exc=exc) from exc
@celery_app.task(name="hot_opening.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30)
@celery_app.task(name="hot_opening.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30, ignore_result=True)
def start_video_prompt_optimize(self, project_id: str, step_id: str | None = None):
"""手动触发后的视频 AI 提词任务。
@@ -97,7 +98,8 @@ if celery_app:
service 内部已经落库为业务失败的情况不会抛出异常,也不会重复 retry。
"""
try:
return run_async(_run_video_prompt(project_id, step_id))
run_async(_run_video_prompt(project_id, step_id))
return None
except Exception as exc:
raise self.retry(exc=exc) from exc
else:
@@ -109,4 +111,4 @@ else:
raise RuntimeError("Celery is disabled")
start_image_prompt_optimize = _DisabledTask()
start_video_prompt_optimize = _DisabledTask()
start_video_prompt_optimize = _DisabledTask()
@@ -2,21 +2,49 @@ from __future__ import annotations
from typing import Any
from app.config import settings
from app.models.base import async_session
from app.services.module_async_recovery_service import recover_module_async_tasks_once
from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock
from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
async def _run_recover_module_async_tasks_once() -> dict[str, Any]:
async with async_session() as db:
return await recover_module_async_tasks_once(db)
redis = await get_registry_redis()
token: str | None = None
if redis is not None:
token = await redis_acquire_lock(
lock_key=settings.MODULE_ASYNC_RECOVERY_LOCK_KEY,
ttl_seconds=int(settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
log_context="module_async_recovery",
)
if not token:
return {"skipped": "lock_held", "lock_key": settings.MODULE_ASYNC_RECOVERY_LOCK_KEY}
try:
async with async_session() as db:
result = await recover_module_async_tasks_once(db)
result["execution_lock"] = "lock_acquired" if token else "redis_unavailable_run_db_fallback"
return result
finally:
if token:
await redis_release_lock(
lock_key=settings.MODULE_ASYNC_RECOVERY_LOCK_KEY,
token=token,
log_context="module_async_recovery",
)
if celery_app:
@celery_app.task(name="module_async.recover_module_async_tasks_once")
def recover_module_async_tasks_once_task() -> dict[str, Any]:
@celery_app.task(
name="module_async.recover_module_async_tasks_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def recover_module_async_tasks_once_task(self) -> dict[str, Any]:
return run_async(_run_recover_module_async_tasks_once())
else:
@@ -21,7 +21,7 @@ from app.tasks.celery_app import celery_app
MODULE = ModuleCodeEnum.SHOT_REPLICATE.value
async def _run_image_prompt(project_id: str, step_id: str | None = None):
async def _run_image_prompt(project_id: str, step_id: str | None = None) -> None:
lock_token: str | None = None
if step_id:
lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id)
@@ -37,17 +37,17 @@ async def _run_image_prompt(project_id: str, step_id: str | None = None):
await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id)
try:
async with async_session() as db:
result = await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id)
await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id)
await db.commit()
if step_id:
await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id)
return result
return None
finally:
if step_id:
await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token)
async def _run_video_prompt(project_id: str, step_id: str | None = None):
async def _run_video_prompt(project_id: str, step_id: str | None = None) -> None:
lock_token: str | None = None
if step_id:
lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id)
@@ -63,11 +63,11 @@ async def _run_video_prompt(project_id: str, step_id: str | None = None):
await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id)
try:
async with async_session() as db:
result = await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id)
await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id)
await db.commit()
if step_id:
await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id)
return result
return None
finally:
if step_id:
await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token)
@@ -75,7 +75,7 @@ async def _run_video_prompt(project_id: str, step_id: str | None = None):
if celery_app:
@celery_app.task(name="shot_replicate.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30)
@celery_app.task(name="shot_replicate.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30, ignore_result=True)
def start_image_prompt_optimize(self, project_id: str, step_id: str | None = None):
"""手动触发后的图片 AI 提词任务。
@@ -84,12 +84,13 @@ if celery_app:
service 内部已经落库为业务失败的情况不会抛出异常,也不会重复 retry。
"""
try:
return run_async(_run_image_prompt(project_id, step_id))
run_async(_run_image_prompt(project_id, step_id))
return None
except Exception as exc:
raise self.retry(exc=exc) from exc
@celery_app.task(name="shot_replicate.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30)
@celery_app.task(name="shot_replicate.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30, ignore_result=True)
def start_video_prompt_optimize(self, project_id: str, step_id: str | None = None):
"""手动触发后的视频 AI 提词任务。
@@ -98,7 +99,8 @@ if celery_app:
service 内部已经落库为业务失败的情况不会抛出异常,也不会重复 retry。
"""
try:
return run_async(_run_video_prompt(project_id, step_id))
run_async(_run_video_prompt(project_id, step_id))
return None
except Exception as exc:
raise self.retry(exc=exc) from exc
@@ -112,4 +114,4 @@ else:
raise RuntimeError("Celery is disabled")
start_image_prompt_optimize = _DisabledTask()
start_video_prompt_optimize = _DisabledTask()
start_video_prompt_optimize = _DisabledTask()
@@ -20,7 +20,7 @@ from app.models.base import async_session
from app.models.shot_replicate_segment import ShotReplicateSegment
from app.models.shot_replicate_task_set import ShotReplicateTaskSet
from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event
from app.services.redis_registry_service import redis_acquire_lock, redis_release_lock
from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock
from app.services.module_async_recovery_service import (
OBJECT_SHOT_SEGMENT_ANALYSIS,
OBJECT_SHOT_SPLIT_SEGMENT,
@@ -561,8 +561,29 @@ async def _run_split_one_segment(segment_id: str) -> None:
async def _run_recover_split_tasks_once() -> dict[str, Any]:
from app.services.shot_replicate_recovery_service import recover_shot_split_tasks_once
async with async_session() as db:
return await recover_shot_split_tasks_once(db)
redis = await get_registry_redis()
token: str | None = None
if redis is not None:
token = await redis_acquire_lock(
lock_key=settings.SHOT_SPLIT_RECOVERY_LOCK_KEY,
ttl_seconds=int(settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
log_context="shot_split_recovery",
)
if not token:
return {"skipped": "lock_held", "lock_key": settings.SHOT_SPLIT_RECOVERY_LOCK_KEY}
try:
async with async_session() as db:
result = await recover_shot_split_tasks_once(db)
result["execution_lock"] = "lock_acquired" if token else "redis_unavailable_run_db_fallback"
return result
finally:
if token:
await redis_release_lock(
lock_key=settings.SHOT_SPLIT_RECOVERY_LOCK_KEY,
token=token,
log_context="shot_split_recovery",
)
if celery_app:
@@ -582,8 +603,13 @@ if celery_app:
return run_async(_run_analyze_custom_segment_video(segment_id))
@celery_app.task(name="shot_replicate.recover_split_tasks_once")
def recover_split_tasks_once() -> dict[str, Any]:
@celery_app.task(
name="shot_replicate.recover_split_tasks_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def recover_split_tasks_once(self) -> dict[str, Any]:
return run_async(_run_recover_split_tasks_once())
else: