ASR-demo/realtime_asr_optimization_demo/tests/test_server.py

389 lines
17 KiB
Python
Raw Normal View History

2026-09-10 05:47:09 +00:00
"""使用模拟 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)
2026-09-10 08:00:36 +00:00
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()
2026-09-10 05:47:09 +00:00
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"])
2026-09-10 08:00:36 +00:00
self.assertFalse(start["speaker_gap_enabled"])
2026-09-10 05:47:09 +00:00
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()