"""使用模拟 VLLM 适配器验证本地 WebSocket 流程。""" from __future__ import annotations from types import SimpleNamespace import asyncio import io import math import struct import unittest import wave import os from unittest.mock import AsyncMock, patch from aiohttp import web from aiohttp.test_utils import AioHTTPTestCase from server import ( AUXILIARY_SERVICE_KEY, EOF, MODEL_SERVICE_KEY, RealtimeSession, config_handler, deployment_model_name, validate_model_service_url, websocket_handler, ) class FakeModelService: native_partial_supported = False config = SimpleNamespace(base_url="http://fake/v1", model="fake-model") async def transcribe(self, audio_bytes: bytes, source: str, file_name: str, partial: bool) -> str: return "partial text" if partial else "final text" class FakeAuxiliaryService: """返回固定时间段的聚类服务,用于验证 final 后的同句 speaker 更新。""" def __init__(self) -> None: self.calls = 0 async def health(self) -> dict[str, object]: """模拟辅助模型服务已完成预加载。""" return {"ready": True, "speaker_embedding_ready": True} async def resolve_speaker( self, audio_bytes: bytes, session_id: str, start_time_ms: float, end_time_ms: float, speaker_verified: bool = True, ) -> dict[str, object]: _ = (audio_bytes, session_id, start_time_ms, end_time_ms, speaker_verified) speaker_id = [0, 1, 0][min(self.calls, 2)] self.calls += 1 return { "speaker_id": speaker_id, "speaker_name": f"说话人 {speaker_id + 1}", "speaker_evidence": "fresh", "speaker_confidence": 0.9, "speaker_strategy": "online_embedding_cluster", } async def reset_speaker_session(self, session_id: str) -> None: _ = session_id class UnhealthyAuxiliaryService(FakeAuxiliaryService): """模拟端口可访问但辅助模型尚未就绪的服务。""" async def health(self) -> dict[str, object]: return {"ready": False, "speaker_embedding_ready": False} class WebSocketFlowTests(AioHTTPTestCase): async def collect(self, ws, until="end"): """限定等待时间,回归测试中的队列卡死必须表现为失败。""" events = [] while True: # 不用 asyncio.timeout:部署环境是 Python 3.10,wait_for 行为等价。 event = await asyncio.wait_for(ws.receive_json(), timeout=10) events.append(event) if event["type"] == until: return events def get_app(self) -> web.Application: app = web.Application() app[MODEL_SERVICE_KEY] = FakeModelService() app[AUXILIARY_SERVICE_KEY] = FakeAuxiliaryService() app.router.add_get("/api/config", config_handler) app.router.add_get("/ws", websocket_handler) return app def test_vllm_url_validation(self) -> None: self.assertEqual(validate_model_service_url(" http://asr.local/v1/ "), "http://asr.local/v1") with self.assertRaises(ValueError): validate_model_service_url("asr.local:8000/v1") with self.assertRaises(ValueError): validate_model_service_url("http://user:password@asr.local/v1") def test_deployment_alias_and_served_name_match_vllm(self): """WebSocket 不能把下载别名直接当成 vLLM 公开模型名。""" with patch.dict(os.environ, {"QWEN3_ASR_MODEL": "0.6b"}, clear=True): self.assertEqual(deployment_model_name(), "Qwen/Qwen3-ASR-0.6B") with patch.dict(os.environ, {"QWEN3_ASR_MODEL": "0.6b", "VLLM_SERVED_MODEL_NAME": "custom-asr"}, clear=True): self.assertEqual(deployment_model_name(), "custom-asr") async def test_empty_final_retracts_partial_instead_of_leaving_pending(self): """最终没有识别文本时撤回临时内容,不留下永远等待声纹的行。""" async def transcribe(*args, partial): return "temporary" if partial else "" self.app[MODEL_SERVICE_KEY].transcribe = transcribe ws = await self.client.ws_connect("/ws") await ws.send_json({"type": "start", "partial_interval_ms": 300}) await ws.receive_json() await ws.send_bytes(b"\xe8\x03" * 16000) await ws.send_json({"type": "eof"}) events = await self.collect(ws) self.assertTrue(any(e["type"] == "sentences" for e in events)) self.assertEqual(events[-1]["sentences"], []) self.assertEqual(events[-1]["display_blocks"], []) async def test_compressed_file_is_rejected_for_streaming_validation(self) -> None: ws = await self.client.ws_connect("/ws") await ws.send_json({"type": "start", "source": "file", "file_name": "meeting.mp3"}) error = await ws.receive_json() self.assertEqual(error["type"], "error") self.assertIn("PCM 或 WAV", error["message"]) await ws.close() async def test_short_interruption_does_not_inherit_or_call_embedding(self): """句尾静音不能凑够声纹时长,长段 A 后的短插话独立保持 pending。""" ws = await self.client.ws_connect("/ws") await ws.send_json({"type": "start", "source": "mic"}) await ws.receive_json() await ws.send_bytes(b"\xe8\x03" * 16000 + b"\x00\x00" * 12800 + b"\xe8\x03" * 8000 + b"\x00\x00" * 12800) await ws.send_json({"type": "stop"}) events = await self.collect(ws) final = events[-1] self.assertEqual(self.app[AUXILIARY_SERVICE_KEY].calls, 1) self.assertEqual([s["speaker_id"] for s in final["sentences"]], [0, -1]) self.assertEqual(final["sentences"][1]["speaker_status"], "insufficient_audio") self.assertEqual(final["sentences"][0]["end_time"], 1000) self.assertEqual(len(final["display_blocks"]), 2) async def test_stop_waits_for_slow_speaker_and_includes_final_snapshot(self): """在 end 前必须收到所有声纹结果,不能复现页面原先五秒断开的行为。""" auxiliary = self.app[AUXILIARY_SERVICE_KEY] original = auxiliary.resolve_speaker async def delayed(*args): await asyncio.sleep(5.1) return await original(*args) auxiliary.resolve_speaker = delayed auxiliary.reset_speaker_session = AsyncMock() ws = await self.client.ws_connect("/ws") await ws.send_json({"type": "start"}) start = await ws.receive_json() await ws.send_bytes(b"\xe8\x03" * 16000) await ws.send_json({"type": "stop"}) events = await self.collect(ws) self.assertEqual(events[-1]["sentences"][0]["speaker_id"], 0) self.assertTrue(any(e["type"] == "draining" for e in events)) self.assertTrue(any(e["type"] == "sentences" and e["sentences"][0]["speaker_status"] == "processing" for e in events)) await ws.receive() # 等待服务端执行 finally 并关闭连接 auxiliary.reset_speaker_session.assert_awaited_once_with(start["session_id"]) async def test_speaker_error_is_visible_on_segment_and_asr_finishes(self): """声纹推理失败不能吞掉转写,且每条失败片段要携带诊断原因。""" self.app[AUXILIARY_SERVICE_KEY].resolve_speaker = AsyncMock(side_effect=RuntimeError("embedding model missing")) ws = await self.client.ws_connect("/ws") await ws.send_json({"type": "start"}) await ws.receive_json() await ws.send_bytes(b"\xe8\x03" * 16000) await ws.send_json({"type": "eof"}) events = await self.collect(ws) segment = events[-1]["sentences"][0] self.assertEqual(segment["sentence"], "final text") self.assertEqual(segment["speaker_status"], "service_error") self.assertIn("embedding model missing", segment["speaker_reason"]) self.assertTrue(any(e["type"] == "speaker_warning" for e in events)) async def test_asr_failure_is_reported_before_stop(self): """音频 worker 抛异常时,接收任务应立即报告,不能等到客户端 stop。""" self.app[MODEL_SERVICE_KEY].transcribe = AsyncMock(side_effect=RuntimeError("vllm unavailable")) ws = await self.client.ws_connect("/ws") await ws.send_json({"type": "start", "partial_interval_ms": 300}) await ws.receive_json() await ws.send_bytes(b"\xe8\x03" * 16000) events = await self.collect(ws, until="error") self.assertIn("vllm unavailable", events[-1]["message"]) async def test_unknown_speaker_response_is_diagnosable(self): """旧服务只回标签、缺少 fresh/confidence 时应说明拒绝原因。""" self.app[AUXILIARY_SERVICE_KEY].resolve_speaker = AsyncMock(return_value={"speaker_id": 0}) ws = await self.client.ws_connect("/ws") await ws.send_json({"type": "start"}) await ws.receive_json() await ws.send_bytes(b"\xe8\x03" * 16000) await ws.send_json({"type": "eof"}) events = await self.collect(ws) self.assertEqual(events[-1]["sentences"][0]["speaker_status"], "evidence_rejected") async def test_abort_and_disconnect_release_cluster_state(self): """清理不应只存在于成功 stop 的路径。""" auxiliary = self.app[AUXILIARY_SERVICE_KEY] auxiliary.reset_speaker_session = AsyncMock() for abort in (True, False): ws = await self.client.ws_connect("/ws") await ws.send_json({"type": "start"}) await ws.receive_json() if abort: await ws.send_json({"type": "abort"}) await ws.receive() else: await ws.close() for _ in range(200): if auxiliary.reset_speaker_session.await_count >= 2: break await asyncio.sleep(0.01) self.assertEqual(auxiliary.reset_speaker_session.await_count, 2) async def test_extended_wav_header_is_removed_before_asr(self): """分片 RIFF/JUNK/fmt/data 头不能混入声纹和 ASR 的 PCM 数据。""" output = io.BytesIO() pcm = b"\xe8\x03" * 16000 with wave.open(output, "wb") as wav: wav.setparams((1, 2, 16000, 0, "NONE", "")) wav.writeframes(pcm) original = output.getvalue() junk = b"JUNK\x04\x00\x00\x00test" payload = b"RIFF" + (len(original) - 8 + len(junk)).to_bytes(4, "little") + original[8:12] + junk + original[12:] transcribe = AsyncMock(return_value="wav text") self.app[MODEL_SERVICE_KEY].transcribe = transcribe ws = await self.client.ws_connect("/ws") await ws.send_json({"type": "start", "source": "file", "file_name": "test.wav"}) await ws.receive_json() for offset in range(0, len(payload), 337): await ws.send_bytes(payload[offset:offset + 337]) await ws.send_json({"type": "eof"}) events = await self.collect(ws) self.assertEqual(transcribe.call_args.args[0], pcm) self.assertEqual(events[-1]["sentences"][0]["end_time"], 1000) async def test_incompatible_wav_is_rejected_immediately(self): """非 16kHz 单声道 WAV 不能被误解释为可识别的 PCM16。""" output = io.BytesIO() with wave.open(output, "wb") as wav: wav.setparams((2, 2, 44100, 0, "NONE", "")) wav.writeframes(b"\x00\x00" * 2000) ws = await self.client.ws_connect("/ws") await ws.send_json({"type": "start", "source": "file", "file_name": "bad.wav"}) await ws.receive_json() await ws.send_bytes(output.getvalue()) events = await self.collect(ws, until="error") self.assertIn("16kHz", events[-1]["message"]) async def test_frontend_can_read_default_vllm_config(self) -> None: response = await self.client.get("/api/config") self.assertEqual(response.status, 200) self.assertEqual(await response.json(), {"model_service_url": "http://fake/v1", "model": "fake-model", "speaker_service_url": None}) async def test_partial_and_final_share_one_sentence_id(self) -> None: ws = await self.client.ws_connect("/ws") await ws.send_json( { "type": "start", "source": "mic", "speaker_diarization": 0, "partial_interval_ms": 300, "max_segment_sec": 12, } ) start = await ws.receive_json() self.assertEqual(start["type"], "start") # 使用幅度足够的 PCM 语音帧,静音帧会被新切句器正确忽略。 await ws.send_bytes(b"\xe8\x03" * 16000) messages = [] while True: message = await ws.receive_json() messages.append(message) if message["type"] == "sentences": self.assertEqual(message["sentences"][0]["sentence_id"], 0) if message["sentences"][0]["sentence_type"] == 0: break await ws.send_json({"type": "eof"}) while True: message = await ws.receive_json() messages.append(message) if message["type"] == "end": break sentence_events = [message for message in messages if message["type"] == "sentences"] self.assertGreaterEqual(len(sentence_events), 2) self.assertTrue(all(event["sentences"][0]["sentence_id"] == 0 for event in sentence_events)) self.assertTrue(all(event["sentences"][0]["sentence_type"] == 0 for event in sentence_events[:-1])) self.assertEqual(sentence_events[-1]["sentences"][0]["sentence_type"], 1) self.assertEqual(sentence_events[-1]["sentences"][0]["sentence"], "final text") await ws.close() async def test_speaker_health_failure_is_reported_without_blocking_asr(self) -> None: """辅助模型未就绪时先报告告警,同时保留 ASR 会话能力。""" self.app[AUXILIARY_SERVICE_KEY] = UnhealthyAuxiliaryService() ws = await self.client.ws_connect("/ws") await ws.send_json({"type": "start", "source": "mic", "speaker_diarization": 1}) start = await ws.receive_json() warning = await ws.receive_json() self.assertEqual(start["type"], "start") self.assertFalse(start["speaker_service_health"]["ready"]) self.assertEqual(warning["type"], "speaker_warning") self.assertIn("未就绪", warning["message"]) await ws.close() async def test_vad_split_and_speaker_update(self) -> None: ws = await self.client.ws_connect("/ws") await ws.send_json( { "type": "start", "source": "mic", "speaker_diarization": 1, "partial_interval_ms": 1200, } ) self.assertEqual((await ws.receive_json())["type"], "start") voiced = b"\xe8\x03" * 16000 # 每段一秒有效语音,满足独立声纹长度要求 silence = b"\x00\x00" * 12800 # 0.8 秒静音,触发当前 turn 提交 await ws.send_bytes(voiced + silence + voiced + silence + voiced) await ws.send_json({"type": "eof"}) events: list[dict[str, object]] = [] while True: message = await ws.receive_json() events.append(message) if message["type"] == "end": break sentence_events = [message for message in events if message["type"] == "sentences"] final_sentences = [ message["sentences"][0] for message in sentence_events if message["sentences"][0]["sentence_type"] == 1 ] latest_by_id = {int(item["sentence_id"]): item for item in final_sentences} self.assertGreaterEqual(len(latest_by_id), 2) latest = [latest_by_id[index] for index in sorted(latest_by_id)[-3:]] self.assertEqual([item["speaker_id"] for item in latest], [0, 1, 0]) self.assertTrue(all(item["speaker_evidence"] == "fresh" for item in latest)) await ws.close() def _constant_pcm(ms: int, amplitude: int) -> bytes: """定幅 PCM16;A 用高幅 6000、B 用低幅 1200,两者都高于 VAD 门限。""" return struct.pack(" bytes: return b"\x00\x00" * (16 * ms) def _speaker_frame_counts(audio: bytes) -> tuple[int, int]: """按帧 RMS 分类计数 (A, B):≥3000 视为说话人 A,450~2999 视为 B。""" counts = [0, 0] usable = len(audio) - len(audio) % 640 for start in range(0, usable, 640): frame = audio[start:start + 640] squares = [int.from_bytes(frame[i:i + 2], "little", signed=True) ** 2 for i in range(0, 640, 2)] rms = math.sqrt(sum(squares) / len(squares)) if rms >= 3000: counts[0] += 1 elif rms >= 450: counts[1] += 1 return counts[0], counts[1] class _DummyWs: closed = True class SwitchModelService: """final 文本携带主导说话人和音频毫秒数,便于断言切段边界。""" native_partial_supported = False config = SimpleNamespace(base_url="http://fake/v1", model="fake-model") async def transcribe(self, audio_bytes: bytes, source: str, file_name: str, partial: bool) -> str: if partial: return "partial text" a, b = _speaker_frame_counts(audio_bytes) if a + b == 0: return "" return f"{'A' if a >= b else 'B'} {round(len(audio_bytes) / 32)}ms" class SwitchAuxiliaryService: """用帧能量成分构造二维 embedding:A→[1,0]、B→[0,1],混合按占比归一。""" speaker_embedding_unsupported = False def __init__(self, fail_windows: bool = False) -> None: self.embedding_calls = 0 self.resolves: list[dict[str, object]] = [] self.fail_windows = fail_windows async def speaker_embedding(self, pcm_bytes: bytes) -> list[float] | None: self.embedding_calls += 1 if self.fail_windows: return None a, b = _speaker_frame_counts(pcm_bytes) if a + b == 0: return None norm = math.hypot(a, b) return [a / norm, b / norm] async def resolve_speaker( self, audio_bytes: bytes, session_id: str, start_time_ms: float, end_time_ms: float, speaker_verified: bool = True, ) -> dict[str, object]: _ = session_id a, b = _speaker_frame_counts(audio_bytes) speaker_id = 0 if a >= b else 1 self.resolves.append({ "audio_bytes": len(audio_bytes), "speaker_verified": speaker_verified, "start_time_ms": start_time_ms, "end_time_ms": end_time_ms, }) return { "speaker_id": speaker_id, "speaker_name": f"说话人 {speaker_id + 1}", "speaker_evidence": "fresh", "speaker_confidence": 0.9, "speaker_strategy": "online_embedding_cluster_match", "speaker_status": "confirmed", } async def reset_speaker_session(self, session_id: str) -> None: _ = session_id class SpeakerSwitchSegmentationTests(unittest.IsolatedAsyncioTestCase): """回归诊断出的核心缺陷:A→B 停顿小于静音阈值时必须在窗口边界强制切段。""" async def _run(self, pcm: bytes, auxiliary: SwitchAuxiliaryService) -> tuple[RealtimeSession, list]: session = RealtimeSession(_DummyWs(), SwitchModelService(), auxiliary, {"source": "mic"}) session.audio_queue.put_nowait(pcm) session.audio_queue.put_nowait(EOF) await asyncio.wait_for(session.process_audio(), timeout=15) jobs = [] while not session.speaker_queue.empty(): jobs.append(session.speaker_queue.get_nowait()) for job in jobs: await session._resolve_speaker(job) return session, jobs @staticmethod def _finals(session: RealtimeSession) -> list[dict]: return [s for s in session.assembler.raw_snapshot() if s.get("sentence_type") == 1] async def test_short_gap_handover_splits_at_speaker_change(self): """A 2.2s + 400ms 停顿(低于 800ms 阈值)+ B 2.2s:必须切成两个 final。""" auxiliary = SwitchAuxiliaryService() pcm = _constant_pcm(2200, 6000) + _silence_pcm(400) + _constant_pcm(2200, 1200) + _silence_pcm(1000) session, jobs = await self._run(pcm, auxiliary) finals = self._finals(session) # 旧行为:整段合成一个 final 归 B;新行为:在 400ms 停顿后的起音处 # (2.6s)切开,A 的整句完整保留,B 从自己的第一个音素开始。 self.assertEqual([s["sentence"] for s in finals], ["A 2600ms", "B 2200ms"]) self.assertEqual(finals[0]["commit_reason"], "speaker_change") self.assertEqual(finals[0]["end_time"], 2600) self.assertEqual(finals[1]["start_time"], 2600) self.assertEqual([s["speaker_id"] for s in finals], [0, 1]) # 头段边界可能混入失配窗内容,排队标记为未验证;尾段正常。 self.assertEqual( [(job.speaker_verified, len(job.audio)) for job in jobs], [(False, 83200), (True, 70400)], ) self.assertGreaterEqual(auxiliary.embedding_calls, 3) async def test_same_speaker_continuation_is_not_split(self): """同一说话人 3s 后正常停顿:只出一个 final,窗比对不误伤。""" auxiliary = SwitchAuxiliaryService() pcm = _constant_pcm(3000, 6000) + _silence_pcm(1000) session, jobs = await self._run(pcm, auxiliary) finals = self._finals(session) self.assertEqual([s["sentence"] for s in finals], ["A 3000ms"]) self.assertEqual(finals[0]["commit_reason"], "silence") self.assertTrue(all(job.speaker_verified for job in jobs)) self.assertGreaterEqual(auxiliary.embedding_calls, 2) async def test_window_failures_disable_checks_without_blocking_asr(self): """窗提取连续失败 3 次后停用并告警一次;ASR 照常产出 final。""" auxiliary = SwitchAuxiliaryService(fail_windows=True) pcm = _constant_pcm(2600, 6000) + _silence_pcm(1000) session, jobs = await self._run(pcm, auxiliary) self.assertFalse(session.window_check_available) self.assertEqual(auxiliary.embedding_calls, 3) self.assertTrue(session.speaker_warning_sent) finals = self._finals(session) self.assertEqual([s["sentence"] for s in finals], ["A 2600ms"]) self.assertEqual([job.speaker_verified for job in jobs], [True]) async def test_unsupported_endpoint_disables_after_single_call(self): """旧版辅助服务没有窗端点:探测一次即停用,不反复冲击服务。""" auxiliary = SwitchAuxiliaryService(fail_windows=True) auxiliary.speaker_embedding_unsupported = True pcm = _constant_pcm(2600, 6000) + _silence_pcm(1000) session, _ = await self._run(pcm, auxiliary) self.assertFalse(session.window_check_available) self.assertEqual(auxiliary.embedding_calls, 1) self.assertEqual([s["sentence"] for s in self._finals(session)], ["A 2600ms"])