747 lines
35 KiB
Python
747 lines
35 KiB
Python
|
|
"""面向浏览器的独立 Qwen3-ASR VLLM WebSocket 编排服务。"""
|
|||
|
|
|
|||
|
|
from __future__ import annotations
|
|||
|
|
|
|||
|
|
import argparse
|
|||
|
|
import asyncio
|
|||
|
|
import json
|
|||
|
|
import math
|
|||
|
|
import logging
|
|||
|
|
import os
|
|||
|
|
import time
|
|||
|
|
import webbrowser
|
|||
|
|
from dataclasses import dataclass
|
|||
|
|
from pathlib import Path
|
|||
|
|
from typing import Any
|
|||
|
|
from urllib.parse import urlparse
|
|||
|
|
from uuid import uuid4
|
|||
|
|
|
|||
|
|
from aiohttp import WSMsgType, web
|
|||
|
|
from dotenv import load_dotenv
|
|||
|
|
|
|||
|
|
from auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig
|
|||
|
|
from model_service import ModelServiceConfig, VLLMTranscriptionService
|
|||
|
|
from speaker_assembler import SegmentAssembler
|
|||
|
|
|
|||
|
|
|
|||
|
|
# 与部署启动器读取同一配置;外部环境变量优先于 demo/.env。
|
|||
|
|
DEPLOY_ROOT = Path(__file__).resolve().parents[1]
|
|||
|
|
load_dotenv(DEPLOY_ROOT / ".env")
|
|||
|
|
|
|||
|
|
# 监听所有网卡,允许同一局域网内的浏览器访问服务器上的 Demo;端口集中在代码
|
|||
|
|
# 变量中维护,便于服务器部署时直接修改并保持页面和 WebSocket 使用一致端口。
|
|||
|
|
WEB_HOST = "0.0.0.0"
|
|||
|
|
WEB_PORT = 8082
|
|||
|
|
WEB_DISPLAY_HOST = os.getenv("WEB_DISPLAY_HOST", "127.0.0.1")
|
|||
|
|
DEFAULT_MODEL_SERVICE_URL = f"http://127.0.0.1:{os.getenv('VLLM_PORT', '9950')}/v1"
|
|||
|
|
PARTIAL_BYTES_PER_SECOND = 16000 * 2
|
|||
|
|
VAD_FRAME_BYTES = 640
|
|||
|
|
VAD_FRAME_MS = 20
|
|||
|
|
VAD_SILENCE_MS = 800
|
|||
|
|
PARAGRAPH_SILENCE_MS = 1400
|
|||
|
|
VAD_RMS_THRESHOLD = 450
|
|||
|
|
MIN_SPEAKER_VOICE_MS = 800
|
|||
|
|
LOGGER = logging.getLogger(__name__)
|
|||
|
|
|
|||
|
|
|
|||
|
|
class EndOfStream:
|
|||
|
|
"""带明确类型的队列结束标记,用于区分控制信号和真实音频字节。"""
|
|||
|
|
|
|||
|
|
|
|||
|
|
EOF = EndOfStream()
|
|||
|
|
MODEL_SERVICE_KEY = web.AppKey("model_service", VLLMTranscriptionService)
|
|||
|
|
AUXILIARY_SERVICE_KEY = web.AppKey("auxiliary_service", AuxiliaryModelService)
|
|||
|
|
|
|||
|
|
|
|||
|
|
@dataclass(frozen=True)
|
|||
|
|
class SpeakerJob:
|
|||
|
|
"""等待辅助服务处理的单个已完成 turn;只保存该 turn 的 PCM 音频。"""
|
|||
|
|
|
|||
|
|
sentence_id: int
|
|||
|
|
audio: bytes
|
|||
|
|
start_time_ms: float
|
|||
|
|
end_time_ms: float
|
|||
|
|
voiced_ms: float = 0.0
|
|||
|
|
|
|||
|
|
|
|||
|
|
def validate_model_service_url(value: str) -> str:
|
|||
|
|
"""只接受用户输入的 HTTP(S) VLLM 地址,并拒绝附带认证和查询参数的地址。"""
|
|||
|
|
candidate = value.strip().rstrip("/")
|
|||
|
|
parsed = urlparse(candidate)
|
|||
|
|
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
|
|||
|
|
raise ValueError("VLLM 地址必须是完整的 http:// 或 https:// URL")
|
|||
|
|
if parsed.username or parsed.password or parsed.query or parsed.fragment:
|
|||
|
|
raise ValueError("VLLM 地址不能包含账号、密码、查询参数或片段")
|
|||
|
|
return candidate
|
|||
|
|
|
|||
|
|
|
|||
|
|
@dataclass
|
|||
|
|
class SessionMetrics:
|
|||
|
|
"""记录 WebSocket 会话耗时和结果修订次数,并在会话结束后展示。"""
|
|||
|
|
|
|||
|
|
started_at: float
|
|||
|
|
audio_bytes: int = 0
|
|||
|
|
input_chunks: int = 0
|
|||
|
|
partial_count: int = 0
|
|||
|
|
partial_revisions: int = 0
|
|||
|
|
first_partial_ms: float | None = None
|
|||
|
|
final_ms: float | None = None
|
|||
|
|
|
|||
|
|
def snapshot(self) -> dict[str, Any]:
|
|||
|
|
"""返回可安全序列化为 JSON 的指标,耗时均相对于会话开始时间计算。"""
|
|||
|
|
now = time.perf_counter()
|
|||
|
|
return {
|
|||
|
|
"audio_bytes": self.audio_bytes,
|
|||
|
|
"input_chunks": self.input_chunks,
|
|||
|
|
"partial_count": self.partial_count,
|
|||
|
|
"partial_revisions": self.partial_revisions,
|
|||
|
|
"first_partial_ms": self.first_partial_ms,
|
|||
|
|
"final_ms": self.final_ms,
|
|||
|
|
"elapsed_ms": round((now - self.started_at) * 1000, 1),
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
|
|||
|
|
class RealtimeSession:
|
|||
|
|
"""串行处理音频快照,并通过同一个 WebSocket 有序推送状态更新。"""
|
|||
|
|
|
|||
|
|
def __init__(
|
|||
|
|
self,
|
|||
|
|
ws: web.WebSocketResponse,
|
|||
|
|
model_service: VLLMTranscriptionService,
|
|||
|
|
auxiliary_service: AuxiliaryModelService | None,
|
|||
|
|
start: dict[str, Any],
|
|||
|
|
) -> None:
|
|||
|
|
self.ws = ws
|
|||
|
|
self.model_service = model_service
|
|||
|
|
self.auxiliary_service = auxiliary_service
|
|||
|
|
self.start = start
|
|||
|
|
# 聚类状态只能属于当前连接,客户端复用 ID 不能串入另一会话的声纹池。
|
|||
|
|
self.session_id = uuid4().hex
|
|||
|
|
self.send_lock = asyncio.Lock()
|
|||
|
|
self.state_revision = 0
|
|||
|
|
self.audio_queue: asyncio.Queue[bytes | EndOfStream] = asyncio.Queue(maxsize=256)
|
|||
|
|
self.speaker_queue: asyncio.Queue[SpeakerJob | EndOfStream] = asyncio.Queue(maxsize=64)
|
|||
|
|
self.assembler = SegmentAssembler()
|
|||
|
|
self.metrics = SessionMetrics(time.perf_counter())
|
|||
|
|
self.source = str(start.get("source") or "mic")
|
|||
|
|
self.file_name = str(start.get("file_name") or "audio.wav")
|
|||
|
|
self.windowed_partial = self.source == "mic" or Path(self.file_name).suffix.lower() in {".pcm", ".wav"}
|
|||
|
|
self.sentence_strategy = int(start.get("sentence_strategy") or 0)
|
|||
|
|
self.silence_limit_ms = PARAGRAPH_SILENCE_MS if self.sentence_strategy == 1 else VAD_SILENCE_MS
|
|||
|
|
self.partial_interval_ms = max(300, int(start.get("partial_interval_ms") or 1200))
|
|||
|
|
self.max_segment_sec = max(2.0, float(start.get("max_segment_sec") or 12.0))
|
|||
|
|
self.merge_adjacent = self._parse_flag(start.get("display_merge"), True)
|
|||
|
|
self.enable_native_partial = self._parse_flag(start.get("enable_native_partial_stream"), True)
|
|||
|
|
self.segment_id = 0
|
|||
|
|
self.segment_audio = bytearray()
|
|||
|
|
self.segment_start_ms = 0.0
|
|||
|
|
self.vad_buffer = bytearray()
|
|||
|
|
self.processed_audio_bytes = 0
|
|||
|
|
self.silence_ms = 0
|
|||
|
|
self.in_speech = False
|
|||
|
|
self.voiced_ms = 0.0
|
|||
|
|
self.pre_roll = bytearray()
|
|||
|
|
self.wav_header_buffer = bytearray()
|
|||
|
|
self.wav_payload_started = self.source != "file" or Path(self.file_name).suffix.lower() != ".wav"
|
|||
|
|
self.wav_riff_read = False
|
|||
|
|
self.wav_format_valid = False
|
|||
|
|
self.wav_data_remaining: int | None = None
|
|||
|
|
self.speaker_warning_sent = False
|
|||
|
|
self.speaker_enabled = self._parse_flag(start.get("speaker_diarization"), True)
|
|||
|
|
self.input_stopped = False
|
|||
|
|
|
|||
|
|
@staticmethod
|
|||
|
|
def _parse_flag(value: Any, default: bool) -> bool:
|
|||
|
|
"""兼容前端传来的 0/1、布尔值和字符串开关,避免字符串 0 被误判为真。"""
|
|||
|
|
if value is None:
|
|||
|
|
return default
|
|||
|
|
if isinstance(value, str):
|
|||
|
|
return value.strip().lower() not in {"", "0", "false", "no", "off"}
|
|||
|
|
return bool(value)
|
|||
|
|
|
|||
|
|
async def emit(self, payload: dict[str, Any]) -> None:
|
|||
|
|
"""在连接仍然有效时发送一条有序事件,避免向已关闭连接写入数据。"""
|
|||
|
|
async with self.send_lock:
|
|||
|
|
if not self.ws.closed:
|
|||
|
|
await self.ws.send_json(payload)
|
|||
|
|
|
|||
|
|
async def emit_state(self, sentence: dict[str, Any] | None = None) -> None:
|
|||
|
|
"""每次状态更新后同时发送原始状态和重新计算的展示快照。"""
|
|||
|
|
# 在首次 await 前冻结快照,音频 worker 与 speaker worker 不会混用两版状态。
|
|||
|
|
self.state_revision += 1
|
|||
|
|
state = {
|
|||
|
|
"type": "display_state", "revision": self.state_revision,
|
|||
|
|
"raw_segments": self.assembler.raw_snapshot(),
|
|||
|
|
"display_blocks": self.assembler.display_blocks(self.merge_adjacent),
|
|||
|
|
"metrics": self.metrics.snapshot(),
|
|||
|
|
}
|
|||
|
|
if sentence is not None:
|
|||
|
|
await self.emit({"type": "sentences", "sentences": [sentence], "metrics": self.metrics.snapshot()})
|
|||
|
|
await self.emit(state)
|
|||
|
|
|
|||
|
|
async def warn_speaker(self, message: str) -> None:
|
|||
|
|
"""只发送一次说话人服务告警,避免辅助服务异常时刷屏。"""
|
|||
|
|
if self.speaker_warning_sent:
|
|||
|
|
return
|
|||
|
|
await self.emit(
|
|||
|
|
{
|
|||
|
|
"type": "speaker_warning",
|
|||
|
|
"session_id": self.session_id,
|
|||
|
|
"speaker_service_url": getattr(getattr(self.auxiliary_service, "config", None), "base_url", None),
|
|||
|
|
"message": message,
|
|||
|
|
}
|
|||
|
|
)
|
|||
|
|
self.speaker_warning_sent = True
|
|||
|
|
|
|||
|
|
def _duration_ms(self) -> float:
|
|||
|
|
"""根据 16 kHz PCM 字节数计算时长,不依赖前端可能漂移的时间戳。"""
|
|||
|
|
return len(self.segment_audio) / PARTIAL_BYTES_PER_SECOND * 1000
|
|||
|
|
|
|||
|
|
async def _transcribe(self, partial: bool) -> str | None:
|
|||
|
|
"""通过 VLLM 适配器转写当前逻辑片段,并保留中间/最终请求的统一入口。"""
|
|||
|
|
return await self.model_service.transcribe(
|
|||
|
|
bytes(self.segment_audio),
|
|||
|
|
"mic",
|
|||
|
|
"turn.pcm",
|
|||
|
|
partial=partial,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
async def _emit_transcription(self, text: str, sentence_type: int, end_ms: float, commit_reason: str | None = None) -> None:
|
|||
|
|
"""写入或更新一条句子,确保中间结果和最终结果不会在前端产生重复行。"""
|
|||
|
|
if not text:
|
|||
|
|
return
|
|||
|
|
sentence = self.assembler.apply_sentence(
|
|||
|
|
{
|
|||
|
|
"sentence_id": self.segment_id,
|
|||
|
|
"sentence": text,
|
|||
|
|
"sentence_type": sentence_type,
|
|||
|
|
"start_time": self.segment_start_ms,
|
|||
|
|
"end_time": end_ms,
|
|||
|
|
"speaker_id": -1,
|
|||
|
|
"speaker_name": "",
|
|||
|
|
"speaker_evidence": "pending",
|
|||
|
|
"speaker_confidence": 0.0,
|
|||
|
|
"speaker_strategy": "vllm_no_speaker_evidence",
|
|||
|
|
"commit_reason": commit_reason,
|
|||
|
|
"speaker_status": ("queued" if sentence_type else "waiting_final") if self.speaker_enabled else "disabled",
|
|||
|
|
"speaker_reason": ("等待声纹处理" if sentence_type else "语音片段结束后识别说话人") if self.speaker_enabled else "说话人分离已关闭",
|
|||
|
|
}
|
|||
|
|
)
|
|||
|
|
if sentence_type == 0:
|
|||
|
|
self.metrics.partial_count += 1
|
|||
|
|
if sentence["revision_count"] > 0:
|
|||
|
|
self.metrics.partial_revisions += 1
|
|||
|
|
if self.metrics.first_partial_ms is None:
|
|||
|
|
self.metrics.first_partial_ms = round((time.perf_counter() - self.metrics.started_at) * 1000, 1)
|
|||
|
|
else:
|
|||
|
|
self.metrics.final_ms = round((time.perf_counter() - self.metrics.started_at) * 1000, 1)
|
|||
|
|
await self.emit_state(sentence)
|
|||
|
|
|
|||
|
|
@staticmethod
|
|||
|
|
def _is_voice_frame(frame: bytes) -> bool:
|
|||
|
|
"""用 PCM 帧的 RMS 判断是否有语音,作为实时低延迟切句触发器。"""
|
|||
|
|
if not frame:
|
|||
|
|
return False
|
|||
|
|
samples = memoryview(frame).cast("h")
|
|||
|
|
if not samples:
|
|||
|
|
return False
|
|||
|
|
square_mean = sum(sample * sample for sample in samples) / len(samples)
|
|||
|
|
return math.sqrt(square_mean) >= VAD_RMS_THRESHOLD
|
|||
|
|
|
|||
|
|
def _strip_wav_header(self, chunk: bytes) -> bytes:
|
|||
|
|
"""增量解析 RIFF chunk;支持扩展头,并拒绝采样率或声道不匹配的 WAV。"""
|
|||
|
|
if self.wav_payload_started:
|
|||
|
|
if self.wav_data_remaining is None:
|
|||
|
|
return chunk
|
|||
|
|
payload = chunk[:self.wav_data_remaining]
|
|||
|
|
self.wav_data_remaining -= len(payload)
|
|||
|
|
return payload
|
|||
|
|
self.wav_header_buffer.extend(chunk)
|
|||
|
|
buffer = self.wav_header_buffer
|
|||
|
|
if not self.wav_riff_read:
|
|||
|
|
if len(buffer) < 12:
|
|||
|
|
return b""
|
|||
|
|
if buffer[:4] != b"RIFF" or buffer[8:12] != b"WAVE":
|
|||
|
|
raise ValueError("文件不是有效的 RIFF/WAV 音频")
|
|||
|
|
del buffer[:12]
|
|||
|
|
self.wav_riff_read = True
|
|||
|
|
while len(buffer) >= 8:
|
|||
|
|
kind = bytes(buffer[:4])
|
|||
|
|
size = int.from_bytes(buffer[4:8], "little")
|
|||
|
|
if kind == b"data":
|
|||
|
|
if not self.wav_format_valid or size % 2:
|
|||
|
|
raise ValueError("WAV 必须为 16kHz、单声道、PCM16")
|
|||
|
|
self.wav_data_remaining = size
|
|||
|
|
self.wav_payload_started = True
|
|||
|
|
payload = bytes(buffer[8:8 + size])
|
|||
|
|
self.wav_data_remaining -= len(payload)
|
|||
|
|
buffer.clear()
|
|||
|
|
return payload
|
|||
|
|
if size > 1024 * 1024:
|
|||
|
|
raise ValueError("WAV 元数据头过大,请转换为标准 PCM WAV")
|
|||
|
|
chunk_size = 8 + size + (size % 2)
|
|||
|
|
if len(buffer) < chunk_size:
|
|||
|
|
return b""
|
|||
|
|
if kind == b"fmt ":
|
|||
|
|
fmt = buffer[8:8 + size]
|
|||
|
|
fields = (int.from_bytes(fmt[0:2], "little"), int.from_bytes(fmt[2:4], "little"),
|
|||
|
|
int.from_bytes(fmt[4:8], "little"), int.from_bytes(fmt[14:16], "little"))
|
|||
|
|
if size < 16 or fields != (1, 1, 16000, 16):
|
|||
|
|
raise ValueError("WAV 必须为 16kHz、单声道、PCM16,请先转换音频")
|
|||
|
|
self.wav_format_valid = True
|
|||
|
|
del buffer[:chunk_size]
|
|||
|
|
return b""
|
|||
|
|
|
|||
|
|
async def _resolve_speaker(self, job: SpeakerJob) -> None:
|
|||
|
|
"""异步解析单个 turn 的说话人,并把结果覆盖回同一个 sentence_id。"""
|
|||
|
|
if not self.speaker_enabled:
|
|||
|
|
return
|
|||
|
|
async def update_status(status: str, reason: str) -> None:
|
|||
|
|
"""将每个失败或等待阶段回写原片段,避免只发一次全局告警。"""
|
|||
|
|
updated = self.assembler.apply_speaker_update({
|
|||
|
|
"sentence_id": job.sentence_id, "speaker_id": -1,
|
|||
|
|
"speaker_evidence": "pending", "speaker_confidence": 0.0,
|
|||
|
|
"speaker_status": status, "speaker_reason": reason,
|
|||
|
|
})
|
|||
|
|
await self.emit_state(updated)
|
|||
|
|
|
|||
|
|
# 按有效有声帧检查长度,不能让句尾 800ms 静音把短插话伪装成长样本。
|
|||
|
|
if job.voiced_ms < MIN_SPEAKER_VOICE_MS:
|
|||
|
|
await update_status("insufficient_audio", f"有效语音不足 {MIN_SPEAKER_VOICE_MS}ms,不继承上一位说话人")
|
|||
|
|
return
|
|||
|
|
if self.auxiliary_service is None:
|
|||
|
|
await update_status("service_unavailable", "未配置说话人辅助模型服务")
|
|||
|
|
await self.warn_speaker("未配置辅助模型服务,无法执行实时说话人分离")
|
|||
|
|
return
|
|||
|
|
await update_status("processing", "正在提取声纹并匹配说话人")
|
|||
|
|
try:
|
|||
|
|
speaker = await self.auxiliary_service.resolve_speaker(
|
|||
|
|
job.audio,
|
|||
|
|
self.session_id,
|
|||
|
|
job.start_time_ms,
|
|||
|
|
job.end_time_ms,
|
|||
|
|
)
|
|||
|
|
except Exception as exc:
|
|||
|
|
# 辅助服务异常不能阻断 ASR;当前片段继续保持 pending,方便定位服务问题。
|
|||
|
|
LOGGER.exception("speaker resolve failed: session=%s sentence=%s", self.session_id, job.sentence_id)
|
|||
|
|
await update_status("service_error", str(exc))
|
|||
|
|
await self.warn_speaker(str(exc))
|
|||
|
|
return
|
|||
|
|
if not speaker:
|
|||
|
|
await update_status("no_embedding", "辅助服务未返回可用声纹结果")
|
|||
|
|
return
|
|||
|
|
update = dict(speaker)
|
|||
|
|
update["sentence_id"] = job.sentence_id
|
|||
|
|
update["speaker_name"] = str(update.get("speaker_name") or "")
|
|||
|
|
updated = self.assembler.apply_speaker_update(update)
|
|||
|
|
if updated is not None:
|
|||
|
|
LOGGER.info("speaker result: session=%s sentence=%s status=%s strategy=%s", self.session_id,
|
|||
|
|
job.sentence_id, updated.get("speaker_status"), updated.get("speaker_strategy"))
|
|||
|
|
await self.emit_state(updated)
|
|||
|
|
|
|||
|
|
async def process_speakers(self) -> None:
|
|||
|
|
"""按 turn 顺序串行访问辅助模型,保证在线聚类中心不会乱序更新。"""
|
|||
|
|
while True:
|
|||
|
|
item = await self.speaker_queue.get()
|
|||
|
|
if isinstance(item, EndOfStream):
|
|||
|
|
return
|
|||
|
|
await self._resolve_speaker(item)
|
|||
|
|
|
|||
|
|
async def _commit_segment(self, reason: str = "final") -> None:
|
|||
|
|
"""在 VAD 检测到一句结束后提交 final,并异步排队当前 turn 的说话人解析。"""
|
|||
|
|
if not self.segment_audio or not self.in_speech:
|
|||
|
|
return
|
|||
|
|
# 去掉句尾触发切段的静音,ASR 与声纹都使用当前片段的真实有效范围。
|
|||
|
|
trailing_bytes = int(self.silence_ms * PARTIAL_BYTES_PER_SECOND / 1000)
|
|||
|
|
if trailing_bytes:
|
|||
|
|
del self.segment_audio[-trailing_bytes:]
|
|||
|
|
final_audio = bytes(self.segment_audio)
|
|||
|
|
final_start_ms = self.segment_start_ms
|
|||
|
|
final_end_ms = self.segment_start_ms + self._duration_ms()
|
|||
|
|
final_sentence_id = self.segment_id
|
|||
|
|
text = await self._transcribe(partial=False)
|
|||
|
|
if text:
|
|||
|
|
await self._emit_transcription(text, 1, final_end_ms, reason)
|
|||
|
|
if self.speaker_enabled:
|
|||
|
|
await self.speaker_queue.put(
|
|||
|
|
SpeakerJob(
|
|||
|
|
sentence_id=final_sentence_id,
|
|||
|
|
audio=final_audio,
|
|||
|
|
start_time_ms=final_start_ms,
|
|||
|
|
end_time_ms=final_end_ms,
|
|||
|
|
voiced_ms=self.voiced_ms,
|
|||
|
|
)
|
|||
|
|
)
|
|||
|
|
elif self.assembler.segments.pop(final_sentence_id, None) is not None:
|
|||
|
|
# final 判定无文本时撤回临时结果,不能留下永远等待声纹的 partial。
|
|||
|
|
await self.emit_state()
|
|||
|
|
self.segment_audio.clear()
|
|||
|
|
self.segment_start_ms = self.processed_audio_bytes / PARTIAL_BYTES_PER_SECOND * 1000
|
|||
|
|
self.segment_id += 1
|
|||
|
|
self.silence_ms = 0
|
|||
|
|
self.in_speech = False
|
|||
|
|
self.voiced_ms = 0.0
|
|||
|
|
|
|||
|
|
async def process_audio(self) -> None:
|
|||
|
|
"""消费音频,以 VAD 静音结束作为切句主逻辑,并按窗口发送 partial。"""
|
|||
|
|
last_partial_bytes = 0
|
|||
|
|
partial_bytes = int(self.partial_interval_ms / 1000 * PARTIAL_BYTES_PER_SECOND)
|
|||
|
|
while True:
|
|||
|
|
item = await self.audio_queue.get()
|
|||
|
|
if isinstance(item, EndOfStream):
|
|||
|
|
break
|
|||
|
|
chunk = self._strip_wav_header(item) if self.source == "file" else item
|
|||
|
|
self.metrics.input_chunks += 1
|
|||
|
|
if not self.windowed_partial:
|
|||
|
|
# websocket_handler 已拒绝压缩文件;这里保留防御分支,避免未来
|
|||
|
|
# 新客户端绕过入口时又悄悄退化成“整段上传后切片”。
|
|||
|
|
raise RuntimeError("实时流式模式只接受 16kHz PCM16 音频")
|
|||
|
|
if not chunk:
|
|||
|
|
continue
|
|||
|
|
self.vad_buffer.extend(chunk)
|
|||
|
|
while len(self.vad_buffer) >= VAD_FRAME_BYTES:
|
|||
|
|
frame = bytes(self.vad_buffer[:VAD_FRAME_BYTES])
|
|||
|
|
del self.vad_buffer[:VAD_FRAME_BYTES]
|
|||
|
|
frame_start_ms = self.processed_audio_bytes / PARTIAL_BYTES_PER_SECOND * 1000
|
|||
|
|
self.processed_audio_bytes += len(frame)
|
|||
|
|
self.metrics.audio_bytes += len(frame)
|
|||
|
|
voiced = self._is_voice_frame(frame)
|
|||
|
|
if voiced and not self.in_speech:
|
|||
|
|
self.in_speech = True
|
|||
|
|
self.segment_start_ms = frame_start_ms - len(self.pre_roll) / PARTIAL_BYTES_PER_SECOND * 1000
|
|||
|
|
self.segment_audio = bytearray(self.pre_roll)
|
|||
|
|
self.pre_roll.clear()
|
|||
|
|
last_partial_bytes = 0
|
|||
|
|
if self.in_speech:
|
|||
|
|
self.segment_audio.extend(frame)
|
|||
|
|
if voiced:
|
|||
|
|
self.voiced_ms += VAD_FRAME_MS
|
|||
|
|
self.silence_ms = 0 if voiced else self.silence_ms + VAD_FRAME_MS
|
|||
|
|
if len(self.segment_audio) - last_partial_bytes >= partial_bytes and self.silence_ms < self.silence_limit_ms:
|
|||
|
|
text = await self._transcribe(partial=True)
|
|||
|
|
if text:
|
|||
|
|
await self._emit_transcription(text, 0, self.segment_start_ms + self._duration_ms())
|
|||
|
|
last_partial_bytes = len(self.segment_audio)
|
|||
|
|
if self.silence_ms >= self.silence_limit_ms or self._duration_ms() >= self.max_segment_sec * 1000:
|
|||
|
|
await self._commit_segment("silence" if self.silence_ms >= self.silence_limit_ms else "max_duration")
|
|||
|
|
last_partial_bytes = 0
|
|||
|
|
else:
|
|||
|
|
# 参考原 WebSocket 保留 200ms 前滚,减少首字低能量音素被裁掉。
|
|||
|
|
self.pre_roll.extend(frame)
|
|||
|
|
del self.pre_roll[:-6400]
|
|||
|
|
|
|||
|
|
if not self.wav_payload_started or (not self.input_stopped and self.wav_data_remaining not in (None, 0)):
|
|||
|
|
raise ValueError("WAV 文件不完整,未收到全部音频数据")
|
|||
|
|
if len(self.vad_buffer) % 2:
|
|||
|
|
raise ValueError("PCM16 音频必须包含完整的双字节采样")
|
|||
|
|
if self.windowed_partial and self.vad_buffer:
|
|||
|
|
tail_ms = len(self.vad_buffer) / PARTIAL_BYTES_PER_SECOND * 1000
|
|||
|
|
tail_voiced = self._is_voice_frame(bytes(self.vad_buffer))
|
|||
|
|
if tail_voiced and not self.in_speech:
|
|||
|
|
self.in_speech = True
|
|||
|
|
self.segment_start_ms = self.processed_audio_bytes / PARTIAL_BYTES_PER_SECOND * 1000
|
|||
|
|
self.processed_audio_bytes += len(self.vad_buffer)
|
|||
|
|
self.metrics.audio_bytes += len(self.vad_buffer)
|
|||
|
|
if self.in_speech:
|
|||
|
|
self.segment_audio.extend(self.vad_buffer)
|
|||
|
|
self.silence_ms = 0 if tail_voiced else self.silence_ms + tail_ms
|
|||
|
|
self.voiced_ms += tail_ms if tail_voiced else 0
|
|||
|
|
self.vad_buffer.clear()
|
|||
|
|
if self.segment_audio:
|
|||
|
|
if not self.windowed_partial:
|
|||
|
|
self.in_speech = True
|
|||
|
|
await self._commit_segment()
|
|||
|
|
await self.emit({"type": "metrics", "metrics": self.metrics.snapshot()})
|
|||
|
|
|
|||
|
|
|
|||
|
|
def deployment_model_name() -> str:
|
|||
|
|
"""把部署脚本的 0.6b/1.7b 别名解析成 vLLM 对外发布的模型名。"""
|
|||
|
|
public_name = os.getenv("VLLM_SERVED_MODEL_NAME")
|
|||
|
|
if public_name:
|
|||
|
|
return public_name
|
|||
|
|
requested = os.getenv("QWEN3_ASR_MODEL", "default")
|
|||
|
|
with (DEPLOY_ROOT / "model_manifest.json").open(encoding="utf-8") as source:
|
|||
|
|
manifest = json.load(source)
|
|||
|
|
if requested == "default":
|
|||
|
|
return str(manifest["default_model"])
|
|||
|
|
for model_id, config in manifest["models"].items():
|
|||
|
|
if requested.lower() == str(config.get("alias", "")).lower():
|
|||
|
|
return model_id
|
|||
|
|
return requested
|
|||
|
|
|
|||
|
|
|
|||
|
|
def parse_args() -> argparse.Namespace:
|
|||
|
|
"""只解析服务选择参数;浏览器服务端口继续由代码内部变量统一维护。"""
|
|||
|
|
parser = argparse.ArgumentParser(description=__doc__)
|
|||
|
|
parser.add_argument("--model-service-url", default=os.getenv("MODEL_SERVICE_URL", DEFAULT_MODEL_SERVICE_URL))
|
|||
|
|
parser.add_argument("--model", default=deployment_model_name())
|
|||
|
|
parser.add_argument("--no-browser", action="store_true")
|
|||
|
|
return parser.parse_args()
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def index_handler(_: web.Request) -> web.FileResponse:
|
|||
|
|
"""返回独立 Demo 测试页面,并避免入口页缓存旧的脚本版本号。"""
|
|||
|
|
# 入口页必须每次重新校验,配合 app.js 的版本号变更,避免用户继续运行旧前端。
|
|||
|
|
return web.FileResponse(
|
|||
|
|
Path(__file__).parent / "static" / "index.html",
|
|||
|
|
headers={"Cache-Control": "no-store"},
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def config_handler(request: web.Request) -> web.Response:
|
|||
|
|
"""暴露服务启动时的默认配置,让页面自动填充 VLLM 地址和模型名。"""
|
|||
|
|
config = request.app[MODEL_SERVICE_KEY].config
|
|||
|
|
auxiliary = request.app.get(AUXILIARY_SERVICE_KEY)
|
|||
|
|
return web.json_response({
|
|||
|
|
"model_service_url": config.base_url, "model": config.model,
|
|||
|
|
"speaker_service_url": getattr(getattr(auxiliary, "config", None), "base_url", None),
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
|
|||
|
|
"""处理一个浏览器会话,每个连接独立保存音频、句子和展示状态。"""
|
|||
|
|
ws = web.WebSocketResponse(max_msg_size=64 * 1024 * 1024)
|
|||
|
|
await ws.prepare(request)
|
|||
|
|
default_model_service: VLLMTranscriptionService = request.app[MODEL_SERVICE_KEY]
|
|||
|
|
model_service = default_model_service
|
|||
|
|
owns_model_service = False
|
|||
|
|
processing: asyncio.Task[None] | None = None
|
|||
|
|
speaker_processing: asyncio.Task[None] | None = None
|
|||
|
|
session: RealtimeSession | None = None
|
|||
|
|
try:
|
|||
|
|
first = await ws.receive()
|
|||
|
|
if first.type != WSMsgType.TEXT:
|
|||
|
|
await ws.send_json({"type": "error", "message": "first message must be JSON start"})
|
|||
|
|
return ws
|
|||
|
|
try:
|
|||
|
|
start = json.loads(first.data)
|
|||
|
|
except json.JSONDecodeError:
|
|||
|
|
await ws.send_json({"type": "error", "message": "invalid start JSON"})
|
|||
|
|
return ws
|
|||
|
|
if not isinstance(start, dict) or start.get("type") != "start":
|
|||
|
|
await ws.send_json({"type": "error", "message": "first message must have type=start"})
|
|||
|
|
return ws
|
|||
|
|
|
|||
|
|
# 本项目用于验证实时流式链路,文件模式只接受可以按 PCM 帧连续处理的
|
|||
|
|
# WAV/PCM;MP3、M4A 等压缩容器只能在文件完整到达后解码,不纳入本次测试。
|
|||
|
|
source = str(start.get("source") or "mic")
|
|||
|
|
file_suffix = Path(str(start.get("file_name") or "")).suffix.lower()
|
|||
|
|
if source == "file" and file_suffix not in {".pcm", ".wav"}:
|
|||
|
|
await ws.send_json(
|
|||
|
|
{
|
|||
|
|
"type": "error",
|
|||
|
|
"message": "实时流式测试的文件模式只支持 PCM 或 WAV,请改用麦克风、PCM 或 WAV",
|
|||
|
|
}
|
|||
|
|
)
|
|||
|
|
return ws
|
|||
|
|
|
|||
|
|
# 页面可以在不重启 WebSocket Demo 的情况下为当前会话切换 VLLM 地址;
|
|||
|
|
# 未切换时继续复用默认服务,避免普通场景为每个连接重复创建 HTTP 会话。
|
|||
|
|
try:
|
|||
|
|
requested_url = validate_model_service_url(
|
|||
|
|
str(start.get("model_service_url") or default_model_service.config.base_url)
|
|||
|
|
)
|
|||
|
|
except ValueError as exc:
|
|||
|
|
await ws.send_json({"type": "error", "message": str(exc)})
|
|||
|
|
return ws
|
|||
|
|
requested_model = str(start.get("model") or default_model_service.config.model).strip()
|
|||
|
|
if (
|
|||
|
|
requested_url != default_model_service.config.base_url
|
|||
|
|
or requested_model != default_model_service.config.model
|
|||
|
|
):
|
|||
|
|
model_service = VLLMTranscriptionService(
|
|||
|
|
ModelServiceConfig(base_url=requested_url, model=requested_model)
|
|||
|
|
)
|
|||
|
|
await model_service.start()
|
|||
|
|
owns_model_service = True
|
|||
|
|
|
|||
|
|
auxiliary_service = request.app.get(AUXILIARY_SERVICE_KEY)
|
|||
|
|
session = RealtimeSession(ws, model_service, auxiliary_service, start)
|
|||
|
|
auxiliary_config = getattr(auxiliary_service, "config", None)
|
|||
|
|
speaker_health: dict[str, Any] | None = None
|
|||
|
|
speaker_health_error: str | None = None
|
|||
|
|
if session.speaker_enabled and auxiliary_service is not None:
|
|||
|
|
# 健康检查只用于尽早暴露辅助服务问题;即使失败也不阻断 ASR,
|
|||
|
|
# 这样可以从同一页面继续观察 ASR 与说话人链路的差异。
|
|||
|
|
health_check = getattr(auxiliary_service, "health", None)
|
|||
|
|
if callable(health_check):
|
|||
|
|
try:
|
|||
|
|
speaker_health = await asyncio.wait_for(health_check(), timeout=5)
|
|||
|
|
if speaker_health.get("speaker_embedding_ready", speaker_health.get("ready")) is False:
|
|||
|
|
speaker_health_error = "辅助模型服务未就绪,请检查 /health 返回的 models 状态"
|
|||
|
|
except Exception as exc:
|
|||
|
|
speaker_health_error = f"说话人辅助服务不可用:{exc}"
|
|||
|
|
await session.emit(
|
|||
|
|
{
|
|||
|
|
"type": "start",
|
|||
|
|
"model_service_url": model_service.config.base_url,
|
|||
|
|
"model": model_service.config.model,
|
|||
|
|
"session_id": session.session_id,
|
|||
|
|
"enable_native_partial_stream": session.enable_native_partial,
|
|||
|
|
"native_partial_supported": model_service.native_partial_supported,
|
|||
|
|
"partial_mode": "http_cumulative_window",
|
|||
|
|
"speaker_diarization_enabled": session.speaker_enabled,
|
|||
|
|
"speaker_service_url": getattr(auxiliary_config, "base_url", None),
|
|||
|
|
"speaker_service_health": speaker_health,
|
|||
|
|
"sentence_strategy": session.sentence_strategy,
|
|||
|
|
"silence_limit_ms": session.silence_limit_ms,
|
|||
|
|
"display_state_supported": True,
|
|||
|
|
}
|
|||
|
|
)
|
|||
|
|
if speaker_health_error:
|
|||
|
|
await session.warn_speaker(speaker_health_error)
|
|||
|
|
processing = asyncio.create_task(session.process_audio())
|
|||
|
|
if session.speaker_enabled:
|
|||
|
|
# speaker worker 与音频处理并行运行;它只消费已经结束的 turn,
|
|||
|
|
# 因此不会阻塞下一帧音频进入队列或影响 ASR partial 输出。
|
|||
|
|
speaker_processing = asyncio.create_task(session.process_speakers())
|
|||
|
|
async def guarded(operation):
|
|||
|
|
"""接收和队列背压同时监听 worker,推理失败立即报错而非永远等 stop。"""
|
|||
|
|
pending = asyncio.create_task(operation)
|
|||
|
|
try:
|
|||
|
|
workers = [task for task in (processing, speaker_processing) if task is not None]
|
|||
|
|
done, _ = await asyncio.wait([pending, *workers], return_when=asyncio.FIRST_COMPLETED)
|
|||
|
|
if pending in done:
|
|||
|
|
return await pending
|
|||
|
|
for worker in workers:
|
|||
|
|
if worker in done:
|
|||
|
|
await worker
|
|||
|
|
raise RuntimeError("实时处理任务意外结束")
|
|||
|
|
return await pending
|
|||
|
|
finally:
|
|||
|
|
if not pending.done():
|
|||
|
|
pending.cancel()
|
|||
|
|
await asyncio.gather(pending, return_exceptions=True)
|
|||
|
|
|
|||
|
|
input_finished = False
|
|||
|
|
while not ws.closed:
|
|||
|
|
message = await guarded(ws.receive())
|
|||
|
|
if message.type == WSMsgType.BINARY:
|
|||
|
|
await guarded(session.audio_queue.put(bytes(message.data)))
|
|||
|
|
continue
|
|||
|
|
if message.type == WSMsgType.TEXT:
|
|||
|
|
try:
|
|||
|
|
control = json.loads(message.data)
|
|||
|
|
except json.JSONDecodeError:
|
|||
|
|
continue
|
|||
|
|
if not isinstance(control, dict):
|
|||
|
|
continue
|
|||
|
|
if control.get("type") in {"eof", "stop"}:
|
|||
|
|
input_finished = True
|
|||
|
|
session.input_stopped = control.get("type") == "stop"
|
|||
|
|
await session.emit({"type": "draining", "message": "正在完成转写和说话人识别"})
|
|||
|
|
await guarded(session.audio_queue.put(EOF))
|
|||
|
|
break
|
|||
|
|
if control.get("type") == "abort":
|
|||
|
|
processing.cancel()
|
|||
|
|
if speaker_processing is not None:
|
|||
|
|
speaker_processing.cancel()
|
|||
|
|
await asyncio.gather(
|
|||
|
|
processing,
|
|||
|
|
*(task for task in [speaker_processing] if task is not None),
|
|||
|
|
return_exceptions=True,
|
|||
|
|
)
|
|||
|
|
return ws
|
|||
|
|
if message.type in {WSMsgType.ERROR, WSMsgType.CLOSE, WSMsgType.CLOSED}:
|
|||
|
|
processing.cancel()
|
|||
|
|
if speaker_processing is not None:
|
|||
|
|
speaker_processing.cancel()
|
|||
|
|
await asyncio.gather(
|
|||
|
|
processing,
|
|||
|
|
*(task for task in [speaker_processing] if task is not None),
|
|||
|
|
return_exceptions=True,
|
|||
|
|
)
|
|||
|
|
return ws
|
|||
|
|
if not input_finished:
|
|||
|
|
# 浏览器或网络断开后,音频生产者已经不存在,不能让处理任务继续等待
|
|||
|
|
# 永远不会到来的 EOF,因此这里主动取消任务并回收异常结果。
|
|||
|
|
processing.cancel()
|
|||
|
|
if speaker_processing is not None:
|
|||
|
|
speaker_processing.cancel()
|
|||
|
|
await asyncio.gather(
|
|||
|
|
processing,
|
|||
|
|
*(task for task in [speaker_processing] if task is not None),
|
|||
|
|
return_exceptions=True,
|
|||
|
|
)
|
|||
|
|
return ws
|
|||
|
|
try:
|
|||
|
|
await processing
|
|||
|
|
if speaker_processing is not None:
|
|||
|
|
await session.speaker_queue.put(EOF)
|
|||
|
|
await speaker_processing
|
|||
|
|
await session.emit_state()
|
|||
|
|
await session.emit({
|
|||
|
|
"type": "end", "metrics": session.metrics.snapshot(),
|
|||
|
|
"sentences": session.assembler.raw_snapshot(),
|
|||
|
|
"display_blocks": session.assembler.display_blocks(session.merge_adjacent),
|
|||
|
|
})
|
|||
|
|
except asyncio.CancelledError:
|
|||
|
|
raise
|
|||
|
|
except Exception as exc:
|
|||
|
|
await session.emit({"type": "error", "message": str(exc)})
|
|||
|
|
except Exception as exc:
|
|||
|
|
LOGGER.exception("WebSocket session failed")
|
|||
|
|
if not ws.closed:
|
|||
|
|
await ws.send_json({"type": "error", "message": str(exc)})
|
|||
|
|
finally:
|
|||
|
|
for task in (processing, speaker_processing):
|
|||
|
|
if task is not None and not task.done():
|
|||
|
|
task.cancel()
|
|||
|
|
pending_tasks = [task for task in (processing, speaker_processing) if task is not None]
|
|||
|
|
if pending_tasks:
|
|||
|
|
await asyncio.gather(*pending_tasks, return_exceptions=True)
|
|||
|
|
# stop、abort、断线和推理异常均释放会话,清理失败不覆盖最终识别结果。
|
|||
|
|
if session is not None:
|
|||
|
|
reset = getattr(session.auxiliary_service, "reset_speaker_session", None)
|
|||
|
|
if reset is not None:
|
|||
|
|
try:
|
|||
|
|
await asyncio.wait_for(reset(session.session_id), timeout=5)
|
|||
|
|
except Exception:
|
|||
|
|
LOGGER.warning("speaker session cleanup failed: %s", session.session_id, exc_info=True)
|
|||
|
|
if owns_model_service:
|
|||
|
|
await model_service.close()
|
|||
|
|
if not ws.closed:
|
|||
|
|
await ws.close()
|
|||
|
|
return ws
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def start_app(model_service_url: str, model: str) -> web.Application:
|
|||
|
|
"""创建 HTTP/WebSocket 应用,并挂载可复用的 VLLM 适配器。"""
|
|||
|
|
app = web.Application()
|
|||
|
|
app[MODEL_SERVICE_KEY] = VLLMTranscriptionService(ModelServiceConfig(base_url=model_service_url, model=model))
|
|||
|
|
app[AUXILIARY_SERVICE_KEY] = AuxiliaryModelService(
|
|||
|
|
AuxiliaryServiceConfig(base_url=os.getenv("AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010"))
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
async def lifecycle(application: web.Application):
|
|||
|
|
await application[MODEL_SERVICE_KEY].start()
|
|||
|
|
await application[AUXILIARY_SERVICE_KEY].start()
|
|||
|
|
yield
|
|||
|
|
await application[AUXILIARY_SERVICE_KEY].close()
|
|||
|
|
await application[MODEL_SERVICE_KEY].close()
|
|||
|
|
|
|||
|
|
app.cleanup_ctx.append(lifecycle)
|
|||
|
|
app.router.add_get("/", index_handler)
|
|||
|
|
app.router.add_get("/api/config", config_handler)
|
|||
|
|
app.router.add_static("/static/", Path(__file__).parent / "static")
|
|||
|
|
app.router.add_get("/ws", websocket_handler)
|
|||
|
|
# 参考腾讯 Demo 将 static 目录挂载到根路径;页面中的 style.css 和 app.js
|
|||
|
|
# 使用相对地址,必须同时提供根路径静态资源路由,否则浏览器会显示无样式页面。
|
|||
|
|
app.router.add_static("/", Path(__file__).parent / "static", show_index=False)
|
|||
|
|
return app
|
|||
|
|
|
|||
|
|
|
|||
|
|
def main() -> None:
|
|||
|
|
"""启动本地测试页面和 WebSocket 服务。"""
|
|||
|
|
args = parse_args()
|
|||
|
|
logging.basicConfig(level=logging.INFO)
|
|||
|
|
if not args.no_browser:
|
|||
|
|
webbrowser.open(f"http://{WEB_DISPLAY_HOST}:{WEB_PORT}/")
|
|||
|
|
print(f"WebSocket demo: http://{WEB_DISPLAY_HOST}:{WEB_PORT}/", flush=True)
|
|||
|
|
print(f"VLLM service: {args.model_service_url} ({args.model})", flush=True)
|
|||
|
|
web.run_app(start_app(args.model_service_url, args.model), host=WEB_HOST, port=WEB_PORT)
|
|||
|
|
|
|||
|
|
|
|||
|
|
if __name__ == "__main__":
|
|||
|
|
main()
|