From e544a60eb3fdf129b581dcfb04dc6016939960ad Mon Sep 17 00:00:00 2001 From: mkbk Date: Wed, 17 Jun 2026 18:21:56 +0800 Subject: [PATCH] =?UTF-8?q?[LLM/TTS=20=E4=B8=8E=E5=AF=B9=E8=AF=9D=E9=97=AD?= =?UTF-8?q?=E7=8E=AF]=EF=BC=9A=E5=AE=8C=E6=88=90=E8=AF=AD=E9=9F=B3?= =?UTF-8?q?=E5=AF=B9=E8=AF=9D=E4=B8=BB=E9=93=BE=E8=B7=AF=EF=BC=8C=E5=8C=85?= =?UTF-8?q?=E5=90=AB=E4=B8=8A=E4=B8=8B=E6=96=87=E7=AE=A1=E7=90=86=E3=80=81?= =?UTF-8?q?NewAPI=20=E5=85=BC=E5=AE=B9=20LLM=E3=80=81=E5=88=86=E5=8F=A5=20?= =?UTF-8?q?TTS=E3=80=81=E7=8A=B6=E6=80=81=E6=9C=BA=E5=92=8C=E9=94=99?= =?UTF-8?q?=E8=AF=AF=E6=81=A2=E5=A4=8D=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../changes/add-voice-pet-pipeline/tasks.md | 12 +- src/owner_voice_pet/__init__.py | 12 ++ src/owner_voice_pet/conversation.py | 48 ++++++ src/owner_voice_pet/llm.py | 149 +++++++++++++++++ src/owner_voice_pet/pipeline.py | 147 +++++++++++++++++ src/owner_voice_pet/tts.py | 123 ++++++++++++++ tests/test_pipeline_llm_tts.py | 151 ++++++++++++++++++ 7 files changed, 636 insertions(+), 6 deletions(-) create mode 100644 src/owner_voice_pet/conversation.py create mode 100644 src/owner_voice_pet/llm.py create mode 100644 src/owner_voice_pet/pipeline.py create mode 100644 src/owner_voice_pet/tts.py create mode 100644 tests/test_pipeline_llm_tts.py diff --git a/openspec/changes/add-voice-pet-pipeline/tasks.md b/openspec/changes/add-voice-pet-pipeline/tasks.md index b18a90a..e16cb51 100644 --- a/openspec/changes/add-voice-pet-pipeline/tasks.md +++ b/openspec/changes/add-voice-pet-pipeline/tasks.md @@ -26,12 +26,12 @@ ## 4. LLM/TTS 与对话闭环 -- [ ] 4.1 实现 ConversationContext;前置条件:Message 模型已定义;验收标准:追加 user/assistant、构造 LLM messages、按预算截断;测试要点:上下文保留 system prompt 和最近轮次;优先级:P0;预计:45 分钟。 -- [ ] 4.2 实现 OpenAI/NewAPI 兼容 LLM Provider;前置条件:配置加载已完成;验收标准:base URL、API key、model 均配置化,支持流式或非流式解析,不记录密钥;测试要点:mock HTTP 和可选真实 smoke 测试通过;优先级:P0;预计:60 分钟。 -- [ ] 4.3 实现分句缓冲;前置条件:LLM delta 模型已定义;验收标准:中文标点、换行和长度阈值触发 TTS chunk;测试要点:空分句、超长句、连续 delta 测试通过;优先级:P0;预计:45 分钟。 -- [ ] 4.4 实现本地 TTS Provider;前置条件:AudioSegment 模型和 Transport 可用;验收标准:提供 macOS `say` 可选 Provider 和 deterministic 测试 Provider;测试要点:非空文本生成非空音频,空文本报错;优先级:P0;预计:60 分钟。 -- [ ] 4.5 实现 Pipeline 状态机和自抑制;前置条件:Wake/VAD/STT、Context、LLM、TTS、Transport 均可测试;验收标准:正常链路、空转写、LLM 失败、TTS 失败、播放期自抑制均可验证;测试要点:状态序列和错误恢复测试通过;优先级:P0;预计:60 分钟。 -- [ ] 4.6 完成“LLM/TTS 与对话闭环”模块提交;前置条件:4.1 至 4.5 已完成;验收标准:先通过相关测试、compileall、真实或 mock LLM smoke 和 OpenSpec strict 校验,再立即执行 Git commit;测试要点:提交信息使用“`[LLM/TTS 与对话闭环]:完成[具体功能描述],包含[关键变更]`”格式;优先级:P0;预计:20 分钟。 +- [x] 4.1 实现 ConversationContext;前置条件:Message 模型已定义;验收标准:追加 user/assistant、构造 LLM messages、按预算截断;测试要点:上下文保留 system prompt 和最近轮次;优先级:P0;预计:45 分钟。 +- [x] 4.2 实现 OpenAI/NewAPI 兼容 LLM Provider;前置条件:配置加载已完成;验收标准:base URL、API key、model 均配置化,支持流式或非流式解析,不记录密钥;测试要点:mock HTTP 和可选真实 smoke 测试通过;优先级:P0;预计:60 分钟。 +- [x] 4.3 实现分句缓冲;前置条件:LLM delta 模型已定义;验收标准:中文标点、换行和长度阈值触发 TTS chunk;测试要点:空分句、超长句、连续 delta 测试通过;优先级:P0;预计:45 分钟。 +- [x] 4.4 实现本地 TTS Provider;前置条件:AudioSegment 模型和 Transport 可用;验收标准:提供 macOS `say` 可选 Provider 和 deterministic 测试 Provider;测试要点:非空文本生成非空音频,空文本报错;优先级:P0;预计:60 分钟。 +- [x] 4.5 实现 Pipeline 状态机和自抑制;前置条件:Wake/VAD/STT、Context、LLM、TTS、Transport 均可测试;验收标准:正常链路、空转写、LLM 失败、TTS 失败、播放期自抑制均可验证;测试要点:状态序列和错误恢复测试通过;优先级:P0;预计:60 分钟。 +- [x] 4.6 完成“LLM/TTS 与对话闭环”模块提交;前置条件:4.1 至 4.5 已完成;验收标准:先通过相关测试、compileall、真实或 mock LLM smoke 和 OpenSpec strict 校验,再立即执行 Git commit;测试要点:提交信息使用“`[LLM/TTS 与对话闭环]:完成[具体功能描述],包含[关键变更]`”格式;优先级:P0;预计:20 分钟。 ## 5. 桌宠 UI 与生图资产 diff --git a/src/owner_voice_pet/__init__.py b/src/owner_voice_pet/__init__.py index a8617f6..1f5cd7c 100644 --- a/src/owner_voice_pet/__init__.py +++ b/src/owner_voice_pet/__init__.py @@ -19,6 +19,10 @@ from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTranspor from .wakeword import KeywordWakeWordProvider from .vad import EnergyVadProvider, VadRecorder from .stt import MetadataSttProvider, is_valid_transcript_text +from .conversation import ConversationContext +from .llm import MockLlmProvider, OpenAICompatibleLlmProvider +from .pipeline import PipelineResult, VoicePipeline +from .tts import MacSayTtsProvider, SentenceBuffer, SineTtsProvider __all__ = [ "AppConfig", @@ -33,6 +37,14 @@ __all__ = [ "VadRecorder", "MetadataSttProvider", "is_valid_transcript_text", + "ConversationContext", + "MockLlmProvider", + "OpenAICompatibleLlmProvider", + "PipelineResult", + "VoicePipeline", + "MacSayTtsProvider", + "SentenceBuffer", + "SineTtsProvider", "Message", "PipelineState", "PlaybackResult", diff --git a/src/owner_voice_pet/conversation.py b/src/owner_voice_pet/conversation.py new file mode 100644 index 0000000..c557994 --- /dev/null +++ b/src/owner_voice_pet/conversation.py @@ -0,0 +1,48 @@ +from __future__ import annotations + +import time +from dataclasses import dataclass, field +from typing import Iterable + +from .models import Message + + +@dataclass(slots=True) +class ConversationContext: + system_prompt: str = "你是一个中文桌宠助手,回答要简洁、自然、适合语音播报。" + max_messages: int = 12 + max_chars: int = 12000 + _messages: list[Message] = field(default_factory=list) + + def append_user(self, text: str) -> None: + self._append("user", text) + + def append_assistant(self, text: str) -> None: + self._append("assistant", text) + + def build_llm_messages(self) -> list[Message]: + system = Message("system", self.system_prompt, 0.0) + return [system, *self._messages] + + def reset(self) -> None: + self._messages.clear() + + def messages(self) -> tuple[Message, ...]: + return tuple(self._messages) + + def _append(self, role: str, text: str) -> None: + clean = text.strip() + if not clean: + return + self._messages.append(Message(role, clean, time.time())) + self.truncate() + + def truncate(self) -> None: + while len(self._messages) > self.max_messages: + self._messages.pop(0) + while self._total_chars(self._messages) > self.max_chars and self._messages: + self._messages.pop(0) + + @staticmethod + def _total_chars(messages: Iterable[Message]) -> int: + return sum(len(message.content) for message in messages) diff --git a/src/owner_voice_pet/llm.py b/src/owner_voice_pet/llm.py new file mode 100644 index 0000000..e771161 --- /dev/null +++ b/src/owner_voice_pet/llm.py @@ -0,0 +1,149 @@ +from __future__ import annotations + +import json +import socket +import urllib.error +import urllib.request +from collections.abc import Callable, Iterable, Sequence +from typing import Any + +from .config import AppConfig +from .models import ErrorCode, Message, ProviderError, ReplyDelta + + +class MockLlmProvider: + def __init__(self, chunks: list[str] | None = None, fail: ProviderError | None = None) -> None: + self.chunks = chunks or ["好的,我听到了。"] + self.fail = fail + self.calls: list[list[Message]] = [] + + def stream_reply(self, messages: Sequence[Message]) -> Iterable[ReplyDelta]: + self.calls.append(list(messages)) + if self.fail: + raise self.fail + for chunk in self.chunks: + yield ReplyDelta(chunk, is_sentence_boundary=_ends_sentence(chunk)) + yield ReplyDelta("", finish_reason="stop") + + +class OpenAICompatibleLlmProvider: + def __init__( + self, + config: AppConfig, + timeout_s: float = 20.0, + urlopen: Callable[..., Any] | None = None, + ) -> None: + self.config = config + self.timeout_s = timeout_s + self.urlopen = urlopen or urllib.request.urlopen + + def stream_reply(self, messages: Sequence[Message]) -> Iterable[ReplyDelta]: + self.config.require_llm_credentials() + if self.config.llm_api_style == "responses": + yield from self._responses(messages) + else: + yield from self._chat_completions(messages) + + def _chat_completions(self, messages: Sequence[Message]) -> Iterable[ReplyDelta]: + payload = { + "model": self.config.llm_model, + "messages": [{"role": m.role, "content": m.content} for m in messages], + "stream": True, + } + yield from self._post_stream( + f"{self.config.llm_base_url}/v1/chat/completions", + payload, + parser=_parse_chat_completion_sse, + ) + + def _responses(self, messages: Sequence[Message]) -> Iterable[ReplyDelta]: + payload = { + "model": self.config.llm_model, + "input": [{"role": m.role, "content": m.content} for m in messages], + "stream": True, + } + yield from self._post_stream( + f"{self.config.llm_base_url}/v1/responses", + payload, + parser=_parse_responses_sse, + ) + + def _post_stream( + self, + url: str, + payload: dict[str, Any], + parser: Callable[[dict[str, Any]], ReplyDelta | None], + ) -> Iterable[ReplyDelta]: + body = json.dumps(payload).encode("utf-8") + request = urllib.request.Request( + url, + data=body, + headers={ + "Authorization": f"Bearer {self.config.llm_api_key}", + "Content-Type": "application/json", + }, + method="POST", + ) + try: + with self.urlopen(request, timeout=self.timeout_s) as response: + emitted = False + for raw_line in response: + line = raw_line.decode("utf-8", errors="replace").strip() + if not line or not line.startswith("data:"): + continue + data = line.removeprefix("data:").strip() + if data == "[DONE]": + yield ReplyDelta("", finish_reason="stop") + return + event = json.loads(data) + delta = parser(event) + if delta is not None: + emitted = True + yield delta + if not emitted: + raise ProviderError( + ErrorCode.LLM_EMPTY_REPLY, + "LLM stream ended without text", + True, + "openai-compatible", + "llm", + ) + except ProviderError: + raise + except urllib.error.HTTPError as exc: + code = ErrorCode.LLM_RATE_LIMITED if exc.code == 429 else ErrorCode.LLM_NETWORK_ERROR + raise ProviderError(code, f"LLM HTTP error {exc.code}", exc.code >= 500, "openai-compatible", "llm") from exc + except (urllib.error.URLError, TimeoutError, socket.timeout) as exc: + raise ProviderError( + ErrorCode.LLM_REQUEST_TIMEOUT, + f"LLM request failed or timed out: {exc}", + True, + "openai-compatible", + "llm", + ) from exc + + +def _parse_chat_completion_sse(event: dict[str, Any]) -> ReplyDelta | None: + choices = event.get("choices") or [] + if not choices: + return None + choice = choices[0] + text = str((choice.get("delta") or {}).get("content") or "") + finish = choice.get("finish_reason") + if not text and not finish: + return None + return ReplyDelta(text, is_sentence_boundary=_ends_sentence(text), finish_reason=finish) + + +def _parse_responses_sse(event: dict[str, Any]) -> ReplyDelta | None: + event_type = event.get("type") + if event_type == "response.output_text.delta": + text = str(event.get("delta") or "") + return ReplyDelta(text, is_sentence_boundary=_ends_sentence(text)) + if event_type in {"response.completed", "response.output_text.done"}: + return ReplyDelta("", finish_reason="stop") + return None + + +def _ends_sentence(text: str) -> bool: + return text.endswith(("。", "!", "?", ".", "!", "?", "\n")) diff --git a/src/owner_voice_pet/pipeline.py b/src/owner_voice_pet/pipeline.py new file mode 100644 index 0000000..c70775c --- /dev/null +++ b/src/owner_voice_pet/pipeline.py @@ -0,0 +1,147 @@ +from __future__ import annotations + +from dataclasses import dataclass, field + +from .conversation import ConversationContext +from .models import AudioSegment, ErrorCode, PipelineState, ProviderError +from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, WakeWordProvider +from .tts import SentenceBuffer +from .vad import VadRecorder + + +@dataclass(slots=True) +class PipelineResult: + success: bool + transcript: str = "" + assistant_text: str = "" + states: list[PipelineState] = field(default_factory=list) + error: ProviderError | None = None + played_segments: int = 0 + + +class VoicePipeline: + def __init__( + self, + transport: AudioTransport, + wakeword: WakeWordProvider, + vad_recorder: VadRecorder, + stt: SttProvider, + context: ConversationContext, + llm: LlmProvider, + tts: TtsProvider, + sentence_buffer: SentenceBuffer | None = None, + ) -> None: + self.transport = transport + self.wakeword = wakeword + self.vad_recorder = vad_recorder + self.stt = stt + self.context = context + self.llm = llm + self.tts = tts + self.sentence_buffer = sentence_buffer or SentenceBuffer() + self.states: list[PipelineState] = [] + self.suppress_input = False + + def load(self) -> None: + self.wakeword.load() + self.vad_recorder.provider.load() + self.stt.load() + self.tts.load() + + def run_once(self, max_frames: int = 1000) -> PipelineResult: + self.states = [] + try: + self._state(PipelineState.WAKE_LISTENING) + self.transport.start_input() + wake_found = False + frames_read = 0 + while frames_read < max_frames: + frames = self.transport.read_frames(timeout_ms=20) + if not frames: + break + for frame in frames: + frames_read += 1 + if self.suppress_input: + continue + if not wake_found: + if self.wakeword.detect(frame): + wake_found = True + self._state(PipelineState.SPEECH_DETECTING) + continue + result = self.vad_recorder.feed(frame) + if isinstance(result, ProviderError): + return self._recover(result) + if isinstance(result, AudioSegment): + return self._handle_segment(result) + return self._recover( + ProviderError( + ErrorCode.VAD_TIMEOUT_NO_SPEECH, + "no complete utterance was captured", + True, + "voice-pipeline", + "pipeline", + ) + ) + finally: + self.transport.stop() + + def _handle_segment(self, segment: AudioSegment) -> PipelineResult: + self._state(PipelineState.TRANSCRIBING) + try: + transcript = self.stt.transcribe(segment) + except ProviderError as exc: + if exc.code == ErrorCode.STT_EMPTY_TRANSCRIPT: + self._state(PipelineState.WAKE_LISTENING) + return PipelineResult(False, states=list(self.states), error=exc) + return self._recover(exc) + text = transcript.normalized_text + self.context.append_user(text) + self._state(PipelineState.THINKING) + assistant_text = "" + played = 0 + try: + for delta in self.llm.stream_reply(self.context.build_llm_messages()): + assistant_text += delta.text_delta + for sentence in self.sentence_buffer.feed(delta.text_delta, bool(delta.finish_reason)): + self._state(PipelineState.SPEAKING) + self.suppress_input = True + audio = self.tts.synthesize(sentence) + playback = self.transport.play_pcm(audio) + self.suppress_input = False + if playback.error: + return self._recover(playback.error) + played += 1 + for sentence in self.sentence_buffer.flush(): + self._state(PipelineState.SPEAKING) + self.suppress_input = True + audio = self.tts.synthesize(sentence) + playback = self.transport.play_pcm(audio) + self.suppress_input = False + if playback.error: + return self._recover(playback.error) + played += 1 + except ProviderError as exc: + self.suppress_input = False + return self._recover(exc) + if not assistant_text.strip(): + return self._recover( + ProviderError( + ErrorCode.LLM_EMPTY_REPLY, + "LLM returned no assistant text", + True, + "voice-pipeline", + "llm", + ) + ) + self.context.append_assistant(assistant_text) + self._state(PipelineState.WAKE_LISTENING) + return PipelineResult(True, text, assistant_text, list(self.states), played_segments=played) + + def _recover(self, error: ProviderError) -> PipelineResult: + self._state(PipelineState.ERROR_RECOVERING) + self.suppress_input = False + self._state(PipelineState.WAKE_LISTENING) + return PipelineResult(False, states=list(self.states), error=error) + + def _state(self, state: PipelineState) -> None: + self.states.append(state) diff --git a/src/owner_voice_pet/tts.py b/src/owner_voice_pet/tts.py new file mode 100644 index 0000000..f8b6146 --- /dev/null +++ b/src/owner_voice_pet/tts.py @@ -0,0 +1,123 @@ +from __future__ import annotations + +import math +import struct +import subprocess +import tempfile +from pathlib import Path + +from .models import AudioSegment, ErrorCode, ProviderError + + +class SentenceBuffer: + def __init__(self, max_chars: int = 80) -> None: + self.max_chars = max_chars + self._buffer = "" + + def feed(self, text: str, force: bool = False) -> list[str]: + if text: + self._buffer += text + chunks: list[str] = [] + while self._should_emit(force): + chunks.append(self._pop_chunk(force)) + force = False + return [chunk for chunk in chunks if chunk.strip()] + + def flush(self) -> list[str]: + return self.feed("", force=True) + + def _should_emit(self, force: bool) -> bool: + stripped = self._buffer.strip() + if not stripped: + return False + return force or stripped.endswith(("。", "!", "?", ".", "!", "?", "\n")) or len(stripped) >= self.max_chars + + def _pop_chunk(self, force: bool) -> str: + stripped = self._buffer.strip() + if force or len(stripped) <= self.max_chars: + self._buffer = "" + return stripped + chunk = stripped[: self.max_chars] + self._buffer = stripped[self.max_chars :] + return chunk + + +class SineTtsProvider: + def __init__(self, sample_rate: int = 16000) -> None: + self.sample_rate = sample_rate + self.loaded = False + + def load(self) -> None: + self.loaded = True + + def synthesize(self, text: str) -> AudioSegment: + if not self.loaded: + raise ProviderError( + ErrorCode.TTS_SYNTHESIS_FAILED, + "TTS provider is not loaded", + False, + "sine-tts", + "tts", + ) + clean = text.strip() + if not clean: + raise ProviderError( + ErrorCode.TTS_EMPTY_AUDIO, + "cannot synthesize empty text", + True, + "sine-tts", + "tts", + ) + duration_ms = max(120, min(1200, len(clean) * 45)) + samples = int(self.sample_rate * duration_ms / 1000) + pcm = bytearray() + for idx in range(samples): + sample = int(math.sin(2 * math.pi * 440 * idx / self.sample_rate) * 12000) + pcm.extend(struct.pack(" None: + self.voice = voice + self.loaded = False + + def load(self) -> None: + self.loaded = True + + def synthesize(self, text: str) -> AudioSegment: + clean = text.strip() + if not clean: + raise ProviderError( + ErrorCode.TTS_EMPTY_AUDIO, + "cannot synthesize empty text", + True, + "macos-say", + "tts", + ) + with tempfile.TemporaryDirectory() as tmp: + output = Path(tmp) / "speech.aiff" + command = ["say", "-o", str(output)] + if self.voice: + command.extend(["-v", self.voice]) + command.append(clean) + try: + subprocess.run(command, check=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE) + except (FileNotFoundError, subprocess.CalledProcessError) as exc: + raise ProviderError( + ErrorCode.TTS_SYNTHESIS_FAILED, + f"macOS say failed: {exc}", + False, + "macos-say", + "tts", + ) from exc + data = output.read_bytes() + if not data: + raise ProviderError( + ErrorCode.TTS_EMPTY_AUDIO, + "macOS say produced empty audio", + True, + "macos-say", + "tts", + ) + return AudioSegment(data, 16000, 1, 0, max(120, len(clean) * 45), {"text": clean, "format": "aiff"}) diff --git a/tests/test_pipeline_llm_tts.py b/tests/test_pipeline_llm_tts.py new file mode 100644 index 0000000..af3ec0a --- /dev/null +++ b/tests/test_pipeline_llm_tts.py @@ -0,0 +1,151 @@ +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 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_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 __iter__(self): + event = {"choices": [{"delta": {"content": "你好。"}, "finish_reason": None}]} + done = {"choices": [{"delta": {}, "finish_reason": "stop"}]} + yield f"data: {json.dumps(event, ensure_ascii=False)}\n".encode() + yield f"data: {json.dumps(done, ensure_ascii=False)}\n".encode() + yield b"data: [DONE]\n" + + requests = [] + + def fake_urlopen(request, timeout): + requests.append(request) + return FakeResponse() + + config = AppConfig( + llm_base_url="https://newapi.mkbk.shop", + 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") + + +if __name__ == "__main__": + unittest.main()