Files
call-center/tests/test_realtime_voice_service.py
T
2026-05-04 23:00:34 +05:00

159 lines
5.0 KiB
Python

from __future__ import annotations
import asyncio
import audioop
import uuid
from realtime_voice_service.core.session import CallSession
from realtime_voice_service.core.vad import SileroVADDetector
from realtime_voice_service.providers.base import BaseLLM, BaseTTS, MockSTT
from realtime_voice_service.transports.audiosocket import (
AUDIO_SOCKET_PACKET_PCM16,
AUDIO_SOCKET_PACKET_UUID,
AudioSocketServer,
encode_audio_packet,
encode_packet,
read_packet,
)
from realtime_voice_service.transports.base import BaseMediaTransport
def _speech_frame() -> bytes:
return (1200).to_bytes(2, "little", signed=True) * 160
def _silence_frame() -> bytes:
return b"\x00\x00" * 160
class _FakeTransport(BaseMediaTransport):
def __init__(self, schedule: list[tuple[float, bytes | None]]) -> None:
super().__init__(transport_id="test-session", sample_rate_hz=8000, frame_duration_ms=20)
self._schedule = list(schedule)
self.sent_audio: list[bytes] = []
self.closed = False
@property
def protocol(self) -> str:
return "fake"
async def receive_audio(self) -> bytes | None:
if not self._schedule:
return None
delay_seconds, payload = self._schedule.pop(0)
if delay_seconds:
await asyncio.sleep(delay_seconds)
return payload
async def send_audio(self, audio_chunk: bytes) -> None:
self.sent_audio.append(audio_chunk)
await asyncio.sleep(0.02)
async def close(self) -> None:
self.closed = True
class _StreamingEchoLLM(BaseLLM):
async def generate_stream(self, transcript: str, context: list):
del context
await asyncio.sleep(0.02)
yield f"reply to {transcript}"
class _LongMockTTS(BaseTTS):
async def synthesize_stream(self, text: str):
del text
await asyncio.sleep(0.01)
frame_count = 14
for _ in range(frame_count):
await asyncio.sleep(0.01)
yield b"\x10\x00" * 160
def _window_prediction(window: bytes) -> float:
return 0.9 if audioop.rms(window, 2) >= 250 else 0.0
def test_call_session_barge_in_increments_generation_epoch():
transport = _FakeTransport(
[
(0.0, _speech_frame()),
(0.0, _speech_frame()),
(0.0, _silence_frame()),
(0.0, _silence_frame()),
(0.0, _silence_frame()),
(0.11, _speech_frame()),
(0.0, _speech_frame()),
(0.0, _silence_frame()),
(0.0, _silence_frame()),
(0.0, _silence_frame()),
(0.8, None),
]
)
session = CallSession(
session_id="test-session",
transport=transport,
vad=SileroVADDetector(
sample_rate_hz=transport.sample_rate_hz,
speech_end_silence_ms=32,
speech_pad_ms=0,
prediction_fn=_window_prediction,
),
stt=MockSTT(
latency_ms=5,
sample_rate_hz=transport.sample_rate_hz,
scripted_transcripts=["first request", "second request"],
),
llm=_StreamingEchoLLM(),
tts=_LongMockTTS(),
)
asyncio.run(session.run())
assert transport.closed is True
assert session.generation_epoch == 1
assert session.interruptions == ["barge-in"]
assert ("user", "first request") in session.conversation
assert ("user", "second request") in session.conversation
assert ("assistant", "reply to first request") not in session.conversation
assert ("assistant", "reply to second request") in session.conversation
assert transport.sent_audio
assert set(session.last_latency_ms) >= {"stt_latency", "ttft", "ttfa"}
def test_audiosocket_server_decodes_uuid_and_echoes_pcm():
received_session_ids: list[str] = []
received_payloads: list[bytes] = []
async def _scenario() -> None:
async def _handler(transport) -> None:
received_session_ids.append(transport.transport_id)
pcm = await transport.receive_audio()
received_payloads.append(pcm or b"")
if pcm is not None:
await transport.send_audio(pcm)
server = AudioSocketServer(host="127.0.0.1", port=0, session_handler=_handler)
await server.start()
media_uuid = str(uuid.uuid4())
reader, writer = await asyncio.open_connection("127.0.0.1", server.bound_port)
writer.write(encode_packet(AUDIO_SOCKET_PACKET_UUID, uuid.UUID(media_uuid).bytes))
writer.write(encode_audio_packet(_speech_frame()))
await writer.drain()
packet_type, payload = await read_packet(reader, timeout_seconds=1.0)
assert packet_type == AUDIO_SOCKET_PACKET_PCM16
assert payload == _speech_frame()
writer.close()
await writer.wait_closed()
await asyncio.sleep(0.05)
await server.stop()
asyncio.run(_scenario())
assert len(received_session_ids) == 1
assert received_session_ids[0]
assert received_payloads == [_speech_frame()]