На Windows процессы поднимаются через spawn и заново импортируют app.main. Хранилище задач создавалось на уровне модуля, поэтому каждый новый процесс при старте помечал чужие выполняющиеся задачи как сорванные - все три записи падали с "сервис был перезапущен во время обработки". Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
284 lines
11 KiB
Python
284 lines
11 KiB
Python
"""Тесты параллельной обработки.
|
|
|
|
Замеры показали: обе стадии упираются в 4 потока, а на 16 работают вчетверо
|
|
медленнее. Значит ядра нужно занимать не шириной одной задачи, а несколькими
|
|
задачами сразу - и тогда очередь обязана быть устойчивой к гонкам.
|
|
"""
|
|
import threading
|
|
|
|
import pytest
|
|
|
|
from app.store import JobStatus, JobStore
|
|
|
|
|
|
@pytest.fixture
|
|
def store(tmp_path):
|
|
return JobStore(tmp_path / "jobs.db")
|
|
|
|
|
|
class TestClaimIsAtomic:
|
|
def test_claim_marks_running(self, store):
|
|
job_id = store.create(filename="a.wav", duration_sec=1.0)
|
|
assert store.claim_next() == job_id
|
|
assert store.get(job_id)["status"] == JobStatus.RUNNING
|
|
|
|
def test_second_claim_gets_nothing(self, store):
|
|
store.create(filename="a.wav", duration_sec=1.0)
|
|
store.claim_next()
|
|
assert store.claim_next() is None
|
|
|
|
def test_each_job_claimed_once_under_load(self, store):
|
|
"""Главное требование: два воркера не должны взять одну задачу."""
|
|
ids = {store.create(filename=f"{i}.wav", duration_sec=1.0) for i in range(50)}
|
|
claimed: list[str] = []
|
|
lock = threading.Lock()
|
|
|
|
def worker():
|
|
while True:
|
|
job_id = store.claim_next()
|
|
if job_id is None:
|
|
return
|
|
with lock:
|
|
claimed.append(job_id)
|
|
|
|
threads = [threading.Thread(target=worker) for _ in range(8)]
|
|
for t in threads:
|
|
t.start()
|
|
for t in threads:
|
|
t.join()
|
|
|
|
assert len(claimed) == len(set(claimed)) == 50
|
|
assert set(claimed) == ids
|
|
|
|
def test_claims_oldest_first(self, store):
|
|
first = store.create(filename="1.wav", duration_sec=1.0)
|
|
store.create(filename="2.wav", duration_sec=1.0)
|
|
assert store.claim_next() == first
|
|
|
|
|
|
class TestWorkerSettings:
|
|
def test_default_threads_is_capped(self, tmp_path, monkeypatch):
|
|
"""Дефолт умеренный: оптимум зависит от процессора и подбирается замером."""
|
|
from app.config import load_settings
|
|
|
|
monkeypatch.setattr("os.cpu_count", lambda: 32)
|
|
config = tmp_path / "config.toml"
|
|
config.write_text('[processing]\nthreads=0\n', encoding="utf-8")
|
|
assert load_settings(config).effective_threads() == 8
|
|
|
|
def test_explicit_threads_respected(self, tmp_path):
|
|
from app.config import load_settings
|
|
|
|
config = tmp_path / "config.toml"
|
|
config.write_text('[processing]\nthreads=6\n', encoding="utf-8")
|
|
assert load_settings(config).effective_threads() == 6
|
|
|
|
def test_workers_derived_from_cores(self, tmp_path, monkeypatch):
|
|
from app.config import load_settings
|
|
|
|
monkeypatch.setattr("os.cpu_count", lambda: 32)
|
|
config = tmp_path / "config.toml"
|
|
config.write_text('[processing]\nthreads=4\nworkers=0\n', encoding="utf-8")
|
|
assert 1 <= load_settings(config).effective_workers() <= 4
|
|
|
|
def test_workers_never_below_one(self, tmp_path, monkeypatch):
|
|
from app.config import load_settings
|
|
|
|
monkeypatch.setattr("os.cpu_count", lambda: 1)
|
|
config = tmp_path / "config.toml"
|
|
config.write_text('[processing]\nthreads=4\nworkers=0\n', encoding="utf-8")
|
|
assert load_settings(config).effective_workers() >= 1
|
|
|
|
def test_explicit_workers_respected(self, tmp_path):
|
|
from app.config import load_settings
|
|
|
|
config = tmp_path / "config.toml"
|
|
config.write_text('[processing]\nworkers=3\n', encoding="utf-8")
|
|
assert load_settings(config).effective_workers() == 3
|
|
|
|
|
|
class TestLazyWarmup:
|
|
"""Воркеры создаются заранее, а модели грузят при первой своей задаче."""
|
|
|
|
def test_transcribe_warms_up_when_needed(self, tmp_path, monkeypatch):
|
|
import sys
|
|
import types
|
|
|
|
for name in ("sherpa_onnx", "onnx_asr", "onnxruntime"):
|
|
monkeypatch.setitem(sys.modules, name, types.ModuleType(name))
|
|
from app.pipeline import Pipeline
|
|
|
|
p = Pipeline(models_dir=tmp_path, threads=1,
|
|
replacements_path=tmp_path / "r.txt", base_dir=tmp_path)
|
|
called = {"warmup": 0}
|
|
monkeypatch.setattr(p, "warmup", lambda: called.__setitem__("warmup", 1))
|
|
# transcribe упадёт дальше на чтении файла, но warmup обязан быть вызван
|
|
with pytest.raises(Exception):
|
|
p.transcribe(tmp_path / "нет.wav")
|
|
assert called["warmup"] == 1
|
|
|
|
def test_no_runtime_error_about_warmup(self, tmp_path, monkeypatch):
|
|
"""Прежде здесь падало «модели не загружены, вызовите warmup()»."""
|
|
import sys
|
|
import types
|
|
|
|
for name in ("sherpa_onnx", "onnx_asr", "onnxruntime"):
|
|
monkeypatch.setitem(sys.modules, name, types.ModuleType(name))
|
|
from app.pipeline import Pipeline
|
|
|
|
p = Pipeline(models_dir=tmp_path, threads=1,
|
|
replacements_path=tmp_path / "r.txt", base_dir=tmp_path)
|
|
monkeypatch.setattr(p, "warmup", lambda: None)
|
|
try:
|
|
p.transcribe(tmp_path / "нет.wav")
|
|
except RuntimeError as exc:
|
|
assert "warmup" not in str(exc)
|
|
except Exception:
|
|
pass
|
|
|
|
|
|
class TestBenchmarkEndpoint:
|
|
def test_benchmark_requires_token(self, tmp_path, monkeypatch):
|
|
import sys
|
|
import types
|
|
|
|
from fastapi.testclient import TestClient
|
|
|
|
config = tmp_path / "config.toml"
|
|
config.write_text('[security]\ntoken="t"\nallow_ips=""\n', encoding="utf-8")
|
|
monkeypatch.setenv("TALKSCORE_ASR_HOME", str(tmp_path))
|
|
for name in ("sherpa_onnx", "onnx_asr", "onnxruntime"):
|
|
monkeypatch.setitem(sys.modules, name, types.ModuleType(name))
|
|
for mod in [m for m in list(sys.modules) if m.startswith("app.")]:
|
|
del sys.modules[mod]
|
|
import app.config as cfg
|
|
monkeypatch.setattr(cfg, "BASE_DIR", tmp_path)
|
|
import app.main as main
|
|
main._worker_stop.set()
|
|
|
|
with TestClient(main.app) as c:
|
|
assert c.post("/v1/benchmark", files={"file": ("a.wav", b"x")}).status_code == 401
|
|
|
|
|
|
class TestSeparationQuality:
|
|
"""Метрика нужна, чтобы было видно, когда разметке по голосам верить нельзя."""
|
|
|
|
def test_distinct_voices_score_high(self):
|
|
import numpy as np
|
|
|
|
from app.acoustics import separation_quality
|
|
|
|
emb = np.vstack([np.random.RandomState(0).randn(8, 64),
|
|
np.random.RandomState(1).randn(8, 64) + 6])
|
|
labels = np.array([0] * 8 + [1] * 8)
|
|
assert separation_quality(emb, labels) > 0.4
|
|
|
|
def test_indistinguishable_voices_score_low(self):
|
|
import numpy as np
|
|
|
|
from app.acoustics import separation_quality
|
|
|
|
emb = np.random.RandomState(2).randn(16, 64)
|
|
labels = np.array([0] * 8 + [1] * 8)
|
|
assert separation_quality(emb, labels) < 0.3
|
|
|
|
def test_too_few_segments_returns_zero(self):
|
|
import numpy as np
|
|
|
|
from app.acoustics import separation_quality
|
|
|
|
assert separation_quality(np.random.randn(2, 64), np.array([0, 1])) == 0.0
|
|
|
|
def test_acoustics_reflect_high_frequency_content(self):
|
|
import numpy as np
|
|
|
|
from app.acoustics import segment_acoustics
|
|
|
|
sr = 16000
|
|
t = np.arange(sr) / sr
|
|
low = np.sin(2 * np.pi * 300 * t).astype(np.float32)
|
|
high = np.sin(2 * np.pi * 5000 * t).astype(np.float32)
|
|
assert segment_acoustics(high)["hf_ratio"] > segment_acoustics(low)["hf_ratio"]
|
|
|
|
def test_acoustics_reflect_loudness(self):
|
|
import numpy as np
|
|
|
|
from app.acoustics import segment_acoustics
|
|
|
|
loud = (np.random.RandomState(0).randn(16000) * 0.3).astype(np.float32)
|
|
quiet = loud * 0.1
|
|
assert segment_acoustics(loud)["loudness_db"] > segment_acoustics(quiet)["loudness_db"]
|
|
|
|
|
|
class TestProcessPool:
|
|
"""Расчёты уходят в процессы: библиотеки держат GIL, в потоках выигрыша нет."""
|
|
|
|
def test_worker_module_does_not_import_main(self):
|
|
"""На Windows дочерний процесс заново импортирует модуль с функцией.
|
|
|
|
Если бы он тянул app.main, в каждом процессе поднимался бы ещё один
|
|
веб-сервер со своей очередью. Проверяем сами импорты, а не текст файла.
|
|
"""
|
|
import ast
|
|
from pathlib import Path
|
|
|
|
src = (Path(__file__).resolve().parent.parent / "app" / "worker.py").read_text(encoding="utf-8")
|
|
imported = set()
|
|
for node in ast.walk(ast.parse(src)):
|
|
if isinstance(node, ast.Import):
|
|
imported.update(a.name for a in node.names)
|
|
elif isinstance(node, ast.ImportFrom) and node.module:
|
|
imported.add(node.module)
|
|
assert "app.main" not in imported
|
|
|
|
def test_pool_failure_falls_back_to_single_process(self, monkeypatch):
|
|
"""Сбой процессов не должен ронять сервис."""
|
|
import sys
|
|
import types
|
|
|
|
for name in ("sherpa_onnx", "onnx_asr", "onnxruntime"):
|
|
monkeypatch.setitem(sys.modules, name, types.ModuleType(name))
|
|
for mod in [m for m in list(sys.modules) if m.startswith("app.")]:
|
|
del sys.modules[mod]
|
|
|
|
import app.main as main
|
|
|
|
monkeypatch.setattr("concurrent.futures.ProcessPoolExecutor",
|
|
lambda **kw: (_ for _ in ()).throw(OSError("нет процессов")))
|
|
assert main._make_pool(4) is None
|
|
|
|
def test_run_job_needs_initialised_worker(self):
|
|
import app.worker as w
|
|
|
|
w._pipeline = None
|
|
try:
|
|
w.run_job("нет.wav", 2, "ffmpeg")
|
|
except RuntimeError as exc:
|
|
assert "не инициализирован" in str(exc)
|
|
else:
|
|
raise AssertionError("ожидалась ошибка о неинициализированном воркере")
|
|
|
|
|
|
class TestSpawnSafety:
|
|
"""На Windows дочерний процесс заново импортирует app.main."""
|
|
|
|
def test_importing_main_does_not_open_store(self, tmp_path, monkeypatch):
|
|
"""Иначе каждый новый процесс объявлял бы чужие задачи сорванными."""
|
|
import sys
|
|
import types
|
|
|
|
config = tmp_path / "config.toml"
|
|
config.write_text('[security]\ntoken="t"\n', encoding="utf-8")
|
|
monkeypatch.setenv("TALKSCORE_ASR_HOME", str(tmp_path))
|
|
for name in ("sherpa_onnx", "onnx_asr", "onnxruntime"):
|
|
monkeypatch.setitem(sys.modules, name, types.ModuleType(name))
|
|
for mod in [m for m in list(sys.modules) if m.startswith("app.")]:
|
|
del sys.modules[mod]
|
|
|
|
import app.config as cfg
|
|
monkeypatch.setattr(cfg, "BASE_DIR", tmp_path)
|
|
import app.main as main
|
|
|
|
assert main.store is None
|
|
assert not (tmp_path / "data" / "jobs.db").exists()
|