[LLM/TTS 与对话闭环]:完成语音对话主链路,包含上下文管理、NewAPI 兼容 LLM、分句 TTS、状态机和错误恢复测试

This commit is contained in:
mkbk
2026-06-17 18:21:56 +08:00
parent 4d6232ed29
commit e544a60eb3
7 changed files with 636 additions and 6 deletions
+12
View File
@@ -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",
+48
View File
@@ -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)
+149
View File
@@ -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"))
+147
View File
@@ -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)
+123
View File
@@ -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"})