410 lines
15 KiB
Python
410 lines
15 KiB
Python
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)
|