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 = 16000, 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_frame(self, frame: bytes) -> None: if self._closed: return await self._websocket.send_bytes(frame) 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)