124 lines
4.0 KiB
Python
124 lines
4.0 KiB
Python
from __future__ import annotations
|
|
|
|
import logging
|
|
import os
|
|
|
|
from realtime_voice_service.providers.base import BaseLLM, BaseSTT, BaseTTS
|
|
from realtime_voice_service.providers.llm import OllamaLLM, OpenAILLM
|
|
from realtime_voice_service.providers.stt import ElevenLabsSTT, FallbackSTT, YandexSpeechKitSTT
|
|
from realtime_voice_service.providers.stt_openai import OpenAISTT
|
|
from realtime_voice_service.providers.tts import ElevenLabsTTS
|
|
|
|
|
|
LOGGER = logging.getLogger("uvicorn.error")
|
|
|
|
|
|
def _first_env(names: tuple[str, ...], default: str = "") -> str:
|
|
for name in names:
|
|
value = os.getenv(name)
|
|
if value is not None and value.strip():
|
|
return value.strip()
|
|
return default
|
|
|
|
|
|
def _stt_provider_name() -> str:
|
|
return _first_env(
|
|
(
|
|
"STT_PROVIDER",
|
|
"REALTIME_VOICE_STT_PROVIDER",
|
|
"AI_VOICE_ASR_PROVIDER",
|
|
),
|
|
"elevenlabs",
|
|
).lower()
|
|
|
|
|
|
def _stt_fallback_provider_name() -> str:
|
|
return _first_env(
|
|
(
|
|
"STT_FALLBACK_PROVIDER",
|
|
"REALTIME_VOICE_STT_FALLBACK_PROVIDER",
|
|
)
|
|
).lower()
|
|
|
|
|
|
def _llm_provider_name() -> str:
|
|
return _first_env(("LLM_PROVIDER", "REALTIME_VOICE_LLM_PROVIDER"), "openai").lower()
|
|
|
|
|
|
def _tts_provider_name() -> str:
|
|
return _first_env(("TTS_PROVIDER", "REALTIME_VOICE_TTS_PROVIDER"), "elevenlabs").lower()
|
|
|
|
|
|
def create_llm_provider() -> BaseLLM:
|
|
provider = _llm_provider_name()
|
|
if provider in {"openai", "openai_chat", "gpt"}:
|
|
LOGGER.info("LLM provider selected: provider=%s", provider)
|
|
return OpenAILLM()
|
|
if provider in {"ollama", "local", "qwen"}:
|
|
LOGGER.info("LLM provider selected: provider=%s", provider)
|
|
return OllamaLLM()
|
|
raise RuntimeError(f"Unsupported LLM provider: {provider}")
|
|
|
|
|
|
def _build_single_stt_provider(
|
|
provider: str,
|
|
*,
|
|
input_sample_rate_hz: int,
|
|
target_sample_rate_hz: int,
|
|
) -> BaseSTT:
|
|
normalized = provider.lower().strip()
|
|
if normalized in {"openai", "whisper", "openai_whisper"}:
|
|
# Whisper accepts the configured WAV container; keep the exact AudioSocket PCM, no resampling.
|
|
return OpenAISTT(input_sample_rate_hz=input_sample_rate_hz)
|
|
if normalized in {"elevenlabs", "eleven_labs", "scribe"}:
|
|
return ElevenLabsSTT(
|
|
input_sample_rate_hz=input_sample_rate_hz,
|
|
target_sample_rate_hz=target_sample_rate_hz,
|
|
)
|
|
if normalized in {"yandex", "yandex_speechkit", "speechkit"}:
|
|
return YandexSpeechKitSTT(
|
|
input_sample_rate_hz=input_sample_rate_hz,
|
|
target_sample_rate_hz=target_sample_rate_hz,
|
|
)
|
|
raise RuntimeError(f"Unsupported STT provider: {provider}")
|
|
|
|
|
|
def create_stt_provider(
|
|
*,
|
|
input_sample_rate_hz: int,
|
|
target_sample_rate_hz: int,
|
|
) -> BaseSTT:
|
|
provider = _stt_provider_name()
|
|
primary = _build_single_stt_provider(
|
|
provider,
|
|
input_sample_rate_hz=input_sample_rate_hz,
|
|
target_sample_rate_hz=target_sample_rate_hz,
|
|
)
|
|
fallback_provider = _stt_fallback_provider_name()
|
|
if not fallback_provider or fallback_provider == provider:
|
|
LOGGER.info("STT provider selected: provider=%s", provider)
|
|
return primary
|
|
|
|
fallback = _build_single_stt_provider(
|
|
fallback_provider,
|
|
input_sample_rate_hz=input_sample_rate_hz,
|
|
target_sample_rate_hz=target_sample_rate_hz,
|
|
)
|
|
LOGGER.info("STT provider selected: provider=%s fallback=%s", provider, fallback_provider)
|
|
return FallbackSTT(primary=primary, fallback=fallback)
|
|
|
|
|
|
def create_tts_provider(
|
|
*,
|
|
target_sample_rate_hz: int,
|
|
output_format: str | None = None,
|
|
) -> BaseTTS:
|
|
provider = _tts_provider_name()
|
|
if provider in {"elevenlabs", "eleven_labs"}:
|
|
LOGGER.info("TTS provider selected: provider=%s", provider)
|
|
return ElevenLabsTTS(
|
|
output_format=output_format or f"pcm_{target_sample_rate_hz}",
|
|
target_sample_rate_hz=target_sample_rate_hz,
|
|
)
|
|
raise RuntimeError(f"Unsupported TTS provider: {provider}")
|