view dictation/webrtc_test.py @ 275:78699f810817

Add Qwen3-VL and WebRTC dictation services Add Bazel targets for the CUDA-backed Qwen3-VL server and a local WebRTC faster-whisper dictation service. Co-authored-by: Copilot <[email protected]> Copilot-Session: e3d8cb06-6c95-4ae0-9757-651d3796ab00
author MrJuneJune <me@mrjunejune.com>
date Mon, 17 Aug 2026 10:58:47 -0700
parents
children 49e9e591c9bb
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):
        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)
        finally:
            await client.close()
            await service.close()
            Path(audio_file.name).unlink(missing_ok=True)
        self.assertEqual(service.sessions, {})


if __name__ == "__main__":
    unittest.main()