91 lines
2.6 KiB
Python
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)
|