mirror of
https://github.com/suitenumerique/meet.git
synced 2026-08-25 09:46:47 +00:00
wip - single websockets
This commit is contained in:
+170
-44
@@ -4,6 +4,8 @@ vLLM exposes Voxtral Realtime over a WebSocket that follows the OpenAI Realtime
|
|||||||
API protocol (not Mistral's proprietary realtime protocol).
|
API protocol (not Mistral's proprietary realtime protocol).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import base64
|
import base64
|
||||||
import json
|
import json
|
||||||
@@ -11,7 +13,7 @@ import logging
|
|||||||
import os
|
import os
|
||||||
import weakref
|
import weakref
|
||||||
from collections import deque
|
from collections import deque
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass, field
|
||||||
|
|
||||||
import websockets
|
import websockets
|
||||||
from livekit.agents import (
|
from livekit.agents import (
|
||||||
@@ -35,6 +37,12 @@ NUM_CHANNELS = 1
|
|||||||
CHUNK_SAMPLES = 1600 # 100 ms @ 16 kHz mono
|
CHUNK_SAMPLES = 1600 # 100 ms @ 16 kHz mono
|
||||||
PREROLL_CHUNKS = 5 # keep 500 ms of audio before START_OF_SPEECH
|
PREROLL_CHUNKS = 5 # keep 500 ms of audio before START_OF_SPEECH
|
||||||
|
|
||||||
|
# Reconnect policy: exponential backoff capped at MAX, give up after MAX_ATTEMPTS
|
||||||
|
# consecutive failures (a successful handshake resets the counter).
|
||||||
|
RECONNECT_BACKOFF_BASE_S = 0.5
|
||||||
|
RECONNECT_BACKOFF_MAX_S = 8.0
|
||||||
|
RECONNECT_MAX_ATTEMPTS = 5
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class _STTOptions:
|
class _STTOptions:
|
||||||
@@ -43,6 +51,22 @@ class _STTOptions:
|
|||||||
api_key: str | None
|
api_key: str | None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _PendingUtterance:
|
||||||
|
"""An utterance in flight on the shared websocket.
|
||||||
|
|
||||||
|
`sent_chunks` holds every chunk we have already enqueued for send on this
|
||||||
|
or a prior connection; on reconnect we replay them before resuming reads
|
||||||
|
from `queue`. vLLM concatenates `input_audio_buffer.append` events into a
|
||||||
|
single audio buffer per generation, so duplicates from a partial prior send
|
||||||
|
are harmless.
|
||||||
|
"""
|
||||||
|
|
||||||
|
queue: asyncio.Queue[bytes | None]
|
||||||
|
sent_chunks: list[bytes] = field(default_factory=list)
|
||||||
|
ended: bool = False # the None sentinel has been drained from `queue`
|
||||||
|
|
||||||
|
|
||||||
class STT(stt.STT):
|
class STT(stt.STT):
|
||||||
"""LiveKit STT speaking the OpenAI Realtime protocol served by vLLM."""
|
"""LiveKit STT speaking the OpenAI Realtime protocol served by vLLM."""
|
||||||
|
|
||||||
@@ -146,11 +170,16 @@ class SpeechStream(stt.RecognizeStream):
|
|||||||
super().__init__(stt=stt, conn_options=conn_options, sample_rate=SAMPLE_RATE)
|
super().__init__(stt=stt, conn_options=conn_options, sample_rate=SAMPLE_RATE)
|
||||||
self._opts = opts
|
self._opts = opts
|
||||||
self._vad = vad_instance
|
self._vad = vad_instance
|
||||||
self._audio_q: asyncio.Queue[bytes | None] = asyncio.Queue()
|
self._utterance_q: asyncio.Queue[bytes | None] | None = None
|
||||||
self._speaking = False
|
self._speaking = False
|
||||||
self._preroll: deque[bytes] = deque(maxlen=PREROLL_CHUNKS)
|
self._preroll: deque[bytes] = deque(maxlen=PREROLL_CHUNKS)
|
||||||
self._idle = asyncio.Event()
|
# Voxtral realtime is strictly sequential: only one generation runs at a
|
||||||
self._idle.set()
|
# time, and a new `commit` is ignored while the previous one is still
|
||||||
|
# producing. We queue per-utterance audio buffers here and let the
|
||||||
|
# pipeline process them one by one on the shared websocket.
|
||||||
|
self._utterance_chan: asyncio.Queue[asyncio.Queue[bytes | None] | None] = (
|
||||||
|
asyncio.Queue()
|
||||||
|
)
|
||||||
|
|
||||||
@utils.log_exceptions(logger=logger)
|
@utils.log_exceptions(logger=logger)
|
||||||
async def _run(self) -> None:
|
async def _run(self) -> None:
|
||||||
@@ -182,51 +211,42 @@ class SpeechStream(stt.RecognizeStream):
|
|||||||
elif ev.type == vad_module.VADEventType.END_OF_SPEECH:
|
elif ev.type == vad_module.VADEventType.END_OF_SPEECH:
|
||||||
self._on_end_of_speech()
|
self._on_end_of_speech()
|
||||||
|
|
||||||
headers: dict[str, str] = {}
|
pipeline_t = asyncio.create_task(self._utterance_pipeline())
|
||||||
if self._opts.api_key:
|
|
||||||
headers["Authorization"] = f"Bearer {self._opts.api_key}"
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
async with websockets.connect(
|
await asyncio.gather(input_task(), vad_task())
|
||||||
self._opts.base_url,
|
# signal end-of-stream; pipeline finishes pending utterances first
|
||||||
additional_headers=headers,
|
self._utterance_chan.put_nowait(None)
|
||||||
open_timeout=self._conn_options.timeout,
|
await pipeline_t
|
||||||
) as ws:
|
|
||||||
request_id = await self._handshake(ws)
|
|
||||||
send_t = asyncio.create_task(self._send_audio(ws, self._audio_q))
|
|
||||||
recv_t = asyncio.create_task(self._receive_events(ws, request_id))
|
|
||||||
try:
|
|
||||||
await asyncio.gather(input_task(), vad_task())
|
|
||||||
# wait for any in-flight utterance to finish transcribing
|
|
||||||
await self._idle.wait()
|
|
||||||
finally:
|
|
||||||
for t in (send_t, recv_t):
|
|
||||||
if not t.done():
|
|
||||||
t.cancel()
|
|
||||||
for t in (send_t, recv_t):
|
|
||||||
try:
|
|
||||||
await t
|
|
||||||
except (asyncio.CancelledError, Exception): # noqa: BLE001
|
|
||||||
pass
|
|
||||||
await vad_stream.aclose()
|
|
||||||
except (APIStatusError, APIConnectionError, asyncio.CancelledError):
|
except (APIStatusError, APIConnectionError, asyncio.CancelledError):
|
||||||
raise
|
raise
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.exception("vLLM realtime stream failed")
|
logger.exception("vLLM realtime stream failed")
|
||||||
raise APIConnectionError() from exc
|
raise APIConnectionError() from exc
|
||||||
|
finally:
|
||||||
|
if not pipeline_t.done():
|
||||||
|
pipeline_t.cancel()
|
||||||
|
try:
|
||||||
|
await pipeline_t
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
except Exception:
|
||||||
|
logger.exception("utterance pipeline failed during finalize")
|
||||||
|
await vad_stream.aclose()
|
||||||
|
|
||||||
def _handle_chunk(self, chunk: bytes) -> None:
|
def _handle_chunk(self, chunk: bytes) -> None:
|
||||||
self._preroll.append(chunk)
|
self._preroll.append(chunk)
|
||||||
if self._speaking:
|
if self._speaking and self._utterance_q is not None:
|
||||||
self._audio_q.put_nowait(chunk)
|
self._utterance_q.put_nowait(chunk)
|
||||||
|
|
||||||
def _on_start_of_speech(self) -> None:
|
def _on_start_of_speech(self) -> None:
|
||||||
if self._speaking:
|
if self._speaking:
|
||||||
return
|
return
|
||||||
self._speaking = True
|
self._speaking = True
|
||||||
self._idle.clear()
|
q: asyncio.Queue[bytes | None] = asyncio.Queue()
|
||||||
for chunk in self._preroll:
|
for chunk in self._preroll:
|
||||||
self._audio_q.put_nowait(chunk)
|
q.put_nowait(chunk)
|
||||||
|
self._utterance_q = q
|
||||||
|
self._utterance_chan.put_nowait(q)
|
||||||
self._event_ch.send_nowait(
|
self._event_ch.send_nowait(
|
||||||
stt.SpeechEvent(type=stt.SpeechEventType.START_OF_SPEECH)
|
stt.SpeechEvent(type=stt.SpeechEventType.START_OF_SPEECH)
|
||||||
)
|
)
|
||||||
@@ -235,7 +255,9 @@ class SpeechStream(stt.RecognizeStream):
|
|||||||
if not self._speaking:
|
if not self._speaking:
|
||||||
return
|
return
|
||||||
self._speaking = False
|
self._speaking = False
|
||||||
self._audio_q.put_nowait(None)
|
if self._utterance_q is not None:
|
||||||
|
self._utterance_q.put_nowait(None)
|
||||||
|
self._utterance_q = None
|
||||||
self._event_ch.send_nowait(
|
self._event_ch.send_nowait(
|
||||||
stt.SpeechEvent(type=stt.SpeechEventType.END_OF_SPEECH)
|
stt.SpeechEvent(type=stt.SpeechEventType.END_OF_SPEECH)
|
||||||
)
|
)
|
||||||
@@ -249,20 +271,122 @@ class SpeechStream(stt.RecognizeStream):
|
|||||||
body=created,
|
body=created,
|
||||||
)
|
)
|
||||||
await ws.send(json.dumps({"type": "session.update", "model": self._opts.model}))
|
await ws.send(json.dumps({"type": "session.update", "model": self._opts.model}))
|
||||||
await ws.send(json.dumps({"type": "input_audio_buffer.commit"}))
|
|
||||||
return created.get("id", "")
|
return created.get("id", "")
|
||||||
|
|
||||||
|
def _auth_headers(self) -> dict[str, str]:
|
||||||
|
if self._opts.api_key:
|
||||||
|
return {"Authorization": f"Bearer {self._opts.api_key}"}
|
||||||
|
return {}
|
||||||
|
|
||||||
|
async def _utterance_pipeline(self) -> None:
|
||||||
|
# Owns the websocket lifecycle. On drop, reopens and resumes the
|
||||||
|
# in-flight utterance (if any) by replaying its already-sent chunks.
|
||||||
|
pending: _PendingUtterance | None = None
|
||||||
|
attempt = 0
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
async with websockets.connect(
|
||||||
|
self._opts.base_url,
|
||||||
|
additional_headers=self._auth_headers(),
|
||||||
|
open_timeout=self._conn_options.timeout,
|
||||||
|
) as ws:
|
||||||
|
request_id = await self._handshake(ws)
|
||||||
|
attempt = 0
|
||||||
|
while True:
|
||||||
|
if pending is None:
|
||||||
|
q = await self._utterance_chan.get()
|
||||||
|
if q is None:
|
||||||
|
return
|
||||||
|
pending = _PendingUtterance(queue=q)
|
||||||
|
await self._process_utterance(ws, pending, request_id)
|
||||||
|
pending = None
|
||||||
|
except (websockets.WebSocketException, OSError, TimeoutError) as exc:
|
||||||
|
attempt += 1
|
||||||
|
if attempt > RECONNECT_MAX_ATTEMPTS:
|
||||||
|
logger.exception(
|
||||||
|
"vLLM realtime: giving up after %d reconnect attempts",
|
||||||
|
RECONNECT_MAX_ATTEMPTS,
|
||||||
|
)
|
||||||
|
raise APIConnectionError() from exc
|
||||||
|
backoff = min(
|
||||||
|
RECONNECT_BACKOFF_BASE_S * (2 ** (attempt - 1)),
|
||||||
|
RECONNECT_BACKOFF_MAX_S,
|
||||||
|
)
|
||||||
|
if pending is None:
|
||||||
|
logger.warning(
|
||||||
|
"vLLM WS connection lost between utterances "
|
||||||
|
"(attempt %d/%d): %s; retrying in %.1fs",
|
||||||
|
attempt,
|
||||||
|
RECONNECT_MAX_ATTEMPTS,
|
||||||
|
exc,
|
||||||
|
backoff,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"vLLM WS dropped mid-utterance (%d chunks buffered, "
|
||||||
|
"ended=%s, attempt %d/%d): %s; retrying in %.1fs",
|
||||||
|
len(pending.sent_chunks),
|
||||||
|
pending.ended,
|
||||||
|
attempt,
|
||||||
|
RECONNECT_MAX_ATTEMPTS,
|
||||||
|
exc,
|
||||||
|
backoff,
|
||||||
|
)
|
||||||
|
await asyncio.sleep(backoff)
|
||||||
|
|
||||||
|
async def _process_utterance(
|
||||||
|
self,
|
||||||
|
ws: websockets.ClientConnection,
|
||||||
|
pending: _PendingUtterance,
|
||||||
|
request_id: str,
|
||||||
|
) -> None:
|
||||||
|
# Start a fresh generation. Safe to send here: the previous utterance's
|
||||||
|
# transcription.done has already been received (we await it below), so
|
||||||
|
# the server-side generation_task is done and won't ignore this commit.
|
||||||
|
await ws.send(json.dumps({"type": "input_audio_buffer.commit"}))
|
||||||
|
send_t = asyncio.create_task(self._send_audio(ws, pending))
|
||||||
|
try:
|
||||||
|
await self._receive_one_transcription(ws, request_id)
|
||||||
|
finally:
|
||||||
|
if not send_t.done():
|
||||||
|
send_t.cancel()
|
||||||
|
try:
|
||||||
|
await send_t
|
||||||
|
except (asyncio.CancelledError, websockets.WebSocketException):
|
||||||
|
pass
|
||||||
|
except Exception:
|
||||||
|
logger.exception("send-audio task failed during finalize")
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def _send_audio(
|
async def _send_audio(
|
||||||
ws: websockets.ClientConnection, q: asyncio.Queue[bytes | None]
|
ws: websockets.ClientConnection, pending: _PendingUtterance
|
||||||
) -> None:
|
) -> None:
|
||||||
|
# Replay anything already sent on a previous (now-dead) connection.
|
||||||
|
# sent_chunks is appended before send, so a chunk that failed to send
|
||||||
|
# last time is still present and gets retried here.
|
||||||
|
for chunk in pending.sent_chunks:
|
||||||
|
await ws.send(
|
||||||
|
json.dumps(
|
||||||
|
{
|
||||||
|
"type": "input_audio_buffer.append",
|
||||||
|
"audio": base64.b64encode(chunk).decode("ascii"),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if pending.ended:
|
||||||
|
await ws.send(
|
||||||
|
json.dumps({"type": "input_audio_buffer.commit", "final": True})
|
||||||
|
)
|
||||||
|
return
|
||||||
while True:
|
while True:
|
||||||
chunk = await q.get()
|
chunk = await pending.queue.get()
|
||||||
if chunk is None:
|
if chunk is None:
|
||||||
|
pending.ended = True
|
||||||
await ws.send(
|
await ws.send(
|
||||||
json.dumps({"type": "input_audio_buffer.commit", "final": True})
|
json.dumps({"type": "input_audio_buffer.commit", "final": True})
|
||||||
)
|
)
|
||||||
continue
|
return
|
||||||
|
pending.sent_chunks.append(chunk)
|
||||||
await ws.send(
|
await ws.send(
|
||||||
json.dumps(
|
json.dumps(
|
||||||
{
|
{
|
||||||
@@ -272,16 +396,19 @@ class SpeechStream(stt.RecognizeStream):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _receive_events(
|
async def _receive_one_transcription(
|
||||||
self, ws: websockets.ClientConnection, request_id: str
|
self, ws: websockets.ClientConnection, request_id: str
|
||||||
) -> None:
|
) -> None:
|
||||||
|
# Use recv() rather than `async for`: the latter swallows
|
||||||
|
# ConnectionClosed on close-mid-iteration, which would let a dropped
|
||||||
|
# WS look like a clean "no transcription" return.
|
||||||
current_text = ""
|
current_text = ""
|
||||||
async for raw in ws:
|
while True:
|
||||||
|
raw = await ws.recv()
|
||||||
data = json.loads(raw)
|
data = json.loads(raw)
|
||||||
event_type = data.get("type")
|
event_type = data.get("type")
|
||||||
|
|
||||||
if event_type == "transcription.delta":
|
if event_type == "transcription.delta":
|
||||||
# TODO: fix this
|
|
||||||
current_text += data.get("delta", "")
|
current_text += data.get("delta", "")
|
||||||
elif event_type == "transcription.done":
|
elif event_type == "transcription.done":
|
||||||
final_text = data.get("text") or current_text
|
final_text = data.get("text") or current_text
|
||||||
@@ -308,8 +435,7 @@ class SpeechStream(stt.RecognizeStream):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
current_text = ""
|
return
|
||||||
self._idle.set()
|
|
||||||
elif event_type == "error":
|
elif event_type == "error":
|
||||||
err = data.get("error")
|
err = data.get("error")
|
||||||
raise APIStatusError(str(err), status_code=500, body=data)
|
raise APIStatusError(str(err), status_code=500, body=data)
|
||||||
|
|||||||
Reference in New Issue
Block a user