Initial commit
This commit is contained in:
@@ -0,0 +1,28 @@
|
||||
__all__ = [
|
||||
"AudioSocketServer",
|
||||
"AudioSocketTransport",
|
||||
"BaseMediaTransport",
|
||||
"WebSocketMediaTransport",
|
||||
]
|
||||
|
||||
|
||||
def __getattr__(name: str):
|
||||
if name == "BaseMediaTransport":
|
||||
from realtime_voice_service.transports.base import BaseMediaTransport
|
||||
|
||||
return BaseMediaTransport
|
||||
if name in {"AudioSocketServer", "AudioSocketTransport"}:
|
||||
from realtime_voice_service.transports.audiosocket import (
|
||||
AudioSocketServer,
|
||||
AudioSocketTransport,
|
||||
)
|
||||
|
||||
return {
|
||||
"AudioSocketServer": AudioSocketServer,
|
||||
"AudioSocketTransport": AudioSocketTransport,
|
||||
}[name]
|
||||
if name == "WebSocketMediaTransport":
|
||||
from realtime_voice_service.transports.websocket import WebSocketMediaTransport
|
||||
|
||||
return WebSocketMediaTransport
|
||||
raise AttributeError(name)
|
||||
@@ -0,0 +1,224 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import struct
|
||||
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_PACKET_PCM16 = 0x10
|
||||
|
||||
|
||||
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 encode_audio_packet(pcm_bytes: bytes) -> bytes:
|
||||
return encode_packet(AUDIO_SOCKET_PACKET_PCM16, 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 = 8000,
|
||||
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._read_timeout_seconds = max(read_timeout_seconds, 1.0)
|
||||
self._closed = False
|
||||
|
||||
@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 = 8000,
|
||||
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.IncompleteReadError, asyncio.TimeoutError, ConnectionError):
|
||||
return None
|
||||
|
||||
if packet_type == AUDIO_SOCKET_PACKET_HANGUP:
|
||||
return None
|
||||
if packet_type == AUDIO_SOCKET_PACKET_PCM16:
|
||||
return payload
|
||||
if packet_type in {AUDIO_SOCKET_PACKET_UUID, AUDIO_SOCKET_PACKET_DTMF}:
|
||||
continue
|
||||
return None
|
||||
|
||||
async def send_audio(self, audio_chunk: bytes) -> None:
|
||||
if self._closed or self._writer.is_closing():
|
||||
return
|
||||
self._writer.write(encode_audio_packet(audio_chunk))
|
||||
await self._writer.drain()
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
if not self._writer.is_closing():
|
||||
self._writer.close()
|
||||
await self._writer.wait_closed()
|
||||
|
||||
|
||||
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 = 8000,
|
||||
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",
|
||||
self._host,
|
||||
self.bound_port,
|
||||
)
|
||||
|
||||
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",
|
||||
transport.transport_id,
|
||||
peer,
|
||||
)
|
||||
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:
|
||||
await transport.close()
|
||||
elif not writer.is_closing():
|
||||
writer.close()
|
||||
await writer.wait_closed()
|
||||
if current_task is not None:
|
||||
self._connection_tasks.discard(current_task)
|
||||
@@ -0,0 +1,49 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
|
||||
class BaseMediaTransport(ABC):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
transport_id: str,
|
||||
sample_rate_hz: int = 8000,
|
||||
frame_duration_ms: int = 20,
|
||||
) -> None:
|
||||
self._transport_id = transport_id
|
||||
self._sample_rate_hz = sample_rate_hz
|
||||
self._frame_duration_ms = frame_duration_ms
|
||||
|
||||
@property
|
||||
def transport_id(self) -> str:
|
||||
return self._transport_id
|
||||
|
||||
@property
|
||||
def sample_rate_hz(self) -> int:
|
||||
return self._sample_rate_hz
|
||||
|
||||
@property
|
||||
def frame_duration_ms(self) -> int:
|
||||
return self._frame_duration_ms
|
||||
|
||||
@property
|
||||
def frame_bytes(self) -> int:
|
||||
return int((self.sample_rate_hz * self.frame_duration_ms / 1000.0) * 2)
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def protocol(self) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def receive_audio(self) -> bytes | None:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def send_audio(self, audio_chunk: bytes) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def close(self) -> None:
|
||||
raise NotImplementedError
|
||||
@@ -0,0 +1,90 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
|
||||
from fastapi import WebSocket
|
||||
from starlette.websockets import WebSocketState
|
||||
|
||||
from realtime_voice_service.transports.base import BaseMediaTransport
|
||||
|
||||
|
||||
class WebSocketMediaTransport(BaseMediaTransport):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
websocket: WebSocket,
|
||||
transport_id: str,
|
||||
sample_rate_hz: int = 8000,
|
||||
frame_duration_ms: int = 20,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
transport_id=transport_id,
|
||||
sample_rate_hz=sample_rate_hz,
|
||||
frame_duration_ms=frame_duration_ms,
|
||||
)
|
||||
self._websocket = websocket
|
||||
self._closed = False
|
||||
|
||||
@property
|
||||
def protocol(self) -> str:
|
||||
return "websocket"
|
||||
|
||||
async def receive_audio(self) -> bytes | None:
|
||||
if self._closed:
|
||||
return None
|
||||
while not self._closed:
|
||||
message = await self._websocket.receive()
|
||||
message_type = message.get("type")
|
||||
if message_type == "websocket.disconnect":
|
||||
return None
|
||||
|
||||
binary_audio = message.get("bytes")
|
||||
if binary_audio is not None:
|
||||
return binary_audio
|
||||
|
||||
text_payload = message.get("text")
|
||||
if text_payload is None:
|
||||
continue
|
||||
decoded = self._decode_text_frame(text_payload)
|
||||
if decoded is None:
|
||||
continue
|
||||
return decoded
|
||||
return None
|
||||
|
||||
async def send_audio(self, audio_chunk: bytes) -> None:
|
||||
if self._closed:
|
||||
return
|
||||
await self._websocket.send_bytes(audio_chunk)
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
if (
|
||||
self._websocket.application_state == WebSocketState.DISCONNECTED
|
||||
or self._websocket.client_state == WebSocketState.DISCONNECTED
|
||||
):
|
||||
return
|
||||
try:
|
||||
await self._websocket.close(code=1000)
|
||||
except RuntimeError:
|
||||
return
|
||||
|
||||
@staticmethod
|
||||
def _decode_text_frame(text_payload: str) -> bytes | None:
|
||||
try:
|
||||
payload = json.loads(text_payload)
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
message_type = str(payload.get("type") or "").strip().lower()
|
||||
if message_type in {"close", "disconnect", "hangup"}:
|
||||
return None
|
||||
if message_type != "audio":
|
||||
return None
|
||||
|
||||
encoded_audio = payload.get("pcm16")
|
||||
if not isinstance(encoded_audio, str):
|
||||
return None
|
||||
return base64.b64decode(encoded_audio)
|
||||
Reference in New Issue
Block a user