fix: 修复取消链路断裂,前端取消后 GPU 仍继续推理
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 2s
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 20s
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 14s
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m9s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 3m16s
AI Code Review / AI Code Review (pull_request) Successful in 7m0s
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Style (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been cancelled

- 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
This commit is contained in:
saas-backend
2026-09-21 22:09:02 +08:00
parent 36e8f91e5e
commit f55ea0100d
7 changed files with 97 additions and 22 deletions
+2 -2
View File
@@ -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 轮询拉任务 ─────────────────────────
+2 -2
View File
@@ -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
+1
View File
@@ -36,6 +36,7 @@ class GpuWorkerRegisterResponse(BaseModel):
ok: bool = True
server_time: datetime
message: str = "ok"
cancel_task: bool = Field(False, description="当前心跳任务是否已被用户取消;为 true 时 Worker 应终止推理")
# ── 轮询任务 ────────────────────────────────────────────────────
+23 -9
View File
@@ -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:
+21 -2
View File
@@ -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()
+8
View File
@@ -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",
+40 -7
View File
@@ -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
# 轮询任务