"""#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 text = "" 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, {})) assert worker._register("task-abc") is True 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 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_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) 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) 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) 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) 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]