Files
2026-05-01 17:50:41 +05:00

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