.
This commit is contained in:
@@ -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()]
|
||||
Reference in New Issue
Block a user