Skip to content

Commit 9165786

Browse files
authored
Merge pull request #46 from quanru/feat/deepgram-transport
feat(asr): add Deepgram streaming transport
2 parents ff70f61 + 7c45b7f commit 9165786

5 files changed

Lines changed: 562 additions & 0 deletions

File tree

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
"""Deepgram recognition provider support."""
Lines changed: 303 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,303 @@
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)
Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,57 @@
1+
"""Owner-only credentials for the Deepgram speech-to-text API."""
2+
3+
from __future__ import annotations
4+
5+
from dataclasses import dataclass
6+
import os
7+
import stat
8+
9+
from doubao_input.settings import config_dir, write_atomic
10+
11+
12+
@dataclass(frozen=True)
13+
class DeepgramCredentials:
14+
api_key: str
15+
16+
def validate(self) -> None:
17+
if (not isinstance(self.api_key, str) or not self.api_key.strip()
18+
or len(self.api_key) > 4096
19+
or any(char in self.api_key for char in "\r\n\x00")):
20+
raise ValueError("Invalid Deepgram API key")
21+
22+
23+
class DeepgramCredentialsStore:
24+
@staticmethod
25+
def path():
26+
return config_dir() / "doubao-say" / "deepgram_api_key"
27+
28+
@classmethod
29+
def load(cls) -> DeepgramCredentials | None:
30+
path = cls.path()
31+
if not path.exists():
32+
return None
33+
info = path.lstat()
34+
if not stat.S_ISREG(info.st_mode) or path.is_symlink():
35+
raise OSError("Unsafe Deepgram API key file")
36+
os.chmod(path, 0o600)
37+
credentials = DeepgramCredentials(path.read_text().strip())
38+
credentials.validate()
39+
return credentials
40+
41+
@classmethod
42+
def save(cls, credentials: DeepgramCredentials | str) -> None:
43+
if isinstance(credentials, str):
44+
credentials = DeepgramCredentials(credentials.strip())
45+
credentials.validate()
46+
write_atomic(cls.path(), (credentials.api_key.strip() + "\n").encode())
47+
48+
@classmethod
49+
def clear(cls) -> None:
50+
cls.path().unlink(missing_ok=True)
51+
52+
@classmethod
53+
def has_saved(cls) -> bool:
54+
try:
55+
return cls.load() is not None
56+
except (OSError, ValueError):
57+
return False

0 commit comments

Comments
 (0)