ASR-demo/backend/realtime_websocket/model_service_qwen_legacy.py

257 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""独立实时 Demo 使用的 OpenAI 兼容 VLLM 服务适配器。"""
from __future__ import annotations
import asyncio
import base64
import io
import json
import wave
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from urllib.parse import urlsplit, urlunsplit
from aiohttp import ClientSession, ClientTimeout, FormData, WSMsgType
@dataclass(frozen=True)
class ModelServiceConfig:
"""一个独立 VLLM 端点所需的连接配置。"""
base_url: str = "http://127.0.0.1:9950/v1"
model: str = "Qwen/Qwen3-ASR-0.6B"
api_key: str = "EMPTY"
timeout_seconds: float = 45.0
realtime_enabled: bool = True
def pcm16_to_wav(pcm_bytes: bytes, sample_rate: int = 16000) -> bytes:
"""将浏览器发送的 PCM16 单声道数据封装为 VLLM 可识别的 WAV 请求。"""
output = io.BytesIO()
with wave.open(output, "wb") as wav_file:
wav_file.setnchannels(1)
wav_file.setsampwidth(2)
wav_file.setframerate(sample_rate)
wav_file.writeframes(pcm_bytes)
return output.getvalue()
def wav_to_pcm16(audio_bytes: bytes) -> bytes:
"""从 WAV 缓冲区提取 PCM 帧,并兼容尚未完整的中间音频数据。"""
try:
with wave.open(io.BytesIO(audio_bytes), "rb") as wav_file:
return wav_file.readframes(wav_file.getnframes())
except (EOFError, wave.Error):
if audio_bytes[:4] == b"RIFF" and audio_bytes[8:12] == b"WAVE" and len(audio_bytes) > 44:
return audio_bytes[44:]
return audio_bytes
def prepare_audio_request(
audio_bytes: bytes,
source: str,
file_name: str,
partial: bool,
) -> tuple[bytes, str, str] | None:
"""将麦克风、PCM 或 WAV 数据转换为 WAV;压缩格式的中间片段延迟到最终帧处理。"""
suffix = Path(file_name).suffix.lower()
if source == "mic" or suffix in {".pcm", ".wav"}:
pcm_bytes = wav_to_pcm16(audio_bytes) if suffix == ".wav" else audio_bytes
return pcm16_to_wav(pcm_bytes), "audio.wav", "audio/wav"
if partial:
# MP3/M4A/OGG 的不断增长前缀通常不是完整容器,不能安全解码,因此只在
# 最终阶段提交压缩文件,避免中间请求产生随机解码错误。
return None
content_type = {
".mp3": "audio/mpeg",
".m4a": "audio/mp4",
".ogg": "audio/ogg",
".opus": "audio/ogg",
}.get(suffix, "application/octet-stream")
return audio_bytes, Path(file_name).name or "audio.bin", content_type
def realtime_ws_url(base_url: str) -> str:
"""Convert the configured OpenAI-compatible base URL to vLLM's realtime URL."""
parsed = urlsplit(base_url.rstrip("/"))
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
raise ValueError(f"invalid VLLM base URL: {base_url}")
scheme = "wss" if parsed.scheme == "https" else "ws"
path = parsed.path.rstrip("/") + "/realtime"
return urlunsplit((scheme, parsed.netloc, path, "", ""))
class VLLMRealtimeStream:
"""One vLLM realtime stream, isolated from the shared HTTP client session."""
def __init__(self, websocket: Any, timeout_seconds: float) -> None:
self._websocket = websocket
self._timeout_seconds = timeout_seconds
self._latest_text = ""
self._error: Exception | None = None
self._done = asyncio.Event()
self._reader = asyncio.create_task(self._read_messages())
async def _read_messages(self) -> None:
"""Collect model deltas continuously so audio ingestion never waits for a snapshot."""
try:
async for message in self._websocket:
if message.type == WSMsgType.TEXT:
try:
payload = json.loads(message.data)
except (TypeError, ValueError):
continue
event_type = payload.get("type")
if event_type == "transcription.delta":
delta = str(payload.get("delta") or "")
if delta:
self._latest_text += delta
elif payload.get("text") is not None:
self._latest_text = str(payload["text"])
elif event_type == "transcription.done":
self._latest_text = str(
payload.get("text") or payload.get("transcript") or self._latest_text
).strip()
self._done.set()
elif event_type == "error":
detail = payload.get("error") or payload.get("message") or "unknown realtime error"
self._error = RuntimeError(str(detail))
self._done.set()
elif message.type == WSMsgType.ERROR:
self._error = self._websocket.exception() or RuntimeError("VLLM realtime WebSocket failed")
self._done.set()
return
elif message.type in {WSMsgType.CLOSE, WSMsgType.CLOSED, WSMsgType.CLOSING}:
if not self._done.is_set():
self._error = RuntimeError("VLLM realtime WebSocket closed before transcription.done")
self._done.set()
return
except asyncio.CancelledError:
raise
except Exception as exc:
self._error = exc
self._done.set()
def latest_text(self) -> str:
"""Return the newest model text already received by the reader task."""
return self._latest_text.strip()
async def append_audio(self, pcm_bytes: bytes) -> None:
"""Push one raw 16 kHz mono PCM16 block without re-uploading old audio."""
if not pcm_bytes:
return
if self._error is not None:
raise self._error
await self._websocket.send_json(
{
"type": "input_audio_buffer.append",
"audio": base64.b64encode(pcm_bytes).decode("ascii"),
}
)
async def finish(self) -> str:
"""Commit the current model turn and return the final realtime transcription."""
if self._error is not None:
raise self._error
await self._websocket.send_json({"type": "input_audio_buffer.commit", "final": True})
try:
await asyncio.wait_for(self._done.wait(), timeout=self._timeout_seconds)
except asyncio.TimeoutError as exc:
raise TimeoutError("VLLM realtime transcription timed out") from exc
if self._error is not None:
raise self._error
return self._latest_text.strip()
async def close(self) -> None:
"""Stop the reader and release the model-side WebSocket."""
if not self._reader.done():
self._reader.cancel()
await asyncio.gather(self._reader, return_exceptions=True)
if not self._websocket.closed:
await self._websocket.close()
class VLLMTranscriptionService:
"""只调用独立项目提供的 VLLM HTTP 接口,不导入原项目应用代码。"""
@property
def native_partial_supported(self) -> bool:
"""Report whether this adapter is configured to use vLLM realtime."""
return self.config.realtime_enabled
def __init__(self, config: ModelServiceConfig) -> None:
self.config = config
self._session: ClientSession | None = None
async def start(self) -> None:
"""创建可复用的 HTTP 会话,供所有中间和最终转写请求共享。"""
self._session = ClientSession(timeout=ClientTimeout(total=self.config.timeout_seconds))
async def close(self) -> None:
"""本地 Demo 退出时释放可复用的 HTTP 会话和底层连接。"""
if self._session is not None:
await self._session.close()
self._session = None
async def open_realtime_stream(self) -> VLLMRealtimeStream:
"""Open a model-native stream; the caller owns and closes the returned turn."""
if not self.config.realtime_enabled:
raise RuntimeError("VLLM realtime streaming is disabled")
if self._session is None:
raise RuntimeError("model service is not started")
endpoint = realtime_ws_url(self.config.base_url)
headers = {"Authorization": f"Bearer {self.config.api_key}"}
connect_timeout = min(5.0, self.config.timeout_seconds)
websocket = await self._session.ws_connect(
endpoint,
headers=headers,
timeout=connect_timeout,
heartbeat=20,
)
try:
created = await asyncio.wait_for(websocket.receive(), timeout=connect_timeout)
if created.type == WSMsgType.TEXT:
payload = json.loads(created.data)
if payload.get("type") == "error":
raise RuntimeError(
str(payload.get("error") or payload.get("message") or "VLLM realtime error")
)
await websocket.send_json({"type": "session.update", "model": self.config.model})
return VLLMRealtimeStream(websocket, self.config.timeout_seconds)
except Exception:
await websocket.close()
raise
async def transcribe(
self,
audio_bytes: bytes,
source: str,
file_name: str,
partial: bool,
) -> str | None:
"""提交一次音频快照并返回文本;返回 None 表示当前格式不支持中间转写。"""
prepared = prepare_audio_request(audio_bytes, source, file_name, partial)
if prepared is None:
return None
payload, upload_name, content_type = prepared
if self._session is None:
raise RuntimeError("model service is not started")
form = FormData()
form.add_field("file", payload, filename=upload_name, content_type=content_type)
form.add_field("model", self.config.model)
form.add_field("response_format", "json")
headers = {"Authorization": f"Bearer {self.config.api_key}"}
endpoint = self.config.base_url.rstrip("/") + "/audio/transcriptions"
async with self._session.post(endpoint, data=form, headers=headers) as response:
body = await response.text()
if response.status >= 400:
raise RuntimeError(f"VLLM transcription failed ({response.status}): {body[:500]}")
try:
decoded: Any = await response.json(content_type=None)
except ValueError:
return body.strip()
if isinstance(decoded, dict):
return str(decoded.get("text") or decoded.get("transcript") or "").strip()
return str(decoded).strip()