From f55ea0100de6d3b2760a6d073ed11fb4b784213d Mon Sep 17 00:00:00 2001 From: saas-backend Date: Mon, 21 Sep 2026 22:09:02 +0800 Subject: [PATCH 1/6] =?UTF-8?q?fix:=20=E4=BF=AE=E5=A4=8D=E5=8F=96=E6=B6=88?= =?UTF-8?q?=E9=93=BE=E8=B7=AF=E6=96=AD=E8=A3=82=EF=BC=8C=E5=89=8D=E7=AB=AF?= =?UTF-8?q?=E5=8F=96=E6=B6=88=E5=90=8E=20GPU=20=E4=BB=8D=E7=BB=A7=E7=BB=AD?= =?UTF-8?q?=E6=8E=A8=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - API 心跳接口 POST /gpu/register 响应增加 cancel_task 字段 - 心跳时检测 GPU 任务是否已被用户取消,通知 Worker 终止推理 - cancel_job 支持 processing 状态,同步标记 GPU 任务为 cancelled - gpu_worker.py TaskHeartbeat 读取心跳响应,检测到取消时调用 /cancel - _handle_task 各阶段检查 cancelled 标志,取消时上报失败而非重试 - report_result 遇到 cancelled 状态的任务保持不变,不回退 pending - wait_for_result 将 cancelled 视为终态,Celery 任务不回退 MediaKit --- apps/api/app/api/routes/gpu_lipsync.py | 4 +- apps/api/app/api/routes/lipsync.py | 4 +- apps/api/app/schemas/gpu_lipsync.py | 1 + apps/api/app/services/gpu_lipsync_service.py | 32 +++++++++---- apps/api/app/services/lipsync_service.py | 23 +++++++++- apps/api/app/tasks/lipsync_gpu.py | 8 ++++ deploy/gpu_worker/gpu_worker.py | 47 +++++++++++++++++--- 7 files changed, 97 insertions(+), 22 deletions(-) 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 # 轮询任务 -- 2.54.0 From 3911050a34596390489fd6d5ac83f081af5e5224 Mon Sep 17 00:00:00 2001 From: CI Bot Date: Mon, 21 Sep 2026 14:27:12 +0000 Subject: [PATCH 2/6] style: auto-format with black + isort + ruff + prettier [skip ci-format-check] --- apps/api/app/services/lipsync_service.py | 1 + 1 file changed, 1 insertion(+) diff --git a/apps/api/app/services/lipsync_service.py b/apps/api/app/services/lipsync_service.py index 7d453271a..f4a248514 100644 --- a/apps/api/app/services/lipsync_service.py +++ b/apps/api/app/services/lipsync_service.py @@ -798,6 +798,7 @@ class LipsyncService: 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" -- 2.54.0 From 828c6703c1167bc20fd8840d96053c1a185b58f5 Mon Sep 17 00:00:00 2001 From: saas-backend Date: Mon, 21 Sep 2026 23:45:49 +0800 Subject: [PATCH 3/6] =?UTF-8?q?test:=20=E9=80=82=E9=85=8D=20=5Fregister/re?= =?UTF-8?q?gister=5Fworker=20=E6=96=B0=E8=BF=94=E5=9B=9E=E5=80=BC=E7=AD=BE?= =?UTF-8?q?=E5=90=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 取消链路修复后 _register 返回 (ok, cancel_task), register_worker 返回 (worker, cancel_task),同步更新单测断言。 --- tests/unit/test_1970_gpu_worker_heartbeat.py | 8 ++++++-- tests/unit/test_gpu_lipsync_service.py | 12 ++++++------ 2 files changed, 12 insertions(+), 8 deletions(-) diff --git a/tests/unit/test_1970_gpu_worker_heartbeat.py b/tests/unit/test_1970_gpu_worker_heartbeat.py index 082149a45..005723e8c 100644 --- a/tests/unit/test_1970_gpu_worker_heartbeat.py +++ b/tests/unit/test_1970_gpu_worker_heartbeat.py @@ -68,7 +68,9 @@ def test_register_payload_includes_task_id_only_when_provided(worker, monkeypatc class _Resp: status_code = 200 - text = "" + + def json(self): + return {"worker_id": captured[-1]["worker_id"], "cancel_task": False} def _fake_post(url, json=None, headers=None, timeout=None): captured.append(json) @@ -77,7 +79,9 @@ def test_register_payload_includes_task_id_only_when_provided(worker, monkeypatc monkeypatch.setattr(worker.requests, "post", _fake_post) monkeypatch.setattr(worker, "_check_musetalk_health", lambda: (True, {})) - assert worker._register("task-abc") is True + _ok, _cancel = worker._register("task-abc") + assert _ok is True + assert _cancel is False assert captured[-1]["task_id"] == "task-abc" assert captured[-1]["worker_id"] diff --git a/tests/unit/test_gpu_lipsync_service.py b/tests/unit/test_gpu_lipsync_service.py index 4d9017e23..6e2e59381 100644 --- a/tests/unit/test_gpu_lipsync_service.py +++ b/tests/unit/test_gpu_lipsync_service.py @@ -173,10 +173,10 @@ def test_timed_out_task_is_redispatched(svc): def test_register_worker_creates_then_updates(svc): - w = svc.register_worker("w-1", hostname="pc1", gpu_name="RTX2060", free_vram_mb=3500) + w, _cancel = svc.register_worker("w-1", hostname="pc1", gpu_name="RTX2060", free_vram_mb=3500) assert w.worker_id == "w-1" assert w.gpu_name == "RTX2060" - w2 = svc.register_worker("w-1", hostname="pc1", gpu_name="RTX2060", free_vram_mb=2000) + w2, _cancel2 = svc.register_worker("w-1", hostname="pc1", gpu_name="RTX2060", free_vram_mb=2000) assert w2.free_vram_mb == 2000 # 更新 assert w2.created_at == w.created_at # 没新建 @@ -195,7 +195,7 @@ def test_register_with_task_id_refreshes_task_heartbeat(svc): {"last_heartbeat_at": old_hb - timedelta(seconds=300)} ) svc.db.commit() - svc.register_worker("w-1", task_id=t.id) + _w, _c = svc.register_worker("w-1", task_id=t.id) svc.db.refresh(t) assert t.last_heartbeat_at > old_hb assert t.status == "processing" # 心跳不改变状态 @@ -213,7 +213,7 @@ def test_register_task_heartbeat_ignores_finished_or_foreign_task(svc): svc.poll_task("w-1") done = svc.report_result(t.id, "w-1", success=True, duration_seconds=10.0) hb_when_done = done.last_heartbeat_at - svc.register_worker("w-1", task_id=t.id) + _w, _c = svc.register_worker("w-1", task_id=t.id) svc.db.refresh(t) assert t.status == "done" assert t.last_heartbeat_at == hb_when_done # 没被改写 @@ -232,14 +232,14 @@ def test_register_task_heartbeat_ignores_finished_or_foreign_task(svc): {"last_heartbeat_at": owner_hb - timedelta(seconds=600)} ) svc.db.commit() - svc.register_worker("w-1", task_id=t2.id) # 旧 worker 迟到心跳 + _w2, _c2 = svc.register_worker("w-1", task_id=t2.id) # 旧 worker 迟到心跳 svc.db.refresh(t2) assert t2.worker_id == "w-2" assert t2.status == "processing" assert t2.last_heartbeat_at == owner_hb # 场景 3:不存在的 task_id 不报错 - svc.register_worker("w-1", task_id="nonexistent-id") + _wn, _cn = svc.register_worker("w-1", task_id="nonexistent-id") assert svc.db.get(GpuLipsyncTaskModel, "nonexistent-id") is None -- 2.54.0 From f42fe18269841fb85c79cf13e88a12d931488cb1 Mon Sep 17 00:00:00 2001 From: saas-backend Date: Tue, 22 Sep 2026 00:46:20 +0800 Subject: [PATCH 4/6] =?UTF-8?q?test:=20=E8=A1=A5=E5=8F=96=E6=B6=88?= =?UTF-8?q?=E9=93=BE=E8=B7=AF=E5=8D=95=E6=B5=8B=20+=20=E4=BF=AE=E5=A4=8D?= =?UTF-8?q?=20SQLite=20=E5=BC=95=E6=93=8E=E8=BF=9E=E6=8E=A5=E6=B1=A0?= =?UTF-8?q?=E5=8F=82=E6=95=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 test_gpu_lipsync_routes.py:register 透传 cancel_task、cancel 接受 processing - gpu_lipsync_service 测试:心跳检测取消、cancelled 不回退 pending、wait_for_result 终态 - lipsync_service 测试(真实 SQLite):cancel_job 同步取消 GPU 任务、非 gpu 前缀不动 GPU 表 - celery 测试:cancelled GPU 任务不回退 MediaKit - worker 测试:TaskHeartbeat 检测 cancel 调 /cancel、取消上报不重试 - 修复 build_engine 给 SQLite 传 pool_size/max_overflow/pool_timeout 导致 TypeError --- packages/adapters/sqlalchemy_impl/session.py | 4 + tests/unit/test_1970_gpu_worker_heartbeat.py | 56 ++++++++- tests/unit/test_gpu_lipsync_routes.py | 89 ++++++++++++++ tests/unit/test_gpu_lipsync_service.py | 44 +++++++ tests/unit/test_lipsync_gpu_async_task.py | 37 ++++++ tests/unit/test_lipsync_gpu_integration.py | 123 +++++++++++++++++++ 6 files changed, 348 insertions(+), 5 deletions(-) create mode 100644 tests/unit/test_gpu_lipsync_routes.py diff --git a/packages/adapters/sqlalchemy_impl/session.py b/packages/adapters/sqlalchemy_impl/session.py index 89f541368..02c0ce4c9 100644 --- a/packages/adapters/sqlalchemy_impl/session.py +++ b/packages/adapters/sqlalchemy_impl/session.py @@ -18,6 +18,10 @@ def build_engine( pool_timeout: int = 30, pool_recycle: int = 3600, ): + # SQLite 不支持 QueuePool 的 pool_size/max_overflow/pool_timeout, + # 传了会在 create_engine 阶段直接 TypeError,这里只对非 SQLite 传连接池参数。 + if _is_sqlite(database_url): + return create_engine(database_url, pool_recycle=pool_recycle) return create_engine( database_url, pool_size=pool_size, diff --git a/tests/unit/test_1970_gpu_worker_heartbeat.py b/tests/unit/test_1970_gpu_worker_heartbeat.py index 005723e8c..bea492320 100644 --- a/tests/unit/test_1970_gpu_worker_heartbeat.py +++ b/tests/unit/test_1970_gpu_worker_heartbeat.py @@ -97,7 +97,7 @@ def test_task_heartbeat_thread_sends_and_stops(worker, monkeypatch): def _fake_register(task_id=None): calls.append(task_id) - return True + return True, False monkeypatch.setattr(worker, "_register", _fake_register) hb = worker.TaskHeartbeat("task-hb1", interval=5) @@ -109,6 +109,52 @@ def test_task_heartbeat_thread_sends_and_stops(worker, monkeypatch): assert calls and all(c == "task-hb1" for c in calls) +def test_task_heartbeat_cancel_calls_musetalk_cancel(worker, monkeypatch): + """心跳响应 cancel_task=True → 调 _cancel_musetalk 并设置 cancelled 标志。""" + cancel_calls = [] + + monkeypatch.setattr(worker, "_register", lambda *a, **k: (True, True)) + monkeypatch.setattr(worker, "_cancel_musetalk", lambda: cancel_calls.append(1)) + + hb = worker.TaskHeartbeat("task-cancel-1", interval=5) + hb.start() + hb.join(timeout=2) # 检测到取消后线程自行 return + assert not hb.is_alive() + assert hb.cancelled is True + assert cancel_calls == [1] + + +def test_handle_task_reports_cancelled_after_musetalk_abort(worker, monkeypatch): + """推理被 /cancel 终止后,hb.cancelled=True → 上报失败而非重试。""" + reports = [] + + monkeypatch.setattr(worker, "_register", lambda *a, **k: (True, False)) + monkeypatch.setattr(worker, "_download", lambda url, path: True) + monkeypatch.setattr(worker, "_probe_duration", lambda path: 12.0) + # 模拟推理被终止(/inference 返回错误) + monkeypatch.setattr( + worker, "_call_musetalk", lambda v, a, o: (False, 0.0, "推理被取消", False) + ) + monkeypatch.setattr( + worker, + "_report_result", + lambda task_id, success, duration=0.0, error_msg="": reports.append(error_msg) or True, + ) + + # 让 TaskHeartbeat 在主线程检查时报告已取消 + orig_hb_init = worker.TaskHeartbeat + + def _hb(task_id, interval): + h = orig_hb_init(task_id, interval) + h.cancelled = True + return h + + monkeypatch.setattr(worker, "TaskHeartbeat", _hb) + + worker._handle_task({"task_id": "t-canceled", "video_url": "u", "audio_url": "u"}) + assert reports == ["用户取消任务"] + + # ── 短视频前置拦截 ───────────────────────────────────────────────── @@ -119,7 +165,7 @@ def test_handle_task_short_video_reports_failed_without_inference(worker, monkey audio.write_bytes(b"fake-audio") reports = [] - monkeypatch.setattr(worker, "_register", lambda *a, **k: True) + monkeypatch.setattr(worker, "_register", lambda *a, **k: (True, False)) monkeypatch.setattr(worker, "_download", lambda url, path: True) # ffprobe 读出 1.2s → 低于 3s 阈值 monkeypatch.setattr(worker, "_probe_duration", lambda path: 1.2) @@ -152,7 +198,7 @@ def test_handle_task_short_video_reports_failed_without_inference(worker, monkey def test_handle_task_probe_failure_does_not_block(worker, monkeypatch): """ffprobe 不可用(duration=0.0)时不能误杀,应继续推理.""" reports = [] - monkeypatch.setattr(worker, "_register", lambda *a, **k: True) + monkeypatch.setattr(worker, "_register", lambda *a, **k: (True, False)) monkeypatch.setattr(worker, "_download", lambda url, path: True) monkeypatch.setattr(worker, "_probe_duration", lambda path: 0.0) monkeypatch.setattr( @@ -218,7 +264,7 @@ def test_handle_task_retries_once_for_transient_then_succeeds(worker, monkeypatc return False, 0.0, "MuseTalk HTTP 503: busy", True return True, 6.5, "", False - monkeypatch.setattr(worker, "_register", lambda *a, **k: True) + monkeypatch.setattr(worker, "_register", lambda *a, **k: (True, False)) monkeypatch.setattr(worker, "_download", lambda url, path: True) monkeypatch.setattr(worker, "_probe_duration", lambda path: 12.0) monkeypatch.setattr(worker, "_call_musetalk", _fake_call) @@ -243,7 +289,7 @@ def test_handle_task_no_retry_for_deterministic_failure(worker, monkeypatch): return False, 0.0, "MuseTalk HTTP 400: bad input", False reports = [] - monkeypatch.setattr(worker, "_register", lambda *a, **k: True) + monkeypatch.setattr(worker, "_register", lambda *a, **k: (True, False)) monkeypatch.setattr(worker, "_download", lambda url, path: True) monkeypatch.setattr(worker, "_probe_duration", lambda path: 12.0) monkeypatch.setattr(worker, "_call_musetalk", _fake_call) diff --git a/tests/unit/test_gpu_lipsync_routes.py b/tests/unit/test_gpu_lipsync_routes.py new file mode 100644 index 000000000..c3d1561b0 --- /dev/null +++ b/tests/unit/test_gpu_lipsync_routes.py @@ -0,0 +1,89 @@ +"""GPU Worker 路由单测 — #2009 取消链路. + +直接调用路由函数(不经 HTTP 栈),显式注入 svc / _token 以跳过 Depends。 +CI 增量映射: gpu_lipsync.py (route) → test_gpu_lipsync_routes.py +""" + +from __future__ import annotations + +import json +from unittest.mock import MagicMock, patch + +import pytest + + +def _payload(**overrides): + from app.schemas.gpu_lipsync import GpuWorkerRegisterRequest + + data = { + "worker_id": "w-1", + "hostname": "gpu-host", + "gpu_name": "RTX3060", + "free_vram_mb": 10000, + "capabilities": json.dumps({"musetalk": True}), + } + data.update(overrides) + return GpuWorkerRegisterRequest(**data) + + +def test_register_returns_cancel_task_true_when_cancelled(): + """心跳接口在任务已取消时必须把 cancel_task=True 透传给 Worker.""" + fake_worker = MagicMock() + fake_worker.worker_id = "w-1" + fake_worker.hostname = "gpu-host" + fake_worker.gpu_name = "RTX3060" + fake_worker.free_vram_mb = 10000 + fake_worker.capabilities = "musetalk" + fake_svc = MagicMock() + fake_svc.register_worker.return_value = (fake_worker, True) + + from app.api.routes.gpu_lipsync import register_worker as route + + resp = route(_payload(task_id="task-cancelled"), svc=fake_svc, _token="t") + + assert resp.cancel_task is True + assert resp.ok is True + fake_svc.register_worker.assert_called_once() + kwargs = fake_svc.register_worker.call_args.kwargs + assert kwargs["task_id"] == "task-cancelled" + + +def test_register_returns_cancel_task_false_normal(): + """正常心跳 cancel_task=False.""" + fake_worker = MagicMock() + fake_worker.worker_id = "w-1" + fake_worker.hostname = "gpu-host" + fake_worker.gpu_name = "RTX3060" + fake_worker.free_vram_mb = 10000 + fake_worker.capabilities = "musetalk" + fake_svc = MagicMock() + fake_svc.register_worker.return_value = (fake_worker, False) + + from app.api.routes.gpu_lipsync import register_worker as route + + resp = route(_payload(), svc=fake_svc, _token="t") + + assert resp.cancel_task is False + + +def test_cancel_route_accepts_processing_status(): + """cancel 路由允许 processing 状态(GPU 推理中),不再 400。""" + fake_job = MagicMock() + fake_job.status = "cancelled" + + svc = MagicMock() + svc.cancel_job.return_value = fake_job + + current_user = MagicMock() + current_user.user.id = "u1" + + from app.api.routes.lipsync import cancel_lipsync_job as route + + result = route("job-1", current_user, svc) + + svc.cancel_job.assert_called_once_with("job-1", "u1") + assert result.status == "cancelled" + + +if __name__ == "__main__": + pytest.main([__file__, "-q"]) diff --git a/tests/unit/test_gpu_lipsync_service.py b/tests/unit/test_gpu_lipsync_service.py index 6e2e59381..651393d91 100644 --- a/tests/unit/test_gpu_lipsync_service.py +++ b/tests/unit/test_gpu_lipsync_service.py @@ -243,6 +243,50 @@ def test_register_task_heartbeat_ignores_finished_or_foreign_task(svc): assert svc.db.get(GpuLipsyncTaskModel, "nonexistent-id") is None +def test_register_task_heartbeat_detects_cancelled(svc): + """取消链路:任务已 cancelled 时,register 心跳必须返回 cancel_task=True.""" + t = svc.create_task(video_url="v", audio_url="a") + svc.poll_task("w-1") + svc.db.refresh(t) + # 用户取消:直接把任务置为 cancelled + t.status = "cancelled" + t.finished_at = datetime.now(UTC) + svc.db.commit() + + _w, cancel_task = svc.register_worker("w-1", task_id=t.id) + assert cancel_task is True + svc.db.refresh(t) + assert t.status == "cancelled" # 心跳不改写已取消状态 + + +def test_report_result_cancelled_stays_cancelled(svc): + """Worker 终止取消任务后上报失败,report_result 必须保持 cancelled 不回退 pending.""" + t = svc.create_task(video_url="v", audio_url="a") + svc.poll_task("w-1") + svc.db.refresh(t) + t.status = "cancelled" + svc.db.commit() + + result = svc.report_result(t.id, "w-1", success=False, error_msg="推理被终止") + assert result.status == "cancelled" + assert result.finished_at is not None + assert "推理被终止" in (result.error_msg or "") + + +def test_wait_for_result_returns_when_cancelled(svc): + """wait_for_result 将 cancelled 视为终态,立即返回,Celery 不回退 MediaKit.""" + t = svc.create_task(video_url="v", audio_url="a") + svc.poll_task("w-1") + svc.db.refresh(t) + t.status = "cancelled" + t.finished_at = datetime.now(UTC) + svc.db.commit() + + result = svc.wait_for_result(t.id, timeout_seconds=5, poll_interval=0.1) + assert result is not None + assert result.status == "cancelled" + + def test_default_gpu_task_timeout_is_900(svc): """#1970 默认超时 300→900,覆盖 RTX2060 长视频推理.""" assert svc.settings.gpu_task_timeout_seconds == 900 diff --git a/tests/unit/test_lipsync_gpu_async_task.py b/tests/unit/test_lipsync_gpu_async_task.py index 59a49f1a5..2ab3a9404 100644 --- a/tests/unit/test_lipsync_gpu_async_task.py +++ b/tests/unit/test_lipsync_gpu_async_task.py @@ -248,3 +248,40 @@ class TestSignMediaUrl: with patch.object(task_mod, "get_shared_storage_service", side_effect=RuntimeError("x")): url = "https://own-bucket.oss-cn-beijing.aliyuncs.com/a.wav" assert task_mod._sign_media_url(url) == url + + +def test_cancelled_gpu_task_does_not_fallback_mediakit(monkeypatch): + """GPU 任务被用户取消 → Celery 任务直接标记 cancelled,不回退 MediaKit。""" + job = MagicMock() + job.id = "job-1" + job.status = "processing" + job.mediakit_task_id = "gpu:gpu-task-1" + + gpu_task = MagicMock() + gpu_task.status = "cancelled" + gpu_task.error_msg = "用户取消" + + fake_gpu_svc = MagicMock() + fake_gpu_svc.wait_for_result.return_value = gpu_task + + fake_db = MagicMock() + fake_db.query.return_value.filter_by.return_value.first.return_value = job + # wait_for_result 直接被 mock 到 gpu_svc,这里仅备查 + + monkeypatch.setattr( + "app.services.gpu_lipsync_service.GpuLipsyncService", + MagicMock(return_value=fake_gpu_svc), + ) + session_factory = MagicMock() + session_factory.return_value = fake_db + # _get_db_session 优先用 worker_app.db(pytest 环境可导入),两个都 patch + monkeypatch.setattr("worker_app.db.SessionLocal", session_factory) + monkeypatch.setattr("app.db.SessionLocal", session_factory) + monkeypatch.setattr("app.tasks.lipsync_gpu.logger", MagicMock()) + + task_mod.lipsync_gpu_process_async.run("job-1", "u1", "gpu-task-1") + + assert job.status == "cancelled" + assert not str(job.mediakit_task_id).startswith("mk-") + fake_db.commit.assert_called() + fake_gpu_svc.wait_for_result.assert_called_once() diff --git a/tests/unit/test_lipsync_gpu_integration.py b/tests/unit/test_lipsync_gpu_integration.py index e8e3a0c0a..5afc81804 100644 --- a/tests/unit/test_lipsync_gpu_integration.py +++ b/tests/unit/test_lipsync_gpu_integration.py @@ -323,3 +323,126 @@ class TestGpuServiceHelpers: svc = GpuLipsyncService(db=fake_db) fake_db.query.return_value.filter.return_value.first.return_value = None assert svc.has_available_worker() is False + + +# ── cancel_job 取消链路 (#2009) ───────────────────────────────────── + + +def _build_sqlite_session(): + from sqlalchemy import create_engine + from sqlalchemy.orm import sessionmaker + + from packages.adapters.sqlalchemy_impl import models as _ # noqa: F401 + from packages.adapters.sqlalchemy_impl.models import Base + + engine = create_engine("sqlite:///:memory:", future=True) + Base.metadata.create_all(engine) + Session = sessionmaker(bind=engine, future=True) + return Session() + + +def _make_real_job(db, *, status="processing", mediakit_task_id="gpu:gpu-task-1"): + import uuid + + from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel + + job = LipsyncJobModel( + id=str(uuid.uuid4()), + user_id="u1", + project_id="p1", + video_url="videos/v.mp4", + audio_url="audios/a.wav", + enable_video_loop=True, + mediakit_task_id=mediakit_task_id, + status=status, + ) + db.add(job) + db.commit() + return job + + +def test_cancel_processing_gpu_job_marks_gpu_task_cancelled(): + """processing 的 GPU job 取消时,关联 GpuLipsyncTask 必须同步置 cancelled.""" + import uuid + + from packages.adapters.sqlalchemy_impl.models import GpuLipsyncTaskModel + + db = _build_sqlite_session() + gpu_task_id = str(uuid.uuid4()) + gpu_task = GpuLipsyncTaskModel( + id=gpu_task_id, + video_url="v", + audio_url="a", + status="processing", + worker_id="w-1", + attempt=1, + ) + db.add(gpu_task) + db.commit() + + job = _make_real_job(db, mediakit_task_id=f"gpu:{gpu_task_id}") + from app.services.lipsync_service import LipsyncService + + svc = LipsyncService(db=db, client=MagicMock()) + result = svc.cancel_job(job.id, "u1") + + assert result.status == "cancelled" + db.refresh(gpu_task) + assert gpu_task.status == "cancelled" + assert gpu_task.error_msg == "用户取消" + assert gpu_task.finished_at is not None + + +def test_cancel_processing_gpu_job_skips_non_processing_gpu_task(): + """GPU task 已不在 processing(如已 done)时,取消 job 不应改它,也不报错.""" + import uuid + + from packages.adapters.sqlalchemy_impl.models import GpuLipsyncTaskModel + + db = _build_sqlite_session() + gpu_task_id = str(uuid.uuid4()) + gpu_task = GpuLipsyncTaskModel( + id=gpu_task_id, video_url="v", audio_url="a", status="done", worker_id="w-1", attempt=1 + ) + db.add(gpu_task) + db.commit() + + job = _make_real_job(db, mediakit_task_id=f"gpu:{gpu_task_id}") + from app.services.lipsync_service import LipsyncService + + svc = LipsyncService(db=db, client=MagicMock()) + result = svc.cancel_job(job.id, "u1") + + assert result.status == "cancelled" + db.refresh(gpu_task) + assert gpu_task.status == "done" # 没被动 + + +def test_cancel_processing_non_gpu_job_does_not_touch_gpu_table(): + """mediakit_task_id 不是 gpu: 前缀(普通 MediaKit 任务)时,不查 GPU task.""" + db = _build_sqlite_session() + job = _make_real_job(db, mediakit_task_id="mk-task-99") + from app.services.lipsync_service import LipsyncService + + svc = LipsyncService(db=db, client=MagicMock()) + result = svc.cancel_job(job.id, "u1") + assert result.status == "cancelled" + + +def test_cancel_completed_job_unchanged(): + """completed 状态不可取消,cancel_job 原样返回.""" + db = _build_sqlite_session() + job = _make_real_job(db, status="completed", mediakit_task_id="gpu:x") + from app.services.lipsync_service import LipsyncService + + svc = LipsyncService(db=db, client=MagicMock()) + result = svc.cancel_job(job.id, "u1") + assert result.status == "completed" + + +def test_cancel_job_not_found_returns_none(): + db = _build_sqlite_session() + from app.services.lipsync_service import LipsyncService + + svc = LipsyncService(db=db, client=MagicMock()) + assert svc.cancel_job("nonexistent", "u1") is None -- 2.54.0 From 3a0f1eba03a07a712f12521936caa3849dbbbd25 Mon Sep 17 00:00:00 2001 From: CI Bot Date: Mon, 21 Sep 2026 16:52:53 +0000 Subject: [PATCH 5/6] style: auto-format with black + isort + ruff + prettier [skip ci-format-check] --- tests/unit/test_1970_gpu_worker_heartbeat.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/tests/unit/test_1970_gpu_worker_heartbeat.py b/tests/unit/test_1970_gpu_worker_heartbeat.py index bea492320..f2cabd38c 100644 --- a/tests/unit/test_1970_gpu_worker_heartbeat.py +++ b/tests/unit/test_1970_gpu_worker_heartbeat.py @@ -132,9 +132,7 @@ def test_handle_task_reports_cancelled_after_musetalk_abort(worker, monkeypatch) monkeypatch.setattr(worker, "_download", lambda url, path: True) monkeypatch.setattr(worker, "_probe_duration", lambda path: 12.0) # 模拟推理被终止(/inference 返回错误) - monkeypatch.setattr( - worker, "_call_musetalk", lambda v, a, o: (False, 0.0, "推理被取消", False) - ) + monkeypatch.setattr(worker, "_call_musetalk", lambda v, a, o: (False, 0.0, "推理被取消", False)) monkeypatch.setattr( worker, "_report_result", -- 2.54.0 From 87dfed2f8a0ab84a0f39af452c6558d17ee1bdfe Mon Sep 17 00:00:00 2001 From: saas-backend Date: Tue, 22 Sep 2026 01:10:22 +0800 Subject: [PATCH 6/6] =?UTF-8?q?test:=20=E4=BF=AE=E5=A4=8D=E5=8F=96?= =?UTF-8?q?=E6=B6=88=E9=93=BE=E8=B7=AF=E6=B5=8B=E8=AF=95=E5=9C=A8=E5=85=A8?= =?UTF-8?q?=E9=87=8F=E8=B7=91=E6=97=B6=E5=8F=97=20sys.modules=20=E6=B1=A1?= =?UTF-8?q?=E6=9F=93=E5=A4=B1=E8=B4=A5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 全量跑时其他测试可能把 app.services.gpu_lipsync_service 换成 MagicMock, 导致任务函数内 from...import 拿到污染对象。改用 monkeypatch.setitem 直接替换 sys.modules 模块项,并 patch _get_db_session 绕开双 import 分支, 确保测试在任何污染场景下都稳定。 --- tests/unit/test_lipsync_gpu_async_task.py | 23 ++++++++++++----------- 1 file changed, 12 insertions(+), 11 deletions(-) diff --git a/tests/unit/test_lipsync_gpu_async_task.py b/tests/unit/test_lipsync_gpu_async_task.py index 2ab3a9404..00b8ece1f 100644 --- a/tests/unit/test_lipsync_gpu_async_task.py +++ b/tests/unit/test_lipsync_gpu_async_task.py @@ -263,21 +263,21 @@ def test_cancelled_gpu_task_does_not_fallback_mediakit(monkeypatch): fake_gpu_svc = MagicMock() fake_gpu_svc.wait_for_result.return_value = gpu_task + gpu_service_cls = MagicMock(return_value=fake_gpu_svc) fake_db = MagicMock() fake_db.query.return_value.filter_by.return_value.first.return_value = job - # wait_for_result 直接被 mock 到 gpu_svc,这里仅备查 - monkeypatch.setattr( - "app.services.gpu_lipsync_service.GpuLipsyncService", - MagicMock(return_value=fake_gpu_svc), - ) - session_factory = MagicMock() - session_factory.return_value = fake_db - # _get_db_session 优先用 worker_app.db(pytest 环境可导入),两个都 patch - monkeypatch.setattr("worker_app.db.SessionLocal", session_factory) - monkeypatch.setattr("app.db.SessionLocal", session_factory) - monkeypatch.setattr("app.tasks.lipsync_gpu.logger", MagicMock()) + # 直接替换 sys.modules 里的 gpu_lipsync_service 模块(全量跑时它可能已被 + # 其他测试换成 MagicMock),保证任务函数内 from...import 一定拿到我们的类; + # 并替换 _get_db_session 绕开 worker_app / app.db 两条 import 分支。 + import sys + from types import SimpleNamespace + + fake_mod = SimpleNamespace(GpuLipsyncService=gpu_service_cls) + monkeypatch.setitem(sys.modules, "app.services.gpu_lipsync_service", fake_mod) + monkeypatch.setattr(task_mod, "_get_db_session", lambda: fake_db) + monkeypatch.setattr(task_mod, "logger", MagicMock()) task_mod.lipsync_gpu_process_async.run("job-1", "u1", "gpu-task-1") @@ -285,3 +285,4 @@ def test_cancelled_gpu_task_does_not_fallback_mediakit(monkeypatch): assert not str(job.mediakit_task_id).startswith("mk-") fake_db.commit.assert_called() fake_gpu_svc.wait_for_result.assert_called_once() + gpu_service_cls.assert_called_once_with(fake_db) -- 2.54.0