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)