Merge origin/main into prep/pr-1031

This commit is contained in:
jamiepine
2026-10-04 00:01:09 +00:00
49 changed files with 811 additions and 123 deletions
+38 -4
View File
@@ -6,6 +6,7 @@ voice prompt combination, and model loading progress tracking.
"""
import logging
import os
import platform
from contextlib import contextmanager
from pathlib import Path
@@ -21,6 +22,24 @@ from ..utils.tasks import get_task_manager
logger = logging.getLogger(__name__)
def has_in_progress_download(blobs_dir: Path) -> bool:
"""
Whether a HuggingFace repo's ``blobs`` dir holds a genuinely in-progress download.
An ``.incomplete`` blob means a download is still in progress -- unless a
completed blob with the same hash already sits next to it, which happens
when a retried/concurrent download leaves a stale ``.incomplete`` behind
after the real transfer already finished. Only orphaned ``.incomplete``
files (no matching completed blob) count as "in progress".
"""
if not blobs_dir.exists():
return False
return any(
not incomplete.with_name(incomplete.name.removesuffix(".incomplete")).exists()
for incomplete in blobs_dir.glob("*.incomplete")
)
def is_model_cached(
hf_repo: str,
*,
@@ -47,10 +66,8 @@ def is_model_cached(
if not repo_cache.exists():
return False
# Incomplete blobs mean a download is still in progress
blobs_dir = repo_cache / "blobs"
if blobs_dir.exists() and any(blobs_dir.glob("*.incomplete")):
logger.debug(f"Found .incomplete files for {hf_repo}")
if has_in_progress_download(repo_cache / "blobs"):
logger.debug(f"Found in-progress .incomplete file for {hf_repo}")
return False
snapshots_dir = repo_cache / "snapshots"
@@ -77,6 +94,13 @@ def is_model_cached(
return False
# Documented escape hatch (docs/content/docs/overview/gpu-acceleration.mdx):
# users whose GPU has no compiled kernels in the bundled PyTorch set this to run
# on CPU instead of crashing at generation time.
FORCE_CPU_ENV_VAR = "VOICEBOX_FORCE_CPU"
FORCE_CPU_ENABLED_VALUE = "1"
def get_torch_device(
*,
allow_xpu: bool = False,
@@ -92,7 +116,17 @@ def get_torch_device(
allow_directml: Check for DirectML (Windows) support.
allow_mps: Allow MPS (Apple Silicon). If False, MPS falls back to CPU.
force_cpu_on_mac: Force CPU on macOS regardless of GPU availability.
The VOICEBOX_FORCE_CPU override wins over every other candidate, and is
resolved before torch is imported so it still works when the installed
build is the reason CPU is wanted.
"""
# Stripped: on Windows, where this override matters most, it is usually set
# through the GUI environment editor.
if os.environ.get(FORCE_CPU_ENV_VAR, "").strip() == FORCE_CPU_ENABLED_VALUE:
logger.info("%s=%s set, forcing CPU device", FORCE_CPU_ENV_VAR, FORCE_CPU_ENABLED_VALUE)
return "cpu"
if force_cpu_on_mac and platform.system() == "Darwin":
return "cpu"
+51 -1
View File
@@ -44,6 +44,7 @@ def run_migrations(engine) -> None:
_migrate_capture_settings(engine, inspector, tables)
_migrate_mcp_bindings(engine, inspector, tables)
_normalize_storage_paths(engine, tables)
_migrate_add_indexes(engine, tables)
# -- helpers ---------------------------------------------------------------
@@ -249,7 +250,7 @@ def _migrate_mcp_bindings(engine, inspector, tables: set[str]) -> None:
"""Drop the legacy ``default_intent`` column and add ``default_personality``.
The intent tri-state (respond / rewrite / compose) has been collapsed
to a boolean: when true, ``voicebox.speak`` rewrites input through the
to a boolean: when true, ``voicebox_speak`` rewrites input through the
profile's personality LLM before TTS.
"""
if "mcp_client_bindings" not in tables:
@@ -292,6 +293,55 @@ def _supports_drop_column(engine) -> bool:
return tuple(int(p) for p in sqlite3.sqlite_version.split(".")[:3]) >= (3, 35, 0)
def _migrate_add_indexes(engine, tables: set[str]) -> None:
"""Create missing indexes on high-traffic foreign keys and sort columns.
SQLite silently ignores ``CREATE INDEX IF NOT EXISTS``, so this is
safe to run on every startup regardless of whether the index already
exists. New installs get the indexes from ``Base.metadata.create_all``
(via the ``index=True`` column flags); this migration brings existing
databases into parity without dropping or recreating any data.
"""
indexes = [
# generations — filtered by profile, ordered/filtered by date, filtered by status
("ix_generations_profile_id", "generations", "profile_id"),
("ix_generations_created_at", "generations", "created_at"),
("ix_generations_status", "generations", "status"),
# story_items — every story lookup filters by story_id; join on generation_id
("ix_story_items_story_id", "story_items", "story_id"),
("ix_story_items_generation_id", "story_items", "generation_id"),
# generation_versions — always filtered/joined on generation_id
("ix_generation_versions_generation_id", "generation_versions", "generation_id"),
# profile_samples — loaded per-profile on every voice prompt build
("ix_profile_samples_profile_id", "profile_samples", "profile_id"),
# captures — ordered by date in list view
("ix_captures_created_at", "captures", "created_at"),
# channel_device_mappings — looked up per channel
("ix_channel_device_mappings_channel_id", "channel_device_mappings", "channel_id"),
]
with engine.connect() as conn:
existing = {
row[0]
for row in conn.execute(text("SELECT name FROM sqlite_master WHERE type = 'index'"))
}
created = []
for index_name, table, column in indexes:
if table not in tables or index_name in existing:
continue
conn.execute(
text(
f"CREATE INDEX IF NOT EXISTS {index_name}"
f" ON {table} ({column})"
)
)
created.append(index_name)
conn.commit()
if created:
logger.info("Created %d missing index(es): %s", len(created), ", ".join(created))
def _normalize_storage_paths(engine, tables: set[str]) -> None:
"""Normalize stored file paths to be relative to the configured data dir."""
from pathlib import Path
+10 -10
View File
@@ -54,7 +54,7 @@ class ProfileSample(Base):
__tablename__ = "profile_samples"
id = Column(String, primary_key=True, default=lambda: str(uuid.uuid4()))
profile_id = Column(String, ForeignKey("profiles.id"), nullable=False)
profile_id = Column(String, ForeignKey("profiles.id"), nullable=False, index=True)
audio_path = Column(String, nullable=False)
reference_text = Column(Text, nullable=False)
@@ -65,7 +65,7 @@ class Generation(Base):
__tablename__ = "generations"
id = Column(String, primary_key=True, default=lambda: str(uuid.uuid4()))
profile_id = Column(String, ForeignKey("profiles.id"), nullable=False)
profile_id = Column(String, ForeignKey("profiles.id"), nullable=False, index=True)
text = Column(Text, nullable=False)
language = Column(String, default="en")
audio_path = Column(String, nullable=True)
@@ -74,7 +74,7 @@ class Generation(Base):
instruct = Column(Text)
engine = Column(String, default="qwen")
model_size = Column(String, nullable=True)
status = Column(String, default="completed")
status = Column(String, default="completed", index=True)
error = Column(Text, nullable=True)
is_favorited = Column(Boolean, default=False)
# Origin of this generation — "manual" for plain /generate calls,
@@ -82,7 +82,7 @@ class Generation(Base):
# profile's personality LLM before TTS. Future sources (bulk import,
# agent replies, etc.) can extend this.
source = Column(String, nullable=False, default="manual")
created_at = Column(DateTime, default=datetime.utcnow)
created_at = Column(DateTime, default=datetime.utcnow, index=True)
class Story(Base):
@@ -103,8 +103,8 @@ class StoryItem(Base):
__tablename__ = "story_items"
id = Column(String, primary_key=True, default=lambda: str(uuid.uuid4()))
story_id = Column(String, ForeignKey("stories.id"), nullable=False)
generation_id = Column(String, ForeignKey("generations.id"), nullable=False)
story_id = Column(String, ForeignKey("stories.id"), nullable=False, index=True)
generation_id = Column(String, ForeignKey("generations.id"), nullable=False, index=True)
version_id = Column(String, ForeignKey("generation_versions.id"), nullable=True)
start_time_ms = Column(Integer, nullable=False, default=0)
track = Column(Integer, nullable=False, default=0)
@@ -132,7 +132,7 @@ class GenerationVersion(Base):
__tablename__ = "generation_versions"
id = Column(String, primary_key=True, default=lambda: str(uuid.uuid4()))
generation_id = Column(String, ForeignKey("generations.id"), nullable=False)
generation_id = Column(String, ForeignKey("generations.id"), nullable=False, index=True)
label = Column(String, nullable=False)
audio_path = Column(String, nullable=False)
effects_chain = Column(Text, nullable=True)
@@ -172,7 +172,7 @@ class ChannelDeviceMapping(Base):
__tablename__ = "channel_device_mappings"
id = Column(String, primary_key=True, default=lambda: str(uuid.uuid4()))
channel_id = Column(String, ForeignKey("audio_channels.id"), nullable=False)
channel_id = Column(String, ForeignKey("audio_channels.id"), nullable=False, index=True)
device_id = Column(String, nullable=False)
@@ -272,7 +272,7 @@ class MCPClientBinding(Base):
label = Column(String, nullable=True) # display name
profile_id = Column(String, ForeignKey("profiles.id"), nullable=True)
default_engine = Column(String, nullable=True)
# When true, voicebox.speak routes through the profile's personality LLM
# When true, voicebox_speak routes through the profile's personality LLM
# (rewrite) before TTS by default. Callers can still override per call.
default_personality = Column(Boolean, nullable=False, default=False)
last_seen_at = Column(DateTime, nullable=True)
@@ -300,4 +300,4 @@ class Capture(Base):
stt_model = Column(String, nullable=True)
llm_model = Column(String, nullable=True)
refinement_flags = Column(Text, nullable=True) # JSON blob
created_at = Column(DateTime, default=datetime.utcnow)
created_at = Column(DateTime, default=datetime.utcnow, index=True)
+16 -1
View File
@@ -3,7 +3,7 @@
import logging
import uuid
from sqlalchemy import create_engine
from sqlalchemy import create_engine, event
from sqlalchemy.orm import sessionmaker
from .. import config
@@ -21,6 +21,7 @@ from .seed import backfill_generation_versions, seed_builtin_presets
logger = logging.getLogger(__name__)
# Initialized by init_db()
engine = None
SessionLocal = None
@@ -39,6 +40,20 @@ def init_db() -> None:
connect_args={"check_same_thread": False},
)
@event.listens_for(engine, "connect")
def _set_sqlite_pragmas(dbapi_connection, _record) -> None:
# Each pooled connection enables WAL journal mode and sets a 5-second
# busy timeout. WAL allows concurrent readers during a write (the
# default DELETE/ROLLBACK journal blocks all readers), which matters
# for voicebox because SSE status polls and history queries run
# concurrently with the generation worker writing to the same db.
# busy_timeout prevents "database is locked" errors when two
# connections briefly contend on the same write slot.
cursor = dbapi_connection.cursor()
cursor.execute("PRAGMA journal_mode=WAL")
cursor.execute("PRAGMA busy_timeout=5000")
cursor.close()
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
run_migrations(engine)
+6 -6
View File
@@ -49,10 +49,10 @@ claude mcp add voicebox \
| Name | Purpose |
|---|---|
| `voicebox.speak` | Speak text in a voice profile. Returns a generation id you can poll. |
| `voicebox.transcribe` | Whisper transcription of a base64 blob or an absolute local path. |
| `voicebox.list_captures` | Recent captures (dictation / recording / file) with transcripts. |
| `voicebox.list_profiles` | Available voice profiles (cloned + preset). |
| `voicebox_speak` | Speak text in a voice profile. Returns a generation id you can poll. |
| `voicebox_transcribe` | Whisper transcription of a base64 blob or an absolute local path. |
| `voicebox_list_captures` | Recent captures (dictation / recording / file) with transcripts. |
| `voicebox_list_profiles` | Available voice profiles (cloned + preset). |
All tools resolve voice profiles in this precedence:
@@ -69,8 +69,8 @@ Settings → MCP.
npx @modelcontextprotocol/inspector http://127.0.0.1:17493/mcp
```
Point it at the URL, hit "List tools," call `voicebox.list_profiles`
first to confirm wiring, then `voicebox.speak` for end-to-end.
Point it at the URL, hit "List tools," call `voicebox_list_profiles`
first to confirm wiring, then `voicebox_speak` for end-to-end.
## Non-MCP REST surface
+3 -3
View File
@@ -33,7 +33,7 @@ current_client_id: ContextVar[str | None] = ContextVar(
)
# Remote address of the in-flight request. Used by tools that gate
# host-filesystem access to loopback callers (see voicebox.transcribe).
# host-filesystem access to loopback callers (see voicebox_transcribe).
current_remote_addr: ContextVar[str | None] = ContextVar(
"current_remote_addr", default=None
)
@@ -61,12 +61,12 @@ def request_is_loopback() -> bool:
# ignored so the Settings UI's "last heard from" column only reflects
# calls that actually acted on the client's bindings.
#
# - /mcp — FastMCP tool calls (voicebox.speak, voicebox.transcribe, …)
# - /mcp — FastMCP tool calls (voicebox_speak, voicebox_transcribe, …)
# and the /mcp/bindings admin surface. The admin surface is never
# called with the header in practice (the frontend manages bindings
# over plain REST), so the `startswith("/mcp")` match doesn't cause
# false stamps.
# - /speak — REST mirror of voicebox.speak for non-MCP agents (shell
# - /speak — REST mirror of voicebox_speak for non-MCP agents (shell
# scripts, ACP, A2A). Uses the same per-client binding lookup, so its
# callers belong in the last-seen list too.
_STAMPED_PATH_PREFIXES: tuple[str, ...] = ("/mcp", "/speak")
+1 -1
View File
@@ -1,6 +1,6 @@
"""In-memory pub/sub for speaking-pill SSE broadcasts.
MCP ``voicebox.speak`` calls and the REST ``POST /speak`` route publish
MCP ``voicebox_speak`` calls and the REST ``POST /speak`` route publish
start/end events that DictateWindow subscribes to via /events/speak, so the
floating pill surfaces whenever an agent is speaking.
"""
+2 -2
View File
@@ -27,8 +27,8 @@ def build_mcp_server() -> FastMCP:
mcp = FastMCP(
name="voicebox",
instructions=(
"Voicebox is a local voice I/O layer. Use `voicebox.speak` to "
"play text in a voice profile, `voicebox.transcribe` for "
"Voicebox is a local voice I/O layer. Use `voicebox_speak` to "
"play text in a voice profile, `voicebox_transcribe` for "
"audio→text, and the `list_*` tools to discover profiles and "
"captures."
),
+10 -9
View File
@@ -1,8 +1,9 @@
"""Voicebox MCP tool implementations.
Thin wrappers over existing services/routes. Tools are registered with dotted
names (``voicebox.speak`` etc.) so they look natural in agent logs —
the Python function name stays snake_case.
Thin wrappers over existing services/routes. Tools are registered with
underscore-separated names (``voicebox_speak`` etc.): MCP clients such as
Claude Desktop validate tool names against ``^[a-zA-Z0-9_-]{1,64}$`` and
reject the whole tool list if any name contains a dot (#790).
"""
from __future__ import annotations
@@ -36,7 +37,7 @@ def register_tools(mcp: FastMCP) -> None:
"""Attach all Voicebox tools to the given FastMCP instance."""
@mcp.tool(
name="voicebox.speak",
name="voicebox_speak",
description=(
"Speak text in a Voicebox voice profile. Returns a generation id "
"the caller can poll at /generate/{id}/status. Audio plays on the "
@@ -104,7 +105,7 @@ def register_tools(mcp: FastMCP) -> None:
profile_name=vp.name,
text=text,
engine=resolved_engine,
language=language,
language=language or vp.language,
personality=use_persona,
model_size=model_size,
db=db,
@@ -113,7 +114,7 @@ def register_tools(mcp: FastMCP) -> None:
db.close()
@mcp.tool(
name="voicebox.transcribe",
name="voicebox_transcribe",
description=(
"Transcribe an audio clip to text using Voicebox's local Whisper. "
"Pass exactly one of `audio_base64` (bytes as base64) or "
@@ -171,7 +172,7 @@ def register_tools(mcp: FastMCP) -> None:
tmp_path.unlink(missing_ok=True)
@mcp.tool(
name="voicebox.list_captures",
name="voicebox_list_captures",
description=(
"List recent voice captures (dictations, recordings, uploads) "
"with their transcripts. Most-recent first."
@@ -199,10 +200,10 @@ def register_tools(mcp: FastMCP) -> None:
db.close()
@mcp.tool(
name="voicebox.list_profiles",
name="voicebox_list_profiles",
description=(
"List available voice profiles (both cloned voices and presets). "
"Use the returned `name` with voicebox.speak(profile=...)."
"Use the returned `name` with voicebox_speak(profile=...)."
),
)
async def voicebox_list_profiles() -> dict[str, Any]:
+3 -3
View File
@@ -85,7 +85,7 @@ class GenerationRequest(BaseModel):
seed: Optional[int] = Field(None, ge=0)
model_size: Optional[str] = Field(default="1.7B", pattern="^(1\\.7B|0\\.6B|1B|3B)$")
instruct: Optional[str] = Field(None, max_length=500)
engine: Optional[str] = Field(default="qwen", pattern="^(qwen|qwen_custom_voice|luxtts|chatterbox|chatterbox_turbo|tada|kokoro)$")
engine: Optional[str] = Field(default=None, pattern="^(qwen|qwen_custom_voice|luxtts|chatterbox|chatterbox_turbo|tada|kokoro)$")
personality: bool = Field(
default=False,
description="When true and the profile has a personality prompt, the input text is rewritten in-character before TTS.",
@@ -309,7 +309,7 @@ class GenerationSettingsUpdate(BaseModel):
class MCPClientBindingResponse(BaseModel):
"""Per-MCP-client voice binding — what voice / engine the server should
use when a given client_id calls voicebox.speak without args, plus an
use when a given client_id calls voicebox_speak without args, plus an
opt-in personality-rewrite default."""
client_id: str
@@ -346,7 +346,7 @@ class MCPClientBindingListResponse(BaseModel):
class SpeakRequest(BaseModel):
"""Body for POST /speak — non-MCP REST surface that mirrors voicebox.speak."""
"""Body for POST /speak — non-MCP REST surface that mirrors voicebox_speak."""
text: str = Field(..., min_length=1, max_length=10000)
profile: Optional[str] = Field(
+1 -1
View File
@@ -64,7 +64,7 @@ pedalboard>=0.9.0
httpx>=0.27.0
# MCP server (Model Context Protocol) — lets local AI agents call
# voicebox.speak / .transcribe / .list_captures / .list_profiles
# voicebox_speak / voicebox_transcribe / voicebox_list_captures / voicebox_list_profiles
fastmcp>=3.0,<4.0
sse-starlette>=2.0
+3 -3
View File
@@ -244,6 +244,7 @@ async def get_model_status():
use_scan_cache = False
from ..backends import get_all_model_configs, check_model_loaded
from ..backends.base import has_in_progress_download
registry_configs = get_all_model_configs()
model_configs = [
@@ -293,8 +294,7 @@ async def get_model_status():
try:
cache_dir = hf_constants.HF_HUB_CACHE
blobs_dir = Path(cache_dir) / ("models--" + repo_id.replace("/", "--")) / "blobs"
if blobs_dir.exists():
has_incomplete = any(blobs_dir.glob("*.incomplete"))
has_incomplete = has_in_progress_download(blobs_dir)
except Exception:
pass
@@ -314,7 +314,7 @@ async def get_model_status():
if repo_cache.exists():
blobs_dir = repo_cache / "blobs"
has_incomplete = blobs_dir.exists() and any(blobs_dir.glob("*.incomplete"))
has_incomplete = has_in_progress_download(blobs_dir)
if not has_incomplete:
snapshots_dir = repo_cache / "snapshots"
+3 -3
View File
@@ -1,4 +1,4 @@
"""POST /speak — REST wrapper around voicebox.speak for non-MCP callers.
"""POST /speak — REST wrapper around voicebox_speak for non-MCP callers.
Shell scripts, ACP, A2A, or any agent that doesn't speak MCP can hit this
endpoint to play text through a cloned voice. Uses the same profile
@@ -30,7 +30,7 @@ async def speak(
request: Request,
db: Session = Depends(get_db),
):
"""Speak text in a voice profile. Mirrors voicebox.speak (MCP).
"""Speak text in a voice profile. Mirrors voicebox_speak (MCP).
Response shape matches POST /generate — a ``GenerationResponse`` with
``status="generating"`` and an ``id`` the caller polls at
@@ -75,7 +75,7 @@ async def speak(
models.GenerationRequest(
profile_id=profile.id,
text=data.text,
language=data.language or "en",
language=data.language or profile.language or "en",
engine=engine,
personality=bool(personality_flag),
),
+37
View File
@@ -27,6 +27,43 @@ if not _is_writable(sys.stdout):
if not _is_writable(sys.stderr):
sys.stderr = open(os.devnull, 'w')
class _PipeSafeStream:
"""Falls back to devnull once the pipe to the Tauri app is gone.
The app reads our stdout/stderr through a pipe. When the server outlives
it (keep-running mode, or a sidecar the next launch reuses), every later
print()/tqdm write raises "[Errno 32] Broken pipe" and fails whatever
request triggered it, e.g. POST /captures.
"""
def __init__(self, stream):
self._stream = stream
def write(self, s):
try:
return self._stream.write(s)
except (OSError, ValueError):
self._stream = open(os.devnull, 'w')
return len(s)
def writelines(self, lines):
for line in lines:
self.write(line)
def flush(self):
try:
self._stream.flush()
except (OSError, ValueError):
self._stream = open(os.devnull, 'w')
def __getattr__(self, name):
return getattr(self._stream, name)
sys.stdout = _PipeSafeStream(sys.stdout)
sys.stderr = _PipeSafeStream(sys.stderr)
# PyInstaller + multiprocessing: child processes re-execute the frozen binary
# with internal arguments. freeze_support() handles this and exits early.
import multiprocessing
+36 -10
View File
@@ -15,21 +15,30 @@ from ..database import Generation as DBGeneration, GenerationVersion as DBGenera
from .. import config
def _get_versions_for_generation(generation_id: str, db: Session) -> tuple:
"""Get versions list and active version ID for a generation."""
def _get_versions_for_generations(generation_ids: list[str], db: Session) -> dict:
"""Fetch versions for many generations in a single query.
Returns a mapping of ``generation_id -> (versions, active_version_id)``
using the same shape as ``_get_versions_for_generation()``, so callers
can batch a whole page of generations without an N+1 query.
"""
import json
ids = list(dict.fromkeys(generation_ids))
if not ids:
return {}
versions_rows = (
db.query(DBGenerationVersion)
.filter_by(generation_id=generation_id)
.filter(DBGenerationVersion.generation_id.in_(ids))
.order_by(DBGenerationVersion.created_at)
.all()
)
if not versions_rows:
return None, None
versions = []
active_version_id = None
versions_by_generation: dict[str, list] = {}
active_by_generation: dict[str, Optional[str]] = {}
for v in versions_rows:
versions = versions_by_generation.setdefault(v.generation_id, [])
effects_chain = None
if v.effects_chain:
try:
@@ -47,9 +56,20 @@ def _get_versions_for_generation(generation_id: str, db: Session) -> tuple:
created_at=v.created_at,
))
if v.is_default:
active_version_id = v.id
active_by_generation[v.generation_id] = v.id
return versions, active_version_id
return {
generation_id: (
versions_by_generation.get(generation_id),
active_by_generation.get(generation_id),
)
for generation_id in ids
}
def _get_versions_for_generation(generation_id: str, db: Session) -> tuple:
"""Get versions list and active version ID for a single generation."""
return _get_versions_for_generations([generation_id], db)[generation_id]
async def create_generation(
@@ -205,10 +225,16 @@ async def list_generations(
# Execute query
results = q.all()
# Fetch versions for every generation on this page with a single
# query instead of one SELECT per generation (N+1).
versions_by_generation = _get_versions_for_generations(
[generation.id for generation, _ in results], db
)
# Convert to HistoryResponse with profile_name
items = []
for generation, profile_name in results:
versions, active_version_id = _get_versions_for_generation(generation.id, db)
versions, active_version_id = versions_by_generation[generation.id]
items.append(HistoryResponse(
id=generation.id,
profile_id=generation.profile_id,
+80
View File
@@ -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")
+74
View File
@@ -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 -1
View File
@@ -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
+158
View File
@@ -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"