[LLM/TTS 与对话闭环]:完成语音对话主链路,包含上下文管理、NewAPI 兼容 LLM、分句 TTS、状态机和错误恢复测试
This commit is contained in:
@@ -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"})
|
||||
Reference in New Issue
Block a user