"""使用模拟 VLLM 适配器验证本地 WebSocket 流程。""" from __future__ import annotations from types import SimpleNamespace import asyncio import io 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, MODEL_SERVICE_KEY, 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, ) -> dict[str, object]: _ = (audio_bytes, session_id, start_time_ms, end_time_ms) 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 = [] async with asyncio.timeout(10): while True: event = await ws.receive_json() 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_short_speaker_gap_splits_turn_before_next_speaker(self): """说话人模式在短交接停顿处切段,保持每段独立进入在线聚类。""" ws = await self.client.ws_connect("/ws") await ws.send_json( { "type": "start", "source": "mic", "speaker_diarization": 1, "speaker_gap_ms": 400, "partial_interval_ms": 1200, } ) self.assertEqual((await ws.receive_json())["type"], "start") voiced = b"\xe8\x03" * 16000 # 每位说话人一秒有效语音 handoff_gap = b"\x00\x00" * 6400 # 400ms,低于原始 800ms 切段阈值 trailing_silence = b"\x00\x00" * 12800 await ws.send_bytes(voiced + handoff_gap + voiced + trailing_silence) 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 final_sentences = [ message["sentences"][0] for message in events if message["type"] == "sentences" and message["sentences"][0]["sentence_type"] == 1 ] latest_by_id = {int(item["sentence_id"]): item for item in final_sentences} latest = [latest_by_id[index] for index in sorted(latest_by_id)] self.assertEqual(len(latest), 2) self.assertEqual([item["speaker_id"] for item in latest], [0, 1]) self.assertEqual(self.app[AUXILIARY_SERVICE_KEY].calls, 2) await ws.close() 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() async with asyncio.timeout(2): while auxiliary.reset_speaker_session.await_count < 2: 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.assertFalse(start["speaker_gap_enabled"]) 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()