Files
audio-engine-hub/app/main.py
stephan fff0252d52 feat: Add OpenAI-compatible TTS endpoint and engines
- 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.
2025-12-09 12:45:17 +01:00

254 lines
10 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
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()