"""独立实时 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()