[LLM/TTS 与对话闭环]:完成语音对话主链路,包含上下文管理、NewAPI 兼容 LLM、分句 TTS、状态机和错误恢复测试
This commit is contained in:
@@ -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 与生图资产
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
@@ -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"))
|
||||
@@ -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)
|
||||
@@ -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("<h", sample))
|
||||
return AudioSegment(bytes(pcm), self.sample_rate, 1, 0, duration_ms, {"text": clean})
|
||||
|
||||
|
||||
class MacSayTtsProvider:
|
||||
def __init__(self, voice: str | None = None) -> 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"})
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user