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()