mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -07:00
Merge remote-tracking branch 'origin/main' into prep/pr-1112
# Conflicts: # CHANGELOG.md
This commit is contained in:
@@ -0,0 +1,80 @@
|
||||
"""
|
||||
Regression tests for the VOICEBOX_FORCE_CPU environment override.
|
||||
|
||||
The docs promise (docs/content/docs/developer/tts-generation.mdx) that
|
||||
get_torch_device() layers "VOICEBOX_FORCE_CPU environment override" ahead of
|
||||
CUDA/XPU/MPS detection, and gpu-acceleration.mdx tells users to set it to fall
|
||||
back to CPU when the bundled PyTorch has no kernels for their GPU.
|
||||
|
||||
torch is stubbed through sys.modules so these run without a torch install and
|
||||
without any GPU.
|
||||
|
||||
Usage:
|
||||
python -m pytest backend/tests/test_force_cpu_env.py -v
|
||||
"""
|
||||
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
from backend.backends.base import get_torch_device
|
||||
|
||||
# The documented public name and value of the override. Pinned here independently
|
||||
# of the production constants so a rename of either fails these tests.
|
||||
FORCE_CPU_ENV_VAR = "VOICEBOX_FORCE_CPU"
|
||||
|
||||
# Sentinel for "the variable is not set at all".
|
||||
UNSET = None
|
||||
|
||||
|
||||
class _FakeCuda:
|
||||
@staticmethod
|
||||
def is_available() -> bool:
|
||||
return True
|
||||
|
||||
|
||||
class _FakeTorch:
|
||||
"""Minimal stand-in for a CUDA-enabled torch install."""
|
||||
|
||||
cuda = _FakeCuda
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def cuda_available(monkeypatch):
|
||||
"""Make torch report a usable CUDA device without installing torch."""
|
||||
monkeypatch.setitem(sys.modules, "torch", _FakeTorch)
|
||||
|
||||
|
||||
def _set_override(monkeypatch, value):
|
||||
if value is UNSET:
|
||||
monkeypatch.delenv(FORCE_CPU_ENV_VAR, raising=False)
|
||||
else:
|
||||
monkeypatch.setenv(FORCE_CPU_ENV_VAR, value)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", ["1", " 1 "])
|
||||
def test_force_cpu_wins_over_available_cuda(monkeypatch, cuda_available, value):
|
||||
"""The documented value must beat an otherwise usable CUDA device.
|
||||
|
||||
Surrounding whitespace is tolerated: on Windows, where this override
|
||||
matters most, it is typically set through the GUI environment editor."""
|
||||
_set_override(monkeypatch, value)
|
||||
|
||||
assert get_torch_device(allow_xpu=True, allow_directml=True, allow_mps=True) == "cpu"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [UNSET, "", "0"])
|
||||
def test_without_override_cuda_is_still_selected(monkeypatch, cuda_available, value):
|
||||
"""Unset or disabled must not disturb normal device detection."""
|
||||
_set_override(monkeypatch, value)
|
||||
|
||||
assert get_torch_device() == "cuda"
|
||||
|
||||
|
||||
def test_force_cpu_does_not_need_torch(monkeypatch):
|
||||
"""The override is honoured before torch is imported, so it works on a
|
||||
broken/incompatible torch install — which is the case it exists for."""
|
||||
_set_override(monkeypatch, "1")
|
||||
monkeypatch.setitem(sys.modules, "torch", None) # makes `import torch` raise
|
||||
|
||||
assert get_torch_device() == "cpu"
|
||||
@@ -0,0 +1,27 @@
|
||||
"""Tests for generation request engine selection."""
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from backend import models
|
||||
|
||||
|
||||
def _request(**kwargs) -> models.GenerationRequest:
|
||||
return models.GenerationRequest(profile_id="profile-1", text="hello", **kwargs)
|
||||
|
||||
|
||||
def test_omitted_engine_does_not_override_profile_default():
|
||||
request = _request()
|
||||
|
||||
assert request.engine is None
|
||||
|
||||
|
||||
def test_explicit_engine_is_preserved():
|
||||
request = _request(engine="chatterbox")
|
||||
|
||||
assert request.engine == "chatterbox"
|
||||
|
||||
|
||||
def test_invalid_explicit_engine_is_rejected():
|
||||
with pytest.raises(ValidationError):
|
||||
_request(engine="invalid")
|
||||
@@ -0,0 +1,74 @@
|
||||
"""
|
||||
Unit tests for ``is_model_cached``'s handling of stale ``.incomplete`` blobs.
|
||||
|
||||
A retried or concurrent download can leave an orphaned ``.incomplete`` file
|
||||
next to its now-completed counterpart (same blob hash, no suffix). Only a
|
||||
genuinely in-progress download -- an ``.incomplete`` with no completed blob
|
||||
alongside it -- should mark the model as not cached.
|
||||
|
||||
``is_model_cached`` is extracted and exec'd standalone (instead of importing
|
||||
``backend.backends.base``) so this test doesn't pull in the module's sibling
|
||||
imports (audio/progress/hf_progress/tasks), which in turn require the full
|
||||
ML stack (torch/transformers/librosa/fastapi/...) this pure filesystem check
|
||||
never touches.
|
||||
"""
|
||||
|
||||
import ast
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
_SOURCE = (Path(__file__).parent.parent / "backends" / "base.py").read_text()
|
||||
_MODULE = ast.parse(_SOURCE)
|
||||
_FUNC_SRC = "\n\n".join(
|
||||
ast.get_source_segment(_SOURCE, node)
|
||||
for node in _MODULE.body
|
||||
if isinstance(node, ast.FunctionDef) and node.name in ("has_in_progress_download", "is_model_cached")
|
||||
)
|
||||
|
||||
_namespace = {
|
||||
"Path": Path,
|
||||
"Optional": Optional,
|
||||
"logger": logging.getLogger("test_is_model_cached"),
|
||||
}
|
||||
exec(_FUNC_SRC, _namespace) # noqa: S102
|
||||
is_model_cached = _namespace["is_model_cached"]
|
||||
|
||||
|
||||
def _make_repo_cache(tmp_path, repo="org/model"):
|
||||
repo_dir = tmp_path / ("models--" + repo.replace("/", "--"))
|
||||
blobs_dir = repo_dir / "blobs"
|
||||
snapshots_dir = repo_dir / "snapshots" / "abc123"
|
||||
blobs_dir.mkdir(parents=True)
|
||||
snapshots_dir.mkdir(parents=True)
|
||||
return repo_dir, blobs_dir, snapshots_dir
|
||||
|
||||
|
||||
def test_orphaned_incomplete_blob_does_not_block_cache_hit(tmp_path, monkeypatch):
|
||||
import huggingface_hub.constants as hf_constants
|
||||
|
||||
monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(tmp_path))
|
||||
|
||||
repo = "org/model"
|
||||
_, blobs_dir, snapshots_dir = _make_repo_cache(tmp_path, repo)
|
||||
|
||||
completed_blob = blobs_dir / "deadbeef"
|
||||
completed_blob.write_bytes(b"weights")
|
||||
(blobs_dir / "deadbeef.incomplete").write_bytes(b"stale partial")
|
||||
(snapshots_dir / "model.safetensors").symlink_to(completed_blob)
|
||||
|
||||
assert is_model_cached(repo) is True
|
||||
|
||||
|
||||
def test_genuinely_in_progress_download_is_not_cached(tmp_path, monkeypatch):
|
||||
import huggingface_hub.constants as hf_constants
|
||||
|
||||
monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(tmp_path))
|
||||
|
||||
repo = "org/model"
|
||||
_, blobs_dir, snapshots_dir = _make_repo_cache(tmp_path, repo)
|
||||
|
||||
(blobs_dir / "feedface.incomplete").write_bytes(b"partial")
|
||||
(snapshots_dir / "config.json").write_text("{}")
|
||||
|
||||
assert is_model_cached(repo) is False
|
||||
@@ -1,4 +1,4 @@
|
||||
"""Tests for the voicebox.speak MCP tool's ``model_size`` plumbing (issue #884).
|
||||
"""Tests for the voicebox_speak MCP tool's ``model_size`` plumbing (issue #884).
|
||||
|
||||
The MCP speak path used to build its ``GenerationRequest`` without a
|
||||
``model_size``, so every agent-triggered generation silently fell back to the
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
"""Tests for voice-profile language fallback on the two speak surfaces.
|
||||
|
||||
Both speak paths built their ``GenerationRequest`` with a hardcoded ``"en"``
|
||||
fallback and never consulted the resolved profile, so a profile created with
|
||||
``language="fr"`` was still synthesised as English unless every caller passed
|
||||
``language=`` explicitly. Agents going through MCP had no way to know the
|
||||
profile's language, so they couldn't pass it either.
|
||||
|
||||
These tests pin the fix: the fallback chain is now explicit argument →
|
||||
resolved profile's language → ``"en"``, matching how ``engine`` and
|
||||
``personality`` already consult the resolved binding.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
import backend.routes.generations as generations
|
||||
import backend.routes.speak as speak_route
|
||||
from backend import models
|
||||
from backend.mcp_server import tools
|
||||
|
||||
|
||||
class _FakeGeneration:
|
||||
"""Minimal stand-in for GenerationResponse consumed by the speak paths."""
|
||||
|
||||
id = "gen-test"
|
||||
status = "generating"
|
||||
|
||||
def model_dump(self, mode="json"):
|
||||
return {"id": self.id, "status": self.status}
|
||||
|
||||
|
||||
class _FakeProfile:
|
||||
def __init__(self, language):
|
||||
self.id = "p1"
|
||||
self.name = "Siwis"
|
||||
self.language = language
|
||||
self.personality = None
|
||||
|
||||
|
||||
class _FakeQuery:
|
||||
def filter(self, *args, **kwargs):
|
||||
return self
|
||||
|
||||
def first(self):
|
||||
# No per-client binding — engine/personality fall through to their
|
||||
# own defaults, leaving language as the only variable under test.
|
||||
return None
|
||||
|
||||
|
||||
class _FakeDB:
|
||||
def query(self, *args, **kwargs):
|
||||
return _FakeQuery()
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
|
||||
class _FakeRequest:
|
||||
"""Stands in for starlette's Request — only headers are read."""
|
||||
|
||||
def __init__(self, client_id=None):
|
||||
self.headers = {"X-Voicebox-Client-Id": client_id} if client_id else {}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def captured_request(monkeypatch):
|
||||
"""Capture the GenerationRequest instead of running a real generation.
|
||||
|
||||
Both speak paths import ``generate_speech`` lazily from
|
||||
``routes.generations``, so patching the attribute on that module
|
||||
intercepts the call on either surface.
|
||||
"""
|
||||
captured = {}
|
||||
|
||||
async def fake_generate_speech(req, db):
|
||||
captured["req"] = req
|
||||
return _FakeGeneration()
|
||||
|
||||
monkeypatch.setattr(generations, "generate_speech", fake_generate_speech)
|
||||
monkeypatch.setattr(speak_route.mcp_events, "publish", lambda *a, **k: None)
|
||||
monkeypatch.setattr(tools.mcp_events, "publish", lambda *a, **k: None)
|
||||
return captured
|
||||
|
||||
|
||||
# ─── REST: POST /speak ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
async def _call_rest(monkeypatch, profile_language, requested_language=None):
|
||||
monkeypatch.setattr(
|
||||
speak_route,
|
||||
"resolve_profile",
|
||||
lambda profile, client_id, db: _FakeProfile(profile_language),
|
||||
)
|
||||
await speak_route.speak(
|
||||
models.SpeakRequest(text="Bonjour", language=requested_language),
|
||||
_FakeRequest(client_id="claude-code"),
|
||||
_FakeDB(),
|
||||
)
|
||||
|
||||
|
||||
async def test_rest_speak_falls_back_to_profile_language(captured_request, monkeypatch):
|
||||
await _call_rest(monkeypatch, profile_language="fr")
|
||||
assert captured_request["req"].language == "fr"
|
||||
|
||||
|
||||
async def test_rest_speak_explicit_language_wins(captured_request, monkeypatch):
|
||||
# An explicit argument still overrides the profile — a French profile can
|
||||
# be asked to read an English string.
|
||||
await _call_rest(monkeypatch, profile_language="fr", requested_language="en")
|
||||
assert captured_request["req"].language == "en"
|
||||
|
||||
|
||||
async def test_rest_speak_defaults_to_en_without_profile_language(captured_request, monkeypatch):
|
||||
# Profiles predating the language column resolve to None; the "en"
|
||||
# backstop keeps their behaviour unchanged.
|
||||
await _call_rest(monkeypatch, profile_language=None)
|
||||
assert captured_request["req"].language == "en"
|
||||
|
||||
|
||||
# ─── MCP: voicebox.speak ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
async def _call_mcp(monkeypatch, profile_language, requested_language=None):
|
||||
# Build the server from the same ``fastmcp`` package production imports so
|
||||
# the registered ``voicebox.speak`` wrapper — where the profile fallback
|
||||
# lives — is the code under test.
|
||||
from fastmcp import FastMCP
|
||||
|
||||
monkeypatch.setattr(
|
||||
tools,
|
||||
"resolve_profile",
|
||||
lambda profile, client_id, db: _FakeProfile(profile_language),
|
||||
)
|
||||
monkeypatch.setattr(tools, "get_db", lambda: iter([_FakeDB()]))
|
||||
|
||||
mcp = FastMCP("test")
|
||||
tools.register_tools(mcp)
|
||||
args = {"text": "Bonjour"}
|
||||
if requested_language is not None:
|
||||
args["language"] = requested_language
|
||||
await mcp.call_tool("voicebox.speak", args)
|
||||
|
||||
|
||||
async def test_mcp_speak_falls_back_to_profile_language(captured_request, monkeypatch):
|
||||
# The agent-facing path matters most: an MCP client can't know the
|
||||
# profile's language, so omitting it must not silently mean English.
|
||||
await _call_mcp(monkeypatch, profile_language="fr")
|
||||
assert captured_request["req"].language == "fr"
|
||||
|
||||
|
||||
async def test_mcp_speak_explicit_language_wins(captured_request, monkeypatch):
|
||||
await _call_mcp(monkeypatch, profile_language="fr", requested_language="en")
|
||||
assert captured_request["req"].language == "en"
|
||||
|
||||
|
||||
async def test_mcp_speak_defaults_to_en_without_profile_language(captured_request, monkeypatch):
|
||||
await _call_mcp(monkeypatch, profile_language=None)
|
||||
assert captured_request["req"].language == "en"
|
||||
Reference in New Issue
Block a user