comparison dictation/webrtc_test.py @ 277:1d99147f520c

Merge Qwen services and infinite canvas heads
author MrJuneJune <me@mrjunejune.com>
date Mon, 17 Aug 2026 17:01:40 -0700
parents 78699f810817
children 49e9e591c9bb
comparison
equal deleted inserted replaced
276:b55c22cff335 277:1d99147f520c
1 import asyncio
2 import json
3 import math
4 from pathlib import Path
5 import tempfile
6 import unittest
7 import wave
8
9 from aiortc import RTCPeerConnection, RTCSessionDescription
10 from aiortc.contrib.media import MediaPlayer
11 import numpy as np
12
13 from dictation.config import DictationConfig
14 from dictation.server import DictationService, Offer
15 from dictation.transcriber import Transcript
16
17
18 class FakeTranscriber:
19 async def warmup(self):
20 return None
21
22 async def transcribe(self, samples, *, final):
23 return Transcript(
24 "final transcript" if final else "partial transcript",
25 "en",
26 0.99,
27 )
28
29
30 def write_test_audio(path: str) -> None:
31 sample_rate = 16000
32 speech = np.array(
33 [
34 int(1600 * math.sin(2 * math.pi * 440 * index / sample_rate))
35 for index in range(sample_rate * 2)
36 ],
37 dtype=np.int16,
38 )
39 silence = np.zeros(sample_rate, dtype=np.int16)
40 with wave.open(path, "wb") as output:
41 output.setnchannels(1)
42 output.setsampwidth(2)
43 output.setframerate(sample_rate)
44 output.writeframes(np.concatenate((speech, silence)).tobytes())
45
46
47 def test_config() -> DictationConfig:
48 return DictationConfig(
49 host="127.0.0.1",
50 port=8090,
51 model_dir=Path("/tmp/model"),
52 compute_type="int8_float16",
53 max_sessions=1,
54 partial_interval_ms=500,
55 silence_ms=200,
56 max_utterance_seconds=5,
57 speech_threshold=0.01,
58 )
59
60
61 class WebRtcTest(unittest.IsolatedAsyncioTestCase):
62 async def test_audio_track_returns_transcript_events(self):
63 service = DictationService(test_config(), FakeTranscriber())
64 await service.start()
65 client = RTCPeerConnection()
66 channel = client.createDataChannel("transcripts")
67 messages = []
68 final_received = asyncio.Event()
69 audio_file = tempfile.NamedTemporaryFile(suffix=".wav", delete=False)
70 audio_file.close()
71 write_test_audio(audio_file.name)
72 player = MediaPlayer(audio_file.name)
73
74 @channel.on("message")
75 def on_message(message):
76 event = json.loads(message)
77 messages.append(event)
78 if event["type"] == "transcript.final":
79 final_received.set()
80
81 client.addTrack(player.audio)
82 offer = await client.createOffer()
83 await client.setLocalDescription(offer)
84 answer = await service.accept_offer(
85 Offer(
86 sdp=client.localDescription.sdp,
87 type=client.localDescription.type,
88 )
89 )
90 await client.setRemoteDescription(
91 RTCSessionDescription(
92 sdp=answer["sdp"],
93 type=answer["type"],
94 )
95 )
96
97 try:
98 try:
99 await asyncio.wait_for(final_received.wait(), timeout=10)
100 except TimeoutError as error:
101 server_states = [
102 {
103 "connection": session.peer.connectionState,
104 "ice": session.peer.iceConnectionState,
105 "channel": (
106 session.channel.readyState
107 if session.channel is not None
108 else None
109 ),
110 "audioTask": (
111 "missing"
112 if session.audio_task is None
113 else (
114 repr(session.audio_task.exception())
115 if session.audio_task.done()
116 and not session.audio_task.cancelled()
117 else "running"
118 )
119 ),
120 }
121 for session in service.sessions.values()
122 ]
123 self.fail(
124 "Timed out waiting for transcript: "
125 f"client={client.connectionState}/"
126 f"{client.iceConnectionState}, "
127 f"server={server_states}, messages={messages}"
128 )
129 event_types = [message["type"] for message in messages]
130 self.assertIn("ready", event_types)
131 self.assertIn("speech.started", event_types)
132 self.assertIn("transcript.partial", event_types)
133 self.assertIn("transcript.final", event_types)
134 finally:
135 await client.close()
136 await service.close()
137 Path(audio_file.name).unlink(missing_ok=True)
138 self.assertEqual(service.sessions, {})
139
140
141 if __name__ == "__main__":
142 unittest.main()