mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-02 16:45:15 -07:00
move CRUD and service modules into services/, platform_detect into utils/
Move 9 business-logic modules from the backend root into services/: channels, effects, history, profiles, stories, versions, export_import, transcribe, tts. Move platform_detect.py into utils/. Backend root now contains only infrastructure (app, main, config, server, models, build_binary) and docs. All 94 routes verified.
This commit is contained in:
+3
-2
@@ -15,9 +15,10 @@ from fastapi import FastAPI
|
|||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
from urllib.parse import quote
|
from urllib.parse import quote
|
||||||
|
|
||||||
from . import __version__, database, tts, transcribe, config
|
from . import __version__, database
|
||||||
|
from .services import tts, transcribe
|
||||||
from .database import get_db
|
from .database import get_db
|
||||||
from .platform_detect import get_backend_type
|
from .utils.platform_detect import get_backend_type
|
||||||
from .utils.progress import get_progress_manager
|
from .utils.progress import get_progress_manager
|
||||||
from .services.task_queue import create_background_task, init_queue
|
from .services.task_queue import create_background_task, init_queue
|
||||||
from .routes import register_routers
|
from .routes import register_routers
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ from typing import Protocol, Optional, Tuple, List
|
|||||||
from typing_extensions import runtime_checkable
|
from typing_extensions import runtime_checkable
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
from ..platform_detect import get_backend_type
|
from ..utils.platform_detect import get_backend_type
|
||||||
|
|
||||||
LANGUAGE_CODE_TO_NAME = {
|
LANGUAGE_CODE_TO_NAME = {
|
||||||
"zh": "chinese",
|
"zh": "chinese",
|
||||||
@@ -375,7 +375,7 @@ async def ensure_model_cached_or_raise(engine: str, model_size: str = "default")
|
|||||||
def unload_model_by_config(config: ModelConfig) -> bool:
|
def unload_model_by_config(config: ModelConfig) -> bool:
|
||||||
"""Unload a model given its config. Returns True if it was loaded, False otherwise."""
|
"""Unload a model given its config. Returns True if it was loaded, False otherwise."""
|
||||||
from . import get_tts_backend_for_engine
|
from . import get_tts_backend_for_engine
|
||||||
from .. import tts, transcribe
|
from ..services import tts, transcribe
|
||||||
|
|
||||||
if config.engine == "whisper":
|
if config.engine == "whisper":
|
||||||
whisper_model = transcribe.get_whisper_model()
|
whisper_model = transcribe.get_whisper_model()
|
||||||
@@ -403,7 +403,7 @@ def unload_model_by_config(config: ModelConfig) -> bool:
|
|||||||
def check_model_loaded(config: ModelConfig) -> bool:
|
def check_model_loaded(config: ModelConfig) -> bool:
|
||||||
"""Check if a model is currently loaded."""
|
"""Check if a model is currently loaded."""
|
||||||
from . import get_tts_backend_for_engine
|
from . import get_tts_backend_for_engine
|
||||||
from .. import tts, transcribe
|
from ..services import tts, transcribe
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if config.engine == "whisper":
|
if config.engine == "whisper":
|
||||||
@@ -424,7 +424,7 @@ def check_model_loaded(config: ModelConfig) -> bool:
|
|||||||
def get_model_load_func(config: ModelConfig):
|
def get_model_load_func(config: ModelConfig):
|
||||||
"""Return a callable that loads/downloads the model."""
|
"""Return a callable that loads/downloads the model."""
|
||||||
from . import get_tts_backend_for_engine
|
from . import get_tts_backend_for_engine
|
||||||
from .. import tts, transcribe
|
from ..services import tts, transcribe
|
||||||
|
|
||||||
if config.engine == "whisper":
|
if config.engine == "whisper":
|
||||||
return lambda: transcribe.get_whisper_model().load_model(config.model_size)
|
return lambda: transcribe.get_whisper_model().load_model(config.model_size)
|
||||||
|
|||||||
@@ -55,11 +55,11 @@ def build_server(cuda=False):
|
|||||||
'--hidden-import', 'backend.config',
|
'--hidden-import', 'backend.config',
|
||||||
'--hidden-import', 'backend.database',
|
'--hidden-import', 'backend.database',
|
||||||
'--hidden-import', 'backend.models',
|
'--hidden-import', 'backend.models',
|
||||||
'--hidden-import', 'backend.profiles',
|
'--hidden-import', 'backend.services.profiles',
|
||||||
'--hidden-import', 'backend.history',
|
'--hidden-import', 'backend.services.history',
|
||||||
'--hidden-import', 'backend.tts',
|
'--hidden-import', 'backend.services.tts',
|
||||||
'--hidden-import', 'backend.transcribe',
|
'--hidden-import', 'backend.services.transcribe',
|
||||||
'--hidden-import', 'backend.platform_detect',
|
'--hidden-import', 'backend.utils.platform_detect',
|
||||||
'--hidden-import', 'backend.backends',
|
'--hidden-import', 'backend.backends',
|
||||||
'--hidden-import', 'backend.backends.pytorch_backend',
|
'--hidden-import', 'backend.backends.pytorch_backend',
|
||||||
'--hidden-import', 'backend.utils.audio',
|
'--hidden-import', 'backend.utils.audio',
|
||||||
@@ -68,9 +68,9 @@ def build_server(cuda=False):
|
|||||||
'--hidden-import', 'backend.utils.hf_progress',
|
'--hidden-import', 'backend.utils.hf_progress',
|
||||||
'--hidden-import', 'backend.utils.validation',
|
'--hidden-import', 'backend.utils.validation',
|
||||||
'--hidden-import', 'backend.services.cuda',
|
'--hidden-import', 'backend.services.cuda',
|
||||||
'--hidden-import', 'backend.effects',
|
'--hidden-import', 'backend.services.effects',
|
||||||
'--hidden-import', 'backend.utils.effects',
|
'--hidden-import', 'backend.utils.effects',
|
||||||
'--hidden-import', 'backend.versions',
|
'--hidden-import', 'backend.services.versions',
|
||||||
'--hidden-import', 'pedalboard',
|
'--hidden-import', 'pedalboard',
|
||||||
'--hidden-import', 'chatterbox',
|
'--hidden-import', 'chatterbox',
|
||||||
'--hidden-import', 'chatterbox.tts_turbo',
|
'--hidden-import', 'chatterbox.tts_turbo',
|
||||||
|
|||||||
@@ -6,7 +6,8 @@ from fastapi import APIRouter, Depends, HTTPException
|
|||||||
from fastapi.responses import FileResponse
|
from fastapi.responses import FileResponse
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from .. import history, models
|
from .. import models
|
||||||
|
from ..services import history
|
||||||
from ..database import get_db
|
from ..database import get_db
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
@@ -15,7 +16,7 @@ router = APIRouter()
|
|||||||
@router.get("/audio/version/{version_id}")
|
@router.get("/audio/version/{version_id}")
|
||||||
async def get_version_audio(version_id: str, db: Session = Depends(get_db)):
|
async def get_version_audio(version_id: str, db: Session = Depends(get_db)):
|
||||||
"""Serve audio for a specific version."""
|
"""Serve audio for a specific version."""
|
||||||
from .. import versions as versions_mod
|
from ..services import versions as versions_mod
|
||||||
|
|
||||||
version = versions_mod.get_version(version_id, db)
|
version = versions_mod.get_version(version_id, db)
|
||||||
if not version:
|
if not version:
|
||||||
|
|||||||
@@ -3,7 +3,8 @@
|
|||||||
from fastapi import APIRouter, Depends, HTTPException
|
from fastapi import APIRouter, Depends, HTTPException
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from .. import channels, models
|
from .. import models
|
||||||
|
from ..services import channels
|
||||||
from ..database import get_db
|
from ..database import get_db
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|||||||
+12
-11
@@ -9,7 +9,8 @@ from fastapi import APIRouter, Depends, HTTPException
|
|||||||
from fastapi.responses import StreamingResponse
|
from fastapi.responses import StreamingResponse
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from .. import config, history, models
|
from .. import config, models
|
||||||
|
from ..services import history
|
||||||
from ..database import Generation as DBGeneration, get_db
|
from ..database import Generation as DBGeneration, get_db
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
@@ -28,7 +29,7 @@ async def preview_effects(
|
|||||||
if (gen.status or "completed") != "completed":
|
if (gen.status or "completed") != "completed":
|
||||||
raise HTTPException(status_code=400, detail="Generation is not completed")
|
raise HTTPException(status_code=400, detail="Generation is not completed")
|
||||||
|
|
||||||
from .. import versions as versions_mod
|
from ..services import versions as versions_mod
|
||||||
from ..utils.effects import apply_effects, validate_effects_chain
|
from ..utils.effects import apply_effects, validate_effects_chain
|
||||||
from ..utils.audio import load_audio
|
from ..utils.audio import load_audio
|
||||||
|
|
||||||
@@ -73,7 +74,7 @@ async def get_available_effects():
|
|||||||
@router.get("/effects/presets", response_model=list[models.EffectPresetResponse])
|
@router.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 ..services import effects as effects_mod
|
||||||
|
|
||||||
return effects_mod.list_presets(db)
|
return effects_mod.list_presets(db)
|
||||||
|
|
||||||
@@ -81,7 +82,7 @@ async def list_effect_presets(db: Session = Depends(get_db)):
|
|||||||
@router.get("/effects/presets/{preset_id}", response_model=models.EffectPresetResponse)
|
@router.get("/effects/presets/{preset_id}", response_model=models.EffectPresetResponse)
|
||||||
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 ..services 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:
|
||||||
@@ -95,7 +96,7 @@ async def create_effect_preset(
|
|||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""Create a new effect preset."""
|
"""Create a new effect preset."""
|
||||||
from .. import effects as effects_mod
|
from ..services import effects as effects_mod
|
||||||
|
|
||||||
try:
|
try:
|
||||||
return effects_mod.create_preset(data, db)
|
return effects_mod.create_preset(data, db)
|
||||||
@@ -110,7 +111,7 @@ async def update_effect_preset(
|
|||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""Update an effect preset."""
|
"""Update an effect preset."""
|
||||||
from .. import effects as effects_mod
|
from ..services 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)
|
||||||
@@ -124,7 +125,7 @@ async def update_effect_preset(
|
|||||||
@router.delete("/effects/presets/{preset_id}")
|
@router.delete("/effects/presets/{preset_id}")
|
||||||
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 ..services 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):
|
||||||
@@ -147,7 +148,7 @@ async def list_generation_versions(
|
|||||||
if not gen:
|
if not gen:
|
||||||
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 ..services import versions as versions_mod
|
||||||
|
|
||||||
return versions_mod.list_versions(generation_id, db)
|
return versions_mod.list_versions(generation_id, db)
|
||||||
|
|
||||||
@@ -168,7 +169,7 @@ async def apply_effects_to_generation(
|
|||||||
if (gen.status or "completed") != "completed":
|
if (gen.status or "completed") != "completed":
|
||||||
raise HTTPException(status_code=400, detail="Generation is not completed")
|
raise HTTPException(status_code=400, detail="Generation is not completed")
|
||||||
|
|
||||||
from .. import versions as versions_mod
|
from ..services import versions as versions_mod
|
||||||
from ..utils.effects import apply_effects, validate_effects_chain
|
from ..utils.effects import apply_effects, validate_effects_chain
|
||||||
from ..utils.audio import load_audio, save_audio
|
from ..utils.audio import load_audio, save_audio
|
||||||
|
|
||||||
@@ -227,7 +228,7 @@ async def set_default_version(
|
|||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""Set a specific version as the default for a generation."""
|
"""Set a specific version as the default for a generation."""
|
||||||
from .. import versions as versions_mod
|
from ..services import versions as versions_mod
|
||||||
|
|
||||||
version = versions_mod.get_version(version_id, db)
|
version = versions_mod.get_version(version_id, db)
|
||||||
if not version or version.generation_id != generation_id:
|
if not version or version.generation_id != generation_id:
|
||||||
@@ -246,7 +247,7 @@ async def delete_generation_version(
|
|||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""Delete a version. Cannot delete the last remaining version."""
|
"""Delete a version. Cannot delete the last remaining version."""
|
||||||
from .. import versions as versions_mod
|
from ..services import versions as versions_mod
|
||||||
|
|
||||||
version = versions_mod.get_version(version_id, db)
|
version = versions_mod.get_version(version_id, db)
|
||||||
if not version or version.generation_id != generation_id:
|
if not version or version.generation_id != generation_id:
|
||||||
|
|||||||
@@ -7,7 +7,8 @@ from fastapi import APIRouter, Depends, HTTPException
|
|||||||
from fastapi.responses import StreamingResponse
|
from fastapi.responses import StreamingResponse
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from .. import history, models, profiles, tts
|
from .. import models
|
||||||
|
from ..services import history, profiles, tts
|
||||||
from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db
|
from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db
|
||||||
from ..services.generation import run_generation
|
from ..services.generation import run_generation
|
||||||
from ..services.task_queue import enqueue_generation
|
from ..services.task_queue import enqueue_generation
|
||||||
|
|||||||
@@ -8,9 +8,10 @@ import torch
|
|||||||
from fastapi import APIRouter, Depends
|
from fastapi import APIRouter, Depends
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from .. import config, models, tts
|
from .. import config, models
|
||||||
|
from ..services import tts
|
||||||
from ..database import get_db
|
from ..database import get_db
|
||||||
from ..platform_detect import get_backend_type
|
from ..utils.platform_detect import get_backend_type
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,8 @@ from fastapi import APIRouter, Depends, File, HTTPException, UploadFile
|
|||||||
from fastapi.responses import FileResponse, StreamingResponse
|
from fastapi.responses import FileResponse, StreamingResponse
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from .. import export_import, history, models
|
from .. import models
|
||||||
|
from ..services import export_import, history
|
||||||
from ..app import safe_content_disposition
|
from ..app import safe_content_disposition
|
||||||
from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db
|
from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db
|
||||||
|
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from fastapi.responses import StreamingResponse
|
|||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from .. import models
|
from .. import models
|
||||||
from ..platform_detect import get_backend_type
|
from ..utils.platform_detect import get_backend_type
|
||||||
from ..services.task_queue import create_background_task
|
from ..services.task_queue import create_background_task
|
||||||
from ..utils.progress import get_progress_manager
|
from ..utils.progress import get_progress_manager
|
||||||
from ..utils.tasks import get_task_manager
|
from ..utils.tasks import get_task_manager
|
||||||
@@ -50,7 +50,7 @@ def _copy_with_progress(src: Path, dst: Path, progress_manager, copied_so_far: i
|
|||||||
@router.post("/models/load")
|
@router.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."""
|
||||||
from .. import tts
|
from ..services import tts
|
||||||
|
|
||||||
try:
|
try:
|
||||||
tts_model = tts.get_tts_model()
|
tts_model = tts.get_tts_model()
|
||||||
@@ -63,7 +63,7 @@ async def load_model(model_size: str = "1.7B"):
|
|||||||
@router.post("/models/unload")
|
@router.post("/models/unload")
|
||||||
async def unload_model():
|
async def unload_model():
|
||||||
"""Unload the default Qwen TTS model to free memory."""
|
"""Unload the default Qwen TTS model to free memory."""
|
||||||
from .. import tts
|
from ..services import tts
|
||||||
|
|
||||||
try:
|
try:
|
||||||
tts.unload_tts_model()
|
tts.unload_tts_model()
|
||||||
|
|||||||
@@ -9,10 +9,11 @@ from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile
|
|||||||
from fastapi.responses import FileResponse, StreamingResponse
|
from fastapi.responses import FileResponse, StreamingResponse
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from .. import channels, config, export_import, models, profiles
|
from .. import config, models
|
||||||
from ..app import safe_content_disposition
|
from ..app import safe_content_disposition
|
||||||
from ..database import VoiceProfile as DBVoiceProfile, get_db
|
from ..database import VoiceProfile as DBVoiceProfile, get_db
|
||||||
from ..profiles import _profile_to_response
|
from ..services import channels, export_import, profiles
|
||||||
|
from ..services.profiles import _profile_to_response
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|
||||||
|
|||||||
@@ -6,7 +6,8 @@ from fastapi import APIRouter, Depends, HTTPException
|
|||||||
from fastapi.responses import StreamingResponse
|
from fastapi.responses import StreamingResponse
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from .. import database, models, stories
|
from .. import database, models
|
||||||
|
from ..services import stories
|
||||||
from ..app import safe_content_disposition
|
from ..app import safe_content_disposition
|
||||||
from ..database import get_db
|
from ..database import get_db
|
||||||
|
|
||||||
|
|||||||
@@ -6,7 +6,8 @@ from pathlib import Path
|
|||||||
|
|
||||||
from fastapi import APIRouter, File, Form, HTTPException, UploadFile
|
from fastapi import APIRouter, File, Form, HTTPException, UploadFile
|
||||||
|
|
||||||
from .. import models, transcribe
|
from .. import models
|
||||||
|
from ..services import transcribe
|
||||||
from ..services.task_queue import create_background_task
|
from ..services.task_queue import create_background_task
|
||||||
from ..utils.tasks import get_task_manager
|
from ..utils.tasks import get_task_manager
|
||||||
|
|
||||||
|
|||||||
@@ -7,14 +7,14 @@ from datetime import datetime
|
|||||||
import uuid
|
import uuid
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from .models import (
|
from ..models import (
|
||||||
AudioChannelCreate,
|
AudioChannelCreate,
|
||||||
AudioChannelUpdate,
|
AudioChannelUpdate,
|
||||||
AudioChannelResponse,
|
AudioChannelResponse,
|
||||||
ChannelVoiceAssignment,
|
ChannelVoiceAssignment,
|
||||||
ProfileChannelAssignment,
|
ProfileChannelAssignment,
|
||||||
)
|
)
|
||||||
from .database import (
|
from ..database import (
|
||||||
AudioChannel as DBAudioChannel,
|
AudioChannel as DBAudioChannel,
|
||||||
ChannelDeviceMapping as DBChannelDeviceMapping,
|
ChannelDeviceMapping as DBChannelDeviceMapping,
|
||||||
ProfileChannelMapping as DBProfileChannelMapping,
|
ProfileChannelMapping as DBProfileChannelMapping,
|
||||||
@@ -11,8 +11,8 @@ from typing import List, Optional
|
|||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
from sqlalchemy.exc import IntegrityError
|
from sqlalchemy.exc import IntegrityError
|
||||||
|
|
||||||
from .database import EffectPreset as DBEffectPreset
|
from ..database import EffectPreset as DBEffectPreset
|
||||||
from .models import EffectPresetResponse, EffectPresetCreate, EffectPresetUpdate, EffectConfig
|
from ..models import EffectPresetResponse, EffectPresetCreate, EffectPresetUpdate, EffectConfig
|
||||||
|
|
||||||
|
|
||||||
def _preset_response(p: DBEffectPreset) -> EffectPresetResponse:
|
def _preset_response(p: DBEffectPreset) -> EffectPresetResponse:
|
||||||
@@ -12,11 +12,11 @@ from pathlib import Path
|
|||||||
from typing import Optional
|
from typing import Optional
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from .models import VoiceProfileResponse
|
from ..models import VoiceProfileResponse
|
||||||
from .database import VoiceProfile as DBVoiceProfile, ProfileSample as DBProfileSample, Generation as DBGeneration, GenerationVersion as DBGenerationVersion
|
from ..database import VoiceProfile as DBVoiceProfile, ProfileSample as DBProfileSample, Generation as DBGeneration, GenerationVersion as DBGenerationVersion
|
||||||
from .profiles import create_profile, add_profile_sample
|
from .profiles import create_profile, add_profile_sample
|
||||||
from .models import VoiceProfileCreate
|
from ..models import VoiceProfileCreate
|
||||||
from . import config
|
from .. import config
|
||||||
|
|
||||||
|
|
||||||
def _get_unique_profile_name(name: str, db: Session) -> str:
|
def _get_unique_profile_name(name: str, db: Session) -> str:
|
||||||
@@ -346,7 +346,7 @@ async def import_generation_from_zip(file_bytes: bytes, db: Session) -> dict:
|
|||||||
import tempfile
|
import tempfile
|
||||||
import shutil
|
import shutil
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from . import config
|
from .. import config
|
||||||
|
|
||||||
zip_buffer = io.BytesIO(file_bytes)
|
zip_buffer = io.BytesIO(file_bytes)
|
||||||
|
|
||||||
@@ -19,7 +19,8 @@ from __future__ import annotations
|
|||||||
import traceback
|
import traceback
|
||||||
from typing import Literal, Optional
|
from typing import Literal, Optional
|
||||||
|
|
||||||
from .. import config, history, profiles
|
from .. import config
|
||||||
|
from . import history, profiles
|
||||||
from ..database import get_db
|
from ..database import get_db
|
||||||
from ..utils.tasks import get_task_manager
|
from ..utils.tasks import get_task_manager
|
||||||
|
|
||||||
@@ -151,7 +152,7 @@ def _save_generate(
|
|||||||
Returns the final audio path (processed if effects were applied,
|
Returns the final audio path (processed if effects were applied,
|
||||||
otherwise clean).
|
otherwise clean).
|
||||||
"""
|
"""
|
||||||
from .. import versions as versions_mod
|
from . import versions as versions_mod
|
||||||
|
|
||||||
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)
|
||||||
@@ -221,7 +222,7 @@ def _save_regenerate(
|
|||||||
|
|
||||||
Returns the audio path.
|
Returns the audio path.
|
||||||
"""
|
"""
|
||||||
from .. import versions as versions_mod
|
from . import versions as versions_mod
|
||||||
|
|
||||||
suffix = version_id[:8] if version_id else generation_id[:8]
|
suffix = version_id[:8] if version_id else generation_id[:8]
|
||||||
audio_path = config.get_generations_dir() / f"{generation_id}_{suffix}.wav"
|
audio_path = config.get_generations_dir() / f"{generation_id}_{suffix}.wav"
|
||||||
|
|||||||
@@ -10,9 +10,9 @@ from pathlib import Path
|
|||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
from sqlalchemy import or_
|
from sqlalchemy import or_
|
||||||
|
|
||||||
from .models import GenerationRequest, GenerationResponse, HistoryQuery, HistoryResponse, HistoryListResponse, GenerationVersionResponse, EffectConfig
|
from ..models import GenerationRequest, GenerationResponse, HistoryQuery, HistoryResponse, HistoryListResponse, GenerationVersionResponse, EffectConfig
|
||||||
from .database import Generation as DBGeneration, GenerationVersion as DBGenerationVersion, VoiceProfile as DBVoiceProfile
|
from ..database import Generation as DBGeneration, GenerationVersion as DBGenerationVersion, VoiceProfile as DBVoiceProfile
|
||||||
from . import config
|
from .. import config
|
||||||
|
|
||||||
|
|
||||||
def _get_versions_for_generation(generation_id: str, db: Session) -> tuple:
|
def _get_versions_for_generation(generation_id: str, db: Session) -> tuple:
|
||||||
@@ -10,23 +10,23 @@ from pathlib import Path
|
|||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
from sqlalchemy import func, select
|
from sqlalchemy import func, select
|
||||||
|
|
||||||
from .models import (
|
from ..models import (
|
||||||
VoiceProfileCreate,
|
VoiceProfileCreate,
|
||||||
VoiceProfileResponse,
|
VoiceProfileResponse,
|
||||||
ProfileSampleCreate,
|
ProfileSampleCreate,
|
||||||
ProfileSampleResponse,
|
ProfileSampleResponse,
|
||||||
)
|
)
|
||||||
from .database import (
|
from ..database import (
|
||||||
VoiceProfile as DBVoiceProfile,
|
VoiceProfile as DBVoiceProfile,
|
||||||
ProfileSample as DBProfileSample,
|
ProfileSample as DBProfileSample,
|
||||||
Generation as DBGeneration,
|
Generation as DBGeneration,
|
||||||
)
|
)
|
||||||
from .models import EffectConfig
|
from ..models import EffectConfig
|
||||||
from .utils.audio import validate_reference_audio, load_audio, save_audio
|
from ..utils.audio import validate_reference_audio, load_audio, save_audio
|
||||||
from .utils.images import validate_image, process_avatar
|
from ..utils.images import validate_image, process_avatar
|
||||||
from .utils.cache import _get_cache_dir, clear_profile_cache
|
from ..utils.cache import _get_cache_dir, clear_profile_cache
|
||||||
from .tts import get_tts_model
|
from .tts import get_tts_model
|
||||||
from . import config
|
from .. import config
|
||||||
import json as _json
|
import json as _json
|
||||||
|
|
||||||
|
|
||||||
@@ -389,7 +389,7 @@ async def create_voice_prompt_for_profile(
|
|||||||
Returns:
|
Returns:
|
||||||
Voice prompt dictionary
|
Voice prompt dictionary
|
||||||
"""
|
"""
|
||||||
from .backends import get_tts_backend_for_engine
|
from ..backends import get_tts_backend_for_engine
|
||||||
|
|
||||||
samples = db.query(DBProfileSample).filter_by(profile_id=profile_id).all()
|
samples = db.query(DBProfileSample).filter_by(profile_id=profile_id).all()
|
||||||
|
|
||||||
@@ -10,7 +10,7 @@ from pathlib import Path
|
|||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
from sqlalchemy import func
|
from sqlalchemy import func
|
||||||
|
|
||||||
from .models import (
|
from ..models import (
|
||||||
StoryCreate,
|
StoryCreate,
|
||||||
StoryResponse,
|
StoryResponse,
|
||||||
StoryDetailResponse,
|
StoryDetailResponse,
|
||||||
@@ -22,14 +22,14 @@ from .models import (
|
|||||||
StoryItemSplit,
|
StoryItemSplit,
|
||||||
StoryItemVersionUpdate,
|
StoryItemVersionUpdate,
|
||||||
)
|
)
|
||||||
from .database import (
|
from ..database import (
|
||||||
Story as DBStory,
|
Story as DBStory,
|
||||||
StoryItem as DBStoryItem,
|
StoryItem as DBStoryItem,
|
||||||
Generation as DBGeneration,
|
Generation as DBGeneration,
|
||||||
VoiceProfile as DBVoiceProfile,
|
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
|
||||||
|
|
||||||
|
|
||||||
@@ -754,7 +754,7 @@ 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 = (
|
version = (
|
||||||
db.query(DBGenerationVersion)
|
db.query(DBGenerationVersion)
|
||||||
@@ -820,7 +820,7 @@ async def export_story_audio(
|
|||||||
# 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:
|
||||||
@@ -3,7 +3,7 @@ STT (Speech-to-Text) module - delegates to backend abstraction layer.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
from .backends import get_stt_backend, STTBackend
|
from ..backends import get_stt_backend, STTBackend
|
||||||
|
|
||||||
|
|
||||||
def get_whisper_model() -> STTBackend:
|
def get_whisper_model() -> STTBackend:
|
||||||
@@ -7,7 +7,7 @@ import numpy as np
|
|||||||
import io
|
import io
|
||||||
import soundfile as sf
|
import soundfile as sf
|
||||||
|
|
||||||
from .backends import get_tts_backend, TTSBackend
|
from ..backends import get_tts_backend, TTSBackend
|
||||||
|
|
||||||
|
|
||||||
def get_tts_model() -> TTSBackend:
|
def get_tts_model() -> TTSBackend:
|
||||||
@@ -14,12 +14,12 @@ from typing import List, Optional
|
|||||||
|
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from .database import (
|
from ..database import (
|
||||||
GenerationVersion as DBGenerationVersion,
|
GenerationVersion as DBGenerationVersion,
|
||||||
Generation as DBGeneration,
|
Generation as DBGeneration,
|
||||||
)
|
)
|
||||||
from .models import GenerationVersionResponse, EffectConfig
|
from ..models import GenerationVersionResponse, EffectConfig
|
||||||
from . import config
|
from .. import config
|
||||||
|
|
||||||
|
|
||||||
def _version_response(v: DBGenerationVersion) -> GenerationVersionResponse:
|
def _version_response(v: DBGenerationVersion) -> GenerationVersionResponse:
|
||||||
Reference in New Issue
Block a user