from __future__ import annotations import asyncio import audioop import logging import os import struct import time import uuid from typing import Awaitable, Callable from realtime_voice_service.transports.base import BaseMediaTransport LOGGER = logging.getLogger("uvicorn.error") AUDIO_SOCKET_PACKET_HANGUP = 0x00 AUDIO_SOCKET_PACKET_UUID = 0x01 AUDIO_SOCKET_PACKET_DTMF = 0x03 _AUDIO_SOCKET_PCM16_PACKET_TYPES_BY_SAMPLE_RATE_HZ: dict[int, int] = { 8000: 0x10, 12000: 0x11, 16000: 0x12, 24000: 0x13, 32000: 0x14, 44100: 0x15, 48000: 0x16, 96000: 0x17, 192000: 0x18, } _AUDIO_SOCKET_SAMPLE_RATE_HZ_BY_PCM16_PACKET_TYPE: dict[int, int] = { packet_type: sample_rate_hz for sample_rate_hz, packet_type in _AUDIO_SOCKET_PCM16_PACKET_TYPES_BY_SAMPLE_RATE_HZ.items() } SUPPORTED_AUDIO_SOCKET_SAMPLE_RATES_HZ = frozenset(_AUDIO_SOCKET_PCM16_PACKET_TYPES_BY_SAMPLE_RATE_HZ) def _float_env(name: str, default: float) -> float: raw = os.getenv(name) if raw is None: return default try: return float(str(raw).strip()) except ValueError: return default def normalize_session_id(value: str | bytes) -> str: if isinstance(value, bytes): try: return str(uuid.UUID(bytes=value)).lower() except (ValueError, AttributeError): return str(uuid.UUID(value.decode("utf-8").strip())).lower() return str(uuid.UUID(str(value).strip())).lower() def encode_packet(packet_type: int, payload: bytes = b"") -> bytes: return struct.pack("!BH", packet_type & 0xFF, len(payload)) + payload def pcm16_packet_type_for_sample_rate(sample_rate_hz: int) -> int: try: return _AUDIO_SOCKET_PCM16_PACKET_TYPES_BY_SAMPLE_RATE_HZ[int(sample_rate_hz)] except KeyError as exc: supported = ", ".join(str(rate) for rate in sorted(SUPPORTED_AUDIO_SOCKET_SAMPLE_RATES_HZ)) raise ValueError(f"unsupported AudioSocket PCM16 sample rate {sample_rate_hz}; supported: {supported}") from exc def sample_rate_for_pcm16_packet_type(packet_type: int) -> int | None: return _AUDIO_SOCKET_SAMPLE_RATE_HZ_BY_PCM16_PACKET_TYPE.get(packet_type) def encode_audio_packet(pcm_bytes: bytes, *, sample_rate_hz: int = 16000) -> bytes: return encode_packet(pcm16_packet_type_for_sample_rate(sample_rate_hz), pcm_bytes) async def read_packet( reader: asyncio.StreamReader, *, timeout_seconds: float = 5.0, ) -> tuple[int, bytes]: header = await asyncio.wait_for(reader.readexactly(3), timeout=timeout_seconds) packet_type, payload_length = struct.unpack("!BH", header) payload = b"" if payload_length: payload = await asyncio.wait_for(reader.readexactly(payload_length), timeout=timeout_seconds) return packet_type, payload class AudioSocketTransport(BaseMediaTransport): def __init__( self, *, transport_id: str, reader: asyncio.StreamReader, writer: asyncio.StreamWriter, sample_rate_hz: int = 16000, frame_duration_ms: int = 20, read_timeout_seconds: float = 30.0, ) -> None: super().__init__( transport_id=transport_id, sample_rate_hz=sample_rate_hz, frame_duration_ms=frame_duration_ms, ) self._reader = reader self._writer = writer self._pcm_packet_type = pcm16_packet_type_for_sample_rate(self.sample_rate_hz) self._read_timeout_seconds = max(read_timeout_seconds, 1.0) self._closed = False self._rx_audio_packet_count = 0 self._rx_audio_bytes = 0 self._rx_ignored_packet_count = 0 self._rx_rms_sum = 0 self._rx_rms_count = 0 self._rx_peak_abs = 0 self._rx_low_level_packet_count = 0 self._rx_silence_rms_threshold = max(int(os.getenv("REALTIME_VOICE_RX_SILENCE_RMS_THRESHOLD", "80")), 0) self._rx_log_interval_seconds = max( _float_env("REALTIME_VOICE_AUDIO_LOG_INTERVAL_SECONDS", 1.0), 0.1, ) self._last_rx_summary_monotonic = time.perf_counter() @property def protocol(self) -> str: return "audiosocket" @classmethod async def from_streams( cls, reader: asyncio.StreamReader, writer: asyncio.StreamWriter, *, handshake_timeout_seconds: float = 5.0, sample_rate_hz: int = 16000, frame_duration_ms: int = 20, ) -> AudioSocketTransport: packet_type, payload = await read_packet(reader, timeout_seconds=handshake_timeout_seconds) if packet_type != AUDIO_SOCKET_PACKET_UUID: raise ValueError(f"expected UUID packet, got {packet_type}") session_id = normalize_session_id(payload) return cls( transport_id=session_id, reader=reader, writer=writer, sample_rate_hz=sample_rate_hz, frame_duration_ms=frame_duration_ms, ) async def receive_audio(self) -> bytes | None: if self._closed: return None while not self._closed: try: packet_type, payload = await read_packet( self._reader, timeout_seconds=self._read_timeout_seconds, ) except asyncio.TimeoutError: LOGGER.info( "AudioSocket receive timeout: session=%s timeout_seconds=%s rx_packets=%s rx_bytes=%s", self.transport_id, self._read_timeout_seconds, self._rx_audio_packet_count, self._rx_audio_bytes, ) return None except asyncio.IncompleteReadError as exc: LOGGER.info( "AudioSocket peer closed stream: session=%s partial_bytes=%s expected_bytes=%s " "rx_packets=%s rx_bytes=%s", self.transport_id, len(exc.partial or b""), exc.expected, self._rx_audio_packet_count, self._rx_audio_bytes, ) return None except ConnectionError: LOGGER.info( "AudioSocket connection error on receive: session=%s rx_packets=%s rx_bytes=%s", self.transport_id, self._rx_audio_packet_count, self._rx_audio_bytes, ) return None if packet_type == AUDIO_SOCKET_PACKET_HANGUP: LOGGER.info( "AudioSocket hangup packet: session=%s rx_packets=%s rx_bytes=%s", self.transport_id, self._rx_audio_packet_count, self._rx_audio_bytes, ) return None packet_sample_rate_hz = sample_rate_for_pcm16_packet_type(packet_type) if packet_sample_rate_hz is not None: if packet_sample_rate_hz != self.sample_rate_hz: LOGGER.error( "AudioSocket sample rate mismatch: session=%s expected_sample_rate=%s expected_packet_type=0x%02x " "received_sample_rate=%s received_packet_type=0x%02x payload_bytes=%s", self.transport_id, self.sample_rate_hz, self._pcm_packet_type, packet_sample_rate_hz, packet_type, len(payload), ) return None self._rx_audio_packet_count += 1 self._rx_audio_bytes += len(payload) self._track_rx_audio_level(payload) self._maybe_log_rx_summary() return payload if packet_type == AUDIO_SOCKET_PACKET_DTMF: self._rx_ignored_packet_count += 1 LOGGER.info( "AudioSocket DTMF packet ignored: session=%s payload=%r", self.transport_id, payload[:16], ) continue if packet_type == AUDIO_SOCKET_PACKET_UUID: self._rx_ignored_packet_count += 1 LOGGER.info("AudioSocket duplicate UUID packet ignored: session=%s", self.transport_id) continue self._rx_ignored_packet_count += 1 LOGGER.warning( "AudioSocket unknown packet ignored: session=%s packet_type=0x%02x payload_bytes=%s", self.transport_id, packet_type, len(payload), ) return None async def _send_frame(self, frame: bytes) -> None: if self._closed or self._writer.is_closing(): return self._writer.write(encode_audio_packet(frame, sample_rate_hz=self.sample_rate_hz)) await self._writer.drain() async def close(self) -> None: if self._closed: return self._closed = True LOGGER.info( "AudioSocket transport closing: session=%s rx_packets=%s rx_bytes=%s ignored_packets=%s " "avg_rms=%s peak_abs=%s low_level_packets=%s", self.transport_id, self._rx_audio_packet_count, self._rx_audio_bytes, self._rx_ignored_packet_count, self._average_rx_rms(), self._rx_peak_abs, self._rx_low_level_packet_count, ) if not self._writer.is_closing(): self._writer.close() await self._writer.wait_closed() def _maybe_log_rx_summary(self, *, force: bool = False) -> None: now = time.perf_counter() if not force and (now - self._last_rx_summary_monotonic) < self._rx_log_interval_seconds: return self._last_rx_summary_monotonic = now LOGGER.info( "AudioSocket audio_rx_summary: session=%s sample_rate=%s frame_ms=%s frame_bytes=%s " "rx_packets=%s rx_bytes=%s ignored_packets=%s avg_rms=%s peak_abs=%s " "low_level_packets=%s low_level_ratio=%.3f", self.transport_id, self.sample_rate_hz, self.frame_duration_ms, self.frame_bytes, self._rx_audio_packet_count, self._rx_audio_bytes, self._rx_ignored_packet_count, self._average_rx_rms(), self._rx_peak_abs, self._rx_low_level_packet_count, self._rx_low_level_ratio(), ) def _track_rx_audio_level(self, payload: bytes) -> None: if not payload: return try: rms = int(audioop.rms(payload, 2)) peak = int(audioop.max(payload, 2)) except Exception: return self._rx_rms_sum += rms self._rx_rms_count += 1 self._rx_peak_abs = max(self._rx_peak_abs, peak) if rms <= self._rx_silence_rms_threshold: self._rx_low_level_packet_count += 1 def _average_rx_rms(self) -> int: if self._rx_rms_count <= 0: return 0 return int(self._rx_rms_sum / self._rx_rms_count) def _rx_low_level_ratio(self) -> float: if self._rx_rms_count <= 0: return 0.0 return self._rx_low_level_packet_count / float(self._rx_rms_count) class AudioSocketServer: def __init__( self, *, host: str, port: int, session_handler: Callable[[AudioSocketTransport], Awaitable[None]], handshake_timeout_seconds: float = 5.0, sample_rate_hz: int = 16000, frame_duration_ms: int = 20, ) -> None: self._host = host self._port = port self._session_handler = session_handler self._handshake_timeout_seconds = max(handshake_timeout_seconds, 1.0) self._sample_rate_hz = sample_rate_hz self._frame_duration_ms = frame_duration_ms self._server: asyncio.base_events.Server | None = None self._connection_tasks: set[asyncio.Task[None]] = set() @property def bound_port(self) -> int: if self._server is None or not self._server.sockets: return self._port return int(self._server.sockets[0].getsockname()[1]) async def start(self) -> None: if self._server is not None: return self._server = await asyncio.start_server( self._handle_connection, self._host, self._port, ) LOGGER.info( "AudioSocket server listening on %s:%s sample_rate=%s frame_ms=%s frame_bytes=%s", self._host, self.bound_port, self._sample_rate_hz, self._frame_duration_ms, int((self._sample_rate_hz * self._frame_duration_ms / 1000.0) * 2), ) async def stop(self) -> None: if self._server is None: return self._server.close() await self._server.wait_closed() self._server = None tasks = list(self._connection_tasks) for task in tasks: task.cancel() for task in tasks: try: await task except asyncio.CancelledError: pass async def _handle_connection( self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter, ) -> None: current_task = asyncio.current_task() if current_task is not None: self._connection_tasks.add(current_task) transport: AudioSocketTransport | None = None peer = writer.get_extra_info("peername") try: transport = await AudioSocketTransport.from_streams( reader, writer, handshake_timeout_seconds=self._handshake_timeout_seconds, sample_rate_hz=self._sample_rate_hz, frame_duration_ms=self._frame_duration_ms, ) LOGGER.info( "AudioSocket client accepted: session=%s peer=%s sample_rate=%s frame_ms=%s frame_bytes=%s", transport.transport_id, peer, transport.sample_rate_hz, transport.frame_duration_ms, transport.frame_bytes, ) await self._session_handler(transport) except asyncio.CancelledError: raise except Exception: LOGGER.exception("AudioSocket connection failed from peer=%s", peer) finally: if transport is not None: transport._maybe_log_rx_summary(force=True) await transport.close() elif not writer.is_closing(): writer.close() await writer.wait_closed() LOGGER.info("AudioSocket connection finished: peer=%s session=%s", peer, transport.transport_id if transport else None) if current_task is not None: self._connection_tasks.discard(current_task)