Files
realtime_voice_service/main.py
T
2026-05-09 09:25:29 +05:00

468 lines
17 KiB
Python

from __future__ import annotations
import asyncio
import inspect
import logging
import os
import uuid
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from fastapi import FastAPI, WebSocket
from realtime_voice_service import crm_client
from realtime_voice_service import sales_sync_client
from realtime_voice_service.core.filler_audio import FillerAudioLibrary
from realtime_voice_service.core.session import CallSession
from realtime_voice_service.core.session import _normalize_voice_pronunciation
from realtime_voice_service.core.vad import SileroVADDetector
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
LOGGER = logging.getLogger("uvicorn.error")
def _int_env(name: str, default: int) -> int:
raw = os.getenv(name)
if raw is None:
return default
try:
return int(str(raw).strip())
except ValueError:
return default
def _float_env(name: str, default: float) -> float:
raw = os.getenv(name)
if raw is None:
return default
try:
return float(str(raw).strip())
except ValueError:
return default
def _optional_float_env(name: str) -> float | None:
raw = os.getenv(name)
if raw is None:
return None
normalized = str(raw).strip()
if not normalized:
return None
try:
return float(normalized)
except ValueError:
return None
def _bool_env(name: str, default: bool = False) -> bool:
raw = os.getenv(name)
if raw is None:
return default
return str(raw).strip().lower() in {"1", "true", "yes", "on"}
def _audiosocket_host() -> str:
return str(os.getenv("REALTIME_VOICE_AUDIOSOCKET_HOST", "0.0.0.0") or "0.0.0.0").strip()
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", 16000), 1)
if configured != 16000:
LOGGER.warning("REALTIME_VOICE_SAMPLE_RATE_HZ=%s is not allowed; using 16000 Hz", configured)
return 16000
def _http_host() -> str:
return str(os.getenv("REALTIME_VOICE_HTTP_HOST", "0.0.0.0") or "0.0.0.0").strip()
def _http_port() -> int:
return max(_int_env("REALTIME_VOICE_HTTP_PORT", 8000), 1)
def _initial_greeting_text() -> str:
raw = str(os.getenv("REALTIME_VOICE_INITIAL_GREETING_TEXT", "")).strip()
if raw and _single_initial_greeting_enabled():
return str(raw).strip()
return " ".join(text for text, _ in _initial_greeting_segments())
def _single_initial_greeting_enabled() -> bool:
raw = str(os.getenv("REALTIME_VOICE_INITIAL_GREETING_USE_SINGLE_TEXT", "")).strip().lower()
return raw in {"1", "true", "yes", "on"}
def _env_text_or_default(name: str, default: str) -> str:
value = os.getenv(name)
if value is None or not str(value).strip():
return default
return str(value).strip()
def _initial_greeting_segments() -> tuple[tuple[str, str | None], ...]:
raw = str(os.getenv("REALTIME_VOICE_INITIAL_GREETING_TEXT", "")).strip()
if raw and _single_initial_greeting_enabled():
language = str(os.getenv("REALTIME_VOICE_INITIAL_GREETING_LANGUAGE_CODE", "")).strip() or None
return ((_normalize_voice_pronunciation(raw), language),)
ru_text = _normalize_voice_pronunciation(
_env_text_or_default(
"REALTIME_VOICE_INITIAL_GREETING_RU_TEXT",
(
"Здравствуйте, мое имя Айну́р. "
"На каком языке Вам удобнее общаться: на казахском или русском?"
),
)
)
kk_text = _normalize_voice_pronunciation(
_env_text_or_default(
"REALTIME_VOICE_INITIAL_GREETING_KK_TEXT",
(
"Сәлеметсіз бе, менің атым Айнұр. "
"Сізге қай тілде сөйлескен ыңғайлы: қазақша ма, орысша ма?"
),
)
)
segments: list[tuple[str, str | None]] = []
if ru_text:
segments.append((ru_text, "ru"))
if kk_text:
segments.append((kk_text, "kk"))
return tuple(segments)
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"
def _utc_now_iso() -> str:
return datetime.now(timezone.utc).replace(microsecond=0).isoformat().replace("+00:00", "Z")
def _conversation_transcript(session: CallSession) -> str | None:
lines = [f"{speaker}: {text.strip()}" for speaker, text in session.conversation if str(text).strip()]
if not lines:
return None
return "\n".join(lines)
def _conversation_summary(session: CallSession) -> str | None:
user_turns = [text.strip() for speaker, text in session.conversation if speaker == "user" and str(text).strip()]
assistant_turns = [text.strip() for speaker, text in session.conversation if speaker == "assistant" and str(text).strip()]
if not user_turns and not assistant_turns:
return None
parts: list[str] = []
if user_turns:
parts.append(f"Customer asked: {user_turns[0][:220]}")
if assistant_turns:
parts.append(f"Assistant response: {assistant_turns[-1][:220]}")
if session.customer_name:
parts.append(f"Customer name: {session.customer_name}")
return " | ".join(parts)[:700]
def _session_sync_metadata(
session: CallSession,
transport: BaseMediaTransport,
*,
failure_reason: str | None = None,
) -> dict[str, object]:
return {
"protocol": transport.protocol,
"sample_rate_hz": transport.sample_rate_hz,
"frame_duration_ms": transport.frame_duration_ms,
"frame_bytes": transport.frame_bytes,
"session_language": session.session_language,
"customer_name": session.customer_name,
"interruptions": list(session.interruptions),
"latency_ms": dict(session.last_latency_ms),
"conversation_entries": len(session.conversation),
"failure_reason": failure_reason,
}
class RealtimeVoiceService:
def __init__(self) -> None:
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_transport=%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_v3"),
os.getenv("ELEVENLABS_TTS_TRANSPORT", "auto"),
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)
@property
def sample_rate_hz(self) -> int:
return self._sample_rate_hz
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_segments = _initial_greeting_segments()
greeting_text = " ".join(text for text, _ in greeting_segments).strip()
if not greeting_segments or not greeting_text:
LOGGER.info("initial greeting audio preload skipped: empty text")
return
try:
LOGGER.info(
"initial greeting audio preload start: segments=%s chars=%s text=%r",
[(language, len(text)) for text, language in greeting_segments],
len(greeting_text),
greeting_text[:120],
)
total_bytes = 0
for index, (segment_text, language_code) in enumerate(greeting_segments, start=1):
segment_clip = await self._filler_audio.synthesize_segments(
self._tts,
((segment_text, language_code),),
)
if not segment_clip:
LOGGER.warning(
"initial greeting segment preload produced empty clip: index=%s language=%s",
index,
language_code,
)
continue
clip_key = f"initial_greeting_{index}_{language_code or 'default'}"
self._filler_audio.add_clip(clip_key, segment_clip)
total_bytes += len(segment_clip)
LOGGER.info(
"initial greeting segment preload done: key=%s language=%s bytes=%s",
clip_key,
language_code,
len(segment_clip),
)
if total_bytes <= 0:
LOGGER.warning("initial greeting audio preload produced no segment clips")
return
LOGGER.info("initial greeting audio preload done: bytes=%s", total_bytes)
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)
async def _run_transport_session(self, transport: BaseMediaTransport) -> None:
session = self._build_session(transport)
started_at = _utc_now_iso()
interaction_id = await crm_client.create_interaction(call_id=session.session_id)
await sales_sync_client.sync_voice_session(
call_id=session.session_id,
interaction_id=interaction_id,
voice_session_id=session.session_id,
ai_state=str(session.state.value).lower(),
telephony_status="connected",
call_status="started",
started_at=started_at,
summary="Realtime voice session started",
metadata=_session_sync_metadata(session, transport),
)
failure_reason: str | None = None
call_status = "completed"
async with self._track_session(session):
try:
LOGGER.info(
"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()
except Exception as exc:
call_status = "failed"
failure_reason = str(exc)[:300]
raise
finally:
await sales_sync_client.sync_voice_session(
call_id=session.session_id,
interaction_id=interaction_id,
caller_name=session.customer_name,
voice_session_id=session.session_id,
ai_state=str(session.state.value).lower(),
telephony_status="ended",
call_status=call_status,
started_at=started_at,
ended_at=_utc_now_iso(),
summary=_conversation_summary(session),
transcript_text=_conversation_transcript(session),
metadata=_session_sync_metadata(session, transport, failure_reason=failure_reason),
)
if interaction_id:
await crm_client.close_interaction(interaction_id)
def _build_session(self, transport: BaseMediaTransport) -> CallSession:
return CallSession(
session_id=transport.transport_id,
transport=transport,
vad=SileroVADDetector(
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", 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),
),
stt=self._stt,
llm=self._llm,
tts=self._tts,
filler_audio=self._filler_audio,
initial_greeting_text=_initial_greeting_text(),
initial_greeting_segments=_initial_greeting_segments(),
)
@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()
@asynccontextmanager
async def lifespan(_: FastAPI):
await service.start()
try:
yield
finally:
await service.stop()
app = FastAPI(title="realtime-voice-service", version="0.1.0", lifespan=lifespan)
@app.get("/health")
async def health() -> dict[str, object]:
return {
"status": "ok",
"service": "realtime-voice-service",
"sample_rate_hz": service.sample_rate_hz,
"active_sessions": service.active_session_count,
"audiosocket_port": _audiosocket_port(),
"sales_sync_enabled": os.getenv("SALES_SYNC_ENABLED", "0"),
"crm_interaction_enabled": os.getenv("CRM_INTERACTION_ENABLED", "0"),
}
@app.websocket("/ws/{client_id}")
async def websocket_media(websocket: WebSocket, client_id: str) -> None:
await service.handle_websocket(websocket, client_id=client_id)
if __name__ == "__main__":
import uvicorn
uvicorn.run(
"realtime_voice_service.main:app",
host=_http_host(),
port=_http_port(),
reload=False,
)