diff --git a/tests/unit/test_lipsync_gpu_async_task.py b/tests/unit/test_lipsync_gpu_async_task.py new file mode 100644 index 000000000..59a49f1a5 --- /dev/null +++ b/tests/unit/test_lipsync_gpu_async_task.py @@ -0,0 +1,250 @@ +"""Celery 任务 lipsync_gpu_process_async 直接单测 (#1978 异步化). + +覆盖 apps/api/app/tasks/lipsync_gpu.py 的全部主路径: +- 成功:wait_for_result 返回 done → 签名 URL → completed +- GPU 超时/失败 → MediaKit 兜底(成功/MediaKitError/其他异常) +- job 不存在 / 状态异常提前返回 +- 主流程异常 → job 标 failed +- _sign_media_url 各分支 +""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import app.tasks.lipsync_gpu as task_mod +import pytest + + +def _make_job(status="processing"): + job = MagicMock() + job.id = "job-1" + job.user_id = "u1" + job.status = status + job.video_url = "videos/v.mp4" + job.audio_url = "audios/a.wav" + job.enable_video_loop = True + return job + + +def _make_gpu_task(status="done", result_url="gpu-lipsync/results/t1.mp4", result_duration=11.2): + t = MagicMock() + t.status = status + t.result_url = result_url + t.result_duration = result_duration + return t + + +@pytest.fixture() +def db_patch(): + """patch _get_db_session 返回 MagicMock,并在任务结束后断言 close.""" + fake_db = MagicMock() + with patch.object(task_mod, "_get_db_session", return_value=fake_db): + yield fake_db + + +def _patch_gpu_service(final_task): + fake_svc = MagicMock() + fake_svc.wait_for_result.return_value = final_task + return patch( + "app.services.gpu_lipsync_service.GpuLipsyncService", + return_value=fake_svc, + ) + + +def _run_task(): + # @shared_task bind=True:直接调用任务对象会自动注入 self + task_mod.lipsync_gpu_process_async("job-1", "u1", "gpu-task-1") + + +class TestHappyPath: + def test_gpu_done_marks_completed(self, db_patch): + job = _make_job() + db_patch.query.return_value.filter_by.return_value.first.return_value = job + gpu_task = _make_gpu_task() + storage = MagicMock() + storage.get_download_url.return_value = "https://signed.example.com/r1.mp4?sig=x" + with ( + _patch_gpu_service(gpu_task), + patch.object(task_mod, "get_shared_storage_service", return_value=storage), + ): + _run_task() + assert job.status == "completed" + assert job.output_video_url == "https://signed.example.com/r1.mp4?sig=x" + assert job.output_duration == 11.2 + assert job.completed_at is not None + db_patch.commit.assert_called_once() + db_patch.close.assert_called_once() + + def test_gpu_done_empty_signed_url_keeps_original(self, db_patch): + job = _make_job() + db_patch.query.return_value.filter_by.return_value.first.return_value = job + gpu_task = _make_gpu_task(result_url="gpu/r2.mp4") + storage = MagicMock() + storage.get_download_url.return_value = "" + with ( + _patch_gpu_service(gpu_task), + patch.object(task_mod, "get_shared_storage_service", return_value=storage), + ): + _run_task() + assert job.status == "completed" + assert job.output_video_url == "gpu/r2.mp4" + + def test_gpu_done_result_duration_none_defaults_zero(self, db_patch): + job = _make_job() + db_patch.query.return_value.filter_by.return_value.first.return_value = job + gpu_task = _make_gpu_task(result_duration=None) + storage = MagicMock() + with ( + _patch_gpu_service(gpu_task), + patch.object(task_mod, "get_shared_storage_service", return_value=storage), + ): + _run_task() + assert job.output_duration == 0.0 + + def test_sign_failure_uses_original_url(self, db_patch): + job = _make_job() + db_patch.query.return_value.filter_by.return_value.first.return_value = job + gpu_task = _make_gpu_task(result_url="gpu/r3.mp4") + with ( + _patch_gpu_service(gpu_task), + patch.object(task_mod, "get_shared_storage_service", side_effect=RuntimeError("oss down")), + ): + _run_task() + assert job.status == "completed" + assert job.output_video_url == "gpu/r3.mp4" + + +class TestJobGuards: + def test_job_not_found_returns(self, db_patch): + db_patch.query.return_value.filter_by.return_value.first.return_value = None + _run_task() + db_patch.commit.assert_not_called() + db_patch.close.assert_called_once() + + def test_job_wrong_status_skipped(self, db_patch): + job = _make_job(status="completed") + db_patch.query.return_value.filter_by.return_value.first.return_value = job + _run_task() + db_patch.commit.assert_not_called() + + +class TestGpuFailureFallback: + def test_gpu_timeout_falls_back_mediakit_success(self, db_patch): + job = _make_job() + db_patch.query.return_value.filter_by.return_value.first.return_value = job + with ( + _patch_gpu_service(None), + patch.object(task_mod, "_fallback_to_mediakit") as fb, + ): + _run_task() + fb.assert_called_once_with(db_patch, job) + + def test_gpu_failed_status_falls_back(self, db_patch): + job = _make_job() + db_patch.query.return_value.filter_by.return_value.first.return_value = job + gpu_task = _make_gpu_task(status="failed") + with ( + _patch_gpu_service(gpu_task), + patch.object(task_mod, "_fallback_to_mediakit") as fb, + ): + _run_task() + fb.assert_called_once_with(db_patch, job) + + +class TestFallbackToMediaKit: + def test_mediakit_success_marks_submitted(self, db_patch): + job = _make_job() + fake_client = MagicMock() + fake_client.submit_lipsync.return_value = {"task_id": "mk-99"} + with ( + patch("app.services.mediakit_client.get_mediakit_client", return_value=fake_client), + patch.object(task_mod, "_sign_media_url", side_effect=lambda u: u + "?s"), + ): + task_mod._fallback_to_mediakit(db_patch, job) + fake_client.submit_lipsync.assert_called_once() + kwargs = fake_client.submit_lipsync.call_args.kwargs + assert kwargs["enable_video_loop"] is True + assert kwargs["client_token"] == "job-1" + assert job.status == "submitted" + assert job.mediakit_task_id == "mk-99" + db_patch.commit.assert_called_once() + + def test_mediakit_error_marks_failed(self, db_patch): + from app.services.mediakit_client import MediaKitError + + job = _make_job() + fake_client = MagicMock() + fake_client.submit_lipsync.side_effect = MediaKitError("api reject", code="MkReject") + with ( + patch("app.services.mediakit_client.get_mediakit_client", return_value=fake_client), + patch.object(task_mod, "_sign_media_url", side_effect=lambda u: u), + ): + task_mod._fallback_to_mediakit(db_patch, job) + assert job.status == "failed" + assert job.error_code == "MkReject" + db_patch.commit.assert_called_once() + + def test_other_exception_marks_failed(self, db_patch): + job = _make_job() + with ( + patch("app.services.mediakit_client.get_mediakit_client", side_effect=RuntimeError("boom")), + patch.object(task_mod, "_sign_media_url", side_effect=lambda u: u), + ): + task_mod._fallback_to_mediakit(db_patch, job) + assert job.status == "failed" + assert job.error_code == "FallbackFailed" + db_patch.commit.assert_called_once() + + +class TestTaskException: + def test_unexpected_exception_marks_job_failed(self, db_patch): + job = _make_job() + # query 第一次返回 job,异常路径里再次 query 也返回 job + db_patch.query.return_value.filter_by.return_value.first.return_value = job + with patch( + "app.services.gpu_lipsync_service.GpuLipsyncService", + side_effect=RuntimeError("svc ctor fail"), + ): + _run_task() + assert job.status == "failed" + assert job.error_code == "GpuAsyncError" + + def test_exception_handler_failure_swallowed(self, db_patch): + # 主流程异常,且异常处理中的 query 也抛异常 → 不应再抛 + db_patch.query.side_effect = RuntimeError("db totally broken") + _run_task() + db_patch.close.assert_called_once() + + +class TestSignMediaUrl: + def test_empty_url_returned_as_is(self): + assert task_mod._sign_media_url("") == "" + + def test_non_own_host_returned_as_is(self): + storage = MagicMock() + storage.public_url = "https://own-bucket.oss-cn-beijing.aliyuncs.com" + with patch.object(task_mod, "get_shared_storage_service", return_value=storage): + url = "https://other.example.com/a.wav" + assert task_mod._sign_media_url(url) == url + + def test_own_host_signed(self): + storage = MagicMock() + storage.public_url = "https://own-bucket.oss-cn-beijing.aliyuncs.com" + storage.get_download_url.return_value = "https://own-bucket.oss-cn-beijing.aliyuncs.com/a?sig=1" + with patch.object(task_mod, "get_shared_storage_service", return_value=storage): + out = task_mod._sign_media_url("https://own-bucket.oss-cn-beijing.aliyuncs.com/a.wav") + assert out.endswith("?sig=1") + storage.get_download_url.assert_called_once() + + def test_missing_public_url_returns_original(self): + storage = MagicMock() + storage.public_url = "" + with patch.object(task_mod, "get_shared_storage_service", return_value=storage): + url = "https://own-bucket.oss-cn-beijing.aliyuncs.com/a.wav" + assert task_mod._sign_media_url(url) == url + + def test_exception_returns_original(self): + 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 diff --git a/tests/unit/test_lipsync_gpu_integration.py b/tests/unit/test_lipsync_gpu_integration.py index 60a1412ba..e8e3a0c0a 100644 --- a/tests/unit/test_lipsync_gpu_integration.py +++ b/tests/unit/test_lipsync_gpu_integration.py @@ -141,6 +141,46 @@ class TestGpuFallback: fake_mediakit.submit_lipsync.assert_called_once() assert job.status == "submitted" + def test_gpu_create_returns_none_falls_back_mediakit(self, fake_db, fake_mediakit): + """_submit_to_gpu_create 返回 None(create_task 失败被内部吞掉)→ rollback + MediaKit.""" + svc = _make_svc(fake_db, fake_mediakit, use_gpu=True) + fake_gpu_svc = MagicMock() + fake_gpu_svc.has_available_worker.return_value = True + with ( + _patch_storage(), + patch("app.services.gpu_lipsync_service.GpuLipsyncService", return_value=fake_gpu_svc), + patch.object(svc, "_submit_to_gpu_create", return_value=None) as m_create, + ): + job = _make_job() + svc._submit_audio_direct(job=job) + m_create.assert_called_once() + fake_db.rollback.assert_called_once() + fake_mediakit.submit_lipsync.assert_called_once() + assert job.status == "submitted" + + def test_submit_to_gpu_wait_timeout_returns(self, fake_db, fake_mediakit): + """降级同步等待:wait_for_result 返回 None → 直接返回,job 保持 processing.""" + svc = _make_svc(fake_db, fake_mediakit, use_gpu=True) + fake_gpu_svc = MagicMock() + fake_gpu_svc.wait_for_result.return_value = None + job = _make_job() + job.status = "processing" + svc._submit_to_gpu_wait(job=job, gpu_svc=fake_gpu_svc, gpu_task=MagicMock(id="gpu-task-x")) + fake_gpu_svc.wait_for_result.assert_called_once_with("gpu-task-x") + fake_db.commit.assert_not_called() + assert job.status == "processing" + + def test_submit_to_gpu_wait_failed_status_returns(self, fake_db, fake_mediakit): + """降级同步等待:final_task.status != done → 直接返回.""" + svc = _make_svc(fake_db, fake_mediakit, use_gpu=True) + fake_gpu_svc = MagicMock() + fake_gpu_svc.wait_for_result.return_value = MagicMock(status="failed", result_url="") + job = _make_job() + job.status = "processing" + svc._submit_to_gpu_wait(job=job, gpu_svc=fake_gpu_svc, gpu_task=MagicMock(id="gpu-task-y")) + fake_db.commit.assert_not_called() + assert job.status == "processing" + def test_gpu_external_audio_persisted_to_own_oss(self, fake_db, fake_mediakit): """Bug2 回归:dashscope 临时音频 URL 在创建 GPU 任务前转存自家 OSS.""" svc = _make_svc(fake_db, fake_mediakit, use_gpu=True) @@ -212,6 +252,53 @@ class TestGpuFallback: assert fake_gpu_svc.create_task.call_args.kwargs["audio_url"] == dashscope_url +class TestRefreshGpuStale: + """refresh_job_status 的 GPU 异步 stale 超时分支.""" + + def test_stale_gpu_job_marked_failed(self, fake_db): + from datetime import UTC, datetime, timedelta + + svc = _make_svc(fake_db, MagicMock(), use_gpu=True) + job = MagicMock() + job.status = "processing" + job.mediakit_task_id = "gpu:gpu-task-stale" + job.updated_at = datetime.now(UTC) - timedelta(minutes=31) + with patch.object(svc, "get_job", return_value=job): + result = svc.refresh_job_status("job-stale", "u1") + assert result is job + assert job.status == "failed" + assert job.error_code == "GpuTimeout" + fake_db.commit.assert_called_once() + + def test_fresh_gpu_job_left_processing(self, fake_db): + from datetime import UTC, datetime, timedelta + + svc = _make_svc(fake_db, MagicMock(), use_gpu=True) + job = MagicMock() + job.status = "processing" + job.mediakit_task_id = "gpu:gpu-task-fresh" + job.updated_at = datetime.now(UTC) - timedelta(minutes=2) + with patch.object(svc, "get_job", return_value=job): + result = svc.refresh_job_status("job-fresh", "u1") + assert result is job + assert job.status == "processing" + fake_db.commit.assert_not_called() + + def test_naive_updated_at_stale_marked_failed(self, fake_db): + """updated_at 为 naive datetime 时按 UTC 补时区后再判定.""" + from datetime import UTC, datetime, timedelta + + svc = _make_svc(fake_db, MagicMock(), use_gpu=True) + job = MagicMock() + job.status = "gpu_processing" + job.mediakit_task_id = "gpu:gpu-task-naive" + job.updated_at = datetime.now(UTC).replace(tzinfo=None) - timedelta(minutes=31) + with patch.object(svc, "get_job", return_value=job): + svc.refresh_job_status("job-naive", "u1") + assert job.status == "failed" + assert job.error_code == "GpuTimeout" + + class TestGpuServiceHelpers: """GpuLipsyncService.has_available_worker 测试."""