352 lines
13 KiB
Python
352 lines
13 KiB
Python
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass, field
|
|
from typing import Protocol
|
|
|
|
from .config import AppConfig
|
|
from .conversation import ConversationContext
|
|
from .events import (
|
|
ACK_STARTED,
|
|
CAPTURE_STARTED,
|
|
LLM_STARTED,
|
|
PLAYBACK_FINISHED,
|
|
QUESTION_PROMPT,
|
|
RECOVERING,
|
|
SPEECH_ENDED,
|
|
SPEECH_STARTED,
|
|
STAGE_ERROR,
|
|
STANDBY_RESUMED,
|
|
STT_STARTED,
|
|
TRANSCRIPT_FINAL,
|
|
TTS_STARTED,
|
|
WAKE_DETECTED,
|
|
WAKE_LISTENING,
|
|
PipelineEvent,
|
|
PipelineEventBus,
|
|
dispatch_pipeline_event,
|
|
)
|
|
from .models import AudioSegment, ErrorCode, PipelineState, ProviderError
|
|
from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, WakeWordProvider
|
|
from .stt import is_valid_transcript_text
|
|
from .tts import SentenceBuffer
|
|
from .vad import VadRecorder
|
|
|
|
|
|
class RuntimeReporter(Protocol):
|
|
def status(self, state: str, message: str, *, turn_id: int | None = None) -> None:
|
|
...
|
|
|
|
def transcript(self, text: str, *, final: bool, turn_id: int | None = None) -> None:
|
|
...
|
|
|
|
def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None:
|
|
...
|
|
|
|
|
|
@dataclass(slots=True)
|
|
class TurnResult:
|
|
success: bool
|
|
transcript: str = ""
|
|
assistant_text: str = ""
|
|
error: ProviderError | None = None
|
|
states: list[PipelineState] = field(default_factory=list)
|
|
|
|
|
|
@dataclass(slots=True)
|
|
class RuntimeSummary:
|
|
completed_turns: int
|
|
failed_turns: int
|
|
interrupted: bool = False
|
|
last_error: ProviderError | None = None
|
|
|
|
|
|
class TurnController:
|
|
def __init__(
|
|
self,
|
|
*,
|
|
config: AppConfig,
|
|
transport: AudioTransport,
|
|
wakeword: WakeWordProvider,
|
|
vad_recorder: VadRecorder,
|
|
stt: SttProvider,
|
|
llm: LlmProvider,
|
|
tts: TtsProvider,
|
|
ack_tts: TtsProvider,
|
|
context: ConversationContext,
|
|
event_bus: PipelineEventBus,
|
|
sentence_buffer: SentenceBuffer | None = None,
|
|
) -> None:
|
|
self.config = config
|
|
self.transport = transport
|
|
self.wakeword = wakeword
|
|
self.vad_recorder = vad_recorder
|
|
self.stt = stt
|
|
self.llm = llm
|
|
self.tts = tts
|
|
self.ack_tts = ack_tts
|
|
self.context = context
|
|
self.event_bus = event_bus
|
|
self.sentence_buffer = sentence_buffer or SentenceBuffer()
|
|
self._states: list[PipelineState] = []
|
|
|
|
def run_turn(self, turn_id: int) -> TurnResult:
|
|
self._states = []
|
|
try:
|
|
self._event(
|
|
WAKE_LISTENING,
|
|
PipelineState.WAKE_LISTENING,
|
|
"待机:等待唤醒词“小杰小杰”",
|
|
turn_id=turn_id,
|
|
)
|
|
user_text = self._wait_for_wake_and_user_text(turn_id)
|
|
if isinstance(user_text, ProviderError):
|
|
return self._recover(user_text, turn_id)
|
|
return self._reply_to_user(user_text, turn_id)
|
|
except ProviderError as exc:
|
|
return self._recover(exc, turn_id)
|
|
|
|
def _wait_for_wake_and_user_text(self, turn_id: int) -> str | ProviderError:
|
|
wake_error = self._wait_for_local_wake()
|
|
if wake_error is not None:
|
|
return wake_error
|
|
self._event(WAKE_DETECTED, PipelineState.SPEECH_DETECTING, "唤醒命中", turn_id=turn_id)
|
|
ack_error = self._acknowledge_wake(turn_id)
|
|
if ack_error is not None:
|
|
return ack_error
|
|
self._event(QUESTION_PROMPT, PipelineState.SPEECH_DETECTING, "请说出问题", turn_id=turn_id)
|
|
user_segment = self._capture_segment(turn_id, state_message="录音中:正在听取问题")
|
|
if isinstance(user_segment, ProviderError):
|
|
return user_segment
|
|
self._event(STT_STARTED, PipelineState.TRANSCRIBING, "转写中:正在识别问题", turn_id=turn_id)
|
|
transcript = self.stt.transcribe(user_segment)
|
|
user_text = transcript.normalized_text
|
|
if not is_valid_transcript_text(user_text):
|
|
return ProviderError(
|
|
ErrorCode.STT_EMPTY_TRANSCRIPT,
|
|
"STT produced no meaningful user text",
|
|
True,
|
|
"voice-assistant-pipeline",
|
|
"stt",
|
|
)
|
|
self._event(TRANSCRIPT_FINAL, PipelineState.TRANSCRIBING, "", turn_id=turn_id, payload={"text": user_text})
|
|
return user_text
|
|
|
|
def _wait_for_local_wake(self) -> ProviderError | None:
|
|
self.wakeword.reset()
|
|
while True:
|
|
frames = self.transport.read_frames(timeout_ms=100)
|
|
if not frames:
|
|
continue
|
|
for frame in frames:
|
|
event = self.wakeword.detect(frame)
|
|
if event is not None:
|
|
self.wakeword.reset()
|
|
return None
|
|
|
|
def _capture_segment(self, turn_id: int, *, state_message: str) -> AudioSegment | ProviderError:
|
|
self.vad_recorder.reset()
|
|
self.vad_recorder.provider.reset()
|
|
self._event(CAPTURE_STARTED, PipelineState.RECORDING, state_message, turn_id=turn_id)
|
|
while True:
|
|
frames = self.transport.read_frames(timeout_ms=100)
|
|
if not frames:
|
|
continue
|
|
for frame in frames:
|
|
was_started = self.vad_recorder.started
|
|
result = self.vad_recorder.feed(frame)
|
|
if not was_started and self.vad_recorder.started:
|
|
self._event(SPEECH_STARTED, PipelineState.RECORDING, "检测到用户语音", turn_id=turn_id)
|
|
if isinstance(result, ProviderError):
|
|
return result
|
|
if isinstance(result, AudioSegment):
|
|
self._event(
|
|
SPEECH_ENDED,
|
|
PipelineState.RECORDING,
|
|
"用户语音结束",
|
|
turn_id=turn_id,
|
|
payload={"end_reason": result.metadata.get("end_reason", "")},
|
|
)
|
|
return result
|
|
|
|
def _acknowledge_wake(self, turn_id: int) -> ProviderError | None:
|
|
text = self.config.wake_ack_text.strip()
|
|
if not text:
|
|
self._drain_input_after_playback()
|
|
return None
|
|
try:
|
|
self._event(ACK_STARTED, PipelineState.SPEAKING, f"应答中:{text}", turn_id=turn_id)
|
|
segment = self.ack_tts.synthesize(text)
|
|
playback = self.transport.play_pcm(segment)
|
|
if playback.error:
|
|
return playback.error
|
|
self._drain_input_after_playback()
|
|
return None
|
|
except ProviderError as exc:
|
|
return exc
|
|
|
|
def _reply_to_user(self, user_text: str, turn_id: int) -> TurnResult:
|
|
self.context.append_user(user_text)
|
|
self._event(LLM_STARTED, PipelineState.THINKING, "思考中:正在生成回复", turn_id=turn_id)
|
|
assistant_text = ""
|
|
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._speak(sentence, turn_id)
|
|
for sentence in self.sentence_buffer.flush():
|
|
self._speak(sentence, turn_id)
|
|
except ProviderError as exc:
|
|
return self._recover(exc, turn_id)
|
|
if not assistant_text.strip():
|
|
return self._recover(
|
|
ProviderError(
|
|
ErrorCode.LLM_EMPTY_REPLY,
|
|
"LLM returned no assistant text",
|
|
True,
|
|
"voice-assistant-pipeline",
|
|
"llm",
|
|
),
|
|
turn_id,
|
|
)
|
|
self.context.append_assistant(assistant_text)
|
|
self._event(STANDBY_RESUMED, PipelineState.WAKE_LISTENING, "恢复待机:可继续唤醒", turn_id=turn_id)
|
|
return TurnResult(True, user_text, assistant_text, states=list(self._states))
|
|
|
|
def _speak(self, sentence: str, turn_id: int) -> None:
|
|
self._event(TTS_STARTED, PipelineState.SPEAKING, "播放中:正在播报回复", turn_id=turn_id)
|
|
segment = self.tts.synthesize(sentence)
|
|
playback = self.transport.play_pcm(segment)
|
|
if playback.error:
|
|
raise playback.error
|
|
self._event(PLAYBACK_FINISHED, PipelineState.SPEAKING, "", turn_id=turn_id)
|
|
self._drain_input_after_playback()
|
|
|
|
def _drain_input_after_playback(self) -> None:
|
|
if self.config.post_playback_drain_ms <= 0:
|
|
return
|
|
self.transport.flush_input()
|
|
remaining_ms = self.config.post_playback_drain_ms
|
|
while remaining_ms > 0:
|
|
timeout_ms = min(50, remaining_ms)
|
|
self.transport.read_frames(timeout_ms=timeout_ms)
|
|
remaining_ms -= timeout_ms
|
|
self.transport.flush_input()
|
|
|
|
def _recover(self, error: ProviderError, turn_id: int) -> TurnResult:
|
|
self._event(STAGE_ERROR, PipelineState.ERROR_RECOVERING, error.message, turn_id=turn_id, payload={"error": error})
|
|
self._event(RECOVERING, PipelineState.ERROR_RECOVERING, "恢复待机:本轮已结束", turn_id=turn_id)
|
|
self._event(STANDBY_RESUMED, PipelineState.WAKE_LISTENING, "恢复待机:可继续唤醒", turn_id=turn_id)
|
|
return TurnResult(False, error=error, states=list(self._states))
|
|
|
|
def _event(
|
|
self,
|
|
event_type: str,
|
|
state: PipelineState,
|
|
message: str,
|
|
*,
|
|
turn_id: int,
|
|
payload: dict[str, object] | None = None,
|
|
) -> None:
|
|
self._states.append(state)
|
|
self.event_bus.emit(event_type, turn_id=turn_id, state=state, message=message, payload=payload)
|
|
|
|
|
|
class VoiceAssistantPipeline:
|
|
def __init__(
|
|
self,
|
|
*,
|
|
config: AppConfig,
|
|
transport: AudioTransport,
|
|
wakeword: WakeWordProvider,
|
|
vad_recorder: VadRecorder,
|
|
stt: SttProvider,
|
|
llm: LlmProvider,
|
|
tts: TtsProvider,
|
|
context: ConversationContext,
|
|
ack_tts: TtsProvider | None = None,
|
|
reporter: RuntimeReporter | None = None,
|
|
event_bus: PipelineEventBus | None = None,
|
|
sentence_buffer: SentenceBuffer | None = None,
|
|
) -> None:
|
|
self.config = config
|
|
self.transport = transport
|
|
self.wakeword = wakeword
|
|
self.vad_recorder = vad_recorder
|
|
self.stt = stt
|
|
self.llm = llm
|
|
self.tts = tts
|
|
self.ack_tts = ack_tts or tts
|
|
self.context = context
|
|
self.reporter = reporter
|
|
self.event_bus = event_bus or PipelineEventBus()
|
|
if reporter is not None:
|
|
self.event_bus.subscribe(self._report_event)
|
|
self.sentence_buffer = sentence_buffer or SentenceBuffer()
|
|
self.controller = TurnController(
|
|
config=config,
|
|
transport=transport,
|
|
wakeword=wakeword,
|
|
vad_recorder=vad_recorder,
|
|
stt=stt,
|
|
llm=llm,
|
|
tts=tts,
|
|
ack_tts=self.ack_tts,
|
|
context=context,
|
|
event_bus=self.event_bus,
|
|
sentence_buffer=self.sentence_buffer,
|
|
)
|
|
|
|
def load(self) -> None:
|
|
self.wakeword.load()
|
|
self.vad_recorder.provider.load()
|
|
self.stt.load()
|
|
self.tts.load()
|
|
if self.ack_tts is not self.tts:
|
|
self.ack_tts.load()
|
|
|
|
def run(self, *, once: bool = False, max_turns: int | None = None) -> RuntimeSummary:
|
|
completed = 0
|
|
failed = 0
|
|
last_error: ProviderError | None = None
|
|
self.load()
|
|
self.transport.start_input(
|
|
device_id=self.config.audio_input_device,
|
|
sample_rate=self.config.sample_rate,
|
|
channels=self.config.channels,
|
|
)
|
|
try:
|
|
while True:
|
|
turn_id = completed + failed + 1
|
|
result = self.run_turn(turn_id)
|
|
if result.success:
|
|
completed += 1
|
|
else:
|
|
failed += 1
|
|
last_error = result.error
|
|
if once:
|
|
break
|
|
if once and completed >= 1:
|
|
break
|
|
if max_turns is not None and completed >= max_turns:
|
|
break
|
|
except KeyboardInterrupt:
|
|
return RuntimeSummary(completed, failed, interrupted=True, last_error=last_error)
|
|
finally:
|
|
self.shutdown()
|
|
return RuntimeSummary(completed, failed, last_error=last_error)
|
|
|
|
def run_turn(self, turn_id: int) -> TurnResult:
|
|
return self.controller.run_turn(turn_id)
|
|
|
|
def shutdown(self) -> None:
|
|
self.transport.stop()
|
|
|
|
def _report_event(self, event: PipelineEvent) -> None:
|
|
if self.reporter is None:
|
|
return
|
|
handler = getattr(self.reporter, "handle_event", None)
|
|
if callable(handler):
|
|
handler(event)
|
|
return
|
|
dispatch_pipeline_event(self.reporter, event)
|