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