celery异步恢复任务独立队列|扩增celery子进程连接池上限配置

This commit is contained in:
2026-06-26 13:15:25 +08:00
parent 06cdab9cbc
commit ee03242e6c
10 changed files with 517 additions and 245 deletions
+21 -12
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,10 +125,10 @@ class Settings(BaseSettings):
DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS: int = 30
DOWNLOAD_TASK_LEASE_SECONDS: int = 10 * 60
DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS: int = 5 * 60
DOWNLOAD_RECOVERY_BATCH_SIZE: int = 100
DOWNLOAD_RECOVERY_BATCH_SIZE: int = 20
DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS: int = 3
# 下载恢复自循环:不依赖 Celery beat,不新增 worker;由 gen_result_download 队列周期扫描 DB/Redis。
DOWNLOAD_RECOVERY_LOOP_ENABLED: bool = True
DOWNLOAD_RECOVERY_LOOP_ENABLED: bool = False
DOWNLOAD_RECOVERY_INTERVAL_SECONDS: int = 60
DOWNLOAD_RECOVERY_LOOP_LOCK_KEY: str = "vg:celery:download_recovery_loop_lock"
DOWNLOAD_RECOVERY_LOOP_LOCK_TTL_SECONDS: int = 55
@@ -141,23 +141,32 @@ class Settings(BaseSettings):
# Celery 生成链路 / provider poll 容灾配置。
# 说明:
# - 不新增 Celery worker;恢复任务仍投递到 gen_result_download
# - worker_ready 每个 worker 都会尝试抢启动恢复锁,只有抢到锁的 worker 投递恢复任务。
# - 启动容灾保留,但恢复扫描独立投递到 CELERY_RECOVERY_QUEUE
# - worker_ready 每个 worker 都会尝试抢启动恢复锁,只有抢到锁的 worker 投递恢复协调任务。
# - poll active 使用独立 Redis key,避免影响稳定的下载 active 注册表。
GENERATION_RECOVERY_BATCH_SIZE: int = 100
GENERATION_RECOVERY_MAX_ROUNDS: int = 5
POLL_RECOVERY_BATCH_SIZE: int = 100
GENERATION_RECOVERY_BATCH_SIZE: int = 20
GENERATION_RECOVERY_MAX_ROUNDS: int = 1
POLL_RECOVERY_BATCH_SIZE: int = 20
POLL_TASK_LEASE_SECONDS: int = 5 * 60
POLL_TASK_QUEUE_TIMEOUT_SECONDS: int = 2 * 60
POLL_ACTIVE_REDIS_HASH_KEY: str = "vg:celery:poll:active"
POLL_ACTIVE_REDIS_ZSET_KEY: str = "vg:celery:poll:active_index"
CELERY_STARTUP_RECOVERY_LOCK_KEY: str = "vg:celery:startup_recovery_lock"
CELERY_STARTUP_RECOVERY_LOCK_TTL_SECONDS: int = 120
CELERY_RECOVERY_QUEUE: str = "gen_recovery"
CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS: int = 10 * 60
CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS: int = 300
CELERY_RECOVERY_TIME_LIMIT_SECONDS: int = 420
CELERY_RECOVERY_STARTUP_TASK_LOCK_KEY: str = "vg:celery:startup_recovery_task_lock"
GENERATION_RECOVERY_LOCK_KEY: str = "vg:celery:generation_recovery_lock"
DOWNLOAD_RECOVERY_LOCK_KEY: str = "vg:celery:download_recovery_lock"
MODULE_ASYNC_RECOVERY_LOCK_KEY: str = "vg:celery:module_async_recovery_lock"
SHOT_SPLIT_RECOVERY_LOCK_KEY: str = "vg:celery:shot_split_recovery_lock"
# 模块异步任务容灾配置。
# 覆盖 ModuleGenerationStep 提词任务、shot 原视频/片段分析、shot ffmpeg 切割 active 注册。
# 不新增 worker 队列:恢复扫描gen_result_download,真实业务任务回到原始队列。
MODULE_ASYNC_RECOVERY_BATCH_SIZE: int = 100
# 恢复扫描走 CELERY_RECOVERY_QUEUE,真实业务任务回到原始队列。
MODULE_ASYNC_RECOVERY_BATCH_SIZE: int = 20
MODULE_ASYNC_QUEUE_TIMEOUT_SECONDS: int = 5 * 60
MODULE_ASYNC_LEASE_SECONDS: int = 10 * 60
MODULE_ASYNC_LOCK_TTL_SECONDS: int = 10 * 60
@@ -201,7 +210,7 @@ class Settings(BaseSettings):
SHOT_SPLIT_RETRY_BACKOFF_SECONDS: int = 30
SHOT_SPLIT_LEASE_SECONDS: int = 10 * 60
SHOT_SPLIT_PENDING_TIMEOUT_SECONDS: int = 5 * 60
SHOT_SPLIT_RECOVERY_BATCH_SIZE: int = 50
SHOT_SPLIT_RECOVERY_BATCH_SIZE: int = 20
SHOT_SPLIT_LOCK_KEY_PREFIX: str = "vg:shot_replicate:split:lock"
SHOT_SPLIT_SEMAPHORE_KEY_PREFIX: str = "vg:shot_replicate:split:semaphore"
@@ -23,10 +23,8 @@ from app.services.celery_download_recovery_service import (
postpone_download_active_check,
remove_download_active,
)
from app.services.generation_log_service import log_provider_call, log_task_event
from app.services.generation_log_service import log_task_event
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
from app.services.generation_provider_service import poll_provider_task
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.redis_registry_service import (
redis_get_due_registry_ids,
@@ -145,10 +143,20 @@ async def recover_one_download_task(
return "clean_final_state"
if task.status != ChatGenerationTaskStatus.GENERATING.value:
await remove_download_active(task.id)
await log_task_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING.value, message=f"{source} 下载恢复跳过:任务不是 generating", detail={"status": task.status, "stage": task.pipeline_stage})
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING.value,
message=f"{source} 下载恢复跳过:任务不是 generating",
detail={"status": task.status, "stage": task.pipeline_stage},
)
return "clean_not_generating"
if not task.remote_result_url:
await log_task_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL.value, message=f"{source} 下载恢复跳过:缺少 remote_result_url", detail={"status": task.status, "stage": task.pipeline_stage})
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL.value,
message=f"{source} 下载恢复跳过:缺少 remote_result_url",
detail={"status": task.status, "stage": task.pipeline_stage},
)
return "skip_no_remote_result_url"
stage = task.pipeline_stage
@@ -362,94 +370,6 @@ async def _mark_failed(
return "mark_failed"
async def _try_final_poll_before_timeout(db: AsyncSession, task: ChatGenerationTask) -> str:
"""超时前最后查一次供应商,避免 Celery 中断导致本地假超时。
如果供应商已经成功,继续进入下载;如果仍 running 或查询失败,再按超时处理。
"""
from app.tasks.generation_download_tasks import enqueue_download_task
if not (task.provider_task_id or task.seedance_task_id):
return await _mark_timeout(db, task)
try:
poll_result = await poll_provider_task(db, task)
status = poll_result.get("status")
response_data = poll_result.get("response_data")
except Exception as exc:
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT_ERROR",
message=str(exc),
)
return await _mark_timeout(db, task)
try:
provider_response = json.loads(response_data or "{}")
except Exception:
provider_response = {"raw": response_data}
snapshot = _engine_snapshot(task)
await log_provider_call(
task,
provider=snapshot.get("provider") or "ark",
api_type=f"{task.gen_type}_final_poll_before_timeout",
model=snapshot.get("model_name"),
engine_id=task.engine_id,
status="success",
provider_task_id=task.seedance_task_id or task.provider_task_id,
response_data=provider_response,
)
if _is_success(status):
if task.gen_type == "image":
task.remote_result_url = poll_result.get("image_url")
task.image_tokens_used = poll_result.get("image_tokens", 0) or 0
else:
task.remote_result_url = poll_result.get("video_url")
task.video_tokens_used = poll_result.get("video_tokens", 0) or 0
task.provider_response_json = response_data
await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=response_data)
if not task.remote_result_url:
return await _mark_failed(
db,
task,
error_message="供应商任务成功但未返回结果URL",
detail=poll_result,
)
task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
task.retry_count = 0
await db.commit()
await _remove_poll_active(task.id)
await log_task_event(
task,
event_type="POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY",
to_stage=ChatGenerationPipelineStage.RESULT_READY.value,
detail=poll_result,
)
await enqueue_download_task(db, task, recover=True, reason="final_poll_before_timeout_success")
return "recover_timeout_success_to_download"
if _is_failed(status):
task.provider_response_json = response_data
return await _mark_failed(
db,
task,
error_message=poll_result.get("error") or f"供应商任务失败: {status}",
detail=poll_result,
)
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT_PENDING",
message=f"status={status}",
detail=poll_result,
)
return await _mark_timeout(db, task)
async def recover_one_generation_task(
db: AsyncSession,
task: ChatGenerationTask,
@@ -457,6 +377,14 @@ async def recover_one_generation_task(
payload: dict[str, Any] | None = None,
source: str = "startup_db",
) -> str:
"""恢复单个生成任务。
分流原则:
1. 已有 remote_result_url:只恢复下载,不 poll,不重新 create。
2. 已有 provider_task_id/seedance_task_id:恢复 poll。
3. 无结果 URL、无供应商任务 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
@@ -472,105 +400,133 @@ async def recover_one_generation_task(
if _is_final_task_state(task):
await _remove_poll_active(task.id)
return "clean_final_state"
if task.status != "generating":
if task.status != ChatGenerationTaskStatus.GENERATING.value:
await _remove_poll_active(task.id)
return "clean_not_generating"
if task.deadline_at and _is_expired(task.deadline_at, current_time):
if task.pipeline_stage in (ChatGenerationPipelineStage.WAITING_REMOTE.value, ChatGenerationPipelineStage.POLLING.value):
return await _try_final_poll_before_timeout(db, task)
return await _mark_timeout(db, task)
has_remote_result = bool(str(task.remote_result_url or "").strip())
has_provider_task_id = bool(str(task.provider_task_id or "").strip() or str(task.seedance_task_id or "").strip())
is_deadline_expired = bool(task.deadline_at and _is_expired(task.deadline_at, current_time))
if task.pipeline_stage in (ChatGenerationPipelineStage.QUEUED.value, ChatGenerationPipelineStage.PREPARING.value, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value):
if task.provider_task_id or task.seedance_task_id:
# 最高优先级:只要远程结果 URL 已经落库,说明生成侧已经成功。
# 不管当前 pipeline_stage 是 queued/creating/waiting/result_ready/download_*,恢复时都不能重复 create 或 poll。
if has_remote_result:
await _remove_poll_active(task.id)
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message=f"{source} 发现任务已存在 remote_result_url,恢复投递下载队列",
detail={
"pipeline_stage": task.pipeline_stage,
"payload": redis_payload,
"deadline_expired": is_deadline_expired,
},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_has_remote_result_url",
)
return "recover_download_has_remote_result"
# 已经过 deadline 且没有结果 URL
# - 有供应商任务 ID:交给 poll worker 做最后一次状态确认;
# - 没有供应商任务 ID:说明没有可查询的远程任务,直接按超时失败处理,不再重新 create。
if is_deadline_expired:
if has_provider_task_id:
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
await db.commit()
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message=f"{source} 发现创建阶段已存在供应商任务ID,恢复投递轮询队列",
message=f"{source} 发现任务已到 deadline 且存在供应商任务ID,投递 poll 队列做最终查询",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
poll_generation_task.apply_async(args=[task.id], queue=POLL_QUEUE, countdown=0)
await register_poll_active(
task,
check_at=_poll_queue_timeout_at(),
reason=f"{source}_create_stage_has_provider_id",
reason=f"{source}_deadline_final_poll",
)
return "recover_poll_from_create_stage"
return "recover_deadline_final_poll"
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message=f"{source} 发现创建阶段任务未完成,恢复投递创建队列",
event_type="GENERATION_RECOVERY_TIMEOUT",
message=f"{source} 发现任务已到 deadline,且没有 remote_result_url/供应商任务ID,按超时失败处理",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
chatapi_create_generation_task.apply_async(
args=[task.id],
queue="gen_chatapi_create",
countdown=0,
)
return "recover_create"
return await _mark_timeout(db, task)
if task.pipeline_stage in (ChatGenerationPipelineStage.WAITING_REMOTE.value, ChatGenerationPipelineStage.POLLING.value):
if task.remote_result_url:
await _remove_poll_active(task.id)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_waiting_remote_has_result",
)
return "recover_waiting_has_result"
if task.provider_task_id or task.seedance_task_id:
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
message=f"{source} 发现远程等待/轮询阶段任务未完成,恢复投递轮询队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
# 未过 deadline:有供应商任务 ID 才允许恢复到 poll 队列。
if has_provider_task_id:
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
await db.commit()
poll_generation_task.apply_async(
args=[task.id],
queue=POLL_QUEUE,
countdown=0,
await log_task_event(
task,
event_type="GENERATION_RECOVERY_ENQUEUE",
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}_recover_poll",
reason=f"{source}_has_provider_task_id",
)
return "recover_poll"
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} 发现任务缺少供应商任务ID,恢复投递创建队列",
message=f"{source} 发现任务未超时且缺少 remote_result_url/供应商任务ID,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
await db.commit()
await _remove_poll_active(task.id)
chatapi_create_generation_task.apply_async(
args=[task.id],
queue="gen_chatapi_create",
countdown=0,
)
return "recover_create_missing_provider_id"
return "recover_create_no_remote_no_provider_before_deadline"
# 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)
if task.remote_result_url:
await enqueue_download_task(
db,
await log_task_event(
task,
recover=True,
reason=f"{source}_generation_result_ready",
event_type="GENERATION_RECOVERY_ENQUEUE",
message=f"{source} 发现 result_ready 但缺少 remote_result_url,未超时,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
return "recover_result_ready"
return "skip_result_ready_no_url"
chatapi_create_generation_task.apply_async(
args=[task.id],
queue="gen_chatapi_create",
countdown=0,
)
return "recover_create_result_ready_no_url_before_deadline"
return f"skip_stage_{task.pipeline_stage}"
@@ -670,11 +626,8 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
if len(tasks) < batch_size or progressed_this_round <= 0:
break
# 下载阶段单独跑 DB fallback。
download_result = await recover_download_tasks_once(db)
return {
"checked": len(checked_ids),
"db_checked": total_db_checked,
"results": results,
"download_recovery": download_result,
}
+81 -3
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,14 +71,16 @@ 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()
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))
loop.run_until_complete(_dispose_async_resources())
loop.run_until_complete(loop.shutdown_asyncgens())
loop.close()
@@ -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()
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)
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():
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():
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
+21 -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
@@ -73,10 +76,12 @@ if broker_url:
"shot_replicate.split_one_segment": {"queue": "gen_result_download"},
"shot_replicate.start_image_prompt_optimize": {"queue": "gen_chatapi_create"},
"shot_replicate.start_video_prompt_optimize": {"queue": "gen_chatapi_create"},
"shot_replicate.recover_split_tasks_once": {"queue": "gen_result_download"},
"generation.recover_download_tasks_once": {"queue": "gen_result_download"},
"generation.recover_generation_tasks_once": {"queue": "gen_result_download"},
"module_async.recover_module_async_tasks_once": {"queue": "gen_result_download"},
# 恢复扫描统一走独立队列,避免占用下载/轮询/创建业务 worker。
"recovery.startup_recovery_once": {"queue": RECOVERY_QUEUE},
"shot_replicate.recover_split_tasks_once": {"queue": RECOVERY_QUEUE},
"generation.recover_download_tasks_once": {"queue": RECOVERY_QUEUE},
"generation.recover_generation_tasks_once": {"queue": RECOVERY_QUEUE},
"module_async.recover_module_async_tasks_once": {"queue": RECOVERY_QUEUE},
"user_oauth.update_oauth_accounts": {"queue": "default"},
"app.tasks.cleanup.*": {"queue": "default"},
},
@@ -86,7 +91,7 @@ else:
async def _try_acquire_startup_recovery_lock() -> bool:
"""任意 worker 启动时都可尝试抢恢复锁,避免依赖 hostname 命名。"""
"""任意 worker 启动时都可尝试抢恢复投递锁,避免依赖 hostname 命名。"""
from app.services.redis_registry_service import redis_acquire_lock
token = await redis_acquire_lock(
@@ -103,9 +108,9 @@ def on_worker_ready(sender=None, **kwargs):
注意:
- 不启用 Celery beat。
- 不要求新增第四条启动命令
- 不再依赖 worker hostname 是否包含 gen_result_download
- 所有 worker 都尝试抢 Redis 锁,只有抢到锁的 worker 投递恢复任务。
- 启动容灾保留,但只投递一个 recovery.startup_recovery_once 协调任务
- 协调任务走独立 gen_recovery 队列,串行扫描并把真实业务任务投回原队列
- 所有 worker 都尝试抢 Redis 投递锁,只有抢到锁的 worker 投递恢复任务。
"""
if celery_app is None:
return
@@ -122,39 +127,21 @@ def on_worker_ready(sender=None, **kwargs):
return
try:
from app.tasks.generation_recovery_tasks import (
recover_download_tasks_once,
recover_generation_tasks_once,
)
from app.tasks.shot_replicate_tasks import recover_split_tasks_once
from app.tasks.module_async_recovery_tasks import recover_module_async_tasks_once_task
from app.tasks.generation_recovery_tasks import startup_recovery_once
countdown = max(0, int(settings.DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS or 0))
recover_generation_tasks_once.apply_async(
startup_recovery_once.apply_async(
countdown=countdown,
queue="gen_result_download",
queue=RECOVERY_QUEUE,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
recover_download_tasks_once.apply_async(
countdown=countdown + 5,
queue="gen_result_download",
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
logger.info(
"启动容灾恢复协调任务已投递。queue=%s countdown=%s",
RECOVERY_QUEUE,
countdown,
)
recover_split_tasks_once.apply_async(
countdown=countdown + 10,
queue="gen_result_download",
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
recover_module_async_tasks_once_task.apply_async(
countdown=countdown + 15,
queue="gen_result_download",
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
logger.info("启动容灾恢复任务已投递。countdown=%s", countdown)
except Exception:
logger.exception("启动容灾恢复任务投递失败")
logger.exception("启动容灾恢复协调任务投递失败")
@worker_process_init.connect
@@ -290,7 +290,13 @@ async def _run(task_id: str):
if celery_app:
@celery_app.task(name="generation.chatapi_create_generation_task", bind=True, max_retries=3, default_retry_delay=30)
def chatapi_create_generation_task(self, task_id: str):
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):
@@ -598,7 +598,13 @@ async def _run(task_id: str):
if celery_app:
@celery_app.task(name="generation.download_generation_result_task", bind=True, max_retries=3, default_retry_delay=30)
def download_generation_result_task(self, task_id: str):
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):
@@ -194,7 +194,8 @@ async def _run(task_id: str):
await remove_poll_active(task.id)
return
if _deadline_expired(task):
final_poll_before_timeout = _deadline_expired(task)
if final_poll_before_timeout and not (task.seedance_task_id or task.provider_task_id):
await _mark_timeout(db, task, message="任务轮询超时")
return
@@ -202,6 +203,14 @@ async def _run(task_id: str):
await _mark_failed(db, task, message="缺少外部任务ID")
return
if final_poll_before_timeout:
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT",
message="任务已到 deadline,执行最后一次供应商查询后再判定超时",
detail={"deadline_at": task.deadline_at, "stage": task.pipeline_stage},
)
# 标记本次正在轮询,并登记 poll lease。
# 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递。
task.pipeline_stage = "polling"
@@ -274,6 +283,16 @@ async def _run(task_id: str):
)
return
if final_poll_before_timeout:
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT_PENDING",
message=f"最终查询后供应商仍未完成,按超时处理。status={status}",
detail=poll_result,
)
await _mark_timeout(db, task, message="任务轮询超时")
return
# 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。
# 同时登记下一次 poll activeCelery countdown 丢失时可由恢复任务拉起。
task.pipeline_stage = "waiting_remote"
@@ -309,6 +328,15 @@ async def _run(task_id: str):
await remove_poll_active(task_id)
return
if final_poll_before_timeout:
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT_ERROR",
message=str(exc),
)
await _mark_timeout(db, task, message="任务轮询超时")
return
task.retry_count = (task.retry_count or 0) + 1
if task.retry_count > settings.CHATAPI_ASYNC_MAX_RETRIES:
@@ -339,7 +367,13 @@ async def _run(task_id: str):
if celery_app:
@celery_app.task(name="generation.poll_generation_task", bind=True, max_retries=3, default_retry_delay=30)
def poll_generation_task(self, task_id: str):
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):
@@ -2,16 +2,19 @@
from __future__ import annotations
import logging
from typing import Any, Dict
from typing import Any, Awaitable, Callable, Dict
from app.config import settings
from app.models.base import async_session
from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock
from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock
from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
logger = logging.getLogger("video_gen")
RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or "gen_recovery"
RecoveryRunner = Callable[[], Awaitable[Dict[str, Any]]]
async def _run_download_once() -> Dict[str, Any]:
from app.services.generation_recovery_service import recover_download_tasks_once
@@ -27,6 +30,55 @@ async def _run_generation_once() -> Dict[str, Any]:
return await recover_generation_tasks_once(db)
async def _run_module_async_once() -> Dict[str, Any]:
from app.services.module_async_recovery_service import recover_module_async_tasks_once
async with async_session() as db:
return await recover_module_async_tasks_once(db)
async def _run_shot_split_once() -> Dict[str, Any]:
from app.services.shot_replicate_recovery_service import recover_shot_split_tasks_once
async with async_session() as db:
return await recover_shot_split_tasks_once(db)
async def _run_with_execution_lock(
*,
lock_key: str,
log_context: str,
runner: RecoveryRunner,
) -> Dict[str, Any]:
"""恢复任务执行锁。
worker_ready 的启动锁只保证“只投递一次”;如果 broker 中残留旧消息,
或者人工手动触发恢复任务,仍可能并发执行。这里再加执行锁,避免多个
恢复扫描同时扫库、抢行锁、抢连接池。
"""
redis = await get_registry_redis()
token: str | None = None
if redis is not None:
token = await redis_acquire_lock(
lock_key=lock_key,
ttl_seconds=int(settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
log_context=log_context,
)
if not token:
return {"skipped": "lock_held", "lock_key": lock_key}
else:
# Redis 不可用时仍允许 DB fallback 执行一次,避免恢复能力彻底失效。
logger.warning("恢复任务执行锁不可用,降级直接执行。context=%s", log_context)
try:
result = await runner()
result["execution_lock"] = "lock_acquired" if token else "redis_unavailable_run_db_fallback"
return result
finally:
if token:
await redis_release_lock(lock_key=lock_key, token=token, log_context=log_context)
async def _acquire_download_recovery_loop_lock() -> tuple[bool, str]:
"""下载恢复循环锁。
@@ -45,36 +97,128 @@ async def _acquire_download_recovery_loop_lock() -> tuple[bool, str]:
def _schedule_next_download_recovery_loop() -> None:
if not celery_app or not bool(getattr(settings, "DOWNLOAD_RECOVERY_LOOP_ENABLED", True)):
if not celery_app or not bool(getattr(settings, "DOWNLOAD_RECOVERY_LOOP_ENABLED", False)):
return
try:
recover_download_tasks_once.apply_async(
countdown=max(1, int(settings.DOWNLOAD_RECOVERY_INTERVAL_SECONDS or 60)),
queue="gen_result_download",
queue=RECOVERY_QUEUE,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
except Exception:
logger.exception("下载恢复循环下一轮投递失败")
async def _run_startup_recovery_once() -> Dict[str, Any]:
"""启动容灾协调器:串行跑恢复扫描。
真实业务任务仍投递回原队列:
- 创建/提词/视频分析 -> gen_chatapi_create
- provider poll -> gen_provider_poll
- 下载/ffmpeg 切片 -> gen_result_download
恢复扫描本身只走 gen_recovery,避免堵住业务 worker。
"""
return await _run_with_execution_lock(
lock_key=settings.CELERY_RECOVERY_STARTUP_TASK_LOCK_KEY,
log_context="startup_recovery_once",
runner=_run_startup_recovery_steps,
)
async def _run_startup_recovery_steps() -> Dict[str, Any]:
results: Dict[str, Any] = {}
steps: list[tuple[str, str, str, RecoveryRunner]] = [
(
"module_async",
settings.MODULE_ASYNC_RECOVERY_LOCK_KEY,
"module_async_recovery",
_run_module_async_once,
),
(
"shot_split",
settings.SHOT_SPLIT_RECOVERY_LOCK_KEY,
"shot_split_recovery",
_run_shot_split_once,
),
(
"generation",
settings.GENERATION_RECOVERY_LOCK_KEY,
"generation_recovery",
_run_generation_once,
),
(
"download",
settings.DOWNLOAD_RECOVERY_LOCK_KEY,
"download_recovery",
_run_download_once,
),
]
for name, lock_key, log_context, runner in steps:
try:
results[name] = await _run_with_execution_lock(
lock_key=lock_key,
log_context=log_context,
runner=runner,
)
except Exception as exc:
logger.exception("启动容灾步骤执行失败。step=%s", name)
results[name] = {"error": str(exc)}
return {"steps": results}
if celery_app:
@celery_app.task(name="generation.recover_download_tasks_once", bind=True)
@celery_app.task(
name="recovery.startup_recovery_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def startup_recovery_once(self) -> Dict[str, Any]:
return run_async(_run_startup_recovery_once())
@celery_app.task(
name="generation.recover_download_tasks_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def recover_download_tasks_once(self) -> Dict[str, Any]:
acquired, reason = run_async(_acquire_download_recovery_loop_lock())
if not acquired:
return {"skipped": reason}
try:
result = run_async(_run_download_once())
result = run_async(
_run_with_execution_lock(
lock_key=settings.DOWNLOAD_RECOVERY_LOCK_KEY,
log_context="download_recovery",
runner=_run_download_once,
)
)
result["loop_lock"] = reason
return result
finally:
_schedule_next_download_recovery_loop()
@celery_app.task(name="generation.recover_generation_tasks_once")
def recover_generation_tasks_once() -> Dict[str, Any]:
return run_async(_run_generation_once())
@celery_app.task(
name="generation.recover_generation_tasks_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def recover_generation_tasks_once(self) -> Dict[str, Any]:
return run_async(
_run_with_execution_lock(
lock_key=settings.GENERATION_RECOVERY_LOCK_KEY,
log_context="generation_recovery",
runner=_run_generation_once,
)
)
else:
@@ -85,5 +229,6 @@ else:
def apply_async(self, *args: Any, **kwargs: Any) -> None:
raise RuntimeError("Celery is disabled")
startup_recovery_once = _DisabledTask()
recover_download_tasks_once = _DisabledTask()
recover_generation_tasks_once = _DisabledTask()
@@ -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]:
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:
return await recover_module_async_tasks_once(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:
@@ -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
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:
return await recover_shot_split_tasks_once(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: