From 72c13fd3fc533de246e12ccfc2a84ee413758b03 Mon Sep 17 00:00:00 2001 From: James Pine Date: Thu, 19 Mar 2026 19:13:28 -0700 Subject: [PATCH] fix: enforce preset profile engine compatibility --- .../Generation/EngineModelSelector.tsx | 6 +- .../components/Generation/GenerationForm.tsx | 27 +++- .../components/VoiceProfiles/ProfileForm.tsx | 9 ++ app/src/lib/api/client.ts | 6 - backend/database/migrations.py | 11 +- backend/requirements.txt | 2 +- backend/routes/generations.py | 25 ++-- backend/routes/profiles.py | 121 ------------------ backend/services/profiles.py | 43 +++++++ backend/voicebox-server.spec | 2 +- 10 files changed, 102 insertions(+), 150 deletions(-) diff --git a/app/src/components/Generation/EngineModelSelector.tsx b/app/src/components/Generation/EngineModelSelector.tsx index 7195a287..7f4f600b 100644 --- a/app/src/components/Generation/EngineModelSelector.tsx +++ b/app/src/components/Generation/EngineModelSelector.tsx @@ -57,7 +57,7 @@ function getSelectValue(engine: string, modelSize?: string): string { return engine; } -function handleEngineChange(form: UseFormReturn, value: string) { +export function applyEngineSelection(form: UseFormReturn, value: string) { if (value.startsWith('qwen_custom_voice:')) { const [, modelSize] = value.split(':'); form.setValue('engine', 'qwen_custom_voice'); @@ -123,7 +123,7 @@ export function EngineModelSelector({ form, compact, selectedProfile }: EngineMo useEffect(() => { if (!currentEngineAvailable && availableOptions.length > 0) { - handleEngineChange(form, availableOptions[0].value); + applyEngineSelection(form, availableOptions[0].value); } }, [availableOptions, currentEngineAvailable, form]); @@ -133,7 +133,7 @@ export function EngineModelSelector({ form, compact, selectedProfile }: EngineMo : undefined; return ( - applyEngineSelection(form, v)}> diff --git a/app/src/components/Generation/GenerationForm.tsx b/app/src/components/Generation/GenerationForm.tsx index 1195e8b1..ef3ff2c0 100644 --- a/app/src/components/Generation/GenerationForm.tsx +++ b/app/src/components/Generation/GenerationForm.tsx @@ -1,3 +1,4 @@ +import { useEffect } from 'react'; import { Loader2, Mic } from 'lucide-react'; import { Button } from '@/components/ui/button'; import { Card, CardContent, CardHeader, CardTitle } from '@/components/ui/card'; @@ -19,19 +20,41 @@ import { SelectValue, } from '@/components/ui/select'; import { Textarea } from '@/components/ui/textarea'; -import { getLanguageOptionsForEngine } from '@/lib/constants/languages'; +import { getLanguageOptionsForEngine, type LanguageCode } from '@/lib/constants/languages'; import { useGenerationForm } from '@/lib/hooks/useGenerationForm'; import { useProfile } from '@/lib/hooks/useProfiles'; import { useUIStore } from '@/stores/uiStore'; -import { EngineModelSelector, getEngineDescription } from './EngineModelSelector'; +import { EngineModelSelector, applyEngineSelection, getEngineDescription } from './EngineModelSelector'; import { ParalinguisticInput } from './ParalinguisticInput'; +function getEngineSelectValue(engine: string): string { + if (engine === 'qwen') return 'qwen:1.7B'; + if (engine === 'qwen_custom_voice') return 'qwen_custom_voice:1.7B'; + if (engine === 'tada') return 'tada:1B'; + return engine; +} + export function GenerationForm() { const selectedProfileId = useUIStore((state) => state.selectedProfileId); const { data: selectedProfile } = useProfile(selectedProfileId || ''); const { form, handleSubmit, isPending } = useGenerationForm(); + useEffect(() => { + if (!selectedProfile) { + return; + } + + if (selectedProfile.language) { + form.setValue('language', selectedProfile.language as LanguageCode); + } + + const preferredEngine = selectedProfile.default_engine || selectedProfile.preset_engine; + if (preferredEngine) { + applyEngineSelection(form, getEngineSelectValue(preferredEngine)); + } + }, [form, selectedProfile]); + async function onSubmit(data: Parameters[0]) { await handleSubmit(data, selectedProfileId); } diff --git a/app/src/components/VoiceProfiles/ProfileForm.tsx b/app/src/components/VoiceProfiles/ProfileForm.tsx index a148bfb5..50b8cb57 100644 --- a/app/src/components/VoiceProfiles/ProfileForm.tsx +++ b/app/src/components/VoiceProfiles/ProfileForm.tsx @@ -375,6 +375,15 @@ export function ProfileForm() { } }, [availableDefaultEngines, defaultEngine]); + useEffect(() => { + if (!selectedPresetVoiceId) { + return; + } + + if (!presetVoices.some((voice: PresetVoice) => voice.voice_id === selectedPresetVoiceId)) { + setSelectedPresetVoiceId(''); + } + }, [presetVoices, selectedPresetVoiceId]); async function handleTranscribe() { const file = form.getValues('sampleFile'); if (!file) { diff --git a/app/src/lib/api/client.ts b/app/src/lib/api/client.ts index 6849b8c7..98a375e3 100644 --- a/app/src/lib/api/client.ts +++ b/app/src/lib/api/client.ts @@ -102,12 +102,6 @@ class ApiClient { return this.request<{ engine: string; voices: PresetVoice[] }>(`/profiles/presets/${engine}`); } - async seedPresetProfiles( - engine: string, - ): Promise<{ engine: string; created: number; total_available: number }> { - return this.request(`/profiles/presets/${engine}/seed`, { method: 'POST' }); - } - async updateProfile(profileId: string, data: VoiceProfileCreate): Promise { return this.request(`/profiles/${profileId}`, { method: 'PUT', diff --git a/backend/database/migrations.py b/backend/database/migrations.py index f4cc5ada..6256d92b 100644 --- a/backend/database/migrations.py +++ b/backend/database/migrations.py @@ -194,6 +194,8 @@ def _resolve_relative_paths(engine, tables: set[str]) -> None: configured data directory. If the path starts with "data/", strip that prefix and prepend get_data_dir(). Otherwise, join the relative path directly under get_data_dir(). + directly under get_data_dir(). If the rebased path still does not exist, + fall back to resolving relative to CWD. """ from pathlib import Path @@ -222,16 +224,13 @@ def _resolve_relative_paths(engine, tables: set[str]) -> None: p = Path(path_val) if p.is_absolute(): continue - - # Try rebasing: "data/generations/abc.wav" → data_dir / "generations/abc.wav" parts = p.parts if parts and parts[0] == "data": - rebased = data_dir / Path(*parts[1:]) + rebased = (data_dir / Path(*parts[1:])).resolve() else: - rebased = data_dir / p - - resolved = rebased.resolve() + rebased = (data_dir / p).resolve() + resolved = rebased if rebased.exists() else p.resolve() if resolved.exists(): conn.execute( text(f"UPDATE {table} SET {column} = :path WHERE id = :id"), diff --git a/backend/requirements.txt b/backend/requirements.txt index c9c65b0f..e916b1d2 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -42,7 +42,7 @@ torchaudio # Kokoro TTS (lightweight 82M-param engine) kokoro>=0.9.4 -misaki[en]>=0.9.4 +misaki[en,ja,zh]>=0.9.4 # spacy model for misaki English G2P — must be pre-installed or misaki # tries spacy.cli.download() at runtime which crashes frozen builds en_core_web_sm @ https://github.com/explosion/spacy-models/releases/download/en_core_web_sm-3.8.0/en_core_web_sm-3.8.0-py3-none-any.whl diff --git a/backend/routes/generations.py b/backend/routes/generations.py index 9f985fcc..2af3832e 100644 --- a/backend/routes/generations.py +++ b/backend/routes/generations.py @@ -20,6 +20,10 @@ from ..utils.tasks import get_task_manager router = APIRouter() +def _resolve_generation_engine(data: models.GenerationRequest, profile) -> str: + return data.engine or getattr(profile, "default_engine", None) or getattr(profile, "preset_engine", None) or "qwen" + + @router.post("/generate", response_model=models.GenerationResponse) async def generate_speech( data: models.GenerationRequest, @@ -35,7 +39,12 @@ async def generate_speech( from ..backends import engine_has_model_sizes - engine = data.engine or "qwen" + engine = _resolve_generation_engine(data, profile) + try: + profiles.validate_profile_engine(profile, engine) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) + model_size = (data.model_size or "1.7B") if engine_has_model_sizes(engine) else None generation = await history.create_generation( @@ -230,15 +239,11 @@ async def stream_speech( if not profile: raise HTTPException(status_code=404, detail="Profile not found") - # Mirror the regular /generate endpoint behavior more closely: - # if the caller doesn't specify an engine, prefer the profile's default - # engine (or preset engine) before falling back to qwen. - engine = ( - data.engine - or getattr(profile, "default_engine", None) - or getattr(profile, "preset_engine", None) - or "qwen" - ) + engine = _resolve_generation_engine(data, profile) + try: + profiles.validate_profile_engine(profile, engine) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) tts_model = get_tts_backend_for_engine(engine) model_size = data.model_size or "1.7B" diff --git a/backend/routes/profiles.py b/backend/routes/profiles.py index 665f6b24..d65a1138 100644 --- a/backend/routes/profiles.py +++ b/backend/routes/profiles.py @@ -4,8 +4,6 @@ import io import json as _json import logging import tempfile -import uuid -from datetime import datetime from pathlib import Path from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile @@ -107,125 +105,6 @@ async def list_preset_voices(engine: str): } return {"engine": engine, "voices": []} - -@router.post("/profiles/presets/{engine}/seed") -async def seed_preset_profiles_route( - engine: str, - db: Session = Depends(get_db), -): - """Seed preset voice profiles for an engine. - - Creates profiles for all available preset voices that don't already exist. - Returns the count of newly created profiles. - """ - if engine == "kokoro": - return _seed_kokoro_presets(db) - if engine == "qwen_custom_voice": - return _seed_qwen_custom_voice_presets(db) - raise HTTPException(status_code=400, detail=f"No presets available for engine: {engine}") - - -def _seed_kokoro_presets(db: Session): - """Seed Kokoro preset profiles.""" - try: - from ..backends.kokoro_backend import KOKORO_VOICES - - created = 0 - for voice_id, display_name, gender, lang in KOKORO_VOICES: - profile_name = display_name - - # Disambiguate duplicate display names across languages - # (e.g. "Alpha" exists in Hindi and Japanese, "Dora" in Spanish and Portuguese) - dupes = [v for v in KOKORO_VOICES if v[1] == display_name] - if len(dupes) > 1: - lang_labels = {"en": "English", "es": "Spanish", "fr": "French", "hi": "Hindi", - "it": "Italian", "pt": "Portuguese", "ja": "Japanese", "zh": "Chinese"} - profile_name = f"{display_name} {lang_labels.get(lang, lang)}" - - # Skip if preset already exists - existing = ( - db.query(DBVoiceProfile) - .filter_by(preset_engine="kokoro", preset_voice_id=voice_id) - .first() - ) - if existing: - continue - - unique_name = profile_name - suffix = 2 - while db.query(DBVoiceProfile).filter_by(name=unique_name).first(): - unique_name = f"{profile_name} {suffix}" - suffix += 1 - - profile = DBVoiceProfile( - id=str(uuid.uuid4()), - name=unique_name, - description=f"Kokoro preset voice — {display_name} ({gender})", - language=lang, - voice_type="preset", - preset_engine="kokoro", - preset_voice_id=voice_id, - created_at=datetime.utcnow(), - updated_at=datetime.utcnow(), - ) - db.add(profile) - created += 1 - - if created > 0: - db.commit() - logger.info(f"Seeded {created} Kokoro preset profiles") - - return {"engine": "kokoro", "created": created, "total_available": len(KOKORO_VOICES)} - except Exception as e: - logger.exception(f"Failed to seed Kokoro profiles: {e}") - raise HTTPException(status_code=500, detail=str(e)) - - -def _seed_qwen_custom_voice_presets(db: Session): - """Seed Qwen CustomVoice preset profiles.""" - try: - from ..backends.qwen_custom_voice_backend import QWEN_CUSTOM_VOICES - - created = 0 - for speaker_id, display_name, gender, lang, description in QWEN_CUSTOM_VOICES: - # Skip if preset already exists - existing = ( - db.query(DBVoiceProfile) - .filter_by(preset_engine="qwen_custom_voice", preset_voice_id=speaker_id) - .first() - ) - if existing: - continue - - # Skip name collisions - if db.query(DBVoiceProfile).filter_by(name=display_name).first(): - continue - - profile = DBVoiceProfile( - id=str(uuid.uuid4()), - name=display_name, - description=f"Qwen CustomVoice — {description}", - language=lang, - voice_type="preset", - preset_engine="qwen_custom_voice", - preset_voice_id=speaker_id, - default_engine="qwen_custom_voice", - created_at=datetime.utcnow(), - updated_at=datetime.utcnow(), - ) - db.add(profile) - created += 1 - - if created > 0: - db.commit() - logger.info(f"Seeded {created} Qwen CustomVoice preset profiles") - - return {"engine": "qwen_custom_voice", "created": created, "total_available": len(QWEN_CUSTOM_VOICES)} - except Exception as e: - logger.exception(f"Failed to seed Qwen CustomVoice profiles: {e}") - raise HTTPException(status_code=500, detail=str(e)) - - @router.get("/profiles/{profile_id}", response_model=models.VoiceProfileResponse) async def get_profile( profile_id: str, diff --git a/backend/services/profiles.py b/backend/services/profiles.py index 16ea7b57..cd418aab 100644 --- a/backend/services/profiles.py +++ b/backend/services/profiles.py @@ -61,6 +61,20 @@ def _profile_to_response( ) +def _get_preset_voice_ids(engine: str) -> set[str]: + if engine == "kokoro": + from ..backends.kokoro_backend import KOKORO_VOICES + + return {voice_id for voice_id, _name, _gender, _lang in KOKORO_VOICES} + + if engine == "qwen_custom_voice": + from ..backends.qwen_custom_voice_backend import QWEN_CUSTOM_VOICES + + return {voice_id for voice_id, _name, _gender, _lang, _desc in QWEN_CUSTOM_VOICES} + + return set() + + def _validate_profile_fields( *, voice_type: str, @@ -74,6 +88,10 @@ def _validate_profile_fields( return "Preset profiles require both preset_engine and preset_voice_id" if default_engine and default_engine != preset_engine: return "Preset profiles must use their preset_engine as default_engine" + + available_voice_ids = _get_preset_voice_ids(preset_engine) + if available_voice_ids and preset_voice_id not in available_voice_ids: + return f"Preset voice '{preset_voice_id}' is not valid for engine '{preset_engine}'" return None if voice_type == "designed": @@ -92,6 +110,30 @@ def _validate_profile_fields( return None +def validate_profile_engine(profile, engine: str) -> None: + voice_type = getattr(profile, "voice_type", None) or "cloned" + + if voice_type == "preset": + preset_engine = getattr(profile, "preset_engine", None) + preset_voice_id = getattr(profile, "preset_voice_id", None) + if not preset_engine or not preset_voice_id: + raise ValueError(f"Preset profile {profile.id} is missing preset engine metadata") + if preset_engine != engine: + raise ValueError( + f"Preset profile {profile.id} only supports engine '{preset_engine}', not '{engine}'" + ) + return + + if voice_type == "designed": + design_prompt = getattr(profile, "design_prompt", None) + if not design_prompt or not design_prompt.strip(): + raise ValueError(f"Designed profile {profile.id} is missing design_prompt") + return + + if engine not in CLONING_ENGINES: + raise ValueError(f"Engine '{engine}' does not support cloned voice profiles") + + async def create_profile( data: VoiceProfileCreate, db: Session, @@ -476,6 +518,7 @@ async def create_voice_prompt_for_profile( raise ValueError(f"Profile not found: {profile_id}") voice_type = getattr(profile, "voice_type", None) or "cloned" + validate_profile_engine(profile, engine) # ── Preset profiles: return engine-specific voice reference ── if voice_type == "preset": diff --git a/backend/voicebox-server.spec b/backend/voicebox-server.spec index 1c208ba3..c1acc541 100644 --- a/backend/voicebox-server.spec +++ b/backend/voicebox-server.spec @@ -5,7 +5,7 @@ from PyInstaller.utils.hooks import copy_metadata datas = [] binaries = [] -hiddenimports = ['backend', 'backend.main', 'backend.config', 'backend.database', 'backend.models', 'backend.services.profiles', 'backend.services.history', 'backend.services.tts', 'backend.services.transcribe', 'backend.utils.platform_detect', 'backend.backends', 'backend.backends.pytorch_backend', 'backend.utils.audio', 'backend.utils.cache', 'backend.utils.progress', 'backend.utils.hf_progress', 'backend.services.cuda', 'backend.services.effects', 'backend.utils.effects', 'backend.services.versions', 'pedalboard', 'chatterbox', 'chatterbox.tts_turbo', 'chatterbox.mtl_tts', 'backend.backends.chatterbox_backend', 'backend.backends.chatterbox_turbo_backend', 'backend.backends.luxtts_backend', 'zipvoice', 'zipvoice.luxvoice', 'torch', 'transformers', 'fastapi', 'uvicorn', 'sqlalchemy', 'soundfile', 'qwen_tts', 'qwen_tts.inference', 'qwen_tts.inference.qwen3_tts_model', 'qwen_tts.inference.qwen3_tts_tokenizer', 'qwen_tts.core', 'qwen_tts.cli', 'requests', 'pkg_resources.extern', 'backend.backends.hume_backend', 'tada', 'tada.modules', 'tada.modules.tada', 'tada.modules.encoder', 'tada.modules.decoder', 'tada.modules.aligner', 'tada.modules.acoustic_spkr_verf', 'tada.nn', 'tada.nn.vibevoice', 'tada.utils', 'tada.utils.gray_code', 'tada.utils.text', 'backend.utils.dac_shim', 'torchaudio', 'backend.backends.kokoro_backend', 'kokoro', 'kokoro.pipeline', 'kokoro.model', 'kokoro.istftnet', 'kokoro.modules', 'kokoro.custom_stft', 'en_core_web_sm', 'loguru', 'backend.backends.mlx_backend', 'mlx', 'mlx.core', 'mlx.nn', 'mlx_audio', 'mlx_audio.tts', 'mlx_audio.stt'] +hiddenimports = ['backend', 'backend.main', 'backend.config', 'backend.database', 'backend.models', 'backend.services.profiles', 'backend.services.history', 'backend.services.tts', 'backend.services.transcribe', 'backend.utils.platform_detect', 'backend.backends', 'backend.backends.pytorch_backend', 'backend.backends.qwen_custom_voice_backend', 'backend.utils.audio', 'backend.utils.cache', 'backend.utils.progress', 'backend.utils.hf_progress', 'backend.services.cuda', 'backend.services.effects', 'backend.utils.effects', 'backend.services.versions', 'pedalboard', 'chatterbox', 'chatterbox.tts_turbo', 'chatterbox.mtl_tts', 'backend.backends.chatterbox_backend', 'backend.backends.chatterbox_turbo_backend', 'backend.backends.luxtts_backend', 'zipvoice', 'zipvoice.luxvoice', 'torch', 'transformers', 'fastapi', 'uvicorn', 'sqlalchemy', 'soundfile', 'qwen_tts', 'qwen_tts.inference', 'qwen_tts.inference.qwen3_tts_model', 'qwen_tts.inference.qwen3_tts_tokenizer', 'qwen_tts.core', 'qwen_tts.cli', 'requests', 'pkg_resources.extern', 'backend.backends.hume_backend', 'tada', 'tada.modules', 'tada.modules.tada', 'tada.modules.encoder', 'tada.modules.decoder', 'tada.modules.aligner', 'tada.modules.acoustic_spkr_verf', 'tada.nn', 'tada.nn.vibevoice', 'tada.utils', 'tada.utils.gray_code', 'tada.utils.text', 'backend.utils.dac_shim', 'torchaudio', 'backend.backends.kokoro_backend', 'kokoro', 'kokoro.pipeline', 'kokoro.model', 'kokoro.istftnet', 'kokoro.modules', 'kokoro.custom_stft', 'en_core_web_sm', 'loguru', 'backend.backends.mlx_backend', 'mlx', 'mlx.core', 'mlx.nn', 'mlx_audio', 'mlx_audio.tts', 'mlx_audio.stt'] datas += copy_metadata('qwen-tts') datas += copy_metadata('requests') datas += copy_metadata('transformers')