102 lines
3.4 KiB
Python
102 lines
3.4 KiB
Python
"""独立音频请求准备逻辑的回归测试。"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import base64
|
|
import json
|
|
from types import SimpleNamespace
|
|
import unittest
|
|
import wave
|
|
from io import BytesIO
|
|
|
|
from model_service import (
|
|
pcm16_to_wav,
|
|
prepare_audio_request,
|
|
realtime_ws_url,
|
|
wav_to_pcm16,
|
|
VLLMRealtimeStream,
|
|
)
|
|
from auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig
|
|
from aiohttp import WSMsgType
|
|
|
|
|
|
|
|
class FakeRealtimeWebSocket:
|
|
def __init__(self) -> None:
|
|
self.messages: asyncio.Queue[object] = asyncio.Queue()
|
|
self.sent: list[dict[str, object]] = []
|
|
self.closed = False
|
|
|
|
def __aiter__(self):
|
|
return self._messages()
|
|
|
|
async def _messages(self):
|
|
while not self.closed:
|
|
yield await self.messages.get()
|
|
|
|
async def send_json(self, payload: dict[str, object]) -> None:
|
|
self.sent.append(payload)
|
|
if payload.get("type") == "input_audio_buffer.commit" and payload.get("final"):
|
|
await self.messages.put(
|
|
SimpleNamespace(
|
|
type=WSMsgType.TEXT,
|
|
data=json.dumps({"type": "transcription.done", "text": "final text"}),
|
|
)
|
|
)
|
|
|
|
def exception(self):
|
|
return None
|
|
|
|
async def close(self) -> None:
|
|
self.closed = True
|
|
|
|
class ModelServiceTests(unittest.TestCase):
|
|
def test_pcm_is_wrapped_as_16k_mono_wav(self) -> None:
|
|
wav_bytes = pcm16_to_wav(b"\x00\x00" * 160)
|
|
with wave.open(BytesIO(wav_bytes), "rb") as wav_file:
|
|
self.assertEqual(wav_file.getframerate(), 16000)
|
|
self.assertEqual(wav_file.getnchannels(), 1)
|
|
self.assertEqual(wav_to_pcm16(wav_bytes), b"\x00\x00" * 160)
|
|
|
|
def test_compressed_partial_is_deferred(self) -> None:
|
|
self.assertIsNone(prepare_audio_request(b"partial", "file", "sample.mp3", partial=True))
|
|
prepared = prepare_audio_request(b"complete", "file", "sample.mp3", partial=False)
|
|
self.assertEqual(prepared[1], "sample.mp3")
|
|
|
|
def test_realtime_url_uses_the_openai_v1_path(self) -> None:
|
|
self.assertEqual(
|
|
realtime_ws_url("https://asr.example/v1/"),
|
|
"wss://asr.example/v1/realtime",
|
|
)
|
|
|
|
def test_realtime_stream_sends_incremental_audio_and_finishes(self) -> None:
|
|
async def exercise() -> None:
|
|
websocket = FakeRealtimeWebSocket()
|
|
stream = VLLMRealtimeStream(websocket, timeout_seconds=1)
|
|
await websocket.messages.put(
|
|
SimpleNamespace(
|
|
type=WSMsgType.TEXT,
|
|
data=json.dumps({"type": "transcription.delta", "delta": "partial"}),
|
|
)
|
|
)
|
|
await stream.append_audio(b"\x01\x02")
|
|
final_text = await stream.finish()
|
|
self.assertEqual(final_text, "final text")
|
|
self.assertEqual(
|
|
base64.b64decode(str(websocket.sent[0]["audio"])),
|
|
b"\x01\x02",
|
|
)
|
|
self.assertEqual(websocket.sent[-1]["type"], "input_audio_buffer.commit")
|
|
await stream.close()
|
|
|
|
asyncio.run(exercise())
|
|
|
|
def test_auxiliary_config_is_independent_from_vllm(self) -> None:
|
|
service = AuxiliaryModelService(AuxiliaryServiceConfig())
|
|
self.assertEqual(service.config.base_url, "http://127.0.0.1:8010")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|