|
| 1 | +"""Session-owned Deepgram streaming speech-to-text transport.""" |
| 2 | + |
| 3 | +from __future__ import annotations |
| 4 | + |
| 5 | +import asyncio |
| 6 | +from collections import deque |
| 7 | +from dataclasses import dataclass, field |
| 8 | +import inspect |
| 9 | +import json |
| 10 | +import logging |
| 11 | +import threading |
| 12 | +from urllib.parse import urlencode |
| 13 | + |
| 14 | +from doubao_input.deepgram.credentials import DeepgramCredentials |
| 15 | + |
| 16 | + |
| 17 | +logger = logging.getLogger(__name__) |
| 18 | +DEEPGRAM_ASR_ENDPOINT = "wss://api.deepgram.com/v1/listen" |
| 19 | +MAX_PENDING_BYTES = 1024 * 1024 |
| 20 | + |
| 21 | + |
| 22 | +def deepgram_asr_url(language: str = "en-US") -> str: |
| 23 | + query = urlencode({ |
| 24 | + "model": "nova-3", |
| 25 | + "language": language, |
| 26 | + "encoding": "linear16", |
| 27 | + "sample_rate": 16000, |
| 28 | + "channels": 1, |
| 29 | + "interim_results": "true", |
| 30 | + "punctuate": "true", |
| 31 | + "smart_format": "true", |
| 32 | + }) |
| 33 | + return f"{DEEPGRAM_ASR_ENDPOINT}?{query}" |
| 34 | + |
| 35 | + |
| 36 | +@dataclass(eq=False) |
| 37 | +class _Session: |
| 38 | + audio: deque = field(default_factory=deque) |
| 39 | + final_segments: list[str] = field(default_factory=list) |
| 40 | + pending_bytes: int = 0 |
| 41 | + sending: bool = False |
| 42 | + accepting: bool = True |
| 43 | + connected: bool = False |
| 44 | + cancelled: bool = False |
| 45 | + callbacks: dict = field(default_factory=dict) |
| 46 | + loop: asyncio.AbstractEventLoop | None = None |
| 47 | + task: asyncio.Task | None = None |
| 48 | + wake: asyncio.Event | None = None |
| 49 | + thread: threading.Thread | None = None |
| 50 | + |
| 51 | + |
| 52 | +class DeepgramASRClient: |
| 53 | + """Thread-safe client matching the application's ASR lifecycle contract.""" |
| 54 | + |
| 55 | + def __init__(self, connect_factory=None, *, language: str = "en-US") -> None: |
| 56 | + self._lock = threading.RLock() |
| 57 | + self._session: _Session | None = None |
| 58 | + self._connect_factory = connect_factory |
| 59 | + self._language = language |
| 60 | + self.on_open = self.on_result = self.on_finish = None |
| 61 | + self.on_error = self.on_auth_error = None |
| 62 | + |
| 63 | + @property |
| 64 | + def is_connected(self) -> bool: |
| 65 | + with self._lock: |
| 66 | + return bool(self._session and self._session.connected) |
| 67 | + |
| 68 | + @property |
| 69 | + def has_pending_audio(self) -> bool: |
| 70 | + with self._lock: |
| 71 | + return bool(self._session and (self._session.audio or self._session.sending)) |
| 72 | + |
| 73 | + def prepare(self) -> None: |
| 74 | + self.disconnect() |
| 75 | + with self._lock: |
| 76 | + self._session = _Session() |
| 77 | + self._snapshot_callbacks(self._session) |
| 78 | + |
| 79 | + def connect(self, credentials: DeepgramCredentials) -> None: |
| 80 | + credentials.validate() |
| 81 | + with self._lock: |
| 82 | + if self._session is None or self._session.thread is not None: |
| 83 | + self.prepare() |
| 84 | + session = self._session |
| 85 | + self._snapshot_callbacks(session) |
| 86 | + session.thread = threading.Thread( |
| 87 | + target=self._run, |
| 88 | + args=(session, credentials), |
| 89 | + name="deepgram-asr", |
| 90 | + daemon=True, |
| 91 | + ) |
| 92 | + session.thread.start() |
| 93 | + |
| 94 | + def _snapshot_callbacks(self, session) -> None: |
| 95 | + session.callbacks = {name: getattr(self, name) for name in ( |
| 96 | + "on_open", "on_result", "on_finish", "on_error", "on_auth_error")} |
| 97 | + |
| 98 | + def _emit(self, session, name, *args) -> None: |
| 99 | + with self._lock: |
| 100 | + if self._session is session and not session.cancelled: |
| 101 | + callback = session.callbacks.get(name) |
| 102 | + if callback: |
| 103 | + callback(*args) |
| 104 | + |
| 105 | + def _run(self, session, credentials) -> None: |
| 106 | + loop = asyncio.new_event_loop() |
| 107 | + asyncio.set_event_loop(loop) |
| 108 | + try: |
| 109 | + with self._lock: |
| 110 | + if session.cancelled: |
| 111 | + return |
| 112 | + session.loop = loop |
| 113 | + session.wake = asyncio.Event() |
| 114 | + session.task = loop.create_task( |
| 115 | + self._listen(session, credentials)) |
| 116 | + loop.run_until_complete(session.task) |
| 117 | + except asyncio.CancelledError: |
| 118 | + pass |
| 119 | + finally: |
| 120 | + pending = asyncio.all_tasks(loop) |
| 121 | + for task in pending: |
| 122 | + task.cancel() |
| 123 | + if pending: |
| 124 | + loop.run_until_complete(asyncio.gather( |
| 125 | + *pending, return_exceptions=True)) |
| 126 | + loop.run_until_complete(loop.shutdown_asyncgens()) |
| 127 | + loop.close() |
| 128 | + with self._lock: |
| 129 | + session.connected = False |
| 130 | + session.accepting = False |
| 131 | + session.audio.clear() |
| 132 | + session.pending_bytes = 0 |
| 133 | + |
| 134 | + async def _listen(self, session, credentials) -> None: |
| 135 | + try: |
| 136 | + factory = self._connect_factory or _load_websockets().connect |
| 137 | + headers = {"Authorization": f"Token {credentials.api_key}"} |
| 138 | + async with factory( |
| 139 | + deepgram_asr_url(self._language), |
| 140 | + open_timeout=5, |
| 141 | + close_timeout=3, |
| 142 | + max_size=2**20, |
| 143 | + **_websocket_header_kwargs(factory, headers), |
| 144 | + ) as websocket: |
| 145 | + with self._lock: |
| 146 | + if session.cancelled: |
| 147 | + return |
| 148 | + session.connected = True |
| 149 | + self._emit(session, "on_open") |
| 150 | + sender = asyncio.create_task(self._send(session, websocket)) |
| 151 | + receiver = asyncio.create_task(self._receive(session, websocket)) |
| 152 | + done, _ = await asyncio.wait( |
| 153 | + (sender, receiver), return_when=asyncio.FIRST_COMPLETED) |
| 154 | + if receiver in done: |
| 155 | + receiver.result() |
| 156 | + if not sender.done(): |
| 157 | + sender.cancel() |
| 158 | + await asyncio.gather(sender, return_exceptions=True) |
| 159 | + else: |
| 160 | + sender.result() |
| 161 | + await receiver |
| 162 | + except asyncio.CancelledError: |
| 163 | + raise |
| 164 | + except Exception as error: |
| 165 | + logger.warning("Deepgram ASR transport failed (%s)", type(error).__name__) |
| 166 | + if _is_auth_error(error): |
| 167 | + self._emit(session, "on_auth_error") |
| 168 | + else: |
| 169 | + self._emit(session, "on_error", RuntimeError( |
| 170 | + "Deepgram ASR connection failed")) |
| 171 | + finally: |
| 172 | + with self._lock: |
| 173 | + session.connected = False |
| 174 | + |
| 175 | + async def _send(self, session, websocket) -> None: |
| 176 | + while True: |
| 177 | + session.wake.clear() |
| 178 | + while True: |
| 179 | + with self._lock: |
| 180 | + if session.cancelled: |
| 181 | + return |
| 182 | + chunk = session.audio.popleft() if session.audio else None |
| 183 | + if chunk is not None: |
| 184 | + session.pending_bytes -= len(chunk) |
| 185 | + session.sending = True |
| 186 | + finished = not session.accepting and not session.audio |
| 187 | + if chunk is None: |
| 188 | + if finished: |
| 189 | + await websocket.send(json.dumps({"type": "Finalize"})) |
| 190 | + return |
| 191 | + break |
| 192 | + try: |
| 193 | + await websocket.send(chunk) |
| 194 | + finally: |
| 195 | + with self._lock: |
| 196 | + session.sending = False |
| 197 | + await session.wake.wait() |
| 198 | + |
| 199 | + async def _receive(self, session, websocket) -> None: |
| 200 | + async for raw_message in websocket: |
| 201 | + if not isinstance(raw_message, str): |
| 202 | + raise ValueError("Deepgram returned a non-JSON response") |
| 203 | + message = json.loads(raw_message) |
| 204 | + kind = message.get("type") |
| 205 | + if kind == "Error": |
| 206 | + code = str(message.get("code", "")) |
| 207 | + if code in {"INVALID_AUTH", "INSUFFICIENT_PERMISSIONS"}: |
| 208 | + self._emit(session, "on_auth_error") |
| 209 | + else: |
| 210 | + self._emit(session, "on_error", RuntimeError( |
| 211 | + "Deepgram rejected the request")) |
| 212 | + return |
| 213 | + if kind != "Results": |
| 214 | + continue |
| 215 | + alternatives = message.get("channel", {}).get("alternatives", []) |
| 216 | + transcript = alternatives[0].get("transcript", "") if alternatives else "" |
| 217 | + if message.get("is_final") and transcript: |
| 218 | + session.final_segments.append(transcript.strip()) |
| 219 | + snapshot = " ".join([ |
| 220 | + *session.final_segments, |
| 221 | + *([transcript.strip()] if transcript and not message.get("is_final") else []), |
| 222 | + ]).strip() |
| 223 | + if snapshot: |
| 224 | + self._emit(session, "on_result", snapshot) |
| 225 | + if message.get("from_finalize"): |
| 226 | + await websocket.send(json.dumps({"type": "CloseStream"})) |
| 227 | + self._emit(session, "on_finish") |
| 228 | + return |
| 229 | + with self._lock: |
| 230 | + finished = not session.accepting |
| 231 | + if finished: |
| 232 | + self._emit(session, "on_finish") |
| 233 | + else: |
| 234 | + raise ConnectionError("Deepgram closed before recording ended") |
| 235 | + |
| 236 | + def send_audio(self, data: bytes) -> None: |
| 237 | + with self._lock: |
| 238 | + session = self._session |
| 239 | + if session is None or session.cancelled or not session.accepting: |
| 240 | + return |
| 241 | + if session.pending_bytes + len(data) > MAX_PENDING_BYTES: |
| 242 | + self._emit(session, "on_error", RuntimeError( |
| 243 | + "ASR audio queue is full")) |
| 244 | + self.disconnect() |
| 245 | + return |
| 246 | + session.audio.append(bytes(data)) |
| 247 | + session.pending_bytes += len(data) |
| 248 | + self._wake(session) |
| 249 | + |
| 250 | + def finish_sending(self) -> None: |
| 251 | + with self._lock: |
| 252 | + if self._session: |
| 253 | + self._session.accepting = False |
| 254 | + self._wake(self._session) |
| 255 | + |
| 256 | + @staticmethod |
| 257 | + def _wake(session) -> None: |
| 258 | + if session.loop and session.wake and not session.loop.is_closed(): |
| 259 | + try: |
| 260 | + session.loop.call_soon_threadsafe(session.wake.set) |
| 261 | + except RuntimeError: |
| 262 | + pass |
| 263 | + |
| 264 | + def disconnect(self) -> None: |
| 265 | + with self._lock: |
| 266 | + session, self._session = self._session, None |
| 267 | + if session is None: |
| 268 | + return |
| 269 | + session.cancelled = True |
| 270 | + session.connected = False |
| 271 | + session.audio.clear() |
| 272 | + session.pending_bytes = 0 |
| 273 | + if session.loop and session.task: |
| 274 | + try: |
| 275 | + session.loop.call_soon_threadsafe(session.task.cancel) |
| 276 | + except RuntimeError: |
| 277 | + pass |
| 278 | + |
| 279 | + |
| 280 | +def _load_websockets(): |
| 281 | + try: |
| 282 | + import websockets |
| 283 | + except ImportError as error: |
| 284 | + raise RuntimeError( |
| 285 | + "Install the Python websockets package before recording") from error |
| 286 | + return websockets |
| 287 | + |
| 288 | + |
| 289 | +def _websocket_header_kwargs(connect_func, headers: dict[str, str]) -> dict: |
| 290 | + try: |
| 291 | + parameters = inspect.signature(connect_func).parameters |
| 292 | + except (TypeError, ValueError): |
| 293 | + return {"additional_headers": headers} |
| 294 | + return {"additional_headers" if "additional_headers" in parameters |
| 295 | + else "extra_headers": headers} |
| 296 | + |
| 297 | + |
| 298 | +def _is_auth_error(error: Exception) -> bool: |
| 299 | + response = getattr(error, "response", None) |
| 300 | + status = getattr(response, "status_code", None) |
| 301 | + if status is None: |
| 302 | + status = getattr(error, "status_code", None) |
| 303 | + return status in (401, 403) |
0 commit comments