"""GpuLipsyncService 单元测试 — 覆盖任务创建、轮询认领、结果上报、超时回退等核心逻辑. 使用 SQLite 内存数据库,mock 掉存储层(不真实调用 OSS)。 """ from __future__ import annotations import os import sys from datetime import UTC, datetime, timedelta from unittest import mock import pytest # 确保 packages / apps/api 可导入 ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")) for p in (ROOT, os.path.join(ROOT, "apps", "api"), os.path.join(ROOT, "packages")): if p not in sys.path: sys.path.insert(0, p) # 强制使用内存 SQLite(避免依赖 PG) os.environ["APP_ENV"] = "development" os.environ["JWT_SECRET_KEY"] = "dev-secret-key-for-testing-00000000" os.environ["DATABASE_URL"] = "sqlite:///:memory:" os.environ["USE_IN_MEMORY_DB"] = "1" os.environ["GPU_WORKER_TOKEN"] = "" # development 空 token 放行 def _build_session(): from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker # 使用 packages 的 Base from packages.adapters.sqlalchemy_impl import models as _ # noqa: F401 # 触发 ORM 注册 from packages.adapters.sqlalchemy_impl.models import Base engine = create_engine("sqlite:///:memory:", future=True) Base.metadata.create_all(engine) Session = sessionmaker(bind=engine, autoflush=False, autocommit=False, future=True) return Session() @pytest.fixture def svc(): from app.services.gpu_lipsync_service import GpuLipsyncService db = _build_session() service = GpuLipsyncService(db) # mock 存储签名(SQLite 测试无 OSS) service.storage = mock.MagicMock() service.storage.get_download_url.side_effect = ( lambda k, expires_seconds=3600: f"https://signed.example.com/download/{k}?e={expires_seconds}" ) service.storage.get_upload_url.side_effect = ( lambda k, expires_seconds=3600, content_type="video/mp4": f"https://signed.example.com/upload/{k}?e={expires_seconds}" ) return service # ── 创建任务 ────────────────────────────────────────────────────── def test_create_task(svc): task = svc.create_task( video_url="uploads/v.mp4", audio_url="uploads/a.mp3", lipsync_job_id="lip-1", user_id="u-1", project_id="p-1", ) assert task.id assert task.status == "pending" assert task.lipsync_job_id == "lip-1" assert task.attempt == 0 assert task.video_url == "uploads/v.mp4" # ── 轮询认领 ────────────────────────────────────────────────────── def test_poll_returns_none_when_empty(svc): assert svc.poll_task("w-1") is None def test_poll_claims_pending_task(svc): svc.create_task(video_url="uploads/v.mp4", audio_url="uploads/a.mp3") claimed = svc.poll_task("w-1") assert claimed is not None assert claimed.status == "processing" assert claimed.worker_id == "w-1" assert claimed.attempt == 1 # 带签名 URL assert claimed._signed_video_url.startswith("https://signed.example.com/download/") assert claimed._signed_upload_url.startswith("https://signed.example.com/upload/") # 再 poll 无任务 assert svc.poll_task("w-1") is None def test_poll_concurrent_claim_only_one_wins(svc): """并发场景:两个 worker 同时 poll 只有一个能拿到任务(借助 update where status=pending)。""" svc.create_task(video_url="v", audio_url="a") t1 = svc.poll_task("w-1") t2 = svc.poll_task("w-2") assert t1 is not None assert t2 is None # ── 结果上报 ────────────────────────────────────────────────────── def test_report_result_success(svc): t = svc.create_task(video_url="v", audio_url="a") svc.poll_task("w-1") # claim done = svc.report_result(t.id, "w-1", success=True, duration_seconds=12.5) assert done.status == "done" assert done.result_duration == 12.5 assert done.result_url.startswith("gpu-lipsync/results/") assert done.finished_at is not None def test_report_result_failure_requeues(svc): t = svc.create_task(video_url="v", audio_url="a") svc.poll_task("w-1") failed = svc.report_result(t.id, "w-1", success=False, error_msg="MuseTalk crash") assert failed.status == "pending" # 仍在重试次数内 → 回队 assert failed.worker_id == "" assert failed.started_at is None assert "MuseTalk crash" in failed.error_msg def test_report_failure_exhausted_goes_failed(svc): """失败达到 MAX_ATTEMPTS 后标记 failed,不再回队. poll 成功会将 attempt 从 0 开始自增; 第 1/2 次失败回队,第 3 次失败(attempt==MAX_ATTEMPTS)置 failed。 """ from app.services import gpu_lipsync_service as mod t = svc.create_task(video_url="v", audio_url="a") # 模拟失败到上限:poll + fail 重复 MAX_ATTEMPTS 次 for i in range(mod.MAX_ATTEMPTS): claimed = svc.poll_task(f"w-{i}") assert claimed is not None, f"第 {i} 次 poll 应能拿到任务" svc.report_result(t.id, claimed.worker_id, success=False, error_msg=f"fail {i}") svc.db.refresh(t) if i == mod.MAX_ATTEMPTS - 1: assert t.status == "failed" else: assert t.status == "pending" # ── 心跳/超时回退 ───────────────────────────────────────────────── def test_timed_out_task_is_redispatched(svc): """processing 超过 gpu_task_timeout_seconds 无心跳 → 回退 pending.""" t = svc.create_task(video_url="v", audio_url="a") svc.poll_task("w-1") svc.db.refresh(t) assert t.status == "processing" # 手动把 last_heartbeat_at 设到很久以前 t.last_heartbeat_at = datetime.now(UTC) - timedelta(seconds=svc.settings.gpu_task_timeout_seconds + 10) svc.db.commit() # 再次 poll 会触发 _recover_timed_out_tasks 把它回队 claimed = svc.poll_task("w-2") assert claimed is not None assert claimed.id == t.id assert claimed.worker_id == "w-2" assert claimed.attempt == 2 # 又认领了一次 # ── Worker 注册 ─────────────────────────────────────────────────── def test_register_worker_creates_then_updates(svc): 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, _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 # 没新建 def test_register_with_task_id_refreshes_task_heartbeat(svc): """#1970 推理期心跳:register(task_id=...) 只刷新本 worker 的 processing 任务.""" from packages.adapters.sqlalchemy_impl.models import GpuWorkerModel t = svc.create_task(video_url="v", audio_url="a") svc.poll_task("w-1") svc.db.refresh(t) old_hb = t.last_heartbeat_at assert t.status == "processing" # 模拟时间流逝后心跳到达 svc.db.query(GpuWorkerModel).filter_by(worker_id="w-1").update( {"last_heartbeat_at": old_hb - timedelta(seconds=300)} ) svc.db.commit() _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" # 心跳不改变状态 # worker 表心跳也被刷新 w = svc.db.query(GpuWorkerModel).filter_by(worker_id="w-1").one() assert w.last_heartbeat_at > old_hb def test_register_task_heartbeat_ignores_finished_or_foreign_task(svc): """任务已 done,或已被超时回收重新派发给别的 worker 时,旧心跳必须忽略.""" from packages.adapters.sqlalchemy_impl.models import GpuLipsyncTaskModel, GpuWorkerModel # 场景 1:任务已完成 → register 带 task_id 不得改写任务心跳 t = svc.create_task(video_url="v", audio_url="a") 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 _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 # 没被改写 # 场景 2:任务超时回收后被 w-2 重新认领,旧 worker w-1 的迟到心跳无效 t2 = svc.create_task(video_url="v2", audio_url="a2") svc.poll_task("w-1") svc.db.refresh(t2) t2.last_heartbeat_at = datetime.now(UTC) - timedelta(days=1) svc.db.commit() claimed = svc.poll_task("w-2") # 触发回收并由 w-2 重新认领 assert claimed is not None and claimed.id == t2.id owner_hb = claimed.last_heartbeat_at # 把 w-2 的 worker 心跳拨早,确认旧心跳不会影响任务归属 svc.db.query(GpuWorkerModel).filter_by(worker_id="w-2").update( {"last_heartbeat_at": owner_hb - timedelta(seconds=600)} ) svc.db.commit() _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 不报错 _wn, _cn = svc.register_worker("w-1", task_id="nonexistent-id") 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 # ── get_by_lipsync_job ───────────────────────────────────────────── def test_get_by_lipsync_job_returns_latest(svc): svc.create_task(video_url="v", audio_url="a", lipsync_job_id="lip-1") svc.create_task(video_url="v", audio_url="a", lipsync_job_id="lip-1") latest = svc.get_by_lipsync_job("lip-1") assert latest is not None