"""#1970 GPU Worker 修复单测. 覆盖 deploy/gpu_worker/gpu_worker.py(独立部署脚本,不在 apps/packages 包内, 按文件路径动态加载): 1. 默认配置:REQUEST_TIMEOUT=900 / TASK_MAX_RETRY=1 / 心跳 30s / 最短 3s; 2. 推理期心跳线程 POST /gpu/register 带 task_id,任务结束能停; 3. <3s 短视频直接上报失败,不调用 MuseTalk; 4. _call_musetalk 仅对 5xx/网络瞬时错误标记 retryable,4xx 不重试; 5. _handle_task 只对 retryable 错误本地重试 1 次。 """ from __future__ import annotations import importlib.util import os import sys import time from pathlib import Path from unittest import mock import pytest ROOT = Path(__file__).resolve().parents[2] WORKER_PATH = ROOT / "deploy" / "gpu_worker" / "gpu_worker.py" def _load_worker_module(): spec = importlib.util.spec_from_file_location("gpu_worker_standalone_1970", WORKER_PATH) mod = importlib.util.module_from_spec(spec) sys.modules[spec.name] = mod spec.loader.exec_module(mod) return mod @pytest.fixture def worker(): return _load_worker_module() # ── 默认配置 ─────────────────────────────────────────────────────── def test_config_defaults_900_and_retry_one(monkeypatch): """CI/本机若显式导出过这些 env,说明是运维覆盖,不应拿默认值断言; 因此只在四个 env 全部缺失时校验脚本内置默认值(#1970:900/1/30/3)。""" keys = ( "REQUEST_TIMEOUT", "TASK_MAX_RETRY", "TASK_HEARTBEAT_INTERVAL", "MIN_VIDEO_DURATION_SECONDS", ) if any(k in os.environ for k in keys): pytest.skip("环境显式设置了 worker 超时/重试变量,跳过默认值断言") for key in keys: monkeypatch.delenv(key, raising=False) mod = _load_worker_module() assert mod.Config.request_timeout == 900.0 assert mod.Config.task_max_retry == 1 assert mod.Config.task_heartbeat_interval == 30.0 assert mod.Config.min_video_duration_seconds == 3.0 # ── register 携带 task_id ────────────────────────────────────────── def test_register_payload_includes_task_id_only_when_provided(worker, monkeypatch): captured = [] class _Resp: status_code = 200 def json(self): return {"worker_id": captured[-1]["worker_id"], "cancel_task": False} def _fake_post(url, json=None, headers=None, timeout=None): captured.append(json) return _Resp() monkeypatch.setattr(worker.requests, "post", _fake_post) monkeypatch.setattr(worker, "_check_musetalk_health", lambda: (True, {})) _ok, _cancel = worker._register("task-abc") assert _ok is True assert _cancel is False assert captured[-1]["task_id"] == "task-abc" assert captured[-1]["worker_id"] worker._register() # 空闲心跳不带 task_id assert "task_id" not in captured[-1] # ── 推理期心跳线程 ───────────────────────────────────────────────── def test_task_heartbeat_thread_sends_and_stops(worker, monkeypatch): calls = [] def _fake_register(task_id=None): calls.append(task_id) return True, False monkeypatch.setattr(worker, "_register", _fake_register) hb = worker.TaskHeartbeat("task-hb1", interval=5) hb.start() time.sleep(0.3) # 启动后立即发一次 hb.stop() hb.join(timeout=2) assert not hb.is_alive() assert calls and all(c == "task-hb1" for c in calls) def test_task_heartbeat_cancel_calls_musetalk_cancel(worker, monkeypatch): """心跳响应 cancel_task=True → 调 _cancel_musetalk 并设置 cancelled 标志。""" cancel_calls = [] monkeypatch.setattr(worker, "_register", lambda *a, **k: (True, True)) monkeypatch.setattr(worker, "_cancel_musetalk", lambda: cancel_calls.append(1)) hb = worker.TaskHeartbeat("task-cancel-1", interval=5) hb.start() hb.join(timeout=2) # 检测到取消后线程自行 return assert not hb.is_alive() assert hb.cancelled is True assert cancel_calls == [1] def test_handle_task_reports_cancelled_after_musetalk_abort(worker, monkeypatch): """推理被 /cancel 终止后,hb.cancelled=True → 上报失败而非重试。""" reports = [] monkeypatch.setattr(worker, "_register", lambda *a, **k: (True, False)) monkeypatch.setattr(worker, "_download", lambda url, path: True) monkeypatch.setattr(worker, "_probe_duration", lambda path: 12.0) # 模拟推理被终止(/inference 返回错误) monkeypatch.setattr(worker, "_call_musetalk", lambda v, a, o: (False, 0.0, "推理被取消", False)) monkeypatch.setattr( worker, "_report_result", lambda task_id, success, duration=0.0, error_msg="": reports.append(error_msg) or True, ) # 让 TaskHeartbeat 在主线程检查时报告已取消 orig_hb_init = worker.TaskHeartbeat def _hb(task_id, interval): h = orig_hb_init(task_id, interval) h.cancelled = True return h monkeypatch.setattr(worker, "TaskHeartbeat", _hb) worker._handle_task({"task_id": "t-canceled", "video_url": "u", "audio_url": "u"}) assert reports == ["用户取消任务"] # ── 短视频前置拦截 ───────────────────────────────────────────────── def test_handle_task_short_video_reports_failed_without_inference(worker, monkeypatch, tmp_path): video = tmp_path / "input.mp4" video.write_bytes(b"fake-mp4-bytes") audio = tmp_path / "input_audio.bin" audio.write_bytes(b"fake-audio") reports = [] monkeypatch.setattr(worker, "_register", lambda *a, **k: (True, False)) monkeypatch.setattr(worker, "_download", lambda url, path: True) # ffprobe 读出 1.2s → 低于 3s 阈值 monkeypatch.setattr(worker, "_probe_duration", lambda path: 1.2) def _boom(*a, **k): raise AssertionError("短视频不应调用 MuseTalk 推理") monkeypatch.setattr(worker, "_call_musetalk", _boom) monkeypatch.setattr( worker, "_report_result", lambda task_id, success, duration=0.0, error_msg="": reports.append((task_id, success, error_msg)) or True, ) task = { "task_id": "task-short", "video_url": "https://example.com/v.mp4", "audio_url": "https://example.com/a.bin", } worker._handle_task(task) assert len(reports) == 1 tid, ok, err = reports[0] assert tid == "task-short" assert ok is False assert "视频过短" in err assert "3" in err def test_handle_task_probe_failure_does_not_block(worker, monkeypatch): """ffprobe 不可用(duration=0.0)时不能误杀,应继续推理.""" reports = [] monkeypatch.setattr(worker, "_register", lambda *a, **k: (True, False)) monkeypatch.setattr(worker, "_download", lambda url, path: True) monkeypatch.setattr(worker, "_probe_duration", lambda path: 0.0) monkeypatch.setattr( worker, "_call_musetalk", lambda v, a, o: (True, 8.0, "", False), ) uploaded = [] monkeypatch.setattr( worker, "_report_success_with_file", lambda task_id, duration, path: uploaded.append((task_id, duration)), ) monkeypatch.setattr(worker, "_report_result", lambda *a, **k: True) worker._handle_task({"task_id": "task-probe0", "video_url": "u", "audio_url": "u"}) assert uploaded == [("task-probe0", 8.0)] assert reports == [] # ── 重试语义:仅瞬时错误重试 ─────────────────────────────────────── def test_call_musetalk_4xx_not_retryable_5xx_retryable(worker, monkeypatch, tmp_path): video = tmp_path / "v.mp4" audio = tmp_path / "a.bin" video.write_bytes(b"v") audio.write_bytes(b"a") out = tmp_path / "o.mp4" class _Resp: def __init__(self, code, body=b"x" * 2048): self.status_code = code self.content = body self.text = "err" # 4xx:确定性失败,不重试 monkeypatch.setattr(worker.requests, "post", lambda *a, **k: _Resp(400)) ok, _, _, retryable = worker._call_musetalk(video, audio, out) assert ok is False and retryable is False monkeypatch.setattr(worker.requests, "post", lambda *a, **k: _Resp(503)) ok, _, _, retryable = worker._call_musetalk(video, audio, out) assert ok is False and retryable is True # 连接异常:瞬时错误,可重试 import requests as _requests def _conn_err(*a, **k): raise _requests.exceptions.ConnectionError("reset") monkeypatch.setattr(worker.requests, "post", _conn_err) ok, _, _, retryable = worker._call_musetalk(video, audio, out) assert ok is False and retryable is True def test_handle_task_retries_once_for_transient_then_succeeds(worker, monkeypatch): calls = [] def _fake_call(v, a, o): calls.append(1) if len(calls) == 1: return False, 0.0, "MuseTalk HTTP 503: busy", True return True, 6.5, "", False monkeypatch.setattr(worker, "_register", lambda *a, **k: (True, False)) monkeypatch.setattr(worker, "_download", lambda url, path: True) monkeypatch.setattr(worker, "_probe_duration", lambda path: 12.0) monkeypatch.setattr(worker, "_call_musetalk", _fake_call) monkeypatch.setattr(worker, "time", mock.MagicMock()) # 重试 sleep 立即返回 uploaded = [] monkeypatch.setattr( worker, "_report_success_with_file", lambda task_id, duration, path: uploaded.append((task_id, duration)), ) worker._handle_task({"task_id": "t-retry", "video_url": "u", "audio_url": "u"}) assert len(calls) == 2 assert uploaded == [("t-retry", 6.5)] def test_handle_task_no_retry_for_deterministic_failure(worker, monkeypatch): calls = [] def _fake_call(v, a, o): calls.append(1) return False, 0.0, "MuseTalk HTTP 400: bad input", False reports = [] monkeypatch.setattr(worker, "_register", lambda *a, **k: (True, False)) monkeypatch.setattr(worker, "_download", lambda url, path: True) monkeypatch.setattr(worker, "_probe_duration", lambda path: 12.0) monkeypatch.setattr(worker, "_call_musetalk", _fake_call) monkeypatch.setattr( worker, "_report_result", lambda task_id, success, duration=0.0, error_msg="": reports.append(error_msg) or True, ) worker._handle_task({"task_id": "t-4xx", "video_url": "u", "audio_url": "u"}) assert len(calls) == 1 # 4xx 本地不重试,直接交服务端决定 assert reports and "400" in reports[0]