diff --git a/apps/api/app/api/routes/gpu_lipsync.py b/apps/api/app/api/routes/gpu_lipsync.py index c00c73ccf..85a7c97a9 100644 --- a/apps/api/app/api/routes/gpu_lipsync.py +++ b/apps/api/app/api/routes/gpu_lipsync.py @@ -93,7 +93,7 @@ def register_worker( svc: GpuLipsyncService = Depends(_get_svc), _token: str = Depends(_verify_gpu_token), ): - svc.register_worker( + worker, cancel_task = svc.register_worker( worker_id=body.worker_id, hostname=body.hostname, gpu_name=body.gpu_name, @@ -101,7 +101,7 @@ def register_worker( capabilities=body.capabilities, task_id=body.task_id, ) - return GpuWorkerRegisterResponse(ok=True, server_time=datetime.now(UTC), message="ok") + return GpuWorkerRegisterResponse(ok=True, server_time=datetime.now(UTC), message="ok", cancel_task=cancel_task) # ── GET /lipsync/poll — Worker 轮询拉任务 ───────────────────────── diff --git a/apps/api/app/api/routes/lipsync.py b/apps/api/app/api/routes/lipsync.py index 12184cf85..6e10608d3 100644 --- a/apps/api/app/api/routes/lipsync.py +++ b/apps/api/app/api/routes/lipsync.py @@ -346,13 +346,13 @@ def cancel_lipsync_job( current_user: AuthenticatedUser = Depends(get_current_user), svc: LipsyncService = Depends(_get_service), ): - """取消对口型任务(仅 pending/tts_processing/submitted 状态可取消).""" + """取消对口型任务(仅 pending/tts_processing/submitted/processing 状态可取消).""" job = svc.cancel_job(job_id, current_user.user.id) if job is None: raise HTTPException(status_code=404, detail="任务不存在") if job.status != "cancelled": raise HTTPException( status_code=400, - detail=f"任务状态 {job.status} 不可取消,仅 pending/tts_processing/submitted 可取消", + detail=f"任务状态 {job.status} 不可取消,仅 pending/tts_processing/submitted/processing 可取消", ) return job diff --git a/apps/api/app/schemas/gpu_lipsync.py b/apps/api/app/schemas/gpu_lipsync.py index fbd26de1c..b9d391ef9 100644 --- a/apps/api/app/schemas/gpu_lipsync.py +++ b/apps/api/app/schemas/gpu_lipsync.py @@ -36,6 +36,7 @@ class GpuWorkerRegisterResponse(BaseModel): ok: bool = True server_time: datetime message: str = "ok" + cancel_task: bool = Field(False, description="当前心跳任务是否已被用户取消;为 true 时 Worker 应终止推理") # ── 轮询任务 ──────────────────────────────────────────────────── diff --git a/apps/api/app/services/gpu_lipsync_service.py b/apps/api/app/services/gpu_lipsync_service.py index 485285dbd..7c8286231 100644 --- a/apps/api/app/services/gpu_lipsync_service.py +++ b/apps/api/app/services/gpu_lipsync_service.py @@ -50,13 +50,16 @@ class GpuLipsyncService: free_vram_mb: int = 0, capabilities: str = "musetalk", task_id: Optional[str] = None, - ) -> GpuWorkerModel: + ) -> 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() @@ -77,10 +80,11 @@ class GpuLipsyncService: worker.free_vram_mb = free_vram_mb worker.capabilities = capabilities or worker.capabilities worker.last_heartbeat_at = now + cancel_task = False if task_id: - self._touch_task_heartbeat(task_id, worker_id, now) + cancel_task = self._touch_task_heartbeat(task_id, worker_id, now) self.db.commit() - return worker + return worker, cancel_task # ── 轮询拉任务(Worker 调用) ────────────────────────────────── @@ -166,6 +170,11 @@ class GpuLipsyncService: 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: @@ -250,15 +259,21 @@ class GpuLipsyncService: 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) -> None: + 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 + 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)", @@ -267,10 +282,11 @@ class GpuLipsyncService: task.status, task.worker_id, ) - return + 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: @@ -371,9 +387,7 @@ class GpuLipsyncService: task = self.db.get(GpuLipsyncTaskModel, task_id) if task is None: return None - if task.status == "done": - return task - if task.status == "failed": + if task.status in ("done", "failed", "cancelled"): return task # pending/processing 继续等 if time.monotonic() >= deadline: diff --git a/apps/api/app/services/lipsync_service.py b/apps/api/app/services/lipsync_service.py index 362e12f59..7d453271a 100644 --- a/apps/api/app/services/lipsync_service.py +++ b/apps/api/app/services/lipsync_service.py @@ -783,12 +783,31 @@ class LipsyncService: # ── 取消任务 ────────────────────────────────────────────────────────── def cancel_job(self, job_id: str, user_id: str) -> Optional[LipsyncJobModel]: - """取消任务(仅 pending/tts_processing/submitted 状态可取消).""" + """取消任务(pending/tts_processing/submitted/processing 状态可取消). + + 当 job 走 GPU 路径(mediakit_task_id 以 "gpu:" 开头)且状态为 processing 时, + 同步将关联的 GpuLipsyncTask 标记为 cancelled,以便 Worker 心跳时检测到取消信号。 + """ job = self.get_job(job_id, user_id) if job is None: return None - if job.status in ("pending", "tts_processing", "submitted"): + if job.status in ("pending", "tts_processing", "submitted", "processing"): + # GPU 路径:同步标记关联的 GPU 任务为 cancelled + if job.status == "processing" and job.mediakit_task_id and job.mediakit_task_id.startswith("gpu:"): + gpu_task_id = job.mediakit_task_id[4:] # 去掉 "gpu:" 前缀 + try: + from packages.adapters.sqlalchemy_impl.models import GpuLipsyncTaskModel + gpu_task = self.db.get(GpuLipsyncTaskModel, gpu_task_id) + if gpu_task and gpu_task.status == "processing": + gpu_task.status = "cancelled" + gpu_task.error_msg = "用户取消" + gpu_task.updated_at = datetime.now(UTC) + gpu_task.finished_at = datetime.now(UTC) + logger.info("GPU 任务 %s 已被用户取消(通过 job_id=%s)", gpu_task_id, job_id) + except Exception as exc: + logger.warning("标记 GPU 任务取消失败(不影响 job 取消): %s", exc) + job.status = "cancelled" job.updated_at = datetime.now(UTC) self.db.commit() diff --git a/apps/api/app/tasks/lipsync_gpu.py b/apps/api/app/tasks/lipsync_gpu.py index 72e3511da..afbb89028 100644 --- a/apps/api/app/tasks/lipsync_gpu.py +++ b/apps/api/app/tasks/lipsync_gpu.py @@ -98,6 +98,14 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str) _fallback_to_mediakit(db, job) return + if final_task.status == "cancelled": + # 用户已取消任务,不回退 MediaKit,直接标记 job 为 cancelled + job.status = "cancelled" + job.updated_at = datetime.now(UTC) + db.commit() + logger.info("[lipsync_gpu_async] GPU 任务已被用户取消: job_id=%s", job_id) + return + if final_task.status != "done": logger.warning( "[lipsync_gpu_async] GPU 失败,回退 MediaKit: job_id=%s gpu_task=%s status=%s", diff --git a/deploy/gpu_worker/gpu_worker.py b/deploy/gpu_worker/gpu_worker.py index d9d53f9b3..6a48349b5 100644 --- a/deploy/gpu_worker/gpu_worker.py +++ b/deploy/gpu_worker/gpu_worker.py @@ -108,11 +108,14 @@ def _check_musetalk_health() -> tuple[bool, dict]: return False, {"error": str(exc)} -def _register(task_id: Optional[str] = None) -> bool: +def _register(task_id: Optional[str] = None) -> tuple[bool, bool]: """向服务端注册 / 心跳,附带 GPU 信息。 推理期间的心跳线程传 task_id:服务端会同步刷新该 processing 任务的 - last_heartbeat_at,防止长推理被误判超时回收。 + last_heartbeat_at,防止长推理被误判超时回收。同时服务端会检查该任务 + 是否已被用户取消,若是则返回 cancel_task=True。 + + 返回 (ok, cancel_task)。 """ ok, info = _check_musetalk_health() if isinstance(info, dict): @@ -142,12 +145,14 @@ def _register(task_id: Optional[str] = None) -> bool: timeout=15, ) if r.status_code == 200: - return True + resp_body = r.json() + cancel_task = resp_body.get("cancel_task", False) + return True, cancel_task logger.error("注册/心跳失败: HTTP %d body=%s", r.status_code, r.text[:300]) - return False + return False, False except Exception as exc: logger.error("注册/心跳异常: %s", exc) - return False + return False, False def _probe_gpu_name() -> str: @@ -329,6 +334,9 @@ class TaskHeartbeat(threading.Thread): 期间无法发送,服务端会因任务 last_heartbeat_at 停滞而误判超时回退 pending。 本线程每 task_heartbeat_interval 秒(默认 30s)POST /gpu/register 并 携带当前 task_id,让服务端持续续期任务心跳;任务处理结束 stop()。 + + 同时检测服务端返回的 cancel_task 信号:若为 True,说明用户已取消任务, + 立即调用 _cancel_musetalk() 终止本地推理,并设置 cancelled 标志供主流程检查。 """ def __init__(self, task_id: str, interval: float): @@ -336,13 +344,21 @@ class TaskHeartbeat(threading.Thread): self.task_id = task_id self.interval = max(5.0, interval) self._stop_event = threading.Event() + self.cancelled = False # 外部可读的取消标志 def run(self) -> None: # 先立即发一次,再按间隔循环(首次心跳失败不影响主流程) while not self._stop_event.is_set(): try: - if _register(self.task_id): + ok, cancel_task = _register(self.task_id) + if ok: logger.debug("任务 %s 心跳已发送", self.task_id) + if cancel_task: + logger.warning("任务 %s 已被用户取消,正在终止本地推理...", self.task_id) + self.cancelled = True + _cancel_musetalk() + self._stop_event.set() + return except Exception as exc: # noqa: BLE001 logger.warning("任务 %s 心跳异常(忽略): %s", self.task_id, exc) self._stop_event.wait(self.interval) @@ -369,9 +385,17 @@ def _handle_task(task: dict) -> None: if not _download(task["video_url"], video_path): _report_result(task_id, False, 0.0, "下载人物视频失败") return + if hb.cancelled: + logger.info("任务 %s 在下载阶段被用户取消", task_id) + _report_result(task_id, False, 0.0, "用户取消任务") + return if not _download(task["audio_url"], audio_path): _report_result(task_id, False, 0.0, "下载驱动音频失败") return + if hb.cancelled: + logger.info("任务 %s 在下载阶段被用户取消", task_id) + _report_result(task_id, False, 0.0, "用户取消任务") + return # 2. 输入时长前置校验:短视频 MuseTalk 会 division by zero, # 直接上报 failed,不浪费 GPU 时间。ffprobe 不可用/读失败(0.0) @@ -392,12 +416,20 @@ def _handle_task(task: dict) -> None: err = "" retryable = False for attempt in range(Config.task_max_retry + 1): + if hb.cancelled: + logger.info("任务 %s 在推理前被用户取消", task_id) + _report_result(task_id, False, 0.0, "用户取消任务") + return if attempt > 0: logger.info("任务 %s 第 %d 次重试(瞬时错误)...", task_id, attempt + 1) time.sleep(2) success, duration, err, retryable = _call_musetalk(video_path, audio_path, out_path) if success or not retryable: break + if hb.cancelled: + logger.info("任务 %s 被用户取消(推理已终止)", task_id) + _report_result(task_id, False, 0.0, "用户取消任务") + return if not success: logger.error("任务 %s 推理失败: %s", task_id, err) _report_result(task_id, False, 0.0, err) @@ -467,7 +499,8 @@ def main() -> int: # 心跳 now = time.time() if now - last_heartbeat >= Config.heartbeat_interval: - if _register(): + ok, _ = _register() + if ok: last_heartbeat = now # 轮询任务