tts-voice-creator-clone-and.../core/routing.py

278 lines
10 KiB
Python

"""TTS routing rules: load/save/resolve routes, language detection, routing log."""
from __future__ import annotations
import json
import re
from fastapi import Request
from core.constants import TTS_ROUTES_FILE, _routing_log_add
from core.database import state_get, state_put, state_updated
# ── Route token / backend validation ─────────────────────────────────────────
_ROUTE_LANGS = {"AUTO", "*", "EN", "DE", "FR", "ES", "IT", "PT", "NL", "PL"}
_ROUTE_BACKENDS = {"voice_clone", "streaming", "voice_design", "nvidia_magpie", "nvidia_zeroshot", "nvidia_flow"}
def _clean_route_token(value: str, default: str = "*") -> str:
value = str(value or "").strip()
if not value:
return default
value = re.sub(r"[^A-Za-z0-9_\-\.\* ]+", "_", value)
return value[:80] or default
def _clean_route_sound(value: str) -> str:
value = str(value or "").strip()
if not value:
return ""
value = re.sub(r"[^A-Za-z0-9_\-\.\* /]+", "_", value)
return value[:240]
def _clean_route_backend(value: str) -> str:
key = re.sub(r"[^a-z0-9]+", "_", str(value or "voice_clone").lower()).strip("_")
aliases = {
"clone": "voice_clone",
"base": "voice_clone",
"tts": "voice_clone",
"standard": "voice_clone",
"voiceclone": "voice_clone",
"voice_clone_base": "voice_clone",
"stream": "streaming",
"tts_streaming": "streaming",
"voice_design": "voice_design",
"voicedesign": "voice_design",
"design": "voice_design",
"nvidia": "nvidia_magpie",
"magpie": "nvidia_magpie",
"nvidia_tts": "nvidia_magpie",
"nvidia_magpie_tts": "nvidia_magpie",
"nvidia_clone": "nvidia_zeroshot",
"nvidia_zeroshot_tts": "nvidia_zeroshot",
"magpie_zeroshot": "nvidia_zeroshot",
"zeroshot": "nvidia_zeroshot",
"zero_shot": "nvidia_zeroshot",
"nvidia_flow_tts": "nvidia_flow",
"magpie_flow": "nvidia_flow",
"flow": "nvidia_flow",
}
key = aliases.get(key, key)
return key if key in _ROUTE_BACKENDS else "voice_clone"
def _normalize_route(rule: dict, idx: int = 0) -> dict:
lang = str(rule.get("language", "*") or "*").strip().upper()
if lang not in _ROUTE_LANGS:
lang = "*"
output_voice = _clean_route_token(rule.get("output_voice", ""), "")
return {
"id": _clean_route_token(rule.get("id", f"route_{idx+1}"), f"route_{idx+1}"),
"enabled": bool(rule.get("enabled", True)),
"app": _clean_route_token(rule.get("app", "Open WebUI"), "Open WebUI"),
"input_voice": _clean_route_token(rule.get("input_voice", "default"), "default"),
"language": lang,
"backend": _clean_route_backend(rule.get("backend", "voice_clone")),
"output_voice": output_voice,
"before_sound": _clean_route_sound(rule.get("before_sound", "")),
"after_sound": _clean_route_sound(rule.get("after_sound", "")),
}
_routes_cache: list[dict] | None = None
_routes_cache_mtime: float = -1.0
_routes_cache_db_updated: str = ""
def _load_tts_routes() -> list[dict]:
global _routes_cache, _routes_cache_mtime, _routes_cache_db_updated
mtime = TTS_ROUTES_FILE.stat().st_mtime if TTS_ROUTES_FILE.exists() else 0.0
db_updated = state_updated("tts_routes")
if (
_routes_cache is not None
and mtime == _routes_cache_mtime
and db_updated == _routes_cache_db_updated
):
return list(_routes_cache)
result: list[dict] = []
try:
if TTS_ROUTES_FILE.exists():
raw = json.loads(TTS_ROUTES_FILE.read_text())
routes = raw.get("routes", raw) if isinstance(raw, dict) else raw
if isinstance(routes, list):
result = [_normalize_route(r, i) for i, r in enumerate(routes) if isinstance(r, dict)]
except Exception:
result = []
try:
db_routes = state_get("tts_routes", None)
if isinstance(db_routes, dict):
db_routes = db_routes.get("routes")
if isinstance(db_routes, list):
routes = db_routes
result = [_normalize_route(r, i) for i, r in enumerate(routes) if isinstance(r, dict)]
except Exception:
pass
if not db_updated and result:
try:
state_put("tts_routes", {"routes": result})
db_updated = state_updated("tts_routes")
except Exception:
pass
_routes_cache = result
_routes_cache_mtime = mtime
_routes_cache_db_updated = db_updated
return list(result)
def _save_tts_routes(routes: list[dict]) -> None:
global _routes_cache, _routes_cache_mtime, _routes_cache_db_updated
TTS_ROUTES_FILE.parent.mkdir(parents=True, exist_ok=True)
clean = [_normalize_route(r, i) for i, r in enumerate(routes)]
state_put("tts_routes", {"routes": clean})
TTS_ROUTES_FILE.write_text(json.dumps({"routes": clean}, indent=2))
_routes_cache = clean
_routes_cache_mtime = TTS_ROUTES_FILE.stat().st_mtime
_routes_cache_db_updated = state_updated("tts_routes")
# ── App name helpers ──────────────────────────────────────────────────────────
def _app_key(value: str) -> str:
return re.sub(r"[^a-z0-9]+", "", str(value or "").lower())
def _canonical_app_name(value: str) -> str:
key = _app_key(value)
if key in {"openwebui", "openwebuiapp"}:
return "Open WebUI"
if key == "sillytavern":
return "SillyTavern"
if key in {"homeassistant", "ha"}:
return "Home Assistant"
return str(value or "").strip() or "Open WebUI"
def _request_app_name(request: Request) -> str:
explicit = request.headers.get("x-tts-app") or request.headers.get("x-client-app")
if explicit:
return _canonical_app_name(explicit)
header = (
request.headers.get("x-openwebui-app")
or request.headers.get("x-openwebui-user-name")
or request.headers.get("referer")
or request.headers.get("origin")
or request.headers.get("user-agent")
or ""
)
h = _app_key(header)
if "openwebui" in h:
return "Open WebUI"
if "sillytavern" in h:
return "SillyTavern"
if "homeassistant" in h:
return "Home Assistant"
return "Open WebUI"
# ── Language detection ────────────────────────────────────────────────────────
def _detect_text_language(text: str) -> str:
low = f" {text.lower()} "
if re.search(r"[äöüß]", low) or re.search(r"\b(und|der|die|das|ich|nicht|mit|ist|ein|eine|auf|für)\b", low):
return "DE"
if re.search(r"[éèêàçùœ]", low) or re.search(r"\b(et|le|la|les|des|une|avec|pour|est|pas|que)\b", low):
return "FR"
if re.search(r"[áéíóúñ¿¡]", low) or re.search(r"\b(el|la|los|las|una|con|para|que|pero|está|hola)\b", low):
return "ES"
if re.search(r"\b(il|lo|gli|una|con|per|che|ciao|grazie|sono|della)\b", low):
return "IT"
if re.search(r"[ãõç]", low) or re.search(r"\b(com|para|uma|que|não|está|obrigado)\b", low):
return "PT"
if re.search(r"\b(het|een|niet|met|voor|zijn|maar|dank|goede)\b", low):
return "NL"
if re.search(r"[ąćęłńóśźż]", low) or re.search(r"\b(jest|nie|tak|dla|oraz|dzień|dziękuję)\b", low):
return "PL"
return "EN"
# ── Route matching ────────────────────────────────────────────────────────────
def _route_specificity(rule: dict, app: str, voice: str, lang: str) -> tuple[int, int, int, int] | None:
if not rule.get("enabled", True) or not rule.get("output_voice"):
return None
r_app = str(rule.get("app", "*"))
r_voice = str(rule.get("input_voice", "*"))
r_lang = str(rule.get("language", "*")).upper()
app_ok = r_app == "*" or _app_key(r_app) == _app_key(app)
voice_ok = r_voice == "*" or r_voice.lower() == voice.lower()
lang_ok = r_lang in {"*", "AUTO"} or r_lang == lang
if not (app_ok and voice_ok and lang_ok):
return None
return (
1 if r_app != "*" else 0,
1 if r_voice != "*" else 0,
1 if r_lang not in {"*", "AUTO"} else 0,
0,
)
def _resolve_tts_route(app: str, voice: str, text: str) -> tuple[str, dict | None]:
lang = _detect_text_language(text)
best: tuple[tuple[int, int, int, int], dict] | None = None
for idx, rule in enumerate(_load_tts_routes()):
spec = _route_specificity(rule, app, voice, lang)
if spec is None:
continue
spec = (spec[0], spec[1], spec[2], -idx)
if best is None or spec > best[0]:
best = (spec, rule)
if not best:
return voice, None
routed = dict(best[1])
routed["detected_language"] = lang
routed["requested_voice"] = voice
routed["app"] = app
return str(routed["output_voice"]), routed
def _route_backend(route: dict | None, voice: str = "") -> str:
from core.presets import _resolve_virtual_voice
if _resolve_virtual_voice(voice):
return "voice_design"
return _clean_route_backend((route or {}).get("backend", "voice_clone"))
# ── Routing log request helper ────────────────────────────────────────────────
def _routing_log_request(
request: Request,
*,
status: str,
app: str,
requested_voice: str,
routed_voice: str,
backend: str,
route: dict | None,
response_format: str,
text: str,
**extra,
) -> None:
route = route or {}
_routing_log_add(
kind="proxy",
status=status,
app=app,
requested_voice=requested_voice,
routed_voice=routed_voice,
backend=backend,
language=str(route.get("detected_language", "")) or (_detect_text_language(text) if text else ""),
matched=bool(route),
route_id=str(route.get("id", "")),
response_format=response_format,
text_preview=text[:160],
client=request.client.host if request.client else "",
user_agent=str(request.headers.get("user-agent", ""))[:160],
**extra,
)