comment cleanup

This commit is contained in:
James Pine
2026-03-16 01:46:19 -07:00
parent fe19a9ca47
commit b7781951df
10 changed files with 738 additions and 694 deletions
+79 -20
View File
@@ -14,9 +14,16 @@ import numpy as np
from ..platform_detect import get_backend_type from ..platform_detect import get_backend_type
LANGUAGE_CODE_TO_NAME = { LANGUAGE_CODE_TO_NAME = {
"zh": "chinese", "en": "english", "ja": "japanese", "ko": "korean", "zh": "chinese",
"de": "german", "fr": "french", "ru": "russian", "pt": "portuguese", "en": "english",
"es": "spanish", "it": "italian", "ja": "japanese",
"ko": "korean",
"de": "german",
"fr": "french",
"ru": "russian",
"pt": "portuguese",
"es": "spanish",
"it": "italian",
} }
WHISPER_HF_REPOS = { WHISPER_HF_REPOS = {
@@ -31,6 +38,7 @@ WHISPER_HF_REPOS = {
@dataclass @dataclass
class ModelConfig: class ModelConfig:
"""Declarative config for a downloadable model variant.""" """Declarative config for a downloadable model variant."""
model_name: str # e.g. "luxtts", "chatterbox-tts" model_name: str # e.g. "luxtts", "chatterbox-tts"
display_name: str # e.g. "LuxTTS (Fast, CPU-friendly)" display_name: str # e.g. "LuxTTS (Fast, CPU-friendly)"
engine: str # e.g. "luxtts", "chatterbox" engine: str # e.g. "luxtts", "chatterbox"
@@ -160,10 +168,6 @@ TTS_ENGINES = {
} }
# ---------------------------------------------------------------------------
# Model config registry
# ---------------------------------------------------------------------------
def _get_qwen_model_configs() -> list[ModelConfig]: def _get_qwen_model_configs() -> list[ModelConfig]:
"""Return Qwen model configs with backend-aware HF repo IDs.""" """Return Qwen model configs with backend-aware HF repo IDs."""
backend_type = get_backend_type() backend_type = get_backend_type()
@@ -220,9 +224,29 @@ def _get_non_qwen_tts_configs() -> list[ModelConfig]:
size_mb=3200, size_mb=3200,
needs_trim=True, needs_trim=True,
languages=[ languages=[
"zh", "en", "ja", "ko", "de", "fr", "ru", "pt", "es", "it", "zh",
"he", "ar", "da", "el", "fi", "hi", "ms", "nl", "no", "pl", "en",
"sv", "sw", "tr", "ja",
"ko",
"de",
"fr",
"ru",
"pt",
"es",
"it",
"he",
"ar",
"da",
"el",
"fi",
"hi",
"ms",
"nl",
"no",
"pl",
"sv",
"sw",
"tr",
], ],
), ),
ModelConfig( ModelConfig(
@@ -240,11 +264,41 @@ def _get_non_qwen_tts_configs() -> list[ModelConfig]:
def _get_whisper_configs() -> list[ModelConfig]: def _get_whisper_configs() -> list[ModelConfig]:
"""Return Whisper STT model configs.""" """Return Whisper STT model configs."""
return [ return [
ModelConfig(model_name="whisper-base", display_name="Whisper Base", engine="whisper", hf_repo_id="openai/whisper-base", model_size="base"), ModelConfig(
ModelConfig(model_name="whisper-small", display_name="Whisper Small", engine="whisper", hf_repo_id="openai/whisper-small", model_size="small"), model_name="whisper-base",
ModelConfig(model_name="whisper-medium", display_name="Whisper Medium", engine="whisper", hf_repo_id="openai/whisper-medium", model_size="medium"), display_name="Whisper Base",
ModelConfig(model_name="whisper-large", display_name="Whisper Large", engine="whisper", hf_repo_id="openai/whisper-large-v3", model_size="large"), engine="whisper",
ModelConfig(model_name="whisper-turbo", display_name="Whisper Turbo", engine="whisper", hf_repo_id="openai/whisper-large-v3-turbo", model_size="turbo"), hf_repo_id="openai/whisper-base",
model_size="base",
),
ModelConfig(
model_name="whisper-small",
display_name="Whisper Small",
engine="whisper",
hf_repo_id="openai/whisper-small",
model_size="small",
),
ModelConfig(
model_name="whisper-medium",
display_name="Whisper Medium",
engine="whisper",
hf_repo_id="openai/whisper-medium",
model_size="medium",
),
ModelConfig(
model_name="whisper-large",
display_name="Whisper Large",
engine="whisper",
hf_repo_id="openai/whisper-large-v3",
model_size="large",
),
ModelConfig(
model_name="whisper-turbo",
display_name="Whisper Turbo",
engine="whisper",
hf_repo_id="openai/whisper-large-v3-turbo",
model_size="turbo",
),
] ]
@@ -260,6 +314,7 @@ def get_tts_model_configs() -> list[ModelConfig]:
# Lookup helpers — these replace the if/elif chains in main.py # Lookup helpers — these replace the if/elif chains in main.py
def get_model_config(model_name: str) -> Optional[ModelConfig]: def get_model_config(model_name: str) -> Optional[ModelConfig]:
"""Look up a model config by model_name.""" """Look up a model config by model_name."""
for cfg in get_all_model_configs(): for cfg in get_all_model_configs():
@@ -294,6 +349,7 @@ async def load_engine_model(engine: str, model_size: str = "default") -> None:
async def ensure_model_cached_or_raise(engine: str, model_size: str = "default") -> None: async def ensure_model_cached_or_raise(engine: str, model_size: str = "default") -> None:
"""Check if a model is cached, raise HTTPException if not. Used by streaming endpoint.""" """Check if a model is cached, raise HTTPException if not. Used by streaming endpoint."""
from fastapi import HTTPException from fastapi import HTTPException
backend = get_tts_backend_for_engine(engine) backend = get_tts_backend_for_engine(engine)
cfg = None cfg = None
for c in get_tts_model_configs(): for c in get_tts_model_configs():
@@ -352,7 +408,7 @@ def check_model_loaded(config: ModelConfig) -> bool:
try: try:
if config.engine == "whisper": if config.engine == "whisper":
whisper_model = transcribe.get_whisper_model() whisper_model = transcribe.get_whisper_model()
return whisper_model.is_loaded() and getattr(whisper_model, 'model_size', None) == config.model_size return whisper_model.is_loaded() and getattr(whisper_model, "model_size", None) == config.model_size
if config.engine == "qwen": if config.engine == "qwen":
tts_model = tts.get_tts_model() tts_model = tts.get_tts_model()
@@ -379,10 +435,6 @@ def get_model_load_func(config: ModelConfig):
return lambda: get_tts_backend_for_engine(config.engine).load_model() return lambda: get_tts_backend_for_engine(config.engine).load_model()
# ---------------------------------------------------------------------------
# Backend factory
# ---------------------------------------------------------------------------
def get_tts_backend() -> TTSBackend: def get_tts_backend() -> TTSBackend:
""" """
Get or create the default (Qwen) TTS backend instance based on platform. Get or create the default (Qwen) TTS backend instance based on platform.
@@ -419,18 +471,23 @@ def get_tts_backend_for_engine(engine: str) -> TTSBackend:
backend_type = get_backend_type() backend_type = get_backend_type()
if backend_type == "mlx": if backend_type == "mlx":
from .mlx_backend import MLXTTSBackend from .mlx_backend import MLXTTSBackend
backend = MLXTTSBackend() backend = MLXTTSBackend()
else: else:
from .pytorch_backend import PyTorchTTSBackend from .pytorch_backend import PyTorchTTSBackend
backend = PyTorchTTSBackend() backend = PyTorchTTSBackend()
elif engine == "luxtts": elif engine == "luxtts":
from .luxtts_backend import LuxTTSBackend from .luxtts_backend import LuxTTSBackend
backend = LuxTTSBackend() backend = LuxTTSBackend()
elif engine == "chatterbox": elif engine == "chatterbox":
from .chatterbox_backend import ChatterboxTTSBackend from .chatterbox_backend import ChatterboxTTSBackend
backend = ChatterboxTTSBackend() backend = ChatterboxTTSBackend()
elif engine == "chatterbox_turbo": elif engine == "chatterbox_turbo":
from .chatterbox_turbo_backend import ChatterboxTurboTTSBackend from .chatterbox_turbo_backend import ChatterboxTurboTTSBackend
backend = ChatterboxTurboTTSBackend() backend = ChatterboxTurboTTSBackend()
else: else:
raise ValueError(f"Unknown TTS engine: {engine}. Supported: {list(TTS_ENGINES.keys())}") raise ValueError(f"Unknown TTS engine: {engine}. Supported: {list(TTS_ENGINES.keys())}")
@@ -453,9 +510,11 @@ def get_stt_backend() -> STTBackend:
if backend_type == "mlx": if backend_type == "mlx":
from .mlx_backend import MLXSTTBackend from .mlx_backend import MLXSTTBackend
_stt_backend = MLXSTTBackend() _stt_backend = MLXSTTBackend()
else: else:
from .pytorch_backend import PyTorchSTTBackend from .pytorch_backend import PyTorchSTTBackend
_stt_backend = PyTorchSTTBackend() _stt_backend = PyTorchSTTBackend()
return _stt_backend return _stt_backend
+4 -23
View File
@@ -21,10 +21,6 @@ from ..utils.tasks import get_task_manager
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# HuggingFace cache checking
# ---------------------------------------------------------------------------
def is_model_cached( def is_model_cached(
hf_repo: str, hf_repo: str,
*, *,
@@ -46,9 +42,7 @@ def is_model_cached(
try: try:
from huggingface_hub import constants as hf_constants from huggingface_hub import constants as hf_constants
repo_cache = Path(hf_constants.HF_HUB_CACHE) / ( repo_cache = Path(hf_constants.HF_HUB_CACHE) / ("models--" + hf_repo.replace("/", "--"))
"models--" + hf_repo.replace("/", "--")
)
if not repo_cache.exists(): if not repo_cache.exists():
return False return False
@@ -83,10 +77,6 @@ def is_model_cached(
return False return False
# ---------------------------------------------------------------------------
# Device detection
# ---------------------------------------------------------------------------
def get_torch_device( def get_torch_device(
*, *,
allow_xpu: bool = False, allow_xpu: bool = False,
@@ -114,6 +104,7 @@ def get_torch_device(
if allow_xpu: if allow_xpu:
try: try:
import intel_extension_for_pytorch # noqa: F401 import intel_extension_for_pytorch # noqa: F401
if hasattr(torch, "xpu") and torch.xpu.is_available(): if hasattr(torch, "xpu") and torch.xpu.is_available():
return "xpu" return "xpu"
except ImportError: except ImportError:
@@ -122,6 +113,7 @@ def get_torch_device(
if allow_directml: if allow_directml:
try: try:
import torch_directml import torch_directml
if torch_directml.device_count() > 0: if torch_directml.device_count() > 0:
return torch_directml.device(0) return torch_directml.device(0)
except ImportError: except ImportError:
@@ -134,10 +126,6 @@ def get_torch_device(
return "cpu" return "cpu"
# ---------------------------------------------------------------------------
# Voice prompt combination
# ---------------------------------------------------------------------------
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],
@@ -169,10 +157,6 @@ async def combine_voice_prompts(
return mixed, combined_text return mixed, combined_text
# ---------------------------------------------------------------------------
# Model loading progress tracking
# ---------------------------------------------------------------------------
@contextmanager @contextmanager
def model_load_progress( def model_load_progress(
model_name: str, model_name: str,
@@ -237,10 +221,6 @@ def model_load_progress(
tracker_context.__exit__(None, None, None) tracker_context.__exit__(None, None, None)
# ---------------------------------------------------------------------------
# Chatterbox f32 dtype patches
# ---------------------------------------------------------------------------
def patch_chatterbox_f32(model) -> None: def patch_chatterbox_f32(model) -> None:
""" """
Patch float64 -> float32 dtype mismatches in upstream chatterbox. Patch float64 -> float32 dtype mismatches in upstream chatterbox.
@@ -261,6 +241,7 @@ def patch_chatterbox_f32(model) -> None:
def _f32_log_mel(self_tokzr, audio, padding=0): def _f32_log_mel(self_tokzr, audio, padding=0):
import torch as _torch import torch as _torch
if _torch.is_tensor(audio): if _torch.is_tensor(audio):
audio = audio.float() audio = audio.float()
return _orig_log_mel(self_tokzr, audio, padding) return _orig_log_mel(self_tokzr, audio, padding)
+145 -155
View File
@@ -41,17 +41,24 @@ def _safe_content_disposition(disposition_type: str, filename: str) -> str:
Uses RFC 5987 ``filename*`` parameter so that browsers can decode Uses RFC 5987 ``filename*`` parameter so that browsers can decode
UTF-8 filenames while the ``filename`` fallback stays ASCII-only. UTF-8 filenames while the ``filename`` fallback stays ASCII-only.
""" """
ascii_name = "".join( ascii_name = "".join(c for c in filename if c.isascii() and (c.isalnum() or c in " -_.")).strip() or "download"
c for c in filename if c.isascii() and (c.isalnum() or c in " -_.")
).strip() or "download"
utf8_name = quote(filename, safe="") utf8_name = quote(filename, safe="")
return ( return f"{disposition_type}; filename=\"{ascii_name}\"; filename*=UTF-8''{utf8_name}"
f'{disposition_type}; filename="{ascii_name}"; '
f"filename*=UTF-8''{utf8_name}"
)
from . import database, models, profiles, history, tts, transcribe, config, export_import, channels, stories, __version__ from . import (
database,
models,
profiles,
history,
tts,
transcribe,
config,
export_import,
channels,
stories,
__version__,
)
from .database import get_db, Generation as DBGeneration, VoiceProfile as DBVoiceProfile from .database import get_db, Generation as DBGeneration, VoiceProfile as DBVoiceProfile
from .profiles import _profile_to_response from .profiles import _profile_to_response
from .utils.progress import get_progress_manager from .utils.progress import get_progress_manager
@@ -92,10 +99,6 @@ app.add_middleware(
) )
# ============================================
# ROOT & HEALTH ENDPOINTS
# ============================================
@app.get("/") @app.get("/")
async def root(): async def root():
"""Root endpoint.""" """Root endpoint."""
@@ -105,6 +108,7 @@ async def root():
@app.post("/shutdown") @app.post("/shutdown")
async def shutdown(): async def shutdown():
"""Gracefully shutdown the server.""" """Gracefully shutdown the server."""
async def shutdown_async(): async def shutdown_async():
await asyncio.sleep(0.1) # Give response time to send await asyncio.sleep(0.1) # Give response time to send
os.kill(os.getpid(), signal.SIGTERM) os.kill(os.getpid(), signal.SIGTERM)
@@ -117,6 +121,7 @@ async def shutdown():
async def watchdog_disable(): async def watchdog_disable():
"""Disable the parent process watchdog so the server keeps running.""" """Disable the parent process watchdog so the server keeps running."""
from backend.server import disable_watchdog from backend.server import disable_watchdog
disable_watchdog() disable_watchdog()
return {"message": "Watchdog disabled"} return {"message": "Watchdog disabled"}
@@ -133,14 +138,15 @@ async def health():
# Check for GPU availability (CUDA, MPS, Intel Arc XPU, or DirectML) # Check for GPU availability (CUDA, MPS, Intel Arc XPU, or DirectML)
has_cuda = torch.cuda.is_available() has_cuda = torch.cuda.is_available()
has_mps = hasattr(torch.backends, 'mps') and torch.backends.mps.is_available() has_mps = hasattr(torch.backends, "mps") and torch.backends.mps.is_available()
# Intel Arc / Intel Xe via intel-extension-for-pytorch (IPEX) # Intel Arc / Intel Xe via intel-extension-for-pytorch (IPEX)
has_xpu = False has_xpu = False
xpu_name = None xpu_name = None
try: try:
import intel_extension_for_pytorch as ipex # noqa: F401 import intel_extension_for_pytorch as ipex # noqa: F401
if hasattr(torch, 'xpu') and torch.xpu.is_available():
if hasattr(torch, "xpu") and torch.xpu.is_available():
has_xpu = True has_xpu = True
try: try:
xpu_name = torch.xpu.get_device_name(0) xpu_name = torch.xpu.get_device_name(0)
@@ -154,6 +160,7 @@ async def health():
directml_name = None directml_name = None
try: try:
import torch_directml import torch_directml
if torch_directml.device_count() > 0: if torch_directml.device_count() > 0:
has_directml = True has_directml = True
try: try:
@@ -190,10 +197,10 @@ async def health():
model_loaded = True model_loaded = True
# Get the actual loaded model size # Get the actual loaded model size
# Check _current_model_size first (more reliable for actually loaded models) # Check _current_model_size first (more reliable for actually loaded models)
model_size = getattr(tts_model, '_current_model_size', None) model_size = getattr(tts_model, "_current_model_size", None)
if not model_size: if not model_size:
# Fallback to model_size attribute (which should be set when model loads) # Fallback to model_size attribute (which should be set when model loads)
model_size = getattr(tts_model, 'model_size', None) model_size = getattr(tts_model, "model_size", None)
except Exception: except Exception:
# If there's an error checking, assume not loaded # If there's an error checking, assume not loaded
model_loaded = False model_loaded = False
@@ -204,12 +211,14 @@ async def health():
try: try:
# Check if the default model (1.7B) is cached # Check if the default model (1.7B) is cached
from .backends import get_model_config from .backends import get_model_config
default_config = get_model_config("qwen-tts-1.7B") default_config = get_model_config("qwen-tts-1.7B")
default_model_id = default_config.hf_repo_id if default_config else "Qwen/Qwen3-TTS-12Hz-1.7B-Base" default_model_id = default_config.hf_repo_id if default_config else "Qwen/Qwen3-TTS-12Hz-1.7B-Base"
# Method 1: Try scan_cache_dir if available # Method 1: Try scan_cache_dir if available
try: try:
from huggingface_hub import scan_cache_dir from huggingface_hub import scan_cache_dir
cache_info = scan_cache_dir() cache_info = scan_cache_dir()
for repo in cache_info.repos: for repo in cache_info.repos:
if repo.repo_id == default_model_id: if repo.repo_id == default_model_id:
@@ -221,11 +230,11 @@ async def health():
repo_cache = Path(cache_dir) / ("models--" + default_model_id.replace("/", "--")) repo_cache = Path(cache_dir) / ("models--" + default_model_id.replace("/", "--"))
if repo_cache.exists(): if repo_cache.exists():
has_model_files = ( has_model_files = (
any(repo_cache.rglob("*.bin")) or any(repo_cache.rglob("*.bin"))
any(repo_cache.rglob("*.safetensors")) or or any(repo_cache.rglob("*.safetensors"))
any(repo_cache.rglob("*.pt")) or or any(repo_cache.rglob("*.pt"))
any(repo_cache.rglob("*.pth")) or or any(repo_cache.rglob("*.pth"))
any(repo_cache.rglob("*.npz")) # MLX models may use npz or any(repo_cache.rglob("*.npz")) # MLX models may use npz
) )
model_downloaded = has_model_files model_downloaded = has_model_files
except Exception: except Exception:
@@ -313,10 +322,6 @@ async def filesystem_health():
) )
# ============================================
# VOICE PROFILE ENDPOINTS
# ============================================
@app.post("/profiles", response_model=models.VoiceProfileResponse) @app.post("/profiles", response_model=models.VoiceProfileResponse)
async def create_profile( async def create_profile(
data: models.VoiceProfileCreate, data: models.VoiceProfileCreate,
@@ -352,8 +357,7 @@ async def import_profile(
if len(content) > MAX_FILE_SIZE: if len(content) > MAX_FILE_SIZE:
raise HTTPException( raise HTTPException(
status_code=400, status_code=400, detail=f"File too large. Maximum size is {MAX_FILE_SIZE / (1024 * 1024)}MB"
detail=f"File too large. Maximum size is {MAX_FILE_SIZE / (1024 * 1024)}MB"
) )
try: try:
@@ -415,9 +419,9 @@ async def add_profile_sample(
"""Add a sample to a voice profile.""" """Add a sample to a voice profile."""
# Preserve the uploaded file's extension so librosa can detect format correctly. # Preserve the uploaded file's extension so librosa can detect format correctly.
# Defaulting to .wav was causing soundfile to reject MP3/WebM content as invalid WAV. # Defaulting to .wav was causing soundfile to reject MP3/WebM content as invalid WAV.
_allowed_audio_exts = {'.wav', '.mp3', '.m4a', '.ogg', '.flac', '.aac', '.webm', '.opus'} _allowed_audio_exts = {".wav", ".mp3", ".m4a", ".ogg", ".flac", ".aac", ".webm", ".opus"}
_uploaded_ext = Path(file.filename or '').suffix.lower() _uploaded_ext = Path(file.filename or "").suffix.lower()
file_suffix = _uploaded_ext if _uploaded_ext in _allowed_audio_exts else '.wav' file_suffix = _uploaded_ext if _uploaded_ext in _allowed_audio_exts else ".wav"
with tempfile.NamedTemporaryFile(suffix=file_suffix, delete=False) as tmp: with tempfile.NamedTemporaryFile(suffix=file_suffix, delete=False) as tmp:
content = await file.read() content = await file.read()
@@ -546,7 +550,7 @@ async def export_profile(
zip_bytes = export_import.export_profile_to_zip(profile_id, db) zip_bytes = export_import.export_profile_to_zip(profile_id, db)
# Create safe filename # Create safe filename
safe_name = "".join(c for c in profile.name if c.isalnum() or c in (' ', '-', '_')).strip() safe_name = "".join(c for c in profile.name if c.isalnum() or c in (" ", "-", "_")).strip()
if not safe_name: if not safe_name:
safe_name = "profile" safe_name = "profile"
filename = f"profile-{safe_name}.voicebox.zip" filename = f"profile-{safe_name}.voicebox.zip"
@@ -555,9 +559,7 @@ async def export_profile(
return StreamingResponse( return StreamingResponse(
io.BytesIO(zip_bytes), io.BytesIO(zip_bytes),
media_type="application/zip", media_type="application/zip",
headers={ headers={"Content-Disposition": _safe_content_disposition("attachment", filename)},
"Content-Disposition": _safe_content_disposition("attachment", filename)
}
) )
except ValueError as e: except ValueError as e:
raise HTTPException(status_code=400, detail=str(e)) raise HTTPException(status_code=400, detail=str(e))
@@ -565,10 +567,6 @@ async def export_profile(
raise HTTPException(status_code=500, detail=str(e)) raise HTTPException(status_code=500, detail=str(e))
# ============================================
# AUDIO CHANNEL ENDPOINTS
# ============================================
@app.get("/channels", response_model=List[models.AudioChannelResponse]) @app.get("/channels", response_model=List[models.AudioChannelResponse])
async def list_channels(db: Session = Depends(get_db)): async def list_channels(db: Session = Depends(get_db)):
"""List all audio channels.""" """List all audio channels."""
@@ -684,10 +682,6 @@ async def set_profile_channels(
raise HTTPException(status_code=400, detail=str(e)) raise HTTPException(status_code=400, detail=str(e))
# ============================================
# GENERATION ENDPOINTS
# ============================================
@app.post("/generate", response_model=models.GenerationResponse) @app.post("/generate", response_model=models.GenerationResponse)
async def generate_speech( async def generate_speech(
data: models.GenerationRequest, data: models.GenerationRequest,
@@ -707,6 +701,7 @@ async def generate_speech(
raise HTTPException(status_code=404, detail="Profile not found") raise HTTPException(status_code=404, detail="Profile not found")
from .backends import engine_has_model_sizes from .backends import engine_has_model_sizes
engine = data.engine or "qwen" engine = data.engine or "qwen"
model_size = data.model_size or "1.7B" model_size = data.model_size or "1.7B"
@@ -740,6 +735,7 @@ async def generate_speech(
else: else:
# Check profile default # Check profile default
import json as _json import json as _json
profile_obj = db.query(DBVoiceProfile).filter_by(id=data.profile_id).first() profile_obj = db.query(DBVoiceProfile).filter_by(id=data.profile_id).first()
if profile_obj and profile_obj.effects_chain: if profile_obj and profile_obj.effects_chain:
try: try:
@@ -748,7 +744,8 @@ async def generate_speech(
pass pass
# Kick off TTS in background # Kick off TTS in background
enqueue_generation(run_generation( enqueue_generation(
run_generation(
generation_id=generation_id, generation_id=generation_id,
profile_id=data.profile_id, profile_id=data.profile_id,
text=data.text, text=data.text,
@@ -762,7 +759,8 @@ async def generate_speech(
mode="generate", mode="generate",
max_chunk_chars=data.max_chunk_chars, max_chunk_chars=data.max_chunk_chars,
crossfade_ms=data.crossfade_ms, crossfade_ms=data.crossfade_ms,
)) )
)
return generation return generation
@@ -792,7 +790,8 @@ async def retry_generation(generation_id: str, db: Session = Depends(get_db)):
text=gen.text, text=gen.text,
) )
enqueue_generation(run_generation( enqueue_generation(
run_generation(
generation_id=generation_id, generation_id=generation_id,
profile_id=gen.profile_id, profile_id=gen.profile_id,
text=gen.text, text=gen.text,
@@ -802,7 +801,8 @@ async def retry_generation(generation_id: str, db: Session = Depends(get_db)):
seed=gen.seed, seed=gen.seed,
instruct=gen.instruct, instruct=gen.instruct,
mode="retry", mode="retry",
)) )
)
return models.GenerationResponse.model_validate(gen) return models.GenerationResponse.model_validate(gen)
@@ -834,7 +834,8 @@ async def regenerate_generation(generation_id: str, db: Session = Depends(get_db
version_id = str(uuid.uuid4()) version_id = str(uuid.uuid4())
enqueue_generation(run_generation( enqueue_generation(
run_generation(
generation_id=generation_id, generation_id=generation_id,
profile_id=gen.profile_id, profile_id=gen.profile_id,
text=gen.text, text=gen.text,
@@ -845,7 +846,8 @@ async def regenerate_generation(generation_id: str, db: Session = Depends(get_db
instruct=gen.instruct, instruct=gen.instruct,
mode="regenerate", mode="regenerate",
version_id=version_id, version_id=version_id,
)) )
)
return models.GenerationResponse.model_validate(gen) return models.GenerationResponse.model_validate(gen)
@@ -914,11 +916,14 @@ async def stream_speech(
model_size = data.model_size or "1.7B" model_size = data.model_size or "1.7B"
from .backends import ensure_model_cached_or_raise, load_engine_model, engine_needs_trim from .backends import ensure_model_cached_or_raise, load_engine_model, engine_needs_trim
await ensure_model_cached_or_raise(engine, model_size) await ensure_model_cached_or_raise(engine, model_size)
await load_engine_model(engine, model_size) await load_engine_model(engine, model_size)
voice_prompt = await profiles.create_voice_prompt_for_profile( voice_prompt = await profiles.create_voice_prompt_for_profile(
data.profile_id, db, engine=engine, data.profile_id,
db,
engine=engine,
) )
from .utils.chunked_tts import generate_chunked from .utils.chunked_tts import generate_chunked
@@ -926,6 +931,7 @@ async def stream_speech(
trim_fn = None trim_fn = None
if engine_needs_trim(engine): if engine_needs_trim(engine):
from .utils.audio import trim_tts_output from .utils.audio import trim_tts_output
trim_fn = trim_tts_output trim_fn = trim_tts_output
audio, sample_rate = await generate_chunked( audio, sample_rate = await generate_chunked(
@@ -942,6 +948,7 @@ async def stream_speech(
if data.normalize: if data.normalize:
from .utils.audio import normalize_audio from .utils.audio import normalize_audio
audio = normalize_audio(audio) audio = normalize_audio(audio)
wav_bytes = tts.audio_to_wav_bytes(audio, sample_rate) wav_bytes = tts.audio_to_wav_bytes(audio, sample_rate)
@@ -959,10 +966,6 @@ async def stream_speech(
) )
# ============================================
# HISTORY ENDPOINTS
# ============================================
@app.get("/history", response_model=models.HistoryListResponse) @app.get("/history", response_model=models.HistoryListResponse)
async def list_history( async def list_history(
profile_id: Optional[str] = None, profile_id: Optional[str] = None,
@@ -1001,8 +1004,7 @@ async def import_generation(
if len(content) > MAX_FILE_SIZE: if len(content) > MAX_FILE_SIZE:
raise HTTPException( raise HTTPException(
status_code=400, status_code=400, detail=f"File too large. Maximum size is {MAX_FILE_SIZE / (1024 * 1024)}MB"
detail=f"File too large. Maximum size is {MAX_FILE_SIZE / (1024 * 1024)}MB"
) )
try: try:
@@ -1021,15 +1023,12 @@ async def get_generation(
): ):
"""Get a generation by ID.""" """Get a generation by ID."""
# Get generation with profile name # Get generation with profile name
result = db.query( result = (
DBGeneration, db.query(DBGeneration, DBVoiceProfile.name.label("profile_name"))
DBVoiceProfile.name.label('profile_name') .join(DBVoiceProfile, DBGeneration.profile_id == DBVoiceProfile.id)
).join( .filter(DBGeneration.id == generation_id)
DBVoiceProfile, .first()
DBGeneration.profile_id == DBVoiceProfile.id )
).filter(
DBGeneration.id == generation_id
).first()
if not result: if not result:
raise HTTPException(status_code=404, detail="Generation not found") raise HTTPException(status_code=404, detail="Generation not found")
@@ -1091,7 +1090,7 @@ async def export_generation(
zip_bytes = export_import.export_generation_to_zip(generation_id, db) zip_bytes = export_import.export_generation_to_zip(generation_id, db)
# Create safe filename from text # Create safe filename from text
safe_text = "".join(c for c in generation.text[:30] if c.isalnum() or c in (' ', '-', '_')).strip() safe_text = "".join(c for c in generation.text[:30] if c.isalnum() or c in (" ", "-", "_")).strip()
if not safe_text: if not safe_text:
safe_text = "generation" safe_text = "generation"
filename = f"generation-{safe_text}.voicebox.zip" filename = f"generation-{safe_text}.voicebox.zip"
@@ -1100,9 +1099,7 @@ async def export_generation(
return StreamingResponse( return StreamingResponse(
io.BytesIO(zip_bytes), io.BytesIO(zip_bytes),
media_type="application/zip", media_type="application/zip",
headers={ headers={"Content-Disposition": _safe_content_disposition("attachment", filename)},
"Content-Disposition": _safe_content_disposition("attachment", filename)
}
) )
except ValueError as e: except ValueError as e:
raise HTTPException(status_code=400, detail=str(e)) raise HTTPException(status_code=400, detail=str(e))
@@ -1125,7 +1122,7 @@ async def export_generation_audio(
raise HTTPException(status_code=404, detail="Audio file not found") raise HTTPException(status_code=404, detail="Audio file not found")
# Create safe filename from text # Create safe filename from text
safe_text = "".join(c for c in generation.text[:30] if c.isalnum() or c in (' ', '-', '_')).strip() safe_text = "".join(c for c in generation.text[:30] if c.isalnum() or c in (" ", "-", "_")).strip()
if not safe_text: if not safe_text:
safe_text = "generation" safe_text = "generation"
filename = f"{safe_text}.wav" filename = f"{safe_text}.wav"
@@ -1133,16 +1130,10 @@ async def export_generation_audio(
return FileResponse( return FileResponse(
audio_path, audio_path,
media_type="audio/wav", media_type="audio/wav",
headers={ headers={"Content-Disposition": _safe_content_disposition("attachment", filename)},
"Content-Disposition": _safe_content_disposition("attachment", filename)
}
) )
# ============================================
# TRANSCRIPTION ENDPOINTS
# ============================================
@app.post("/transcribe", response_model=models.TranscriptionResponse) @app.post("/transcribe", response_model=models.TranscriptionResponse)
async def transcribe_audio( async def transcribe_audio(
file: UploadFile = File(...), file: UploadFile = File(...),
@@ -1158,6 +1149,7 @@ async def transcribe_audio(
try: try:
# Get audio duration # Get audio duration
from .utils.audio import load_audio 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
@@ -1175,6 +1167,7 @@ async def transcribe_audio(
# Check if model is cached # Check if model is cached
from huggingface_hub import constants as hf_constants from huggingface_hub import constants as hf_constants
repo_cache = Path(hf_constants.HF_HUB_CACHE) / ("models--" + model_name.replace("/", "--")) repo_cache = Path(hf_constants.HF_HUB_CACHE) / ("models--" + model_name.replace("/", "--"))
if not repo_cache.exists(): if not repo_cache.exists():
# Start download in background # Start download in background
@@ -1195,8 +1188,8 @@ async def transcribe_audio(
detail={ detail={
"message": f"Whisper model {model_size} is being downloaded. Please wait and try again.", "message": f"Whisper model {model_size} is being downloaded. Please wait and try again.",
"model_name": progress_model_name, "model_name": progress_model_name,
"downloading": True "downloading": True,
} },
) )
text = await whisper_model.transcribe(tmp_path, language) text = await whisper_model.transcribe(tmp_path, language)
@@ -1213,10 +1206,6 @@ async def transcribe_audio(
Path(tmp_path).unlink(missing_ok=True) Path(tmp_path).unlink(missing_ok=True)
# ============================================
# STORY ENDPOINTS
# ============================================
@app.get("/stories", response_model=List[models.StoryResponse]) @app.get("/stories", response_model=List[models.StoryResponse])
async def list_stories(db: Session = Depends(get_db)): async def list_stories(db: Session = Depends(get_db)):
"""List all stories.""" """List all stories."""
@@ -1320,7 +1309,9 @@ async def reorder_story_items(
"""Reorder story items and recalculate timecodes.""" """Reorder story items and recalculate timecodes."""
items = await stories.reorder_story_items(story_id, data.generation_ids, db) items = await stories.reorder_story_items(story_id, data.generation_ids, db)
if items is None: if items is None:
raise HTTPException(status_code=400, detail="Invalid reorder request - ensure all generation IDs belong to this story") raise HTTPException(
status_code=400, detail="Invalid reorder request - ensure all generation IDs belong to this story"
)
return items return items
@@ -1411,7 +1402,7 @@ async def export_story_audio(
raise HTTPException(status_code=400, detail="Story has no audio items") raise HTTPException(status_code=400, detail="Story has no audio items")
# Create safe filename # Create safe filename
safe_name = "".join(c for c in story.name if c.isalnum() or c in (' ', '-', '_')).strip() safe_name = "".join(c for c in story.name if c.isalnum() or c in (" ", "-", "_")).strip()
if not safe_name: if not safe_name:
safe_name = "story" safe_name = "story"
filename = f"{safe_name}.wav" filename = f"{safe_name}.wav"
@@ -1420,9 +1411,7 @@ async def export_story_audio(
return StreamingResponse( return StreamingResponse(
io.BytesIO(audio_bytes), io.BytesIO(audio_bytes),
media_type="audio/wav", media_type="audio/wav",
headers={ headers={"Content-Disposition": _safe_content_disposition("attachment", filename)},
"Content-Disposition": _safe_content_disposition("attachment", filename)
}
) )
except HTTPException: except HTTPException:
raise raise
@@ -1430,10 +1419,6 @@ async def export_story_audio(
raise HTTPException(status_code=500, detail=str(e)) raise HTTPException(status_code=500, detail=str(e))
# ============================================
# EFFECTS & VERSIONS
# ============================================
@app.post("/effects/preview/{generation_id}") @app.post("/effects/preview/{generation_id}")
async def preview_effects( async def preview_effects(
generation_id: str, generation_id: str,
@@ -1473,6 +1458,7 @@ async def preview_effects(
# Write to in-memory buffer # Write to in-memory buffer
import soundfile as sf import soundfile as sf
buf = io.BytesIO() buf = io.BytesIO()
await asyncio.to_thread(lambda: sf.write(buf, processed, sample_rate, format="WAV")) await asyncio.to_thread(lambda: sf.write(buf, processed, sample_rate, format="WAV"))
buf.seek(0) buf.seek(0)
@@ -1491,15 +1477,15 @@ async def preview_effects(
async def get_available_effects(): async def get_available_effects():
"""List all available effect types with parameter definitions.""" """List all available effect types with parameter definitions."""
from .utils.effects import get_available_effects as _get_effects from .utils.effects import get_available_effects as _get_effects
return models.AvailableEffectsResponse(effects=[
models.AvailableEffect(**e) for e in _get_effects() return models.AvailableEffectsResponse(effects=[models.AvailableEffect(**e) for e in _get_effects()])
])
@app.get("/effects/presets", response_model=List[models.EffectPresetResponse]) @app.get("/effects/presets", response_model=List[models.EffectPresetResponse])
async def list_effect_presets(db: Session = Depends(get_db)): async def list_effect_presets(db: Session = Depends(get_db)):
"""List all effect presets (built-in + user-created).""" """List all effect presets (built-in + user-created)."""
from . import effects as effects_mod from . import effects as effects_mod
return effects_mod.list_presets(db) return effects_mod.list_presets(db)
@@ -1507,6 +1493,7 @@ async def list_effect_presets(db: Session = Depends(get_db)):
async def get_effect_preset(preset_id: str, db: Session = Depends(get_db)): async def get_effect_preset(preset_id: str, db: Session = Depends(get_db)):
"""Get a specific effect preset.""" """Get a specific effect preset."""
from . import effects as effects_mod from . import effects as effects_mod
preset = effects_mod.get_preset(preset_id, db) preset = effects_mod.get_preset(preset_id, db)
if not preset: if not preset:
raise HTTPException(status_code=404, detail="Preset not found") raise HTTPException(status_code=404, detail="Preset not found")
@@ -1520,6 +1507,7 @@ async def create_effect_preset(
): ):
"""Create a new effect preset.""" """Create a new effect preset."""
from . import effects as effects_mod from . import effects as effects_mod
try: try:
return effects_mod.create_preset(data, db) return effects_mod.create_preset(data, db)
except ValueError as e: except ValueError as e:
@@ -1534,6 +1522,7 @@ async def update_effect_preset(
): ):
"""Update an effect preset.""" """Update an effect preset."""
from . import effects as effects_mod from . import effects as effects_mod
try: try:
result = effects_mod.update_preset(preset_id, data, db) result = effects_mod.update_preset(preset_id, data, db)
if not result: if not result:
@@ -1547,6 +1536,7 @@ async def update_effect_preset(
async def delete_effect_preset(preset_id: str, db: Session = Depends(get_db)): async def delete_effect_preset(preset_id: str, db: Session = Depends(get_db)):
"""Delete a user effect preset.""" """Delete a user effect preset."""
from . import effects as effects_mod from . import effects as effects_mod
try: try:
if not effects_mod.delete_preset(preset_id, db): if not effects_mod.delete_preset(preset_id, db):
raise HTTPException(status_code=404, detail="Preset not found") raise HTTPException(status_code=404, detail="Preset not found")
@@ -1569,6 +1559,7 @@ async def list_generation_versions(
raise HTTPException(status_code=404, detail="Generation not found") raise HTTPException(status_code=404, detail="Generation not found")
from . import versions as versions_mod from . import versions as versions_mod
return versions_mod.list_versions(generation_id, db) return versions_mod.list_versions(generation_id, db)
@@ -1602,16 +1593,12 @@ async def apply_effects_to_generation(
all_versions = versions_mod.list_versions(generation_id, db) all_versions = versions_mod.list_versions(generation_id, db)
source_version_id = data.source_version_id source_version_id = data.source_version_id
if source_version_id: if source_version_id:
source_version = next( source_version = next((v for v in all_versions if v.id == source_version_id), None)
(v for v in all_versions if v.id == source_version_id), None
)
if not source_version: if not source_version:
raise HTTPException(status_code=404, detail="Source version not found") raise HTTPException(status_code=404, detail="Source version not found")
source_path = source_version.audio_path source_path = source_version.audio_path
else: else:
clean_version = next( clean_version = next((v for v in all_versions if v.effects_chain is None), None)
(v for v in all_versions if v.effects_chain is None), None
)
if not clean_version: if not clean_version:
source_path = gen.audio_path source_path = gen.audio_path
else: else:
@@ -1724,6 +1711,7 @@ async def update_profile_effects(
if data.effects_chain is not None: if data.effects_chain is not None:
from .utils.effects import validate_effects_chain from .utils.effects import 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)
if error: if error:
@@ -1739,10 +1727,6 @@ async def update_profile_effects(
return _profile_to_response(profile) return _profile_to_response(profile)
# ============================================
# FILE SERVING
# ============================================
@app.get("/audio/{generation_id}") @app.get("/audio/{generation_id}")
async def get_audio(generation_id: str, db: Session = Depends(get_db)): async def get_audio(generation_id: str, db: Session = Depends(get_db)):
"""Serve generated audio file (serves the default version).""" """Serve generated audio file (serves the default version)."""
@@ -1781,10 +1765,6 @@ async def get_sample_audio(sample_id: str, db: Session = Depends(get_db)):
) )
# ============================================
# MODEL MANAGEMENT
# ============================================
@app.post("/models/load") @app.post("/models/load")
async def load_model(model_size: str = "1.7B"): async def load_model(model_size: str = "1.7B"):
"""Manually load TTS model.""" """Manually load TTS model."""
@@ -1851,6 +1831,7 @@ async def get_model_progress(model_name: str):
async def get_models_cache_dir(): async def get_models_cache_dir():
"""Get the path to the HuggingFace model cache directory.""" """Get the path to the HuggingFace model cache directory."""
from huggingface_hub import constants as hf_constants from huggingface_hub import constants as hf_constants
return {"path": str(Path(hf_constants.HF_HUB_CACHE))} return {"path": str(Path(hf_constants.HF_HUB_CACHE))}
@@ -1866,6 +1847,7 @@ def _get_dir_size(path: Path) -> int:
def _copy_with_progress(src: Path, dst: Path, progress_manager, copied_so_far: int, total_bytes: int) -> int: def _copy_with_progress(src: Path, dst: Path, progress_manager, copied_so_far: int, total_bytes: int) -> int:
"""Copy a directory tree with byte-level progress tracking.""" """Copy a directory tree with byte-level progress tracking."""
import shutil import shutil
dst.mkdir(parents=True, exist_ok=True) dst.mkdir(parents=True, exist_ok=True)
for item in src.iterdir(): for item in src.iterdir():
dest_item = dst / item.name dest_item = dst / item.name
@@ -1876,8 +1858,11 @@ def _copy_with_progress(src: Path, dst: Path, progress_manager, copied_so_far: i
shutil.copy2(str(item), str(dest_item)) shutil.copy2(str(item), str(dest_item))
copied_so_far += size copied_so_far += size
progress_manager.update_progress( progress_manager.update_progress(
"migration", copied_so_far, total_bytes, "migration",
filename=item.name, status="downloading", copied_so_far,
total_bytes,
filename=item.name,
status="downloading",
) )
return copied_so_far return copied_so_far
@@ -1924,15 +1909,20 @@ async def migrate_models(request: models.ModelMigrateRequest):
shutil.move(str(item), str(dest_item)) shutil.move(str(item), str(dest_item))
moved += 1 moved += 1
progress_manager.update_progress( progress_manager.update_progress(
"migration", i + 1, total, "migration",
filename=item.name, status="downloading", i + 1,
total,
filename=item.name,
status="downloading",
) )
except Exception as e: except Exception as e:
errors.append(f"{item.name}: {str(e)}") errors.append(f"{item.name}: {str(e)}")
else: else:
# Cross-filesystem: copy with byte-level progress, then delete source # Cross-filesystem: copy with byte-level progress, then delete source
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("migration", 0, total_bytes, filename="Calculating...", status="downloading") progress_manager.update_progress(
"migration", 0, total_bytes, filename="Calculating...", status="downloading"
)
copied = 0 copied = 0
for item in model_dirs: for item in model_dirs:
@@ -1997,6 +1987,7 @@ async def get_model_status():
# Try to import scan_cache_dir (might not be available in older versions) # Try to import scan_cache_dir (might not be available in older versions)
try: try:
from huggingface_hub import scan_cache_dir from huggingface_hub import scan_cache_dir
use_scan_cache = True use_scan_cache = True
except ImportError: except ImportError:
use_scan_cache = False use_scan_cache = False
@@ -2050,7 +2041,7 @@ async def get_model_status():
for rev in repo.revisions: for rev in repo.revisions:
for f in rev.files: for f in rev.files:
fname = f.file_name.lower() fname = f.file_name.lower()
if fname.endswith(('.safetensors', '.bin', '.pt', '.pth', '.npz')): if fname.endswith((".safetensors", ".bin", ".pt", ".pth", ".npz")):
has_model_weights = True has_model_weights = True
break break
if has_model_weights: if has_model_weights:
@@ -2095,11 +2086,11 @@ async def get_model_status():
has_model_files = False has_model_files = False
if snapshots_dir.exists(): if snapshots_dir.exists():
has_model_files = ( has_model_files = (
any(snapshots_dir.rglob("*.bin")) or any(snapshots_dir.rglob("*.bin"))
any(snapshots_dir.rglob("*.safetensors")) or or any(snapshots_dir.rglob("*.safetensors"))
any(snapshots_dir.rglob("*.pt")) or or any(snapshots_dir.rglob("*.pt"))
any(snapshots_dir.rglob("*.pth")) or or any(snapshots_dir.rglob("*.pth"))
any(snapshots_dir.rglob("*.npz")) or any(snapshots_dir.rglob("*.npz"))
) )
if has_model_files: if has_model_files:
@@ -2107,8 +2098,9 @@ async def get_model_status():
# Calculate size (exclude .incomplete files) # Calculate size (exclude .incomplete files)
try: try:
total_size = sum( total_size = sum(
f.stat().st_size for f in repo_cache.rglob("*") f.stat().st_size
if f.is_file() and not f.name.endswith('.incomplete') for f in repo_cache.rglob("*")
if f.is_file() and not f.name.endswith(".incomplete")
) )
size_mb = total_size / (1024 * 1024) size_mb = total_size / (1024 * 1024)
except Exception: except Exception:
@@ -2133,7 +2125,8 @@ async def get_model_status():
downloaded = False downloaded = False
size_mb = None # Don't show partial size during download size_mb = None # Don't show partial size during download
statuses.append(models.ModelStatus( statuses.append(
models.ModelStatus(
model_name=config["model_name"], model_name=config["model_name"],
display_name=config["display_name"], display_name=config["display_name"],
hf_repo_id=config["hf_repo_id"], hf_repo_id=config["hf_repo_id"],
@@ -2141,7 +2134,8 @@ async def get_model_status():
downloading=is_downloading, downloading=is_downloading,
size_mb=size_mb, size_mb=size_mb,
loaded=loaded, loaded=loaded,
)) )
)
except Exception as e: except Exception as e:
# If check fails, try to at least check if loaded # If check fails, try to at least check if loaded
try: try:
@@ -2152,7 +2146,8 @@ async def get_model_status():
# Check if this model (or its shared repo) is currently being downloaded # Check if this model (or its shared repo) is currently being downloaded
is_downloading = config["hf_repo_id"] in active_download_repos is_downloading = config["hf_repo_id"] in active_download_repos
statuses.append(models.ModelStatus( statuses.append(
models.ModelStatus(
model_name=config["model_name"], model_name=config["model_name"],
display_name=config["display_name"], display_name=config["display_name"],
hf_repo_id=config["hf_repo_id"], hf_repo_id=config["hf_repo_id"],
@@ -2160,7 +2155,8 @@ async def get_model_status():
downloading=is_downloading, downloading=is_downloading,
size_mb=None, size_mb=None,
loaded=loaded, loaded=loaded,
)) )
)
return models.ModelStatusListResponse(models=statuses) return models.ModelStatusListResponse(models=statuses)
@@ -2281,10 +2277,7 @@ 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( raise HTTPException(status_code=500, detail=f"Failed to delete model cache directory: {str(e)}")
status_code=500,
detail=f"Failed to delete model cache directory: {str(e)}"
)
return {"message": f"Model {model_name} deleted successfully"} return {"message": f"Model {model_name} deleted successfully"}
@@ -2307,10 +2300,6 @@ async def clear_cache():
raise HTTPException(status_code=500, detail=f"Failed to clear cache: {str(e)}") raise HTTPException(status_code=500, detail=f"Failed to clear cache: {str(e)}")
# ============================================
# TASK MANAGEMENT
# ============================================
@app.get("/tasks/active", response_model=models.ActiveTasksResponse) @app.get("/tasks/active", response_model=models.ActiveTasksResponse)
async def get_active_tasks(): async def get_active_tasks():
"""Return all currently active downloads and generations.""" """Return all currently active downloads and generations."""
@@ -2349,7 +2338,8 @@ async def get_active_tasks():
pm_data = progress_manager._progress.get(model_name) pm_data = progress_manager._progress.get(model_name)
if pm_data: if pm_data:
prog = pm_data prog = pm_data
active_downloads.append(models.ActiveDownloadTask( active_downloads.append(
models.ActiveDownloadTask(
model_name=model_name, model_name=model_name,
status=task.status, status=task.status,
started_at=task.started_at, started_at=task.started_at,
@@ -2358,19 +2348,21 @@ async def get_active_tasks():
current=prog.get("current"), current=prog.get("current"),
total=prog.get("total"), total=prog.get("total"),
filename=prog.get("filename"), filename=prog.get("filename"),
)) )
)
elif progress: elif progress:
# Progress exists but no task - create from progress data # Progress exists but no task - create from progress data
timestamp_str = progress.get("timestamp") timestamp_str = progress.get("timestamp")
if timestamp_str: if timestamp_str:
try: try:
started_at = datetime.fromisoformat(timestamp_str.replace('Z', '+00:00')) started_at = datetime.fromisoformat(timestamp_str.replace("Z", "+00:00"))
except (ValueError, AttributeError): except (ValueError, AttributeError):
started_at = datetime.utcnow() started_at = datetime.utcnow()
else: else:
started_at = datetime.utcnow() started_at = datetime.utcnow()
active_downloads.append(models.ActiveDownloadTask( active_downloads.append(
models.ActiveDownloadTask(
model_name=model_name, model_name=model_name,
status=progress.get("status", "downloading"), status=progress.get("status", "downloading"),
started_at=started_at, started_at=started_at,
@@ -2379,17 +2371,20 @@ async def get_active_tasks():
current=progress.get("current"), current=progress.get("current"),
total=progress.get("total"), total=progress.get("total"),
filename=progress.get("filename"), filename=progress.get("filename"),
)) )
)
# Get active generations # Get active generations
active_generations = [] active_generations = []
for gen_task in task_manager.get_active_generations(): for gen_task in task_manager.get_active_generations():
active_generations.append(models.ActiveGenerationTask( active_generations.append(
models.ActiveGenerationTask(
task_id=gen_task.task_id, task_id=gen_task.task_id,
profile_id=gen_task.profile_id, profile_id=gen_task.profile_id,
text_preview=gen_task.text_preview, text_preview=gen_task.text_preview,
started_at=gen_task.started_at, started_at=gen_task.started_at,
)) )
)
return models.ActiveTasksResponse( return models.ActiveTasksResponse(
downloads=active_downloads, downloads=active_downloads,
@@ -2397,14 +2392,11 @@ async def get_active_tasks():
) )
# ============================================
# CUDA BACKEND MANAGEMENT
# ============================================
@app.get("/backend/cuda-status") @app.get("/backend/cuda-status")
async def get_cuda_status(): async def get_cuda_status():
"""Get CUDA backend download/availability status.""" """Get CUDA backend download/availability status."""
from . import cuda_download from . import cuda_download
return cuda_download.get_cuda_status() return cuda_download.get_cuda_status()
@@ -2422,6 +2414,7 @@ async def download_cuda_backend():
await cuda_download.download_cuda_binary() await cuda_download.download_cuda_binary()
except Exception as e: except Exception as e:
import logging import logging
logging.getLogger(__name__).error(f"CUDA download failed: {e}") logging.getLogger(__name__).error(f"CUDA download failed: {e}")
create_background_task(_download()) create_background_task(_download())
@@ -2466,21 +2459,17 @@ async def get_cuda_download_progress():
) )
# ============================================
# STARTUP & SHUTDOWN
# ============================================
def _get_gpu_status() -> str: def _get_gpu_status() -> str:
"""Get GPU availability status.""" """Get GPU availability status."""
backend_type = get_backend_type() backend_type = get_backend_type()
if torch.cuda.is_available(): if torch.cuda.is_available():
device_name = torch.cuda.get_device_name(0) device_name = torch.cuda.get_device_name(0)
# Check if this is ROCm (AMD) or CUDA (NVIDIA) # Check if this is ROCm (AMD) or CUDA (NVIDIA)
is_rocm = hasattr(torch.version, 'hip') and torch.version.hip is not None is_rocm = hasattr(torch.version, "hip") and torch.version.hip is not None
if is_rocm: if is_rocm:
return f"ROCm ({device_name})" return f"ROCm ({device_name})"
return f"CUDA ({device_name})" return f"CUDA ({device_name})"
elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
return "MPS (Apple Silicon)" return "MPS (Apple Silicon)"
elif backend_type == "mlx": elif backend_type == "mlx":
return "Metal (Apple Silicon via MLX)" return "Metal (Apple Silicon via MLX)"
@@ -2501,9 +2490,12 @@ async def startup_event():
# from a previous process that was killed mid-generation # from a previous process that was killed mid-generation
try: try:
from sqlalchemy import text as sa_text from sqlalchemy import text as sa_text
db = next(get_db()) db = next(get_db())
result = db.execute( result = db.execute(
sa_text("UPDATE generations SET status = 'failed', error = 'Server was shut down during generation' WHERE status = 'generating'") sa_text(
"UPDATE generations SET status = 'failed', error = 'Server was shut down during generation' WHERE status = 'generating'"
)
) )
if result.rowcount > 0: if result.rowcount > 0:
print(f"Marked {result.rowcount} stale generation(s) as failed") print(f"Marked {result.rowcount} stale generation(s) as failed")
@@ -2517,6 +2509,7 @@ async def startup_event():
# Auto-update CUDA binary if installed but outdated # Auto-update CUDA binary if installed but outdated
from .cuda_download import check_and_update_cuda_binary from .cuda_download import check_and_update_cuda_binary
create_background_task(check_and_update_cuda_binary()) create_background_task(check_and_update_cuda_binary())
# Initialize progress manager with main event loop for thread-safe operations # Initialize progress manager with main event loop for thread-safe operations
@@ -2530,6 +2523,7 @@ async def startup_event():
# Ensure HuggingFace cache directory exists # Ensure HuggingFace cache directory exists
try: try:
from huggingface_hub import constants as hf_constants from huggingface_hub import constants as hf_constants
cache_dir = Path(hf_constants.HF_HUB_CACHE) cache_dir = Path(hf_constants.HF_HUB_CACHE)
cache_dir.mkdir(parents=True, exist_ok=True) cache_dir.mkdir(parents=True, exist_ok=True)
print(f"HuggingFace cache directory: {cache_dir}") print(f"HuggingFace cache directory: {cache_dir}")
@@ -2547,10 +2541,6 @@ async def shutdown_event():
transcribe.unload_whisper_model() transcribe.unload_whisper_model()
# ============================================
# MAIN
# ============================================
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser(description="voicebox backend server") parser = argparse.ArgumentParser(description="voicebox backend server")
parser.add_argument( parser.add_argument(
+65 -9
View File
@@ -9,13 +9,17 @@ from datetime import datetime
class VoiceProfileCreate(BaseModel): 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: Optional[str] = Field(None, max_length=500)
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)$"
)
class VoiceProfileResponse(BaseModel): class VoiceProfileResponse(BaseModel):
"""Response model for voice profile.""" """Response model for voice profile."""
id: str id: str
name: str name: str
description: Optional[str] description: Optional[str]
@@ -33,16 +37,19 @@ class VoiceProfileResponse(BaseModel):
class ProfileSampleCreate(BaseModel): class ProfileSampleCreate(BaseModel):
"""Request model for adding a sample to a profile.""" """Request model for adding a sample to a profile."""
reference_text: str = Field(..., min_length=1, max_length=1000) reference_text: str = Field(..., min_length=1, max_length=1000)
class ProfileSampleUpdate(BaseModel): class ProfileSampleUpdate(BaseModel):
"""Request model for updating a profile sample.""" """Request model for updating a profile sample."""
reference_text: str = Field(..., min_length=1, max_length=1000) reference_text: str = Field(..., min_length=1, max_length=1000)
class ProfileSampleResponse(BaseModel): class ProfileSampleResponse(BaseModel):
"""Response model for profile sample.""" """Response model for profile sample."""
id: str id: str
profile_id: str profile_id: str
audio_path: str audio_path: str
@@ -54,6 +61,7 @@ class ProfileSampleResponse(BaseModel):
class GenerationRequest(BaseModel): class GenerationRequest(BaseModel):
"""Request model for voice generation.""" """Request model for voice generation."""
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)$") language: str = Field(default="en", pattern="^(zh|en|ja|ko|de|fr|ru|pt|es|it|he)$")
@@ -61,14 +69,21 @@ class GenerationRequest(BaseModel):
model_size: Optional[str] = Field(default="1.7B", pattern="^(1\\.7B|0\\.6B)$") model_size: Optional[str] = Field(default="1.7B", pattern="^(1\\.7B|0\\.6B)$")
instruct: Optional[str] = Field(None, max_length=500) instruct: Optional[str] = Field(None, max_length=500)
engine: Optional[str] = Field(default="qwen", pattern="^(qwen|luxtts|chatterbox|chatterbox_turbo)$") engine: Optional[str] = Field(default="qwen", pattern="^(qwen|luxtts|chatterbox|chatterbox_turbo)$")
max_chunk_chars: int = Field(default=800, ge=100, le=5000, description="Max characters per chunk for long text splitting") max_chunk_chars: int = Field(
crossfade_ms: int = Field(default=50, ge=0, le=500, description="Crossfade duration in ms between chunks (0 for hard cut)") default=800, ge=100, le=5000, description="Max characters per chunk for long text splitting"
)
crossfade_ms: int = Field(
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(None, description="Effects chain to apply after generation (overrides profile default)") effects_chain: Optional[List["EffectConfig"]] = Field(
None, description="Effects chain to apply after generation (overrides profile default)"
)
class GenerationResponse(BaseModel): class GenerationResponse(BaseModel):
"""Response model for voice generation.""" """Response model for voice generation."""
id: str id: str
profile_id: str profile_id: str
text: str text: str
@@ -92,6 +107,7 @@ 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: Optional[str] = None
search: Optional[str] = None search: Optional[str] = None
limit: int = Field(default=50, ge=1, le=100) limit: int = Field(default=50, ge=1, le=100)
@@ -100,6 +116,7 @@ class HistoryQuery(BaseModel):
class HistoryResponse(BaseModel): class HistoryResponse(BaseModel):
"""Response model for history entry (includes profile name).""" """Response model for history entry (includes profile name)."""
id: str id: str
profile_id: str profile_id: str
profile_name: str profile_name: str
@@ -124,23 +141,27 @@ 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)$") language: Optional[str] = Field(None, pattern="^(en|zh)$")
class TranscriptionResponse(BaseModel): class TranscriptionResponse(BaseModel):
"""Response model for transcription.""" """Response model for transcription."""
text: str text: str
duration: float duration: float
class HealthResponse(BaseModel): class HealthResponse(BaseModel):
"""Response model for health check.""" """Response model for health check."""
status: str status: str
model_loaded: bool model_loaded: bool
model_downloaded: Optional[bool] = None # Whether model is cached/downloaded model_downloaded: Optional[bool] = None # Whether model is cached/downloaded
@@ -154,6 +175,7 @@ class HealthResponse(BaseModel):
class DirectoryCheck(BaseModel): class DirectoryCheck(BaseModel):
"""Health status for a single directory.""" """Health status for a single directory."""
path: str path: str
exists: bool exists: bool
writable: bool writable: bool
@@ -162,6 +184,7 @@ class DirectoryCheck(BaseModel):
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: Optional[float] = None
disk_total_mb: Optional[float] = None disk_total_mb: Optional[float] = None
@@ -170,6 +193,7 @@ class FilesystemHealthResponse(BaseModel):
class ModelStatus(BaseModel): class ModelStatus(BaseModel):
"""Response model for model status.""" """Response model for model status."""
model_name: str model_name: str
display_name: str display_name: str
hf_repo_id: Optional[str] = None # HuggingFace repository ID hf_repo_id: Optional[str] = None # HuggingFace repository ID
@@ -181,21 +205,25 @@ class ModelStatus(BaseModel):
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):
"""Request model for triggering model download.""" """Request model for triggering model download."""
model_name: str model_name: str
class ModelMigrateRequest(BaseModel): class ModelMigrateRequest(BaseModel):
"""Request model for migrating models to a new directory.""" """Request model for migrating models to a new directory."""
destination: str destination: str
class ActiveDownloadTask(BaseModel): class ActiveDownloadTask(BaseModel):
"""Response model for active download task.""" """Response model for active download task."""
model_name: str model_name: str
status: str status: str
started_at: datetime started_at: datetime
@@ -208,6 +236,7 @@ class ActiveDownloadTask(BaseModel):
class ActiveGenerationTask(BaseModel): class ActiveGenerationTask(BaseModel):
"""Response model for active generation task.""" """Response model for active generation task."""
task_id: str task_id: str
profile_id: str profile_id: str
text_preview: str text_preview: str
@@ -216,24 +245,28 @@ 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: Optional[str] = Field(None, min_length=1, max_length=100)
device_ids: Optional[List[str]] = None device_ids: Optional[List[str]] = None
class AudioChannelResponse(BaseModel): class AudioChannelResponse(BaseModel):
"""Response model for audio channel.""" """Response model for audio channel."""
id: str id: str
name: str name: str
is_default: bool is_default: bool
@@ -246,22 +279,26 @@ 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: Optional[str] = Field(None, max_length=500)
class StoryResponse(BaseModel): class StoryResponse(BaseModel):
"""Response model for story (list view).""" """Response model for story (list view)."""
id: str id: str
name: str name: str
description: Optional[str] description: Optional[str]
@@ -275,6 +312,7 @@ class StoryResponse(BaseModel):
class StoryItemDetail(BaseModel): class StoryItemDetail(BaseModel):
"""Detail model for story item with generation info.""" """Detail model for story item with generation info."""
id: str id: str
story_id: str story_id: str
generation_id: str generation_id: str
@@ -304,6 +342,7 @@ class StoryItemDetail(BaseModel):
class StoryDetailResponse(BaseModel): class StoryDetailResponse(BaseModel):
"""Response model for story with items.""" """Response model for story with items."""
id: str id: str
name: str name: str
description: Optional[str] description: Optional[str]
@@ -317,6 +356,7 @@ class StoryDetailResponse(BaseModel):
class StoryItemCreate(BaseModel): 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: Optional[int] = None # If not provided, will be calculated automatically
track: Optional[int] = 0 # Track number (0 = main track) track: Optional[int] = 0 # Track number (0 = main track)
@@ -324,48 +364,52 @@ class StoryItemCreate(BaseModel):
class StoryItemUpdateTime(BaseModel): class StoryItemUpdateTime(BaseModel):
"""Request model for updating a story item's timecode.""" """Request model for updating a story item's timecode."""
generation_id: str generation_id: str
start_time_ms: int = Field(..., ge=0) start_time_ms: int = Field(..., ge=0)
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):
"""Request model for moving a story item (position and/or track).""" """Request model for moving a story item (position and/or track)."""
start_time_ms: int = Field(..., ge=0) start_time_ms: int = Field(..., ge=0)
track: int = 0 track: int = 0
class StoryItemTrim(BaseModel): class StoryItemTrim(BaseModel):
"""Request model for trimming a story item.""" """Request model for trimming a story item."""
trim_start_ms: int = Field(..., ge=0) trim_start_ms: int = Field(..., ge=0)
trim_end_ms: int = Field(..., ge=0) trim_end_ms: int = Field(..., ge=0)
class StoryItemSplit(BaseModel): class StoryItemSplit(BaseModel):
"""Request model for splitting a story item.""" """Request model for splitting a story item."""
split_time_ms: int = Field(..., ge=0) # Time within the clip to split at (relative to clip start) split_time_ms: int = Field(..., ge=0) # Time within the clip to split at (relative to clip start)
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: Optional[str] = None # null = use generation default
# ============================================
# Effects & Versions
# ============================================
class EffectConfig(BaseModel): class EffectConfig(BaseModel):
"""A single effect in an effects chain.""" """A single effect in an effects chain."""
type: str type: str
enabled: bool = True enabled: bool = True
params: dict = Field(default_factory=dict) params: dict = Field(default_factory=dict)
@@ -373,11 +417,13 @@ 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: Optional[str] = Field(None, max_length=500)
effects_chain: List[EffectConfig] effects_chain: List[EffectConfig]
@@ -385,6 +431,7 @@ class EffectPresetCreate(BaseModel):
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: Optional[str] = Field(None, min_length=1, max_length=100)
description: Optional[str] = None description: Optional[str] = None
effects_chain: Optional[List[EffectConfig]] = None effects_chain: Optional[List[EffectConfig]] = None
@@ -392,6 +439,7 @@ class EffectPresetUpdate(BaseModel):
class EffectPresetResponse(BaseModel): class EffectPresetResponse(BaseModel):
"""Response model for effect preset.""" """Response model for effect preset."""
id: str id: str
name: str name: str
description: Optional[str] = None description: Optional[str] = None
@@ -405,6 +453,7 @@ class EffectPresetResponse(BaseModel):
class GenerationVersionResponse(BaseModel): class GenerationVersionResponse(BaseModel):
"""Response model for a generation version.""" """Response model for a generation version."""
id: str id: str
generation_id: str generation_id: str
label: str label: str
@@ -420,19 +469,24 @@ 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(None, description="Version to use as source audio (defaults to clean/original)") source_version_id: Optional[str] = Field(
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: Optional[str] = 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: Optional[List[EffectConfig]] = Field(None, description="Effects chain (null to remove)")
class AvailableEffectParam(BaseModel): class AvailableEffectParam(BaseModel):
"""Description of a single effect parameter.""" """Description of a single effect parameter."""
default: float default: float
min: float min: float
max: float max: float
@@ -442,6 +496,7 @@ class AvailableEffectParam(BaseModel):
class AvailableEffect(BaseModel): class AvailableEffect(BaseModel):
"""Description of an available effect type.""" """Description of an available effect type."""
type: str type: str
label: str label: str
description: str description: str
@@ -450,4 +505,5 @@ 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]
+10 -47
View File
@@ -43,6 +43,7 @@ def _profile_to_response(
effects_chain = [EffectConfig(**e) for e in raw] effects_chain = [EffectConfig(**e) for e in raw]
except Exception as e: except Exception as e:
import logging import logging
logging.warning(f"Failed to parse effects_chain for profile {profile.id}: {e}") logging.warning(f"Failed to parse effects_chain for profile {profile.id}: {e}")
return VoiceProfileResponse( return VoiceProfileResponse(
id=profile.id, id=profile.id,
@@ -75,12 +76,10 @@ async def create_profile(
Raises: Raises:
ValueError: If a profile with the same name already exists ValueError: If a profile with the same name already exists
""" """
# Check if profile name already exists
existing_profile = db.query(DBVoiceProfile).filter_by(name=data.name).first() existing_profile = db.query(DBVoiceProfile).filter_by(name=data.name).first()
if existing_profile: if existing_profile:
raise ValueError(f"A profile with the name '{data.name}' already exists. Please choose a different name.") raise ValueError(f"A profile with the name '{data.name}' already exists. Please choose a different name.")
# Create profile in database
db_profile = DBVoiceProfile( db_profile = DBVoiceProfile(
id=str(uuid.uuid4()), id=str(uuid.uuid4()),
name=data.name, name=data.name,
@@ -94,7 +93,6 @@ async def create_profile(
db.commit() db.commit()
db.refresh(db_profile) db.refresh(db_profile)
# Create profile directory
profile_dir = config.get_profiles_dir() / db_profile.id profile_dir = config.get_profiles_dir() / db_profile.id
profile_dir.mkdir(parents=True, exist_ok=True) profile_dir.mkdir(parents=True, exist_ok=True)
@@ -119,27 +117,22 @@ async def add_profile_sample(
Returns: Returns:
Created sample Created sample
""" """
# Validate profile exists
profile = db.query(DBVoiceProfile).filter_by(id=profile_id).first() profile = db.query(DBVoiceProfile).filter_by(id=profile_id).first()
if not profile: if not profile:
raise ValueError(f"Profile {profile_id} not found") raise ValueError(f"Profile {profile_id} not found")
# Validate audio
is_valid, error_msg = validate_reference_audio(audio_path) is_valid, error_msg = validate_reference_audio(audio_path)
if not is_valid: if not is_valid:
raise ValueError(f"Invalid reference audio: {error_msg}") raise ValueError(f"Invalid reference audio: {error_msg}")
# Create sample ID and directory
sample_id = str(uuid.uuid4()) sample_id = str(uuid.uuid4())
profile_dir = config.get_profiles_dir() / profile_id profile_dir = config.get_profiles_dir() / profile_id
profile_dir.mkdir(parents=True, exist_ok=True) profile_dir.mkdir(parents=True, exist_ok=True)
# Copy audio file to profile directory
dest_path = profile_dir / f"{sample_id}.wav" dest_path = profile_dir / f"{sample_id}.wav"
audio, sr = load_audio(audio_path) audio, sr = load_audio(audio_path)
save_audio(audio, str(dest_path), sr) save_audio(audio, str(dest_path), sr)
# Create database entry
db_sample = DBProfileSample( db_sample = DBProfileSample(
id=sample_id, id=sample_id,
profile_id=profile_id, profile_id=profile_id,
@@ -149,7 +142,6 @@ async def add_profile_sample(
db.add(db_sample) db.add(db_sample)
# Update profile timestamp
profile.updated_at = datetime.utcnow() profile.updated_at = datetime.utcnow()
db.commit() db.commit()
@@ -211,26 +203,20 @@ async def list_profiles(db: Session) -> List[VoiceProfileResponse]:
Returns: Returns:
List of profiles List of profiles
""" """
profiles = db.query(DBVoiceProfile).order_by( profiles = db.query(DBVoiceProfile).order_by(DBVoiceProfile.created_at.desc()).all()
DBVoiceProfile.created_at.desc()
).all()
if not profiles: if not profiles:
return [] return []
# Batch-fetch generation counts # Batch-fetch generation counts
gen_counts_rows = ( gen_counts_rows = (
db.query(DBGeneration.profile_id, func.count(DBGeneration.id)) db.query(DBGeneration.profile_id, func.count(DBGeneration.id)).group_by(DBGeneration.profile_id).all()
.group_by(DBGeneration.profile_id)
.all()
) )
gen_counts = {row[0]: row[1] for row in gen_counts_rows} gen_counts = {row[0]: row[1] for row in gen_counts_rows}
# Batch-fetch sample counts # Batch-fetch sample counts
sample_counts_rows = ( sample_counts_rows = (
db.query(DBProfileSample.profile_id, func.count(DBProfileSample.id)) db.query(DBProfileSample.profile_id, func.count(DBProfileSample.id)).group_by(DBProfileSample.profile_id).all()
.group_by(DBProfileSample.profile_id)
.all()
) )
sample_counts = {row[0]: row[1] for row in sample_counts_rows} sample_counts = {row[0]: row[1] for row in sample_counts_rows}
@@ -267,13 +253,11 @@ async def update_profile(
if not profile: if not profile:
return None return None
# Check if the new name conflicts with another profile
if profile.name != data.name: if profile.name != data.name:
existing_profile = db.query(DBVoiceProfile).filter_by(name=data.name).first() existing_profile = db.query(DBVoiceProfile).filter_by(name=data.name).first()
if existing_profile: if existing_profile:
raise ValueError(f"A profile with the name '{data.name}' already exists. Please choose a different name.") raise ValueError(f"A profile with the name '{data.name}' already exists. Please choose a different name.")
# Update fields
profile.name = data.name profile.name = data.name
profile.description = data.description profile.description = data.description
profile.language = data.language profile.language = data.language
@@ -303,14 +287,11 @@ async def delete_profile(
if not profile: if not profile:
return False return False
# Delete samples from database
db.query(DBProfileSample).filter_by(profile_id=profile_id).delete() db.query(DBProfileSample).filter_by(profile_id=profile_id).delete()
# Delete profile from database
db.delete(profile) db.delete(profile)
db.commit() db.commit()
# Delete profile directory
profile_dir = config.get_profiles_dir() / profile_id profile_dir = config.get_profiles_dir() / profile_id
if profile_dir.exists(): if profile_dir.exists():
shutil.rmtree(profile_dir) shutil.rmtree(profile_dir)
@@ -342,12 +323,10 @@ async def delete_profile_sample(
# Store profile_id before deleting # Store profile_id before deleting
profile_id = sample.profile_id profile_id = sample.profile_id
# Delete audio file
audio_path = Path(sample.audio_path) audio_path = Path(sample.audio_path)
if audio_path.exists(): if audio_path.exists():
audio_path.unlink() audio_path.unlink()
# Delete from database
db.delete(sample) db.delete(sample)
db.commit() db.commit()
@@ -412,7 +391,6 @@ async def create_voice_prompt_for_profile(
""" """
from .backends import get_tts_backend_for_engine from .backends import get_tts_backend_for_engine
# Get all samples for profile
samples = db.query(DBProfileSample).filter_by(profile_id=profile_id).all() samples = db.query(DBProfileSample).filter_by(profile_id=profile_id).all()
if not samples: if not samples:
@@ -421,7 +399,6 @@ async def create_voice_prompt_for_profile(
tts_model = get_tts_backend_for_engine(engine) tts_model = get_tts_backend_for_engine(engine)
if len(samples) == 1: if len(samples) == 1:
# Single sample - use directly
sample = samples[0] sample = samples[0]
voice_prompt, _ = await tts_model.create_voice_prompt( voice_prompt, _ = await tts_model.create_voice_prompt(
sample.audio_path, sample.audio_path,
@@ -430,11 +407,9 @@ async def create_voice_prompt_for_profile(
) )
return voice_prompt return voice_prompt
else: else:
# Multiple samples - combine them
audio_paths = [s.audio_path for s in samples] audio_paths = [s.audio_path for s in samples]
reference_texts = [s.reference_text for s in samples] reference_texts = [s.reference_text for s in samples]
# Combine audio
combined_audio, combined_text = await tts_model.combine_voice_prompts( combined_audio, combined_text = await tts_model.combine_voice_prompts(
audio_paths, audio_paths,
reference_texts, reference_texts,
@@ -443,18 +418,16 @@ async def create_voice_prompt_for_profile(
# Save combined audio to cache directory (persistent) # Save combined audio to cache directory (persistent)
# Create a hash of sample IDs to identify this specific combination # Create a hash of sample IDs to identify this specific combination
import hashlib import hashlib
sample_ids_str = "-".join(sorted([s.id for s in samples])) sample_ids_str = "-".join(sorted([s.id for s in samples]))
combination_hash = hashlib.md5(sample_ids_str.encode()).hexdigest()[:12] combination_hash = hashlib.md5(sample_ids_str.encode()).hexdigest()[:12]
# Store in cache directory
cache_dir = _get_cache_dir() cache_dir = _get_cache_dir()
cache_dir.mkdir(parents=True, exist_ok=True) cache_dir.mkdir(parents=True, exist_ok=True)
combined_path = cache_dir / f"combined_{profile_id}_{combination_hash}.wav" combined_path = cache_dir / f"combined_{profile_id}_{combination_hash}.wav"
# Save combined audio
save_audio(combined_audio, str(combined_path), 24000) save_audio(combined_audio, str(combined_path), 24000)
# Create prompt from combined audio
voice_prompt, _ = await tts_model.create_voice_prompt( voice_prompt, _ = await tts_model.create_voice_prompt(
str(combined_path), str(combined_path),
combined_text, combined_text,
@@ -479,17 +452,14 @@ async def upload_avatar(
Returns: Returns:
Updated profile Updated profile
""" """
# Validate profile exists
profile = db.query(DBVoiceProfile).filter_by(id=profile_id).first() profile = db.query(DBVoiceProfile).filter_by(id=profile_id).first()
if not profile: if not profile:
raise ValueError(f"Profile {profile_id} not found") raise ValueError(f"Profile {profile_id} not found")
# Validate image
is_valid, error_msg = validate_image(image_path) is_valid, error_msg = validate_image(image_path)
if not is_valid: if not is_valid:
raise ValueError(error_msg) raise ValueError(error_msg)
# Delete existing avatar if present
if profile.avatar_path: if profile.avatar_path:
old_avatar = Path(profile.avatar_path) old_avatar = Path(profile.avatar_path)
if old_avatar.exists(): if old_avatar.exists():
@@ -497,27 +467,22 @@ async def upload_avatar(
# Determine file extension from uploaded file # Determine file extension from uploaded file
from PIL import Image from PIL import Image
with Image.open(image_path) as img: with Image.open(image_path) as img:
# Normalize JPEG variants (MPO is multi-picture format from some cameras) # Normalize JPEG variants (MPO is multi-picture format from some cameras)
img_format = img.format img_format = img.format
if img_format in ('MPO', 'JPG'): if img_format in ("MPO", "JPG"):
img_format = 'JPEG' img_format = "JPEG"
ext_map = { ext_map = {"PNG": ".png", "JPEG": ".jpg", "WEBP": ".webp"}
'PNG': '.png', ext = ext_map.get(img_format, ".png")
'JPEG': '.jpg',
'WEBP': '.webp'
}
ext = ext_map.get(img_format, '.png')
# Save processed image to profile directory
profile_dir = config.get_profiles_dir() / profile_id profile_dir = config.get_profiles_dir() / profile_id
profile_dir.mkdir(parents=True, exist_ok=True) profile_dir.mkdir(parents=True, exist_ok=True)
output_path = profile_dir / f"avatar{ext}" output_path = profile_dir / f"avatar{ext}"
process_avatar(image_path, str(output_path)) process_avatar(image_path, str(output_path))
# Update database
profile.avatar_path = str(output_path) profile.avatar_path = str(output_path)
profile.updated_at = datetime.utcnow() profile.updated_at = datetime.utcnow()
@@ -545,12 +510,10 @@ async def delete_avatar(
if not profile or not profile.avatar_path: if not profile or not profile.avatar_path:
return False return False
# Delete avatar file
avatar_path = Path(profile.avatar_path) avatar_path = Path(profile.avatar_path)
if avatar_path.exists(): if avatar_path.exists():
avatar_path.unlink() avatar_path.unlink()
# Update database
profile.avatar_path = None profile.avatar_path = None
profile.updated_at = datetime.utcnow() profile.updated_at = datetime.utcnow()
+3 -14
View File
@@ -81,9 +81,7 @@ async def run_generation(
if crossfade_ms is not None: if crossfade_ms is not None:
gen_kwargs["crossfade_ms"] = crossfade_ms gen_kwargs["crossfade_ms"] = crossfade_ms
audio, sample_rate = await generate_chunked( audio, sample_rate = await generate_chunked(tts_model, text, voice_prompt, **gen_kwargs)
tts_model, text, voice_prompt, **gen_kwargs
)
# --- Normalize (generate and regenerate always; retry skips) ----- # --- Normalize (generate and regenerate always; retry skips) -----
if normalize or mode == "regenerate": if normalize or mode == "regenerate":
@@ -139,11 +137,6 @@ async def run_generation(
bg_db.close() bg_db.close()
# ---------------------------------------------------------------------
# Mode-specific save helpers (sync, return final audio path)
# ---------------------------------------------------------------------
def _save_generate( def _save_generate(
*, *,
generation_id: str, generation_id: str,
@@ -163,9 +156,7 @@ def _save_generate(
clean_audio_path = config.get_generations_dir() / f"{generation_id}.wav" clean_audio_path = config.get_generations_dir() / f"{generation_id}.wav"
save_audio(audio, str(clean_audio_path), sample_rate) save_audio(audio, str(clean_audio_path), sample_rate)
has_effects = effects_chain and any( has_effects = effects_chain and any(e.get("enabled", True) for e in effects_chain)
e.get("enabled", True) for e in effects_chain
)
versions_mod.create_version( versions_mod.create_version(
generation_id=generation_id, generation_id=generation_id,
@@ -186,9 +177,7 @@ def _save_generate(
print(f"Warning: invalid effects chain, skipping: {error_msg}") print(f"Warning: invalid effects chain, skipping: {error_msg}")
else: else:
processed_audio = apply_effects(audio, sample_rate, effects_chain) processed_audio = apply_effects(audio, sample_rate, effects_chain)
processed_path = ( processed_path = config.get_generations_dir() / f"{generation_id}_processed.wav"
config.get_generations_dir() / f"{generation_id}_processed.wav"
)
save_audio(processed_audio, str(processed_path), sample_rate) save_audio(processed_audio, str(processed_path), sample_rate)
final_audio_path = str(processed_path) final_audio_path = str(processed_path)
versions_mod.create_version( versions_mod.create_version(
+111 -103
View File
@@ -22,7 +22,12 @@ from .models import (
StoryItemSplit, StoryItemSplit,
StoryItemVersionUpdate, StoryItemVersionUpdate,
) )
from .database import Story as DBStory, StoryItem as DBStoryItem, Generation as DBGeneration, VoiceProfile as DBVoiceProfile from .database import (
Story as DBStory,
StoryItem as DBStoryItem,
Generation as DBGeneration,
VoiceProfile as DBVoiceProfile,
)
from .history import _get_versions_for_generation from .history import _get_versions_for_generation
from .utils.audio import load_audio, save_audio from .utils.audio import load_audio, save_audio
import numpy as np import numpy as np
@@ -49,11 +54,11 @@ def _build_item_detail(
id=item.id, id=item.id,
story_id=item.story_id, story_id=item.story_id,
generation_id=item.generation_id, generation_id=item.generation_id,
version_id=getattr(item, 'version_id', None), version_id=getattr(item, "version_id", None),
start_time_ms=item.start_time_ms, start_time_ms=item.start_time_ms,
track=item.track, track=item.track,
trim_start_ms=getattr(item, 'trim_start_ms', 0), trim_start_ms=getattr(item, "trim_start_ms", 0),
trim_end_ms=getattr(item, 'trim_end_ms', 0), trim_end_ms=getattr(item, "trim_end_ms", 0),
created_at=item.created_at, created_at=item.created_at,
profile_id=generation.profile_id, profile_id=generation.profile_id,
profile_name=profile_name, profile_name=profile_name,
@@ -95,10 +100,7 @@ async def create_story(
db.commit() db.commit()
db.refresh(db_story) db.refresh(db_story)
# Get item count item_count = db.query(func.count(DBStoryItem.id)).filter(DBStoryItem.story_id == db_story.id).scalar()
item_count = db.query(func.count(DBStoryItem.id)).filter(
DBStoryItem.story_id == db_story.id
).scalar()
response = StoryResponse.model_validate(db_story) response = StoryResponse.model_validate(db_story)
response.item_count = item_count response.item_count = item_count
@@ -121,9 +123,7 @@ async def list_stories(
result = [] result = []
for story in stories: for story in stories:
item_count = db.query(func.count(DBStoryItem.id)).filter( item_count = db.query(func.count(DBStoryItem.id)).filter(DBStoryItem.story_id == story.id).scalar()
DBStoryItem.story_id == story.id
).scalar()
response = StoryResponse.model_validate(story) response = StoryResponse.model_validate(story)
response.item_count = item_count response.item_count = item_count
@@ -150,22 +150,15 @@ async def get_story(
if not story: if not story:
return None return None
# Get all items ordered by start_time_ms items = (
items = db.query( db.query(DBStoryItem, DBGeneration, DBVoiceProfile.name.label("profile_name"))
DBStoryItem, .join(DBGeneration, DBStoryItem.generation_id == DBGeneration.id)
DBGeneration, .join(DBVoiceProfile, DBGeneration.profile_id == DBVoiceProfile.id)
DBVoiceProfile.name.label('profile_name') .filter(DBStoryItem.story_id == story_id)
).join( .order_by(DBStoryItem.start_time_ms)
DBGeneration, .all()
DBStoryItem.generation_id == DBGeneration.id )
).join(
DBVoiceProfile,
DBGeneration.profile_id == DBVoiceProfile.id
).filter(
DBStoryItem.story_id == story_id
).order_by(DBStoryItem.start_time_ms).all()
# Build item details
item_details = [] item_details = []
for item, generation, profile_name in items: for item, generation, profile_name in items:
item_details.append(_build_item_detail(item, generation, profile_name, db)) item_details.append(_build_item_detail(item, generation, profile_name, db))
@@ -202,10 +195,7 @@ async def update_story(
db.commit() db.commit()
db.refresh(story) db.refresh(story)
# Get item count item_count = db.query(func.count(DBStoryItem.id)).filter(DBStoryItem.story_id == story.id).scalar()
item_count = db.query(func.count(DBStoryItem.id)).filter(
DBStoryItem.story_id == story.id
).scalar()
response = StoryResponse.model_validate(story) response = StoryResponse.model_validate(story)
response.item_count = item_count response.item_count = item_count
@@ -267,10 +257,7 @@ async def add_item_to_story(
return None return None
# Check if generation is already in story # Check if generation is already in story
existing = db.query(DBStoryItem).filter_by( existing = db.query(DBStoryItem).filter_by(story_id=story_id, generation_id=data.generation_id).first()
story_id=story_id,
generation_id=data.generation_id
).first()
if existing: if existing:
# Return existing item # Return existing item
profile = db.query(DBVoiceProfile).filter_by(id=generation.profile_id).first() profile = db.query(DBVoiceProfile).filter_by(id=generation.profile_id).first()
@@ -283,17 +270,15 @@ async def add_item_to_story(
if data.start_time_ms is not None: if data.start_time_ms is not None:
start_time_ms = data.start_time_ms start_time_ms = data.start_time_ms
else: else:
# Find the maximum end time on the target track only existing_items = (
existing_items = db.query( db.query(DBStoryItem, DBGeneration)
DBStoryItem, .join(DBGeneration, DBStoryItem.generation_id == DBGeneration.id)
DBGeneration .filter(
).join(
DBGeneration,
DBStoryItem.generation_id == DBGeneration.id
).filter(
DBStoryItem.story_id == story_id, DBStoryItem.story_id == story_id,
DBStoryItem.track == track, DBStoryItem.track == track,
).all() )
.all()
)
if not existing_items: if not existing_items:
start_time_ms = 0 start_time_ms = 0
@@ -349,10 +334,14 @@ async def move_story_item(
Updated item detail or None if not found Updated item detail or None if not found
""" """
# Get the item # Get the item
item = db.query(DBStoryItem).filter_by( item = (
db.query(DBStoryItem)
.filter_by(
id=item_id, id=item_id,
story_id=story_id, story_id=story_id,
).first() )
.first()
)
if not item: if not item:
return None return None
@@ -395,10 +384,14 @@ async def remove_item_from_story(
Returns: Returns:
True if removed, False if not found True if removed, False if not found
""" """
item = db.query(DBStoryItem).filter_by( item = (
db.query(DBStoryItem)
.filter_by(
id=item_id, id=item_id,
story_id=story_id, story_id=story_id,
).first() )
.first()
)
if not item: if not item:
return False return False
@@ -433,10 +426,14 @@ async def trim_story_item(
Updated item detail or None if not found Updated item detail or None if not found
""" """
# Get the item # Get the item
item = db.query(DBStoryItem).filter_by( item = (
db.query(DBStoryItem)
.filter_by(
id=item_id, id=item_id,
story_id=story_id, story_id=story_id,
).first() )
.first()
)
if not item: if not item:
return None return None
@@ -487,10 +484,14 @@ async def split_story_item(
List of two updated item details (original and new) or None if not found/invalid List of two updated item details (original and new) or None if not found/invalid
""" """
# Get the item # Get the item
item = db.query(DBStoryItem).filter_by( item = (
db.query(DBStoryItem)
.filter_by(
id=item_id, id=item_id,
story_id=story_id, story_id=story_id,
).first() )
.first()
)
if not item: if not item:
return None return None
@@ -500,8 +501,8 @@ async def split_story_item(
return None return None
# Calculate effective duration and validate split point # Calculate effective duration and validate split point
current_trim_start = getattr(item, 'trim_start_ms', 0) current_trim_start = getattr(item, "trim_start_ms", 0)
current_trim_end = getattr(item, 'trim_end_ms', 0) current_trim_end = getattr(item, "trim_end_ms", 0)
original_duration_ms = int(generation.duration * 1000) original_duration_ms = int(generation.duration * 1000)
effective_duration_ms = original_duration_ms - current_trim_start - current_trim_end effective_duration_ms = original_duration_ms - current_trim_start - current_trim_end
@@ -520,7 +521,7 @@ async def split_story_item(
id=str(uuid.uuid4()), id=str(uuid.uuid4()),
story_id=story_id, story_id=story_id,
generation_id=item.generation_id, # Same generation, different trim generation_id=item.generation_id, # Same generation, different trim
version_id=getattr(item, 'version_id', None), # Preserve pinned version version_id=getattr(item, "version_id", None), # Preserve pinned version
start_time_ms=item.start_time_ms + data.split_time_ms, start_time_ms=item.start_time_ms + data.split_time_ms,
track=item.track, track=item.track,
trim_start_ms=absolute_split_ms, trim_start_ms=absolute_split_ms,
@@ -566,10 +567,14 @@ async def duplicate_story_item(
New item detail or None if not found New item detail or None if not found
""" """
# Get the original item # Get the original item
original_item = db.query(DBStoryItem).filter_by( original_item = (
db.query(DBStoryItem)
.filter_by(
id=item_id, id=item_id,
story_id=story_id, story_id=story_id,
).first() )
.first()
)
if not original_item: if not original_item:
return None return None
@@ -579,8 +584,8 @@ async def duplicate_story_item(
return None return None
# Calculate effective duration # Calculate effective duration
current_trim_start = getattr(original_item, 'trim_start_ms', 0) current_trim_start = getattr(original_item, "trim_start_ms", 0)
current_trim_end = getattr(original_item, 'trim_end_ms', 0) current_trim_end = getattr(original_item, "trim_end_ms", 0)
original_duration_ms = int(generation.duration * 1000) original_duration_ms = int(generation.duration * 1000)
effective_duration_ms = original_duration_ms - current_trim_start - current_trim_end effective_duration_ms = original_duration_ms - current_trim_start - current_trim_end
@@ -589,7 +594,7 @@ async def duplicate_story_item(
id=str(uuid.uuid4()), id=str(uuid.uuid4()),
story_id=story_id, story_id=story_id,
generation_id=original_item.generation_id, # Same generation as original generation_id=original_item.generation_id, # Same generation as original
version_id=getattr(original_item, 'version_id', None), # Preserve pinned version version_id=getattr(original_item, "version_id", None), # Preserve pinned version
start_time_ms=original_item.start_time_ms + effective_duration_ms + 200, # 200ms gap start_time_ms=original_item.start_time_ms + effective_duration_ms + 200, # 200ms gap
track=original_item.track, track=original_item.track,
trim_start_ms=current_trim_start, trim_start_ms=current_trim_start,
@@ -673,19 +678,13 @@ async def reorder_story_items(
return None return None
# Get all items for this story with their generation data # Get all items for this story with their generation data
items_with_gen = db.query( items_with_gen = (
DBStoryItem, db.query(DBStoryItem, DBGeneration, DBVoiceProfile.name.label("profile_name"))
DBGeneration, .join(DBGeneration, DBStoryItem.generation_id == DBGeneration.id)
DBVoiceProfile.name.label('profile_name') .join(DBVoiceProfile, DBGeneration.profile_id == DBVoiceProfile.id)
).join( .filter(DBStoryItem.story_id == story_id)
DBGeneration, .all()
DBStoryItem.generation_id == DBGeneration.id )
).join(
DBVoiceProfile,
DBGeneration.profile_id == DBVoiceProfile.id
).filter(
DBStoryItem.story_id == story_id
).all()
# Create maps for quick lookup # Create maps for quick lookup
item_map = {item.generation_id: (item, gen, profile_name) for item, gen, profile_name in items_with_gen} item_map = {item.generation_id: (item, gen, profile_name) for item, gen, profile_name in items_with_gen}
@@ -738,10 +737,14 @@ async def set_story_item_version(
Returns: Returns:
Updated item detail or None if not found Updated item detail or None if not found
""" """
item = db.query(DBStoryItem).filter_by( item = (
db.query(DBStoryItem)
.filter_by(
id=item_id, id=item_id,
story_id=story_id, story_id=story_id,
).first() )
.first()
)
if not item: if not item:
return None return None
@@ -752,10 +755,15 @@ async def set_story_item_version(
# Validate version_id belongs to this generation if provided # Validate version_id belongs to this generation if provided
if data.version_id: if data.version_id:
from .database import GenerationVersion as DBGenerationVersion from .database import GenerationVersion as DBGenerationVersion
version = db.query(DBGenerationVersion).filter_by(
version = (
db.query(DBGenerationVersion)
.filter_by(
id=data.version_id, id=data.version_id,
generation_id=item.generation_id, generation_id=item.generation_id,
).first() )
.first()
)
if not version: if not version:
return None return None
@@ -793,15 +801,13 @@ async def export_story_audio(
return None return None
# Get all items ordered by start_time_ms # Get all items ordered by start_time_ms
items = db.query( items = (
DBStoryItem, db.query(DBStoryItem, DBGeneration)
DBGeneration .join(DBGeneration, DBStoryItem.generation_id == DBGeneration.id)
).join( .filter(DBStoryItem.story_id == story_id)
DBGeneration, .order_by(DBStoryItem.start_time_ms)
DBStoryItem.generation_id == DBGeneration.id .all()
).filter( )
DBStoryItem.story_id == story_id
).order_by(DBStoryItem.start_time_ms).all()
if not items: if not items:
return None return None
@@ -813,8 +819,9 @@ async def export_story_audio(
for item, generation in items: for item, generation in items:
# Resolve audio path: use pinned version if set, otherwise generation default # Resolve audio path: use pinned version if set, otherwise generation default
resolved_audio_path = generation.audio_path resolved_audio_path = generation.audio_path
if getattr(item, 'version_id', None): if getattr(item, "version_id", None):
from .database import GenerationVersion as DBGenerationVersion from .database import GenerationVersion as DBGenerationVersion
version = db.query(DBGenerationVersion).filter_by(id=item.version_id).first() version = db.query(DBGenerationVersion).filter_by(id=item.version_id).first()
if version: if version:
resolved_audio_path = version.audio_path resolved_audio_path = version.audio_path
@@ -828,8 +835,8 @@ async def export_story_audio(
sample_rate = sr # Use actual sample rate from first file sample_rate = sr # Use actual sample rate from first file
# Get trim values # Get trim values
trim_start_ms = getattr(item, 'trim_start_ms', 0) trim_start_ms = getattr(item, "trim_start_ms", 0)
trim_end_ms = getattr(item, 'trim_end_ms', 0) trim_end_ms = getattr(item, "trim_end_ms", 0)
# Calculate effective duration # Calculate effective duration
original_duration_ms = int(generation.duration * 1000) original_duration_ms = int(generation.duration * 1000)
@@ -841,18 +848,22 @@ async def export_story_audio(
# Extract the trimmed portion # Extract the trimmed portion
if trim_end_ms > 0: if trim_end_ms > 0:
trimmed_audio = audio[trim_start_sample:-trim_end_sample] if trim_end_sample > 0 else audio[trim_start_sample:] trimmed_audio = (
audio[trim_start_sample:-trim_end_sample] if trim_end_sample > 0 else audio[trim_start_sample:]
)
else: else:
trimmed_audio = audio[trim_start_sample:] trimmed_audio = audio[trim_start_sample:]
# Store audio with its timecode info # Store audio with its timecode info
start_time_ms = item.start_time_ms start_time_ms = item.start_time_ms
audio_data.append({ audio_data.append(
'audio': trimmed_audio, {
'start_time_ms': start_time_ms, "audio": trimmed_audio,
'duration_ms': effective_duration_ms, "start_time_ms": start_time_ms,
}) "duration_ms": effective_duration_ms,
}
)
except Exception: except Exception:
# Skip files that can't be loaded # Skip files that can't be loaded
continue continue
@@ -861,10 +872,7 @@ async def export_story_audio(
return None return None
# Calculate total duration: max(start_time_ms + duration_ms) # Calculate total duration: max(start_time_ms + duration_ms)
max_end_time_ms = max( max_end_time_ms = max((data["start_time_ms"] + data["duration_ms"] for data in audio_data), default=0)
(data['start_time_ms'] + data['duration_ms'] for data in audio_data),
default=0
)
# Convert to samples # Convert to samples
total_samples = int((max_end_time_ms / 1000.0) * sample_rate) total_samples = int((max_end_time_ms / 1000.0) * sample_rate)
@@ -874,8 +882,8 @@ async def export_story_audio(
# Mix each audio segment at its timecode position # Mix each audio segment at its timecode position
for data in audio_data: for data in audio_data:
audio = data['audio'] audio = data["audio"]
start_time_ms = data['start_time_ms'] start_time_ms = data["start_time_ms"]
# Calculate start sample index # Calculate start sample index
start_sample = int((start_time_ms / 1000.0) * sample_rate) start_sample = int((start_time_ms / 1000.0) * sample_rate)
@@ -886,7 +894,7 @@ async def export_story_audio(
if start_sample < total_samples: if start_sample < total_samples:
# Trim audio if it extends beyond buffer # Trim audio if it extends beyond buffer
audio_to_mix = audio[:end_sample - start_sample] audio_to_mix = audio[: end_sample - start_sample]
# Mix: add audio to existing buffer (overlapping audio will sum) # Mix: add audio to existing buffer (overlapping audio will sum)
# Normalize to prevent clipping (simple approach: divide by max) # Normalize to prevent clipping (simple approach: divide by max)
@@ -898,14 +906,14 @@ async def export_story_audio(
final_audio = final_audio / max_val final_audio = final_audio / max_val
# Save to temporary file # Save to temporary file
with tempfile.NamedTemporaryFile(suffix='.wav', delete=False) as tmp: with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
tmp_path = tmp.name tmp_path = tmp.name
try: try:
save_audio(final_audio, tmp_path, sample_rate) save_audio(final_audio, tmp_path, sample_rate)
# Read file bytes # Read file bytes
with open(tmp_path, 'rb') as f: with open(tmp_path, "rb") as f:
audio_bytes = f.read() audio_bytes = f.read()
return audio_bytes return audio_bytes
+9 -25
View File
@@ -37,11 +37,10 @@ async def monitor_sse_stream(model_name: str, timeout: int = 120):
if line.startswith("data: "): if line.startswith("data: "):
try: try:
data = json.loads(line[6:]) data = json.loads(line[6:])
print(f"[{timestamp}] → SSE Event: {data['status']:12} {data.get('progress', 0):6.1f}% {data.get('filename', '')}") print(
events.append({ f"[{timestamp}] → SSE Event: {data['status']:12} {data.get('progress', 0):6.1f}% {data.get('filename', '')}"
**data, )
"_timestamp": timestamp events.append({**data, "_timestamp": timestamp})
})
# Stop if complete or error # Stop if complete or error
if data.get("status") in ("complete", "error"): if data.get("status") in ("complete", "error"):
@@ -74,12 +73,15 @@ async def trigger_generation(profile_id: str, text: str, model_size: str = "1.7B
try: try:
async with httpx.AsyncClient(timeout=120) as client: async with httpx.AsyncClient(timeout=120) as client:
response = await client.post(url, json={ response = await client.post(
url,
json={
"profile_id": profile_id, "profile_id": profile_id,
"text": text, "text": text,
"language": "en", "language": "en",
"model_size": model_size, "model_size": model_size,
}) },
)
print(f"[{_timestamp()}] Response: {response.status_code}") print(f"[{_timestamp()}] Response: {response.status_code}")
@@ -292,24 +294,6 @@ async def main():
print(" Users see progress events even when the model is already cached,") print(" Users see progress events even when the model is already cached,")
print(" making them think the model is downloading again.") print(" making them think the model is downloading again.")
# Test Case 2: Fresh download (optional, commented out by default)
# Uncomment if you want to test download progress
# print("\n" + "🧪 " * 20)
# events_download = await test_generation_with_fresh_download()
#
# print("\n" + "=" * 80)
# print("TEST CASE 2 RESULTS: Generation with Model Download")
# print("=" * 80)
#
# if not events_download:
# print("ℹ Model was already cached, no download occurred")
# else:
# print(f"✓ Received {len(events_download)} download progress events")
# print("\nDownload Timeline:")
# for i, event in enumerate(events_download, 1):
# timestamp = event.pop("_timestamp", "??:??:??.???")
# print(f" {i}. [{timestamp}] {event}")
print("\n" + "=" * 80) print("\n" + "=" * 80)
print("Test Complete!") print("Test Complete!")
print("=" * 80) print("=" * 80)
+16 -19
View File
@@ -58,11 +58,6 @@ _ABBREVIATIONS = frozenset(
_PARA_TAG_RE = re.compile(r"\[[^\]]*\]") _PARA_TAG_RE = re.compile(r"\[[^\]]*\]")
# ---------------------------------------------------------------------------
# Text splitting
# ---------------------------------------------------------------------------
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*.
@@ -174,11 +169,6 @@ def _safe_hard_cut(segment: str, max_chars: int) -> int:
return cut return cut
# ---------------------------------------------------------------------------
# Audio concatenation
# ---------------------------------------------------------------------------
def concatenate_audio_chunks( def concatenate_audio_chunks(
chunks: List[np.ndarray], chunks: List[np.ndarray],
sample_rate: int, sample_rate: int,
@@ -211,11 +201,6 @@ def concatenate_audio_chunks(
return result return result
# ---------------------------------------------------------------------------
# Engine-agnostic chunked generation
# ---------------------------------------------------------------------------
async def generate_chunked( async def generate_chunked(
backend, backend,
text: str, text: str,
@@ -264,7 +249,11 @@ async def generate_chunked(
if len(chunks) <= 1: if len(chunks) <= 1:
# Short text — single-shot fast path # Short text — single-shot fast path
audio, sample_rate = await backend.generate( audio, sample_rate = await backend.generate(
text, voice_prompt, language, seed, instruct, text,
voice_prompt,
language,
seed,
instruct,
) )
if trim_fn is not None: if trim_fn is not None:
audio = trim_fn(audio, sample_rate) audio = trim_fn(audio, sample_rate)
@@ -273,7 +262,9 @@ async def generate_chunked(
# Long text — chunked generation # Long text — chunked generation
logger.info( logger.info(
"Splitting %d chars into %d chunks (max %d chars each)", "Splitting %d chars into %d chunks (max %d chars each)",
len(text), len(chunks), max_chunk_chars, len(text),
len(chunks),
max_chunk_chars,
) )
audio_chunks: List[np.ndarray] = [] audio_chunks: List[np.ndarray] = []
sample_rate: int | None = None sample_rate: int | None = None
@@ -281,7 +272,9 @@ async def generate_chunked(
for i, chunk_text in enumerate(chunks): for i, chunk_text in enumerate(chunks):
logger.info( logger.info(
"Generating chunk %d/%d (%d chars)", "Generating chunk %d/%d (%d chars)",
i + 1, len(chunks), len(chunk_text), i + 1,
len(chunks),
len(chunk_text),
) )
# Vary the seed per chunk to avoid correlated RNG artefacts, # Vary the seed per chunk to avoid correlated RNG artefacts,
# but keep it deterministic so the same (text, seed) pair # but keep it deterministic so the same (text, seed) pair
@@ -289,7 +282,11 @@ async def generate_chunked(
chunk_seed = (seed + i) if seed is not None else None chunk_seed = (seed + i) if seed is not None else None
chunk_audio, chunk_sr = await backend.generate( chunk_audio, chunk_sr = await backend.generate(
chunk_text, voice_prompt, language, chunk_seed, instruct, chunk_text,
voice_prompt,
language,
chunk_seed,
instruct,
) )
if trim_fn is not None: if trim_fn is not None:
chunk_audio = trim_fn(chunk_audio, chunk_sr) chunk_audio = trim_fn(chunk_audio, chunk_sr)
+40 -23
View File
@@ -35,10 +35,6 @@ from pedalboard import (
) )
# ---------------------------------------------------------------------------
# Effect registry: maps type names -> (pedalboard class, param definitions)
# ---------------------------------------------------------------------------
# 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": {
@@ -49,7 +45,13 @@ EFFECT_REGISTRY: Dict[str, Dict[str, Any]] = {
"rate_hz": {"default": 1.0, "min": 0.01, "max": 20.0, "step": 0.01, "description": "LFO speed (Hz)"}, "rate_hz": {"default": 1.0, "min": 0.01, "max": 20.0, "step": 0.01, "description": "LFO speed (Hz)"},
"depth": {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "description": "Modulation depth"}, "depth": {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "description": "Modulation depth"},
"feedback": {"default": 0.0, "min": 0.0, "max": 0.95, "step": 0.01, "description": "Feedback amount"}, "feedback": {"default": 0.0, "min": 0.0, "max": 0.95, "step": 0.01, "description": "Feedback amount"},
"centre_delay_ms": {"default": 7.0, "min": 0.5, "max": 50.0, "step": 0.1, "description": "Centre delay (ms)"}, "centre_delay_ms": {
"default": 7.0,
"min": 0.5,
"max": 50.0,
"step": 0.1,
"description": "Centre delay (ms)",
},
"mix": {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "description": "Wet/dry mix"}, "mix": {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "description": "Wet/dry mix"},
}, },
}, },
@@ -70,7 +72,13 @@ EFFECT_REGISTRY: Dict[str, Dict[str, Any]] = {
"label": "Delay", "label": "Delay",
"description": "Echo / delay line.", "description": "Echo / delay line.",
"params": { "params": {
"delay_seconds": {"default": 0.3, "min": 0.01, "max": 2.0, "step": 0.01, "description": "Delay time (seconds)"}, "delay_seconds": {
"default": 0.3,
"min": 0.01,
"max": 2.0,
"step": 0.01,
"description": "Delay time (seconds)",
},
"feedback": {"default": 0.3, "min": 0.0, "max": 0.95, "step": 0.01, "description": "Feedback amount"}, "feedback": {"default": 0.3, "min": 0.0, "max": 0.95, "step": 0.01, "description": "Feedback amount"},
"mix": {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01, "description": "Wet/dry mix"}, "mix": {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01, "description": "Wet/dry mix"},
}, },
@@ -83,7 +91,13 @@ EFFECT_REGISTRY: Dict[str, Dict[str, Any]] = {
"threshold_db": {"default": -20.0, "min": -60.0, "max": 0.0, "step": 0.5, "description": "Threshold (dB)"}, "threshold_db": {"default": -20.0, "min": -60.0, "max": 0.0, "step": 0.5, "description": "Threshold (dB)"},
"ratio": {"default": 4.0, "min": 1.0, "max": 20.0, "step": 0.1, "description": "Compression ratio"}, "ratio": {"default": 4.0, "min": 1.0, "max": 20.0, "step": 0.1, "description": "Compression ratio"},
"attack_ms": {"default": 10.0, "min": 0.1, "max": 100.0, "step": 0.1, "description": "Attack time (ms)"}, "attack_ms": {"default": 10.0, "min": 0.1, "max": 100.0, "step": 0.1, "description": "Attack time (ms)"},
"release_ms": {"default": 100.0, "min": 10.0, "max": 1000.0,"step": 1.0, "description": "Release time (ms)"}, "release_ms": {
"default": 100.0,
"min": 10.0,
"max": 1000.0,
"step": 1.0,
"description": "Release time (ms)",
},
}, },
}, },
"gain": { "gain": {
@@ -99,7 +113,13 @@ EFFECT_REGISTRY: Dict[str, Dict[str, Any]] = {
"label": "High-Pass Filter", "label": "High-Pass Filter",
"description": "Removes frequencies below the cutoff.", "description": "Removes frequencies below the cutoff.",
"params": { "params": {
"cutoff_frequency_hz": {"default": 80.0, "min": 20.0, "max": 8000.0, "step": 1.0, "description": "Cutoff frequency (Hz)"}, "cutoff_frequency_hz": {
"default": 80.0,
"min": 20.0,
"max": 8000.0,
"step": 1.0,
"description": "Cutoff frequency (Hz)",
},
}, },
}, },
"lowpass": { "lowpass": {
@@ -107,7 +127,13 @@ EFFECT_REGISTRY: Dict[str, Dict[str, Any]] = {
"label": "Low-Pass Filter", "label": "Low-Pass Filter",
"description": "Removes frequencies above the cutoff.", "description": "Removes frequencies above the cutoff.",
"params": { "params": {
"cutoff_frequency_hz": {"default": 8000.0, "min": 200.0, "max": 20000.0, "step": 1.0, "description": "Cutoff frequency (Hz)"}, "cutoff_frequency_hz": {
"default": 8000.0,
"min": 200.0,
"max": 20000.0,
"step": 1.0,
"description": "Cutoff frequency (Hz)",
},
}, },
}, },
"pitch_shift": { "pitch_shift": {
@@ -121,10 +147,6 @@ EFFECT_REGISTRY: Dict[str, Dict[str, Any]] = {
} }
# ---------------------------------------------------------------------------
# Built-in presets
# ---------------------------------------------------------------------------
BUILTIN_PRESETS: Dict[str, Dict[str, Any]] = { BUILTIN_PRESETS: Dict[str, Dict[str, Any]] = {
"robotic": { "robotic": {
"name": "Robotic", "name": "Robotic",
@@ -233,10 +255,6 @@ BUILTIN_PRESETS: Dict[str, Dict[str, Any]] = {
} }
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
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.
@@ -244,15 +262,14 @@ def get_available_effects() -> List[Dict[str, Any]]:
""" """
result = [] result = []
for effect_type, info in EFFECT_REGISTRY.items(): for effect_type, info in EFFECT_REGISTRY.items():
result.append({ result.append(
{
"type": effect_type, "type": effect_type,
"label": info["label"], "label": info["label"],
"description": info["description"], "description": info["description"],
"params": { "params": {name: {k: v for k, v in pdef.items()} for name, pdef in info["params"].items()},
name: {k: v for k, v in pdef.items()} }
for name, pdef in info["params"].items() )
},
})
return result return result