"""面向浏览器的独立 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 # 说话人分离开启时,用比普通 VAD 更短的静音作为候选话轮边界。 # 这个边界只负责把 A→B 的短交接停顿拆开,不调用声纹模型,因此不会 # 把窗口推理延迟叠加到音频帧处理路径;同一说话人的拆分片段仍由聚类合并。 SPEAKER_GAP_MS = max(20, int(os.getenv("SPEAKER_GAP_MS", "400"))) 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.speaker_gap_enabled = self.speaker_enabled and auxiliary_service is not None raw_speaker_gap = start.get("speaker_gap_ms") if raw_speaker_gap is None: self.speaker_gap_ms = SPEAKER_GAP_MS else: try: self.speaker_gap_ms = max(20, int(float(raw_speaker_gap))) except (TypeError, ValueError, OverflowError): self.speaker_gap_ms = SPEAKER_GAP_MS 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) # 短交接停顿优先于普通 800/1400ms 静音切段,但必须先有 # 至少 800ms 有效语音,避免把咳嗽、噪声或极短插话送去聚类。 short_speaker_gap = ( self.speaker_gap_enabled and self.voiced_ms >= MIN_SPEAKER_VOICE_MS and self.silence_ms >= self.speaker_gap_ms ) if short_speaker_gap or self.silence_ms >= self.silence_limit_ms or self._duration_ms() >= self.max_segment_sec * 1000: reason = ( "speaker_gap" if short_speaker_gap else "silence" if self.silence_ms >= self.silence_limit_ms else "max_duration" ) await self._commit_segment(reason) 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}" if speaker_health_error: # 辅助服务未就绪时先保留原始 VAD 切段,避免在没有声纹结果的 # 情况下增加大量短片段;服务恢复后由新的会话重新启用。 session.speaker_gap_enabled = False 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, "speaker_gap_enabled": session.speaker_gap_enabled, "speaker_gap_ms": session.speaker_gap_ms, "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()