ASR-demo/realtime_asr_optimization_demo/tests/test_server.py

522 lines
23 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
2026-09-10 06:23:18 +00:00
import math
import struct
import unittest
2026-09-10 05:47:09 +00:00
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,
2026-09-10 06:23:18 +00:00
EOF,
2026-09-10 05:47:09 +00:00
MODEL_SERVICE_KEY,
2026-09-10 06:23:18 +00:00
RealtimeSession,
2026-09-10 05:47:09 +00:00
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,
2026-09-10 06:23:18 +00:00
speaker_verified: bool = True,
2026-09-10 05:47:09 +00:00
) -> dict[str, object]:
2026-09-10 06:23:18 +00:00
_ = (audio_bytes, session_id, start_time_ms, end_time_ms, speaker_verified)
2026-09-10 05:47:09 +00:00
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 = []
2026-09-10 06:23:18 +00:00
while True:
# 不用 asyncio.timeout部署环境是 Python 3.10wait_for 行为等价。
event = await asyncio.wait_for(ws.receive_json(), timeout=10)
events.append(event)
if event["type"] == until:
return events
2026-09-10 05:47:09 +00:00
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()
2026-09-10 06:23:18 +00:00
for _ in range(200):
if auxiliary.reset_speaker_session.await_count >= 2:
break
await asyncio.sleep(0.01)
2026-09-10 05:47:09 +00:00
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()
2026-09-10 06:23:18 +00:00
def _constant_pcm(ms: int, amplitude: int) -> bytes:
"""定幅 PCM16A 用高幅 6000、B 用低幅 1200两者都高于 VAD 门限。"""
return struct.pack("<h", amplitude) * (16 * ms)
def _silence_pcm(ms: int) -> bytes:
return b"\x00\x00" * (16 * ms)
def _speaker_frame_counts(audio: bytes) -> tuple[int, int]:
"""按帧 RMS 分类计数 (A, B)≥3000 视为说话人 A450~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:
"""用帧能量成分构造二维 embeddingA→[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"])