mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -07:00
Merge origin/main into prep/pr-1031
This commit is contained in:
@@ -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"
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,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.
|
||||
"""
|
||||
|
||||
@@ -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."
|
||||
),
|
||||
|
||||
@@ -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
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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),
|
||||
),
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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