Mercurial
view dictation/webrtc_test.py @ 280:49e9e591c9bb
Add persistent dictation, prewarmed WebRTC speech input, Copilot SDK routing, animated conversation lifecycle controls, parking, and architecture coverage.
| author | MrJuneJune <me@mrjunejune.com> |
|---|---|
| date | Tue, 18 Aug 2026 19:14:53 -0700 |
| parents | 78699f810817 |
| children |
line wrap: on
line source
import asyncio import json import math from pathlib import Path import tempfile import unittest import wave from aiortc import RTCPeerConnection, RTCSessionDescription from aiortc.contrib.media import MediaPlayer import numpy as np from dictation.config import DictationConfig from dictation.server import DictationService, Offer from dictation.transcriber import Transcript class FakeTranscriber: async def warmup(self): return None async def transcribe(self, samples, *, final): return Transcript( "final transcript" if final else "partial transcript", "en", 0.99, ) def write_test_audio(path: str) -> None: sample_rate = 16000 speech = np.array( [ int(1600 * math.sin(2 * math.pi * 440 * index / sample_rate)) for index in range(sample_rate * 2) ], dtype=np.int16, ) silence = np.zeros(sample_rate, dtype=np.int16) with wave.open(path, "wb") as output: output.setnchannels(1) output.setsampwidth(2) output.setframerate(sample_rate) output.writeframes(np.concatenate((speech, silence)).tobytes()) def test_config() -> DictationConfig: return DictationConfig( host="127.0.0.1", port=8090, model_dir=Path("/tmp/model"), compute_type="int8_float16", max_sessions=1, partial_interval_ms=500, silence_ms=200, max_utterance_seconds=5, speech_threshold=0.01, ) class WebRtcTest(unittest.IsolatedAsyncioTestCase): async def test_audio_track_returns_transcript_events(self): loop_errors = [] loop = asyncio.get_running_loop() previous_exception_handler = loop.get_exception_handler() loop.set_exception_handler( lambda _loop, context: loop_errors.append(context) ) self.addCleanup( loop.set_exception_handler, previous_exception_handler, ) service = DictationService(test_config(), FakeTranscriber()) await service.start() client = RTCPeerConnection() channel = client.createDataChannel("transcripts") messages = [] final_received = asyncio.Event() audio_file = tempfile.NamedTemporaryFile(suffix=".wav", delete=False) audio_file.close() write_test_audio(audio_file.name) player = MediaPlayer(audio_file.name) @channel.on("message") def on_message(message): event = json.loads(message) messages.append(event) if event["type"] == "transcript.final": final_received.set() client.addTrack(player.audio) offer = await client.createOffer() await client.setLocalDescription(offer) answer = await service.accept_offer( Offer( sdp=client.localDescription.sdp, type=client.localDescription.type, ) ) await client.setRemoteDescription( RTCSessionDescription( sdp=answer["sdp"], type=answer["type"], ) ) try: try: await asyncio.wait_for(final_received.wait(), timeout=10) except TimeoutError as error: server_states = [ { "connection": session.peer.connectionState, "ice": session.peer.iceConnectionState, "channel": ( session.channel.readyState if session.channel is not None else None ), "audioTask": ( "missing" if session.audio_task is None else ( repr(session.audio_task.exception()) if session.audio_task.done() and not session.audio_task.cancelled() else "running" ) ), } for session in service.sessions.values() ] self.fail( "Timed out waiting for transcript: " f"client={client.connectionState}/" f"{client.iceConnectionState}, " f"server={server_states}, messages={messages}" ) event_types = [message["type"] for message in messages] self.assertIn("ready", event_types) self.assertIn("speech.started", event_types) self.assertIn("transcript.partial", event_types) self.assertIn("transcript.final", event_types) channel.send("commit") await asyncio.sleep(0.2) self.assertNotEqual(client.connectionState, "closed") self.assertIn(answer["sessionId"], service.sessions) finally: await client.close() await service.close() await asyncio.sleep(0.6) Path(audio_file.name).unlink(missing_ok=True) self.assertEqual(service.sessions, {}) self.assertEqual(loop_errors, []) if __name__ == "__main__": unittest.main()