- Implements POST /v1/audio/speech endpoint (OpenAI API compatible). - Integrates Kokoro and XTTS engines (including dependencies and implementations). - Updates main application to register new engines and router. - Adds unit tests for OpenAI compatibility. - Updates requirements.txt for new engines.
254 lines
10 KiB
Python
254 lines
10 KiB
Python
"""
|
||
NovaAi – TTS-Engine-Hub
|
||
main.py
|
||
Version: v0.1.0
|
||
|
||
Description:
|
||
Refactored main.py: uses utils modules for chunking, concat, and cache key generation.
|
||
All endpoints, features, and logic as before, but cleaner and more modular.
|
||
|
||
Author: Abby (ChatGPT)
|
||
Date: 2025-07-23
|
||
Canvas: main.py
|
||
"""
|
||
|
||
import asyncio
|
||
from fastapi import FastAPI, HTTPException, Query
|
||
from fastapi.responses import JSONResponse, FileResponse
|
||
from pydantic import BaseModel
|
||
import os
|
||
import base64
|
||
import shutil
|
||
import uvicorn
|
||
import logging
|
||
|
||
from app.config import settings
|
||
from app.engines.piper import PiperEngine
|
||
from app.engines.styletts import StyleTTSEngine
|
||
from app.engines.chattts import ChatTTSEngine
|
||
from app.engines.f5_tts import F5TTSEngine
|
||
from app.engines.kokoro import KokoroEngine
|
||
from app.engines.xtts import XTTSEngine
|
||
from app.utils.text import chunk_text
|
||
from app.utils.audio import concat_audio
|
||
from app.utils.cache import build_cache_key
|
||
from app.routers import openai_compatible
|
||
|
||
# Configure logging based on settings
|
||
logging.basicConfig(level=settings.LOG_LEVEL, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
|
||
logger = logging.getLogger(__name__)
|
||
|
||
# Explicitly configure uvicorn loggers
|
||
logging.getLogger("uvicorn.access").setLevel(settings.LOG_LEVEL)
|
||
logging.getLogger("uvicorn.error").setLevel(settings.LOG_LEVEL)
|
||
logging.getLogger("uvicorn.server").setLevel(settings.LOG_LEVEL)
|
||
|
||
# --- Master list of all possible engine classes. ---
|
||
ALL_ENGINES = {
|
||
"piper": PiperEngine,
|
||
"styletts": StyleTTSEngine,
|
||
"chattts": ChatTTSEngine,
|
||
"f5-tts": F5TTSEngine,
|
||
"kokoro": KokoroEngine,
|
||
"xtts": XTTSEngine,
|
||
}
|
||
|
||
def create_app():
|
||
app = FastAPI(
|
||
title="NovaAi – TTS-Engine-Hub",
|
||
version="0.3.0",
|
||
description="Local-first, modular multi-engine TTS server for your homelab and automation."
|
||
)
|
||
|
||
# Dynamically build the registry of active engines based on settings.
|
||
# This registry is local to the app instance created by this function.
|
||
app.ENGINE_REGISTRY = {}
|
||
for engine_name in settings.ACTIVE_ENGINES:
|
||
if engine_name in ALL_ENGINES:
|
||
logger.info(f"Activating engine: {engine_name}")
|
||
app.ENGINE_REGISTRY[engine_name] = ALL_ENGINES[engine_name]()
|
||
else:
|
||
logger.warning(f"Engine '{engine_name}' requested in config but not found in ALL_ENGINES.")
|
||
|
||
# Ensure the audio asset/cache directory exists.
|
||
os.makedirs(settings.AUDIO_CACHE_DIR, exist_ok=True)
|
||
|
||
# Register Routers
|
||
app.include_router(openai_compatible.router)
|
||
|
||
class TTSRequest(BaseModel):
|
||
text: str
|
||
engine: str
|
||
model: str = None
|
||
speaker: str = None
|
||
format: str = "ogg"
|
||
chunking: bool = False
|
||
|
||
@app.post("/tts")
|
||
async def tts_endpoint(req: TTSRequest, as_base64: bool = Query(False, alias="as")):
|
||
# --- 1. Check for engine and handle health ---
|
||
engine = app.ENGINE_REGISTRY.get(req.engine.lower())
|
||
if not engine:
|
||
raise HTTPException(status_code=404, detail=f"Engine '{req.engine}' not found.")
|
||
|
||
health = engine.healthcheck()
|
||
if health.get("status") != "ok":
|
||
raise HTTPException(status_code=503, detail=f"Engine '{req.engine}' is not available. Status: {health.get('status')}")
|
||
|
||
# --- 2. Input validation ---
|
||
available_models = engine.list_models()
|
||
if req.model and available_models and req.model not in available_models:
|
||
raise HTTPException(status_code=400, detail=f"Model '{req.model}' not found for engine '{req.engine}'. Available models: {available_models}")
|
||
|
||
available_voices = engine.list_voices(req.model)
|
||
if req.speaker and available_voices and req.speaker not in available_voices:
|
||
raise HTTPException(status_code=400, detail=f"Speaker '{req.speaker}' not found for model '{req.model}'. Available speakers: {available_voices}")
|
||
|
||
# --- 3. Check cache ---
|
||
cache_key = build_cache_key(req)
|
||
ext = f'.{req.format.lower()}'
|
||
output_filename = f"tts_{cache_key}{ext}"
|
||
output_filepath = os.path.join(settings.AUDIO_CACHE_DIR, output_filename)
|
||
|
||
is_cached = await asyncio.to_thread(os.path.isfile, output_filepath)
|
||
if is_cached:
|
||
if as_base64:
|
||
audio_bytes = await asyncio.to_thread(lambda: open(output_filepath, "rb").read())
|
||
audio_b64 = base64.b64encode(audio_bytes).decode("utf-8")
|
||
return JSONResponse({
|
||
"engine": req.engine, "model": req.model, "speaker": req.speaker, "format": req.format,
|
||
"audio_base64": audio_b64, "chunking": req.chunking, "message": "Audio from cache, base64 included"
|
||
})
|
||
return JSONResponse({
|
||
"engine": req.engine, "model": req.model, "speaker": req.speaker, "format": req.format,
|
||
"audio_url": f"/audio/{output_filename}", "cached": True, "chunking": req.chunking,
|
||
"message": "Audio served from cache. Download from audio_url"
|
||
})
|
||
|
||
# --- 4. Synthesize audio ---
|
||
try:
|
||
if req.chunking and len(req.text) > 250:
|
||
chunks = chunk_text(req.text, maxlen=250)
|
||
synthesis_tasks = [engine.synthesize(c, speaker=req.speaker, model=req.model, fmt=req.format) for c in chunks]
|
||
chunk_files = await asyncio.gather(*synthesis_tasks)
|
||
synthesized_path = await asyncio.to_thread(concat_audio, chunk_files, req.format.lower())
|
||
else:
|
||
synthesized_path = await engine.synthesize(req.text, speaker=req.speaker, model=req.model, fmt=req.format)
|
||
except (RuntimeError, ValueError, FileNotFoundError) as e:
|
||
raise HTTPException(status_code=500, detail=f"Error during synthesis: {e}")
|
||
except Exception as e:
|
||
raise HTTPException(status_code=500, detail=f"An unexpected error occurred: {e}")
|
||
|
||
# --- 5. Cache and return result ---
|
||
await asyncio.to_thread(shutil.copy, synthesized_path, output_filepath)
|
||
|
||
# If the synth created a temp file in a different directory, clean it up
|
||
if settings.AUDIO_CACHE_DIR not in os.path.abspath(synthesized_path):
|
||
await asyncio.to_thread(os.remove, synthesized_path)
|
||
|
||
if as_base64:
|
||
audio_bytes = await asyncio.to_thread(lambda: open(output_filepath, "rb").read())
|
||
audio_b64 = base64.b64encode(audio_bytes).decode("utf-8")
|
||
return JSONResponse({
|
||
"engine": req.engine, "model": req.model, "speaker": req.speaker, "format": req.format,
|
||
"audio_base64": audio_b64, "chunking": req.chunking, "message": "Audio from synth, base64 included"
|
||
})
|
||
|
||
return JSONResponse({
|
||
"engine": req.engine, "model": req.model, "speaker": req.speaker, "format": req.format,
|
||
"audio_url": f"/audio/{output_filename}", "cached": False, "chunking": req.chunking,
|
||
"message": "Synthesized new audio. Download from audio_url"
|
||
})
|
||
|
||
@app.get("/audio/{filename}")
|
||
async def audio_file(filename: str): # Made async
|
||
fpath = os.path.join(settings.AUDIO_CACHE_DIR, filename)
|
||
fpath_abs = os.path.abspath(fpath)
|
||
cache_dir_abs = os.path.abspath(settings.AUDIO_CACHE_DIR)
|
||
if not await asyncio.to_thread(os.path.isfile, fpath) or not fpath_abs.startswith(cache_dir_abs):
|
||
raise HTTPException(status_code=404, detail="Audio file not found")
|
||
media_type = "audio/wav" if filename.endswith(".wav") else (
|
||
"audio/ogg" if filename.endswith(".ogg") else "audio/mpeg"
|
||
)
|
||
return FileResponse(fpath, media_type=media_type, filename=filename)
|
||
|
||
@app.get("/engines")
|
||
def engines_endpoint():
|
||
engines = {}
|
||
for name, engine in app.ENGINE_REGISTRY.items():
|
||
engines[name] = engine.healthcheck()
|
||
return engines
|
||
|
||
@app.get("/models")
|
||
def models_endpoint():
|
||
result = {}
|
||
for name, engine in app.ENGINE_REGISTRY.items():
|
||
try:
|
||
result[name] = engine.list_models()
|
||
except Exception as e:
|
||
result[name] = {"error": str(e)}
|
||
return result
|
||
|
||
@app.get("/speakers")
|
||
def speakers_endpoint(engine: str, model: str = None):
|
||
e = app.ENGINE_REGISTRY.get(engine.lower())
|
||
if not e:
|
||
raise HTTPException(status_code=404, detail=f"Engine '{engine}' not found.")
|
||
try:
|
||
speakers = e.list_voices(model)
|
||
except Exception as err:
|
||
speakers = []
|
||
return {"engine": engine, "model": model, "speakers": speakers}
|
||
|
||
@app.get("/version")
|
||
def version():
|
||
return {"version": app.version}
|
||
|
||
@app.get("/health")
|
||
def health():
|
||
status = {name: engine.healthcheck()["status"] for name, engine in app.ENGINE_REGISTRY.items()}
|
||
return {"status": status, "detail": "API and engines loaded"}
|
||
|
||
@app.on_event("startup")
|
||
async def startup_event():
|
||
"""
|
||
On startup, check for the existence of the models directory.
|
||
This helps prevent race conditions with volume mounts.
|
||
"""
|
||
if os.getenv("SKIP_MODEL_CHECK", "false").lower() == "true":
|
||
logger.info("Skipping model directory check (SKIP_MODEL_CHECK=true)")
|
||
return
|
||
|
||
model_path = "/models/piper"
|
||
max_retries = 10
|
||
retry_delay = 2 # seconds
|
||
|
||
for i in range(max_retries):
|
||
if os.path.exists(model_path) and os.listdir(model_path):
|
||
print(f"Models directory '{model_path}' found and is not empty.")
|
||
return
|
||
print(f"Waiting for models directory '{model_path}' to be available... (Attempt {i+1}/{max_retries})")
|
||
await asyncio.sleep(retry_delay)
|
||
|
||
print(f"CRITICAL: Models directory '{model_path}' not found or is empty after {max_retries * retry_delay} seconds. Shutting down.")
|
||
# This will cause the server to exit if run with --lifespan on,
|
||
# or at least log a critical failure.
|
||
# In a real production setup, this should trigger a process manager to restart or alert.
|
||
raise RuntimeError("Models not found on startup")
|
||
|
||
return app
|
||
|
||
# If main.py is executed directly, create the app and run uvicorn
|
||
if __name__ == "__main__":
|
||
app_instance = create_app()
|
||
uvicorn.run(
|
||
app_instance, # Pass the app instance
|
||
host=settings.HOST,
|
||
port=settings.PORT,
|
||
reload=True
|
||
)
|
||
|
||
app = create_app()
|
||
|
||
|