celery异步恢复任务独立队列|扩增celery子进程连接池上限配置
This commit is contained in:
+21
-12
@@ -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、无供应商任务 ID:deadline 未过才恢复 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,
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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 active,Celery 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:
|
||||
|
||||
Reference in New Issue
Block a user