32473485d7
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 1s
CI/CD Pipeline / Check push changed paths (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 Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
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 / 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
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 2m0s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 2m44s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 2m51s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 2m57s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 6m34s
AI Code Review / AI Code Review (pull_request) Successful in 6m45s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 8m2s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 8m9s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 9m43s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 10m36s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 20m59s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 1s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 23s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 1m18s
服务端新建 deploy/gpu_worker/musetalk_server.py(替代原 ~/projects/MuseTalk/worker.py): 1. Flask app.run(threaded=True):推理阻塞时 /health 仍可达 2. _get_video_fps 兜底:ffprobe 返回 0 或失败时 fallback 到 default_fps(25) 3. _run_ffmpeg 统一封装:subprocess.run(check=True) + timeout,失败/超时抛 RuntimeError 4. inference_lock 并发锁:多请求同时到达时第二请求立即 503 5. 推理超时控制:thread.join(timeout=inference_timeout) 默认 600s,超时返回 504 6. finally 块清理临时目录:成功/失败/超时都删除 task_dir 7. 无人脸检测兜底:_run_inference 中帧提取后校验,无帧直接抛错返回 500 8. 上传大小限制:视频 <=100MB / 音频 <=20MB,超限返回 413 9. 新增 POST /cancel 端点:终止当前推理、清理临时文件、释放锁 10. GET /health 返回 GPU 显存信息(nvidia-smi)+ 当前任务状态 客户端 deploy/gpu_worker/gpu_worker.py 配套: - _call_musetalk 超时后 POST /cancel 终止服务端僵尸推理 - _call_musetalk 返回 (ok, duration, err, retryable) 四元组 - _handle_task 仅 retryable=True 时重试,4xx/短视频等确定性失败直接上报 - 新增 _cancel_musetalk_task 辅助方法 测试:新增 15 个单测覆盖服务端全部修复点;全量 15854 passed / 28 skipped 部署提醒:用户需在 RTX2060 上 wget 新 musetalk_server.py 替换旧 worker.py 并重启服务。
324 lines
12 KiB
Python
324 lines
12 KiB
Python
"""#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
|