"""#1970 MuseTalk Flask 服务端 8 项工程 bug 修复单测. 覆盖 deploy/gpu_worker/musetalk_server.py(独立部署脚本,按文件路径动态加载): 1. threaded=True 启动,/health 在推理阻塞时仍可达 2. fps 兜底:ffprobe 返回 0 或失败时使用 default_fps 3. ffmpeg 走 subprocess.run(check=True),失败抛 RuntimeError 4. 并发锁:推理期间第二请求立即 503 5. 推理超时:超过 MUSE_INFERENCE_TIMEOUT 返回 504 6. 结果文件清理:临时目录在请求结束(成功/失败)后删除 7. 文件大小限制:超过限制返回 413,空文件返回 400 8. /cancel 端点:终止当前推理,清理临时文件 """ from __future__ import annotations import importlib.util import io import os import shutil import sys import threading import time from pathlib import Path from unittest import mock import pytest # 检查 Flask 是否可用(CI 环境可能没装) try: import flask # noqa: F401 HAS_FLASK = True except ImportError: HAS_FLASK = False pytestmark = pytest.mark.skipif(not HAS_FLASK, reason="Flask 未安装(gpu_worker 独立部署依赖)") ROOT = Path(__file__).resolve().parents[2] SERVER_PATH = ROOT / "deploy" / "gpu_worker" / "musetalk_server.py" def _load_server_module(name: str = "musetalk_server_test"): """加载 musetalk_server.py 为独立模块.""" # 避免重复注册 if name in sys.modules: del sys.modules[name] spec = importlib.util.spec_from_file_location(name, SERVER_PATH) mod = importlib.util.module_from_spec(spec) sys.modules[name] = mod spec.loader.exec_module(mod) return mod @pytest.fixture def server(tmp_path, monkeypatch): """加载一个干净的 musetalk_server 模块,使用独立临时目录和端口.""" if not HAS_FLASK: pytest.skip("Flask 未安装(gpu_worker 独立部署依赖)") monkeypatch.setenv("MUSE_TEMP_DIR", str(tmp_path / "musetalk_temp")) monkeypatch.setenv("MUSE_PORT", "0") monkeypatch.setenv("MUSE_INFERENCE_TIMEOUT", "2") monkeypatch.setenv("MUSE_VIDEO_MAX_MB", "1") monkeypatch.setenv("MUSE_AUDIO_MAX_MB", "1") monkeypatch.setenv("MUSE_DEFAULT_FPS", "25.0") mod_name = f"musetalk_server_test_{os.getpid()}_{id(tmp_path)}" mod = _load_server_module(mod_name) # 确保配置已更新 mod.Config.temp_dir = str(tmp_path / "musetalk_temp") mod.Config.inference_timeout = 2.0 mod.Config.video_max_mb = 1 mod.Config.audio_max_mb = 1 mod.Config.default_fps = 25.0 Path(mod.Config.temp_dir).mkdir(parents=True, exist_ok=True) # 重置全局状态 mod.inference_lock = threading.Lock() mod.current_task = {"task_id": None, "process": None, "start_time": 0.0} return mod # ── 1. Flask threaded=True ────────────────────────────────────────── def test_flask_run_uses_threaded(server): """验证 app.run 调用时 threaded=True.""" with mock.patch.object(server.app, "run") as mock_run: server.main() mock_run.assert_called_once() call_kwargs = mock_run.call_args assert call_kwargs.kwargs.get("threaded") is True # ── 2. fps=0 兜底 ─────────────────────────────────────────────────── def test_get_video_fps_fallback_on_zero(server, tmp_path): """ffprobe 返回 0/1 时兜底为 default_fps.""" fake_video = tmp_path / "fake.mp4" fake_video.write_bytes(b"fake") with mock.patch("subprocess.check_output", return_value=b"0/1"): fps = server._get_video_fps(fake_video) assert fps == 25.0 def test_get_video_fps_normal(server, tmp_path): """正常 fps 解析.""" fake_video = tmp_path / "fake.mp4" fake_video.write_bytes(b"fake") with mock.patch("subprocess.check_output", return_value=b"30/1"): fps = server._get_video_fps(fake_video) assert abs(fps - 30.0) < 0.01 def test_get_video_fps_exception_fallback(server, tmp_path): """ffprobe 异常时兜底 default_fps.""" fake_video = tmp_path / "fake.mp4" fake_video.write_bytes(b"fake") with mock.patch("subprocess.check_output", side_effect=Exception("no ffprobe")): fps = server._get_video_fps(fake_video) assert fps == 25.0 # ── 3. ffmpeg 错误检查 ────────────────────────────────────────────── def test_run_ffmpeg_raises_on_nonzero_exit(server): """ffmpeg 返回非零应抛 RuntimeError.""" import subprocess with mock.patch( "subprocess.run", side_effect=subprocess.CalledProcessError(1, "ffmpeg", stderr=b"decode error"), ): with pytest.raises(RuntimeError, match="ffmpeg 失败"): server._run_ffmpeg(["ffmpeg", "-i", "in", "out"]) def test_run_ffmpeg_raises_on_timeout(server): """ffmpeg 超时应抛 RuntimeError.""" import subprocess with mock.patch("subprocess.run", side_effect=subprocess.TimeoutExpired("ffmpeg", 10)): with pytest.raises(RuntimeError, match="ffmpeg 超时"): server._run_ffmpeg(["ffmpeg", "-i", "in", "out"], timeout=10) # ── 4. 并发锁 503 ────────────────────────────────────────────────── def test_inference_returns_503_when_busy(server): """推理期间第二请求立即 503.""" server.inference_lock.acquire() server.current_task["task_id"] = "task-busy" server.current_task["start_time"] = time.time() try: with server.app.test_client() as c: resp = c.post( "/inference", data={ "video": (io.BytesIO(b"v" * 100), "v.mp4"), "audio": (io.BytesIO(b"a" * 100), "a.wav"), }, content_type="multipart/form-data", ) assert resp.status_code == 503 assert resp.get_json()["status"] == "busy" finally: server.inference_lock.release() server.current_task = {"task_id": None, "process": None, "start_time": 0.0} # ── 5. 推理超时 504 ───────────────────────────────────────────────── def test_inference_timeout_returns_504(server): """推理超时返回 504.""" def slow_inference(*args, **kwargs): time.sleep(10) # 远超 2s 超时 with mock.patch.object(server, "_run_inference", side_effect=slow_inference): with server.app.test_client() as c: resp = c.post( "/inference", data={ "video": (io.BytesIO(b"v" * 100), "v.mp4"), "audio": (io.BytesIO(b"a" * 100), "a.wav"), }, content_type="multipart/form-data", ) assert resp.status_code == 504 assert "超时" in resp.get_json()["error"] # ── 6. 临时文件清理 ────────────────────────────────────────────────── def test_temp_files_cleaned_after_success(server, tmp_path): """推理成功后临时目录被清理.""" def fake_inference(video_path, audio_path, output_path): output_path.parent.mkdir(parents=True, exist_ok=True) output_path.write_bytes(b"v" * 2048) with mock.patch.object(server, "_run_inference", side_effect=fake_inference): with server.app.test_client() as c: resp = c.post( "/inference", data={ "video": (io.BytesIO(b"v" * 100), "v.mp4"), "audio": (io.BytesIO(b"a" * 100), "a.wav"), "task_id": "task-cleanup-ok", }, content_type="multipart/form-data", ) # send_file 返回 200 或推理异常 500 assert resp.status_code in (200, 500) task_dir = Path(server.Config.temp_dir) / "task-cleanup-ok" assert not task_dir.exists(), f"临时目录 {task_dir} 应被清理" def test_temp_files_cleaned_after_failure(server, tmp_path): """推理失败后临时目录也被清理.""" def failing_inference(*args, **kwargs): raise RuntimeError("MuseTalk crash") with mock.patch.object(server, "_run_inference", side_effect=failing_inference): with server.app.test_client() as c: resp = c.post( "/inference", data={ "video": (io.BytesIO(b"v" * 100), "v.mp4"), "audio": (io.BytesIO(b"a" * 100), "a.wav"), "task_id": "task-cleanup-fail", }, content_type="multipart/form-data", ) assert resp.status_code == 500 task_dir = Path(server.Config.temp_dir) / "task-cleanup-fail" assert not task_dir.exists() # ── 7. 文件大小限制 ───────────────────────────────────────────────── def test_oversize_video_returns_413(server): """视频超过大小限制返回 413.""" big_video = b"v" * (2 * 1024 * 1024) # 2MB > 1MB limit with server.app.test_client() as c: resp = c.post( "/inference", data={ "video": (io.BytesIO(big_video), "v.mp4"), "audio": (io.BytesIO(b"a" * 100), "a.wav"), }, content_type="multipart/form-data", ) assert resp.status_code == 413 assert "超过限制" in resp.get_json()["error"] def test_empty_file_returns_400(server): """空文件返回 400.""" with server.app.test_client() as c: resp = c.post( "/inference", data={ "video": (io.BytesIO(b""), "v.mp4"), "audio": (io.BytesIO(b"a" * 100), "a.wav"), }, content_type="multipart/form-data", ) assert resp.status_code in (400, 413) assert "为空" in resp.get_json().get("error", "") or "超过限制" in resp.get_json().get("error", "") def test_missing_file_returns_400(server): """缺少必要文件返回 400.""" with server.app.test_client() as c: resp = c.post( "/inference", data={"video": (io.BytesIO(b"v" * 100), "v.mp4")}, content_type="multipart/form-data", ) assert resp.status_code == 400 # ── 8. /cancel 端点 ───────────────────────────────────────────────── def test_cancel_no_running_task(server): """无任务时 /cancel 返回提示.""" with server.app.test_client() as c: resp = c.post("/cancel") assert resp.status_code == 200 assert "无正在运行" in resp.get_json()["message"] def test_cancel_terminates_running_task(server, tmp_path): """有任务时 /cancel 清理临时目录并重置状态.""" task_dir = Path(server.Config.temp_dir) / "task-cancel" task_dir.mkdir(parents=True, exist_ok=True) (task_dir / "some_file.txt").write_text("temp") server.current_task["task_id"] = "task-cancel" server.current_task["start_time"] = time.time() server.current_task["process"] = "inference_thread" with server.app.test_client() as c: resp = c.post("/cancel") assert resp.status_code == 200 assert "已取消" in resp.get_json()["message"] assert not task_dir.exists() assert server.current_task["task_id"] is None assert server.current_task["process"] is None assert server.current_task["start_time"] == 0.0