ASR-demo/realtime_asr_optimization_demo/tests/test_server.py

389 lines
17 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""使用模拟 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()