mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-28 06:35:18 -07:00
Compare commits
24
Commits
ui-testing
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
51f49dea19 | ||
|
|
80610d880e | ||
|
|
397051ba44 | ||
|
|
1ba935e83b | ||
|
|
6a6f4643da | ||
|
|
44ef8daba3 | ||
|
|
68ece25a80 | ||
|
|
1db0fdf645 | ||
|
|
2a001fd63f | ||
|
|
e5813304ef | ||
|
|
a5773807a5 | ||
|
|
ed54347e81 | ||
|
|
624f6a2140 | ||
|
|
669f85024f | ||
|
|
52f8d8dd38 | ||
|
|
fb1e16d2ce | ||
|
|
f750596364 | ||
|
|
91cd6df108 | ||
|
|
484a39ad9f | ||
|
|
3bfcbdc819 | ||
|
|
190bc5e8a8 | ||
|
|
80af641b61 | ||
|
|
6936789a88 | ||
|
|
f3eca34d33 |
+2
-1
@@ -8,7 +8,8 @@ tauri/
|
|||||||
landing/
|
landing/
|
||||||
docs/
|
docs/
|
||||||
mlx-test/
|
mlx-test/
|
||||||
scripts/
|
scripts/*
|
||||||
|
!scripts/rocm-entrypoint.sh
|
||||||
|
|
||||||
# Dependencies & build artifacts (rebuilt in Docker)
|
# Dependencies & build artifacts (rebuilt in Docker)
|
||||||
node_modules/
|
node_modules/
|
||||||
|
|||||||
@@ -0,0 +1,2 @@
|
|||||||
|
package.json text eol=lf
|
||||||
|
scripts/*.sh text eol=lf
|
||||||
+1
-1
@@ -91,7 +91,7 @@ On Windows, to build with CUDA support for local testing:
|
|||||||
just build-local # Build CPU + CUDA server binaries + Tauri installer
|
just build-local # Build CPU + CUDA server binaries + Tauri installer
|
||||||
```
|
```
|
||||||
|
|
||||||
This builds the CPU sidecar (bundled with the app), the CUDA binary (placed in `%APPDATA%/com.voicebox.app/backends/` for runtime GPU switching), and the installable Tauri app.
|
This builds the CPU sidecar (bundled with the app), the CUDA binary (placed in `%APPDATA%/sh.voicebox.app/backends/` for runtime GPU switching), and the installable Tauri app.
|
||||||
|
|
||||||
Creates platform-specific installers (`.dmg`, `.msi`, `.AppImage`) in `tauri/src-tauri/target/release/bundle/`.
|
Creates platform-specific installers (`.dmg`, `.msi`, `.AppImage`) in `tauri/src-tauri/target/release/bundle/`.
|
||||||
|
|
||||||
|
|||||||
+10
-3
@@ -20,8 +20,11 @@ COPY package.json bun.lock CHANGELOG.md ./
|
|||||||
COPY app/ ./app/
|
COPY app/ ./app/
|
||||||
COPY web/ ./web/
|
COPY web/ ./web/
|
||||||
|
|
||||||
# Strip workspaces not needed for web build, and fix trailing comma
|
# Normalize line endings first (a Windows CRLF checkout would otherwise
|
||||||
RUN sed -i '/"tauri"/d; /"landing"/d' package.json && \
|
# defeat the `-z 's/,\n ]/…/'` match below, since it's LF-anchored), then
|
||||||
|
# strip workspaces not needed for web build, and fix trailing comma
|
||||||
|
RUN sed -i 's/\r$//' package.json && \
|
||||||
|
sed -i '/"tauri"/d; /"landing"/d' package.json && \
|
||||||
sed -i -z 's/,\n ]/\n ]/' package.json
|
sed -i -z 's/,\n ]/\n ]/' package.json
|
||||||
RUN bun install --no-save
|
RUN bun install --no-save
|
||||||
# Build frontend (skip tsc — upstream has pre-existing type errors)
|
# Build frontend (skip tsc — upstream has pre-existing type errors)
|
||||||
@@ -100,7 +103,11 @@ EXPOSE 17493
|
|||||||
HEALTHCHECK --interval=30s --timeout=10s --retries=3 --start-period=60s \
|
HEALTHCHECK --interval=30s --timeout=10s --retries=3 --start-period=60s \
|
||||||
CMD curl -f http://localhost:17493/health || exit 1
|
CMD curl -f http://localhost:17493/health || exit 1
|
||||||
|
|
||||||
# Entrypoint joins GPU groups then drops to the voicebox user
|
# Entrypoint joins GPU groups then drops to the voicebox user.
|
||||||
|
# Normalize CRLF (a Windows checkout otherwise leaves the shebang as
|
||||||
|
# `#!/bin/sh\r`, which Linux can't resolve — reported as a misleading
|
||||||
|
# "no such file or directory" even though the file exists).
|
||||||
COPY --chmod=755 scripts/rocm-entrypoint.sh /usr/local/bin/entrypoint.sh
|
COPY --chmod=755 scripts/rocm-entrypoint.sh /usr/local/bin/entrypoint.sh
|
||||||
|
RUN sed -i 's/\r$//' /usr/local/bin/entrypoint.sh
|
||||||
ENTRYPOINT ["/usr/local/bin/entrypoint.sh"]
|
ENTRYPOINT ["/usr/local/bin/entrypoint.sh"]
|
||||||
CMD ["uvicorn", "backend.main:app", "--host", "0.0.0.0", "--port", "17493"]
|
CMD ["uvicorn", "backend.main:app", "--host", "0.0.0.0", "--port", "17493"]
|
||||||
|
|||||||
@@ -139,7 +139,7 @@ export function EngineModelSelector({ form, compact, selectedProfile }: EngineMo
|
|||||||
<SelectValue />
|
<SelectValue />
|
||||||
</SelectTrigger>
|
</SelectTrigger>
|
||||||
</FormControl>
|
</FormControl>
|
||||||
<SelectContent>
|
<SelectContent side={compact ? 'top' : undefined}>
|
||||||
{availableOptions.map((opt) => (
|
{availableOptions.map((opt) => (
|
||||||
<SelectItem key={opt.value} value={opt.value} className={itemClass}>
|
<SelectItem key={opt.value} value={opt.value} className={itemClass}>
|
||||||
{opt.label}
|
{opt.label}
|
||||||
|
|||||||
@@ -555,7 +555,7 @@ export function FloatingGenerateBox({
|
|||||||
<SelectTrigger className="h-8 text-xs bg-card border-border rounded-full hover:bg-background/50 transition-all w-full">
|
<SelectTrigger className="h-8 text-xs bg-card border-border rounded-full hover:bg-background/50 transition-all w-full">
|
||||||
<SelectValue placeholder={t('generation.voiceSelector.placeholder')} />
|
<SelectValue placeholder={t('generation.voiceSelector.placeholder')} />
|
||||||
</SelectTrigger>
|
</SelectTrigger>
|
||||||
<SelectContent>
|
<SelectContent side="top">
|
||||||
{profiles?.map((profile) => (
|
{profiles?.map((profile) => (
|
||||||
<SelectItem key={profile.id} value={profile.id} className="text-xs">
|
<SelectItem key={profile.id} value={profile.id} className="text-xs">
|
||||||
{profile.name}
|
{profile.name}
|
||||||
@@ -582,7 +582,7 @@ export function FloatingGenerateBox({
|
|||||||
<SelectValue />
|
<SelectValue />
|
||||||
</SelectTrigger>
|
</SelectTrigger>
|
||||||
</FormControl>
|
</FormControl>
|
||||||
<SelectContent>
|
<SelectContent side="top">
|
||||||
{engineLangs.map((lang) => (
|
{engineLangs.map((lang) => (
|
||||||
<SelectItem key={lang.value} value={lang.value} className="text-xs">
|
<SelectItem key={lang.value} value={lang.value} className="text-xs">
|
||||||
{lang.label}
|
{lang.label}
|
||||||
@@ -610,7 +610,7 @@ export function FloatingGenerateBox({
|
|||||||
<SelectTrigger className="h-8 text-xs bg-card border-border rounded-full hover:bg-background/50 transition-all">
|
<SelectTrigger className="h-8 text-xs bg-card border-border rounded-full hover:bg-background/50 transition-all">
|
||||||
<SelectValue placeholder={t('generation.effects.none')} />
|
<SelectValue placeholder={t('generation.effects.none')} />
|
||||||
</SelectTrigger>
|
</SelectTrigger>
|
||||||
<SelectContent>
|
<SelectContent side="top">
|
||||||
<SelectItem value="none" className="text-xs">
|
<SelectItem value="none" className="text-xs">
|
||||||
{t('generation.effects.none')}
|
{t('generation.effects.none')}
|
||||||
</SelectItem>
|
</SelectItem>
|
||||||
|
|||||||
@@ -47,12 +47,14 @@ export function useExportGeneration() {
|
|||||||
mutationFn: async ({ generationId, text }: { generationId: string; text: string }) => {
|
mutationFn: async ({ generationId, text }: { generationId: string; text: string }) => {
|
||||||
const blob = await apiClient.exportGeneration(generationId);
|
const blob = await apiClient.exportGeneration(generationId);
|
||||||
|
|
||||||
// Create safe filename from text
|
// Create safe filename from text. Append a short id so exports of
|
||||||
|
// similarly-worded generations don't collide on the same filename
|
||||||
|
// (the first 30 chars are frequently identical).
|
||||||
const safeText = text
|
const safeText = text
|
||||||
.substring(0, 30)
|
.substring(0, 30)
|
||||||
.replace(/[^a-z0-9]/gi, '-')
|
.replace(/[^a-z0-9]/gi, '-')
|
||||||
.toLowerCase();
|
.toLowerCase();
|
||||||
const filename = `generation-${safeText}.voicebox.zip`;
|
const filename = `generation-${safeText}-${generationId.substring(0, 8)}.voicebox.zip`;
|
||||||
|
|
||||||
await platform.filesystem.saveFile(filename, blob, [
|
await platform.filesystem.saveFile(filename, blob, [
|
||||||
{
|
{
|
||||||
@@ -73,12 +75,14 @@ export function useExportGenerationAudio() {
|
|||||||
mutationFn: async ({ generationId, text }: { generationId: string; text: string }) => {
|
mutationFn: async ({ generationId, text }: { generationId: string; text: string }) => {
|
||||||
const blob = await apiClient.exportGenerationAudio(generationId);
|
const blob = await apiClient.exportGenerationAudio(generationId);
|
||||||
|
|
||||||
// Create safe filename from text
|
// Create safe filename from text. Append a short id so exports of
|
||||||
|
// similarly-worded generations don't collide on the same filename
|
||||||
|
// (the first 30 chars are frequently identical).
|
||||||
const safeText = text
|
const safeText = text
|
||||||
.substring(0, 30)
|
.substring(0, 30)
|
||||||
.replace(/[^a-z0-9]/gi, '-')
|
.replace(/[^a-z0-9]/gi, '-')
|
||||||
.toLowerCase();
|
.toLowerCase();
|
||||||
const filename = `${safeText}.wav`;
|
const filename = `${safeText}-${generationId.substring(0, 8)}.wav`;
|
||||||
|
|
||||||
await platform.filesystem.saveFile(filename, blob, [
|
await platform.filesystem.saveFile(filename, blob, [
|
||||||
{
|
{
|
||||||
|
|||||||
+14
-13
@@ -25,27 +25,28 @@ function getDateLocale() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
export function formatDate(date: string | Date): string {
|
// Backend timestamps are naive UTC — append `Z` so JS doesn't parse a
|
||||||
let dateObj: Date;
|
// timezone-less date-time string as local time.
|
||||||
if (typeof date === 'string') {
|
function parseServerDate(date: string | Date): Date {
|
||||||
const dateStr = date.trim();
|
if (typeof date !== 'string') {
|
||||||
if (!dateStr.includes('Z') && !dateStr.match(/[+-]\d{2}:\d{2}$/)) {
|
return date;
|
||||||
dateObj = new Date(`${dateStr}Z`);
|
|
||||||
} else {
|
|
||||||
dateObj = new Date(dateStr);
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
dateObj = date;
|
|
||||||
}
|
}
|
||||||
|
const dateStr = date.trim();
|
||||||
|
if (!dateStr.includes('Z') && !dateStr.match(/[+-]\d{2}:\d{2}$/)) {
|
||||||
|
return new Date(`${dateStr}Z`);
|
||||||
|
}
|
||||||
|
return new Date(dateStr);
|
||||||
|
}
|
||||||
|
|
||||||
return formatDistance(dateObj, new Date(), {
|
export function formatDate(date: string | Date): string {
|
||||||
|
return formatDistance(parseServerDate(date), new Date(), {
|
||||||
addSuffix: true,
|
addSuffix: true,
|
||||||
locale: getDateLocale(),
|
locale: getDateLocale(),
|
||||||
}).replace(/^about /i, '');
|
}).replace(/^about /i, '');
|
||||||
}
|
}
|
||||||
|
|
||||||
export function formatAbsoluteDate(date: string | Date): string {
|
export function formatAbsoluteDate(date: string | Date): string {
|
||||||
const dateObj = typeof date === 'string' ? new Date(date) : date;
|
const dateObj = parseServerDate(date);
|
||||||
return dateObj.toLocaleString(i18n.language, {
|
return dateObj.toLocaleString(i18n.language, {
|
||||||
month: 'short',
|
month: 'short',
|
||||||
day: 'numeric',
|
day: 'numeric',
|
||||||
|
|||||||
@@ -38,6 +38,13 @@ logging.basicConfig(
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# An empty HSA_OVERRIDE_GFX_VERSION poisons the ROCm HSA runtime. It is
|
||||||
|
# treated as "force-empty" and no GPU is detected, even natively supported
|
||||||
|
# ones (e.g. gfx1201 / RX 9070 on ROCm 7.2). docker-compose can't
|
||||||
|
# conditionally omit an env var, so we clean it up here before torch loads.
|
||||||
|
if not os.environ.get("HSA_OVERRIDE_GFX_VERSION"):
|
||||||
|
os.environ.pop("HSA_OVERRIDE_GFX_VERSION", None)
|
||||||
|
|
||||||
# AMD GPU environment variables must be set before torch import
|
# AMD GPU environment variables must be set before torch import
|
||||||
# Only set HSA_OVERRIDE_GFX_VERSION for older GPUs that need it.
|
# Only set HSA_OVERRIDE_GFX_VERSION for older GPUs that need it.
|
||||||
# RDNA 3+ (gfx1100+) and RDNA 4 (gfx1200+) are natively supported by ROCm
|
# RDNA 3+ (gfx1100+) and RDNA 4 (gfx1200+) are natively supported by ROCm
|
||||||
|
|||||||
@@ -56,6 +56,7 @@ class ModelConfig:
|
|||||||
model_size: str = "default"
|
model_size: str = "default"
|
||||||
size_mb: int = 0
|
size_mb: int = 0
|
||||||
needs_trim: bool = False
|
needs_trim: bool = False
|
||||||
|
retries_runaway: bool = False
|
||||||
supports_instruct: bool = False
|
supports_instruct: bool = False
|
||||||
languages: list[str] = field(default_factory=lambda: ["en"])
|
languages: list[str] = field(default_factory=lambda: ["en"])
|
||||||
|
|
||||||
@@ -232,6 +233,10 @@ def _get_qwen_model_configs() -> list[ModelConfig]:
|
|||||||
repo_1_7b = "Qwen/Qwen3-TTS-12Hz-1.7B-Base"
|
repo_1_7b = "Qwen/Qwen3-TTS-12Hz-1.7B-Base"
|
||||||
repo_0_6b = "Qwen/Qwen3-TTS-12Hz-0.6B-Base"
|
repo_0_6b = "Qwen/Qwen3-TTS-12Hz-0.6B-Base"
|
||||||
|
|
||||||
|
# mlx-audio can continue after an EOS miss with silence followed by
|
||||||
|
# codec noise. Retry only the affected text as smaller chunks.
|
||||||
|
retries_runaway = backend_type == "mlx"
|
||||||
|
|
||||||
return [
|
return [
|
||||||
ModelConfig(
|
ModelConfig(
|
||||||
model_name="qwen-tts-1.7B",
|
model_name="qwen-tts-1.7B",
|
||||||
@@ -240,6 +245,7 @@ def _get_qwen_model_configs() -> list[ModelConfig]:
|
|||||||
hf_repo_id=repo_1_7b,
|
hf_repo_id=repo_1_7b,
|
||||||
model_size="1.7B",
|
model_size="1.7B",
|
||||||
size_mb=3500,
|
size_mb=3500,
|
||||||
|
retries_runaway=retries_runaway,
|
||||||
supports_instruct=False, # Base model drops instruct silently
|
supports_instruct=False, # Base model drops instruct silently
|
||||||
languages=["zh", "en", "ja", "ko", "de", "fr", "ru", "pt", "es", "it"],
|
languages=["zh", "en", "ja", "ko", "de", "fr", "ru", "pt", "es", "it"],
|
||||||
),
|
),
|
||||||
@@ -250,6 +256,7 @@ def _get_qwen_model_configs() -> list[ModelConfig]:
|
|||||||
hf_repo_id=repo_0_6b,
|
hf_repo_id=repo_0_6b,
|
||||||
model_size="0.6B",
|
model_size="0.6B",
|
||||||
size_mb=1200,
|
size_mb=1200,
|
||||||
|
retries_runaway=retries_runaway,
|
||||||
supports_instruct=False,
|
supports_instruct=False,
|
||||||
languages=["zh", "en", "ja", "ko", "de", "fr", "ru", "pt", "es", "it"],
|
languages=["zh", "en", "ja", "ko", "de", "fr", "ru", "pt", "es", "it"],
|
||||||
),
|
),
|
||||||
@@ -504,6 +511,14 @@ def engine_needs_trim(engine: str) -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def engine_retries_runaway(engine: str) -> bool:
|
||||||
|
"""Whether unstable output should be retried in smaller chunks."""
|
||||||
|
for cfg in get_tts_model_configs():
|
||||||
|
if cfg.engine == engine:
|
||||||
|
return cfg.retries_runaway
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def engine_has_model_sizes(engine: str) -> bool:
|
def engine_has_model_sizes(engine: str) -> bool:
|
||||||
"""Whether this engine supports multiple model sizes (only Qwen currently)."""
|
"""Whether this engine supports multiple model sizes (only Qwen currently)."""
|
||||||
configs = [c for c in get_tts_model_configs() if c.engine == engine]
|
configs = [c for c in get_tts_model_configs() if c.engine == engine]
|
||||||
|
|||||||
@@ -248,9 +248,13 @@ class HumeTadaBackend:
|
|||||||
audio = audio.T # (samples, channels) -> (channels, samples)
|
audio = audio.T # (samples, channels) -> (channels, samples)
|
||||||
audio = audio.to(device)
|
audio = audio.to(device)
|
||||||
|
|
||||||
# Encode with forced alignment
|
# Encode with forced alignment.
|
||||||
|
# Must run under inference_mode: encoder params still require
|
||||||
|
# grad by default, and an autograd graph across the DAC/Snake
|
||||||
|
# stack can balloon VRAM far past the model footprint (#890).
|
||||||
text_arg = [reference_text] if reference_text else None
|
text_arg = [reference_text] if reference_text else None
|
||||||
prompt = self.encoder(audio, text=text_arg, sample_rate=sr)
|
with torch.inference_mode():
|
||||||
|
prompt = self.encoder(audio, text=text_arg, sample_rate=sr)
|
||||||
|
|
||||||
# Serialize EncoderOutput to a dict of CPU tensors for caching
|
# Serialize EncoderOutput to a dict of CPU tensors for caching
|
||||||
prompt_dict = {}
|
prompt_dict = {}
|
||||||
|
|||||||
@@ -330,6 +330,9 @@ def build_server(cuda=False, rocm=False):
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if sys.version_info >= (3, 13):
|
||||||
|
args.extend(["--hidden-import", "audioop"])
|
||||||
|
|
||||||
# Add CUDA/ROCm-specific hidden imports
|
# Add CUDA/ROCm-specific hidden imports
|
||||||
if cuda or rocm:
|
if cuda or rocm:
|
||||||
variant = "ROCm" if rocm else "CUDA"
|
variant = "ROCm" if rocm else "CUDA"
|
||||||
|
|||||||
@@ -80,6 +80,11 @@ def resolve_storage_path(path: str | Path | None) -> Path | None:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
stored_path = Path(path)
|
stored_path = Path(path)
|
||||||
|
# Empty paths (e.g. failed generations) must not resolve to the data
|
||||||
|
# dir itself, which exists and would defeat the callers' 404 guards.
|
||||||
|
# Path("") is truthy, so check parts rather than the raw value.
|
||||||
|
if not stored_path.parts:
|
||||||
|
return None
|
||||||
if stored_path.is_absolute():
|
if stored_path.is_absolute():
|
||||||
rebased_path = _path_relative_to_any_data_dir(stored_path)
|
rebased_path = _path_relative_to_any_data_dir(stored_path)
|
||||||
if rebased_path is not None:
|
if rebased_path is not None:
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ import base64 as b64
|
|||||||
import logging
|
import logging
|
||||||
import tempfile
|
import tempfile
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any, Literal
|
||||||
|
|
||||||
from fastmcp import FastMCP
|
from fastmcp import FastMCP
|
||||||
|
|
||||||
@@ -49,6 +49,7 @@ def register_tools(mcp: FastMCP) -> None:
|
|||||||
engine: str | None = None,
|
engine: str | None = None,
|
||||||
personality: bool | None = None,
|
personality: bool | None = None,
|
||||||
language: str | None = None,
|
language: str | None = None,
|
||||||
|
model_size: Literal["1.7B", "0.6B", "1B", "3B"] | None = None,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""Speak ``text`` in a voice profile.
|
"""Speak ``text`` in a voice profile.
|
||||||
|
|
||||||
@@ -61,6 +62,12 @@ def register_tools(mcp: FastMCP) -> None:
|
|||||||
LLM before TTS. When omitted, the per-client binding's
|
LLM before TTS. When omitted, the per-client binding's
|
||||||
``default_personality`` flag decides; when that is unset, the
|
``default_personality`` flag decides; when that is unset, the
|
||||||
default is plain TTS.
|
default is plain TTS.
|
||||||
|
|
||||||
|
``model_size`` selects a model variant for engines that ship more
|
||||||
|
than one — ``qwen`` and ``qwen_custom_voice`` accept "1.7B" (default)
|
||||||
|
or "0.6B"; ``tada`` accepts "1B" or "3B". Other engines ignore it.
|
||||||
|
Omit to use the engine default. Requesting a smaller variant (e.g.
|
||||||
|
"0.6B") is faster and avoids reloading a heavier model between calls.
|
||||||
"""
|
"""
|
||||||
from ..database.models import MCPClientBinding
|
from ..database.models import MCPClientBinding
|
||||||
|
|
||||||
@@ -99,6 +106,7 @@ def register_tools(mcp: FastMCP) -> None:
|
|||||||
engine=resolved_engine,
|
engine=resolved_engine,
|
||||||
language=language,
|
language=language,
|
||||||
personality=use_persona,
|
personality=use_persona,
|
||||||
|
model_size=model_size,
|
||||||
db=db,
|
db=db,
|
||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
@@ -228,18 +236,23 @@ async def _speak(
|
|||||||
engine: str | None,
|
engine: str | None,
|
||||||
language: str | None,
|
language: str | None,
|
||||||
personality: bool,
|
personality: bool,
|
||||||
|
model_size: str | None = None,
|
||||||
db,
|
db,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""Delegate to POST /generate — the route handles personality-rewrite
|
"""Delegate to POST /generate — the route handles personality-rewrite
|
||||||
internally when ``personality=true`` and the profile has a prompt."""
|
internally when ``personality=true`` and the profile has a prompt."""
|
||||||
from ..routes.generations import generate_speech
|
from ..routes.generations import generate_speech
|
||||||
|
|
||||||
|
# model_size=None is intentional: generate_speech normalizes it to the
|
||||||
|
# engine default (see routes/generations.py), so an omitted size behaves
|
||||||
|
# exactly like the REST /generate endpoint with no model_size in the body.
|
||||||
req = models.GenerationRequest(
|
req = models.GenerationRequest(
|
||||||
profile_id=profile_id,
|
profile_id=profile_id,
|
||||||
text=text,
|
text=text,
|
||||||
language=language or "en",
|
language=language or "en",
|
||||||
engine=engine,
|
engine=engine,
|
||||||
personality=personality,
|
personality=personality,
|
||||||
|
model_size=model_size,
|
||||||
)
|
)
|
||||||
generation = await generate_speech(req, db)
|
generation = await generate_speech(req, db)
|
||||||
return _speak_response(generation, profile_name, source="mcp")
|
return _speak_response(generation, profile_name, source="mcp")
|
||||||
|
|||||||
@@ -16,7 +16,8 @@ miniaudio>=1.59
|
|||||||
# mlx_audio.stt.load) works fine on transformers 4.57.x in practice.
|
# mlx_audio.stt.load) works fine on transformers 4.57.x in practice.
|
||||||
#
|
#
|
||||||
# Install it via `pip install --no-deps mlx-audio==0.4.1` after this file
|
# Install it via `pip install --no-deps mlx-audio==0.4.1` after this file
|
||||||
# (see .github/workflows/release.yml). Most other mlx-audio runtime deps
|
# (see .github/workflows/release.yml and the setup-python recipe in the
|
||||||
|
# justfile). Most other mlx-audio runtime deps
|
||||||
# (huggingface_hub, librosa, mlx-lm, numba, numpy, protobuf, pyloudnorm,
|
# (huggingface_hub, librosa, mlx-lm, numba, numpy, protobuf, pyloudnorm,
|
||||||
# sounddevice, tqdm) are already in requirements.txt or pulled in by
|
# sounddevice, tqdm) are already in requirements.txt or pulled in by
|
||||||
# other engines.
|
# other engines.
|
||||||
|
|||||||
@@ -53,6 +53,7 @@ en_core_web_sm @ https://github.com/explosion/spacy-models/releases/download/en_
|
|||||||
unidic-lite>=1.0.8
|
unidic-lite>=1.0.8
|
||||||
|
|
||||||
# Audio processing
|
# Audio processing
|
||||||
|
audioop-lts>=0.2.1; python_version >= "3.13"
|
||||||
librosa>=0.10.0
|
librosa>=0.10.0
|
||||||
soundfile>=0.12.0
|
soundfile>=0.12.0
|
||||||
numpy>=1.24.0,<2.0
|
numpy>=1.24.0,<2.0
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ async def get_version_audio(version_id: str, db: Session = Depends(get_db)):
|
|||||||
raise HTTPException(status_code=404, detail="Version not found")
|
raise HTTPException(status_code=404, detail="Version not found")
|
||||||
|
|
||||||
audio_path = config.resolve_storage_path(version.audio_path)
|
audio_path = config.resolve_storage_path(version.audio_path)
|
||||||
if audio_path is None or not audio_path.exists():
|
if audio_path is None or not audio_path.is_file():
|
||||||
raise HTTPException(status_code=404, detail="Audio file not found")
|
raise HTTPException(status_code=404, detail="Audio file not found")
|
||||||
|
|
||||||
return FileResponse(
|
return FileResponse(
|
||||||
@@ -52,8 +52,13 @@ async def get_audio(generation_id: str, db: Session = Depends(get_db)):
|
|||||||
raise HTTPException(status_code=404, detail="Generation not found")
|
raise HTTPException(status_code=404, detail="Generation not found")
|
||||||
|
|
||||||
audio_path = config.resolve_storage_path(generation.audio_path)
|
audio_path = config.resolve_storage_path(generation.audio_path)
|
||||||
if audio_path is None or not audio_path.exists():
|
if audio_path is None or not audio_path.is_file():
|
||||||
raise HTTPException(status_code=404, detail="Audio file not found")
|
detail = (
|
||||||
|
"Generation failed; no audio available"
|
||||||
|
if generation.status == "failed"
|
||||||
|
else "Audio file not found"
|
||||||
|
)
|
||||||
|
raise HTTPException(status_code=404, detail=detail)
|
||||||
|
|
||||||
return FileResponse(
|
return FileResponse(
|
||||||
audio_path,
|
audio_path,
|
||||||
@@ -72,7 +77,7 @@ async def get_sample_audio(sample_id: str, db: Session = Depends(get_db)):
|
|||||||
raise HTTPException(status_code=404, detail="Sample not found")
|
raise HTTPException(status_code=404, detail="Sample not found")
|
||||||
|
|
||||||
audio_path = config.resolve_storage_path(sample.audio_path)
|
audio_path = config.resolve_storage_path(sample.audio_path)
|
||||||
if audio_path is None or not audio_path.exists():
|
if audio_path is None or not audio_path.is_file():
|
||||||
raise HTTPException(status_code=404, detail="Audio file not found")
|
raise HTTPException(status_code=404, detail="Audio file not found")
|
||||||
|
|
||||||
return FileResponse(
|
return FileResponse(
|
||||||
|
|||||||
@@ -321,7 +321,13 @@ async def stream_speech(
|
|||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""Generate speech and stream the WAV audio directly without saving to disk."""
|
"""Generate speech and stream the WAV audio directly without saving to disk."""
|
||||||
from ..backends import get_tts_backend_for_engine, ensure_model_cached_or_raise, load_engine_model, engine_needs_trim
|
from ..backends import (
|
||||||
|
engine_needs_trim,
|
||||||
|
engine_retries_runaway,
|
||||||
|
ensure_model_cached_or_raise,
|
||||||
|
get_tts_backend_for_engine,
|
||||||
|
load_engine_model,
|
||||||
|
)
|
||||||
|
|
||||||
profile = await profiles.get_profile(data.profile_id, db)
|
profile = await profiles.get_profile(data.profile_id, db)
|
||||||
if not profile:
|
if not profile:
|
||||||
@@ -347,10 +353,15 @@ async def stream_speech(
|
|||||||
from ..utils.chunked_tts import generate_chunked
|
from ..utils.chunked_tts import generate_chunked
|
||||||
|
|
||||||
trim_fn = None
|
trim_fn = None
|
||||||
|
runaway_detector = None
|
||||||
if engine_needs_trim(engine):
|
if engine_needs_trim(engine):
|
||||||
from ..utils.audio import trim_tts_output
|
from ..utils.audio import trim_tts_output
|
||||||
|
|
||||||
trim_fn = trim_tts_output
|
trim_fn = trim_tts_output
|
||||||
|
if engine_retries_runaway(engine):
|
||||||
|
from ..utils.audio import has_tts_runaway
|
||||||
|
|
||||||
|
runaway_detector = has_tts_runaway
|
||||||
|
|
||||||
audio, sample_rate = await generate_chunked(
|
audio, sample_rate = await generate_chunked(
|
||||||
tts_model,
|
tts_model,
|
||||||
@@ -362,6 +373,7 @@ async def stream_speech(
|
|||||||
max_chunk_chars=data.max_chunk_chars,
|
max_chunk_chars=data.max_chunk_chars,
|
||||||
crossfade_ms=data.crossfade_ms,
|
crossfade_ms=data.crossfade_ms,
|
||||||
trim_fn=trim_fn,
|
trim_fn=trim_fn,
|
||||||
|
runaway_detector=runaway_detector,
|
||||||
)
|
)
|
||||||
|
|
||||||
effects_chain_config = None
|
effects_chain_config = None
|
||||||
|
|||||||
@@ -151,7 +151,9 @@ async def export_generation(
|
|||||||
safe_text = "".join(c for c in generation.text[:30] if c.isalnum() or c in (" ", "-", "_")).strip()
|
safe_text = "".join(c for c in generation.text[:30] if c.isalnum() or c in (" ", "-", "_")).strip()
|
||||||
if not safe_text:
|
if not safe_text:
|
||||||
safe_text = "generation"
|
safe_text = "generation"
|
||||||
filename = f"generation-{safe_text}.voicebox.zip"
|
# Append a short id so exports of similarly-worded generations don't collide
|
||||||
|
# on the same filename (the first 30 chars are frequently identical).
|
||||||
|
filename = f"generation-{safe_text}-{generation_id[:8]}.voicebox.zip"
|
||||||
|
|
||||||
return StreamingResponse(
|
return StreamingResponse(
|
||||||
io.BytesIO(zip_bytes),
|
io.BytesIO(zip_bytes),
|
||||||
@@ -180,7 +182,9 @@ async def export_generation_audio(
|
|||||||
safe_text = "".join(c for c in generation.text[:30] if c.isalnum() or c in (" ", "-", "_")).strip()
|
safe_text = "".join(c for c in generation.text[:30] if c.isalnum() or c in (" ", "-", "_")).strip()
|
||||||
if not safe_text:
|
if not safe_text:
|
||||||
safe_text = "generation"
|
safe_text = "generation"
|
||||||
filename = f"{safe_text}.wav"
|
# Append a short id so exports of similarly-worded generations don't collide
|
||||||
|
# on the same filename (the first 30 chars are frequently identical).
|
||||||
|
filename = f"{safe_text}-{generation_id[:8]}.wav"
|
||||||
|
|
||||||
return FileResponse(
|
return FileResponse(
|
||||||
audio_path,
|
audio_path,
|
||||||
|
|||||||
@@ -232,7 +232,7 @@ async def upload_profile_avatar(
|
|||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""Upload or update avatar image for a profile."""
|
"""Upload or update avatar image for a profile."""
|
||||||
with tempfile.NamedTemporaryFile(delete=False, suffix=Path(file.filename).suffix) as tmp:
|
with tempfile.NamedTemporaryFile(delete=False, suffix=Path(file.filename or "").suffix) as tmp:
|
||||||
content = await file.read()
|
content = await file.read()
|
||||||
tmp.write(content)
|
tmp.write(content)
|
||||||
tmp_path = tmp.name
|
tmp_path = tmp.name
|
||||||
|
|||||||
@@ -35,13 +35,25 @@ async def transcribe_audio(
|
|||||||
tmp.write(chunk)
|
tmp.write(chunk)
|
||||||
tmp_path = tmp.name
|
tmp_path = tmp.name
|
||||||
|
|
||||||
|
stt_path = tmp_path
|
||||||
try:
|
try:
|
||||||
from ..utils.audio import load_audio
|
from ..utils.audio import load_audio, save_audio
|
||||||
from ..backends import WHISPER_HF_REPOS
|
from ..backends import WHISPER_HF_REPOS
|
||||||
|
|
||||||
audio, sr = await asyncio.to_thread(load_audio, tmp_path)
|
audio, sr = await asyncio.to_thread(load_audio, tmp_path)
|
||||||
duration = len(audio) / sr
|
duration = len(audio) / sr
|
||||||
|
|
||||||
|
# The STT backend (mlx_audio.stt -> miniaudio) only decodes
|
||||||
|
# WAV/FLAC/MP3/Vorbis, so browser recordings uploaded as WebM/Opus
|
||||||
|
# fail with "unsupported file format" (issue: web-mode dictation).
|
||||||
|
# librosa already decoded the file above (it falls back to
|
||||||
|
# audioread/ffmpeg for exotic containers), so re-encode that PCM to a
|
||||||
|
# temp WAV and hand *that* to Whisper. WAV inputs pass through
|
||||||
|
# unchanged.
|
||||||
|
if file_suffix != ".wav":
|
||||||
|
stt_path = f"{tmp_path}.stt.wav"
|
||||||
|
await asyncio.to_thread(save_audio, audio, stt_path, sr)
|
||||||
|
|
||||||
whisper_model = transcribe.get_whisper_model()
|
whisper_model = transcribe.get_whisper_model()
|
||||||
model_size = model if model else whisper_model.model_size
|
model_size = model if model else whisper_model.model_size
|
||||||
|
|
||||||
@@ -76,7 +88,7 @@ async def transcribe_audio(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
text = await whisper_model.transcribe(tmp_path, language, model_size)
|
text = await whisper_model.transcribe(stt_path, language, model_size)
|
||||||
|
|
||||||
return models.TranscriptionResponse(
|
return models.TranscriptionResponse(
|
||||||
text=text,
|
text=text,
|
||||||
@@ -89,3 +101,5 @@ async def transcribe_audio(
|
|||||||
raise HTTPException(status_code=500, detail=str(e))
|
raise HTTPException(status_code=500, detail=str(e))
|
||||||
finally:
|
finally:
|
||||||
Path(tmp_path).unlink(missing_ok=True)
|
Path(tmp_path).unlink(missing_ok=True)
|
||||||
|
if stt_path != tmp_path:
|
||||||
|
Path(stt_path).unlink(missing_ok=True)
|
||||||
|
|||||||
@@ -48,9 +48,14 @@ async def run_generation(
|
|||||||
This is the single entry point for all background generation work.
|
This is the single entry point for all background generation work.
|
||||||
It is designed to be enqueued via ``services.task_queue.enqueue_generation``.
|
It is designed to be enqueued via ``services.task_queue.enqueue_generation``.
|
||||||
"""
|
"""
|
||||||
from ..backends import load_engine_model, get_tts_backend_for_engine, engine_needs_trim
|
from ..backends import (
|
||||||
|
engine_needs_trim,
|
||||||
|
engine_retries_runaway,
|
||||||
|
get_tts_backend_for_engine,
|
||||||
|
load_engine_model,
|
||||||
|
)
|
||||||
from ..utils.chunked_tts import generate_chunked
|
from ..utils.chunked_tts import generate_chunked
|
||||||
from ..utils.audio import normalize_audio, save_audio, trim_tts_output
|
from ..utils.audio import has_tts_runaway, normalize_audio, save_audio, trim_tts_output
|
||||||
|
|
||||||
task_manager = get_task_manager()
|
task_manager = get_task_manager()
|
||||||
bg_db = next(get_db())
|
bg_db = next(get_db())
|
||||||
@@ -72,12 +77,14 @@ async def run_generation(
|
|||||||
|
|
||||||
await history.update_generation_status(generation_id, "generating", bg_db)
|
await history.update_generation_status(generation_id, "generating", bg_db)
|
||||||
trim_fn = trim_tts_output if engine_needs_trim(engine) else None
|
trim_fn = trim_tts_output if engine_needs_trim(engine) else None
|
||||||
|
runaway_detector = has_tts_runaway if engine_retries_runaway(engine) else None
|
||||||
|
|
||||||
gen_kwargs: dict = dict(
|
gen_kwargs: dict = dict(
|
||||||
language=language,
|
language=language,
|
||||||
seed=seed if mode != "regenerate" else None,
|
seed=seed if mode != "regenerate" else None,
|
||||||
instruct=instruct,
|
instruct=instruct,
|
||||||
trim_fn=trim_fn,
|
trim_fn=trim_fn,
|
||||||
|
runaway_detector=runaway_detector,
|
||||||
)
|
)
|
||||||
if max_chunk_chars is not None:
|
if max_chunk_chars is not None:
|
||||||
gen_kwargs["max_chunk_chars"] = max_chunk_chars
|
gen_kwargs["max_chunk_chars"] = max_chunk_chars
|
||||||
@@ -267,9 +274,14 @@ async def generate_audio_sync(
|
|||||||
normalize, then encodes in-memory via :func:`tts.audio_to_wav_bytes`
|
normalize, then encodes in-memory via :func:`tts.audio_to_wav_bytes`
|
||||||
(same helper ``/generate/stream`` uses).
|
(same helper ``/generate/stream`` uses).
|
||||||
"""
|
"""
|
||||||
from ..backends import load_engine_model, get_tts_backend_for_engine, engine_needs_trim
|
from ..backends import (
|
||||||
|
engine_needs_trim,
|
||||||
|
engine_retries_runaway,
|
||||||
|
get_tts_backend_for_engine,
|
||||||
|
load_engine_model,
|
||||||
|
)
|
||||||
from ..utils.chunked_tts import generate_chunked
|
from ..utils.chunked_tts import generate_chunked
|
||||||
from ..utils.audio import normalize_audio, trim_tts_output
|
from ..utils.audio import has_tts_runaway, normalize_audio, trim_tts_output
|
||||||
from . import tts
|
from . import tts
|
||||||
|
|
||||||
bg_db = next(get_db())
|
bg_db = next(get_db())
|
||||||
@@ -287,12 +299,14 @@ async def generate_audio_sync(
|
|||||||
bg_db.close()
|
bg_db.close()
|
||||||
|
|
||||||
trim_fn = trim_tts_output if engine_needs_trim(engine) else None
|
trim_fn = trim_tts_output if engine_needs_trim(engine) else None
|
||||||
|
runaway_detector = has_tts_runaway if engine_retries_runaway(engine) else None
|
||||||
|
|
||||||
gen_kwargs: dict = dict(
|
gen_kwargs: dict = dict(
|
||||||
language=language,
|
language=language,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
instruct=instruct,
|
instruct=instruct,
|
||||||
trim_fn=trim_fn,
|
trim_fn=trim_fn,
|
||||||
|
runaway_detector=runaway_detector,
|
||||||
)
|
)
|
||||||
if max_chunk_chars is not None:
|
if max_chunk_chars is not None:
|
||||||
gen_kwargs["max_chunk_chars"] = max_chunk_chars
|
gen_kwargs["max_chunk_chars"] = max_chunk_chars
|
||||||
|
|||||||
@@ -125,12 +125,24 @@ async def list_stories(
|
|||||||
"""
|
"""
|
||||||
stories = db.query(DBStory).order_by(DBStory.updated_at.desc()).all()
|
stories = db.query(DBStory).order_by(DBStory.updated_at.desc()).all()
|
||||||
|
|
||||||
|
if not stories:
|
||||||
|
return []
|
||||||
|
|
||||||
|
# Batch-fetch all story item counts in one query to avoid an N+1 pattern
|
||||||
|
# (previously there was one COUNT query per story in the loop below).
|
||||||
|
story_ids = [s.id for s in stories]
|
||||||
|
count_rows = (
|
||||||
|
db.query(DBStoryItem.story_id, func.count(DBStoryItem.id).label("cnt"))
|
||||||
|
.filter(DBStoryItem.story_id.in_(story_ids))
|
||||||
|
.group_by(DBStoryItem.story_id)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
item_counts = {row.story_id: row.cnt for row in count_rows}
|
||||||
|
|
||||||
result = []
|
result = []
|
||||||
for story in stories:
|
for story in stories:
|
||||||
item_count = db.query(func.count(DBStoryItem.id)).filter(DBStoryItem.story_id == story.id).scalar()
|
|
||||||
|
|
||||||
response = StoryResponse.model_validate(story)
|
response = StoryResponse.model_validate(story)
|
||||||
response.item_count = item_count
|
response.item_count = item_counts.get(story.id, 0)
|
||||||
result.append(response)
|
result.append(response)
|
||||||
|
|
||||||
return result
|
return result
|
||||||
|
|||||||
@@ -0,0 +1,164 @@
|
|||||||
|
"""
|
||||||
|
Regression tests for GET /audio/{generation_id} on failed generations.
|
||||||
|
|
||||||
|
A failed generation stores an empty ``audio_path``. Previously,
|
||||||
|
``config.resolve_storage_path("")`` resolved to the data directory itself,
|
||||||
|
which exists, so the route's 404 guard passed and ``FileResponse`` raised
|
||||||
|
``RuntimeError: File at path .../data is not a file`` — a 500 instead of
|
||||||
|
a clean 404.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python -m pytest backend/tests/test_audio_failed_generation.py -v
|
||||||
|
"""
|
||||||
|
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from fastapi import FastAPI
|
||||||
|
from sqlalchemy import create_engine
|
||||||
|
from sqlalchemy.orm import sessionmaker
|
||||||
|
from starlette.testclient import TestClient
|
||||||
|
|
||||||
|
# Repo root on sys.path so ``backend`` imports as a package (the audio
|
||||||
|
# routes use package-relative imports).
|
||||||
|
sys.path.insert(0, str(Path(__file__).parent.parent.parent))
|
||||||
|
|
||||||
|
from backend import config
|
||||||
|
from backend.database import (
|
||||||
|
Base,
|
||||||
|
Generation,
|
||||||
|
GenerationVersion,
|
||||||
|
ProfileSample,
|
||||||
|
VoiceProfile,
|
||||||
|
get_db,
|
||||||
|
)
|
||||||
|
from backend.routes.audio import router as audio_router
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_storage_path_empty_returns_none():
|
||||||
|
"""An empty stored path must not resolve to the data dir itself."""
|
||||||
|
assert config.resolve_storage_path("") is None
|
||||||
|
assert config.resolve_storage_path(None) is None
|
||||||
|
# Path("") is truthy, so it must be rejected via its (empty) parts.
|
||||||
|
assert config.resolve_storage_path(Path("")) is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def client(tmp_path, monkeypatch):
|
||||||
|
"""Minimal app with only the audio routes and a temp sqlite DB."""
|
||||||
|
monkeypatch.setattr(config, "_data_dir", tmp_path)
|
||||||
|
# An existing directory that a stored audio_path may wrongly point to.
|
||||||
|
(tmp_path / "somedir").mkdir()
|
||||||
|
|
||||||
|
engine = create_engine(
|
||||||
|
f"sqlite:///{tmp_path / 'test.db'}",
|
||||||
|
connect_args={"check_same_thread": False},
|
||||||
|
)
|
||||||
|
Base.metadata.create_all(bind=engine)
|
||||||
|
testing_session_local = sessionmaker(autocommit=False, autoflush=False, bind=engine)
|
||||||
|
|
||||||
|
session = testing_session_local()
|
||||||
|
profile = VoiceProfile(id="profile-1", name="Test Profile")
|
||||||
|
session.add(profile)
|
||||||
|
|
||||||
|
session.add_all(
|
||||||
|
[
|
||||||
|
Generation(
|
||||||
|
id="gen-failed-empty",
|
||||||
|
profile_id="profile-1",
|
||||||
|
text="failed generation",
|
||||||
|
audio_path="",
|
||||||
|
status="failed",
|
||||||
|
error="engine exploded",
|
||||||
|
),
|
||||||
|
Generation(
|
||||||
|
id="gen-failed-null",
|
||||||
|
profile_id="profile-1",
|
||||||
|
text="failed generation",
|
||||||
|
audio_path=None,
|
||||||
|
status="failed",
|
||||||
|
),
|
||||||
|
Generation(
|
||||||
|
id="gen-missing-file",
|
||||||
|
profile_id="profile-1",
|
||||||
|
text="completed but file deleted",
|
||||||
|
audio_path="generations/does-not-exist.wav",
|
||||||
|
status="completed",
|
||||||
|
),
|
||||||
|
Generation(
|
||||||
|
id="gen-with-version",
|
||||||
|
profile_id="profile-1",
|
||||||
|
text="generation with a broken version",
|
||||||
|
audio_path="somedir",
|
||||||
|
status="completed",
|
||||||
|
),
|
||||||
|
GenerationVersion(
|
||||||
|
id="version-dir",
|
||||||
|
generation_id="gen-with-version",
|
||||||
|
label="original",
|
||||||
|
audio_path="somedir",
|
||||||
|
),
|
||||||
|
ProfileSample(
|
||||||
|
id="sample-dir",
|
||||||
|
profile_id="profile-1",
|
||||||
|
audio_path="somedir",
|
||||||
|
reference_text="sample pointing at a directory",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
session.commit()
|
||||||
|
session.close()
|
||||||
|
|
||||||
|
app = FastAPI()
|
||||||
|
app.include_router(audio_router)
|
||||||
|
|
||||||
|
def override_get_db():
|
||||||
|
db = testing_session_local()
|
||||||
|
try:
|
||||||
|
yield db
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
app.dependency_overrides[get_db] = override_get_db
|
||||||
|
return TestClient(app)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("generation_id", ["gen-failed-empty", "gen-failed-null"])
|
||||||
|
def test_failed_generation_returns_404(client, generation_id):
|
||||||
|
"""Failed generations (empty/null audio_path) get a clean 404, not a 500."""
|
||||||
|
response = client.get(f"/audio/{generation_id}")
|
||||||
|
assert response.status_code == 404
|
||||||
|
assert response.json()["detail"] == "Generation failed; no audio available"
|
||||||
|
|
||||||
|
|
||||||
|
def test_missing_audio_file_returns_404(client):
|
||||||
|
"""A completed generation whose file vanished still 404s."""
|
||||||
|
response = client.get("/audio/gen-missing-file")
|
||||||
|
assert response.status_code == 404
|
||||||
|
assert response.json()["detail"] == "Audio file not found"
|
||||||
|
|
||||||
|
|
||||||
|
def test_unknown_generation_returns_404(client):
|
||||||
|
response = client.get("/audio/no-such-generation")
|
||||||
|
assert response.status_code == 404
|
||||||
|
assert response.json()["detail"] == "Generation not found"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"url",
|
||||||
|
[
|
||||||
|
"/audio/gen-with-version",
|
||||||
|
"/audio/version/version-dir",
|
||||||
|
"/samples/sample-dir",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_audio_path_pointing_at_directory_returns_404(client, url):
|
||||||
|
"""A stored path resolving to an existing directory must 404, not 500.
|
||||||
|
|
||||||
|
Guards the is_file() checks: a directory passes exists() and would
|
||||||
|
crash FileResponse.
|
||||||
|
"""
|
||||||
|
response = client.get(url)
|
||||||
|
assert response.status_code == 404
|
||||||
|
assert response.json()["detail"] == "Audio file not found"
|
||||||
@@ -0,0 +1,123 @@
|
|||||||
|
"""
|
||||||
|
Regression tests for issue #852: audioop removed from Python 3.13 stdlib.
|
||||||
|
|
||||||
|
Voice sample validation imports audioop transitively (librosa → audioread).
|
||||||
|
The audioop-lts backport must be declared in requirements and bundled in
|
||||||
|
PyInstaller builds on 3.13+.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import re
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||||
|
|
||||||
|
from build_binary import build_server
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def backend_dir():
|
||||||
|
return Path(__file__).parent.parent
|
||||||
|
|
||||||
|
|
||||||
|
class TestAudioopRequirements:
|
||||||
|
def test_requirements_declare_audioop_lts_for_python_313(self, backend_dir):
|
||||||
|
content = (backend_dir / "requirements.txt").read_text()
|
||||||
|
assert re.search(
|
||||||
|
r"^audioop-lts.*python_version\s*>=\s*['\"]3\.13['\"]",
|
||||||
|
content,
|
||||||
|
re.MULTILINE,
|
||||||
|
), "requirements.txt must pin audioop-lts for Python 3.13+"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(sys.version_info < (3, 13), reason="Python 3.13+ only")
|
||||||
|
class TestAudioopRuntime:
|
||||||
|
def test_audioop_importable(self):
|
||||||
|
import audioop # noqa: F401
|
||||||
|
|
||||||
|
def test_validate_reference_wav_does_not_fail_on_missing_audioop(self, tmp_path):
|
||||||
|
import numpy as np
|
||||||
|
import soundfile as sf
|
||||||
|
from utils.audio import validate_and_load_reference_audio
|
||||||
|
|
||||||
|
sr = 24000
|
||||||
|
t = np.arange(int(sr * 3), dtype=np.float32) / sr
|
||||||
|
audio = (0.3 * np.sin(2 * np.pi * 220 * t)).astype(np.float32)
|
||||||
|
path = tmp_path / "reference.wav"
|
||||||
|
sf.write(str(path), audio, sr)
|
||||||
|
|
||||||
|
ok, err, out_audio, out_sr = validate_and_load_reference_audio(str(path))
|
||||||
|
|
||||||
|
assert ok, err
|
||||||
|
assert out_audio is not None
|
||||||
|
assert out_sr == sr
|
||||||
|
assert "audioop" not in (err or "").lower()
|
||||||
|
|
||||||
|
|
||||||
|
class TestAudioopBuildArgs:
|
||||||
|
@staticmethod
|
||||||
|
def _hidden_imports(args):
|
||||||
|
imports = []
|
||||||
|
for i, arg in enumerate(args):
|
||||||
|
if arg == "--hidden-import" and i + 1 < len(args):
|
||||||
|
imports.append(args[i + 1])
|
||||||
|
return imports
|
||||||
|
|
||||||
|
def test_pyinstaller_includes_audioop_on_python_313(self):
|
||||||
|
class FakeVersionInfo(tuple):
|
||||||
|
@property
|
||||||
|
def major(self):
|
||||||
|
return self[0]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def minor(self):
|
||||||
|
return self[1]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def micro(self):
|
||||||
|
return self[2]
|
||||||
|
|
||||||
|
fake_313 = FakeVersionInfo((3, 13, 0, "final", 0))
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("build_binary.PyInstaller.__main__.run") as mock_run,
|
||||||
|
patch("build_binary.platform.system", return_value="Linux"),
|
||||||
|
patch("build_binary.is_apple_silicon", return_value=False),
|
||||||
|
patch("build_binary.os.chdir"),
|
||||||
|
patch("build_binary.sys.version_info", fake_313),
|
||||||
|
):
|
||||||
|
build_server()
|
||||||
|
args = mock_run.call_args[0][0]
|
||||||
|
|
||||||
|
assert "audioop" in self._hidden_imports(args)
|
||||||
|
|
||||||
|
def test_pyinstaller_omits_audioop_on_python_312(self):
|
||||||
|
class FakeVersionInfo(tuple):
|
||||||
|
@property
|
||||||
|
def major(self):
|
||||||
|
return self[0]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def minor(self):
|
||||||
|
return self[1]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def micro(self):
|
||||||
|
return self[2]
|
||||||
|
|
||||||
|
fake_312 = FakeVersionInfo((3, 12, 0, "final", 0))
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("build_binary.PyInstaller.__main__.run") as mock_run,
|
||||||
|
patch("build_binary.platform.system", return_value="Linux"),
|
||||||
|
patch("build_binary.is_apple_silicon", return_value=False),
|
||||||
|
patch("build_binary.os.chdir"),
|
||||||
|
patch("build_binary.sys.version_info", fake_312),
|
||||||
|
):
|
||||||
|
build_server()
|
||||||
|
args = mock_run.call_args[0][0]
|
||||||
|
|
||||||
|
assert "audioop" not in self._hidden_imports(args)
|
||||||
@@ -0,0 +1,68 @@
|
|||||||
|
"""Ensure TADA voice-prompt encoding disables autograd (#890)."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from unittest.mock import AsyncMock
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
import soundfile as sf
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from backend.backends.hume_backend import HumeTadaBackend
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _FakeEncoderOutput:
|
||||||
|
emb: torch.Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class _GradTrackingEncoder:
|
||||||
|
"""Raises unless called under torch.inference_mode()."""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.called_under_inference_mode = False
|
||||||
|
|
||||||
|
def __call__(self, audio, text=None, sample_rate=None):
|
||||||
|
self.called_under_inference_mode = torch.is_inference_mode_enabled()
|
||||||
|
if not self.called_under_inference_mode:
|
||||||
|
raise AssertionError("encoder forward must run under inference_mode")
|
||||||
|
# Touch a requires_grad tensor the way Snake1d alpha would.
|
||||||
|
alpha = torch.nn.Parameter(torch.ones(1, device=audio.device))
|
||||||
|
_ = audio.mean() * alpha
|
||||||
|
return _FakeEncoderOutput(emb=torch.zeros(1, 4, device=audio.device))
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_create_voice_prompt_runs_encoder_under_inference_mode(tmp_path, monkeypatch):
|
||||||
|
wav = tmp_path / "ref.wav"
|
||||||
|
sf.write(str(wav), np.zeros(24000, dtype=np.float32), 24000)
|
||||||
|
|
||||||
|
backend = HumeTadaBackend()
|
||||||
|
backend.model = object() # mark loaded
|
||||||
|
backend.model_size = "1B"
|
||||||
|
backend._device = "cpu"
|
||||||
|
encoder = _GradTrackingEncoder()
|
||||||
|
backend.encoder = encoder
|
||||||
|
|
||||||
|
monkeypatch.setattr(backend, "load_model", AsyncMock(return_value=None))
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"backend.backends.hume_backend.get_cached_voice_prompt",
|
||||||
|
lambda key: None,
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"backend.backends.hume_backend.cache_voice_prompt",
|
||||||
|
lambda key, value: None,
|
||||||
|
)
|
||||||
|
|
||||||
|
prompt, from_cache = await backend.create_voice_prompt(
|
||||||
|
str(wav),
|
||||||
|
reference_text="hello world",
|
||||||
|
use_cache=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert from_cache is False
|
||||||
|
assert encoder.called_under_inference_mode is True
|
||||||
|
assert isinstance(prompt["emb"], torch.Tensor)
|
||||||
|
assert prompt["emb"].device.type == "cpu"
|
||||||
@@ -0,0 +1,91 @@
|
|||||||
|
"""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
|
||||||
|
schema default ("1.7B") — there was no way to reach 0.6B (or TADA's 1B/3B)
|
||||||
|
through MCP. These tests pin the fix: ``_speak`` now forwards ``model_size``
|
||||||
|
straight into the request, matching the REST ``/generate`` surface.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from pydantic import ValidationError
|
||||||
|
|
||||||
|
import backend.routes.generations as generations
|
||||||
|
from backend.mcp_server import tools
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeGeneration:
|
||||||
|
"""Minimal stand-in for GenerationResponse consumed by ``_speak_response``."""
|
||||||
|
|
||||||
|
def model_dump(self, mode="json"):
|
||||||
|
return {"id": "gen-test", "status": "generating"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def captured_request(monkeypatch):
|
||||||
|
"""Replace the real (torch-backed) generate_speech with a capturing stub.
|
||||||
|
|
||||||
|
``_speak`` imports ``generate_speech`` lazily from ``routes.generations``,
|
||||||
|
so patching the attribute on that module intercepts the call and lets us
|
||||||
|
inspect the ``GenerationRequest`` it would have run.
|
||||||
|
"""
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
async def fake_generate_speech(req, db):
|
||||||
|
captured["req"] = req
|
||||||
|
return _FakeGeneration()
|
||||||
|
|
||||||
|
monkeypatch.setattr(generations, "generate_speech", fake_generate_speech)
|
||||||
|
# Isolate the unit from the MCP event bus — _speak_response fires a
|
||||||
|
# speak-start event we don't care about here.
|
||||||
|
monkeypatch.setattr(tools.mcp_events, "publish", lambda *a, **k: None)
|
||||||
|
return captured
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_speak_forwards_explicit_model_size(captured_request):
|
||||||
|
await tools._speak(
|
||||||
|
profile_id="p1",
|
||||||
|
profile_name="Morgan",
|
||||||
|
text="hello",
|
||||||
|
engine="qwen",
|
||||||
|
language="en",
|
||||||
|
personality=False,
|
||||||
|
model_size="0.6B",
|
||||||
|
db=None,
|
||||||
|
)
|
||||||
|
assert captured_request["req"].model_size == "0.6B"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_speak_omitted_model_size_is_none(captured_request):
|
||||||
|
# Omitted → None; generate_speech normalizes None to the engine default,
|
||||||
|
# so this reproduces the pre-fix behaviour for callers that don't ask.
|
||||||
|
await tools._speak(
|
||||||
|
profile_id="p1",
|
||||||
|
profile_name="Morgan",
|
||||||
|
text="hello",
|
||||||
|
engine="qwen",
|
||||||
|
language="en",
|
||||||
|
personality=False,
|
||||||
|
db=None,
|
||||||
|
)
|
||||||
|
assert captured_request["req"].model_size is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_speak_rejects_invalid_model_size(captured_request):
|
||||||
|
# The GenerationRequest schema pattern is the single source of truth for
|
||||||
|
# valid sizes; a bad value is rejected before any generation runs.
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
await tools._speak(
|
||||||
|
profile_id="p1",
|
||||||
|
profile_name="Morgan",
|
||||||
|
text="hello",
|
||||||
|
engine="qwen",
|
||||||
|
language="en",
|
||||||
|
personality=False,
|
||||||
|
model_size="9B",
|
||||||
|
db=None,
|
||||||
|
)
|
||||||
|
assert "req" not in captured_request
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
"""
|
||||||
|
Smoke test for the MLX backend dependencies on Apple Silicon.
|
||||||
|
|
||||||
|
Guards the `--no-deps` install of mlx-audio/mlx-lm done by `just setup-python`
|
||||||
|
and release.yml: those packages skip their declared dependencies (transformers
|
||||||
|
>=5.x conflict), so a missing transitive dep only surfaces at import time.
|
||||||
|
This test fails fast if the MLX STT/TTS entry points the backend uses stop
|
||||||
|
importing (e.g. the `miniaudio` regression from issue #505).
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python -m pytest backend/tests/test_mlx_smoke.py -v
|
||||||
|
"""
|
||||||
|
|
||||||
|
import platform
|
||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
pytestmark = pytest.mark.skipif(
|
||||||
|
not (sys.platform == "darwin" and platform.machine() == "arm64"),
|
||||||
|
reason="MLX packages are only installed on Apple Silicon macOS",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_mlx_core_runs():
|
||||||
|
"""The MLX runtime itself works (Metal array op)."""
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
assert mx.array([1, 2]).sum().item() == 3
|
||||||
|
|
||||||
|
|
||||||
|
def test_mlx_audio_tts_entry_point():
|
||||||
|
"""`from mlx_audio.tts import load` — used by MLXBackend.load_model_async."""
|
||||||
|
from mlx_audio.tts import load
|
||||||
|
|
||||||
|
assert callable(load)
|
||||||
|
|
||||||
|
|
||||||
|
def test_mlx_audio_stt_entry_point():
|
||||||
|
"""`from mlx_audio.stt import load` — used by the Whisper MLX STT path.
|
||||||
|
|
||||||
|
Importing mlx_audio.stt also pulls in miniaudio, so this catches the
|
||||||
|
ModuleNotFoundError from issue #505 on fresh installs.
|
||||||
|
"""
|
||||||
|
from mlx_audio.stt import load
|
||||||
|
|
||||||
|
assert callable(load)
|
||||||
|
|
||||||
|
|
||||||
|
def test_mlx_lm_entry_points():
|
||||||
|
"""`mlx_lm.load` / `mlx_lm.generate` — used by qwen_llm_backend."""
|
||||||
|
from mlx_lm import generate, load
|
||||||
|
|
||||||
|
assert callable(load)
|
||||||
|
assert callable(generate)
|
||||||
@@ -0,0 +1,117 @@
|
|||||||
|
"""Regression coverage for runaway MLX Qwen TTS output."""
|
||||||
|
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from backend.backends import engine_needs_trim, engine_retries_runaway
|
||||||
|
from backend.utils.audio import has_tts_runaway
|
||||||
|
from backend.utils.chunked_tts import generate_chunked
|
||||||
|
|
||||||
|
SAMPLE_RATE = 1000
|
||||||
|
|
||||||
|
|
||||||
|
def test_mlx_qwen_enables_runaway_retry_without_aggressive_trim():
|
||||||
|
with patch("backend.backends.get_backend_type", return_value="mlx"):
|
||||||
|
assert engine_needs_trim("qwen") is False
|
||||||
|
assert engine_retries_runaway("qwen") is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_pytorch_qwen_keeps_runaway_retry_disabled():
|
||||||
|
with patch("backend.backends.get_backend_type", return_value="pytorch"):
|
||||||
|
assert engine_needs_trim("qwen") is False
|
||||||
|
assert engine_retries_runaway("qwen") is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_detector_flags_long_internal_silence():
|
||||||
|
speech = np.full(2 * SAMPLE_RATE, 0.2, dtype=np.float32)
|
||||||
|
runaway_gap = np.zeros(2500, dtype=np.float32)
|
||||||
|
hallucinated_noise = np.full(2 * SAMPLE_RATE, 0.8, dtype=np.float32)
|
||||||
|
audio = np.concatenate([speech, runaway_gap, hallucinated_noise])
|
||||||
|
|
||||||
|
assert has_tts_runaway(audio, SAMPLE_RATE) is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_detector_ignores_normal_internal_pause():
|
||||||
|
speech = np.full(SAMPLE_RATE, 0.2, dtype=np.float32)
|
||||||
|
normal_pause = np.zeros(1200, dtype=np.float32)
|
||||||
|
audio = np.concatenate([speech, normal_pause, speech])
|
||||||
|
|
||||||
|
assert has_tts_runaway(audio, SAMPLE_RATE) is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_trailing_silence_is_not_a_runaway():
|
||||||
|
speech = np.full(SAMPLE_RATE, 0.2, dtype=np.float32)
|
||||||
|
trailing_silence = np.zeros(2 * SAMPLE_RATE, dtype=np.float32)
|
||||||
|
|
||||||
|
assert (
|
||||||
|
has_tts_runaway(
|
||||||
|
np.concatenate([speech, trailing_silence]),
|
||||||
|
SAMPLE_RATE,
|
||||||
|
)
|
||||||
|
is False
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_runaway_chunk_is_retried_as_smaller_chunks():
|
||||||
|
class FakeBackend:
|
||||||
|
def __init__(self):
|
||||||
|
self.calls = []
|
||||||
|
|
||||||
|
async def generate(self, text, *_args):
|
||||||
|
self.calls.append(text)
|
||||||
|
if len(text) > 200:
|
||||||
|
speech = np.full(SAMPLE_RATE, 0.2, dtype=np.float32)
|
||||||
|
silence = np.zeros(2500, dtype=np.float32)
|
||||||
|
noise = np.full(SAMPLE_RATE, 0.8, dtype=np.float32)
|
||||||
|
return np.concatenate([speech, silence, noise]), SAMPLE_RATE
|
||||||
|
return np.full(SAMPLE_RATE, 0.2, dtype=np.float32), SAMPLE_RATE
|
||||||
|
|
||||||
|
backend = FakeBackend()
|
||||||
|
text = f"{'A' * 119}. {'B' * 119}."
|
||||||
|
|
||||||
|
audio, sample_rate = await generate_chunked(
|
||||||
|
backend,
|
||||||
|
text,
|
||||||
|
{},
|
||||||
|
max_chunk_chars=800,
|
||||||
|
crossfade_ms=50,
|
||||||
|
runaway_detector=has_tts_runaway,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert sample_rate == SAMPLE_RATE
|
||||||
|
assert backend.calls == [text, f"{'A' * 119}.", f"{'B' * 119}."]
|
||||||
|
assert len(audio) == 1950
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_persistent_runaway_fails_instead_of_returning_corrupt_audio():
|
||||||
|
class AlwaysRunawayBackend:
|
||||||
|
def __init__(self):
|
||||||
|
self.calls = []
|
||||||
|
|
||||||
|
async def generate(self, text, *_args):
|
||||||
|
self.calls.append(text)
|
||||||
|
speech = np.full(SAMPLE_RATE, 0.2, dtype=np.float32)
|
||||||
|
silence = np.zeros(2500, dtype=np.float32)
|
||||||
|
noise = np.full(SAMPLE_RATE, 0.8, dtype=np.float32)
|
||||||
|
return np.concatenate([speech, silence, noise]), SAMPLE_RATE
|
||||||
|
|
||||||
|
backend = AlwaysRunawayBackend()
|
||||||
|
text = f"{'A' * 119}. {'B' * 119}."
|
||||||
|
|
||||||
|
with pytest.raises(
|
||||||
|
RuntimeError,
|
||||||
|
match="remained unstable after retrying smaller text chunks",
|
||||||
|
):
|
||||||
|
await generate_chunked(
|
||||||
|
backend,
|
||||||
|
text,
|
||||||
|
{},
|
||||||
|
max_chunk_chars=800,
|
||||||
|
runaway_detector=has_tts_runaway,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert [len(call) for call in backend.calls] == [241, 120, 100]
|
||||||
@@ -110,6 +110,43 @@ def save_audio(
|
|||||||
raise OSError(f"Failed to save audio to {path}: {e}") from e
|
raise OSError(f"Failed to save audio to {path}: {e}") from e
|
||||||
|
|
||||||
|
|
||||||
|
def has_tts_runaway(
|
||||||
|
audio: np.ndarray,
|
||||||
|
sample_rate: int = 24000,
|
||||||
|
frame_ms: int = 20,
|
||||||
|
silence_threshold_db: float = -40.0,
|
||||||
|
max_internal_silence_ms: int = 2000,
|
||||||
|
) -> bool:
|
||||||
|
"""Detect speech followed by a long silence and then more output.
|
||||||
|
|
||||||
|
This shape is a reliable signal that a TTS model missed EOS and resumed
|
||||||
|
with hallucinated speech or codec noise. Leading and trailing silence do
|
||||||
|
not count because they are not bounded by non-silent audio.
|
||||||
|
"""
|
||||||
|
frame_len = int(sample_rate * frame_ms / 1000)
|
||||||
|
if frame_len == 0 or len(audio) < frame_len:
|
||||||
|
return False
|
||||||
|
|
||||||
|
n_frames = len(audio) // frame_len
|
||||||
|
threshold_linear = 10 ** (silence_threshold_db / 20)
|
||||||
|
max_silence_frames = int(max_internal_silence_ms / frame_ms)
|
||||||
|
seen_speech = False
|
||||||
|
consecutive_silence = 0
|
||||||
|
|
||||||
|
for i in range(n_frames):
|
||||||
|
frame = audio[i * frame_len : (i + 1) * frame_len]
|
||||||
|
is_speech = np.sqrt(np.mean(frame**2)) >= threshold_linear
|
||||||
|
if is_speech:
|
||||||
|
if seen_speech and consecutive_silence >= max_silence_frames:
|
||||||
|
return True
|
||||||
|
seen_speech = True
|
||||||
|
consecutive_silence = 0
|
||||||
|
elif seen_speech:
|
||||||
|
consecutive_silence += 1
|
||||||
|
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def trim_tts_output(
|
def trim_tts_output(
|
||||||
audio: np.ndarray,
|
audio: np.ndarray,
|
||||||
sample_rate: int = 24000,
|
sample_rate: int = 24000,
|
||||||
|
|||||||
@@ -20,6 +20,8 @@ logger = logging.getLogger("voicebox.chunked-tts")
|
|||||||
# Default chunk size in characters. Can be overridden per-request via
|
# Default chunk size in characters. Can be overridden per-request via
|
||||||
# the ``max_chunk_chars`` field on GenerationRequest.
|
# the ``max_chunk_chars`` field on GenerationRequest.
|
||||||
DEFAULT_MAX_CHUNK_CHARS = 800
|
DEFAULT_MAX_CHUNK_CHARS = 800
|
||||||
|
MAX_RUNAWAY_RETRIES = 2
|
||||||
|
MIN_RUNAWAY_RETRY_CHARS = 100
|
||||||
|
|
||||||
# Common abbreviations that should NOT be treated as sentence endings.
|
# Common abbreviations that should NOT be treated as sentence endings.
|
||||||
# Lowercase for case-insensitive matching.
|
# Lowercase for case-insensitive matching.
|
||||||
@@ -211,6 +213,7 @@ async def generate_chunked(
|
|||||||
max_chunk_chars: int = DEFAULT_MAX_CHUNK_CHARS,
|
max_chunk_chars: int = DEFAULT_MAX_CHUNK_CHARS,
|
||||||
crossfade_ms: int = 50,
|
crossfade_ms: int = 50,
|
||||||
trim_fn=None,
|
trim_fn=None,
|
||||||
|
runaway_detector=None,
|
||||||
) -> Tuple[np.ndarray, int]:
|
) -> Tuple[np.ndarray, int]:
|
||||||
"""Generate audio with automatic chunking for long text.
|
"""Generate audio with automatic chunking for long text.
|
||||||
|
|
||||||
@@ -239,25 +242,75 @@ async def generate_chunked(
|
|||||||
Optional ``(audio, sample_rate) -> audio`` post-processing
|
Optional ``(audio, sample_rate) -> audio`` post-processing
|
||||||
function applied to each chunk before concatenation (e.g.
|
function applied to each chunk before concatenation (e.g.
|
||||||
``trim_tts_output`` for Chatterbox engines).
|
``trim_tts_output`` for Chatterbox engines).
|
||||||
|
runaway_detector : callable | None
|
||||||
|
Optional ``(audio, sample_rate) -> bool`` detector. When it flags
|
||||||
|
unstable output, the affected text is split in half and retried.
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
(audio, sample_rate) : Tuple[np.ndarray, int]
|
(audio, sample_rate) : Tuple[np.ndarray, int]
|
||||||
"""
|
"""
|
||||||
|
async def generate_one(
|
||||||
|
chunk_text: str,
|
||||||
|
chunk_seed: int | None,
|
||||||
|
retry_depth: int = 0,
|
||||||
|
) -> tuple[np.ndarray, int]:
|
||||||
|
chunk_audio, chunk_sr = await backend.generate(
|
||||||
|
chunk_text,
|
||||||
|
voice_prompt,
|
||||||
|
language,
|
||||||
|
chunk_seed,
|
||||||
|
instruct,
|
||||||
|
)
|
||||||
|
|
||||||
|
if runaway_detector is not None and runaway_detector(chunk_audio, chunk_sr):
|
||||||
|
if retry_depth >= MAX_RUNAWAY_RETRIES or len(chunk_text) <= MIN_RUNAWAY_RETRY_CHARS:
|
||||||
|
raise RuntimeError(
|
||||||
|
"TTS output remained unstable after retrying smaller text chunks"
|
||||||
|
)
|
||||||
|
|
||||||
|
retry_max_chars = max(MIN_RUNAWAY_RETRY_CHARS, len(chunk_text) // 2)
|
||||||
|
retry_chunks = split_text_into_chunks(chunk_text, retry_max_chars)
|
||||||
|
if len(retry_chunks) <= 1:
|
||||||
|
raise RuntimeError("Unable to split unstable TTS output for retry")
|
||||||
|
|
||||||
|
logger.warning(
|
||||||
|
"Detected unstable TTS output for %d chars; retrying as %d smaller chunks",
|
||||||
|
len(chunk_text),
|
||||||
|
len(retry_chunks),
|
||||||
|
)
|
||||||
|
retry_audio: list[np.ndarray] = []
|
||||||
|
for i, retry_text in enumerate(retry_chunks):
|
||||||
|
retry_seed = (
|
||||||
|
chunk_seed + ((retry_depth + 1) * 1000) + i
|
||||||
|
if chunk_seed is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
audio, sample_rate = await generate_one(
|
||||||
|
retry_text,
|
||||||
|
retry_seed,
|
||||||
|
retry_depth + 1,
|
||||||
|
)
|
||||||
|
retry_audio.append(np.asarray(audio, dtype=np.float32))
|
||||||
|
|
||||||
|
return (
|
||||||
|
concatenate_audio_chunks(
|
||||||
|
retry_audio,
|
||||||
|
sample_rate,
|
||||||
|
crossfade_ms=crossfade_ms,
|
||||||
|
),
|
||||||
|
sample_rate,
|
||||||
|
)
|
||||||
|
|
||||||
|
if trim_fn is not None:
|
||||||
|
chunk_audio = trim_fn(chunk_audio, chunk_sr)
|
||||||
|
return np.asarray(chunk_audio, dtype=np.float32), chunk_sr
|
||||||
|
|
||||||
chunks = split_text_into_chunks(text, max_chunk_chars)
|
chunks = split_text_into_chunks(text, max_chunk_chars)
|
||||||
|
|
||||||
if len(chunks) <= 1:
|
if len(chunks) <= 1:
|
||||||
# Short text — single-shot fast path
|
# Short text — single-shot fast path
|
||||||
audio, sample_rate = await backend.generate(
|
return await generate_one(text, seed)
|
||||||
text,
|
|
||||||
voice_prompt,
|
|
||||||
language,
|
|
||||||
seed,
|
|
||||||
instruct,
|
|
||||||
)
|
|
||||||
if trim_fn is not None:
|
|
||||||
audio = trim_fn(audio, sample_rate)
|
|
||||||
return audio, sample_rate
|
|
||||||
|
|
||||||
# Long text — chunked generation
|
# Long text — chunked generation
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -281,17 +334,12 @@ async def generate_chunked(
|
|||||||
# always produces the same output.
|
# always produces the same output.
|
||||||
chunk_seed = (seed + i) if seed is not None else None
|
chunk_seed = (seed + i) if seed is not None else None
|
||||||
|
|
||||||
chunk_audio, chunk_sr = await backend.generate(
|
chunk_audio, chunk_sr = await generate_one(
|
||||||
chunk_text,
|
chunk_text,
|
||||||
voice_prompt,
|
|
||||||
language,
|
|
||||||
chunk_seed,
|
chunk_seed,
|
||||||
instruct,
|
|
||||||
)
|
)
|
||||||
if trim_fn is not None:
|
|
||||||
chunk_audio = trim_fn(chunk_audio, chunk_sr)
|
|
||||||
|
|
||||||
audio_chunks.append(np.asarray(chunk_audio, dtype=np.float32))
|
audio_chunks.append(chunk_audio)
|
||||||
if sample_rate is None:
|
if sample_rate is None:
|
||||||
sample_rate = chunk_sr
|
sample_rate = chunk_sr
|
||||||
|
|
||||||
|
|||||||
@@ -34,3 +34,15 @@ services:
|
|||||||
|
|
||||||
# Tune the ROCm memory allocator
|
# Tune the ROCm memory allocator
|
||||||
- PYTORCH_HIP_ALLOC_CONF=garbage_collection_threshold:0.8,max_split_size_mb:512
|
- PYTORCH_HIP_ALLOC_CONF=garbage_collection_threshold:0.8,max_split_size_mb:512
|
||||||
|
|
||||||
|
# Redirect MIOpen kernel cache to a writable, persistent directory.
|
||||||
|
# Without this, MIOpen may fail to write its cache and throw
|
||||||
|
# miopenStatusUnknownError on fresh containers.
|
||||||
|
- MIOPEN_USER_DB_PATH=/app/data/cache/miopen_db
|
||||||
|
- MIOPEN_CUSTOM_CACHE_DIR=/app/data/cache/miopen_cache
|
||||||
|
|
||||||
|
# Use fast heuristics for kernel selection instead of exhaustive
|
||||||
|
# benchmarking. On RDNA4, exhaustive mode tries kernels that fail to
|
||||||
|
# allocate workspace memory (ptr: 0 size: 0), causing system stuttering
|
||||||
|
# on every generation even when the cache is present.
|
||||||
|
- MIOPEN_FIND_MODE=FAST
|
||||||
|
|||||||
@@ -49,6 +49,7 @@ class ModelConfig:
|
|||||||
model_size: str = "default"
|
model_size: str = "default"
|
||||||
size_mb: int = 0
|
size_mb: int = 0
|
||||||
needs_trim: bool = False
|
needs_trim: bool = False
|
||||||
|
retries_runaway: bool = False
|
||||||
supports_instruct: bool = False
|
supports_instruct: bool = False
|
||||||
languages: list[str] = field(default_factory=lambda: ["en"])
|
languages: list[str] = field(default_factory=lambda: ["en"])
|
||||||
```
|
```
|
||||||
@@ -59,6 +60,7 @@ Registry helpers in `backends/__init__.py` replace what used to be per-engine `i
|
|||||||
- `get_tts_model_configs()` — only TTS variants
|
- `get_tts_model_configs()` — only TTS variants
|
||||||
- `get_model_config(model_name)` — lookup by name
|
- `get_model_config(model_name)` — lookup by name
|
||||||
- `engine_needs_trim(engine)` — whether output should run through `trim_tts_output()`
|
- `engine_needs_trim(engine)` — whether output should run through `trim_tts_output()`
|
||||||
|
- `engine_retries_runaway(engine)` — whether unstable output should be retried as smaller chunks
|
||||||
- `load_engine_model(engine, model_size)` — downloads + loads, handles engines with multiple sizes
|
- `load_engine_model(engine, model_size)` — downloads + loads, handles engines with multiple sizes
|
||||||
- `get_tts_backend_for_engine(engine)` — thread-safe backend factory with double-checked locking
|
- `get_tts_backend_for_engine(engine)` — thread-safe backend factory with double-checked locking
|
||||||
|
|
||||||
@@ -152,7 +154,7 @@ The request path from frontend to audio file:
|
|||||||
|
|
||||||
6. **Inference** — the engine's `generate()` returns `(audio_array, sample_rate)`.
|
6. **Inference** — the engine's `generate()` returns `(audio_array, sample_rate)`.
|
||||||
|
|
||||||
7. **Post-process** — if `engine_needs_trim(engine)` is True, `trim_tts_output()` strips trailing silence. Effects chains (if any) are applied per generation version, not the clean version.
|
7. **Validate and post-process** — engines with `retries_runaway=True` retry unstable output as smaller chunks. If `engine_needs_trim(engine)` is True, `trim_tts_output()` strips trailing silence. Effects chains (if any) are applied per generation version, not the clean version.
|
||||||
|
|
||||||
8. **Persist** — audio is written to the generations directory, a row is inserted into the `generations` table, and the response includes the generation metadata.
|
8. **Persist** — audio is written to the generations directory, a row is inserted into the `generations` table, and the response includes the generation metadata.
|
||||||
|
|
||||||
|
|||||||
@@ -14,12 +14,12 @@ Make sure you have [installed Voicebox](/overview/installation) and launched the
|
|||||||
Voice profiles are the foundation of Voicebox. Each profile contains voice samples that the AI uses to clone the voice.
|
Voice profiles are the foundation of Voicebox. Each profile contains voice samples that the AI uses to clone the voice.
|
||||||
|
|
||||||
<Steps>
|
<Steps>
|
||||||
<Step title="Navigate to Profiles">
|
<Step title="Navigate to Voices">
|
||||||
Click the **Profiles** tab in the sidebar
|
Click the **Voices** tab in the sidebar
|
||||||
</Step>
|
</Step>
|
||||||
|
|
||||||
<Step title="Create New Profile">
|
<Step title="Create New Voice">
|
||||||
Click the **+ New Profile** button
|
Click the **+ New Voice** button
|
||||||
|
|
||||||
Fill in the details:
|
Fill in the details:
|
||||||
- **Name:** A descriptive name (e.g., "John Smith")
|
- **Name:** A descriptive name (e.g., "John Smith")
|
||||||
|
|||||||
@@ -72,6 +72,12 @@ setup-python:
|
|||||||
if [ "$(uname -m)" = "arm64" ] && [ "$(uname)" = "Darwin" ]; then
|
if [ "$(uname -m)" = "arm64" ] && [ "$(uname)" = "Darwin" ]; then
|
||||||
echo "Detected Apple Silicon — installing MLX dependencies..."
|
echo "Detected Apple Silicon — installing MLX dependencies..."
|
||||||
{{ pip }} install -r {{ backend_dir }}/requirements-mlx.txt
|
{{ pip }} install -r {{ backend_dir }}/requirements-mlx.txt
|
||||||
|
# mlx-lm and mlx-audio declare transformers>=5.x, which conflicts with
|
||||||
|
# our transformers<=4.57.x cap, so install them --no-deps (their other
|
||||||
|
# runtime deps are covered by requirements.txt / requirements-mlx.txt —
|
||||||
|
# see the note in requirements-mlx.txt and .github/workflows/release.yml)
|
||||||
|
{{ pip }} install --no-deps mlx-lm==0.31.1
|
||||||
|
{{ pip }} install --no-deps mlx-audio==0.4.1
|
||||||
fi
|
fi
|
||||||
{{ pip }} install git+https://github.com/QwenLM/Qwen3-TTS.git
|
{{ pip }} install git+https://github.com/QwenLM/Qwen3-TTS.git
|
||||||
{{ pip }} install pyinstaller ruff pytest pytest-asyncio -q
|
{{ pip }} install pyinstaller ruff pytest pytest-asyncio -q
|
||||||
@@ -89,10 +95,10 @@ setup-python:
|
|||||||
}
|
}
|
||||||
Write-Host "Installing Python dependencies..."
|
Write-Host "Installing Python dependencies..."
|
||||||
& "{{ python }}" -m pip install --upgrade pip -q
|
& "{{ python }}" -m pip install --upgrade pip -q
|
||||||
$gpus = Get-CimInstance Win32_VideoController | Select-Object -ExpandProperty Name
|
$gpus = Get-CimInstance Win32_VideoController | Select-Object -ExpandProperty Name; \
|
||||||
Write-Host "Detected GPUs: $($gpus -join ', ')"
|
Write-Host "Detected GPUs: $($gpus -join ', ')"; \
|
||||||
$hasNvidia = ($gpus | Where-Object { $_ -match 'NVIDIA' }).Count -gt 0
|
$hasNvidia = ($gpus | Where-Object { $_ -match 'NVIDIA' }).Count -gt 0; \
|
||||||
$hasIntelArc = ($gpus | Where-Object { $_ -match 'Arc' }).Count -gt 0
|
$hasIntelArc = ($gpus | Where-Object { $_ -match 'Arc' }).Count -gt 0; \
|
||||||
if ($hasNvidia) { \
|
if ($hasNvidia) { \
|
||||||
Write-Host "NVIDIA GPU detected — installing PyTorch with CUDA support..."; \
|
Write-Host "NVIDIA GPU detected — installing PyTorch with CUDA support..."; \
|
||||||
& "{{ pip }}" install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu128; \
|
& "{{ pip }}" install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu128; \
|
||||||
@@ -226,12 +232,16 @@ build-server: _ensure-venv
|
|||||||
build-server: _ensure-venv
|
build-server: _ensure-venv
|
||||||
$ErrorActionPreference = "Stop"; \
|
$ErrorActionPreference = "Stop"; \
|
||||||
$env:PATH = "{{ venv_bin }};$env:PATH"; \
|
$env:PATH = "{{ venv_bin }};$env:PATH"; \
|
||||||
& "{{ python }}" backend/build_binary.py; \
|
|
||||||
if ($LASTEXITCODE -ne 0) { throw "build_binary.py failed with exit code $LASTEXITCODE" }; \
|
|
||||||
$triple = (rustc --print host-tuple); \
|
$triple = (rustc --print host-tuple); \
|
||||||
New-Item -ItemType Directory -Path "{{ tauri_dir }}/src-tauri/binaries" -Force | Out-Null; \
|
New-Item -ItemType Directory -Path "{{ tauri_dir }}/src-tauri/binaries" -Force | Out-Null; \
|
||||||
|
& "{{ python }}" backend/build_binary.py; \
|
||||||
|
if ($LASTEXITCODE -ne 0) { throw "build_binary.py failed with exit code $LASTEXITCODE" }; \
|
||||||
Copy-Item "backend/dist/voicebox-server.exe" "{{ tauri_dir }}/src-tauri/binaries/voicebox-server-$triple.exe" -Force; \
|
Copy-Item "backend/dist/voicebox-server.exe" "{{ tauri_dir }}/src-tauri/binaries/voicebox-server-$triple.exe" -Force; \
|
||||||
Write-Host "Copied sidecar: voicebox-server-$triple.exe"
|
Write-Host "Copied sidecar: voicebox-server-$triple.exe"; \
|
||||||
|
& "{{ python }}" backend/build_binary.py --shim; \
|
||||||
|
if ($LASTEXITCODE -ne 0) { throw "build_binary.py --shim failed with exit code $LASTEXITCODE" }; \
|
||||||
|
Copy-Item "backend/dist/voicebox-mcp.exe" "{{ tauri_dir }}/src-tauri/binaries/voicebox-mcp-$triple.exe" -Force; \
|
||||||
|
Write-Host "Copied sidecar: voicebox-mcp-$triple.exe"
|
||||||
|
|
||||||
# Build CUDA server binary and place in app data dir for local testing
|
# Build CUDA server binary and place in app data dir for local testing
|
||||||
[windows]
|
[windows]
|
||||||
|
|||||||
@@ -66,13 +66,46 @@ fn find_monitor_source_via_pactl() -> Option<String> {
|
|||||||
None
|
None
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Select the capture device: prefer an exact match against the monitor
|
||||||
|
/// source name reported by `pactl`, then fall back to any device whose name
|
||||||
|
/// contains "monitor", then the host's default input device.
|
||||||
|
fn select_capture_device(host: &cpal::Host, monitor_source: Option<&str>) -> Option<cpal::Device> {
|
||||||
|
let devices: Vec<cpal::Device> = host.input_devices().ok()?.collect();
|
||||||
|
|
||||||
|
if let Some(target) = monitor_source {
|
||||||
|
if let Some(pos) = devices
|
||||||
|
.iter()
|
||||||
|
.position(|d| d.name().map(|n| n == target).unwrap_or(false))
|
||||||
|
{
|
||||||
|
eprintln!(
|
||||||
|
"Linux audio capture: Using pactl monitor device: {}",
|
||||||
|
target
|
||||||
|
);
|
||||||
|
return devices.into_iter().nth(pos);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if let Some(pos) = devices.iter().position(|d| {
|
||||||
|
d.name()
|
||||||
|
.map(|n| n.to_lowercase().contains("monitor"))
|
||||||
|
.unwrap_or(false)
|
||||||
|
}) {
|
||||||
|
let name = devices[pos].name().unwrap_or_default();
|
||||||
|
eprintln!("Linux audio capture: Found monitor device by name: {}", name);
|
||||||
|
return devices.into_iter().nth(pos);
|
||||||
|
}
|
||||||
|
|
||||||
|
eprintln!("Linux audio capture: No monitor device found, falling back to default input");
|
||||||
|
host.default_input_device()
|
||||||
|
}
|
||||||
|
|
||||||
/// Start capturing system audio on Linux using PulseAudio monitor sources.
|
/// Start capturing system audio on Linux using PulseAudio monitor sources.
|
||||||
///
|
///
|
||||||
/// On modern Linux with PulseAudio or PipeWire, we first try to detect the
|
/// On modern Linux with PulseAudio or PipeWire, we first try to detect the
|
||||||
/// monitor source via `pactl` and set the `PULSE_SOURCE` environment variable.
|
/// monitor source via `pactl`, then select the matching cpal input device by
|
||||||
/// This tells PulseAudio's ALSA plugin to use the monitor as the default input
|
/// name. This avoids mutating the process environment (`PULSE_SOURCE`), which
|
||||||
/// source for this process. If `pactl` is unavailable, we fall back to searching
|
/// is not thread-safe and would affect every thread in the process. If `pactl`
|
||||||
/// cpal device names for "monitor".
|
/// is unavailable, we fall back to searching cpal device names for "monitor".
|
||||||
pub async fn start_capture(
|
pub async fn start_capture(
|
||||||
state: &AudioCaptureState,
|
state: &AudioCaptureState,
|
||||||
max_duration_secs: u32,
|
max_duration_secs: u32,
|
||||||
@@ -101,73 +134,16 @@ pub async fn start_capture(
|
|||||||
|
|
||||||
// Spawn capture on a dedicated thread
|
// Spawn capture on a dedicated thread
|
||||||
thread::spawn(move || {
|
thread::spawn(move || {
|
||||||
// Try to set PULSE_SOURCE to a monitor before initializing cpal.
|
|
||||||
// This tells PulseAudio/PipeWire's ALSA plugin to use the monitor
|
|
||||||
// as the default input source for this process.
|
|
||||||
let monitor_source = find_monitor_source_via_pactl();
|
|
||||||
if let Some(ref source_name) = monitor_source {
|
|
||||||
eprintln!(
|
|
||||||
"Linux audio capture: Setting PULSE_SOURCE={}",
|
|
||||||
source_name
|
|
||||||
);
|
|
||||||
std::env::set_var("PULSE_SOURCE", source_name);
|
|
||||||
}
|
|
||||||
|
|
||||||
let host = cpal::default_host();
|
let host = cpal::default_host();
|
||||||
|
let monitor_source = find_monitor_source_via_pactl();
|
||||||
|
|
||||||
// Select the capture device.
|
let device = match select_capture_device(&host, monitor_source.as_deref()) {
|
||||||
// If PULSE_SOURCE was set, the default input device IS the monitor.
|
Some(d) => d,
|
||||||
// Otherwise, fall back to searching device names for "monitor".
|
None => {
|
||||||
let device = if monitor_source.is_some() {
|
let error_msg = "No audio input device available".to_string();
|
||||||
// PULSE_SOURCE was set — default input IS the monitor now
|
eprintln!("{}", error_msg);
|
||||||
match host.default_input_device() {
|
*error_arc.lock().unwrap() = Some(error_msg);
|
||||||
Some(d) => {
|
return;
|
||||||
let name = d.name().unwrap_or_default();
|
|
||||||
eprintln!(
|
|
||||||
"Linux audio capture: Using PULSE_SOURCE monitor device: {}",
|
|
||||||
name
|
|
||||||
);
|
|
||||||
d
|
|
||||||
}
|
|
||||||
None => {
|
|
||||||
let error_msg = "No audio input device available".to_string();
|
|
||||||
eprintln!("{}", error_msg);
|
|
||||||
*error_arc.lock().unwrap() = Some(error_msg);
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
// pactl not available — try to find monitor by name (original approach)
|
|
||||||
let mut monitor_device = None;
|
|
||||||
if let Ok(devices) = host.input_devices() {
|
|
||||||
for d in devices {
|
|
||||||
if let Ok(name) = d.name() {
|
|
||||||
let name_lower = name.to_lowercase();
|
|
||||||
if name_lower.contains("monitor") {
|
|
||||||
eprintln!(
|
|
||||||
"Linux audio capture: Found monitor device by name: {}",
|
|
||||||
name
|
|
||||||
);
|
|
||||||
monitor_device = Some(d);
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
match monitor_device {
|
|
||||||
Some(d) => d,
|
|
||||||
None => {
|
|
||||||
eprintln!("Linux audio capture: No monitor device found, falling back to default input");
|
|
||||||
match host.default_input_device() {
|
|
||||||
Some(d) => d,
|
|
||||||
None => {
|
|
||||||
let error_msg = "No audio input device available".to_string();
|
|
||||||
eprintln!("{}", error_msg);
|
|
||||||
*error_arc.lock().unwrap() = Some(error_msg);
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ pub fn key_from_str(name: &str) -> Option<Key> {
|
|||||||
"ShiftLeft" => Key::ShiftLeft,
|
"ShiftLeft" => Key::ShiftLeft,
|
||||||
"ShiftRight" => Key::ShiftRight,
|
"ShiftRight" => Key::ShiftRight,
|
||||||
"CapsLock" => Key::CapsLock,
|
"CapsLock" => Key::CapsLock,
|
||||||
|
"Function" => Key::Function,
|
||||||
|
|
||||||
// Whitespace / navigation
|
// Whitespace / navigation
|
||||||
"Space" => Key::Space,
|
"Space" => Key::Space,
|
||||||
|
|||||||
@@ -19,19 +19,23 @@
|
|||||||
//! regardless of the active layout — most Windows apps treat that as
|
//! regardless of the active layout — most Windows apps treat that as
|
||||||
//! Ctrl+V. AutoHotkey relies on the same behaviour.
|
//! Ctrl+V. AutoHotkey relies on the same behaviour.
|
||||||
|
|
||||||
|
#[cfg(target_os = "macos")]
|
||||||
use std::sync::atomic::{AtomicU16, Ordering};
|
use std::sync::atomic::{AtomicU16, Ordering};
|
||||||
|
|
||||||
/// `kVK_ANSI_V` — the keycode for the physical V key on a US QWERTY
|
/// `kVK_ANSI_V` — the keycode for the physical V key on a US QWERTY
|
||||||
/// layout. Used as the fallback whenever live resolution can't produce a
|
/// layout. Used as the fallback whenever live resolution can't produce a
|
||||||
/// better answer (no Unicode key layout data, lookup failure, non-macOS).
|
/// better answer (no Unicode key layout data, lookup failure, non-macOS).
|
||||||
|
#[cfg(target_os = "macos")]
|
||||||
const FALLBACK_V_KEYCODE: u16 = 9;
|
const FALLBACK_V_KEYCODE: u16 = 9;
|
||||||
|
|
||||||
|
#[cfg(target_os = "macos")]
|
||||||
static V_KEYCODE: AtomicU16 = AtomicU16::new(FALLBACK_V_KEYCODE);
|
static V_KEYCODE: AtomicU16 = AtomicU16::new(FALLBACK_V_KEYCODE);
|
||||||
|
|
||||||
/// Returns the keycode whose current-layout translation is `'v'`. Falls
|
/// Returns the keycode whose current-layout translation is `'v'`. Falls
|
||||||
/// back to `kVK_ANSI_V` when resolution hasn't run, the active input
|
/// back to `kVK_ANSI_V` when resolution hasn't run, the active input
|
||||||
/// source carries no Unicode key layout data, or no keycode in the layout
|
/// source carries no Unicode key layout data, or no keycode in the layout
|
||||||
/// produces `v`.
|
/// produces `v`.
|
||||||
|
#[cfg(target_os = "macos")]
|
||||||
pub fn paste_keycode_v() -> u16 {
|
pub fn paste_keycode_v() -> u16 {
|
||||||
V_KEYCODE.load(Ordering::Relaxed)
|
V_KEYCODE.load(Ordering::Relaxed)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,10 +4,14 @@
|
|||||||
//! pipeline so the focused app performs its native paste action against
|
//! pipeline so the focused app performs its native paste action against
|
||||||
//! whatever the clipboard module has just staged.
|
//! whatever the clipboard module has just staged.
|
||||||
//!
|
//!
|
||||||
//! - **macOS** — Cmd down, V down with Cmd flag, V up with Cmd flag, Cmd
|
//! - **macOS** — Cmd down with Cmd flag, V down with Cmd flag, V up with
|
||||||
//! up via `CGEventPost` at `kCGHIDEventTap`. Accessibility permission is
|
//! Cmd flag, Cmd up via `CGEventPost` at `kCGHIDEventTap`. The Cmd-down
|
||||||
//! load-bearing: without it the system swallows the events silently, so
|
//! event carries the Command flag so its `flagsChanged` representation
|
||||||
//! callers must gate on [`crate::accessibility::is_trusted`].
|
//! matches hardware — Electron/Chromium tracks modifier state from that
|
||||||
|
//! flag and drops the paste otherwise (see the note on the event table).
|
||||||
|
//! Accessibility permission is load-bearing: without it the system
|
||||||
|
//! swallows the events silently, so callers must gate on
|
||||||
|
//! [`crate::accessibility::is_trusted`].
|
||||||
//! - **Windows** — Ctrl down, V down, V up, Ctrl up via `SendInput`. No
|
//! - **Windows** — Ctrl down, V down, V up, Ctrl up via `SendInput`. No
|
||||||
//! permission gate, but UAC/UIPI blocks delivery into elevated target
|
//! permission gate, but UAC/UIPI blocks delivery into elevated target
|
||||||
//! windows when we run non-elevated — nothing we can do short of also
|
//! windows when we run non-elevated — nothing we can do short of also
|
||||||
@@ -101,7 +105,18 @@ pub fn send_paste() -> Result<(), String> {
|
|||||||
let _source_guard = scopeguard::guard(source, |s| CFRelease(s as *const c_void));
|
let _source_guard = scopeguard::guard(source, |s| CFRelease(s as *const c_void));
|
||||||
|
|
||||||
let events = [
|
let events = [
|
||||||
(KEYCODE_LEFT_CMD, true, 0),
|
// The Cmd-down event must carry the Command flag itself. On real
|
||||||
|
// hardware the Cmd keyDown is a flagsChanged event whose flags
|
||||||
|
// already include Command; Chromium/Electron builds its tracked
|
||||||
|
// modifier state from that flag. Posting Cmd-down with flags = 0
|
||||||
|
// leaves that tracker showing "Command up", so the following V —
|
||||||
|
// even though its own flags carry Command — matches neither the
|
||||||
|
// Cmd+V accelerator (tracker says no modifier) nor plain-text
|
||||||
|
// insertion (event flags say Command), and Electron drops it
|
||||||
|
// silently. AppKit reads the V event's own flags and pastes
|
||||||
|
// regardless, which is why native apps worked but Electron
|
||||||
|
// targets (Slack, VS Code) silently no-op'd.
|
||||||
|
(KEYCODE_LEFT_CMD, true, K_CG_EVENT_FLAG_MASK_COMMAND),
|
||||||
(v_keycode, true, K_CG_EVENT_FLAG_MASK_COMMAND),
|
(v_keycode, true, K_CG_EVENT_FLAG_MASK_COMMAND),
|
||||||
(v_keycode, false, K_CG_EVENT_FLAG_MASK_COMMAND),
|
(v_keycode, false, K_CG_EVENT_FLAG_MASK_COMMAND),
|
||||||
(KEYCODE_LEFT_CMD, false, 0),
|
(KEYCODE_LEFT_CMD, false, 0),
|
||||||
|
|||||||
Reference in New Issue
Block a user