diff --git a/video-gen-api/app/tasks/celery_app.py b/video-gen-api/app/tasks/celery_app.py index 93184ed6..f7e87e71 100644 --- a/video-gen-api/app/tasks/celery_app.py +++ b/video-gen-api/app/tasks/celery_app.py @@ -57,8 +57,103 @@ def _derive_redis_db(url: str, db_no: int) -> str: return url.rstrip("/") + f"/{db_no}" +def _parse_schedule_to_celery(schedule_str: str): + """将 schedule 字符串解析为 Celery 可识别的调度值。 + + - 纯数字:视为间隔秒数(返回 int) + - cron 表达式 (5 字段空格分隔):返回 crontab 对象 + """ + from celery.schedules import crontab + + s = (schedule_str or "").strip() + if not s: + return None + # 纯数字 → 间隔秒数 + if s.isdigit(): + return int(s) + # cron 表达式 (分 时 日 月 周) + parts = s.split() + if len(parts) == 5: + try: + return crontab( + minute=parts[0], + hour=parts[1], + day_of_month=parts[2], + month_of_year=parts[3], + day_of_week=parts[4], + ) + except Exception: + logger.exception("解析 cron 表达式失败: %s", s) + return None + logger.warning("无法解析 schedule 表达式: %s", s) + return None + + +# 同步引擎(用于 Beat 启动时加载数据库定时任务) +_sync_engine = None + + +def _get_sync_engine(): + global _sync_engine + if _sync_engine is None: + from sqlalchemy import create_engine + from app.config import settings + + db_url = settings.DATABASE_URL + # 异步 URL → 同步 URL(asyncpg → psycopg2 / aiosqlite → sqlite3) + if db_url.startswith("postgresql+asyncpg"): + db_url = db_url.replace("postgresql+asyncpg", "postgresql+psycopg2", 1) + elif db_url.startswith("sqlite+aiosqlite"): + db_url = db_url.replace("sqlite+aiosqlite", "sqlite", 1) + elif db_url.startswith("mysql+aiomysql"): + db_url = db_url.replace("mysql+aiomysql", "mysql+pymysql", 1) + + _sync_engine = create_engine(db_url, pool_pre_ping=True) + return _sync_engine + + +def _load_dynamic_beat_tasks() -> dict: + """从数据库加载活跃定时任务,返回 beat_schedule 格式的字典。 + + 在 Celery 配置阶段同步调用,确保 Beat 启动时能读取到动态任务。 + """ + from sqlalchemy import select + from sqlalchemy.orm import Session + + from app.models.scheduled_task import ScheduledTask + + dynamic_schedule = {} + try: + engine = _get_sync_engine() + with Session(engine) as session: + result = session.execute( + select(ScheduledTask).where(ScheduledTask.is_active.is_(True)) + ) + tasks = result.scalars().all() + + for task in tasks: + schedule_val = _parse_schedule_to_celery(task.schedule) + if schedule_val is None: + logger.warning("定时任务 %s schedule 无效,跳过注册: %s", task.id, task.schedule) + continue + beat_key = f"dynamic-scheduled-task-{task.id}" + dynamic_schedule[beat_key] = { + "task": "execute_scheduled_task", + "schedule": schedule_val, + "args": (task.id,), + "options": {"queue": RECOVERY_QUEUE}, + } + logger.info( + "动态注册定时任务到 Beat: %s (%s) schedule=%s", + task.name, task.id, task.schedule, + ) + except Exception: + logger.exception("加载动态定时任务失败,跳过") + return dynamic_schedule + + def _beat_schedule() -> dict: - schedule: dict = {} + schedule: dict = _load_dynamic_beat_tasks() if bool(getattr(settings, "POLL_DUE_DISPATCH_ENABLED", True)): schedule["dispatch-due-poll-tasks-every-minute"] = { "task": CeleryTaskName.DISPATCH_DUE_POLL.value, @@ -293,122 +388,6 @@ else: celery_app = None -def _parse_schedule_to_celery(schedule_str: str): - """将 schedule 字符串解析为 Celery 可识别的调度值。 - - - 纯数字:视为间隔秒数(返回 int) - - cron 表达式 (5 字段空格分隔):返回 crontab 对象 - """ - from celery.schedules import crontab - - s = (schedule_str or "").strip() - if not s: - return None - # 纯数字 → 间隔秒数 - if s.isdigit(): - return int(s) - # cron 表达式 (分 时 日 月 周) - parts = s.split() - if len(parts) == 5: - try: - return crontab( - minute=parts[0], - hour=parts[1], - day_of_month=parts[2], - month_of_year=parts[3], - day_of_week=parts[4], - ) - except Exception: - logger.exception("解析 cron 表达式失败: %s", s) - return None - logger.warning("无法解析 schedule 表达式: %s", s) - return None - - -def _setup_dynamic_beat_tasks(sender, **kwargs): - """从数据库加载活跃定时任务并注册到 Beat 调度。 - - 通过 @celery_app.on_after_configure.connect 在 Celery 配置完成后执行, - 适用于 Worker 和 Beat 启动场景。 - - 注意:此信号可能运行在已有 event loop 的上下文中(Celery Beat 主 loop), - 不能使用 run_until_complete 或 run_async,改用同步数据库查询。 - """ - if celery_app is None: - return - - try: - active_tasks = _load_active_tasks_sync() - except Exception: - logger.exception("加载定时任务失败,跳过动态 Beat 注册") - return - - for task in active_tasks: - schedule_val = _parse_schedule_to_celery(task.schedule) - if schedule_val is None: - logger.warning("定时任务 %s schedule 无效,跳过注册: %s", task.id, task.schedule) - continue - beat_key = f"dynamic-scheduled-task-{task.id}" - sender.conf.beat_schedule[beat_key] = { - "task": "execute_scheduled_task", - "schedule": schedule_val, - "args": (task.id,), - "options": {"queue": RECOVERY_QUEUE}, - } - logger.info( - "动态注册定时任务到 Beat: %s (%s) schedule=%s", - task.name, task.id, task.schedule, - ) - - -# 同步引擎(仅在 Beat 动态注册定时任务时使用) -_sync_engine = None - - -def _get_sync_engine(): - global _sync_engine - if _sync_engine is None: - from sqlalchemy import create_engine - from app.config import settings - - db_url = settings.DATABASE_URL - # 异步 URL → 同步 URL(asyncpg → psycopg2 / aiosqlite → sqlite3) - if db_url.startswith("postgresql+asyncpg"): - db_url = db_url.replace("postgresql+asyncpg", "postgresql+psycopg2", 1) - elif db_url.startswith("sqlite+aiosqlite"): - db_url = db_url.replace("sqlite+aiosqlite", "sqlite", 1) - elif db_url.startswith("mysql+aiomysql"): - db_url = db_url.replace("mysql+aiomysql", "mysql+pymysql", 1) - - _sync_engine = create_engine(db_url, pool_pre_ping=True) - return _sync_engine - - -def _load_active_tasks_sync(): - """同步查询活跃的定时任务。""" - from sqlalchemy import select - from sqlalchemy.orm import Session - - from app.models.scheduled_task import ScheduledTask - - engine = _get_sync_engine() - with Session(engine) as session: - result = session.execute( - select(ScheduledTask).where(ScheduledTask.is_active.is_(True)) - ) - # detach objects so they can be used after session closes - tasks = result.scalars().all() - # Expunge all to detach from session - for t in tasks: - session.expunge(t) - return tasks - - -# 注册信号(放在函数定义之后,确保函数已存在) -if celery_app is not None: - celery_app.on_after_configure.connect(_setup_dynamic_beat_tasks) - - async def _try_acquire_startup_recovery_lock() -> bool: """任意 worker 启动时都可尝试抢恢复投递锁,避免依赖 hostname 命名。""" from app.services.redis_registry_service import redis_acquire_lock