view dictation/webrtc_test.py @ 278:8d560f50ed4c

Improve infinite canvas interactions and browser chrome Render Lucide icons directly with Raylib, add searchable icon browsing, robust text editing, entity lifecycle animations, z-order-safe input, semantic themes, and animated editable browser controls. Document rendering, pinning, context, and component extension for future agents. Co-authored-by: Copilot <[email protected]> Copilot-Session: f68442b1-fa8f-46a0-9689-81710613bbd4
author MrJuneJune <me@mrjunejune.com>
date Mon, 17 Aug 2026 22:16:14 -0700
parents 78699f810817
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()