219 lines
8.5 KiB
Python
219 lines
8.5 KiB
Python
from __future__ import annotations
|
|
|
|
import io
|
|
import base64
|
|
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 json.dumps(
|
|
{
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"audio": {"data": base64.b64encode(b"fake-wav").decode("ascii")}
|
|
}
|
|
}
|
|
]
|
|
}
|
|
).encode()
|
|
|
|
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"], "wav")
|
|
self.assertEqual(segment.pcm, b"fake-wav")
|
|
body = json.loads(requests[0].data.decode())
|
|
self.assertEqual(body["model"], "mimo-v2.5-tts")
|
|
self.assertEqual(body["messages"][0]["role"], "assistant")
|
|
self.assertEqual(body["messages"][0]["content"], "你好")
|
|
self.assertEqual(body["audio"]["format"], "wav")
|
|
self.assertIn("/v1/chat/completions", 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()
|