257 lines
11 KiB
Python
257 lines
11 KiB
Python
"""独立实时 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()
|