feat: migrate realtime voice service to OpenAI provider pipeline
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import inspect
|
||||
import logging
|
||||
import os
|
||||
import uuid
|
||||
@@ -8,11 +9,10 @@ from contextlib import asynccontextmanager
|
||||
|
||||
from fastapi import FastAPI, WebSocket
|
||||
|
||||
from realtime_voice_service.core.filler_audio import FillerAudioLibrary
|
||||
from realtime_voice_service.core.session import CallSession
|
||||
from realtime_voice_service.core.vad import SileroVADDetector
|
||||
from realtime_voice_service.providers.llm import OpenAILLM
|
||||
from realtime_voice_service.providers.stt import ElevenLabsSTT
|
||||
from realtime_voice_service.providers.tts import ElevenLabsTTS
|
||||
from realtime_voice_service.providers.factory import create_llm_provider, create_stt_provider, create_tts_provider
|
||||
from realtime_voice_service.transports.audiosocket import AudioSocketServer
|
||||
from realtime_voice_service.transports.base import BaseMediaTransport
|
||||
from realtime_voice_service.transports.websocket import WebSocketMediaTransport
|
||||
@@ -69,6 +69,11 @@ def _audiosocket_port() -> int:
|
||||
return max(_int_env("REALTIME_VOICE_AUDIOSOCKET_PORT", 9092), 1)
|
||||
|
||||
|
||||
def _sample_rate_hz() -> int:
|
||||
configured = max(_int_env("REALTIME_VOICE_SAMPLE_RATE_HZ", 8000), 1)
|
||||
return configured if configured in {8000, 16000, 24000} else 8000
|
||||
|
||||
|
||||
def _http_host() -> str:
|
||||
return str(os.getenv("REALTIME_VOICE_HTTP_HOST", "0.0.0.0") or "0.0.0.0").strip()
|
||||
|
||||
@@ -77,38 +82,119 @@ def _http_port() -> int:
|
||||
return max(_int_env("REALTIME_VOICE_HTTP_PORT", 8000), 1)
|
||||
|
||||
|
||||
def _initial_greeting_text() -> str:
|
||||
raw = os.getenv("REALTIME_VOICE_INITIAL_GREETING_TEXT")
|
||||
if raw is None:
|
||||
return "Здравствуйте! Чем могу помочь?"
|
||||
return str(raw).strip()
|
||||
|
||||
def _llm_provider_name() -> str:
|
||||
return str(os.getenv("LLM_PROVIDER") or os.getenv("REALTIME_VOICE_LLM_PROVIDER") or "openai").strip().lower()
|
||||
|
||||
|
||||
def _llm_model_name() -> str:
|
||||
provider = _llm_provider_name()
|
||||
if provider in {"ollama", "local", "qwen"}:
|
||||
return str(os.getenv("OLLAMA_LLM_MODEL", "qwen2.5:1.5b")).strip() or "qwen2.5:1.5b"
|
||||
return str(os.getenv("OPENAI_LLM_MODEL", "gpt-4o-mini")).strip() or "gpt-4o-mini"
|
||||
|
||||
|
||||
|
||||
class RealtimeVoiceService:
|
||||
def __init__(self) -> None:
|
||||
self._stt = ElevenLabsSTT()
|
||||
self._llm = OpenAILLM()
|
||||
self._tts = ElevenLabsTTS()
|
||||
self._sample_rate_hz = _sample_rate_hz()
|
||||
self._stt = create_stt_provider(
|
||||
input_sample_rate_hz=self._sample_rate_hz,
|
||||
target_sample_rate_hz=self._sample_rate_hz,
|
||||
)
|
||||
self._llm = create_llm_provider()
|
||||
self._tts = create_tts_provider(
|
||||
output_format=f"pcm_{self._sample_rate_hz}",
|
||||
target_sample_rate_hz=self._sample_rate_hz,
|
||||
)
|
||||
self._filler_audio = FillerAudioLibrary(sample_rate_hz=self._sample_rate_hz)
|
||||
self._audiosocket_server = AudioSocketServer(
|
||||
host=_audiosocket_host(),
|
||||
port=_audiosocket_port(),
|
||||
sample_rate_hz=self._sample_rate_hz,
|
||||
session_handler=self._run_transport_session,
|
||||
)
|
||||
self._active_sessions: dict[str, CallSession] = {}
|
||||
self._session_lock = asyncio.Lock()
|
||||
LOGGER.info(
|
||||
"realtime voice service config: sample_rate=%s audiosocket=%s:%s http=%s:%s "
|
||||
"llm_provider=%s llm_model=%s tools_enabled=%s serper_configured=%s stt_provider=%s stt_model=%s stt_realtime_model=%s "
|
||||
"tts_voice_id=%s tts_model=%s tts_format=%s tts_speed=%s",
|
||||
self._sample_rate_hz,
|
||||
_audiosocket_host(),
|
||||
_audiosocket_port(),
|
||||
_http_host(),
|
||||
_http_port(),
|
||||
_llm_provider_name(),
|
||||
_llm_model_name(),
|
||||
os.getenv("OPENAI_LLM_ENABLE_TOOLS", "true"),
|
||||
bool(os.getenv("SERPER_API_KEY")),
|
||||
os.getenv("STT_PROVIDER") or os.getenv("REALTIME_VOICE_STT_PROVIDER") or os.getenv("AI_VOICE_ASR_PROVIDER", "elevenlabs"),
|
||||
os.getenv("ELEVENLABS_STT_MODEL_ID", "scribe_v2"),
|
||||
os.getenv("ELEVENLABS_STT_REALTIME_MODEL_ID", "scribe_v2_realtime"),
|
||||
os.getenv("ELEVENLABS_TTS_VOICE_ID", ""),
|
||||
os.getenv("ELEVENLABS_TTS_MODEL_ID", "eleven_turbo_v2_5"),
|
||||
f"pcm_{self._sample_rate_hz}",
|
||||
os.getenv("ELEVENLABS_TTS_SPEED", "1.0"),
|
||||
)
|
||||
|
||||
@property
|
||||
def active_session_count(self) -> int:
|
||||
return len(self._active_sessions)
|
||||
|
||||
async def start(self) -> None:
|
||||
LOGGER.info("realtime voice service starting")
|
||||
await self._audiosocket_server.start()
|
||||
try:
|
||||
LOGGER.info("preloading filler audio library")
|
||||
await self._filler_audio.preload(tts=self._tts)
|
||||
await self._preload_initial_greeting_audio()
|
||||
LOGGER.info("filler audio library preload completed")
|
||||
except Exception:
|
||||
LOGGER.exception("failed to preload filler audio clips")
|
||||
|
||||
async def _preload_initial_greeting_audio(self) -> None:
|
||||
greeting_text = _initial_greeting_text()
|
||||
if not greeting_text:
|
||||
LOGGER.info("initial greeting audio preload skipped: empty text")
|
||||
return
|
||||
try:
|
||||
LOGGER.info(
|
||||
"initial greeting audio preload start: chars=%s text=%r",
|
||||
len(greeting_text),
|
||||
greeting_text[:120],
|
||||
)
|
||||
clip = await self._filler_audio.synthesize_text(self._tts, greeting_text)
|
||||
if not clip:
|
||||
LOGGER.warning("initial greeting audio preload produced empty clip")
|
||||
return
|
||||
self._filler_audio.add_clip("initial_greeting", clip)
|
||||
LOGGER.info("initial greeting audio preload done: bytes=%s", len(clip))
|
||||
except Exception:
|
||||
LOGGER.exception("failed to preload initial greeting audio")
|
||||
|
||||
async def stop(self) -> None:
|
||||
LOGGER.info("realtime voice service stopping active_sessions=%s", len(self._active_sessions))
|
||||
await self._audiosocket_server.stop()
|
||||
sessions = list(self._active_sessions.values())
|
||||
for session in sessions:
|
||||
await session.stop()
|
||||
self._active_sessions.clear()
|
||||
await self._close_provider(self._stt)
|
||||
await self._close_provider(self._llm)
|
||||
await self._close_provider(self._tts)
|
||||
|
||||
async def handle_websocket(self, websocket: WebSocket, *, client_id: str | None = None) -> None:
|
||||
await websocket.accept()
|
||||
transport = WebSocketMediaTransport(
|
||||
websocket=websocket,
|
||||
transport_id=client_id or str(uuid.uuid4()),
|
||||
sample_rate_hz=self._sample_rate_hz,
|
||||
)
|
||||
await self._run_transport_session(transport)
|
||||
|
||||
@@ -116,9 +202,12 @@ class RealtimeVoiceService:
|
||||
session = self._build_session(transport)
|
||||
async with self._track_session(session):
|
||||
LOGGER.info(
|
||||
"starting realtime session %s via %s",
|
||||
"starting realtime session %s via %s sample_rate=%s frame_ms=%s frame_bytes=%s",
|
||||
session.session_id,
|
||||
transport.protocol,
|
||||
transport.sample_rate_hz,
|
||||
transport.frame_duration_ms,
|
||||
transport.frame_bytes,
|
||||
)
|
||||
await session.run()
|
||||
|
||||
@@ -130,7 +219,7 @@ class RealtimeVoiceService:
|
||||
sample_rate_hz=transport.sample_rate_hz,
|
||||
threshold=_float_env("VAD_THRESHOLD", 0.5),
|
||||
negative_threshold=_optional_float_env("VAD_NEGATIVE_THRESHOLD"),
|
||||
speech_end_silence_ms=_int_env("VAD_SILENCE_TIMEOUT_MS", 1600),
|
||||
speech_end_silence_ms=_int_env("VAD_SILENCE_TIMEOUT_MS", 550),
|
||||
speech_pad_ms=_int_env("VAD_SPEECH_PAD_MS", 64),
|
||||
min_speech_duration_ms=_int_env("VAD_MIN_SPEECH_DURATION_MS", 0),
|
||||
use_onnx=_bool_env("VAD_USE_ONNX", False),
|
||||
@@ -138,17 +227,38 @@ class RealtimeVoiceService:
|
||||
stt=self._stt,
|
||||
llm=self._llm,
|
||||
tts=self._tts,
|
||||
filler_audio=self._filler_audio,
|
||||
initial_greeting_text=_initial_greeting_text(),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _close_provider(provider: object) -> None:
|
||||
close = getattr(provider, "close", None)
|
||||
if close is None:
|
||||
return
|
||||
result = close()
|
||||
if inspect.isawaitable(result):
|
||||
await result
|
||||
|
||||
@asynccontextmanager
|
||||
async def _track_session(self, session: CallSession):
|
||||
async with self._session_lock:
|
||||
self._active_sessions[session.session_id] = session
|
||||
LOGGER.info(
|
||||
"session tracked: session=%s active_sessions=%s",
|
||||
session.session_id,
|
||||
len(self._active_sessions),
|
||||
)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
async with self._session_lock:
|
||||
self._active_sessions.pop(session.session_id, None)
|
||||
LOGGER.info(
|
||||
"session untracked: session=%s active_sessions=%s",
|
||||
session.session_id,
|
||||
len(self._active_sessions),
|
||||
)
|
||||
|
||||
|
||||
service = RealtimeVoiceService()
|
||||
|
||||
Reference in New Issue
Block a user