From d617b6908c9bd65cefd329f09192bab95afd8a25 Mon Sep 17 00:00:00 2001 From: Magzhan Zhumabayev Date: Mon, 4 May 2026 23:00:34 +0500 Subject: [PATCH] . --- tests/test_realtime_voice_service.py | 158 +++++++++++++++++++++++++++ 1 file changed, 158 insertions(+) create mode 100644 tests/test_realtime_voice_service.py diff --git a/tests/test_realtime_voice_service.py b/tests/test_realtime_voice_service.py new file mode 100644 index 0000000..6e3633f --- /dev/null +++ b/tests/test_realtime_voice_service.py @@ -0,0 +1,158 @@ +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()]