mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 09:35:16 -07:00
comment cleanup
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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()
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user