chore(backend): repair test suite and bring ruff to green

The suite hadn't run green since the routes refactor:
- test_profile_duplicate_names.py imported the pre-refactor module
  layout and broke collection; now imports backend.services.profiles
- tests/conftest.py puts the repo root and backend dir on sys.path so
  files collect standalone instead of depending on run order
- test_cors.py tested a hand-copied mirror of the origin list that had
  drifted from app.py (missing http://tauri.localhost); it now builds
  the app via the real create_app() factory
- test_progress.py simulated a 1KB download, below the tracker's 1MB
  reporting threshold; simulation raised to 5MB
- slow/timeout markers registered in pyproject

Ruff: ~900 violations auto-fixed (typing modernization, import
sorting, unused imports, whitespace). The remaining rules are baselined
in pyproject.toml with per-rule counts to burn down, plus per-file
carve-outs for deliberate env-before-import ordering. ruff check is
now clean; suite is 134 passed, 2 skipped.
This commit is contained in:
Jamie Pine
2026-07-26 23:16:09 -07:00
parent 766c51a8a1
commit b434db22f6
82 changed files with 970 additions and 999 deletions
+10 -9
View File
@@ -95,18 +95,19 @@ if not os.environ.get("HSA_OVERRIDE_GFX_VERSION"):
if not os.environ.get("MIOPEN_LOG_LEVEL"): if not os.environ.get("MIOPEN_LOG_LEVEL"):
os.environ["MIOPEN_LOG_LEVEL"] = "4" os.environ["MIOPEN_LOG_LEVEL"] = "4"
from urllib.parse import quote
import torch import torch
from fastapi import FastAPI from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from urllib.parse import quote
from . import __version__, config, database from . import __version__, config, database
from .services import tts, transcribe, llm
from .database import get_db from .database import get_db
from .routes import register_routers
from .services import llm, transcribe, tts
from .services.task_queue import create_background_task, init_queue
from .utils.platform_detect import get_backend_type from .utils.platform_detect import get_backend_type
from .utils.progress import get_progress_manager from .utils.progress import get_progress_manager
from .services.task_queue import create_background_task, init_queue
from .routes import register_routers
def safe_content_disposition(disposition_type: str, filename: str) -> str: def safe_content_disposition(disposition_type: str, filename: str) -> str:
@@ -122,8 +123,8 @@ def safe_content_disposition(disposition_type: str, filename: str) -> str:
def create_app() -> FastAPI: def create_app() -> FastAPI:
"""Create and configure the FastAPI application.""" """Create and configure the FastAPI application."""
from .mcp_server.server import build_mcp_server, compose_lifespan
from .mcp_server.context import ClientIdMiddleware from .mcp_server.context import ClientIdMiddleware
from .mcp_server.server import build_mcp_server, compose_lifespan
# Build the MCP app up-front so we can wire its lifespan into FastAPI's — # Build the MCP app up-front so we can wire its lifespan into FastAPI's —
# FastMCP's Streamable HTTP transport only works if its session manager # FastMCP's Streamable HTTP transport only works if its session manager
@@ -202,8 +203,8 @@ def _mount_frontend(application: FastAPI) -> None:
if not frontend_dir.is_dir(): if not frontend_dir.is_dir():
return return
from fastapi.staticfiles import StaticFiles
from fastapi.responses import FileResponse from fastapi.responses import FileResponse
from fastapi.staticfiles import StaticFiles
# Mount hashed assets (JS, CSS, images) that Vite places under /assets # Mount hashed assets (JS, CSS, images) that Vite places under /assets
assets_dir = frontend_dir / "assets" assets_dir = frontend_dir / "assets"
@@ -243,9 +244,9 @@ def _get_gpu_status() -> str:
if not compatible: if not compatible:
label += " [UNSUPPORTED - see logs]" label += " [UNSUPPORTED - see logs]"
return label return label
elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
return "MPS (Apple Silicon)" return "MPS (Apple Silicon)"
elif backend_type == "mlx": if backend_type == "mlx":
return "Metal (Apple Silicon via MLX)" return "Metal (Apple Silicon via MLX)"
# Intel XPU (Arc / Data Center) via IPEX # Intel XPU (Arc / Data Center) via IPEX
@@ -302,7 +303,7 @@ async def _run_startup(application: FastAPI) -> None:
if result.rowcount > 0: if result.rowcount > 0:
logger.info("Marked %d stale generation(s) as failed", result.rowcount) logger.info("Marked %d stale generation(s) as failed", result.rowcount)
from .database import VoiceProfile as DBVoiceProfile, Generation as DBGeneration from .database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile
profile_count = db.query(DBVoiceProfile).count() profile_count = db.query(DBVoiceProfile).count()
generation_count = db.query(DBGeneration).count() generation_count = db.query(DBGeneration).count()
+8 -9
View File
@@ -9,13 +9,12 @@ import logging
import platform import platform
from contextlib import contextmanager from contextlib import contextmanager
from pathlib import Path from pathlib import Path
from typing import Callable, List, Optional, Tuple
import numpy as np import numpy as np
from ..utils.audio import normalize_audio, load_audio from ..utils.audio import load_audio, normalize_audio
from ..utils.progress import get_progress_manager
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
from ..utils.progress import get_progress_manager
from ..utils.tasks import get_task_manager from ..utils.tasks import get_task_manager
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -25,7 +24,7 @@ def is_model_cached(
hf_repo: str, hf_repo: str,
*, *,
weight_extensions: tuple[str, ...] = (".safetensors", ".bin"), weight_extensions: tuple[str, ...] = (".safetensors", ".bin"),
required_files: Optional[list[str]] = None, required_files: list[str] | None = None,
) -> bool: ) -> bool:
""" """
Check if a HuggingFace model is fully cached locally. Check if a HuggingFace model is fully cached locally.
@@ -201,11 +200,11 @@ def manual_seed(seed: int, device: str) -> None:
async def combine_voice_prompts( async def combine_voice_prompts(
audio_paths: List[str], audio_paths: list[str],
reference_texts: List[str], reference_texts: list[str],
*, *,
sample_rate: Optional[int] = None, sample_rate: int | None = None,
) -> Tuple[np.ndarray, str]: ) -> tuple[np.ndarray, str]:
""" """
Combine multiple reference audio samples into one. Combine multiple reference audio samples into one.
@@ -235,7 +234,7 @@ async def combine_voice_prompts(
def model_load_progress( def model_load_progress(
model_name: str, model_name: str,
is_cached: bool, is_cached: bool,
filter_non_downloads: Optional[bool] = None, filter_non_downloads: bool | None = None,
): ):
""" """
Context manager for model loading with HF download progress tracking. Context manager for model loading with HF download progress tracking.
+12 -13
View File
@@ -10,17 +10,16 @@ import asyncio
import logging import logging
import threading import threading
from pathlib import Path from pathlib import Path
from typing import ClassVar, List, Optional, Tuple from typing import ClassVar
import numpy as np import numpy as np
from . import TTSBackend
from .base import ( from .base import (
is_model_cached,
get_torch_device,
empty_device_cache,
manual_seed,
combine_voice_prompts as _combine_voice_prompts, combine_voice_prompts as _combine_voice_prompts,
empty_device_cache,
get_torch_device,
is_model_cached,
manual_seed,
model_load_progress, model_load_progress,
patch_chatterbox_f32, patch_chatterbox_f32,
) )
@@ -127,7 +126,7 @@ class ChatterboxTTSBackend:
audio_path: str, audio_path: str,
reference_text: str, reference_text: str,
use_cache: bool = True, use_cache: bool = True,
) -> Tuple[dict, bool]: ) -> tuple[dict, bool]:
""" """
Create voice prompt from reference audio. Create voice prompt from reference audio.
@@ -143,9 +142,9 @@ class ChatterboxTTSBackend:
async def combine_voice_prompts( async def combine_voice_prompts(
self, self,
audio_paths: List[str], audio_paths: list[str],
reference_texts: List[str], reference_texts: list[str],
) -> Tuple[np.ndarray, str]: ) -> tuple[np.ndarray, str]:
return await _combine_voice_prompts(audio_paths, reference_texts) return await _combine_voice_prompts(audio_paths, reference_texts)
# Per-language generation defaults. Lower temp + higher cfg = clearer speech. # Per-language generation defaults. Lower temp + higher cfg = clearer speech.
@@ -169,9 +168,9 @@ class ChatterboxTTSBackend:
text: str, text: str,
voice_prompt: dict, voice_prompt: dict,
language: str = "en", language: str = "en",
seed: Optional[int] = None, seed: int | None = None,
instruct: Optional[str] = None, instruct: str | None = None,
) -> Tuple[np.ndarray, int]: ) -> tuple[np.ndarray, int]:
""" """
Generate audio using Chatterbox Multilingual TTS. Generate audio using Chatterbox Multilingual TTS.
+13 -14
View File
@@ -10,17 +10,16 @@ import asyncio
import logging import logging
import threading import threading
from pathlib import Path from pathlib import Path
from typing import ClassVar, List, Optional, Tuple from typing import ClassVar
import numpy as np import numpy as np
from . import TTSBackend
from .base import ( from .base import (
is_model_cached,
get_torch_device,
empty_device_cache,
manual_seed,
combine_voice_prompts as _combine_voice_prompts, combine_voice_prompts as _combine_voice_prompts,
empty_device_cache,
get_torch_device,
is_model_cached,
manual_seed,
model_load_progress, model_load_progress,
patch_chatterbox_f32, patch_chatterbox_f32,
) )
@@ -81,8 +80,8 @@ class ChatterboxTurboTTSBackend:
logger.info(f"Loading Chatterbox Turbo TTS on {device}...") logger.info(f"Loading Chatterbox Turbo TTS on {device}...")
import torch import torch
from huggingface_hub import snapshot_download
from chatterbox.tts_turbo import ChatterboxTurboTTS from chatterbox.tts_turbo import ChatterboxTurboTTS
from huggingface_hub import snapshot_download
local_path = snapshot_download( local_path = snapshot_download(
repo_id=CHATTERBOX_TURBO_HF_REPO, repo_id=CHATTERBOX_TURBO_HF_REPO,
@@ -126,7 +125,7 @@ class ChatterboxTurboTTSBackend:
audio_path: str, audio_path: str,
reference_text: str, reference_text: str,
use_cache: bool = True, use_cache: bool = True,
) -> Tuple[dict, bool]: ) -> tuple[dict, bool]:
""" """
Create voice prompt from reference audio. Create voice prompt from reference audio.
@@ -141,9 +140,9 @@ class ChatterboxTurboTTSBackend:
async def combine_voice_prompts( async def combine_voice_prompts(
self, self,
audio_paths: List[str], audio_paths: list[str],
reference_texts: List[str], reference_texts: list[str],
) -> Tuple[np.ndarray, str]: ) -> tuple[np.ndarray, str]:
return await _combine_voice_prompts(audio_paths, reference_texts) return await _combine_voice_prompts(audio_paths, reference_texts)
async def generate( async def generate(
@@ -151,9 +150,9 @@ class ChatterboxTurboTTSBackend:
text: str, text: str,
voice_prompt: dict, voice_prompt: dict,
language: str = "en", language: str = "en",
seed: Optional[int] = None, seed: int | None = None,
instruct: Optional[str] = None, instruct: str | None = None,
) -> Tuple[np.ndarray, int]: ) -> tuple[np.ndarray, int]:
""" """
Generate audio using Chatterbox Turbo TTS. Generate audio using Chatterbox Turbo TTS.
+16 -19
View File
@@ -16,20 +16,19 @@ causal LM generates speech via flow-matching diffusion.
import asyncio import asyncio
import logging import logging
import threading import threading
from typing import ClassVar, List, Optional, Tuple from typing import ClassVar
import numpy as np import numpy as np
from . import TTSBackend from ..utils.cache import cache_voice_prompt, get_cache_key, get_cached_voice_prompt
from .base import ( from .base import (
is_model_cached,
get_torch_device,
empty_device_cache,
manual_seed,
combine_voice_prompts as _combine_voice_prompts, combine_voice_prompts as _combine_voice_prompts,
empty_device_cache,
get_torch_device,
is_model_cached,
manual_seed,
model_load_progress, model_load_progress,
) )
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -182,7 +181,7 @@ class HumeTadaBackend:
# getattr(config, "tokenizer_name", "meta-llama/Llama-3.2-1B") # getattr(config, "tokenizer_name", "meta-llama/Llama-3.2-1B")
# which hits the gated repo. Pre-load the config from HF, # which hits the gated repo. Pre-load the config from HF,
# inject the local tokenizer path, then pass it in. # inject the local tokenizer path, then pass it in.
from tada.modules.tada import TadaForCausalLM, TadaConfig from tada.modules.tada import TadaConfig, TadaForCausalLM
logger.info(f"Loading TADA {model_size} model...") logger.info(f"Loading TADA {model_size} model...")
config = TadaConfig.from_pretrained(repo) config = TadaConfig.from_pretrained(repo)
@@ -214,7 +213,7 @@ class HumeTadaBackend:
audio_path: str, audio_path: str,
reference_text: str, reference_text: str,
use_cache: bool = True, use_cache: bool = True,
) -> Tuple[dict, bool]: ) -> tuple[dict, bool]:
""" """
Create voice prompt from reference audio using TADA's encoder. Create voice prompt from reference audio using TADA's encoder.
@@ -234,8 +233,8 @@ class HumeTadaBackend:
return cached, True return cached, True
def _encode_sync(): def _encode_sync():
import torch
import soundfile as sf import soundfile as sf
import torch
device = self._device device = self._device
@@ -258,9 +257,7 @@ class HumeTadaBackend:
val = getattr(prompt, field_name) val = getattr(prompt, field_name)
if isinstance(val, torch.Tensor): if isinstance(val, torch.Tensor):
prompt_dict[field_name] = val.detach().cpu() prompt_dict[field_name] = val.detach().cpu()
elif isinstance(val, list): elif isinstance(val, (list, int, float)):
prompt_dict[field_name] = val
elif isinstance(val, (int, float)):
prompt_dict[field_name] = val prompt_dict[field_name] = val
else: else:
prompt_dict[field_name] = val prompt_dict[field_name] = val
@@ -275,9 +272,9 @@ class HumeTadaBackend:
async def combine_voice_prompts( async def combine_voice_prompts(
self, self,
audio_paths: List[str], audio_paths: list[str],
reference_texts: List[str], reference_texts: list[str],
) -> Tuple[np.ndarray, str]: ) -> tuple[np.ndarray, str]:
return await _combine_voice_prompts(audio_paths, reference_texts, sample_rate=24000) return await _combine_voice_prompts(audio_paths, reference_texts, sample_rate=24000)
async def generate( async def generate(
@@ -285,9 +282,9 @@ class HumeTadaBackend:
text: str, text: str,
voice_prompt: dict, voice_prompt: dict,
language: str = "en", language: str = "en",
seed: Optional[int] = None, seed: int | None = None,
instruct: Optional[str] = None, instruct: str | None = None,
) -> Tuple[np.ndarray, int]: ) -> tuple[np.ndarray, int]:
""" """
Generate audio from text using HumeAI TADA. Generate audio from text using HumeAI TADA.
+22 -22
View File
@@ -2,25 +2,25 @@
PyTorch backend implementation for TTS and STT. PyTorch backend implementation for TTS and STT.
""" """
from typing import Optional, List, Tuple
import asyncio import asyncio
import logging import logging
import torch
import numpy as np import numpy as np
import torch
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS from ..utils.audio import load_audio
from ..utils.cache import cache_voice_prompt, get_cache_key, get_cached_voice_prompt
from . import LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS
from .base import ( from .base import (
is_model_cached,
get_torch_device,
empty_device_cache,
manual_seed,
combine_voice_prompts as _combine_voice_prompts, combine_voice_prompts as _combine_voice_prompts,
empty_device_cache,
get_torch_device,
is_model_cached,
manual_seed,
model_load_progress, model_load_progress,
) )
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
from ..utils.audio import load_audio
class PyTorchTTSBackend: class PyTorchTTSBackend:
@@ -63,7 +63,7 @@ class PyTorchTTSBackend:
def _is_model_cached(self, model_size: str) -> bool: def _is_model_cached(self, model_size: str) -> bool:
return is_model_cached(self._get_model_path(model_size)) return is_model_cached(self._get_model_path(model_size))
async def load_model_async(self, model_size: Optional[str] = None): async def load_model_async(self, model_size: str | None = None):
""" """
Lazy load the TTS model with automatic downloading from HuggingFace Hub. Lazy load the TTS model with automatic downloading from HuggingFace Hub.
@@ -140,7 +140,7 @@ class PyTorchTTSBackend:
audio_path: str, audio_path: str,
reference_text: str, reference_text: str,
use_cache: bool = True, use_cache: bool = True,
) -> Tuple[dict, bool]: ) -> tuple[dict, bool]:
""" """
Create voice prompt from reference audio. Create voice prompt from reference audio.
@@ -165,7 +165,7 @@ class PyTorchTTSBackend:
# For PyTorch backend, the dict should contain tensors, not file paths # For PyTorch backend, the dict should contain tensors, not file paths
# So we can safely return it # So we can safely return it
return cached_prompt, True return cached_prompt, True
elif isinstance(cached_prompt, torch.Tensor): if isinstance(cached_prompt, torch.Tensor):
# Legacy cache format - convert to dict # Legacy cache format - convert to dict
# This shouldn't happen in practice, but handle it # This shouldn't happen in practice, but handle it
return {"prompt": cached_prompt}, True return {"prompt": cached_prompt}, True
@@ -194,9 +194,9 @@ class PyTorchTTSBackend:
async def combine_voice_prompts( async def combine_voice_prompts(
self, self,
audio_paths: List[str], audio_paths: list[str],
reference_texts: List[str], reference_texts: list[str],
) -> Tuple[np.ndarray, str]: ) -> tuple[np.ndarray, str]:
return await _combine_voice_prompts(audio_paths, reference_texts) return await _combine_voice_prompts(audio_paths, reference_texts)
async def generate( async def generate(
@@ -204,9 +204,9 @@ class PyTorchTTSBackend:
text: str, text: str,
voice_prompt: dict, voice_prompt: dict,
language: str = "en", language: str = "en",
seed: Optional[int] = None, seed: int | None = None,
instruct: Optional[str] = None, instruct: str | None = None,
) -> Tuple[np.ndarray, int]: ) -> tuple[np.ndarray, int]:
""" """
Generate audio from text using voice prompt. Generate audio from text using voice prompt.
@@ -266,7 +266,7 @@ class PyTorchSTTBackend:
hf_repo = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}") hf_repo = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
return is_model_cached(hf_repo) return is_model_cached(hf_repo)
async def load_model_async(self, model_size: Optional[str] = None): async def load_model_async(self, model_size: str | None = None):
""" """
Lazy load the Whisper model. Lazy load the Whisper model.
@@ -290,7 +290,7 @@ class PyTorchSTTBackend:
is_cached = self._is_model_cached(model_size) is_cached = self._is_model_cached(model_size)
with model_load_progress(progress_model_name, is_cached): with model_load_progress(progress_model_name, is_cached):
from transformers import WhisperProcessor, WhisperForConditionalGeneration from transformers import WhisperForConditionalGeneration, WhisperProcessor
model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}") model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
logger.info("Loading Whisper model %s on %s...", model_size, self.device) logger.info("Loading Whisper model %s on %s...", model_size, self.device)
@@ -317,8 +317,8 @@ class PyTorchSTTBackend:
async def transcribe( async def transcribe(
self, self,
audio_path: str, audio_path: str,
language: Optional[str] = None, language: str | None = None,
model_size: Optional[str] = None, model_size: str | None = None,
) -> str: ) -> str:
""" """
Transcribe audio to text. Transcribe audio to text.
@@ -16,16 +16,15 @@ Languages supported: zh, en, ja, ko, de, fr, ru, pt, es, it
import asyncio import asyncio
import logging import logging
from typing import Optional
import numpy as np import numpy as np
import torch import torch
from . import TTSBackend, LANGUAGE_CODE_TO_NAME from . import LANGUAGE_CODE_TO_NAME
from .base import ( from .base import (
is_model_cached,
get_torch_device,
combine_voice_prompts as _combine_voice_prompts, combine_voice_prompts as _combine_voice_prompts,
get_torch_device,
is_model_cached,
model_load_progress, model_load_progress,
) )
@@ -62,7 +61,7 @@ class QwenCustomVoiceBackend:
self.model = None self.model = None
self.model_size = model_size self.model_size = model_size
self.device = self._get_device() self.device = self._get_device()
self._current_model_size: Optional[str] = None self._current_model_size: str | None = None
def _get_device(self) -> str: def _get_device(self) -> str:
return get_torch_device(allow_xpu=True, allow_directml=True) return get_torch_device(allow_xpu=True, allow_directml=True)
@@ -75,11 +74,11 @@ class QwenCustomVoiceBackend:
raise ValueError(f"Unknown model size: {model_size}") raise ValueError(f"Unknown model size: {model_size}")
return QWEN_CV_HF_REPOS[model_size] return QWEN_CV_HF_REPOS[model_size]
def _is_model_cached(self, model_size: Optional[str] = None) -> bool: def _is_model_cached(self, model_size: str | None = None) -> bool:
size = model_size or self.model_size size = model_size or self.model_size
return is_model_cached(self._get_model_path(size)) return is_model_cached(self._get_model_path(size))
async def load_model_async(self, model_size: Optional[str] = None) -> None: async def load_model_async(self, model_size: str | None = None) -> None:
if model_size is None: if model_size is None:
model_size = self.model_size model_size = self.model_size
@@ -164,8 +163,8 @@ class QwenCustomVoiceBackend:
text: str, text: str,
voice_prompt: dict, voice_prompt: dict,
language: str = "en", language: str = "en",
seed: Optional[int] = None, seed: int | None = None,
instruct: Optional[str] = None, instruct: str | None = None,
) -> tuple[np.ndarray, int]: ) -> tuple[np.ndarray, int]:
""" """
Generate audio using Qwen CustomVoice. Generate audio using Qwen CustomVoice.
+22 -24
View File
@@ -9,18 +9,16 @@ and STT engines.
import asyncio import asyncio
import logging import logging
from typing import Optional
from . import LLMBackend, DEFAULT_LLM_MAX_TOKENS, DEFAULT_LLM_TEMPERATURE from ..services.mlx_thread import clear_mlx_cache, run_on_mlx_thread
from ..utils.hf_offline_patch import force_offline_if_cached
from . import DEFAULT_LLM_MAX_TOKENS, DEFAULT_LLM_TEMPERATURE
from .base import ( from .base import (
is_model_cached,
get_torch_device,
empty_device_cache, empty_device_cache,
manual_seed, get_torch_device,
is_model_cached,
model_load_progress, model_load_progress,
) )
from ..services.mlx_thread import run_on_mlx_thread, clear_mlx_cache
from ..utils.hf_offline_patch import force_offline_if_cached
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -44,8 +42,8 @@ def _progress_name(model_size: str) -> str:
def _build_messages( def _build_messages(
prompt: str, prompt: str,
system: Optional[str], system: str | None,
examples: Optional[list[tuple[str, str]]] = None, examples: list[tuple[str, str]] | None = None,
) -> list[dict]: ) -> list[dict]:
messages: list[dict] = [] messages: list[dict] = []
if system: if system:
@@ -65,7 +63,7 @@ class PyTorchQwenLLMBackend:
self.model = None self.model = None
self.tokenizer = None self.tokenizer = None
self.model_size = model_size self.model_size = model_size
self._current_model_size: Optional[str] = None self._current_model_size: str | None = None
self.device = self._get_device() self.device = self._get_device()
def _get_device(self) -> str: def _get_device(self) -> str:
@@ -82,7 +80,7 @@ class PyTorchQwenLLMBackend:
def _is_model_cached(self, model_size: str) -> bool: def _is_model_cached(self, model_size: str) -> bool:
return is_model_cached(self._get_model_path(model_size)) return is_model_cached(self._get_model_path(model_size))
async def load_model(self, model_size: Optional[str] = None) -> None: async def load_model(self, model_size: str | None = None) -> None:
if model_size is None: if model_size is None:
model_size = self.model_size model_size = self.model_size
@@ -132,11 +130,11 @@ class PyTorchQwenLLMBackend:
async def generate( async def generate(
self, self,
prompt: str, prompt: str,
system: Optional[str] = None, system: str | None = None,
max_tokens: int = DEFAULT_LLM_MAX_TOKENS, max_tokens: int = DEFAULT_LLM_MAX_TOKENS,
temperature: float = DEFAULT_LLM_TEMPERATURE, temperature: float = DEFAULT_LLM_TEMPERATURE,
model_size: Optional[str] = None, model_size: str | None = None,
examples: Optional[list[tuple[str, str]]] = None, examples: list[tuple[str, str]] | None = None,
) -> str: ) -> str:
await self.load_model(model_size) await self.load_model(model_size)
return await asyncio.to_thread( return await asyncio.to_thread(
@@ -146,10 +144,10 @@ class PyTorchQwenLLMBackend:
def _generate_sync( def _generate_sync(
self, self,
prompt: str, prompt: str,
system: Optional[str], system: str | None,
max_tokens: int, max_tokens: int,
temperature: float, temperature: float,
examples: Optional[list[tuple[str, str]]] = None, examples: list[tuple[str, str]] | None = None,
) -> str: ) -> str:
import torch import torch
@@ -187,7 +185,7 @@ class MLXQwenLLMBackend:
self.model = None self.model = None
self.tokenizer = None self.tokenizer = None
self.model_size = model_size self.model_size = model_size
self._current_model_size: Optional[str] = None self._current_model_size: str | None = None
def is_loaded(self) -> bool: def is_loaded(self) -> bool:
return self.model is not None return self.model is not None
@@ -203,7 +201,7 @@ class MLXQwenLLMBackend:
weight_extensions=(".safetensors", ".bin", ".npz"), weight_extensions=(".safetensors", ".bin", ".npz"),
) )
def _ensure_loaded_sync(self, model_size: Optional[str]) -> None: def _ensure_loaded_sync(self, model_size: str | None) -> None:
"""Load the model if the requested size isn't already resident. """Load the model if the requested size isn't already resident.
Runs on the MLX worker thread so it stays serialized with generation. Runs on the MLX worker thread so it stays serialized with generation.
@@ -219,7 +217,7 @@ class MLXQwenLLMBackend:
self._load_model_sync(model_size) self._load_model_sync(model_size)
async def load_model(self, model_size: Optional[str] = None) -> None: async def load_model(self, model_size: str | None = None) -> None:
await run_on_mlx_thread(self._ensure_loaded_sync, model_size) await run_on_mlx_thread(self._ensure_loaded_sync, model_size)
async def unload(self) -> None: async def unload(self) -> None:
@@ -261,11 +259,11 @@ class MLXQwenLLMBackend:
async def generate( async def generate(
self, self,
prompt: str, prompt: str,
system: Optional[str] = None, system: str | None = None,
max_tokens: int = DEFAULT_LLM_MAX_TOKENS, max_tokens: int = DEFAULT_LLM_MAX_TOKENS,
temperature: float = DEFAULT_LLM_TEMPERATURE, temperature: float = DEFAULT_LLM_TEMPERATURE,
model_size: Optional[str] = None, model_size: str | None = None,
examples: Optional[list[tuple[str, str]]] = None, examples: list[tuple[str, str]] | None = None,
) -> str: ) -> str:
# Load-if-needed and inference run as one job on the MLX worker so a # Load-if-needed and inference run as one job on the MLX worker so a
# concurrent unload or different-size load can't land between them. # concurrent unload or different-size load can't land between them.
@@ -278,10 +276,10 @@ class MLXQwenLLMBackend:
def _generate_sync( def _generate_sync(
self, self,
prompt: str, prompt: str,
system: Optional[str], system: str | None,
max_tokens: int, max_tokens: int,
temperature: float, temperature: float,
examples: Optional[list[tuple[str, str]]] = None, examples: list[tuple[str, str]] | None = None,
) -> str: ) -> str:
from mlx_lm import generate as mlx_generate from mlx_lm import generate as mlx_generate
from mlx_lm.sample_utils import make_sampler from mlx_lm.sample_utils import make_sampler
+6 -6
View File
@@ -6,8 +6,8 @@ without changing any importers.
""" """
from .models import ( from .models import (
Base,
AudioChannel, AudioChannel,
Base,
Capture, Capture,
CaptureSettings, CaptureSettings,
ChannelDeviceMapping, ChannelDeviceMapping,
@@ -24,12 +24,12 @@ from .models import (
StoryItem, StoryItem,
VoiceProfile, VoiceProfile,
) )
from .session import engine, SessionLocal, _db_path, init_db, get_db from .session import SessionLocal, _db_path, engine, get_db, init_db
__all__ = [ __all__ = [
"AudioChannel",
# Models # Models
"Base", "Base",
"AudioChannel",
"Capture", "Capture",
"CaptureSettings", "CaptureSettings",
"ChannelDeviceMapping", "ChannelDeviceMapping",
@@ -42,13 +42,13 @@ __all__ = [
"ProfileChannelMapping", "ProfileChannelMapping",
"ProfileSample", "ProfileSample",
"Project", "Project",
"SessionLocal",
"Story", "Story",
"StoryItem", "StoryItem",
"VoiceProfile", "VoiceProfile",
"_db_path",
# Session # Session
"engine", "engine",
"SessionLocal",
"_db_path",
"init_db",
"get_db", "get_db",
"init_db",
] ]
+1 -1
View File
@@ -303,7 +303,7 @@ def _normalize_storage_paths(engine, tables: set[str]) -> None:
"""Normalize stored file paths to be relative to the configured data dir.""" """Normalize stored file paths to be relative to the configured data dir."""
from pathlib import Path from pathlib import Path
from ..config import get_data_dir, to_storage_path, resolve_storage_path from ..config import get_data_dir, resolve_storage_path, to_storage_path
data_dir = get_data_dir() data_dir = get_data_dir()
+2 -2
View File
@@ -1,9 +1,9 @@
"""ORM model definitions for the voicebox SQLite database.""" """ORM model definitions for the voicebox SQLite database."""
from datetime import datetime
import uuid import uuid
from datetime import datetime
from sqlalchemy import Column, String, Integer, Float, DateTime, Text, ForeignKey, Boolean, JSON from sqlalchemy import JSON, Boolean, Column, DateTime, Float, ForeignKey, Integer, String, Text
from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.ext.declarative import declarative_base
from ..utils.capture_chords import ( from ..utils.capture_chords import (
+2 -1
View File
@@ -5,10 +5,11 @@ entry point for development.
""" """
import argparse import argparse
import uvicorn import uvicorn
from .app import app # noqa: F401 -- re-export for uvicorn "backend.main:app"
from . import config, database from . import config, database
from .app import app # noqa: F401 -- re-export for uvicorn "backend.main:app"
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser(description="voicebox backend server") parser = argparse.ArgumentParser(description="voicebox backend server")
+2 -3
View File
@@ -11,14 +11,13 @@ import asyncio
import ipaddress import ipaddress
import logging import logging
from contextvars import ContextVar from contextvars import ContextVar
from datetime import datetime, timezone from datetime import UTC, datetime
from starlette.middleware.base import BaseHTTPMiddleware from starlette.middleware.base import BaseHTTPMiddleware
from starlette.requests import Request from starlette.requests import Request
from starlette.responses import Response from starlette.responses import Response
from starlette.types import ASGIApp from starlette.types import ASGIApp
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Strong refs to in-flight stamp tasks so asyncio.create_task results # Strong refs to in-flight stamp tasks so asyncio.create_task results
@@ -141,7 +140,7 @@ def _stamp_last_seen(client_id: str) -> None:
if row is None: if row is None:
row = MCPClientBinding(client_id=client_id) row = MCPClientBinding(client_id=client_id)
db.add(row) db.add(row)
row.last_seen_at = datetime.now(timezone.utc) row.last_seen_at = datetime.now(UTC)
db.commit() db.commit()
except Exception: except Exception:
logger.debug( logger.debug(
-1
View File
@@ -8,7 +8,6 @@ floating pill surfaces whenever an agent is speaking.
import asyncio import asyncio
from typing import Any from typing import Any
# Each subscriber gets its own queue. Bounded to drop oldest if a client lags. # Each subscriber gets its own queue. Bounded to drop oldest if a client lags.
_subscribers: set[asyncio.Queue[dict[str, Any]]] = set() _subscribers: set[asyncio.Queue[dict[str, Any]]] = set()
+1 -1
View File
@@ -30,7 +30,7 @@ def resolve_profile(
if client_id: if client_id:
# Per-client binding. Imported lazily so this module stays importable # Per-client binding. Imported lazily so this module stays importable
# even before the migration adds the table on first boot. # even before the migration adds the table on first boot.
from ..database.models import MCPClientBinding # noqa: WPS433 from ..database.models import MCPClientBinding
binding = ( binding = (
db.query(MCPClientBinding) db.query(MCPClientBinding)
+1 -2
View File
@@ -9,8 +9,8 @@ binary bundled with the desktop app.
from __future__ import annotations from __future__ import annotations
import logging import logging
from contextlib import AsyncExitStack, asynccontextmanager
from collections.abc import Callable from collections.abc import Callable
from contextlib import AsyncExitStack, asynccontextmanager
from fastapi import FastAPI from fastapi import FastAPI
from fastmcp import FastMCP from fastmcp import FastMCP
@@ -18,7 +18,6 @@ from fastmcp import FastMCP
from .context import ClientIdMiddleware from .context import ClientIdMiddleware
from .tools import register_tools from .tools import register_tools
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
+1 -3
View File
@@ -18,13 +18,11 @@ from fastmcp import FastMCP
from .. import models from .. import models
from ..database import get_db from ..database import get_db
from ..services import captures as captures_service from ..services import captures as captures_service, profiles as profiles_service
from ..services import profiles as profiles_service
from . import events as mcp_events from . import events as mcp_events
from .context import current_client_id, request_is_loopback from .context import current_client_id, request_is_loopback
from .resolve import resolve_profile from .resolve import resolve_profile
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Absolute-path transcribes are bounded to keep a bad client from # Absolute-path transcribes are bounded to keep a bad client from
-1
View File
@@ -23,7 +23,6 @@ from typing import Any
import httpx import httpx
CLIENT_ID_HEADER = "X-Voicebox-Client-Id" CLIENT_ID_HEADER = "X-Voicebox-Client-Id"
SESSION_HEADER = "mcp-session-id" SESSION_HEADER = "mcp-session-id"
HEALTH_TIMEOUT_S = 30.0 HEALTH_TIMEOUT_S = 30.0
+153 -153
View File
@@ -2,10 +2,10 @@
Pydantic models for request/response validation. Pydantic models for request/response validation.
""" """
from pydantic import BaseModel, Field
from typing import Optional, List
from datetime import datetime from datetime import datetime
from pydantic import BaseModel, Field
from .utils.capture_chords import ( from .utils.capture_chords import (
default_push_to_talk_chord, default_push_to_talk_chord,
default_toggle_to_talk_chord, default_toggle_to_talk_chord,
@@ -16,16 +16,16 @@ class VoiceProfileCreate(BaseModel):
"""Request model for creating a voice profile.""" """Request model for creating a voice profile."""
name: str = Field(..., min_length=1, max_length=100) name: str = Field(..., min_length=1, max_length=100)
description: Optional[str] = Field(None, max_length=500) description: str | None = Field(None, max_length=500)
language: str = Field( language: str = Field(
default="en", pattern="^(zh|en|ja|ko|de|fr|ru|pt|es|it|he|ar|da|el|fi|hi|ms|nl|no|pl|sv|sw|tr)$" default="en", pattern="^(zh|en|ja|ko|de|fr|ru|pt|es|it|he|ar|da|el|fi|hi|ms|nl|no|pl|sv|sw|tr)$"
) )
voice_type: Optional[str] = Field(default="cloned", pattern="^(cloned|preset|designed)$") voice_type: str | None = Field(default="cloned", pattern="^(cloned|preset|designed)$")
preset_engine: Optional[str] = Field(None, max_length=50) preset_engine: str | None = Field(None, max_length=50)
preset_voice_id: Optional[str] = Field(None, max_length=100) preset_voice_id: str | None = Field(None, max_length=100)
design_prompt: Optional[str] = Field(None, max_length=2000) design_prompt: str | None = Field(None, max_length=2000)
default_engine: Optional[str] = Field(None, max_length=50) default_engine: str | None = Field(None, max_length=50)
personality: Optional[str] = Field(None, max_length=2000) personality: str | None = Field(None, max_length=2000)
class VoiceProfileResponse(BaseModel): class VoiceProfileResponse(BaseModel):
@@ -33,16 +33,16 @@ class VoiceProfileResponse(BaseModel):
id: str id: str
name: str name: str
description: Optional[str] description: str | None
language: str language: str
avatar_path: Optional[str] = None avatar_path: str | None = None
effects_chain: Optional[List["EffectConfig"]] = None effects_chain: list["EffectConfig"] | None = None
voice_type: str = "cloned" voice_type: str = "cloned"
preset_engine: Optional[str] = None preset_engine: str | None = None
preset_voice_id: Optional[str] = None preset_voice_id: str | None = None
design_prompt: Optional[str] = None design_prompt: str | None = None
default_engine: Optional[str] = None default_engine: str | None = None
personality: Optional[str] = None personality: str | None = None
generation_count: int = 0 generation_count: int = 0
sample_count: int = 0 sample_count: int = 0
created_at: datetime created_at: datetime
@@ -82,10 +82,10 @@ class GenerationRequest(BaseModel):
profile_id: str profile_id: str
text: str = Field(..., min_length=1, max_length=50000) text: str = Field(..., min_length=1, max_length=50000)
language: str = Field(default="en", pattern="^(zh|en|ja|ko|de|fr|ru|pt|es|it|he|ar|da|el|fi|hi|ms|nl|no|pl|sv|sw|tr)$") language: str = Field(default="en", pattern="^(zh|en|ja|ko|de|fr|ru|pt|es|it|he|ar|da|el|fi|hi|ms|nl|no|pl|sv|sw|tr)$")
seed: Optional[int] = Field(None, ge=0) seed: int | None = Field(None, ge=0)
model_size: Optional[str] = Field(default="1.7B", pattern="^(1\\.7B|0\\.6B|1B|3B)$") model_size: str | None = Field(default="1.7B", pattern="^(1\\.7B|0\\.6B|1B|3B)$")
instruct: Optional[str] = Field(None, max_length=500) instruct: str | None = Field(None, max_length=500)
engine: Optional[str] = Field(default="qwen", pattern="^(qwen|qwen_custom_voice|luxtts|chatterbox|chatterbox_turbo|tada|kokoro)$") engine: str | None = Field(default="qwen", pattern="^(qwen|qwen_custom_voice|luxtts|chatterbox|chatterbox_turbo|tada|kokoro)$")
personality: bool = Field( personality: bool = Field(
default=False, default=False,
description="When true and the profile has a personality prompt, the input text is rewritten in-character before TTS.", description="When true and the profile has a personality prompt, the input text is rewritten in-character before TTS.",
@@ -97,7 +97,7 @@ class GenerationRequest(BaseModel):
default=50, ge=0, le=500, description="Crossfade duration in ms between chunks (0 for hard cut)" default=50, ge=0, le=500, description="Crossfade duration in ms between chunks (0 for hard cut)"
) )
normalize: bool = Field(default=True, description="Normalize output audio volume") normalize: bool = Field(default=True, description="Normalize output audio volume")
effects_chain: Optional[List["EffectConfig"]] = Field( effects_chain: list["EffectConfig"] | None = Field(
None, description="Effects chain to apply after generation (overrides profile default)" None, description="Effects chain to apply after generation (overrides profile default)"
) )
@@ -109,19 +109,19 @@ class GenerationResponse(BaseModel):
profile_id: str profile_id: str
text: str text: str
language: str language: str
audio_path: Optional[str] = None audio_path: str | None = None
duration: Optional[float] = None duration: float | None = None
seed: Optional[int] = None seed: int | None = None
instruct: Optional[str] = None instruct: str | None = None
engine: Optional[str] = "qwen" engine: str | None = "qwen"
model_size: Optional[str] = None model_size: str | None = None
status: str = "completed" status: str = "completed"
error: Optional[str] = None error: str | None = None
is_favorited: bool = False is_favorited: bool = False
source: str = "manual" source: str = "manual"
created_at: datetime created_at: datetime
versions: Optional[List["GenerationVersionResponse"]] = None versions: list["GenerationVersionResponse"] | None = None
active_version_id: Optional[str] = None active_version_id: str | None = None
class Config: class Config:
from_attributes = True from_attributes = True
@@ -130,8 +130,8 @@ class GenerationResponse(BaseModel):
class HistoryQuery(BaseModel): class HistoryQuery(BaseModel):
"""Query model for generation history.""" """Query model for generation history."""
profile_id: Optional[str] = None profile_id: str | None = None
search: Optional[str] = None search: str | None = None
limit: int = Field(default=50, ge=1, le=100) limit: int = Field(default=50, ge=1, le=100)
offset: int = Field(default=0, ge=0) offset: int = Field(default=0, ge=0)
@@ -144,18 +144,18 @@ class HistoryResponse(BaseModel):
profile_name: str profile_name: str
text: str text: str
language: str language: str
audio_path: Optional[str] = None audio_path: str | None = None
duration: Optional[float] = None duration: float | None = None
seed: Optional[int] = None seed: int | None = None
instruct: Optional[str] = None instruct: str | None = None
engine: Optional[str] = "qwen" engine: str | None = "qwen"
model_size: Optional[str] = None model_size: str | None = None
status: str = "completed" status: str = "completed"
error: Optional[str] = None error: str | None = None
is_favorited: bool = False is_favorited: bool = False
created_at: datetime created_at: datetime
versions: Optional[List["GenerationVersionResponse"]] = None versions: list["GenerationVersionResponse"] | None = None
active_version_id: Optional[str] = None active_version_id: str | None = None
class Config: class Config:
from_attributes = True from_attributes = True
@@ -164,15 +164,15 @@ class HistoryResponse(BaseModel):
class HistoryListResponse(BaseModel): class HistoryListResponse(BaseModel):
"""Response model for history list.""" """Response model for history list."""
items: List[HistoryResponse] items: list[HistoryResponse]
total: int total: int
class TranscriptionRequest(BaseModel): class TranscriptionRequest(BaseModel):
"""Request model for audio transcription.""" """Request model for audio transcription."""
language: Optional[str] = Field(None, pattern="^(en|zh|ja|ko|de|fr|ru|pt|es|it)$") language: str | None = Field(None, pattern="^(en|zh|ja|ko|de|fr|ru|pt|es|it)$")
model: Optional[str] = Field(None, pattern="^(base|small|medium|large|turbo)$") model: str | None = Field(None, pattern="^(base|small|medium|large|turbo)$")
class TranscriptionResponse(BaseModel): class TranscriptionResponse(BaseModel):
@@ -196,13 +196,13 @@ class CaptureResponse(BaseModel):
id: str id: str
audio_path: str audio_path: str
source: str source: str
language: Optional[str] = None language: str | None = None
duration_ms: Optional[int] = None duration_ms: int | None = None
transcript_raw: str transcript_raw: str
transcript_refined: Optional[str] = None transcript_refined: str | None = None
stt_model: Optional[str] = None stt_model: str | None = None
llm_model: Optional[str] = None llm_model: str | None = None
refinement_flags: Optional[RefinementFlagsModel] = None refinement_flags: RefinementFlagsModel | None = None
created_at: datetime created_at: datetime
class Config: class Config:
@@ -212,7 +212,7 @@ class CaptureResponse(BaseModel):
class CaptureListResponse(BaseModel): class CaptureListResponse(BaseModel):
"""Response model for paginated capture list.""" """Response model for paginated capture list."""
items: List[CaptureResponse] items: list[CaptureResponse]
total: int total: int
@@ -234,15 +234,15 @@ class CaptureCreateResponse(CaptureResponse):
class CaptureRefineRequest(BaseModel): class CaptureRefineRequest(BaseModel):
"""Request to refine a capture's transcript via the LLM.""" """Request to refine a capture's transcript via the LLM."""
flags: Optional[RefinementFlagsModel] = None flags: RefinementFlagsModel | None = None
model_size: Optional[str] = Field(default=None, pattern="^(0\\.6B|1\\.7B|4B)$") model_size: str | None = Field(default=None, pattern="^(0\\.6B|1\\.7B|4B)$")
class CaptureRetranscribeRequest(BaseModel): class CaptureRetranscribeRequest(BaseModel):
"""Request to re-run STT on a capture's audio with a different model.""" """Request to re-run STT on a capture's audio with a different model."""
model: Optional[str] = Field(None, pattern="^(base|small|medium|large|turbo)$") model: str | None = Field(None, pattern="^(base|small|medium|large|turbo)$")
language: Optional[str] = Field(None, pattern="^(en|zh|ja|ko|de|fr|ru|pt|es|it)$") language: str | None = Field(None, pattern="^(en|zh|ja|ko|de|fr|ru|pt|es|it)$")
class CaptureSettingsResponse(BaseModel): class CaptureSettingsResponse(BaseModel):
@@ -256,13 +256,13 @@ class CaptureSettingsResponse(BaseModel):
self_correction: bool = True self_correction: bool = True
preserve_technical: bool = True preserve_technical: bool = True
allow_auto_paste: bool = True allow_auto_paste: bool = True
default_playback_voice_id: Optional[str] = None default_playback_voice_id: str | None = None
hotkey_enabled: bool = False hotkey_enabled: bool = False
keep_mic_warm: bool = False keep_mic_warm: bool = False
chord_push_to_talk_keys: List[str] = Field( chord_push_to_talk_keys: list[str] = Field(
default_factory=default_push_to_talk_chord default_factory=default_push_to_talk_chord
) )
chord_toggle_to_talk_keys: List[str] = Field( chord_toggle_to_talk_keys: list[str] = Field(
default_factory=default_toggle_to_talk_chord default_factory=default_toggle_to_talk_chord
) )
@@ -273,19 +273,19 @@ class CaptureSettingsResponse(BaseModel):
class CaptureSettingsUpdate(BaseModel): class CaptureSettingsUpdate(BaseModel):
"""Partial update for capture settings — every field is optional.""" """Partial update for capture settings — every field is optional."""
stt_model: Optional[str] = Field(default=None, pattern="^(base|small|medium|large|turbo)$") stt_model: str | None = Field(default=None, pattern="^(base|small|medium|large|turbo)$")
language: Optional[str] = None language: str | None = None
auto_refine: Optional[bool] = None auto_refine: bool | None = None
llm_model: Optional[str] = Field(default=None, pattern="^(0\\.6B|1\\.7B|4B)$") llm_model: str | None = Field(default=None, pattern="^(0\\.6B|1\\.7B|4B)$")
smart_cleanup: Optional[bool] = None smart_cleanup: bool | None = None
self_correction: Optional[bool] = None self_correction: bool | None = None
preserve_technical: Optional[bool] = None preserve_technical: bool | None = None
allow_auto_paste: Optional[bool] = None allow_auto_paste: bool | None = None
default_playback_voice_id: Optional[str] = None default_playback_voice_id: str | None = None
hotkey_enabled: Optional[bool] = None hotkey_enabled: bool | None = None
keep_mic_warm: Optional[bool] = None keep_mic_warm: bool | None = None
chord_push_to_talk_keys: Optional[List[str]] = Field(default=None, min_length=1, max_length=6) chord_push_to_talk_keys: list[str] | None = Field(default=None, min_length=1, max_length=6)
chord_toggle_to_talk_keys: Optional[List[str]] = Field(default=None, min_length=1, max_length=6) chord_toggle_to_talk_keys: list[str] | None = Field(default=None, min_length=1, max_length=6)
class GenerationSettingsResponse(BaseModel): class GenerationSettingsResponse(BaseModel):
@@ -303,10 +303,10 @@ class GenerationSettingsResponse(BaseModel):
class GenerationSettingsUpdate(BaseModel): class GenerationSettingsUpdate(BaseModel):
"""Partial update for generation settings — every field is optional.""" """Partial update for generation settings — every field is optional."""
max_chunk_chars: Optional[int] = Field(default=None, ge=100, le=5000) max_chunk_chars: int | None = Field(default=None, ge=100, le=5000)
crossfade_ms: Optional[int] = Field(default=None, ge=0, le=500) crossfade_ms: int | None = Field(default=None, ge=0, le=500)
normalize_audio: Optional[bool] = None normalize_audio: bool | None = None
autoplay_on_generate: Optional[bool] = None autoplay_on_generate: bool | None = None
class MCPClientBindingResponse(BaseModel): class MCPClientBindingResponse(BaseModel):
@@ -315,14 +315,14 @@ class MCPClientBindingResponse(BaseModel):
opt-in personality-rewrite default.""" opt-in personality-rewrite default."""
client_id: str client_id: str
label: Optional[str] = None label: str | None = None
profile_id: Optional[str] = None profile_id: str | None = None
default_engine: Optional[str] = Field( default_engine: str | None = Field(
None, None,
pattern="^(qwen|qwen_custom_voice|luxtts|chatterbox|chatterbox_turbo|tada|kokoro)$", pattern="^(qwen|qwen_custom_voice|luxtts|chatterbox|chatterbox_turbo|tada|kokoro)$",
) )
default_personality: bool = False default_personality: bool = False
last_seen_at: Optional[datetime] = None last_seen_at: datetime | None = None
created_at: datetime created_at: datetime
updated_at: datetime updated_at: datetime
@@ -334,9 +334,9 @@ class MCPClientBindingUpsert(BaseModel):
"""Create or update a binding. Matched by ``client_id``.""" """Create or update a binding. Matched by ``client_id``."""
client_id: str = Field(..., min_length=1, max_length=64) client_id: str = Field(..., min_length=1, max_length=64)
label: Optional[str] = Field(None, max_length=128) label: str | None = Field(None, max_length=128)
profile_id: Optional[str] = None profile_id: str | None = None
default_engine: Optional[str] = Field( default_engine: str | None = Field(
None, None,
pattern="^(qwen|qwen_custom_voice|luxtts|chatterbox|chatterbox_turbo|tada|kokoro)$", pattern="^(qwen|qwen_custom_voice|luxtts|chatterbox|chatterbox_turbo|tada|kokoro)$",
) )
@@ -344,26 +344,26 @@ class MCPClientBindingUpsert(BaseModel):
class MCPClientBindingListResponse(BaseModel): class MCPClientBindingListResponse(BaseModel):
items: List[MCPClientBindingResponse] items: list[MCPClientBindingResponse]
class SpeakRequest(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) text: str = Field(..., min_length=1, max_length=10000)
profile: Optional[str] = Field( profile: str | None = Field(
None, None,
description="Voice profile name or id. Falls back to per-client binding, then default.", description="Voice profile name or id. Falls back to per-client binding, then default.",
) )
engine: Optional[str] = Field( engine: str | None = Field(
None, None,
pattern="^(qwen|qwen_custom_voice|luxtts|chatterbox|chatterbox_turbo|tada|kokoro)$", pattern="^(qwen|qwen_custom_voice|luxtts|chatterbox|chatterbox_turbo|tada|kokoro)$",
) )
personality: Optional[bool] = Field( personality: bool | None = Field(
None, None,
description="When true and the profile has a personality prompt, the input text is rewritten in-character before TTS. When null, the per-client binding's default_personality flag decides.", description="When true and the profile has a personality prompt, the input text is rewritten in-character before TTS. When null, the per-client binding's default_personality flag decides.",
) )
language: Optional[str] = Field( language: str | None = Field(
None, None,
pattern="^(zh|en|ja|ko|de|fr|ru|pt|es|it|he|ar|da|el|fi|hi|ms|nl|no|pl|sv|sw|tr)$", pattern="^(zh|en|ja|ko|de|fr|ru|pt|es|it|he|ar|da|el|fi|hi|ms|nl|no|pl|sv|sw|tr)$",
) )
@@ -373,15 +373,15 @@ class LLMGenerateRequest(BaseModel):
"""Request model for LLM text generation.""" """Request model for LLM text generation."""
prompt: str = Field(..., min_length=1, max_length=50000) prompt: str = Field(..., min_length=1, max_length=50000)
system: Optional[str] = Field(None, max_length=4000) system: str | None = Field(None, max_length=4000)
model_size: Optional[str] = Field(default="0.6B", pattern="^(0\\.6B|1\\.7B|4B)$") model_size: str | None = Field(default="0.6B", pattern="^(0\\.6B|1\\.7B|4B)$")
max_tokens: int = Field(default=512, ge=1, le=4096) max_tokens: int = Field(default=512, ge=1, le=4096)
temperature: float = Field(default=0.7, ge=0.0, le=2.0) temperature: float = Field(default=0.7, ge=0.0, le=2.0)
# Few-shot (user, assistant) pairs prepended as real chat turns. # Few-shot (user, assistant) pairs prepended as real chat turns.
# Used by the refinement service to pin tricky rules (imperatives # Used by the refinement service to pin tricky rules (imperatives
# staying imperatives, technical-term punctuation) that small models # staying imperatives, technical-term punctuation) that small models
# lose when the examples live inline in the system prompt. # lose when the examples live inline in the system prompt.
examples: Optional[List[List[str]]] = Field(default=None, max_length=8) examples: list[list[str]] | None = Field(default=None, max_length=8)
class LLMGenerateResponse(BaseModel): class LLMGenerateResponse(BaseModel):
@@ -418,7 +418,7 @@ class ModelReadiness(BaseModel):
model_name: str model_name: str
display_name: str display_name: str
size: str size: str
size_mb: Optional[int] = None size_mb: int | None = None
class CaptureReadinessResponse(BaseModel): class CaptureReadinessResponse(BaseModel):
@@ -438,15 +438,15 @@ class HealthResponse(BaseModel):
status: str status: str
model_loaded: bool model_loaded: bool
model_downloaded: Optional[bool] = None # Whether model is cached/downloaded model_downloaded: bool | None = None # Whether model is cached/downloaded
model_size: Optional[str] = None # Current model size if loaded model_size: str | None = None # Current model size if loaded
gpu_available: bool gpu_available: bool
gpu_type: Optional[str] = None # GPU type (CUDA, MPS, or None) gpu_type: str | None = None # GPU type (CUDA, MPS, or None)
vram_used_mb: Optional[float] = None vram_used_mb: float | None = None
backend_type: Optional[str] = None # Backend type (mlx or pytorch) backend_type: str | None = None # Backend type (mlx or pytorch)
backend_variant: Optional[str] = None # Binary variant (cpu, cuda, or rocm) backend_variant: str | None = None # Binary variant (cpu, cuda, or rocm)
supports_rocm: bool = False # AMD GPU on Windows — the ROCm backend is applicable supports_rocm: bool = False # AMD GPU on Windows — the ROCm backend is applicable
gpu_compatibility_warning: Optional[str] = None # Warning if GPU arch unsupported gpu_compatibility_warning: str | None = None # Warning if GPU arch unsupported
class DirectoryCheck(BaseModel): class DirectoryCheck(BaseModel):
@@ -455,16 +455,16 @@ class DirectoryCheck(BaseModel):
path: str path: str
exists: bool exists: bool
writable: bool writable: bool
error: Optional[str] = None error: str | None = None
class FilesystemHealthResponse(BaseModel): class FilesystemHealthResponse(BaseModel):
"""Response model for filesystem health check.""" """Response model for filesystem health check."""
healthy: bool healthy: bool
disk_free_mb: Optional[float] = None disk_free_mb: float | None = None
disk_total_mb: Optional[float] = None disk_total_mb: float | None = None
directories: List[DirectoryCheck] directories: list[DirectoryCheck]
class ModelStatus(BaseModel): class ModelStatus(BaseModel):
@@ -472,17 +472,17 @@ class ModelStatus(BaseModel):
model_name: str model_name: str
display_name: str display_name: str
hf_repo_id: Optional[str] = None # HuggingFace repository ID hf_repo_id: str | None = None # HuggingFace repository ID
downloaded: bool downloaded: bool
downloading: bool = False # True if download is in progress downloading: bool = False # True if download is in progress
size_mb: Optional[float] = None size_mb: float | None = None
loaded: bool = False loaded: bool = False
class ModelStatusListResponse(BaseModel): class ModelStatusListResponse(BaseModel):
"""Response model for model status list.""" """Response model for model status list."""
models: List[ModelStatus] models: list[ModelStatus]
class ModelDownloadRequest(BaseModel): class ModelDownloadRequest(BaseModel):
@@ -503,11 +503,11 @@ class ActiveDownloadTask(BaseModel):
model_name: str model_name: str
status: str status: str
started_at: datetime started_at: datetime
error: Optional[str] = None error: str | None = None
progress: Optional[float] = None # 0-100 percentage progress: float | None = None # 0-100 percentage
current: Optional[int] = None # bytes downloaded current: int | None = None # bytes downloaded
total: Optional[int] = None # total bytes total: int | None = None # total bytes
filename: Optional[str] = None # current file being downloaded filename: str | None = None # current file being downloaded
class ActiveGenerationTask(BaseModel): class ActiveGenerationTask(BaseModel):
@@ -522,22 +522,22 @@ class ActiveGenerationTask(BaseModel):
class ActiveTasksResponse(BaseModel): class ActiveTasksResponse(BaseModel):
"""Response model for active tasks.""" """Response model for active tasks."""
downloads: List[ActiveDownloadTask] downloads: list[ActiveDownloadTask]
generations: List[ActiveGenerationTask] generations: list[ActiveGenerationTask]
class AudioChannelCreate(BaseModel): class AudioChannelCreate(BaseModel):
"""Request model for creating an audio channel.""" """Request model for creating an audio channel."""
name: str = Field(..., min_length=1, max_length=100) name: str = Field(..., min_length=1, max_length=100)
device_ids: List[str] = Field(default_factory=list) device_ids: list[str] = Field(default_factory=list)
class AudioChannelUpdate(BaseModel): class AudioChannelUpdate(BaseModel):
"""Request model for updating an audio channel.""" """Request model for updating an audio channel."""
name: Optional[str] = Field(None, min_length=1, max_length=100) name: str | None = Field(None, min_length=1, max_length=100)
device_ids: Optional[List[str]] = None device_ids: list[str] | None = None
class AudioChannelResponse(BaseModel): class AudioChannelResponse(BaseModel):
@@ -546,7 +546,7 @@ class AudioChannelResponse(BaseModel):
id: str id: str
name: str name: str
is_default: bool is_default: bool
device_ids: List[str] device_ids: list[str]
created_at: datetime created_at: datetime
class Config: class Config:
@@ -556,20 +556,20 @@ class AudioChannelResponse(BaseModel):
class ChannelVoiceAssignment(BaseModel): class ChannelVoiceAssignment(BaseModel):
"""Request model for assigning voices to a channel.""" """Request model for assigning voices to a channel."""
profile_ids: List[str] profile_ids: list[str]
class ProfileChannelAssignment(BaseModel): class ProfileChannelAssignment(BaseModel):
"""Request model for assigning channels to a profile.""" """Request model for assigning channels to a profile."""
channel_ids: List[str] channel_ids: list[str]
class StoryCreate(BaseModel): class StoryCreate(BaseModel):
"""Request model for creating a story.""" """Request model for creating a story."""
name: str = Field(..., min_length=1, max_length=100) name: str = Field(..., min_length=1, max_length=100)
description: Optional[str] = Field(None, max_length=500) description: str | None = Field(None, max_length=500)
class StoryResponse(BaseModel): class StoryResponse(BaseModel):
@@ -577,7 +577,7 @@ class StoryResponse(BaseModel):
id: str id: str
name: str name: str
description: Optional[str] description: str | None
created_at: datetime created_at: datetime
updated_at: datetime updated_at: datetime
item_count: int = 0 item_count: int = 0
@@ -592,7 +592,7 @@ class StoryItemDetail(BaseModel):
id: str id: str
story_id: str story_id: str
generation_id: str generation_id: str
version_id: Optional[str] = None version_id: str | None = None
start_time_ms: int start_time_ms: int
track: int = 0 track: int = 0
trim_start_ms: int = 0 trim_start_ms: int = 0
@@ -605,14 +605,14 @@ class StoryItemDetail(BaseModel):
language: str language: str
audio_path: str audio_path: str
duration: float duration: float
seed: Optional[int] seed: int | None
instruct: Optional[str] instruct: str | None
engine: Optional[str] = None engine: str | None = None
volume: float = 1.0 volume: float = 1.0
generation_created_at: datetime generation_created_at: datetime
# Versions available for this generation # Versions available for this generation
versions: Optional[List["GenerationVersionResponse"]] = None versions: list["GenerationVersionResponse"] | None = None
active_version_id: Optional[str] = None active_version_id: str | None = None
class Config: class Config:
from_attributes = True from_attributes = True
@@ -623,10 +623,10 @@ class StoryDetailResponse(BaseModel):
id: str id: str
name: str name: str
description: Optional[str] description: str | None
created_at: datetime created_at: datetime
updated_at: datetime updated_at: datetime
items: List[StoryItemDetail] = [] items: list[StoryItemDetail] = []
class Config: class Config:
from_attributes = True from_attributes = True
@@ -636,8 +636,8 @@ class StoryItemCreate(BaseModel):
"""Request model for adding a generation to a story.""" """Request model for adding a generation to a story."""
generation_id: str generation_id: str
start_time_ms: Optional[int] = None # If not provided, will be calculated automatically start_time_ms: int | None = None # If not provided, will be calculated automatically
track: Optional[int] = 0 # Track number (0 = main track) track: int | None = 0 # Track number (0 = main track)
class StoryItemUpdateTime(BaseModel): class StoryItemUpdateTime(BaseModel):
@@ -650,13 +650,13 @@ class StoryItemUpdateTime(BaseModel):
class StoryItemBatchUpdate(BaseModel): class StoryItemBatchUpdate(BaseModel):
"""Request model for batch updating story item timecodes.""" """Request model for batch updating story item timecodes."""
updates: List[StoryItemUpdateTime] updates: list[StoryItemUpdateTime]
class StoryItemReorder(BaseModel): class StoryItemReorder(BaseModel):
"""Request model for reordering story items.""" """Request model for reordering story items."""
generation_ids: List[str] = Field(..., min_length=1) generation_ids: list[str] = Field(..., min_length=1)
class StoryItemMove(BaseModel): class StoryItemMove(BaseModel):
@@ -682,7 +682,7 @@ class StoryItemSplit(BaseModel):
class StoryItemVersionUpdate(BaseModel): class StoryItemVersionUpdate(BaseModel):
"""Request model for setting a story item's pinned version.""" """Request model for setting a story item's pinned version."""
version_id: Optional[str] = None # null = use generation default version_id: str | None = None # null = use generation default
class StoryItemVolumeUpdate(BaseModel): class StoryItemVolumeUpdate(BaseModel):
@@ -707,23 +707,23 @@ class EffectConfig(BaseModel):
class EffectsChain(BaseModel): class EffectsChain(BaseModel):
"""An ordered list of effects to apply.""" """An ordered list of effects to apply."""
effects: List[EffectConfig] = Field(default_factory=list) effects: list[EffectConfig] = Field(default_factory=list)
class EffectPresetCreate(BaseModel): class EffectPresetCreate(BaseModel):
"""Request model for creating an effect preset.""" """Request model for creating an effect preset."""
name: str = Field(..., min_length=1, max_length=100) name: str = Field(..., min_length=1, max_length=100)
description: Optional[str] = Field(None, max_length=500) description: str | None = Field(None, max_length=500)
effects_chain: List[EffectConfig] effects_chain: list[EffectConfig]
class EffectPresetUpdate(BaseModel): class EffectPresetUpdate(BaseModel):
"""Request model for updating an effect preset.""" """Request model for updating an effect preset."""
name: Optional[str] = Field(None, min_length=1, max_length=100) name: str | None = Field(None, min_length=1, max_length=100)
description: Optional[str] = None description: str | None = None
effects_chain: Optional[List[EffectConfig]] = None effects_chain: list[EffectConfig] | None = None
class EffectPresetResponse(BaseModel): class EffectPresetResponse(BaseModel):
@@ -731,8 +731,8 @@ class EffectPresetResponse(BaseModel):
id: str id: str
name: str name: str
description: Optional[str] = None description: str | None = None
effects_chain: List[EffectConfig] effects_chain: list[EffectConfig]
is_builtin: bool = False is_builtin: bool = False
created_at: datetime created_at: datetime
@@ -747,8 +747,8 @@ class GenerationVersionResponse(BaseModel):
generation_id: str generation_id: str
label: str label: str
audio_path: str audio_path: str
effects_chain: Optional[List[EffectConfig]] = None effects_chain: list[EffectConfig] | None = None
source_version_id: Optional[str] = None source_version_id: str | None = None
is_default: bool is_default: bool
created_at: datetime created_at: datetime
@@ -759,18 +759,18 @@ class GenerationVersionResponse(BaseModel):
class ApplyEffectsRequest(BaseModel): class ApplyEffectsRequest(BaseModel):
"""Request to apply effects to an existing generation.""" """Request to apply effects to an existing generation."""
effects_chain: List[EffectConfig] effects_chain: list[EffectConfig]
source_version_id: Optional[str] = Field( source_version_id: str | None = Field(
None, description="Version to use as source audio (defaults to clean/original)" None, description="Version to use as source audio (defaults to clean/original)"
) )
label: Optional[str] = Field(None, max_length=100, description="Label for this version (auto-generated if omitted)") label: str | None = Field(None, max_length=100, description="Label for this version (auto-generated if omitted)")
set_as_default: bool = Field(default=True, description="Set this version as the default") set_as_default: bool = Field(default=True, description="Set this version as the default")
class ProfileEffectsUpdate(BaseModel): class ProfileEffectsUpdate(BaseModel):
"""Request to update the default effects chain on a profile.""" """Request to update the default effects chain on a profile."""
effects_chain: Optional[List[EffectConfig]] = Field(None, description="Effects chain (null to remove)") effects_chain: list[EffectConfig] | None = Field(None, description="Effects chain (null to remove)")
class AvailableEffectParam(BaseModel): class AvailableEffectParam(BaseModel):
@@ -795,7 +795,7 @@ class AvailableEffect(BaseModel):
class AvailableEffectsResponse(BaseModel): class AvailableEffectsResponse(BaseModel):
"""Response listing all available effect types.""" """Response listing all available effect types."""
effects: List[AvailableEffect] effects: list[AvailableEffect]
# ─── Cloud (backup & sync) ────────────────────────────────────────────── # ─── Cloud (backup & sync) ──────────────────────────────────────────────
@@ -812,8 +812,8 @@ class CloudStatusResponse(BaseModel):
"""Current link between this device and a Voicebox Cloud account.""" """Current link between this device and a Voicebox Cloud account."""
connected: bool connected: bool
device_name: Optional[str] = None device_name: str | None = None
account_user_id: Optional[str] = None account_user_id: str | None = None
key_prefix: Optional[str] = None key_prefix: str | None = None
connected_at: Optional[datetime] = None connected_at: datetime | None = None
dashboard_url: str dashboard_url: str
@@ -56,7 +56,6 @@ import sys
import tempfile import tempfile
import types import types
# Diagnostics — log hook activity to a file alongside the bundle so we can # Diagnostics — log hook activity to a file alongside the bundle so we can
# see what's happening when the server is run as a sidecar (no stdout for # see what's happening when the server is run as a sidecar (no stdout for
# runtime hook prints). Safe no-op if the file can't be written. # runtime hook prints). Safe no-op if the file can't be written.
+32 -4
View File
@@ -1,6 +1,6 @@
[project] [project]
name = "voicebox-backend" name = "voicebox-backend"
version = "0.2.3" version = "0.5.0"
requires-python = ">=3.12" requires-python = ">=3.12"
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -49,19 +49,43 @@ ignore = [
"SIM108", # use ternary operator (sometimes less readable) "SIM108", # use ternary operator (sometimes less readable)
"B008", # function call in default argument (FastAPI Depends() pattern) "B008", # function call in default argument (FastAPI Depends() pattern)
"UP007", # use X | Y for union (auto-fixed by UP, but noisy on big diffs) "UP007", # use X | Y for union (auto-fixed by UP, but noisy on big diffs)
# Existing-violation baseline so ruff can gate CI. Remove entries from
# this list as the remaining occurrences are fixed; counts are as of
# 2026-07-26 after the auto-fix pass.
"B904", # raise without `from` inside except (49) -- needs per-site from err/from None
"SIM105", # try/except/pass instead of contextlib.suppress (14)
"N806", # non-lowercase variable in function (9)
"RUF002", # ambiguous unicode in docstring (6)
"F841", # unused variable (5)
"N803", # invalid argument name (5)
"B007", # unused loop control variable (4)
"ERA001", # commented-out code (4)
"SIM102", # collapsible if (4)
"SIM117", # multiple with statements (4)
"SIM115", # open() without context manager (3)
"RUF001", # ambiguous unicode in string (2)
"RUF012", # mutable class default (2)
"SIM110", # reimplemented builtin (2)
"RUF006", # asyncio dangling task (1)
"RUF034", # useless if-else (1)
] ]
# Per-file rule overrides. # Per-file rule overrides.
[tool.ruff.lint.per-file-ignores] [tool.ruff.lint.per-file-ignores]
# Tests can use assert, print, and magic values freely. # Tests can use assert, print, magic values, and script-style setup freely.
"tests/**" = ["S101", "T201", "PLR2004", "ERA001"] "tests/**" = ["S101", "T201", "PLR2004", "ERA001", "E402", "PT011", "PT018", "PT019"]
# __init__.py re-exports are expected to have unused imports. # __init__.py re-exports are expected to have unused imports.
"**/__init__.py" = ["F401"] "**/__init__.py" = ["F401"]
# Entry points and scripts legitimately use print. # Entry points and scripts legitimately use print.
"server.py" = ["T201"]
"main.py" = ["T201"] "main.py" = ["T201"]
# AMD GPU env vars must be set before torch import. # AMD GPU env vars must be set before torch import.
"app.py" = ["E402"] "app.py" = ["E402"]
# Environment and stdout hardening must run before heavy imports.
"server.py" = ["T201", "E402"]
"backends/__init__.py" = ["E402"]
"backends/mlx_backend.py" = ["E402"]
"backends/pytorch_backend.py" = ["E402"]
[tool.ruff.lint.isort] [tool.ruff.lint.isort]
known-first-party = ["backend"] known-first-party = ["backend"]
@@ -81,3 +105,7 @@ docstring-code-format = true
[tool.pytest.ini_options] [tool.pytest.ini_options]
testpaths = ["tests"] testpaths = ["tests"]
asyncio_mode = "auto" asyncio_mode = "auto"
markers = [
"slow: long-running tests, deselect with '-m \"not slow\"'",
"timeout: per-test timeout in seconds (enforced only when pytest-timeout is installed)",
]
+18 -18
View File
@@ -5,26 +5,26 @@ from fastapi import FastAPI
def register_routers(app: FastAPI) -> None: def register_routers(app: FastAPI) -> None:
"""Include all domain routers on the application.""" """Include all domain routers on the application."""
from .health import router as health_router
from .profiles import router as profiles_router
from .channels import router as channels_router
from .generations import router as generations_router
from .history import router as history_router
from .transcription import router as transcription_router
from .llm import router as llm_router
from .captures import router as captures_router
from .stories import router as stories_router
from .effects import router as effects_router
from .audio import router as audio_router from .audio import router as audio_router
from .models import router as models_router from .captures import router as captures_router
from .settings import router as settings_router from .channels import router as channels_router
from .tasks import router as tasks_router
from .cuda import router as cuda_router
from .rocm import router as rocm_router
from .speak import router as speak_router
from .mcp_bindings import router as mcp_bindings_router
from .events import router as events_router
from .cloud import router as cloud_router from .cloud import router as cloud_router
from .cuda import router as cuda_router
from .effects import router as effects_router
from .events import router as events_router
from .generations import router as generations_router
from .health import router as health_router
from .history import router as history_router
from .llm import router as llm_router
from .mcp_bindings import router as mcp_bindings_router
from .models import router as models_router
from .profiles import router as profiles_router
from .rocm import router as rocm_router
from .settings import router as settings_router
from .speak import router as speak_router
from .stories import router as stories_router
from .tasks import router as tasks_router
from .transcription import router as transcription_router
app.include_router(health_router) app.include_router(health_router)
app.include_router(profiles_router) app.include_router(profiles_router)
+2 -2
View File
@@ -7,9 +7,9 @@ from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import FileResponse from fastapi.responses import FileResponse
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .. import config, models from .. import config
from ..services import history
from ..database import get_db from ..database import get_db
from ..services import history
router = APIRouter() router = APIRouter()
+1 -2
View File
@@ -10,8 +10,7 @@ from .. import config, models
from ..backends import get_llm_model_configs, get_stt_model_configs from ..backends import get_llm_model_configs, get_stt_model_configs
from ..backends.base import is_model_cached from ..backends.base import is_model_cached
from ..database import Capture as DBCapture, get_db from ..database import Capture as DBCapture, get_db
from ..services import captures as captures_service from ..services import captures as captures_service, settings as settings_service
from ..services import settings as settings_service
from ..services.refinement import RefinementFlags from ..services.refinement import RefinementFlags
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
+1 -1
View File
@@ -4,8 +4,8 @@ from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .. import models from .. import models
from ..services import channels
from ..database import get_db from ..database import get_db
from ..services import channels
router = APIRouter() router = APIRouter()
+3 -3
View File
@@ -9,8 +9,8 @@ from fastapi.responses import StreamingResponse
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .. import config, models from .. import config, models
from ..services import history
from ..database import Generation as DBGeneration, get_db from ..database import Generation as DBGeneration, get_db
from ..services import history
router = APIRouter() router = APIRouter()
@@ -29,8 +29,8 @@ async def preview_effects(
raise HTTPException(status_code=400, detail="Generation is not completed") raise HTTPException(status_code=400, detail="Generation is not completed")
from ..services import versions as versions_mod from ..services import versions as versions_mod
from ..utils.effects import apply_effects, validate_effects_chain
from ..utils.audio import load_audio from ..utils.audio import load_audio
from ..utils.effects import apply_effects, validate_effects_chain
chain_dicts = [e.model_dump() for e in data.effects_chain] chain_dicts = [e.model_dump() for e in data.effects_chain]
error = validate_effects_chain(chain_dicts) error = validate_effects_chain(chain_dicts)
@@ -170,8 +170,8 @@ async def apply_effects_to_generation(
raise HTTPException(status_code=400, detail="Generation is not completed") raise HTTPException(status_code=400, detail="Generation is not completed")
from ..services import versions as versions_mod from ..services import versions as versions_mod
from ..utils.effects import apply_effects, validate_effects_chain
from ..utils.audio import load_audio, save_audio from ..utils.audio import load_audio, save_audio
from ..utils.effects import apply_effects, validate_effects_chain
chain_dicts = [e.model_dump() for e in data.effects_chain] chain_dicts = [e.model_dump() for e in data.effects_chain]
error = validate_effects_chain(chain_dicts) error = validate_effects_chain(chain_dicts)
-1
View File
@@ -14,7 +14,6 @@ from sse_starlette.sse import EventSourceResponse
from ..mcp_server import events as mcp_events from ..mcp_server import events as mcp_events
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
router = APIRouter() router = APIRouter()
+7 -2
View File
@@ -10,8 +10,8 @@ from fastapi.responses import StreamingResponse
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .. import config, models from .. import config, models
from ..services import history, personality, profiles, tts
from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db
from ..services import history, personality, profiles, tts
from ..services.generation import run_generation from ..services.generation import run_generation
from ..services.task_queue import cancel_generation as cancel_generation_job, enqueue_generation from ..services.task_queue import cancel_generation as cancel_generation_job, enqueue_generation
from ..utils.audio import load_audio from ..utils.audio import load_audio
@@ -321,7 +321,12 @@ 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,
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:
+3 -4
View File
@@ -6,13 +6,11 @@ import signal
from pathlib import Path from pathlib import Path
import torch import torch
from fastapi import APIRouter, Depends from fastapi import APIRouter
from fastapi.responses import FileResponse from fastapi.responses import FileResponse
from sqlalchemy.orm import Session
from .. import config, models from .. import config, models
from ..services import tts from ..services import tts
from ..database import get_db
from ..utils.platform_detect import get_backend_type, is_amd_gpu_windows from ..utils.platform_detect import get_backend_type, is_amd_gpu_windows
router = APIRouter() router = APIRouter()
@@ -56,9 +54,10 @@ async def watchdog_disable():
@router.get("/health", response_model=models.HealthResponse) @router.get("/health", response_model=models.HealthResponse)
async def health(): async def health():
"""Health check endpoint.""" """Health check endpoint."""
from huggingface_hub import constants as hf_constants
from pathlib import Path from pathlib import Path
from huggingface_hub import constants as hf_constants
tts_model = tts.get_tts_model() tts_model = tts.get_tts_model()
backend_type = get_backend_type() backend_type = get_backend_type()
+1 -1
View File
@@ -7,9 +7,9 @@ from fastapi.responses import FileResponse, StreamingResponse
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .. import config, models from .. import config, models
from ..services import export_import, history
from ..app import safe_content_disposition from ..app import safe_content_disposition
from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db
from ..services import export_import, history
router = APIRouter() router = APIRouter()
+2 -3
View File
@@ -6,7 +6,7 @@ column is the same value the MCP client sends in ``X-Voicebox-Client-Id``
(or the stdio shim pulls from ``VOICEBOX_CLIENT_ID``). (or the stdio shim pulls from ``VOICEBOX_CLIENT_ID``).
""" """
from datetime import datetime, timezone from datetime import UTC, datetime
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@@ -15,7 +15,6 @@ from .. import models
from ..database import get_db from ..database import get_db
from ..database.models import MCPClientBinding from ..database.models import MCPClientBinding
router = APIRouter() router = APIRouter()
@@ -56,7 +55,7 @@ async def upsert_mcp_binding(
row.profile_id = data.profile_id row.profile_id = data.profile_id
row.default_engine = data.default_engine row.default_engine = data.default_engine
row.default_personality = data.default_personality row.default_personality = data.default_personality
row.updated_at = datetime.now(timezone.utc) row.updated_at = datetime.now(UTC)
db.commit() db.commit()
db.refresh(row) db.refresh(row)
return models.MCPClientBindingResponse.model_validate(row) return models.MCPClientBindingResponse.model_validate(row)
+8 -8
View File
@@ -4,13 +4,12 @@ import asyncio
import shutil import shutil
from pathlib import Path from pathlib import Path
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, HTTPException
from fastapi.responses import StreamingResponse from fastapi.responses import StreamingResponse
from sqlalchemy.orm import Session
from .. import models from .. import models
from ..utils.platform_detect import get_backend_type
from ..services.task_queue import create_background_task from ..services.task_queue import create_background_task
from ..utils.platform_detect import get_backend_type
from ..utils.progress import get_progress_manager from ..utils.progress import get_progress_manager
from ..utils.tasks import get_task_manager from ..utils.tasks import get_task_manager
@@ -171,7 +170,7 @@ async def migrate_models(request: models.ModelMigrateRequest):
status="downloading", status="downloading",
) )
except Exception as e: except Exception as e:
errors.append(f"{item.name}: {str(e)}") errors.append(f"{item.name}: {e!s}")
else: else:
total_bytes = sum(_get_dir_size(d) for d in model_dirs) total_bytes = sum(_get_dir_size(d) for d in model_dirs)
progress_manager.update_progress( progress_manager.update_progress(
@@ -190,7 +189,7 @@ async def migrate_models(request: models.ModelMigrateRequest):
await asyncio.to_thread(shutil.rmtree, str(item)) await asyncio.to_thread(shutil.rmtree, str(item))
moved += 1 moved += 1
except Exception as e: except Exception as e:
errors.append(f"{item.name}: {str(e)}") errors.append(f"{item.name}: {e!s}")
progress_manager.update_progress("migration", 1, 1, status="complete") progress_manager.update_progress("migration", 1, 1, status="complete")
progress_manager.mark_complete("migration") progress_manager.mark_complete("migration")
@@ -240,7 +239,7 @@ async def get_model_status():
except ImportError: except ImportError:
use_scan_cache = False use_scan_cache = False
from ..backends import get_all_model_configs, check_model_loaded from ..backends import check_model_loaded, get_all_model_configs
registry_configs = get_all_model_configs() registry_configs = get_all_model_configs()
model_configs = [ model_configs = [
@@ -445,6 +444,7 @@ async def cancel_model_download(request: models.ModelDownloadRequest):
async def delete_model(model_name: str): async def delete_model(model_name: str):
"""Delete a downloaded model from the HuggingFace cache.""" """Delete a downloaded model from the HuggingFace cache."""
from huggingface_hub import constants as hf_constants from huggingface_hub import constants as hf_constants
from ..backends import get_model_config, unload_model_by_config from ..backends import get_model_config, unload_model_by_config
config = get_model_config(model_name) config = get_model_config(model_name)
@@ -465,11 +465,11 @@ async def delete_model(model_name: str):
try: try:
shutil.rmtree(repo_cache_dir) shutil.rmtree(repo_cache_dir)
except OSError as e: except OSError as e:
raise HTTPException(status_code=500, detail=f"Failed to delete model cache directory: {str(e)}") raise HTTPException(status_code=500, detail=f"Failed to delete model cache directory: {e!s}")
return {"message": f"Model {model_name} deleted successfully"} return {"message": f"Model {model_name} deleted successfully"}
except HTTPException: except HTTPException:
raise raise
except Exception as e: except Exception as e:
raise HTTPException(status_code=500, detail=f"Failed to delete model: {str(e)}") raise HTTPException(status_code=500, detail=f"Failed to delete model: {e!s}")
+1 -1
View File
@@ -186,7 +186,7 @@ async def add_profile_sample(
except ValueError as e: except ValueError as e:
raise HTTPException(status_code=400, detail=str(e)) raise HTTPException(status_code=400, detail=str(e))
except Exception as e: except Exception as e:
raise HTTPException(status_code=500, detail=f"Failed to process audio file: {str(e)}") raise HTTPException(status_code=500, detail=f"Failed to process audio file: {e!s}")
finally: finally:
Path(tmp_path).unlink(missing_ok=True) Path(tmp_path).unlink(missing_ok=True)
-1
View File
@@ -18,7 +18,6 @@ from ..database import MCPClientBinding, get_db
from ..mcp_server import events as mcp_events from ..mcp_server import events as mcp_events
from ..mcp_server.resolve import resolve_profile from ..mcp_server.resolve import resolve_profile
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
router = APIRouter() router = APIRouter()
+1 -1
View File
@@ -7,9 +7,9 @@ from fastapi.responses import StreamingResponse
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .. import database, models from .. import database, models
from ..services import stories
from ..app import safe_content_disposition from ..app import safe_content_disposition
from ..database import get_db from ..database import get_db
from ..services import stories
router = APIRouter() router = APIRouter()
+2 -3
View File
@@ -2,13 +2,12 @@
from datetime import datetime from datetime import datetime
from fastapi import APIRouter from fastapi import APIRouter, HTTPException
from .. import models from .. import models
from ..utils.cache import clear_voice_prompt_cache from ..utils.cache import clear_voice_prompt_cache
from ..utils.progress import get_progress_manager from ..utils.progress import get_progress_manager
from ..utils.tasks import get_task_manager from ..utils.tasks import get_task_manager
from fastapi import HTTPException
router = APIRouter() router = APIRouter()
@@ -39,7 +38,7 @@ async def clear_cache():
"files_deleted": deleted_count, "files_deleted": deleted_count,
} }
except Exception as e: except Exception as e:
raise HTTPException(status_code=500, detail=f"Failed to clear cache: {str(e)}") raise HTTPException(status_code=500, detail=f"Failed to clear cache: {e!s}")
@router.get("/tasks/active", response_model=models.ActiveTasksResponse) @router.get("/tasks/active", response_model=models.ActiveTasksResponse)
+1 -1
View File
@@ -29,8 +29,8 @@ async def transcribe_audio(
tmp_path = tmp.name tmp_path = tmp.name
try: try:
from ..utils.audio import load_audio
from ..backends import WHISPER_HF_REPOS from ..backends import WHISPER_HF_REPOS
from ..utils.audio import load_audio
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
+5 -4
View File
@@ -5,9 +5,10 @@ This module provides an entry point that works with PyInstaller by using
absolute imports instead of relative imports. absolute imports instead of relative imports.
""" """
import sys
import os import os
import re import re
import sys
# On Windows with --noconsole (PyInstaller), sys.stdout/stderr are None. # On Windows with --noconsole (PyInstaller), sys.stdout/stderr are None.
# They can also be broken file objects in some edge cases. # They can also be broken file objects in some edge cases.
@@ -30,6 +31,7 @@ if not _is_writable(sys.stderr):
# PyInstaller + multiprocessing: child processes re-execute the frozen binary # PyInstaller + multiprocessing: child processes re-execute the frozen binary
# with internal arguments. freeze_support() handles this and exits early. # with internal arguments. freeze_support() handles this and exits early.
import multiprocessing import multiprocessing
multiprocessing.freeze_support() multiprocessing.freeze_support()
# In frozen builds, piper_phonemize's espeak-ng C library falls back to # In frozen builds, piper_phonemize's espeak-ng C library falls back to
@@ -167,9 +169,8 @@ def _start_parent_watchdog(parent_pid, data_dir=None):
return True # process exists, we just can't open it return True # process exists, we just can't open it
watchdog_logger.info(f"PID {pid}: OpenProcess failed, error={error}") watchdog_logger.info(f"PID {pid}: OpenProcess failed, error={error}")
return False return False
else: os.kill(pid, 0)
os.kill(pid, 0) return True
return True
except (OSError, PermissionError): except (OSError, PermissionError):
return False return False
+9 -10
View File
@@ -12,7 +12,6 @@ import json
import logging import logging
import uuid import uuid
from pathlib import Path from pathlib import Path
from typing import Optional
import soundfile as sf import soundfile as sf
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@@ -35,7 +34,7 @@ WHISPER_NATIVE_FORMATS = (".wav", ".mp3", ".flac", ".ogg")
def _to_response(row: DBCapture) -> CaptureResponse: def _to_response(row: DBCapture) -> CaptureResponse:
flags_model: Optional[RefinementFlagsModel] = None flags_model: RefinementFlagsModel | None = None
if row.refinement_flags: if row.refinement_flags:
try: try:
flags_model = RefinementFlagsModel(**json.loads(row.refinement_flags)) flags_model = RefinementFlagsModel(**json.loads(row.refinement_flags))
@@ -62,8 +61,8 @@ async def create_capture(
audio_bytes: bytes, audio_bytes: bytes,
filename: str, filename: str,
source: str, source: str,
language: Optional[str], language: str | None,
stt_model: Optional[str], stt_model: str | None,
db: Session, db: Session,
) -> CaptureResponse: ) -> CaptureResponse:
"""Persist raw audio, run STT, store the row.""" """Persist raw audio, run STT, store the row."""
@@ -159,7 +158,7 @@ def list_captures(db: Session, limit: int = 50, offset: int = 0) -> tuple[list[C
return [_to_response(r) for r in rows], total return [_to_response(r) for r in rows], total
def get_capture(capture_id: str, db: Session) -> Optional[CaptureResponse]: def get_capture(capture_id: str, db: Session) -> CaptureResponse | None:
row = db.query(DBCapture).filter(DBCapture.id == capture_id).first() row = db.query(DBCapture).filter(DBCapture.id == capture_id).first()
return _to_response(row) if row else None return _to_response(row) if row else None
@@ -184,9 +183,9 @@ def delete_capture(capture_id: str, db: Session) -> bool:
async def refine_capture( async def refine_capture(
capture_id: str, capture_id: str,
flags: RefinementFlags, flags: RefinementFlags,
model_size: Optional[str], model_size: str | None,
db: Session, db: Session,
) -> Optional[CaptureResponse]: ) -> CaptureResponse | None:
row = db.query(DBCapture).filter(DBCapture.id == capture_id).first() row = db.query(DBCapture).filter(DBCapture.id == capture_id).first()
if not row: if not row:
return None return None
@@ -207,10 +206,10 @@ async def refine_capture(
async def retranscribe_capture( async def retranscribe_capture(
capture_id: str, capture_id: str,
stt_model: Optional[str], stt_model: str | None,
language: Optional[str], language: str | None,
db: Session, db: Session,
) -> Optional[CaptureResponse]: ) -> CaptureResponse | None:
row = db.query(DBCapture).filter(DBCapture.id == capture_id).first() row = db.query(DBCapture).filter(DBCapture.id == capture_id).first()
if not row: if not row:
return None return None
+14 -14
View File
@@ -2,27 +2,27 @@
Audio channel management module. Audio channel management module.
""" """
from typing import List, Optional
from datetime import datetime
import uuid import uuid
from datetime import datetime
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from ..models import (
AudioChannelCreate,
AudioChannelUpdate,
AudioChannelResponse,
ChannelVoiceAssignment,
ProfileChannelAssignment,
)
from ..database import ( from ..database import (
AudioChannel as DBAudioChannel, AudioChannel as DBAudioChannel,
ChannelDeviceMapping as DBChannelDeviceMapping, ChannelDeviceMapping as DBChannelDeviceMapping,
ProfileChannelMapping as DBProfileChannelMapping, ProfileChannelMapping as DBProfileChannelMapping,
VoiceProfile as DBVoiceProfile, VoiceProfile as DBVoiceProfile,
) )
from ..models import (
AudioChannelCreate,
AudioChannelResponse,
AudioChannelUpdate,
ChannelVoiceAssignment,
ProfileChannelAssignment,
)
async def list_channels(db: Session) -> List[AudioChannelResponse]: async def list_channels(db: Session) -> list[AudioChannelResponse]:
"""List all audio channels.""" """List all audio channels."""
channels = db.query(DBAudioChannel).all() channels = db.query(DBAudioChannel).all()
result = [] result = []
@@ -45,7 +45,7 @@ async def list_channels(db: Session) -> List[AudioChannelResponse]:
return result return result
async def get_channel(channel_id: str, db: Session) -> Optional[AudioChannelResponse]: async def get_channel(channel_id: str, db: Session) -> AudioChannelResponse | None:
"""Get a channel by ID.""" """Get a channel by ID."""
channel = db.query(DBAudioChannel).filter_by(id=channel_id).first() channel = db.query(DBAudioChannel).filter_by(id=channel_id).first()
if not channel: if not channel:
@@ -111,7 +111,7 @@ async def update_channel(
channel_id: str, channel_id: str,
data: AudioChannelUpdate, data: AudioChannelUpdate,
db: Session, db: Session,
) -> Optional[AudioChannelResponse]: ) -> AudioChannelResponse | None:
"""Update an audio channel.""" """Update an audio channel."""
channel = db.query(DBAudioChannel).filter_by(id=channel_id).first() channel = db.query(DBAudioChannel).filter_by(id=channel_id).first()
if not channel: if not channel:
@@ -185,7 +185,7 @@ async def delete_channel(channel_id: str, db: Session) -> bool:
return True return True
async def get_channel_voices(channel_id: str, db: Session) -> List[str]: async def get_channel_voices(channel_id: str, db: Session) -> list[str]:
"""Get list of profile IDs assigned to a channel.""" """Get list of profile IDs assigned to a channel."""
mappings = db.query(DBProfileChannelMapping).filter_by( mappings = db.query(DBProfileChannelMapping).filter_by(
channel_id=channel_id channel_id=channel_id
@@ -224,7 +224,7 @@ async def set_channel_voices(
db.commit() db.commit()
async def get_profile_channels(profile_id: str, db: Session) -> List[str]: async def get_profile_channels(profile_id: str, db: Session) -> list[str]:
"""Get list of channel IDs assigned to a profile.""" """Get list of channel IDs assigned to a profile."""
mappings = db.query(DBProfileChannelMapping).filter_by( mappings = db.query(DBProfileChannelMapping).filter_by(
profile_id=profile_id profile_id=profile_id
+9 -13
View File
@@ -19,11 +19,10 @@ import os
import sys import sys
import tarfile import tarfile
from pathlib import Path from pathlib import Path
from typing import Optional
from .. import __version__
from ..config import get_data_dir from ..config import get_data_dir
from ..utils.progress import get_progress_manager from ..utils.progress import get_progress_manager
from .. import __version__
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -63,7 +62,7 @@ def get_cuda_exe_name() -> str:
return "voicebox-server-cuda" return "voicebox-server-cuda"
def get_cuda_binary_path() -> Optional[Path]: def get_cuda_binary_path() -> Path | None:
"""Return path to the CUDA executable if it exists inside the onedir.""" """Return path to the CUDA executable if it exists inside the onedir."""
p = get_cuda_dir() / get_cuda_exe_name() p = get_cuda_dir() / get_cuda_exe_name()
if p.exists(): if p.exists():
@@ -76,7 +75,7 @@ def get_cuda_libs_manifest_path() -> Path:
return get_cuda_dir() / "cuda-libs.json" return get_cuda_dir() / "cuda-libs.json"
def get_installed_cuda_libs_version() -> Optional[str]: def get_installed_cuda_libs_version() -> str | None:
"""Read the installed CUDA libs version from cuda-libs.json, or None.""" """Read the installed CUDA libs version from cuda-libs.json, or None."""
manifest_path = get_cuda_libs_manifest_path() manifest_path = get_cuda_libs_manifest_path()
if not manifest_path.exists(): if not manifest_path.exists():
@@ -114,7 +113,7 @@ def get_cuda_status() -> dict:
} }
def _needs_server_download(version: Optional[str] = None) -> bool: def _needs_server_download(version: str | None = None) -> bool:
"""Check if the server core archive needs to be (re)downloaded.""" """Check if the server core archive needs to be (re)downloaded."""
cuda_path = get_cuda_binary_path() cuda_path = get_cuda_binary_path()
if not cuda_path: if not cuda_path:
@@ -138,7 +137,7 @@ def _needs_cuda_libs_download() -> bool:
async def _download_and_extract_archive( async def _download_and_extract_archive(
client, client,
url: str, url: str,
sha256_url: Optional[str], sha256_url: str | None,
dest_dir: Path, dest_dir: Path,
label: str, label: str,
progress_offset: int, progress_offset: int,
@@ -223,10 +222,7 @@ async def _download_and_extract_archive(
status="downloading", status="downloading",
) )
with tarfile.open(temp_path, "r:gz") as tar: with tarfile.open(temp_path, "r:gz") as tar:
if sys.version_info >= (3, 12): tar.extractall(path=dest_dir, filter="data")
tar.extractall(path=dest_dir, filter="data")
else:
tar.extractall(path=dest_dir)
logger.info(f"{label}: extracted to {dest_dir}") logger.info(f"{label}: extracted to {dest_dir}")
finally: finally:
@@ -235,7 +231,7 @@ async def _download_and_extract_archive(
return downloaded return downloaded
async def download_cuda_binary(version: Optional[str] = None): async def download_cuda_binary(version: str | None = None):
"""Download the CUDA backend (server core + CUDA libs if needed). """Download the CUDA backend (server core + CUDA libs if needed).
Downloads both archives from GitHub Releases, extracts them into Downloads both archives from GitHub Releases, extracts them into
@@ -255,7 +251,7 @@ async def download_cuda_binary(version: Optional[str] = None):
await _download_cuda_binary_locked(version) await _download_cuda_binary_locked(version)
async def _download_cuda_binary_locked(version: Optional[str] = None): async def _download_cuda_binary_locked(version: str | None = None):
"""Inner implementation of download_cuda_binary, called under _download_lock.""" """Inner implementation of download_cuda_binary, called under _download_lock."""
import httpx import httpx
@@ -353,7 +349,7 @@ async def _download_cuda_binary_locked(version: Optional[str] = None):
raise raise
def get_cuda_binary_version() -> Optional[str]: def get_cuda_binary_version() -> str | None:
"""Get the version of the installed CUDA binary, or None if not installed.""" """Get the version of the installed CUDA binary, or None if not installed."""
import subprocess import subprocess
+7 -9
View File
@@ -6,15 +6,13 @@ from __future__ import annotations
import json import json
import uuid import uuid
from typing import List, Optional
from sqlalchemy.orm import Session
from sqlalchemy.exc import IntegrityError from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from ..utils.effects import validate_effects_chain
from ..database import EffectPreset as DBEffectPreset from ..database import EffectPreset as DBEffectPreset
from ..models import EffectPresetResponse, EffectPresetCreate, EffectPresetUpdate, EffectConfig from ..models import EffectConfig, EffectPresetCreate, EffectPresetResponse, EffectPresetUpdate
from ..utils.effects import validate_effects_chain
def _preset_response(p: DBEffectPreset) -> EffectPresetResponse: def _preset_response(p: DBEffectPreset) -> EffectPresetResponse:
@@ -30,13 +28,13 @@ def _preset_response(p: DBEffectPreset) -> EffectPresetResponse:
) )
def list_presets(db: Session) -> List[EffectPresetResponse]: def list_presets(db: Session) -> list[EffectPresetResponse]:
"""List all effect presets (built-in + user-created).""" """List all effect presets (built-in + user-created)."""
presets = db.query(DBEffectPreset).order_by(DBEffectPreset.sort_order, DBEffectPreset.name).all() presets = db.query(DBEffectPreset).order_by(DBEffectPreset.sort_order, DBEffectPreset.name).all()
return [_preset_response(p) for p in presets] return [_preset_response(p) for p in presets]
def get_preset(preset_id: str, db: Session) -> Optional[EffectPresetResponse]: def get_preset(preset_id: str, db: Session) -> EffectPresetResponse | None:
"""Get a preset by ID.""" """Get a preset by ID."""
p = db.query(DBEffectPreset).filter_by(id=preset_id).first() p = db.query(DBEffectPreset).filter_by(id=preset_id).first()
if not p: if not p:
@@ -44,7 +42,7 @@ def get_preset(preset_id: str, db: Session) -> Optional[EffectPresetResponse]:
return _preset_response(p) return _preset_response(p)
def get_preset_by_name(name: str, db: Session) -> Optional[EffectPresetResponse]: def get_preset_by_name(name: str, db: Session) -> EffectPresetResponse | None:
"""Get a preset by name.""" """Get a preset by name."""
p = db.query(DBEffectPreset).filter_by(name=name).first() p = db.query(DBEffectPreset).filter_by(name=name).first()
if not p: if not p:
@@ -82,7 +80,7 @@ def create_preset(data: EffectPresetCreate, db: Session) -> EffectPresetResponse
return _preset_response(preset) return _preset_response(preset)
def update_preset(preset_id: str, data: EffectPresetUpdate, db: Session) -> Optional[EffectPresetResponse]: def update_preset(preset_id: str, data: EffectPresetUpdate, db: Session) -> EffectPresetResponse | None:
"""Update a user effect preset. Cannot modify built-in presets.""" """Update a user effect preset. Cannot modify built-in presets."""
preset = db.query(DBEffectPreset).filter_by(id=preset_id).first() preset = db.query(DBEffectPreset).filter_by(id=preset_id).first()
if not preset: if not preset:
+16 -11
View File
@@ -5,18 +5,22 @@ Handles exporting profiles to ZIP archives and importing them back.
Also handles exporting individual generations. Also handles exporting individual generations.
""" """
import io
import json import json
import zipfile import zipfile
import io
from pathlib import Path from pathlib import Path
from typing import Optional
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from ..models import VoiceProfileResponse
from ..database import VoiceProfile as DBVoiceProfile, ProfileSample as DBProfileSample, Generation as DBGeneration, GenerationVersion as DBGenerationVersion
from .profiles import create_profile, add_profile_sample
from ..models import VoiceProfileCreate
from .. import config from .. import config
from ..database import (
Generation as DBGeneration,
GenerationVersion as DBGenerationVersion,
ProfileSample as DBProfileSample,
VoiceProfile as DBVoiceProfile,
)
from ..models import VoiceProfileCreate, VoiceProfileResponse
from .profiles import add_profile_sample, create_profile
def _get_unique_profile_name(name: str, db: Session) -> str: def _get_unique_profile_name(name: str, db: Session) -> str:
@@ -197,7 +201,7 @@ async def import_profile_from_zip(file_bytes: bytes, db: Session) -> VoiceProfil
await upload_avatar(profile.id, tmp_path, db) await upload_avatar(profile.id, tmp_path, db)
finally: finally:
Path(tmp_path).unlink(missing_ok=True) Path(tmp_path).unlink(missing_ok=True)
except Exception as e: except Exception:
# Avatar import is optional - continue even if it fails # Avatar import is optional - continue even if it fails
pass pass
@@ -239,7 +243,7 @@ async def import_profile_from_zip(file_bytes: bytes, db: Session) -> VoiceProfil
except Exception as e: except Exception as e:
if isinstance(e, ValueError): if isinstance(e, ValueError):
raise raise
raise ValueError(f"Error importing profile: {str(e)}") raise ValueError(f"Error importing profile: {e!s}")
def export_generation_to_zip(generation_id: str, db: Session) -> bytes: def export_generation_to_zip(generation_id: str, db: Session) -> bytes:
@@ -344,10 +348,11 @@ async def import_generation_from_zip(file_bytes: bytes, db: Session) -> dict:
Raises: Raises:
ValueError: If ZIP is invalid or missing required files ValueError: If ZIP is invalid or missing required files
""" """
from pathlib import Path
import tempfile
import shutil import shutil
import tempfile
from datetime import datetime from datetime import datetime
from pathlib import Path
from .. import config from .. import config
zip_buffer = io.BytesIO(file_bytes) zip_buffer = io.BytesIO(file_bytes)
@@ -458,4 +463,4 @@ async def import_generation_from_zip(file_bytes: bytes, db: Session) -> dict:
except Exception as e: except Exception as e:
if isinstance(e, ValueError): if isinstance(e, ValueError):
raise raise
raise ValueError(f"Error importing generation: {str(e)}") raise ValueError(f"Error importing generation: {e!s}")
+20 -20
View File
@@ -18,12 +18,12 @@ from __future__ import annotations
import asyncio import asyncio
import traceback import traceback
from typing import Literal, Optional from typing import Literal
from .. import config from .. import config
from . import history, profiles
from ..database import get_db from ..database import get_db
from ..utils.tasks import get_task_manager from ..utils.tasks import get_task_manager
from . import history, profiles
async def run_generation( async def run_generation(
@@ -34,23 +34,23 @@ async def run_generation(
language: str, language: str,
engine: str, engine: str,
model_size: str, model_size: str,
seed: Optional[int], seed: int | None,
normalize: bool = False, normalize: bool = False,
effects_chain: Optional[list] = None, effects_chain: list | None = None,
instruct: Optional[str] = None, instruct: str | None = None,
mode: Literal["generate", "retry", "regenerate"], mode: Literal["generate", "retry", "regenerate"],
max_chunk_chars: Optional[int] = None, max_chunk_chars: int | None = None,
crossfade_ms: Optional[int] = None, crossfade_ms: int | None = None,
version_id: Optional[str] = None, version_id: str | None = None,
) -> None: ) -> None:
"""Execute TTS inference and persist the result. """Execute TTS inference and persist the result.
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, get_tts_backend_for_engine, load_engine_model
from ..utils.chunked_tts import generate_chunked
from ..utils.audio import normalize_audio, save_audio, trim_tts_output from ..utils.audio import normalize_audio, save_audio, trim_tts_output
from ..utils.chunked_tts import generate_chunked
task_manager = get_task_manager() task_manager = get_task_manager()
bg_db = next(get_db()) bg_db = next(get_db())
@@ -170,7 +170,7 @@ def _save_generate(
generation_id: str, generation_id: str,
audio, audio,
sample_rate: int, sample_rate: int,
effects_chain: Optional[list], effects_chain: list | None,
save_audio, save_audio,
db, db,
) -> str: ) -> str:
@@ -249,11 +249,11 @@ async def generate_audio_sync(
language: str, language: str,
engine: str, engine: str,
model_size: str, model_size: str,
seed: Optional[int] = None, seed: int | None = None,
instruct: Optional[str] = None, instruct: str | None = None,
normalize: bool = True, normalize: bool = True,
max_chunk_chars: Optional[int] = None, max_chunk_chars: int | None = None,
crossfade_ms: Optional[int] = None, crossfade_ms: int | None = None,
) -> bytes: ) -> bytes:
"""Run a TTS generation synchronously and return the resulting wav bytes. """Run a TTS generation synchronously and return the resulting wav bytes.
@@ -267,9 +267,9 @@ 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, get_tts_backend_for_engine, load_engine_model
from ..utils.chunked_tts import generate_chunked
from ..utils.audio import normalize_audio, trim_tts_output from ..utils.audio import normalize_audio, trim_tts_output
from ..utils.chunked_tts import generate_chunked
from . import tts from . import tts
bg_db = next(get_db()) bg_db = next(get_db())
@@ -312,7 +312,7 @@ async def generate_audio_sync(
def _save_regenerate( def _save_regenerate(
*, *,
generation_id: str, generation_id: str,
version_id: Optional[str], version_id: str | None,
audio, audio,
sample_rate: int, sample_rate: int,
save_audio, save_audio,
@@ -322,10 +322,10 @@ def _save_regenerate(
Returns the audio path. Returns the audio path.
""" """
from . import versions as versions_mod
import uuid as _uuid import uuid as _uuid
from . import versions as versions_mod
suffix = _uuid.uuid4().hex[:8] suffix = _uuid.uuid4().hex[:8]
audio_path = config.get_generations_dir() / f"{generation_id}_{suffix}.wav" audio_path = config.get_generations_dir() / f"{generation_id}_{suffix}.wav"
save_audio(audio, str(audio_path), sample_rate) save_audio(audio, str(audio_path), sample_rate)
+26 -18
View File
@@ -2,17 +2,25 @@
Generation history management module. Generation history management module.
""" """
from typing import List, Optional, Tuple
from datetime import datetime
import uuid import uuid
import shutil from datetime import datetime
from pathlib import Path
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from sqlalchemy import or_
from ..models import GenerationRequest, GenerationResponse, HistoryQuery, HistoryResponse, HistoryListResponse, GenerationVersionResponse, EffectConfig
from ..database import Generation as DBGeneration, GenerationVersion as DBGenerationVersion, VoiceProfile as DBVoiceProfile
from .. import config from .. import config
from ..database import (
Generation as DBGeneration,
GenerationVersion as DBGenerationVersion,
VoiceProfile as DBVoiceProfile,
)
from ..models import (
EffectConfig,
GenerationResponse,
GenerationVersionResponse,
HistoryListResponse,
HistoryQuery,
HistoryResponse,
)
def _get_versions_for_generation(generation_id: str, db: Session) -> tuple: def _get_versions_for_generation(generation_id: str, db: Session) -> tuple:
@@ -58,13 +66,13 @@ async def create_generation(
language: str, language: str,
audio_path: str, audio_path: str,
duration: float, duration: float,
seed: Optional[int], seed: int | None,
db: Session, db: Session,
instruct: Optional[str] = None, instruct: str | None = None,
generation_id: Optional[str] = None, generation_id: str | None = None,
status: str = "completed", status: str = "completed",
engine: Optional[str] = "qwen", engine: str | None = "qwen",
model_size: Optional[str] = None, model_size: str | None = None,
source: str = "manual", source: str = "manual",
) -> GenerationResponse: ) -> GenerationResponse:
""" """
@@ -118,10 +126,10 @@ async def update_generation_status(
generation_id: str, generation_id: str,
status: str, status: str,
db: Session, db: Session,
audio_path: Optional[str] = None, audio_path: str | None = None,
duration: Optional[float] = None, duration: float | None = None,
error: Optional[str] = None, error: str | None = None,
) -> Optional[GenerationResponse]: ) -> GenerationResponse | None:
"""Update the status of a generation (used by async generation flow).""" """Update the status of a generation (used by async generation flow)."""
generation = db.query(DBGeneration).filter_by(id=generation_id).first() generation = db.query(DBGeneration).filter_by(id=generation_id).first()
if not generation: if not generation:
@@ -143,7 +151,7 @@ async def update_generation_status(
async def get_generation( async def get_generation(
generation_id: str, generation_id: str,
db: Session, db: Session,
) -> Optional[GenerationResponse]: ) -> GenerationResponse | None:
""" """
Get a generation by ID. Get a generation by ID.
-1
View File
@@ -24,7 +24,6 @@ from dataclasses import dataclass
from . import llm as llm_service from . import llm as llm_service
from .refinement import collapse_repetitive_artifacts from .refinement import collapse_repetitive_artifacts
# Shared rules block embedded in every mode-specific system prompt. Kept # Shared rules block embedded in every mode-specific system prompt. Kept
# short because small LLMs (0.6B) degrade when the system prompt is long, # short because small LLMs (0.6B) degrade when the system prompt is long,
# and because the per-mode instructions downstream carry the specifics. # and because the per-mode instructions downstream carry the specifics.
-1
View File
@@ -5,7 +5,6 @@ import logging
import shutil import shutil
import uuid import uuid
from datetime import datetime from datetime import datetime
from pathlib import Path
from sqlalchemy import func from sqlalchemy import func
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
-1
View File
@@ -13,7 +13,6 @@ from dataclasses import dataclass
from . import llm as llm_service from . import llm as llm_service
# A run that repeats this many times gets collapsed before the LLM sees # A run that repeats this many times gets collapsed before the LLM sees
# the transcript. Whisper occasionally loops content hundreds of times # the transcript. Whisper occasionally loops content hundreds of times
# when audio trails off — "URL URL URL…" (single word), "thanks for # when audio trails off — "URL URL URL…" (single word), "thanks for
+8 -9
View File
@@ -20,11 +20,10 @@ import shutil
import sys import sys
import tarfile import tarfile
from pathlib import Path from pathlib import Path
from typing import Optional
from .. import __version__
from ..config import get_data_dir from ..config import get_data_dir
from ..utils.progress import get_progress_manager from ..utils.progress import get_progress_manager
from .. import __version__
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -64,7 +63,7 @@ def get_rocm_exe_name() -> str:
return "voicebox-server-rocm" return "voicebox-server-rocm"
def get_rocm_binary_path() -> Optional[Path]: def get_rocm_binary_path() -> Path | None:
"""Return path to the ROCm executable if it exists inside the onedir.""" """Return path to the ROCm executable if it exists inside the onedir."""
p = get_rocm_dir() / get_rocm_exe_name() p = get_rocm_dir() / get_rocm_exe_name()
if p.exists(): if p.exists():
@@ -77,7 +76,7 @@ def get_rocm_libs_manifest_path() -> Path:
return get_rocm_dir() / "rocm-libs.json" return get_rocm_dir() / "rocm-libs.json"
def get_installed_rocm_libs_version() -> Optional[str]: def get_installed_rocm_libs_version() -> str | None:
"""Read the installed ROCm libs version from rocm-libs.json, or None.""" """Read the installed ROCm libs version from rocm-libs.json, or None."""
manifest_path = get_rocm_libs_manifest_path() manifest_path = get_rocm_libs_manifest_path()
if not manifest_path.exists(): if not manifest_path.exists():
@@ -115,7 +114,7 @@ def get_rocm_status() -> dict:
} }
def _needs_server_download(version: Optional[str] = None) -> bool: def _needs_server_download(version: str | None = None) -> bool:
"""Check if the server core archive needs to be (re)downloaded.""" """Check if the server core archive needs to be (re)downloaded."""
rocm_path = get_rocm_binary_path() rocm_path = get_rocm_binary_path()
if not rocm_path: if not rocm_path:
@@ -139,7 +138,7 @@ def _needs_rocm_libs_download() -> bool:
async def _download_and_extract_archive( async def _download_and_extract_archive(
client, client,
url: str, url: str,
sha256_url: Optional[str], sha256_url: str | None,
dest_dir: Path, dest_dir: Path,
label: str, label: str,
progress_offset: int, progress_offset: int,
@@ -233,7 +232,7 @@ async def _download_and_extract_archive(
return downloaded return downloaded
async def download_rocm_binary(version: Optional[str] = None): async def download_rocm_binary(version: str | None = None):
"""Download the ROCm backend (server core + ROCm libs if needed). """Download the ROCm backend (server core + ROCm libs if needed).
Downloads both archives from GitHub Releases, extracts them into Downloads both archives from GitHub Releases, extracts them into
@@ -253,7 +252,7 @@ async def download_rocm_binary(version: Optional[str] = None):
await _download_rocm_binary_locked(version) await _download_rocm_binary_locked(version)
async def _download_rocm_binary_locked(version: Optional[str] = None): async def _download_rocm_binary_locked(version: str | None = None):
"""Inner implementation of download_rocm_binary, called under _download_lock.""" """Inner implementation of download_rocm_binary, called under _download_lock."""
import httpx import httpx
@@ -394,7 +393,7 @@ async def _download_rocm_binary_locked(version: Optional[str] = None):
raise raise
def get_rocm_binary_version() -> Optional[str]: def get_rocm_binary_version() -> str | None:
"""Get the version of the installed ROCm binary, or None if not installed.""" """Get the version of the installed ROCm binary, or None if not installed."""
import subprocess import subprocess
+1 -3
View File
@@ -11,14 +11,12 @@ from typing import Any
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from ..database import CaptureSettings as DBCaptureSettings from ..database import CaptureSettings as DBCaptureSettings, GenerationSettings as DBGenerationSettings
from ..database import GenerationSettings as DBGenerationSettings
from ..utils.capture_chords import ( from ..utils.capture_chords import (
default_push_to_talk_chord, default_push_to_talk_chord,
default_toggle_to_talk_chord, default_toggle_to_talk_chord,
) )
SINGLETON_ID = 1 SINGLETON_ID = 1
+33 -33
View File
@@ -2,37 +2,37 @@
Story management module. Story management module.
""" """
from typing import List, Optional
from datetime import datetime
import uuid
import tempfile import tempfile
import uuid
from datetime import datetime
from pathlib import Path from pathlib import Path
from sqlalchemy.orm import Session
import numpy as np
from sqlalchemy import func from sqlalchemy import func
from sqlalchemy.orm import Session
from .. import config from .. import config
from ..models import (
StoryCreate,
StoryResponse,
StoryDetailResponse,
StoryItemDetail,
StoryItemCreate,
StoryItemBatchUpdate,
StoryItemMove,
StoryItemTrim,
StoryItemVolumeUpdate,
StoryItemSplit,
StoryItemVersionUpdate,
)
from ..database import ( from ..database import (
Generation as DBGeneration,
Story as DBStory, Story as DBStory,
StoryItem as DBStoryItem, StoryItem as DBStoryItem,
Generation as DBGeneration,
VoiceProfile as DBVoiceProfile, VoiceProfile as DBVoiceProfile,
) )
from .history import _get_versions_for_generation from ..models import (
StoryCreate,
StoryDetailResponse,
StoryItemBatchUpdate,
StoryItemCreate,
StoryItemDetail,
StoryItemMove,
StoryItemSplit,
StoryItemTrim,
StoryItemVersionUpdate,
StoryItemVolumeUpdate,
StoryResponse,
)
from ..utils.audio import load_audio, save_audio from ..utils.audio import load_audio, save_audio
import numpy as np from .history import _get_versions_for_generation
def _build_item_detail( def _build_item_detail(
@@ -113,7 +113,7 @@ async def create_story(
async def list_stories( async def list_stories(
db: Session, db: Session,
) -> List[StoryResponse]: ) -> list[StoryResponse]:
""" """
List all stories. List all stories.
@@ -139,7 +139,7 @@ async def list_stories(
async def get_story( async def get_story(
story_id: str, story_id: str,
db: Session, db: Session,
) -> Optional[StoryDetailResponse]: ) -> StoryDetailResponse | None:
""" """
Get a story with all its items. Get a story with all its items.
@@ -176,7 +176,7 @@ async def update_story(
story_id: str, story_id: str,
data: StoryCreate, data: StoryCreate,
db: Session, db: Session,
) -> Optional[StoryResponse]: ) -> StoryResponse | None:
""" """
Update a story. Update a story.
@@ -238,7 +238,7 @@ async def add_item_to_story(
story_id: str, story_id: str,
data: StoryItemCreate, data: StoryItemCreate,
db: Session, db: Session,
) -> Optional[StoryItemDetail]: ) -> StoryItemDetail | None:
""" """
Add a generation to a story. Add a generation to a story.
@@ -324,7 +324,7 @@ async def move_story_item(
item_id: str, item_id: str,
data: StoryItemMove, data: StoryItemMove,
db: Session, db: Session,
) -> Optional[StoryItemDetail]: ) -> StoryItemDetail | None:
""" """
Move a story item (update position and/or track). Move a story item (update position and/or track).
@@ -416,7 +416,7 @@ async def trim_story_item(
item_id: str, item_id: str,
data: StoryItemTrim, data: StoryItemTrim,
db: Session, db: Session,
) -> Optional[StoryItemDetail]: ) -> StoryItemDetail | None:
""" """
Trim a story item (update trim_start_ms and trim_end_ms). Trim a story item (update trim_start_ms and trim_end_ms).
@@ -474,7 +474,7 @@ async def update_story_item_volume(
item_id: str, item_id: str,
data: StoryItemVolumeUpdate, data: StoryItemVolumeUpdate,
db: Session, db: Session,
) -> Optional[StoryItemDetail]: ) -> StoryItemDetail | None:
"""Update a story item's playback volume (per-clip linear gain).""" """Update a story item's playback volume (per-clip linear gain)."""
item = ( item = (
db.query(DBStoryItem) db.query(DBStoryItem)
@@ -505,7 +505,7 @@ async def split_story_item(
item_id: str, item_id: str,
data: StoryItemSplit, data: StoryItemSplit,
db: Session, db: Session,
) -> Optional[List[StoryItemDetail]]: ) -> list[StoryItemDetail] | None:
""" """
Split a story item at a given time, creating two clips. Split a story item at a given time, creating two clips.
@@ -592,7 +592,7 @@ async def duplicate_story_item(
story_id: str, story_id: str,
item_id: str, item_id: str,
db: Session, db: Session,
) -> Optional[StoryItemDetail]: ) -> StoryItemDetail | None:
""" """
Duplicate a story item, creating a copy with all properties. Duplicate a story item, creating a copy with all properties.
@@ -696,10 +696,10 @@ async def update_story_item_times(
async def reorder_story_items( async def reorder_story_items(
story_id: str, story_id: str,
generation_ids: List[str], generation_ids: list[str],
db: Session, db: Session,
gap_ms: int = 200, gap_ms: int = 200,
) -> Optional[List[StoryItemDetail]]: ) -> list[StoryItemDetail] | None:
""" """
Reorder story items and recalculate timecodes. Reorder story items and recalculate timecodes.
@@ -763,7 +763,7 @@ async def set_story_item_version(
item_id: str, item_id: str,
data: StoryItemVersionUpdate, data: StoryItemVersionUpdate,
db: Session, db: Session,
) -> Optional[StoryItemDetail]: ) -> StoryItemDetail | None:
""" """
Pin a story item to a specific generation version. Pin a story item to a specific generation version.
@@ -824,7 +824,7 @@ async def set_story_item_version(
async def export_story_audio( async def export_story_audio(
story_id: str, story_id: str,
db: Session, db: Session,
) -> Optional[bytes]: ) -> bytes | None:
""" """
Export story as single mixed audio file with timecode-based mixing. Export story as single mixed audio file with timecode-based mixing.
+2 -1
View File
@@ -5,8 +5,9 @@ to avoid GPU contention.
import asyncio import asyncio
import traceback import traceback
from collections.abc import Coroutine
from dataclasses import dataclass from dataclasses import dataclass
from typing import Coroutine, Literal from typing import Literal
# Keep references to fire-and-forget background tasks to prevent GC # Keep references to fire-and-forget background tasks to prevent GC
_background_tasks: set = set() _background_tasks: set = set()
+11 -13
View File
@@ -9,17 +9,15 @@ from __future__ import annotations
import json import json
import uuid import uuid
from pathlib import Path
from typing import List, Optional
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from ..database import (
GenerationVersion as DBGenerationVersion,
Generation as DBGeneration,
)
from ..models import GenerationVersionResponse, EffectConfig
from .. import config from .. import config
from ..database import (
Generation as DBGeneration,
GenerationVersion as DBGenerationVersion,
)
from ..models import EffectConfig, GenerationVersionResponse
def _version_response(v: DBGenerationVersion) -> GenerationVersionResponse: def _version_response(v: DBGenerationVersion) -> GenerationVersionResponse:
@@ -40,7 +38,7 @@ def _version_response(v: DBGenerationVersion) -> GenerationVersionResponse:
) )
def list_versions(generation_id: str, db: Session) -> List[GenerationVersionResponse]: def list_versions(generation_id: str, db: Session) -> list[GenerationVersionResponse]:
"""List all versions for a generation.""" """List all versions for a generation."""
versions = ( versions = (
db.query(DBGenerationVersion) db.query(DBGenerationVersion)
@@ -51,7 +49,7 @@ def list_versions(generation_id: str, db: Session) -> List[GenerationVersionResp
return [_version_response(v) for v in versions] return [_version_response(v) for v in versions]
def get_version(version_id: str, db: Session) -> Optional[GenerationVersionResponse]: def get_version(version_id: str, db: Session) -> GenerationVersionResponse | None:
"""Get a specific version by ID.""" """Get a specific version by ID."""
v = db.query(DBGenerationVersion).filter_by(id=version_id).first() v = db.query(DBGenerationVersion).filter_by(id=version_id).first()
if not v: if not v:
@@ -59,7 +57,7 @@ def get_version(version_id: str, db: Session) -> Optional[GenerationVersionRespo
return _version_response(v) return _version_response(v)
def get_default_version(generation_id: str, db: Session) -> Optional[GenerationVersionResponse]: def get_default_version(generation_id: str, db: Session) -> GenerationVersionResponse | None:
"""Get the default version for a generation.""" """Get the default version for a generation."""
v = ( v = (
db.query(DBGenerationVersion) db.query(DBGenerationVersion)
@@ -84,9 +82,9 @@ def create_version(
label: str, label: str,
audio_path: str, audio_path: str,
db: Session, db: Session,
effects_chain: Optional[List[dict]] = None, effects_chain: list[dict] | None = None,
is_default: bool = False, is_default: bool = False,
source_version_id: Optional[str] = None, source_version_id: str | None = None,
) -> GenerationVersionResponse: ) -> GenerationVersionResponse:
"""Create a new version for a generation. """Create a new version for a generation.
@@ -119,7 +117,7 @@ def create_version(
return _version_response(version) return _version_response(version)
def set_default_version(version_id: str, db: Session) -> Optional[GenerationVersionResponse]: def set_default_version(version_id: str, db: Session) -> GenerationVersionResponse | None:
"""Set a version as the default for its generation.""" """Set a version as the default for its generation."""
version = db.query(DBGenerationVersion).filter_by(id=version_id).first() version = db.query(DBGenerationVersion).filter_by(id=version_id).first()
if not version: if not version:
+17
View File
@@ -0,0 +1,17 @@
"""Shared test setup.
The suite mixes flat imports (``from utils.progress import ...``) with
package imports (``from backend import config``). Both the repo root and
the backend directory go on ``sys.path`` here so every test file collects
on its own, regardless of which file loads first.
"""
import sys
from pathlib import Path
BACKEND_DIR = Path(__file__).resolve().parent.parent
REPO_ROOT = BACKEND_DIR.parent
for _path in (str(BACKEND_DIR), str(REPO_ROOT)):
if _path not in sys.path:
sys.path.insert(0, _path)
+25 -27
View File
@@ -25,14 +25,12 @@ import tempfile
import threading import threading
import time import time
from collections import deque from collections import deque
from dataclasses import asdict, dataclass, field from dataclasses import asdict, dataclass
from datetime import datetime, timezone from datetime import UTC, datetime
from pathlib import Path from pathlib import Path
from typing import Optional
import httpx import httpx
REPO_ROOT = Path(__file__).resolve().parents[2] REPO_ROOT = Path(__file__).resolve().parents[2]
BACKEND_DIR = REPO_ROOT / "backend" BACKEND_DIR = REPO_ROOT / "backend"
DIST_DIR = BACKEND_DIR / "dist" DIST_DIR = BACKEND_DIR / "dist"
@@ -46,7 +44,7 @@ RESULTS_DIR = Path(__file__).resolve().parent / "results"
class MatrixRow: class MatrixRow:
label: str # human-readable (appears in report) label: str # human-readable (appears in report)
engine: str # /generate engine engine: str # /generate engine
model_size: Optional[str] # /generate model_size (None = omit) model_size: str | None # /generate model_size (None = omit)
profile_kind: str # "cloned" | "preset_kokoro" | "preset_qwen_cv" profile_kind: str # "cloned" | "preset_kokoro" | "preset_qwen_cv"
model_name: str # /models/status key for cache lookup model_name: str # /models/status key for cache lookup
@@ -76,22 +74,22 @@ HEALTH_TIMEOUT = 120
class ModelResult: class ModelResult:
label: str label: str
engine: str engine: str
model_size: Optional[str] model_size: str | None
status: str # "passed" | "failed" | "timeout" status: str # "passed" | "failed" | "timeout"
was_cached: Optional[bool] = None was_cached: bool | None = None
generation_id: Optional[str] = None generation_id: str | None = None
elapsed_seconds: float = 0.0 elapsed_seconds: float = 0.0
audio_duration: Optional[float] = None audio_duration: float | None = None
audio_path: Optional[str] = None audio_path: str | None = None
audio_bytes: Optional[int] = None audio_bytes: int | None = None
error: Optional[str] = None error: str | None = None
http_status: Optional[int] = None http_status: int | None = None
server_log_tail: Optional[list[str]] = None server_log_tail: list[str] | None = None
# ── Binary resolution ──────────────────────────────────────────────── # ── Binary resolution ────────────────────────────────────────────────
def find_binary() -> Optional[Path]: def find_binary() -> Path | None:
"""Return the first existing binary in priority order, or None.""" """Return the first existing binary in priority order, or None."""
is_win = platform.system() == "Windows" is_win = platform.system() == "Windows"
exe = ".exe" if is_win else "" exe = ".exe" if is_win else ""
@@ -129,9 +127,9 @@ class ServerProcess:
self.port = port self.port = port
self.data_dir = data_dir self.data_dir = data_dir
self.log_path = log_path self.log_path = log_path
self.proc: Optional[subprocess.Popen] = None self.proc: subprocess.Popen | None = None
self._log_buffer: deque[str] = deque(maxlen=500) self._log_buffer: deque[str] = deque(maxlen=500)
self._reader_thread: Optional[threading.Thread] = None self._reader_thread: threading.Thread | None = None
def start(self) -> None: def start(self) -> None:
args = [ args = [
@@ -227,7 +225,7 @@ def wait_for_health(base_url: str, server: ServerProcess, timeout: int) -> None:
raise TimeoutError(f"Server did not become healthy within {timeout}s") raise TimeoutError(f"Server did not become healthy within {timeout}s")
def get_model_cached(client: httpx.Client, base_url: str, model_name: str) -> Optional[bool]: def get_model_cached(client: httpx.Client, base_url: str, model_name: str) -> bool | None:
try: try:
r = client.get(f"{base_url}/models/status", timeout=30.0) r = client.get(f"{base_url}/models/status", timeout=30.0)
r.raise_for_status() r.raise_for_status()
@@ -331,7 +329,7 @@ def run_one_generation(
def fetch_audio_info( def fetch_audio_info(
client: httpx.Client, base_url: str, generation_id: str, data_dir: Path client: httpx.Client, base_url: str, generation_id: str, data_dir: Path
) -> tuple[Optional[str], Optional[int]]: ) -> tuple[str | None, int | None]:
"""Return (audio_path, audio_bytes) for a completed generation. """Return (audio_path, audio_bytes) for a completed generation.
Server stores audio_path relative to data_dir; resolve it to get a size. Server stores audio_path relative to data_dir; resolve it to get a size.
@@ -504,8 +502,8 @@ def main() -> int:
# Reference audio (only required if any cloning row is in the matrix) # Reference audio (only required if any cloning row is in the matrix)
needs_reference = any(r.profile_kind == "cloned" for r in rows) needs_reference = any(r.profile_kind == "cloned" for r in rows)
ref_wav: Optional[Path] = None ref_wav: Path | None = None
ref_text: Optional[str] = None ref_text: str | None = None
if needs_reference: if needs_reference:
try: try:
ref_wav, ref_text = resolve_reference(args) ref_wav, ref_text = resolve_reference(args)
@@ -518,14 +516,14 @@ def main() -> int:
# Tempdir + log path # Tempdir + log path
data_dir = Path(tempfile.mkdtemp(prefix="voicebox-e2e-")) data_dir = Path(tempfile.mkdtemp(prefix="voicebox-e2e-"))
args.output_dir.mkdir(parents=True, exist_ok=True) args.output_dir.mkdir(parents=True, exist_ok=True)
ts = datetime.now(timezone.utc).strftime("%Y%m%d-%H%M%S") ts = datetime.now(UTC).strftime("%Y%m%d-%H%M%S")
log_path = args.output_dir / f"server-{ts}.log" log_path = args.output_dir / f"server-{ts}.log"
port = args.port or pick_free_port() port = args.port or pick_free_port()
base_url = f"http://127.0.0.1:{port}" base_url = f"http://127.0.0.1:{port}"
server = ServerProcess(binary=binary, port=port, data_dir=data_dir, log_path=log_path) server = ServerProcess(binary=binary, port=port, data_dir=data_dir, log_path=log_path)
started_at = datetime.now(timezone.utc) started_at = datetime.now(UTC)
results: list[ModelResult] = [] results: list[ModelResult] = []
try: try:
@@ -536,9 +534,9 @@ def main() -> int:
with httpx.Client(timeout=30.0) as client: with httpx.Client(timeout=30.0) as client:
# Profile setup (only create what's needed) # Profile setup (only create what's needed)
cloned_profile_id: Optional[str] = None cloned_profile_id: str | None = None
kokoro_profile_id: Optional[str] = None kokoro_profile_id: str | None = None
qwen_cv_profile_id: Optional[str] = None qwen_cv_profile_id: str | None = None
needed_kinds = {r.profile_kind for r in rows} needed_kinds = {r.profile_kind for r in rows}
if "cloned" in needed_kinds: if "cloned" in needed_kinds:
assert ref_wav is not None and ref_text is not None assert ref_wav is not None and ref_text is not None
@@ -608,7 +606,7 @@ def main() -> int:
+ (f" ({result.error})" if result.error else ""), flush=True) + (f" ({result.error})" if result.error else ""), flush=True)
results.append(result) results.append(result)
finally: finally:
finished_at = datetime.now(timezone.utc) finished_at = datetime.now(UTC)
server.stop() server.stop()
if not args.keep_data_dir: if not args.keep_data_dir:
shutil.rmtree(data_dir, ignore_errors=True) shutil.rmtree(data_dir, ignore_errors=True)
+1 -2
View File
@@ -14,12 +14,11 @@ import soundfile as sf
sys.path.insert(0, str(Path(__file__).parent.parent)) sys.path.insert(0, str(Path(__file__).parent.parent))
from utils.audio import ( # noqa: E402 from utils.audio import (
preprocess_reference_audio, preprocess_reference_audio,
validate_and_load_reference_audio, validate_and_load_reference_audio,
) )
SR = 24000 SR = 24000
+26 -58
View File
@@ -4,64 +4,34 @@ Tests for CORS origin restrictions.
Validates that the CORS middleware only allows known local origins Validates that the CORS middleware only allows known local origins
and respects the VOICEBOX_CORS_ORIGINS environment variable. and respects the VOICEBOX_CORS_ORIGINS environment variable.
Uses a minimal FastAPI app that mirrors the exact CORS configuration Builds the app via the real ``backend.app.create_app`` factory so the
from backend/main.py, so tests run without heavy ML dependencies. tests exercise the actual CORS configuration rather than a copy of it.
Usage:
pip install httpx pytest fastapi starlette
python -m pytest backend/tests/test_cors.py -v
""" """
import os
import pytest import pytest
from unittest.mock import patch
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from starlette.testclient import TestClient from starlette.testclient import TestClient
from backend.app import create_app
def _build_app(env_origins: str = "") -> FastAPI:
"""
Build a minimal FastAPI app with the same CORS logic as backend/main.py.
This mirrors the exact code in main.py so the test validates the real
configuration without needing torch/numpy/transformers installed.
"""
app = FastAPI()
_default_origins = [
"http://localhost:5173",
"http://127.0.0.1:5173",
"http://localhost:17493",
"http://127.0.0.1:17493",
"tauri://localhost",
"https://tauri.localhost",
]
_cors_origins = _default_origins + [o.strip() for o in env_origins.split(",") if o.strip()]
app.add_middleware(
CORSMiddleware,
allow_origins=_cors_origins,
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
@app.get("/health")
async def health():
return {"status": "ok"}
return app
@pytest.fixture() def _build_client(monkeypatch, env_origins: str | None = None) -> TestClient:
def client(): if env_origins is None:
return TestClient(_build_app()) monkeypatch.delenv("VOICEBOX_CORS_ORIGINS", raising=False)
else:
monkeypatch.setenv("VOICEBOX_CORS_ORIGINS", env_origins)
# Plain TestClient (no context manager) skips lifespan startup, so no
# model scans or queue workers run — only middleware is exercised.
return TestClient(create_app())
@pytest.fixture() @pytest.fixture
def client_with_custom_origins(): def client(monkeypatch):
return TestClient(_build_app("https://custom.example.com,https://other.example.com")) return _build_client(monkeypatch)
@pytest.fixture
def client_with_custom_origins(monkeypatch):
return _build_client(monkeypatch, "https://custom.example.com,https://other.example.com")
def _get_with_origin(client: TestClient, origin: str) -> dict: def _get_with_origin(client: TestClient, origin: str) -> dict:
@@ -92,6 +62,7 @@ class TestCORSDefaultOrigins:
"http://127.0.0.1:17493", "http://127.0.0.1:17493",
"tauri://localhost", "tauri://localhost",
"https://tauri.localhost", "https://tauri.localhost",
"http://tauri.localhost",
]) ])
def test_allowed_origins(self, client, origin): def test_allowed_origins(self, client, origin):
headers = _get_with_origin(client, origin) headers = _get_with_origin(client, origin)
@@ -143,20 +114,17 @@ class TestCORSCustomOrigins:
class TestCORSEnvVarParsing: class TestCORSEnvVarParsing:
"""Edge cases for VOICEBOX_CORS_ORIGINS parsing.""" """Edge cases for VOICEBOX_CORS_ORIGINS parsing."""
def test_empty_env_var(self): def test_empty_env_var(self, monkeypatch):
app = _build_app("") client = _build_client(monkeypatch, "")
client = TestClient(app)
headers = _get_with_origin(client, "http://evil.com") headers = _get_with_origin(client, "http://evil.com")
assert "access-control-allow-origin" not in headers assert "access-control-allow-origin" not in headers
def test_whitespace_trimmed(self): def test_whitespace_trimmed(self, monkeypatch):
app = _build_app(" https://spaced.example.com ") client = _build_client(monkeypatch, " https://spaced.example.com ")
client = TestClient(app)
headers = _get_with_origin(client, "https://spaced.example.com") headers = _get_with_origin(client, "https://spaced.example.com")
assert headers.get("access-control-allow-origin") == "https://spaced.example.com" assert headers.get("access-control-allow-origin") == "https://spaced.example.com"
def test_trailing_comma_ignored(self): def test_trailing_comma_ignored(self, monkeypatch):
app = _build_app("https://one.example.com,") client = _build_client(monkeypatch, "https://one.example.com,")
client = TestClient(app)
headers = _get_with_origin(client, "https://one.example.com") headers = _get_with_origin(client, "https://one.example.com")
assert headers.get("access-control-allow-origin") == "https://one.example.com" assert headers.get("access-control-allow-origin") == "https://one.example.com"
+7 -8
View File
@@ -7,14 +7,14 @@ the model is already cached.
import asyncio import asyncio
import json import json
import httpx
from typing import List, Dict, Optional
from datetime import datetime from datetime import datetime
import httpx
async def monitor_sse_stream(model_name: str, timeout: int = 120): async def monitor_sse_stream(model_name: str, timeout: int = 120):
"""Monitor SSE stream for a model during generation.""" """Monitor SSE stream for a model during generation."""
events: List[Dict] = [] events: list[dict] = []
url = f"http://localhost:8000/models/progress/{model_name}" url = f"http://localhost:8000/models/progress/{model_name}"
print(f"[{_timestamp()}] Connecting to SSE endpoint: {url}") print(f"[{_timestamp()}] Connecting to SSE endpoint: {url}")
@@ -54,7 +54,7 @@ async def monitor_sse_stream(model_name: str, timeout: int = 120):
elif line.startswith(": heartbeat"): elif line.startswith(": heartbeat"):
print(f"[{timestamp}] ♥ heartbeat") print(f"[{timestamp}] ♥ heartbeat")
except asyncio.TimeoutError: except TimeoutError:
print(f"[{_timestamp()}] SSE monitoring timed out") print(f"[{_timestamp()}] SSE monitoring timed out")
except Exception as e: except Exception as e:
print(f"[{_timestamp()}] SSE error: {e}") print(f"[{_timestamp()}] SSE error: {e}")
@@ -91,15 +91,14 @@ async def trigger_generation(profile_id: str, text: str, model_size: str = "1.7B
print(f" Generation ID: {result.get('id')}") print(f" Generation ID: {result.get('id')}")
print(f" Duration: {result.get('duration', 0):.2f}s") print(f" Duration: {result.get('duration', 0):.2f}s")
return True, result return True, result
elif response.status_code == 202: if response.status_code == 202:
# Model is being downloaded # Model is being downloaded
result = response.json() result = response.json()
print(f"[{_timestamp()}] → Model download in progress") print(f"[{_timestamp()}] → Model download in progress")
print(f" Detail: {result}") print(f" Detail: {result}")
return False, result return False, result
else: print(f"[{_timestamp()}] ✗ Error: {response.text}")
print(f"[{_timestamp()}] ✗ Error: {response.text}") return False, None
return False, None
except Exception as e: except Exception as e:
print(f"[{_timestamp()}] ✗ Exception: {e}") print(f"[{_timestamp()}] ✗ Exception: {e}")
+4 -4
View File
@@ -21,7 +21,7 @@ import pytest
sys.path.insert(0, str(Path(__file__).parent.parent)) sys.path.insert(0, str(Path(__file__).parent.parent))
from utils.hf_offline_patch import force_offline_if_cached # noqa: E402 from utils.hf_offline_patch import force_offline_if_cached
def _hf_const(): def _hf_const():
@@ -53,7 +53,7 @@ def test_mutates_cached_transformers_constant():
def test_sets_env_variable(): def test_sets_env_variable():
original = os.environ.get("HF_HUB_OFFLINE") original = os.environ.get("HF_HUB_OFFLINE")
with force_offline_if_cached(True, "t"): with force_offline_if_cached(True, "t"):
assert "1" == os.environ.get("HF_HUB_OFFLINE") assert os.environ.get("HF_HUB_OFFLINE") == "1"
assert original == os.environ.get("HF_HUB_OFFLINE") assert original == os.environ.get("HF_HUB_OFFLINE")
@@ -88,14 +88,14 @@ def test_concurrent_threads_share_offline_window():
barrier.wait(timeout=5) barrier.wait(timeout=5)
assert fast_exited.wait(timeout=5), "fast thread did not exit" assert fast_exited.wait(timeout=5), "fast thread did not exit"
observations.append(_hf_const().HF_HUB_OFFLINE) observations.append(_hf_const().HF_HUB_OFFLINE)
except Exception as exc: # noqa: BLE001 except Exception as exc:
errors.append(exc) errors.append(exc)
def fast(): def fast():
try: try:
with force_offline_if_cached(True, "fast"): with force_offline_if_cached(True, "fast"):
barrier.wait(timeout=5) barrier.wait(timeout=5)
except Exception as exc: # noqa: BLE001 except Exception as exc:
errors.append(exc) errors.append(exc)
finally: finally:
fast_exited.set() fast_exited.set()
+3 -3
View File
@@ -18,10 +18,10 @@ import pytest
sys.path.insert(0, str(Path(__file__).parent.parent)) sys.path.insert(0, str(Path(__file__).parent.parent))
from huggingface_hub.errors import OfflineModeIsEnabled # noqa: E402 from huggingface_hub.errors import OfflineModeIsEnabled
from transformers.tokenization_utils_base import PreTrainedTokenizerBase # noqa: E402 from transformers.tokenization_utils_base import PreTrainedTokenizerBase
import utils.hf_offline_patch as hf_offline_patch # noqa: E402 import utils.hf_offline_patch as hf_offline_patch
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
+4 -6
View File
@@ -32,11 +32,9 @@ import sys
import time import time
from dataclasses import asdict, dataclass, field from dataclasses import asdict, dataclass, field
from pathlib import Path from pathlib import Path
from typing import Optional
import httpx import httpx
REPO_ROOT = Path(__file__).resolve().parents[2] REPO_ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(REPO_ROOT)) sys.path.insert(0, str(REPO_ROOT))
@@ -126,13 +124,13 @@ class Scorecard:
refined: str refined: str
latency_ms: int latency_ms: int
length_chars: int = 0 length_chars: int = 0
prompt_leak: Optional[str] = None prompt_leak: str | None = None
refusal: Optional[str] = None refusal: str | None = None
stage_directions: list[str] = field(default_factory=list) stage_directions: list[str] = field(default_factory=list)
flags: list[str] = field(default_factory=list) flags: list[str] = field(default_factory=list)
def first_match(patterns, text: str) -> Optional[str]: def first_match(patterns, text: str) -> str | None:
s = text.lstrip() s = text.lstrip()
for pat in patterns: for pat in patterns:
m = pat.search(s) m = pat.search(s)
@@ -186,7 +184,7 @@ known-shipping Kokoro voice so the throwaway profile satisfies the
preset-engine validator on creation.""" preset-engine validator on creation."""
def detect_backend_port(hint: Optional[int]) -> int: def detect_backend_port(hint: int | None) -> int:
candidates: list[int] = [] candidates: list[int] = []
if hint is not None: if hint is not None:
candidates.append(hint) candidates.append(hint)
@@ -5,20 +5,17 @@ This test suite verifies that the application correctly handles
duplicate profile names and provides user-friendly error messages. duplicate profile names and provides user-friendly error messages.
""" """
import pytest
import tempfile
import shutil import shutil
import tempfile
from pathlib import Path from pathlib import Path
import pytest
from sqlalchemy import create_engine from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker from sqlalchemy.orm import sessionmaker
# Add parent directory to path to import backend modules from backend.database import Base
import sys from backend.models import VoiceProfileCreate
sys.path.insert(0, str(Path(__file__).parent.parent)) from backend.services.profiles import create_profile, update_profile
from database import Base, VoiceProfile as DBVoiceProfile
from models import VoiceProfileCreate
from profiles import create_profile, update_profile
@pytest.fixture @pytest.fixture
+13 -13
View File
@@ -4,9 +4,8 @@ Test script to debug model download progress tracking.
import asyncio import asyncio
import json import json
import time
from typing import List, Dict
import logging import logging
import time
# Set up logging to see what's happening # Set up logging to see what's happening
logging.basicConfig( logging.basicConfig(
@@ -14,8 +13,8 @@ logging.basicConfig(
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s' format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
) )
from utils.progress import ProgressManager, get_progress_manager
from utils.hf_progress import HFProgressTracker, create_hf_progress_callback from utils.hf_progress import HFProgressTracker, create_hf_progress_callback
from utils.progress import ProgressManager, get_progress_manager
def test_progress_manager_basic(): def test_progress_manager_basic():
@@ -61,7 +60,7 @@ async def test_progress_manager_sse():
print("=" * 60) print("=" * 60)
pm = ProgressManager() pm = ProgressManager()
collected_events: List[Dict] = [] collected_events: list[dict] = []
# Simulate SSE client # Simulate SSE client
async def sse_client(): async def sse_client():
@@ -123,7 +122,7 @@ def test_hf_progress_tracker():
print("Test 3: HFProgressTracker tqdm Patching") print("Test 3: HFProgressTracker tqdm Patching")
print("=" * 60) print("=" * 60)
captured_progress: List[tuple] = [] captured_progress: list[tuple] = []
def progress_callback(downloaded: int, total: int, filename: str): def progress_callback(downloaded: int, total: int, filename: str):
"""Capture progress updates.""" """Capture progress updates."""
@@ -137,12 +136,14 @@ def test_hf_progress_tracker():
try: try:
from tqdm import tqdm from tqdm import tqdm
# Simulate downloading a file # Simulate downloading a file. The tracker only reports once the
# combined total crosses MIN_TOTAL_BYTES (1 MB), so the simulated
# file must be larger than that.
print(" Simulating download with tqdm...") print(" Simulating download with tqdm...")
total_size = 1000 total_size = 5_000_000
with tqdm(total=total_size, desc="model.bin", unit="B", unit_scale=True) as pbar: with tqdm(total=total_size, desc="model.bin", unit="B", unit_scale=True) as pbar:
for chunk in range(0, total_size, 100): for chunk in range(0, total_size, 500_000):
pbar.update(100) pbar.update(500_000)
time.sleep(0.01) time.sleep(0.01)
print(f" Captured {len(captured_progress)} progress updates") print(f" Captured {len(captured_progress)} progress updates")
@@ -170,7 +171,7 @@ async def test_full_integration():
print("=" * 60) print("=" * 60)
pm = get_progress_manager() pm = get_progress_manager()
collected_events: List[Dict] = [] collected_events: list[dict] = []
# SSE client # SSE client
async def sse_client(): async def sse_client():
@@ -244,9 +245,8 @@ async def test_full_integration():
assert collected_events[-1]["status"] == "complete", "Should end with 'complete'" assert collected_events[-1]["status"] == "complete", "Should end with 'complete'"
print("✓ Test 4 PASSED\n") print("✓ Test 4 PASSED\n")
return True return True
else: print("✗ Test 4 FAILED - No events received\n")
print("✗ Test 4 FAILED - No events received\n") return False
return False
async def main(): async def main():
+13 -14
View File
@@ -15,12 +15,12 @@ Prerequisites:
import asyncio import asyncio
import json import json
import httpx
import time import time
from typing import List, Dict, Optional
import httpx
async def monitor_sse_stream(model_name: str, timeout: int = 600) -> List[Dict]: async def monitor_sse_stream(model_name: str, timeout: int = 600) -> list[dict]:
""" """
Monitor SSE stream for a model download. Monitor SSE stream for a model download.
@@ -31,7 +31,7 @@ async def monitor_sse_stream(model_name: str, timeout: int = 600) -> List[Dict]:
Returns: Returns:
List of SSE events received List of SSE events received
""" """
events: List[Dict] = [] events: list[dict] = []
url = f"http://localhost:8000/models/progress/{model_name}" url = f"http://localhost:8000/models/progress/{model_name}"
last_progress = -1 last_progress = -1
@@ -72,7 +72,7 @@ async def monitor_sse_stream(model_name: str, timeout: int = 600) -> List[Dict]:
# Stop if complete or error # Stop if complete or error
if status in ("complete", "error"): if status in ("complete", "error"):
if status == "complete": if status == "complete":
print(f" ✅ Download complete!") print(" ✅ Download complete!")
else: else:
print(f" ❌ Download error: {data.get('error', 'unknown')}") print(f" ❌ Download error: {data.get('error', 'unknown')}")
break break
@@ -119,20 +119,19 @@ async def delete_model(model_name: str) -> bool:
async with httpx.AsyncClient(timeout=30) as client: async with httpx.AsyncClient(timeout=30) as client:
response = await client.delete(url) response = await client.delete(url)
if response.status_code == 200: if response.status_code == 200:
print(f" ✅ Model deleted") print(" ✅ Model deleted")
return True return True
elif response.status_code == 404: if response.status_code == 404:
print(f" ℹ️ Model not found (already deleted)") print(" ℹ️ Model not found (already deleted)")
return True return True
else: print(f" ⚠️ Delete response: {response.status_code} - {response.text}")
print(f" ⚠️ Delete response: {response.status_code} - {response.text}") return False
return False
except Exception as e: except Exception as e:
print(f" ❌ Error deleting model: {e}") print(f" ❌ Error deleting model: {e}")
return False return False
async def check_model_status(model_name: str) -> Optional[Dict]: async def check_model_status(model_name: str) -> dict | None:
"""Check the status of a model.""" """Check the status of a model."""
try: try:
async with httpx.AsyncClient(timeout=10) as client: async with httpx.AsyncClient(timeout=10) as client:
@@ -271,11 +270,11 @@ async def main():
first_event = events[0] first_event = events[0]
last_event = events[-1] last_event = events[-1]
print(f"\n📊 First event:") print("\n📊 First event:")
print(f" Status: {first_event.get('status')}") print(f" Status: {first_event.get('status')}")
print(f" Progress: {first_event.get('progress', 0):.1f}%") print(f" Progress: {first_event.get('progress', 0):.1f}%")
print(f"\n📊 Last event:") print("\n📊 Last event:")
print(f" Status: {last_event.get('status')}") print(f" Status: {last_event.get('status')}")
print(f" Progress: {last_event.get('progress', 0):.1f}%") print(f" Progress: {last_event.get('progress', 0):.1f}%")
@@ -10,7 +10,6 @@ character-level pass added.
from backend.services.refinement import collapse_repetitive_artifacts from backend.services.refinement import collapse_repetitive_artifacts
# ── single-word loops (word-level pass) ───────────────────────────────── # ── single-word loops (word-level pass) ─────────────────────────────────
+9 -12
View File
@@ -32,28 +32,25 @@ import re
import socket import socket
import sys import sys
import time import time
from collections.abc import Iterable
from dataclasses import asdict, dataclass, field from dataclasses import asdict, dataclass, field
from pathlib import Path from pathlib import Path
from collections.abc import Iterable
from typing import Optional
import httpx import httpx
REPO_ROOT = Path(__file__).resolve().parents[2] REPO_ROOT = Path(__file__).resolve().parents[2]
# Point sys.path at the repo root so ``backend.services.refinement`` resolves # Point sys.path at the repo root so ``backend.services.refinement`` resolves
# as a package. Using backend/ as root breaks the service's own # as a package. Using backend/ as root breaks the service's own
# ``from ..backends import …`` relative imports. # ``from ..backends import …`` relative imports.
sys.path.insert(0, str(REPO_ROOT)) sys.path.insert(0, str(REPO_ROOT))
from backend.services.refinement import ( # noqa: E402 from backend.services.refinement import (
build_refinement_prompt,
collapse_repetitive_artifacts,
REFINEMENT_EXAMPLES, REFINEMENT_EXAMPLES,
RefinementFlags, RefinementFlags,
build_refinement_prompt,
collapse_repetitive_artifacts,
) )
# ── Sample inputs ───────────────────────────────────────────────────── # ── Sample inputs ─────────────────────────────────────────────────────
@@ -222,8 +219,8 @@ class Scorecard:
filler_count_refined: int = 0 filler_count_refined: int = 0
length_ratio: float = 0.0 length_ratio: float = 0.0
has_loop_artifact: bool = False has_loop_artifact: bool = False
prompt_leak: Optional[str] = None prompt_leak: str | None = None
answer_leak: Optional[str] = None answer_leak: str | None = None
missing_substrings: list[str] = field(default_factory=list) missing_substrings: list[str] = field(default_factory=list)
missing_question_mark: bool = False missing_question_mark: bool = False
flags: list[str] = field(default_factory=list) flags: list[str] = field(default_factory=list)
@@ -242,7 +239,7 @@ def has_loop_run(text: str, threshold: int = 6) -> bool:
if len(tokens) < threshold: if len(tokens) < threshold:
return False return False
run = 1 run = 1
prev: Optional[str] = None prev: str | None = None
for tok in tokens: for tok in tokens:
key = re.sub(r"[^\w]", "", tok).lower() key = re.sub(r"[^\w]", "", tok).lower()
if key and key == prev: if key and key == prev:
@@ -255,7 +252,7 @@ def has_loop_run(text: str, threshold: int = 6) -> bool:
return False return False
def first_match(patterns: Iterable[re.Pattern[str]], text: str) -> Optional[str]: def first_match(patterns: Iterable[re.Pattern[str]], text: str) -> str | None:
stripped = text.lstrip() stripped = text.lstrip()
for pat in patterns: for pat in patterns:
m = pat.search(stripped) m = pat.search(stripped)
@@ -319,7 +316,7 @@ def score(sample: Sample, model: str, refined: str, latency_ms: int) -> Scorecar
DEFAULT_PORTS = (8000, 8765, 8899, 17493) DEFAULT_PORTS = (8000, 8765, 8899, 17493)
def detect_backend_port(hint: Optional[int]) -> int: def detect_backend_port(hint: int | None) -> int:
"""Return a port that answers /health, preferring the hint.""" """Return a port that answers /health, preferring the hint."""
candidates: list[int] = [] candidates: list[int] = []
if hint is not None: if hint is not None:
+27 -32
View File
@@ -10,8 +10,6 @@ Usage:
from unittest.mock import patch from unittest.mock import patch
import pytest
class TestCheckCudaCompatibility: class TestCheckCudaCompatibility:
"""Unit tests for check_cuda_compatibility with ROCm awareness.""" """Unit tests for check_cuda_compatibility with ROCm awareness."""
@@ -28,41 +26,38 @@ class TestCheckCudaCompatibility:
"""On ROCm, the NVIDIA compute-capability check should be skipped.""" """On ROCm, the NVIDIA compute-capability check should be skipped."""
from backend.backends.base import check_cuda_compatibility from backend.backends.base import check_cuda_compatibility
with patch("torch.cuda.is_available", return_value=True): with patch("torch.cuda.is_available", return_value=True), patch("torch.version.hip", "6.2.41133"):
with patch("torch.version.hip", "6.2.41133"): compatible, warning = check_cuda_compatibility()
compatible, warning = check_cuda_compatibility() assert compatible is True
assert compatible is True assert warning is None
assert warning is None
def test_cuda_compatible_arch(self): def test_cuda_compatible_arch(self):
from backend.backends.base import check_cuda_compatibility from backend.backends.base import check_cuda_compatibility
with patch("torch.cuda.is_available", return_value=True): with patch("torch.cuda.is_available", return_value=True), patch("torch.version.hip", None):
with patch("torch.version.hip", None): with patch("torch.cuda.get_device_capability", return_value=(8, 6)):
with patch("torch.cuda.get_device_capability", return_value=(8, 6)): with patch("torch.cuda.get_device_name", return_value="NVIDIA GeForce RTX 3060"):
with patch("torch.cuda.get_device_name", return_value="NVIDIA GeForce RTX 3060"): with patch.object(
with patch.object( __import__("torch").cuda, "_get_arch_list",
__import__("torch").cuda, "_get_arch_list", return_value=["sm_80", "sm_86", "sm_89"],
return_value=["sm_80", "sm_86", "sm_89"], create=True,
create=True, ):
): compatible, warning = check_cuda_compatibility()
compatible, warning = check_cuda_compatibility() assert compatible is True
assert compatible is True assert warning is None
assert warning is None
def test_cuda_incompatible_arch(self): def test_cuda_incompatible_arch(self):
from backend.backends.base import check_cuda_compatibility from backend.backends.base import check_cuda_compatibility
with patch("torch.cuda.is_available", return_value=True): with patch("torch.cuda.is_available", return_value=True), patch("torch.version.hip", None):
with patch("torch.version.hip", None): with patch("torch.cuda.get_device_capability", return_value=(9, 0)):
with patch("torch.cuda.get_device_capability", return_value=(9, 0)): with patch("torch.cuda.get_device_name", return_value="NVIDIA GeForce RTX 4090"):
with patch("torch.cuda.get_device_name", return_value="NVIDIA GeForce RTX 4090"): with patch.object(
with patch.object( __import__("torch").cuda, "_get_arch_list",
__import__("torch").cuda, "_get_arch_list", return_value=["sm_80", "sm_86"],
return_value=["sm_80", "sm_86"], create=True,
create=True, ):
): compatible, warning = check_cuda_compatibility()
compatible, warning = check_cuda_compatibility() assert compatible is False
assert compatible is False assert warning is not None
assert warning is not None assert "not supported" in warning
assert "not supported" in warning
+1 -1
View File
@@ -78,7 +78,7 @@ class TestRocmBuildCli:
build_server(cuda=True, rocm=True) build_server(cuda=True, rocm=True)
@pytest.mark.slow() @pytest.mark.slow
@pytest.mark.skipif(sys.platform != "win32", reason="ROCm build E2E only runs on Windows") @pytest.mark.skipif(sys.platform != "win32", reason="ROCm build E2E only runs on Windows")
class TestRocmBuildE2E: class TestRocmBuildE2E:
""" """
+1 -2
View File
@@ -7,10 +7,9 @@ without hitting the network.
import json import json
import tarfile import tarfile
import tempfile
from io import BytesIO from io import BytesIO
from pathlib import Path from pathlib import Path
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import patch
import pytest import pytest
+1 -1
View File
@@ -40,7 +40,7 @@ def _has_amd_hardware():
return False return False
@pytest.fixture() @pytest.fixture
def backend_dir(): def backend_dir():
return Path(__file__).parent.parent return Path(__file__).parent.parent
+26 -27
View File
@@ -4,48 +4,47 @@ Test real model download with SSE progress monitoring.
import asyncio import asyncio
import json import json
import httpx import httpx
import time
from typing import List, Dict
async def monitor_sse_stream(model_name: str, timeout: int = 300): async def monitor_sse_stream(model_name: str, timeout: int = 300):
"""Monitor SSE stream for a model download.""" """Monitor SSE stream for a model download."""
events: List[Dict] = [] events: list[dict] = []
url = f"http://localhost:8000/models/progress/{model_name}" url = f"http://localhost:8000/models/progress/{model_name}"
print(f"Connecting to SSE endpoint: {url}") print(f"Connecting to SSE endpoint: {url}")
async with httpx.AsyncClient(timeout=timeout) as client: async with httpx.AsyncClient(timeout=timeout) as client, client.stream("GET", url) as response:
async with client.stream("GET", url) as response: print(f"SSE connected, status: {response.status_code}")
print(f"SSE connected, status: {response.status_code}")
if response.status_code != 200: if response.status_code != 200:
print(f"Error: SSE endpoint returned {response.status_code}") print(f"Error: SSE endpoint returned {response.status_code}")
return events return events
async for line in response.aiter_lines(): async for line in response.aiter_lines():
if not line: if not line:
continue continue
print(f" Raw SSE: {line[:100]}...") # Print first 100 chars print(f" Raw SSE: {line[:100]}...") # Print first 100 chars
if line.startswith("data: "): if line.startswith("data: "):
try: try:
data = json.loads(line[6:]) data = json.loads(line[6:])
print(f" → {data['status']:12} {data.get('progress', 0):6.1f}% {data.get('filename', '')}") print(f" → {data['status']:12} {data.get('progress', 0):6.1f}% {data.get('filename', '')}")
events.append(data) events.append(data)
# Stop if complete or error # Stop if complete or error
if data.get("status") in ("complete", "error"): if data.get("status") in ("complete", "error"):
print(f" Download {data['status']}!") print(f" Download {data['status']}!")
break break
except json.JSONDecodeError as e: except json.JSONDecodeError as e:
print(f" Error parsing JSON: {e}") print(f" Error parsing JSON: {e}")
print(f" Line was: {line}") print(f" Line was: {line}")
elif line.startswith(": heartbeat"): elif line.startswith(": heartbeat"):
print(" ♥ heartbeat") print(" ♥ heartbeat")
return events return events
+7 -7
View File
@@ -2,10 +2,10 @@
Audio processing utilities. Audio processing utilities.
""" """
import librosa
import numpy as np import numpy as np
import soundfile as sf import soundfile as sf
import librosa
from typing import Tuple, Optional
def normalize_audio( def normalize_audio(
@@ -48,7 +48,7 @@ def load_audio(
path: str, path: str,
sample_rate: int = 24000, sample_rate: int = 24000,
mono: bool = True, mono: bool = True,
) -> Tuple[np.ndarray, int]: ) -> tuple[np.ndarray, int]:
""" """
Load audio file with normalization. Load audio file with normalization.
@@ -84,8 +84,8 @@ def save_audio(
Raises: Raises:
OSError: If file cannot be written OSError: If file cannot be written
""" """
from pathlib import Path
import os import os
from pathlib import Path
temp_path = f"{path}.tmp" temp_path = f"{path}.tmp"
try: try:
@@ -264,7 +264,7 @@ def validate_reference_audio(
min_duration: float = 2.0, min_duration: float = 2.0,
max_duration: float = 30.0, max_duration: float = 30.0,
min_rms: float = 0.01, min_rms: float = 0.01,
) -> Tuple[bool, Optional[str]]: ) -> tuple[bool, str | None]:
""" """
Validate reference audio for voice cloning. Validate reference audio for voice cloning.
@@ -288,7 +288,7 @@ def validate_and_load_reference_audio(
min_duration: float = 2.0, min_duration: float = 2.0,
max_duration: float = 30.0, max_duration: float = 30.0,
min_rms: float = 0.01, min_rms: float = 0.01,
) -> Tuple[bool, Optional[str], Optional[np.ndarray], Optional[int]]: ) -> tuple[bool, str | None, np.ndarray | None, int | None]:
""" """
Validate and load reference audio in a single pass. Validate and load reference audio in a single pass.
@@ -315,4 +315,4 @@ def validate_and_load_reference_audio(
return True, None, audio, sr return True, None, audio, sr
except Exception as e: except Exception as e:
return False, f"Error validating audio: {str(e)}", None, None return False, f"Error validating audio: {e!s}", None, None
+6 -5
View File
@@ -4,9 +4,10 @@ Voice prompt caching utilities.
import hashlib import hashlib
import logging import logging
import torch
from pathlib import Path from pathlib import Path
from typing import Optional, Union, Dict, Any from typing import Any, Union
import torch
from .. import config from .. import config
@@ -19,7 +20,7 @@ def _get_cache_dir() -> Path:
# In-memory cache - can store dict (voice prompt) or tensor (legacy) # In-memory cache - can store dict (voice prompt) or tensor (legacy)
_memory_cache: dict[str, Union[torch.Tensor, Dict[str, Any]]] = {} _memory_cache: dict[str, Union[torch.Tensor, dict[str, Any]]] = {}
def get_cache_key(audio_path: str, reference_text: str) -> str: def get_cache_key(audio_path: str, reference_text: str) -> str:
@@ -46,7 +47,7 @@ def get_cache_key(audio_path: str, reference_text: str) -> str:
def get_cached_voice_prompt( def get_cached_voice_prompt(
cache_key: str, cache_key: str,
) -> Optional[Union[torch.Tensor, Dict[str, Any]]]: ) -> Union[torch.Tensor, dict[str, Any]] | None:
""" """
Get cached voice prompt if available. Get cached voice prompt if available.
@@ -76,7 +77,7 @@ def get_cached_voice_prompt(
def cache_voice_prompt( def cache_voice_prompt(
cache_key: str, cache_key: str,
voice_prompt: Union[torch.Tensor, Dict[str, Any]], voice_prompt: Union[torch.Tensor, dict[str, Any]],
) -> None: ) -> None:
""" """
Cache voice prompt to memory and disk. Cache voice prompt to memory and disk.
-1
View File
@@ -4,7 +4,6 @@ from __future__ import annotations
import sys import sys
MAC_PUSH_TO_TALK = ["MetaRight", "AltGr"] MAC_PUSH_TO_TALK = ["MetaRight", "AltGr"]
MAC_TOGGLE_TO_TALK = ["MetaRight", "AltGr", "Space"] MAC_TOGGLE_TO_TALK = ["MetaRight", "AltGr", "Space"]
NON_MAC_PUSH_TO_TALK = ["ControlRight", "ShiftRight"] NON_MAC_PUSH_TO_TALK = ["ControlRight", "ShiftRight"]
+5 -6
View File
@@ -11,7 +11,6 @@ overhead.
import logging import logging
import re import re
from typing import List, Tuple
import numpy as np import numpy as np
@@ -58,7 +57,7 @@ _ABBREVIATIONS = frozenset(
_PARA_TAG_RE = re.compile(r"\[[^\]]*\]") _PARA_TAG_RE = re.compile(r"\[[^\]]*\]")
def split_text_into_chunks(text: str, max_chars: int = DEFAULT_MAX_CHUNK_CHARS) -> List[str]: def split_text_into_chunks(text: str, max_chars: int = DEFAULT_MAX_CHUNK_CHARS) -> list[str]:
"""Split *text* at natural boundaries into chunks of at most *max_chars*. """Split *text* at natural boundaries into chunks of at most *max_chars*.
Priority: sentence-end (``.!?`` not preceded by an abbreviation and not Priority: sentence-end (``.!?`` not preceded by an abbreviation and not
@@ -73,7 +72,7 @@ def split_text_into_chunks(text: str, max_chars: int = DEFAULT_MAX_CHUNK_CHARS)
if len(text) <= max_chars: if len(text) <= max_chars:
return [text] return [text]
chunks: List[str] = [] chunks: list[str] = []
remaining = text remaining = text
while remaining: while remaining:
@@ -170,7 +169,7 @@ def _safe_hard_cut(segment: str, max_chars: int) -> int:
def concatenate_audio_chunks( def concatenate_audio_chunks(
chunks: List[np.ndarray], chunks: list[np.ndarray],
sample_rate: int, sample_rate: int,
crossfade_ms: int = 50, crossfade_ms: int = 50,
) -> np.ndarray: ) -> np.ndarray:
@@ -211,7 +210,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,
) -> Tuple[np.ndarray, int]: ) -> tuple[np.ndarray, int]:
"""Generate audio with automatic chunking for long text. """Generate audio with automatic chunking for long text.
For text shorter than *max_chunk_chars* this is a thin wrapper around For text shorter than *max_chunk_chars* this is a thin wrapper around
@@ -266,7 +265,7 @@ async def generate_chunked(
len(chunks), len(chunks),
max_chunk_chars, max_chunk_chars,
) )
audio_chunks: List[np.ndarray] = [] audio_chunks: list[np.ndarray] = []
sample_rate: int | None = None sample_rate: int | None = None
for i, chunk_text in enumerate(chunks): for i, chunk_text in enumerate(chunks):
-1
View File
@@ -21,7 +21,6 @@ import types
import torch import torch
import torch.nn as nn import torch.nn as nn
# ── Snake activation (from dac/nn/layers.py) ──────────────────────── # ── Snake activation (from dac/nn/layers.py) ────────────────────────
# NOTE: The original DAC code uses @torch.jit.script here for a 1.4x # NOTE: The original DAC code uses @torch.jit.script here for a 1.4x
+12 -13
View File
@@ -19,24 +19,23 @@ Supported effect types:
from __future__ import annotations from __future__ import annotations
import numpy as np from typing import Any
from typing import Any, Dict, List, Optional
import numpy as np
from pedalboard import ( from pedalboard import (
Pedalboard,
Chorus, Chorus,
Reverb,
Compressor, Compressor,
Delay,
Gain, Gain,
HighpassFilter, HighpassFilter,
LowpassFilter, LowpassFilter,
Delay, Pedalboard,
PitchShift, PitchShift,
Reverb,
) )
# Each param definition: (default, min, max, description) # Each param definition: (default, min, max, description)
EFFECT_REGISTRY: Dict[str, Dict[str, Any]] = { EFFECT_REGISTRY: dict[str, dict[str, Any]] = {
"chorus": { "chorus": {
"cls": Chorus, "cls": Chorus,
"label": "Chorus / Flanger", "label": "Chorus / Flanger",
@@ -147,7 +146,7 @@ EFFECT_REGISTRY: Dict[str, Dict[str, Any]] = {
} }
BUILTIN_PRESETS: Dict[str, Dict[str, Any]] = { BUILTIN_PRESETS: dict[str, dict[str, Any]] = {
"robotic": { "robotic": {
"name": "Robotic", "name": "Robotic",
"sort_order": 0, "sort_order": 0,
@@ -255,7 +254,7 @@ BUILTIN_PRESETS: Dict[str, Dict[str, Any]] = {
} }
def get_available_effects() -> List[Dict[str, Any]]: def get_available_effects() -> list[dict[str, Any]]:
"""Return the list of available effect types with their parameter definitions. """Return the list of available effect types with their parameter definitions.
Used by the frontend to build the effects chain editor UI. Used by the frontend to build the effects chain editor UI.
@@ -273,12 +272,12 @@ def get_available_effects() -> List[Dict[str, Any]]:
return result return result
def get_builtin_presets() -> Dict[str, Dict[str, Any]]: def get_builtin_presets() -> dict[str, dict[str, Any]]:
"""Return all built-in effect presets.""" """Return all built-in effect presets."""
return BUILTIN_PRESETS return BUILTIN_PRESETS
def validate_effects_chain(effects_chain: List[Dict[str, Any]]) -> Optional[str]: def validate_effects_chain(effects_chain: list[dict[str, Any]]) -> str | None:
"""Validate an effects chain configuration. """Validate an effects chain configuration.
Returns None if valid, or an error message string. Returns None if valid, or an error message string.
@@ -315,7 +314,7 @@ def validate_effects_chain(effects_chain: List[Dict[str, Any]]) -> Optional[str]
return None return None
def build_pedalboard(effects_chain: List[Dict[str, Any]]) -> Pedalboard: def build_pedalboard(effects_chain: list[dict[str, Any]]) -> Pedalboard:
"""Build a Pedalboard instance from an effects chain config. """Build a Pedalboard instance from an effects chain config.
Skips effects where ``enabled`` is ``False``. Skips effects where ``enabled`` is ``False``.
@@ -342,7 +341,7 @@ def build_pedalboard(effects_chain: List[Dict[str, Any]]) -> Pedalboard:
def apply_effects( def apply_effects(
audio: np.ndarray, audio: np.ndarray,
sample_rate: int, sample_rate: int,
effects_chain: List[Dict[str, Any]], effects_chain: list[dict[str, Any]],
) -> np.ndarray: ) -> np.ndarray:
"""Apply an effects chain to audio data. """Apply an effects chain to audio data.
+8 -8
View File
@@ -9,7 +9,7 @@ import os
import threading import threading
from contextlib import contextmanager from contextlib import contextmanager
from pathlib import Path from pathlib import Path
from typing import Optional, Union from typing import Union
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -25,9 +25,9 @@ logger = logging.getLogger(__name__)
_offline_lock = threading.RLock() _offline_lock = threading.RLock()
_offline_refcount = 0 _offline_refcount = 0
_saved_env: Optional[str] = None _saved_env: str | None = None
_saved_hf_const: Optional[bool] = None _saved_hf_const: bool | None = None
_saved_transformers_const: Optional[bool] = None _saved_transformers_const: bool | None = None
@contextmanager @contextmanager
@@ -61,8 +61,8 @@ def force_offline_if_cached(is_cached: bool, model_label: str = ""):
# bumping the refcount — a persistent offline leak that outlives # bumping the refcount — a persistent offline leak that outlives
# the process and is miserable to debug. # the process and is miserable to debug.
prev_env = os.environ.get("HF_HUB_OFFLINE") prev_env = os.environ.get("HF_HUB_OFFLINE")
prev_hf: Optional[bool] = None prev_hf: bool | None = None
prev_tf: Optional[bool] = None prev_tf: bool | None = None
try: try:
try: try:
import huggingface_hub.constants as hf_const import huggingface_hub.constants as hf_const
@@ -206,8 +206,8 @@ def patch_huggingface_hub_offline():
repo_id: str, repo_id: str,
filename: str, filename: str,
cache_dir: Union[str, Path, None] = None, cache_dir: Union[str, Path, None] = None,
revision: Optional[str] = None, revision: str | None = None,
repo_type: Optional[str] = None, repo_type: str | None = None,
): ):
result = original_try_load( result = original_try_load(
repo_id=repo_id, repo_id=repo_id,
+4 -4
View File
@@ -2,11 +2,11 @@
HuggingFace Hub download progress tracking. HuggingFace Hub download progress tracking.
""" """
from typing import Optional, Callable
from contextlib import contextmanager
import logging import logging
import threading
import sys import sys
import threading
from collections.abc import Callable
from contextlib import contextmanager
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -14,7 +14,7 @@ logger = logging.getLogger(__name__)
class HFProgressTracker: class HFProgressTracker:
"""Tracks HuggingFace Hub download progress by intercepting tqdm.""" """Tracks HuggingFace Hub download progress by intercepting tqdm."""
def __init__(self, progress_callback: Optional[Callable] = None, filter_non_downloads: bool = False): def __init__(self, progress_callback: Callable | None = None, filter_non_downloads: bool = False):
self.progress_callback = progress_callback self.progress_callback = progress_callback
self.filter_non_downloads = filter_non_downloads # Only filter if True self.filter_non_downloads = filter_non_downloads # Only filter if True
self._original_tqdm_class = None self._original_tqdm_class = None
+4 -4
View File
@@ -1,7 +1,7 @@
"""Image processing utilities for avatar uploads.""" """Image processing utilities for avatar uploads."""
from pathlib import Path from pathlib import Path
from typing import Optional, Tuple
from PIL import Image from PIL import Image
# JPEG can be reported as 'JPEG' or 'MPO' (for multi-picture format from some cameras) # JPEG can be reported as 'JPEG' or 'MPO' (for multi-picture format from some cameras)
@@ -10,7 +10,7 @@ MAX_SIZE = 512
MAX_FILE_SIZE = 5 * 1024 * 1024 # 5MB MAX_FILE_SIZE = 5 * 1024 * 1024 # 5MB
def validate_image(file_path: str) -> Tuple[bool, Optional[str]]: def validate_image(file_path: str) -> tuple[bool, str | None]:
""" """
Validate image format and file size. Validate image format and file size.
@@ -41,7 +41,7 @@ def validate_image(file_path: str) -> Tuple[bool, Optional[str]]:
return True, None return True, None
except Exception as e: except Exception as e:
return False, f"Invalid image file: {str(e)}" return False, f"Invalid image file: {e!s}"
def process_avatar(input_path: str, output_path: str, max_size: int = MAX_SIZE) -> None: def process_avatar(input_path: str, output_path: str, max_size: int = MAX_SIZE) -> None:
@@ -59,7 +59,7 @@ def process_avatar(input_path: str, output_path: str, max_size: int = MAX_SIZE)
# Handle EXIF orientation for JPEG images # Handle EXIF orientation for JPEG images
try: try:
from PIL import ExifTags from PIL import ExifTags
for orientation in ExifTags.TAGS.keys(): for orientation in ExifTags.TAGS:
if ExifTags.TAGS[orientation] == 'Orientation': if ExifTags.TAGS[orientation] == 'Orientation':
break break
exif = img._getexif() exif = img._getexif()
+13 -15
View File
@@ -2,8 +2,6 @@
Progress tracking for model downloads using Server-Sent Events. Progress tracking for model downloads using Server-Sent Events.
""" """
from typing import Optional, Callable, Dict, List
from fastapi.responses import StreamingResponse
import asyncio import asyncio
import json import json
import threading import threading
@@ -21,18 +19,18 @@ class ProgressManager:
THROTTLE_PROGRESS_DELTA = 1.0 # Minimum progress change (%) to force update THROTTLE_PROGRESS_DELTA = 1.0 # Minimum progress change (%) to force update
def __init__(self): def __init__(self):
self._progress: Dict[str, Dict] = {} self._progress: dict[str, dict] = {}
self._listeners: Dict[str, list] = {} self._listeners: dict[str, list] = {}
self._lock = threading.Lock() # Thread-safe lock for progress dict self._lock = threading.Lock() # Thread-safe lock for progress dict
self._main_loop: Optional[asyncio.AbstractEventLoop] = None self._main_loop: asyncio.AbstractEventLoop | None = None
self._last_notify_time: Dict[str, float] = {} # Last notification time per model self._last_notify_time: dict[str, float] = {} # Last notification time per model
self._last_notify_progress: Dict[str, float] = {} # Last notified progress per model self._last_notify_progress: dict[str, float] = {} # Last notified progress per model
def _set_main_loop(self, loop: asyncio.AbstractEventLoop): def _set_main_loop(self, loop: asyncio.AbstractEventLoop):
"""Set the main event loop for thread-safe operations.""" """Set the main event loop for thread-safe operations."""
self._main_loop = loop self._main_loop = loop
def _notify_listeners_threadsafe(self, model_name: str, progress_data: Dict): def _notify_listeners_threadsafe(self, model_name: str, progress_data: dict):
"""Notify listeners in a thread-safe manner.""" """Notify listeners in a thread-safe manner."""
import logging import logging
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -66,7 +64,7 @@ class ProgressManager:
model_name: str, model_name: str,
current: int, current: int,
total: int, total: int,
filename: Optional[str] = None, filename: str | None = None,
status: str = "downloading", status: str = "downloading",
): ):
""" """
@@ -143,13 +141,13 @@ class ProgressManager:
else: else:
logger.debug(f"No listeners for {model_name}, progress update stored: {progress_pct:.1f}%") logger.debug(f"No listeners for {model_name}, progress update stored: {progress_pct:.1f}%")
def get_progress(self, model_name: str) -> Optional[Dict]: def get_progress(self, model_name: str) -> dict | None:
"""Get current progress for a model. Thread-safe.""" """Get current progress for a model. Thread-safe."""
with self._lock: with self._lock:
progress = self._progress.get(model_name) progress = self._progress.get(model_name)
return progress.copy() if progress else None return progress.copy() if progress else None
def get_all_active(self) -> List[Dict]: def get_all_active(self) -> list[dict]:
"""Get all active downloads (status is 'downloading' or 'extracting'). Thread-safe.""" """Get all active downloads (status is 'downloading' or 'extracting'). Thread-safe."""
active = [] active = []
with self._lock: with self._lock:
@@ -159,7 +157,7 @@ class ProgressManager:
active.append(progress.copy()) active.append(progress.copy())
return active return active
def create_progress_callback(self, model_name: str, filename: Optional[str] = None): def create_progress_callback(self, model_name: str, filename: str | None = None):
""" """
Create a progress callback function for HuggingFace downloads. Create a progress callback function for HuggingFace downloads.
@@ -170,7 +168,7 @@ class ProgressManager:
Returns: Returns:
Callback function Callback function
""" """
def callback(progress: Dict): def callback(progress: dict):
"""HuggingFace Hub progress callback.""" """HuggingFace Hub progress callback."""
if "total" in progress and "current" in progress: if "total" in progress and "current" in progress:
current = progress.get("current", 0) current = progress.get("current", 0)
@@ -242,7 +240,7 @@ class ProgressManager:
if progress.get("status") in ("complete", "error"): if progress.get("status") in ("complete", "error"):
logger.info(f"Download {progress.get('status')} for {model_name}, closing SSE connection") logger.info(f"Download {progress.get('status')} for {model_name}, closing SSE connection")
break break
except asyncio.TimeoutError: except TimeoutError:
# Send heartbeat # Send heartbeat
yield ": heartbeat\n\n" yield ": heartbeat\n\n"
continue continue
@@ -304,7 +302,7 @@ class ProgressManager:
# Global progress manager instance # Global progress manager instance
_progress_manager: Optional[ProgressManager] = None _progress_manager: ProgressManager | None = None
def get_progress_manager() -> ProgressManager: def get_progress_manager() -> ProgressManager:
+7 -8
View File
@@ -2,9 +2,8 @@
Task tracking for active downloads and generations. Task tracking for active downloads and generations.
""" """
from typing import Optional, Dict, List
from datetime import datetime
from dataclasses import dataclass, field from dataclasses import dataclass, field
from datetime import datetime
@dataclass @dataclass
@@ -13,7 +12,7 @@ class DownloadTask:
model_name: str model_name: str
status: str = "downloading" # downloading, extracting, complete, error status: str = "downloading" # downloading, extracting, complete, error
started_at: datetime = field(default_factory=datetime.utcnow) started_at: datetime = field(default_factory=datetime.utcnow)
error: Optional[str] = None error: str | None = None
@dataclass @dataclass
@@ -29,8 +28,8 @@ class TaskManager:
"""Manages active downloads and generations.""" """Manages active downloads and generations."""
def __init__(self): def __init__(self):
self._active_downloads: Dict[str, DownloadTask] = {} self._active_downloads: dict[str, DownloadTask] = {}
self._active_generations: Dict[str, GenerationTask] = {} self._active_generations: dict[str, GenerationTask] = {}
def start_download(self, model_name: str) -> None: def start_download(self, model_name: str) -> None:
"""Mark a download as started.""" """Mark a download as started."""
@@ -64,11 +63,11 @@ class TaskManager:
if task_id in self._active_generations: if task_id in self._active_generations:
del self._active_generations[task_id] del self._active_generations[task_id]
def get_active_downloads(self) -> List[DownloadTask]: def get_active_downloads(self) -> list[DownloadTask]:
"""Get all active downloads.""" """Get all active downloads."""
return list(self._active_downloads.values()) return list(self._active_downloads.values())
def get_active_generations(self) -> List[GenerationTask]: def get_active_generations(self) -> list[GenerationTask]:
"""Get all active generations.""" """Get all active generations."""
return list(self._active_generations.values()) return list(self._active_generations.values())
@@ -91,7 +90,7 @@ class TaskManager:
# Global task manager instance # Global task manager instance
_task_manager: Optional[TaskManager] = None _task_manager: TaskManager | None = None
def get_task_manager() -> TaskManager: def get_task_manager() -> TaskManager: