"""GPU MuseTalk 口型同步服务 — 反向轮询模式. 职责: 1. 创建任务(由 lipsync 业务流程调用),为输入/输出生成预签名 URL,任务入队; 2. Worker 心跳注册(register):登记/刷新 worker 状态; 3. Worker 轮询拉任务(poll):原子地 CLAIM 一条 pending 任务,返回预签名 URL; 4. Worker 上报结果(report_result):标记 done/failed,失败可重试; 5. 业务侧查询状态(get_status)。 """ from __future__ import annotations import logging import uuid from datetime import UTC, datetime, timedelta from typing import Optional from app.core.storage import get_storage_service from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import GpuLipsyncTaskModel, GpuWorkerModel from packages.config import get_api_settings logger = logging.getLogger(__name__) # 任务在 processing 超过此时长仍未完成 → 超时回退 pending 或置 failed MAX_ATTEMPTS = 3 class GpuLipsyncService: """GPU 口型同步服务(无状态方法,每次调用从 DI 拿 db/storage).""" RESULT_PREFIX = "gpu-lipsync/results/" INPUT_SIGN_EXPIRES_PAD = 600 # 输入预签名 URL 在任务超时基础上再加 10min 余量 # ── 公共入口 ──────────────────────────────────────────────────── def __init__(self, db: Session): self.db = db self.settings = get_api_settings() self.storage = get_storage_service() # ── Worker 注册/心跳 ──────────────────────────────────────────── def register_worker( self, worker_id: str, hostname: str = "", gpu_name: str = "", free_vram_mb: int = 0, capabilities: str = "musetalk", task_id: Optional[str] = None, ) -> tuple[GpuWorkerModel, bool]: """Worker 注册/心跳。 task_id 非空时(Worker 推理期间的任务级心跳),同步把对应 processing 任务的 last_heartbeat_at 续到当前时间,使长推理不会被 ``_recover_timed_out_tasks`` 误回退。任务已结束 / 不属于该 worker (如已被超时回收重新派发)时忽略,不报错。 返回 ``(worker, cancel_task)``:当心跳任务已被用户取消时 ``cancel_task=True``,Worker 应尽快终止推理并释放 GPU。 """ now = datetime.now(UTC) worker = self.db.query(GpuWorkerModel).filter(GpuWorkerModel.worker_id == worker_id).one_or_none() if worker is None: worker = GpuWorkerModel( worker_id=worker_id, hostname=hostname, gpu_name=gpu_name, free_vram_mb=free_vram_mb, capabilities=capabilities, last_heartbeat_at=now, created_at=now, ) self.db.add(worker) else: worker.hostname = hostname or worker.hostname worker.gpu_name = gpu_name or worker.gpu_name worker.free_vram_mb = free_vram_mb worker.capabilities = capabilities or worker.capabilities worker.last_heartbeat_at = now cancel_task = False if task_id: cancel_task = self._touch_task_heartbeat(task_id, worker_id, now) self.db.commit() return worker, cancel_task # ── 轮询拉任务(Worker 调用) ────────────────────────────────── def poll_task(self, worker_id: str) -> Optional[GpuLipsyncTaskModel]: """原子地认领一条最早的 pending 任务,返回给 worker;无任务返回 None. 同时会: - 把 processing 状态且真正超时(任务心跳停滞超过 gpu_task_timeout_seconds;Worker 推理期会通过 register(task_id=...) 续心跳,长推理不会误判)的任务回退为 pending(attempt++,超过 MAX_ATTEMPTS 置 failed),让其它 worker 认领。 - 刷新 worker 心跳。 """ now = datetime.now(UTC) self._recover_timed_out_tasks(now) # 更新 worker 心跳 self._touch_worker(worker_id, now) # 选一条最早 pending 任务(FOR UPDATE SKIP LOCKED 语义:简单起见先查再锁状态) task = ( self.db.query(GpuLipsyncTaskModel) .filter(GpuLipsyncTaskModel.status == "pending") .order_by(GpuLipsyncTaskModel.created_at.asc()) .first() ) if task is None: self.db.commit() return None # 原子 claim:用 UPDATE WHERE status=pending 避免并发 upd_rows = ( self.db.query(GpuLipsyncTaskModel) .filter( GpuLipsyncTaskModel.id == task.id, GpuLipsyncTaskModel.status == "pending", ) .update( { GpuLipsyncTaskModel.status: "processing", GpuLipsyncTaskModel.worker_id: worker_id, GpuLipsyncTaskModel.started_at: now, GpuLipsyncTaskModel.last_heartbeat_at: now, GpuLipsyncTaskModel.attempt: GpuLipsyncTaskModel.attempt + 1, GpuLipsyncTaskModel.updated_at: now, }, synchronize_session=False, ) ) self.db.commit() if upd_rows == 0: # 被其它 worker 抢先了 return None self.db.refresh(task) # 生成预签名输入/输出 URL(在 claim 时动态生成,避免长时间过期) expires = self.settings.gpu_task_timeout_seconds + self.INPUT_SIGN_EXPIRES_PAD task._signed_video_url = self.storage.get_download_url(task.video_url, expires_seconds=expires) task._signed_audio_url = self.storage.get_download_url(task.audio_url, expires_seconds=expires) task._signed_upload_url = self.storage.get_upload_url( self._result_key(task.id), expires_seconds=expires, content_type="video/mp4", ) task._upload_expires_at = now + timedelta(seconds=expires) return task # ── 上报结果 ────────────────────────────────────────────────── def report_result( self, task_id: str, worker_id: str, success: bool, duration_seconds: float = 0.0, error_msg: str = "", ) -> GpuLipsyncTaskModel: task = self.db.get(GpuLipsyncTaskModel, task_id) if task is None: raise KeyError(f"task {task_id} not found") now = datetime.now(UTC) if success: task.status = "done" task.result_url = self._result_key(task_id) task.result_duration = duration_seconds or 0.0 task.error_msg = "" task.finished_at = now elif task.status == "cancelled": # 用户已取消的任务,Worker 终止后上报失败,保持 cancelled 状态不回退 task.finished_at = now task.error_msg = (error_msg or "用户取消")[:2000] logger.info("GPU 任务 %s 已被用户取消,保持 cancelled 状态", task_id) else: # 失败:若仍可重试(已尝试次数 < MAX_ATTEMPTS)→ 回退 pending;否则 → failed if task.attempt < MAX_ATTEMPTS: task.status = "pending" task.worker_id = "" task.started_at = None task.error_msg = error_msg[:2000] logger.warning( "GPU 任务 %s 在 worker %s 上失败,回退 pending 等待重试(attempt=%d): %s", task_id, worker_id, task.attempt, error_msg[:200], ) else: task.status = "failed" task.error_msg = error_msg[:2000] task.finished_at = now logger.error( "GPU 任务 %s 失败达到最大重试次数 %d,置为 failed: %s", task_id, MAX_ATTEMPTS, error_msg[:200], ) task.updated_at = now task.last_heartbeat_at = now self._touch_worker(worker_id, now) self.db.commit() self.db.refresh(task) return task # ── 业务侧查询 ──────────────────────────────────────────────── def get_task(self, task_id: str) -> Optional[GpuLipsyncTaskModel]: return self.db.get(GpuLipsyncTaskModel, task_id) def get_by_lipsync_job(self, lipsync_job_id: str) -> Optional[GpuLipsyncTaskModel]: return ( self.db.query(GpuLipsyncTaskModel) .filter(GpuLipsyncTaskModel.lipsync_job_id == lipsync_job_id) .order_by(GpuLipsyncTaskModel.created_at.desc()) .first() ) # ── 创建任务(业务侧调用) ──────────────────────────────────── def create_task( self, video_url: str, audio_url: str, lipsync_job_id: str = "", user_id: str = "", project_id: str = "", ) -> GpuLipsyncTaskModel: task_id = str(uuid.uuid4()) now = datetime.now(UTC) task = GpuLipsyncTaskModel( id=task_id, lipsync_job_id=lipsync_job_id, user_id=user_id, project_id=project_id, video_url=video_url, audio_url=audio_url, status="pending", attempt=0, created_at=now, updated_at=now, ) self.db.add(task) self.db.commit() self.db.refresh(task) logger.info( "创建 GPU 口型任务 %s (lipsync_job=%s, user=%s)", task_id, lipsync_job_id, user_id, ) return task # ── 内部辅助 ────────────────────────────────────────────────── def _result_key(self, task_id: str) -> str: return f"{self.RESULT_PREFIX}{task_id}.mp4" def _touch_task_heartbeat(self, task_id: str, worker_id: str, now: datetime) -> bool: """Worker 推理期间的任务级心跳:只刷新属于该 worker 且仍在 processing 的任务。 任务不存在 / 已被超时回收重新派发 / 已完成 → 静默忽略(此时旧 worker 的 结果上报会被结果接口按最终态处理)。 返回 ``cancel_task``:任务已被用户取消时为 True,Worker 应终止推理。 """ task = self.db.get(GpuLipsyncTaskModel, task_id) if task is None: return False # 任务已被用户取消 → 通知 Worker 终止推理 if task.status == "cancelled": logger.info("任务心跳检测到已取消 task=%s worker=%s,通知 Worker 终止", task_id, worker_id) return True if task.status != "processing" or task.worker_id != worker_id: logger.info( "忽略过期任务心跳 task=%s worker=%s(status=%s owner=%s)", task_id, worker_id, task.status, task.worker_id, ) return False task.last_heartbeat_at = now task.updated_at = now self.db.flush() return False def _touch_worker(self, worker_id: str, now: datetime) -> None: if not worker_id: return worker = self.db.query(GpuWorkerModel).filter(GpuWorkerModel.worker_id == worker_id).one_or_none() if worker is not None: worker.last_heartbeat_at = now self.db.flush() else: # 自注册(poll 时允许自动建一个空 worker 记录,运维可见) worker = GpuWorkerModel( worker_id=worker_id, hostname="", gpu_name="", free_vram_mb=0, capabilities="musetalk", last_heartbeat_at=now, created_at=now, ) self.db.add(worker) self.db.flush() def _recover_timed_out_tasks(self, now: datetime) -> None: """扫描 processing 状态且真正超时的任务,回退 pending 或失败。 判定只看任务自身 last_heartbeat_at:claim 时写入,Worker 推理期间通过 /gpu/register(task_id=...) 每 30s 续期。因此仅在 Worker 崩溃/断网 (任务心跳停滞超过 gpu_task_timeout_seconds)时才回收, 不会因 Worker 主循环忙于推理而误回退。 """ timeout = self.settings.gpu_task_timeout_seconds cutoff = now - timedelta(seconds=timeout) stuck_tasks = ( self.db.query(GpuLipsyncTaskModel) .filter( GpuLipsyncTaskModel.status == "processing", GpuLipsyncTaskModel.last_heartbeat_at < cutoff, ) .all() ) for t in stuck_tasks: if t.attempt >= MAX_ATTEMPTS: t.status = "failed" t.error_msg = f"worker 心跳超时({timeout}s),重试次数已耗尽" t.finished_at = now else: t.status = "pending" t.worker_id = "" t.started_at = None t.error_msg = f"worker 心跳超时({timeout}s),等待重试" logger.warning("GPU 任务 %s 心跳超时,回退 pending(attempt=%d)", t.id, t.attempt) t.updated_at = now if stuck_tasks: self.db.flush() # ── 业务侧辅助 ────────────────────────────────────────────────── def has_available_worker(self) -> bool: """判断是否有 Worker 在心跳新鲜窗口内可用.""" stale_cutoff = datetime.now(UTC) - timedelta(seconds=self.settings.gpu_worker_stale_seconds) return ( self.db.query(GpuWorkerModel).filter(GpuWorkerModel.last_heartbeat_at >= stale_cutoff).first() is not None ) def wait_for_result( self, task_id: str, timeout_seconds: Optional[int] = None, poll_interval: Optional[float] = None, ) -> Optional[GpuLipsyncTaskModel]: """同步轮询等待 GPU 任务完成。 Args: task_id: 任务 ID(由 create_task 返回) timeout_seconds: 总超时,默认取 settings.gpu_lipsync_wait_timeout poll_interval: 轮询间隔秒,默认取 settings.gpu_lipsync_poll_interval Returns: 终态 task(status=done/failed);超时返回 None(此时调用方应回退 MediaKit)。 等待期间会自动调用 _recover_timed_out_tasks 做超时回收。 """ import time timeout = timeout_seconds if timeout_seconds is not None else self.settings.gpu_lipsync_wait_timeout interval = poll_interval if poll_interval is not None else self.settings.gpu_lipsync_poll_interval deadline = time.monotonic() + timeout while True: now = datetime.now(UTC) # 顺手回收超时任务 try: self._recover_timed_out_tasks(now) self.db.commit() except Exception as exc: # noqa: BLE001 - 回收失败不阻塞主流程 logger.warning("wait_for_result 回收超时任务异常: %s", exc) self.db.rollback() task = self.db.get(GpuLipsyncTaskModel, task_id) if task is None: return None if task.status in ("done", "failed", "cancelled"): return task # pending/processing 继续等 if time.monotonic() >= deadline: logger.warning("GPU 任务 %s 等待超时(%ds),回退 MediaKit", task_id, timeout) return None time.sleep(interval)