159 lines
5.0 KiB
Python
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()]
|