"""AI 数字人口型 TTS Celery 异步任务 — 单元测试. 覆盖 lipsync_tts.py 的全部主要分支: - Job 不存在/cancelled/正常/异常路径 - TTS 合成、音频下载、OSS 上传、MediaKit 提交 - CosyVoiceError/ValueError/MediaKitError/顶层异常等错误码 """ import os import sys import types from types import ModuleType from unittest.mock import MagicMock, patch os.environ.setdefault("JWT_SECRET_KEY", "dev-secret-key-for-testing") class _FakeQuery: """模拟 SQLAlchemy query.filter().first() 链式调用.""" def __init__(self, job): self._job = job def filter(self, *args, **kwargs): return self def first(self): return self._job def _make_fake_job(**kwargs): """构造可 setattr 的 job 记录.""" job = MagicMock() job.id = kwargs.get("job_id", "job-1") job.user_id = kwargs.get("user_id", "user-1") job.status = kwargs.get("status", "tts_processing") job.audio_url = kwargs.get("audio_url", "") job.video_url = kwargs.get("video_url", "https://oss/video.mp4") job.mediakit_task_id = kwargs.get("mediakit_task_id", "") job.enable_video_loop = kwargs.get("enable_video_loop", False) job.error_code = "" job.error_message = "" job.submitted_at = None job.updated_at = None return job def _build_session(job): """构造 mock DB session + factory. 返回 (session, factory).""" session = MagicMock() session.query.return_value = _FakeQuery(job) session.commit = MagicMock() session.close = MagicMock() factory = MagicMock(return_value=session) return session, factory def _apply_all_patches( *, job=None, cosyvoice_service=None, cosyvoice_side_effect=None, cosyvoice_error=None, download_bytes=b"AUDIO", download_error=None, storage=None, mk_client=None, mk_submit_return=None, mk_submit_error=None, ): """统一构造测试需要的 patch 列表. lipsync_tts.run() 在函数体内部懒 import 多个模块,通过 sys.modules 注入 伪造包路径避免真实导入;对存在的模块用 patch() 替换返回值/side_effect。 """ # SessionLocal 通过懒探测获取(Worker 用 worker_app.db,API 用 app.db), # 测试环境里两个模块都能被真实导入,必须同时 mock 保证用的是 fake session。 fake_app_db = ModuleType("app.db") fake_worker_db = ModuleType("worker_app.db") session, factory = _build_session(job) fake_app_db.SessionLocal = factory fake_worker_db.SessionLocal = factory patches = [ patch.dict(sys.modules, {"app.db": fake_app_db, "worker_app.db": fake_worker_db}), patch( "app.tasks.lipsync_tts._sign_media_url", side_effect=lambda url: url + "?signed" if url else url, ), ] # CosyVoice if cosyvoice_service is not None: cosy_instance = cosyvoice_service else: cosy_instance = MagicMock() if cosyvoice_side_effect is not None: cosy_instance.submit_synthesize_task.side_effect = cosyvoice_side_effect elif cosyvoice_error is not None: cosy_instance.submit_synthesize_task.side_effect = cosyvoice_error else: cosy_instance.submit_synthesize_task.return_value = {"audio_url": "https://tts/raw.mp3"} patches.append(patch("packages.application.cosyvoice_service.CosyVoiceService", return_value=cosy_instance)) # safe_download_bytes if download_error is not None: patches.append(patch("packages.shared.url_security.safe_download_bytes", side_effect=download_error)) else: patches.append(patch("packages.shared.url_security.safe_download_bytes", return_value=download_bytes)) # Storage if storage is None: storage = MagicMock() storage.public_url = "https://oss.example.com" storage.upload_file.return_value = "https://oss.example.com/tts.mp3" patches.append(patch("packages.shared.storage.get_shared_storage_service", return_value=storage)) # MediaKit client if mk_client is not None: patches.append(patch("app.services.mediakit_client.get_mediakit_client", return_value=mk_client)) else: client = MagicMock() if mk_submit_error is not None: client.submit_lipsync.side_effect = mk_submit_error else: client.submit_lipsync.return_value = mk_submit_return or {"task_id": "mk-1"} patches.append(patch("app.services.mediakit_client.get_mediakit_client", return_value=client)) return session, patches class TestTtsSynthesizeAndSubmit: """测试 Celery 任务 tts_synthesize_and_submit.run 的所有分支.""" def test_job_not_found_returns_early(self): """Job 不存在 → 日志报错直接返回,不抛异常.""" from app.tasks.lipsync_tts import tts_synthesize_and_submit session, patches = _apply_all_patches(job=None) entered = [p.__enter__() for p in patches] try: tts_synthesize_and_submit.run("missing-job", "user-1", "v1", "你好", 1.0, "") finally: for p in reversed(patches): p.__exit__(None, None, None) session.commit.assert_not_called() session.close.assert_called_once() def test_cancelled_job_skipped(self): """Job 已 cancelled → 跳过不处理,不调用 TTS/MediaKit.""" from app.tasks.lipsync_tts import tts_synthesize_and_submit job = _make_fake_job(status="cancelled") session, patches = _apply_all_patches(job=job) entered = [p.__enter__() for p in patches] try: tts_synthesize_and_submit.run("job-1", "user-1", "v1", "你好", 1.0, "") finally: for p in reversed(patches): p.__exit__(None, None, None) # cancelled 不应 commit,不应触发 TTS/MediaKit session.commit.assert_not_called() session.close.assert_called_once() def test_happy_path_tts_to_mediakit(self): """完整正常流程:TTS 合成 → OSS 上传 → 签名 → 提交 MediaKit → submitted.""" from app.tasks.lipsync_tts import tts_synthesize_and_submit job = _make_fake_job() mk_client = MagicMock() mk_client.submit_lipsync.return_value = {"task_id": "mk-999"} session, patches = _apply_all_patches(job=job, mk_client=mk_client) entered = [p.__enter__() for p in patches] try: tts_synthesize_and_submit.run("job-1", "user-1", "v1", "你好世界", 1.0, "happy") finally: for p in reversed(patches): p.__exit__(None, None, None) assert job.status == "submitted" assert job.mediakit_task_id == "mk-999" assert job.error_code == "" mk_client.submit_lipsync.assert_called_once() call_kwargs = mk_client.submit_lipsync.call_args.kwargs assert call_kwargs["client_token"] == "job-1" # CosyVoice 临时 URL 经 _sign_media_url 透传(mock 统一追加 ?signed), # 自家 OSS 才会被重签,外部 URL 原样透传;job.audio_url 存原始临时 URL assert call_kwargs["audio_url"] == "https://tts/raw.mp3?signed" assert job.audio_url == "https://tts/raw.mp3" session.commit.assert_called() session.close.assert_called_once() def test_cosyvoice_error_marks_tts_synthesis_failed(self): """CosyVoiceError → failed, error_code=TTSSynthesisFailed.""" from app.tasks.lipsync_tts import tts_synthesize_and_submit from packages.application.cosyvoice_service import CosyVoiceError job = _make_fake_job() session, patches = _apply_all_patches(job=job, cosyvoice_error=CosyVoiceError("tts boom")) entered = [p.__enter__() for p in patches] try: tts_synthesize_and_submit.run("job-1", "user-1", "v1", "你好", 1.0, "") finally: for p in reversed(patches): p.__exit__(None, None, None) assert job.status == "failed" assert job.error_code == "TTSSynthesisFailed" session.close.assert_called_once() def test_value_error_marks_tts_invalid_param(self): """ValueError(参数错误)→ failed, error_code=TTSInvalidParam.""" from app.tasks.lipsync_tts import tts_synthesize_and_submit job = _make_fake_job() session, patches = _apply_all_patches(job=job, cosyvoice_side_effect=ValueError("bad param")) entered = [p.__enter__() for p in patches] try: tts_synthesize_and_submit.run("job-1", "user-1", "v1", "你好", -1.0, "") finally: for p in reversed(patches): p.__exit__(None, None, None) assert job.status == "failed" assert job.error_code == "TTSInvalidParam" session.close.assert_called_once() def test_no_audio_url_marks_tts_no_audio(self): """TTS 返回空 audio_url → failed, error_code=TTSNoAudio.""" from app.tasks.lipsync_tts import tts_synthesize_and_submit job = _make_fake_job() cosy = MagicMock() cosy.submit_synthesize_task.return_value = {"audio_url": ""} session, patches = _apply_all_patches(job=job, cosyvoice_service=cosy) entered = [p.__enter__() for p in patches] try: tts_synthesize_and_submit.run("job-1", "user-1", "v1", "你好", 1.0, "") finally: for p in reversed(patches): p.__exit__(None, None, None) assert job.status == "failed" assert job.error_code == "TTSNoAudio" session.close.assert_called_once() def test_oss_upload_failure_falls_back_to_temp_url(self): """OSS 上传失败 → 回退临时 URL,仍然 submitted.""" from app.tasks.lipsync_tts import tts_synthesize_and_submit job = _make_fake_job() storage = MagicMock() storage.public_url = "https://oss.example.com" storage.upload_file.side_effect = RuntimeError("oss down") mk_client = MagicMock() mk_client.submit_lipsync.return_value = {"task_id": "mk-7"} session, patches = _apply_all_patches(job=job, storage=storage, mk_client=mk_client) entered = [p.__enter__() for p in patches] try: tts_synthesize_and_submit.run("job-1", "user-1", "v1", "你好", 1.0, "") finally: for p in reversed(patches): p.__exit__(None, None, None) # 上传失败后 audio_url 回退为临时 TTS URL,仍继续提交到 MediaKit assert job.audio_url == "https://tts/raw.mp3" assert job.status == "submitted" assert job.mediakit_task_id == "mk-7" mk_client.submit_lipsync.assert_called_once() session.close.assert_called_once() def test_mediakit_error_marks_mediakit_unavailable(self): """MediaKit 提交失败 → failed, error_code=MediaKitUnavailable.""" from app.services.mediakit_client import MediaKitError from app.tasks.lipsync_tts import tts_synthesize_and_submit job = _make_fake_job() mk_err = MediaKitError("mk down", code="MediaKitUnavailable") session, patches = _apply_all_patches(job=job, mk_submit_error=mk_err) entered = [p.__enter__() for p in patches] try: tts_synthesize_and_submit.run("job-1", "user-1", "v1", "你好", 1.0, "") finally: for p in reversed(patches): p.__exit__(None, None, None) assert job.status == "failed" assert job.error_code == "MediaKitUnavailable" session.close.assert_called_once() def test_top_level_exception_marks_async_task_error(self): """顶层未预期异常 → failed, error_code=AsyncTaskError.""" from app.tasks.lipsync_tts import tts_synthesize_and_submit job = _make_fake_job() fake_app_db = ModuleType("app.db") fake_worker_db = ModuleType("worker_app.db") session, factory = _build_session(job) fake_app_db.SessionLocal = factory fake_worker_db.SessionLocal = factory # CosyVoiceService 在 __init__ 抛 RuntimeError(非 CosyVoiceError/ValueError) fake_cosy_mod = ModuleType("packages.application.cosyvoice_service") class _CosyVoiceErrorForTest(Exception): pass class _BoomService: def __init__(self): raise RuntimeError("top-level boom") fake_cosy_mod.CosyVoiceError = _CosyVoiceErrorForTest fake_cosy_mod.CosyVoiceService = _BoomService with patch.dict( sys.modules, { "app.db": fake_app_db, "worker_app.db": fake_worker_db, "packages.application.cosyvoice_service": fake_cosy_mod, }, ): tts_synthesize_and_submit.run("job-1", "user-1", "v1", "你好", 1.0, "") assert job.status == "failed" assert job.error_code == "AsyncTaskError" session.close.assert_called() class TestSignMediaUrl: """覆盖模块内 _sign_media_url 的所有分支(CI 增量覆盖率需要).""" def test_empty_url_returns_empty(self): from app.tasks.lipsync_tts import _sign_media_url assert _sign_media_url("") == "" assert _sign_media_url(None) is None def test_own_oss_url_signed(self): """自家 OSS URL → 调用 storage.get_download_url 签名.""" from app.tasks.lipsync_tts import _sign_media_url fake_storage = MagicMock() fake_storage.public_url = "https://oss.example.com/" fake_storage.get_download_url.return_value = "https://oss.example.com/a?sig=xyz" with patch("packages.shared.storage.get_shared_storage_service", return_value=fake_storage): result = _sign_media_url("https://oss.example.com/lipsync/a.mp3") assert result == "https://oss.example.com/a?sig=xyz" fake_storage.get_download_url.assert_called_once() def test_external_url_passthrough(self): """外部 URL(不是自家 OSS host)→ 原样透传,不签名.""" from app.tasks.lipsync_tts import _sign_media_url fake_storage = MagicMock() fake_storage.public_url = "https://oss.example.com/" with patch("packages.shared.storage.get_shared_storage_service", return_value=fake_storage): result = _sign_media_url("https://tts.example.com/raw.mp3") assert result == "https://tts.example.com/raw.mp3" fake_storage.get_download_url.assert_not_called() def test_storage_exception_falls_back(self): """storage 调用异常 → 降级原样返回,不抛错.""" from app.tasks.lipsync_tts import _sign_media_url with patch( "packages.shared.storage.get_shared_storage_service", side_effect=RuntimeError("storage down"), ): result = _sign_media_url("https://oss.example.com/a.mp3") assert result == "https://oss.example.com/a.mp3" def test_no_public_url_passthrough(self): """storage.public_url 为空 → 原样透传.""" from app.tasks.lipsync_tts import _sign_media_url fake_storage = MagicMock() fake_storage.public_url = "" with patch("packages.shared.storage.get_shared_storage_service", return_value=fake_storage): result = _sign_media_url("https://anything.example.com/a.mp3") assert result == "https://anything.example.com/a.mp3" fake_storage.get_download_url.assert_not_called() class TestPersistOutputVideoTask: """persist_output_video_task:下载 MediaKit 临时视频 → 上传自有 OSS → 更新 DB.""" def _make_persist_job(self, **kwargs): job = MagicMock() job.id = kwargs.get("job_id", "job-1") job.user_id = kwargs.get("user_id", "user-1") job.output_video_url = kwargs.get("output_video_url", "https://temp.mk/output.mp4") job.updated_at = None return job def _persist_patches(self, *, job, video_bytes=b"FAKEMP4", download_side_effect=None, upload_url=None): """统一 patch:SessionLocal、httpx.Client、storage、_sign_media_url.""" fake_app_db = ModuleType("app.db") fake_worker_db = ModuleType("worker_app.db") session, factory = _build_session(job) fake_app_db.SessionLocal = factory fake_worker_db.SessionLocal = factory # httpx.Client 上下文管理器 fake_response = MagicMock() fake_response.content = video_bytes fake_response.raise_for_status = MagicMock() fake_client = MagicMock() fake_client.get.return_value = fake_response fake_client_cm = MagicMock() fake_client_cm.__enter__ = MagicMock(return_value=fake_client) fake_client_cm.__exit__ = MagicMock(return_value=False) FakeHttpxClient = MagicMock(return_value=fake_client_cm) if download_side_effect is not None: fake_client.get.side_effect = download_side_effect # storage storage = MagicMock() storage.public_url = "https://oss.example.com/" storage.upload_file.return_value = upload_url or "https://oss.example.com/lipsync-outputs/user-1/job-1.mp4" # _sign_media_url 内部会调 storage.get_download_url,必须mock返回字符串 _upload_url = upload_url or "https://oss.example.com/lipsync-outputs/user-1/job-1.mp4" storage.get_download_url.return_value = _upload_url + "?signed" fake_httpx = ModuleType("httpx") fake_httpx.Client = FakeHttpxClient patches = [ patch.dict( sys.modules, {"app.db": fake_app_db, "worker_app.db": fake_worker_db, "httpx": fake_httpx}, ), patch("packages.shared.storage.get_shared_storage_service", return_value=storage), patch("app.tasks.lipsync_tts._sign_media_url", side_effect=lambda url: url + "?signed" if url else url), ] return session, fake_client, storage, patches def test_success_download_upload_updates_db(self): """正常路径:下载 temp_url → 上传 OSS → 签名 → 写回 DB commit.""" from app.tasks.lipsync_tts import persist_output_video_task job = self._make_persist_job(output_video_url="https://temp.mk/x.mp4") session, fake_client, storage, patches = self._persist_patches( job=job, video_bytes=b"VIDEODATA", upload_url="https://oss.example.com/lipsync-outputs/user-1/job-1.mp4" ) entered = [p.__enter__() for p in patches] try: persist_output_video_task("job-1", "user-1", "https://temp.mk/x.mp4") finally: for p in reversed(patches): p.__exit__(None, None, None) fake_client.get.assert_called_once_with("https://temp.mk/x.mp4") storage.upload_file.assert_called_once() call_args = storage.upload_file.call_args.args # 上传的 key 必须是 lipsync-outputs/{user_id}/{job_id}.mp4 assert call_args[1] == "lipsync-outputs/user-1/job-1.mp4" # upload_file 返回永久 URL,再被 _sign_media_url 追加 ?signed assert job.output_video_url == "https://oss.example.com/lipsync-outputs/user-1/job-1.mp4?signed" assert job.updated_at is not None session.commit.assert_called_once() session.close.assert_called_once() def test_download_failure_keeps_temp_url_no_commit(self): """下载失败(raise)→ 记录 warning、保留 temp_url、不抛异常.""" from app.tasks.lipsync_tts import persist_output_video_task job = self._make_persist_job(output_video_url="https://temp.mk/x.mp4") session, fake_client, storage, patches = self._persist_patches( job=job, download_side_effect=RuntimeError("network down") ) entered = [p.__enter__() for p in patches] try: persist_output_video_task("job-1", "user-1", "https://temp.mk/x.mp4") finally: for p in reversed(patches): p.__exit__(None, None, None) storage.upload_file.assert_not_called() # output_video_url 保持原值(temp_url) assert job.output_video_url == "https://temp.mk/x.mp4" # 内层 except 不会 commit # 注:若内部发生 commit 说明测试失败 session.close.assert_called_once() def test_empty_temp_url_skips_persist(self): """temp_url 为空 → 直接返回,不下载不上传.""" from app.tasks.lipsync_tts import persist_output_video_task job = self._make_persist_job(output_video_url="") session, fake_client, storage, patches = self._persist_patches(job=job) entered = [p.__enter__() for p in patches] try: persist_output_video_task("job-1", "user-1", "") finally: for p in reversed(patches): p.__exit__(None, None, None) fake_client.get.assert_not_called() storage.upload_file.assert_not_called() session.commit.assert_not_called() session.close.assert_called_once() def test_job_not_found_returns_early(self): """DB 中找不到 job → 直接返回,不抛错.""" from app.tasks.lipsync_tts import persist_output_video_task session, fake_client, storage, patches = self._persist_patches(job=None) entered = [p.__enter__() for p in patches] try: persist_output_video_task("missing", "user-1", "https://temp.mk/x.mp4") finally: for p in reversed(patches): p.__exit__(None, None, None) fake_client.get.assert_not_called() storage.upload_file.assert_not_called() session.commit.assert_not_called() session.close.assert_called_once() class TestLipsyncServiceRefreshCompletedAsyncPersist: """refresh_job_status 在 completed 分支异步转存的单元测试(补 0% 覆盖的 316~335 行).""" def test_refresh_completed_dispatches_persist_task(self): """completed 分支:设置 temp_url → commit → dispatch persist_output_video_task.apply_async.""" from app.services.lipsync_service import LipsyncService mock_job = MagicMock() mock_job.id = "job-1" mock_job.user_id = "user-1" mock_job.mediakit_task_id = "mk-1" mock_job.status = "submitted" mock_job.output_video_url = "" mock_job.output_duration = 0.0 mock_db = MagicMock() mock_query = MagicMock() mock_filter = MagicMock() mock_filter.first.return_value = mock_job mock_query.filter.return_value = mock_filter mock_db.query.return_value = mock_query mock_client = MagicMock() mock_client.get_task_status.return_value = { "status": "completed", "result": {"video_url": "https://temp.mk/out.mp4", "duration": 25.5}, } fake_persist_task = MagicMock() svc = LipsyncService(mock_db, client=mock_client, cosyvoice_service=MagicMock()) with patch.dict("sys.modules", {}): # 直接 patch 懒 import 路径 with patch("app.tasks.lipsync_tts.persist_output_video_task", fake_persist_task, create=False): # 但懒 import 发生在函数内部 from app.tasks.lipsync_tts import persist_output_video_task # 通过 patch sys.modules 的方式提供 import sys as _sys fake_mod = MagicMock() fake_mod.persist_output_video_task = fake_persist_task _sys.modules["app.tasks.lipsync_tts"] = fake_mod try: result = svc.refresh_job_status("job-1", "user-1") finally: _sys.modules.pop("app.tasks.lipsync_tts", None) assert result.status == "completed" assert result.output_video_url == "https://temp.mk/out.mp4" assert result.output_duration == 25.5 mock_db.commit.assert_called() # 必须在 commit 之后 dispatch fake_persist_task.apply_async.assert_called_once() kwargs = fake_persist_task.apply_async.call_args.kwargs assert kwargs["args"] == ("job-1", "user-1", "https://temp.mk/out.mp4") def test_refresh_completed_dispatch_exception_does_not_break_return(self): """apply_async 抛异常(如 Celery 不可用)→ 捕获 warning,仍返回 completed job.""" from app.services.lipsync_service import LipsyncService mock_job = MagicMock() mock_job.id = "job-2" mock_job.user_id = "user-1" mock_job.mediakit_task_id = "mk-2" mock_job.status = "submitted" mock_job.output_video_url = "" mock_job.output_duration = 0.0 mock_db = MagicMock() mock_query = MagicMock() mock_filter = MagicMock() mock_filter.first.return_value = mock_job mock_query.filter.return_value = mock_filter mock_db.query.return_value = mock_query mock_client = MagicMock() mock_client.get_task_status.return_value = { "status": "completed", "result": {"video_url": "https://temp.mk/out2.mp4", "duration": 10.0}, } fake_persist_task = MagicMock() fake_persist_task.apply_async.side_effect = ConnectionError("celery down") svc = LipsyncService(mock_db, client=mock_client, cosyvoice_service=MagicMock()) import sys as _sys fake_mod = MagicMock() fake_mod.persist_output_video_task = fake_persist_task _sys.modules["app.tasks.lipsync_tts"] = fake_mod try: result = svc.refresh_job_status("job-2", "user-1") finally: _sys.modules.pop("app.tasks.lipsync_tts", None) # 即便 dispatch 失败,主流程不受影响:仍然返回 completed + temp_url assert result.status == "completed" assert result.output_video_url == "https://temp.mk/out2.mp4" fake_persist_task.apply_async.assert_called_once() class TestRefreshJobStatusDbRollback: """#1867 三层防御 ②:refresh_job_status DB commit 失败时 rollback 不 raise.""" def test_refresh_db_commit_failure_rollback_and_return(self): """commit 抛异常 → rollback,不 raise,返回当前 job 对象让下次轮询重试.""" from app.services.lipsync_service import LipsyncService mock_job = MagicMock() mock_job.id = "job-rollback" mock_job.user_id = "u" mock_job.mediakit_task_id = "mk" mock_job.status = "submitted" mock_job.output_video_url = "" mock_job.output_duration = 0.0 mock_db = MagicMock() mock_query = MagicMock() mock_filter = MagicMock() mock_filter.first.return_value = mock_job mock_query.filter.return_value = mock_filter mock_db.query.return_value = mock_query mock_db.commit.side_effect = Exception("DB connection lost") mock_client = MagicMock() mock_client.get_task_status.return_value = { "status": "completed", "result": {"video_url": "https://temp.mk/out.mp4", "duration": 5.0}, } # refresh_job_status 内部懒 import persist_output_video_task # 用 patch 替换模块属性,create=True 因为原模块里它是 shared_task,patch 后是 MagicMock with patch("app.tasks.lipsync_tts.persist_output_video_task", MagicMock(), create=True): svc = LipsyncService(mock_db, client=mock_client, cosyvoice_service=MagicMock()) result_job = svc.refresh_job_status("job-rollback", "u") assert result_job is mock_job mock_db.rollback.assert_called() def test_refresh_intermediate_status_commits(self): """中间状态(running/processing 等)→ 更新 status 和 updated_at,正常 commit.""" from app.services.lipsync_service import LipsyncService mock_job = MagicMock() mock_job.id = "job-mid" mock_job.user_id = "u" mock_job.mediakit_task_id = "mk" mock_job.status = "submitted" mock_job.output_video_url = "" mock_job.output_duration = 0.0 mock_db = MagicMock() mock_query = MagicMock() mock_filter = MagicMock() mock_filter.first.return_value = mock_job mock_query.filter.return_value = mock_filter mock_db.query.return_value = mock_query mock_client = MagicMock() mock_client.get_task_status.return_value = {"status": "processing"} svc = LipsyncService(mock_db, client=mock_client, cosyvoice_service=MagicMock()) result_job = svc.refresh_job_status("job-mid", "u") assert result_job is mock_job assert mock_job.status == "processing" mock_db.commit.assert_called() mock_db.rollback.assert_not_called() def test_refresh_failed_commits_error_info(self): """failed 状态 → 写入 error_message/error_code/completed_at,commit.""" from app.services.lipsync_service import LipsyncService mock_job = MagicMock() mock_job.id = "job-fail" mock_job.user_id = "u" mock_job.mediakit_task_id = "mk" mock_job.status = "submitted" mock_job.output_video_url = "" mock_job.output_duration = 0.0 mock_db = MagicMock() mock_query = MagicMock() mock_filter = MagicMock() mock_filter.first.return_value = mock_job mock_query.filter.return_value = mock_filter mock_db.query.return_value = mock_query mock_client = MagicMock() mock_client.get_task_status.return_value = { "status": "failed", "error": {"code": "BadAudio", "message": "音频损坏"}, } svc = LipsyncService(mock_db, client=mock_client, cosyvoice_service=MagicMock()) result_job = svc.refresh_job_status("job-fail", "u") assert result_job.status == "failed" assert mock_job.error_message == "音频损坏" assert mock_job.error_code == "BadAudio" mock_db.commit.assert_called() class TestPollMediakitStatus: """#1867 三层防御 ③:Celery 兜底轮询任务单元测试.""" def _make_session(self, job): mock_db = MagicMock() mock_query = MagicMock() mock_filter = MagicMock() mock_filter.first.return_value = job mock_query.filter.return_value = mock_filter mock_db.query.return_value = mock_query return mock_db def _patch_sessions(self, mock_db): """用 patch.dict 注入 fake SessionLocal 到 worker_app.db 和 app.db.""" from types import ModuleType fake_app_db = ModuleType("app.db") fake_worker_db = ModuleType("worker_app.db") fake_sl = MagicMock(return_value=mock_db) fake_app_db.SessionLocal = fake_sl fake_worker_db.SessionLocal = fake_sl fake_worker_pkg = ModuleType("worker_app") return patch.dict( sys.modules, { "worker_app": fake_worker_pkg, "worker_app.db": fake_worker_db, "app.db": fake_app_db, }, ) def _run(self, mock_db, job_id="j", user_id="u1", fake_self=None): from app.tasks.lipsync_tts import poll_mediakit_status self_ref = fake_self if fake_self is not None else MagicMock() real = poll_mediakit_status._get_current_object() return real.run.__func__(self_ref, job_id=job_id, user_id=user_id), self_ref def test_poll_job_not_found_returns(self): mock_db = self._make_session(None) with self._patch_sessions(mock_db): self._run(mock_db, job_id="missing") mock_db.close.assert_called_once() def test_poll_terminal_completed_skips_mediakit(self): mock_job = MagicMock(status="completed", mediakit_task_id="mk-1") mock_db = self._make_session(mock_job) with self._patch_sessions(mock_db), patch("app.services.mediakit_client.get_mediakit_client") as fg: self._run(mock_db) fg.assert_not_called() mock_db.close.assert_called_once() def test_poll_terminal_failed_skips(self): mock_job = MagicMock(status="failed", mediakit_task_id="mk-1") mock_db = self._make_session(mock_job) with self._patch_sessions(mock_db), patch("app.services.mediakit_client.get_mediakit_client") as fg: self._run(mock_db, job_id="j1b") fg.assert_not_called() def test_poll_terminal_cancelled_skips(self): mock_job = MagicMock(status="cancelled", mediakit_task_id="mk-1") mock_db = self._make_session(mock_job) with self._patch_sessions(mock_db), patch("app.services.mediakit_client.get_mediakit_client") as fg: self._run(mock_db, job_id="j1c") fg.assert_not_called() def test_poll_no_mk_task_id_returns(self): mock_job = MagicMock(status="submitted", mediakit_task_id=None) mock_db = self._make_session(mock_job) with self._patch_sessions(mock_db): self._run(mock_db, job_id="j6") mock_db.close.assert_called_once() def test_poll_succeeded_marks_completed(self): mock_job = MagicMock(status="submitted", mediakit_task_id="mk-2", output_video_url="", output_duration=0.0) mock_db = self._make_session(mock_job) mock_client = MagicMock() mock_client.get_task_status.return_value = { "status": "succeeded", "result": {"video_url": "https://temp/x.mp4", "duration": 12.0}, } with ( self._patch_sessions(mock_db), patch("app.services.mediakit_client.get_mediakit_client", return_value=mock_client), patch("app.services.lipsync_service.LipsyncService") as MS, ): svc = MS.return_value svc._persist_output_video.return_value = "https://oss/x.mp4" self._run(mock_db, job_id="j2") assert mock_job.status == "completed" assert mock_job.output_video_url == "https://oss/x.mp4" assert mock_job.output_duration == 12.0 svc._persist_output_video.assert_called_once_with("https://temp/x.mp4", "j2", "u1") mock_db.commit.assert_called() def test_poll_succeeded_persist_fail_fallback(self): mock_job = MagicMock(status="submitted", mediakit_task_id="mk-3", output_video_url="", output_duration=0.0) mock_db = self._make_session(mock_job) mock_client = MagicMock() mock_client.get_task_status.return_value = { "status": "succeeded", "result": {"video_url": "https://temp/y.mp4", "duration": 8.0}, } with ( self._patch_sessions(mock_db), patch("app.services.mediakit_client.get_mediakit_client", return_value=mock_client), patch("app.services.lipsync_service.LipsyncService") as MS, ): svc = MS.return_value svc._persist_output_video.side_effect = Exception("OSS down") self._run(mock_db, job_id="j3") assert mock_job.status == "completed" assert mock_job.output_video_url == "https://temp/y.mp4" def test_poll_failed_marks_failed(self): mock_job = MagicMock(status="running", mediakit_task_id="mk-4", error_message="", error_code="") mock_db = self._make_session(mock_job) mock_client = MagicMock() mock_client.get_task_status.return_value = { "status": "failed", "error": {"code": "TaskFailed", "message": "音频错误"}, } with ( self._patch_sessions(mock_db), patch("app.services.mediakit_client.get_mediakit_client", return_value=mock_client), ): self._run(mock_db, job_id="j4") assert mock_job.status == "failed" assert mock_job.error_message == "音频错误" assert mock_job.error_code == "TaskFailed" def test_poll_error_status_marks_failed(self): mock_job = MagicMock(status="submitted", mediakit_task_id="mk-4b", error_message="", error_code="") mock_db = self._make_session(mock_job) mock_client = MagicMock() mock_client.get_task_status.return_value = { "status": "error", "error": {"code": "Internal", "message": "内部错误"}, } with ( self._patch_sessions(mock_db), patch("app.services.mediakit_client.get_mediakit_client", return_value=mock_client), ): self._run(mock_db, job_id="j4b") assert mock_job.status == "failed" assert mock_job.error_message == "内部错误" def test_poll_running_triggers_retry(self): mock_job = MagicMock(status="submitted", mediakit_task_id="mk-5") mock_db = self._make_session(mock_job) mock_client = MagicMock() mock_client.get_task_status.return_value = {"status": "running"} class _Retry(Exception): pass fake_self = MagicMock() fake_self.retry.side_effect = _Retry() with ( self._patch_sessions(mock_db), patch("app.services.mediakit_client.get_mediakit_client", return_value=mock_client), patch("app.services.lipsync_service.LipsyncService"), ): try: self._run(mock_db, job_id="j5", fake_self=fake_self) assert False, "expected retry" except _Retry: pass assert mock_job.updated_at is not None assert mock_job.status == "running" assert fake_self.retry.call_count >= 1 class TestTtsSubmitDispatchesPoll: """#1867 三层防御 ③:tts_synthesize_and_submit 成功后链式触发 poll.""" def _run_tts(self, job, mk_client, fake_poll): """运行 tts_synthesize_and_submit.run,直接在用户函数 __globals__ 替换 poll_mediakit_status. pytest --import-mode=importlib 下,patch("app.tasks.lipsync_tts.xxx") 拿到的模块对象 可能与 celery autoretry 包装下用户函数的 __globals__ 不一致(之前的 Poll 测试通过 patch.dict(sys.modules, ...) 注入 fake 模块,进一步加剧了模块命名空间不一致)。 直接 patch _orig_run.__func__.__globals__ 是最稳的方式。 """ from app.tasks.lipsync_tts import tts_synthesize_and_submit session, patches = _apply_all_patches(job=job, mk_client=mk_client) for p in patches: p.__enter__() # 定位 autoretry 包装下的真实用户函数 real_tts = tts_synthesize_and_submit._get_current_object() fn = real_tts._orig_run while hasattr(fn, "__func__"): fn = fn.__func__ g = fn.__globals__ orig_poll = g.get("poll_mediakit_status") g["poll_mediakit_status"] = fake_poll try: tts_synthesize_and_submit.run("job-1", "user-1", "v1", "你好", 1.0, "") finally: if orig_poll is not None: g["poll_mediakit_status"] = orig_poll else: g.pop("poll_mediakit_status", None) for p in reversed(patches): p.__exit__(None, None, None) return session def test_submit_success_dispatches_poll(self): """提交 MediaKit 成功后链式触发 poll_mediakit_status.apply_async (countdown=10).""" job = _make_fake_job() job.emotion = "" mk_client = MagicMock() mk_client.submit_lipsync.return_value = {"success": True, "task_id": "mk-poll-1"} fake_poll = MagicMock() self._run_tts(job, mk_client, fake_poll) assert job.status == "submitted" assert job.mediakit_task_id == "mk-poll-1" fake_poll.apply_async.assert_called_once() kwargs = fake_poll.apply_async.call_args.kwargs assert kwargs["kwargs"] == {"job_id": "job-1", "user_id": "user-1"} assert kwargs["countdown"] == 10 def test_submit_failure_no_poll(self): from app.services.mediakit_client import MediaKitError job = _make_fake_job() job.emotion = "" mk_client = MagicMock() mk_client.submit_lipsync.side_effect = MediaKitError("bad", code="X") fake_poll = MagicMock() self._run_tts(job, mk_client, fake_poll) assert job.status == "failed" fake_poll.apply_async.assert_not_called() def test_poll_dispatch_fail_does_not_break(self): job = _make_fake_job() job.emotion = "" mk_client = MagicMock() mk_client.submit_lipsync.return_value = {"success": True, "task_id": "mk-poll-2"} fake_poll = MagicMock() fake_poll.apply_async.side_effect = ConnectionError("celery down") session = self._run_tts(job, mk_client, fake_poll) assert job.status == "submitted" assert job.mediakit_task_id == "mk-poll-2" session.commit.assert_called()