Files
Owner/tests/test_pipeline_llm_tts.py
T

205 lines
7.9 KiB
Python

from __future__ import annotations
import io
import json
import unittest
from owner_voice_pet.config import AppConfig
from owner_voice_pet.conversation import ConversationContext
from owner_voice_pet.llm import MockLlmProvider, OpenAICompatibleLlmProvider
from owner_voice_pet.models import AudioFrame, ErrorCode, Message, PipelineState, ProviderError
from owner_voice_pet.pipeline import VoicePipeline
from owner_voice_pet.stt import MetadataSttProvider
from owner_voice_pet.transport import MemoryAudioTransport
from owner_voice_pet.tts import CloudTtsProvider, SentenceBuffer, SineTtsProvider
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder
from owner_voice_pet.wakeword import KeywordWakeWordProvider
def frame(idx: int, timestamp_ms: int, metadata: dict[str, object]) -> AudioFrame:
speech = bool(metadata.get("speech"))
return AudioFrame(
b"\xff\xff" if speech else b"\x80\x80",
16000,
1,
timestamp_ms,
idx,
{"duration_ms": 20, **metadata},
)
def make_pipeline(frames: list[AudioFrame], llm: MockLlmProvider | None = None) -> tuple[VoicePipeline, MemoryAudioTransport]:
transport = MemoryAudioTransport(frames)
pipeline = VoicePipeline(
transport=transport,
wakeword=KeywordWakeWordProvider(),
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
stt=MetadataSttProvider(),
context=ConversationContext(max_messages=4, max_chars=200),
llm=llm or MockLlmProvider(["你好,我在。"]),
tts=SineTtsProvider(),
)
pipeline.load()
return pipeline, transport
class PipelineLlmTtsTests(unittest.TestCase):
def test_context_truncates_old_messages(self) -> None:
context = ConversationContext(max_messages=2, max_chars=100)
context.append_user("")
context.append_assistant("")
context.append_user("")
self.assertEqual([m.content for m in context.messages()], ["", ""])
self.assertEqual(context.build_llm_messages()[0].role, "system")
def test_sentence_buffer_chunks_on_chinese_punctuation(self) -> None:
buffer = SentenceBuffer(max_chars=20)
self.assertEqual(buffer.feed("你好"), [])
self.assertEqual(buffer.feed(""), ["你好。"])
self.assertEqual(buffer.feed("剩余").copy(), [])
self.assertEqual(buffer.flush(), ["剩余"])
def test_sine_tts_generates_non_empty_audio(self) -> None:
provider = SineTtsProvider()
provider.load()
segment = provider.synthesize("你好")
self.assertGreater(len(segment.pcm), 0)
self.assertGreater(segment.duration_ms, 0)
def test_cloud_tts_posts_speech_request(self) -> None:
class FakeResponse:
def __enter__(self):
return self
def __exit__(self, *args) -> None:
return None
def read(self) -> bytes:
return b"fake-mp3"
requests = []
def fake_urlopen(request, timeout):
requests.append(request)
return FakeResponse()
provider = CloudTtsProvider(
AppConfig(llm_api_key="secret", tts_model="mimo-v2.5-tts"),
urlopen=fake_urlopen,
)
provider.load()
segment = provider.synthesize("你好")
self.assertEqual(segment.metadata["format"], "mp3")
self.assertEqual(segment.pcm, b"fake-mp3")
body = json.loads(requests[0].data.decode())
self.assertEqual(body["model"], "mimo-v2.5-tts")
self.assertIn("/v1/audio/speech", requests[0].full_url)
def test_pipeline_runs_from_wake_to_playback(self) -> None:
frames = [
frame(0, 0, {"wake_word": "小杰小杰", "wake_confidence": 0.95}),
frame(1, 20, {"speech": True, "transcript": "你好"}),
frame(2, 40, {"speech": True}),
frame(3, 60, {"speech": False}),
frame(4, 80, {"speech": False}),
]
pipeline, transport = make_pipeline(frames)
result = pipeline.run_once()
self.assertTrue(result.success)
self.assertEqual(result.transcript, "你好")
self.assertIn(PipelineState.SPEAKING, result.states)
self.assertEqual(len(transport.played_segments), 1)
self.assertEqual(pipeline.context.messages()[-1].role, "assistant")
def test_pipeline_skips_llm_on_empty_transcript(self) -> None:
llm = MockLlmProvider(["不应调用"])
frames = [
frame(0, 0, {"wake": True}),
frame(1, 20, {"speech": True, "transcript": "?!"}),
frame(2, 40, {"speech": True}),
frame(3, 60, {"speech": False}),
frame(4, 80, {"speech": False}),
]
pipeline, _ = make_pipeline(frames, llm)
result = pipeline.run_once()
self.assertFalse(result.success)
self.assertEqual(result.error.code, ErrorCode.STT_EMPTY_TRANSCRIPT)
self.assertEqual(llm.calls, [])
def test_pipeline_recovers_from_llm_failure(self) -> None:
error = ProviderError(ErrorCode.LLM_NETWORK_ERROR, "boom", True, "mock", "llm")
frames = [
frame(0, 0, {"wake": True}),
frame(1, 20, {"speech": True, "transcript": "你好"}),
frame(2, 40, {"speech": True}),
frame(3, 60, {"speech": False}),
frame(4, 80, {"speech": False}),
]
pipeline, _ = make_pipeline(frames, MockLlmProvider(fail=error))
result = pipeline.run_once()
self.assertFalse(result.success)
self.assertIn(PipelineState.ERROR_RECOVERING, result.states)
self.assertEqual(result.error.code, ErrorCode.LLM_NETWORK_ERROR)
def test_openai_chat_completion_sse_parser(self) -> None:
class FakeResponse:
def __enter__(self) -> "FakeResponse":
return self
def __exit__(self, *args: object) -> None:
return None
def read(self):
event = {"choices": [{"delta": {"content": "你好。"}, "finish_reason": None}]}
done = {"choices": [{"delta": {}, "finish_reason": "stop"}]}
return (
f"data: {json.dumps(event, ensure_ascii=False)}\n"
f"data: {json.dumps(done, ensure_ascii=False)}\n"
"data: [DONE]\n"
).encode()
requests = []
def fake_urlopen(request, timeout):
requests.append(request)
return FakeResponse()
config = AppConfig(
llm_base_url="https://token-plan-cn.xiaomimimo.com/v1",
llm_api_key="secret",
llm_model="test-model",
llm_api_style="chat_completions",
)
provider = OpenAICompatibleLlmProvider(config, urlopen=fake_urlopen)
deltas = list(provider.stream_reply([Message("user", "你好", 1.0)]))
self.assertEqual(deltas[0].text_delta, "你好。")
body = json.loads(requests[0].data.decode())
self.assertTrue(body["stream"])
self.assertEqual(body["model"], "test-model")
def test_openai_non_stream_json_parser(self) -> None:
class FakeResponse:
def __enter__(self) -> "FakeResponse":
return self
def __exit__(self, *args: object) -> None:
return None
def read(self):
return json.dumps({"choices": [{"message": {"content": "非流式回复。"}}]}).encode()
config = AppConfig(
llm_base_url="https://token-plan-cn.xiaomimimo.com/v1",
llm_api_key="secret",
llm_model="test-model",
llm_api_style="chat_completions",
llm_stream=False,
)
provider = OpenAICompatibleLlmProvider(config, urlopen=lambda request, timeout: FakeResponse())
self.assertEqual(list(provider.stream_reply([Message("user", "你好", 1.0)]))[0].text_delta, "非流式回复。")
if __name__ == "__main__":
unittest.main()