Files

91 lines
2.6 KiB
Python

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)