from __future__ import annotations from dataclasses import dataclass, field from typing import Protocol from .audio_preprocess import NoopAudioPreprocessor from .config import AppConfig from .continuation import ContinuationDecision, ContinuationDecisionProvider, build_continuation_decider from .conversation import ConversationContext from .events import ( ACK_STARTED, BARGE_IN_DETECTED, CAPTURE_STARTED, CONTINUATION_DECISION_MADE, CONTINUATION_DECISION_STARTED, CONTINUOUS_SESSION_ENDED, FOLLOWUP_LISTENING, FOLLOWUP_TIMEOUT, LLM_STARTED, PLAYBACK_FINISHED, PLAYBACK_INTERRUPTED, QUESTION_PROMPT, RECOVERING, SPEECH_ENDED, SPEECH_STARTED, STAGE_ERROR, STANDBY_RESUMED, STT_STARTED, TRANSCRIPT_FINAL, TRANSCRIPT_PARTIAL, TTS_STARTED, WAKE_DETECTED, WAKE_LISTENING, PipelineEvent, PipelineEventBus, dispatch_pipeline_event, ) from .models import AudioFrame, AudioSegment, ErrorCode, PipelineState, ProviderError from .protocols import ( AudioPreprocessor, AudioTransport, LlmProvider, RealtimeSttProvider, RealtimeTranscriptSession, 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) completed_turns: int = 0 failed_turns: int = 0 @dataclass(slots=True) class RuntimeSummary: completed_turns: int failed_turns: int interrupted: bool = False last_error: ProviderError | None = None @dataclass(slots=True) class SpeakResult: spoken_text: str interrupted: bool = False class TurnController: def __init__( self, *, config: AppConfig, transport: AudioTransport, wakeword: WakeWordProvider, vad_recorder: VadRecorder, audio_preprocessor: AudioPreprocessor, stt: SttProvider, realtime_stt: RealtimeSttProvider | None, llm: LlmProvider, tts: TtsProvider, ack_tts: TtsProvider, context: ConversationContext, event_bus: PipelineEventBus, continuation_decider: ContinuationDecisionProvider, sentence_buffer: SentenceBuffer | None = None, ) -> None: self.config = config self.transport = transport self.wakeword = wakeword self.vad_recorder = vad_recorder self.audio_preprocessor = audio_preprocessor self.stt = stt self.realtime_stt = realtime_stt self.llm = llm self.tts = tts self.ack_tts = ack_tts self.context = context self.event_bus = event_bus self.continuation_decider = continuation_decider self.sentence_buffer = sentence_buffer or SentenceBuffer() self._states: list[PipelineState] = [] self._pending_capture_frames: list[AudioFrame] = [] 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) return self._capture_and_transcribe(turn_id, state_message="录音中:正在听取问题") def _capture_and_transcribe( self, turn_id: int, *, state_message: str, no_speech_timeout_ms: int | None = None, ) -> str | ProviderError: user_segment = self._capture_segment( turn_id, state_message=state_message, no_speech_timeout_ms=no_speech_timeout_ms, ) 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, no_speech_timeout_ms: int | None = None, ) -> AudioSegment | ProviderError: original_no_speech_timeout_ms = self.vad_recorder.no_speech_timeout_ms if no_speech_timeout_ms is not None: self.vad_recorder.no_speech_timeout_ms = no_speech_timeout_ms try: self.vad_recorder.reset() self.vad_recorder.provider.reset() self.audio_preprocessor.reset() realtime_session = self._start_realtime_transcript() last_partial_ms: int | None = None self._event(CAPTURE_STARTED, PipelineState.RECORDING, state_message, turn_id=turn_id) while True: frames = self._read_capture_frames(timeout_ms=100) if not frames: continue for frame in frames: try: frame = self.audio_preprocessor.process_frame(frame) except ProviderError as exc: return exc 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 self.vad_recorder.started and realtime_session is not None: if self._emit_realtime_transcript(realtime_session, frame, turn_id): last_partial_ms = frame.timestamp_ms if isinstance(result, AudioSegment): if realtime_session is not None: self._finish_realtime_transcript(realtime_session, turn_id) self._event( SPEECH_ENDED, PipelineState.RECORDING, "用户语音结束", turn_id=turn_id, payload={"end_reason": result.metadata.get("end_reason", "")}, ) return result if self._should_end_after_realtime_idle(last_partial_ms, frame.timestamp_ms): if realtime_session is not None: self._finish_realtime_transcript(realtime_session, turn_id) result = self.vad_recorder.finish("partial_transcript_idle") self._event( SPEECH_ENDED, PipelineState.RECORDING, "用户语音结束", turn_id=turn_id, payload={"end_reason": result.metadata.get("end_reason", "")}, ) return result finally: self.vad_recorder.no_speech_timeout_ms = original_no_speech_timeout_ms def _read_capture_frames(self, *, timeout_ms: int) -> list[AudioFrame]: if self._pending_capture_frames: frames = list(self._pending_capture_frames) self._pending_capture_frames.clear() return frames return self.transport.read_frames(timeout_ms=timeout_ms) def _start_realtime_transcript(self) -> RealtimeTranscriptSession | None: if not self.config.realtime_transcript_enabled or self.realtime_stt is None: return None return self.realtime_stt.start_stream() def _emit_realtime_transcript( self, realtime_session: RealtimeTranscriptSession, frame: AudioFrame, turn_id: int, ) -> bool: transcript = realtime_session.accept_frame(frame) if transcript is None: return False text = transcript.normalized_text if not is_valid_transcript_text(text): return False self._event(TRANSCRIPT_PARTIAL, PipelineState.RECORDING, "", turn_id=turn_id, payload={"text": text}) return True def _should_end_after_realtime_idle(self, last_partial_ms: int | None, current_ms: int) -> bool: timeout_ms = self.config.realtime_transcript_idle_timeout_ms return timeout_ms > 0 and last_partial_ms is not None and current_ms - last_partial_ms >= timeout_ms def _finish_realtime_transcript( self, realtime_session: RealtimeTranscriptSession, turn_id: int, ) -> None: transcript = realtime_session.finish() if transcript is None: return text = transcript.normalized_text if is_valid_transcript_text(text): self._event(TRANSCRIPT_PARTIAL, PipelineState.RECORDING, "", turn_id=turn_id, payload={"text": text}) def _acknowledge_wake(self, turn_id: int) -> ProviderError | None: text = self.config.wake_ack_text.strip() if not text: 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: completed_turns = 0 current_user_text = user_text current_turn_id = turn_id last_assistant_text = "" while True: reply_result = self._reply_once(current_user_text, current_turn_id, completed_turns=completed_turns) completed_turns += reply_result.completed_turns if not reply_result.success: return reply_result last_assistant_text = reply_result.assistant_text if reply_result.error is not None: return reply_result if reply_result.states and reply_result.states[-1] == PipelineState.INTERRUPTED: followup = self._listen_for_followup(current_turn_id + 1, interrupted=True) else: decision = self._decide_continuation(current_user_text, last_assistant_text, current_turn_id) if not decision.should_continue: self._event( CONTINUOUS_SESSION_ENDED, PipelineState.WAKE_LISTENING, "", turn_id=current_turn_id, payload={"decision": decision.action, "reason": decision.reason}, ) self._event( STANDBY_RESUMED, PipelineState.WAKE_LISTENING, "恢复待机:可继续唤醒", turn_id=current_turn_id, ) return TurnResult( True, current_user_text, last_assistant_text, states=list(self._states), completed_turns=completed_turns, ) followup = self._listen_for_followup(current_turn_id + 1, interrupted=False) if followup is None: return TurnResult( True, current_user_text, last_assistant_text, states=list(self._states), completed_turns=completed_turns, ) if isinstance(followup, ProviderError): return self._recover(followup, current_turn_id + 1, completed_turns=completed_turns) current_turn_id += 1 current_user_text = followup def _reply_once(self, user_text: str, turn_id: int, *, completed_turns: int) -> TurnResult: self.context.append_user(user_text) self._event(LLM_STARTED, PipelineState.THINKING, "思考中:正在生成回复", turn_id=turn_id) assistant_text = "" spoken_parts: list[str] = [] interrupted = False 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)): speak_result = self._speak(sentence, turn_id) if speak_result.interrupted: interrupted = True break spoken_parts.append(speak_result.spoken_text) if interrupted: break if interrupted: self.sentence_buffer.flush() else: for sentence in self.sentence_buffer.flush(): speak_result = self._speak(sentence, turn_id) if speak_result.interrupted: interrupted = True break spoken_parts.append(speak_result.spoken_text) except ProviderError as exc: return self._recover(exc, turn_id, completed_turns=completed_turns) if not assistant_text.strip() and not "".join(spoken_parts).strip(): return self._recover( ProviderError( ErrorCode.LLM_EMPTY_REPLY, "LLM returned no assistant text", True, "voice-assistant-pipeline", "llm", ), turn_id, completed_turns=completed_turns, ) spoken_text = "".join(spoken_parts) if spoken_text.strip(): self.context.append_assistant(spoken_text) states = list(self._states) if interrupted: states.append(PipelineState.INTERRUPTED) return TurnResult( True, user_text, spoken_text or assistant_text, states=states, completed_turns=1, ) def _decide_continuation(self, user_text: str, assistant_text: str, turn_id: int) -> ContinuationDecision: if not self.config.continuous_dialog_enabled: return ContinuationDecision("standby", 1.0, "continuous dialog disabled", "config") self._event(CONTINUATION_DECISION_STARTED, PipelineState.THINKING, "", turn_id=turn_id) decision = self.continuation_decider.decide( user_text=user_text, assistant_text=assistant_text, history=self.context.messages(), ) self._event( CONTINUATION_DECISION_MADE, PipelineState.THINKING, "", turn_id=turn_id, payload={ "action": decision.action, "confidence": decision.confidence, "reason": decision.reason, "provider": decision.provider, }, ) return decision def _listen_for_followup(self, turn_id: int, *, interrupted: bool) -> str | ProviderError | None: if not interrupted: seconds = max(1, round(self.config.followup_listen_timeout_ms / 1000)) self._event( FOLLOWUP_LISTENING, PipelineState.RECORDING, f"继续对话:{seconds}秒内可直接回答", turn_id=turn_id, ) user_text = self._capture_and_transcribe( turn_id, state_message="录音中:正在听取追问" if not interrupted else "录音中:正在听取打断内容", no_speech_timeout_ms=self.config.followup_listen_timeout_ms, ) if isinstance(user_text, ProviderError) and user_text.code == ErrorCode.VAD_TIMEOUT_NO_SPEECH: self._event( FOLLOWUP_TIMEOUT, PipelineState.WAKE_LISTENING, "追问超时:未检测到用户回答", turn_id=turn_id, ) self._event(CONTINUOUS_SESSION_ENDED, PipelineState.WAKE_LISTENING, "", turn_id=turn_id) self._event(STANDBY_RESUMED, PipelineState.WAKE_LISTENING, "恢复待机:可继续唤醒", turn_id=turn_id) return None return user_text def _speak(self, sentence: str, turn_id: int) -> SpeakResult: self._event(TTS_STARTED, PipelineState.SPEAKING, "播放中:正在播报回复", turn_id=turn_id) segment = self.tts.synthesize(sentence) if not self._can_interrupt_playback(segment): 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() return SpeakResult(sentence) if self._play_interruptible(segment, turn_id=turn_id): return SpeakResult("", interrupted=True) self._event(PLAYBACK_FINISHED, PipelineState.SPEAKING, "", turn_id=turn_id) self._drain_input_after_playback() return SpeakResult(sentence) def _can_interrupt_playback(self, segment: AudioSegment) -> bool: return ( self.config.barge_in_enabled and self.realtime_stt is not None and segment.duration_ms > self.config.barge_in_echo_guard_ms and not segment.metadata.get("format") ) def _play_interruptible(self, segment: AudioSegment, *, turn_id: int) -> bool: elapsed_ms = 0 guard_cleared = False self.vad_recorder.provider.reset() realtime_session = self._start_realtime_transcript() speech_ms = 0 partial_seen = False pending_frames: list[AudioFrame] = [] for chunk in _audio_chunks(segment, chunk_ms=100): playback = self.transport.play_pcm(chunk) if playback.error: raise playback.error elapsed_ms += chunk.duration_ms if elapsed_ms < self.config.barge_in_echo_guard_ms: continue if not guard_cleared: self.transport.flush_input() guard_cleared = True continue detected, speech_ms, partial_seen, new_frames = self._detect_barge_in( realtime_session, turn_id=turn_id, speech_ms=speech_ms, partial_seen=partial_seen, ) pending_frames.extend(new_frames) if detected: self._pending_capture_frames.extend(pending_frames) self._event(BARGE_IN_DETECTED, PipelineState.INTERRUPTED, "检测到用户打断", turn_id=turn_id) self._event(PLAYBACK_INTERRUPTED, PipelineState.INTERRUPTED, "播报已打断", turn_id=turn_id) if realtime_session is not None: realtime_session.finish() return True if realtime_session is not None: realtime_session.finish() return False def _detect_barge_in( self, realtime_session: RealtimeTranscriptSession | None, *, turn_id: int, speech_ms: int, partial_seen: bool, ) -> tuple[bool, int, bool, list[AudioFrame]]: frames = self.transport.read_frames(timeout_ms=0) if not frames: return False, speech_ms, partial_seen, [] for frame in frames: result = self.vad_recorder.provider.analyze(frame) if result.is_speech: speech_ms += int(frame.metadata.get("duration_ms", 20)) if realtime_session is not None and self._emit_realtime_transcript(realtime_session, frame, turn_id): partial_seen = True else: speech_ms = 0 detected = speech_ms >= self.config.barge_in_min_speech_ms and partial_seen return detected, speech_ms, partial_seen, frames def _drain_input_after_playback(self) -> None: self.transport.flush_input() if self.config.post_playback_drain_ms <= 0: return 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, *, completed_turns: int = 0) -> 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), completed_turns=completed_turns, failed_turns=1) 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, audio_preprocessor: AudioPreprocessor | None = None, realtime_stt: RealtimeSttProvider | None = None, ack_tts: TtsProvider | None = None, reporter: RuntimeReporter | None = None, event_bus: PipelineEventBus | None = None, continuation_decider: ContinuationDecisionProvider | None = None, sentence_buffer: SentenceBuffer | None = None, ) -> None: self.config = config self.transport = transport self.wakeword = wakeword self.vad_recorder = vad_recorder self.audio_preprocessor = audio_preprocessor or NoopAudioPreprocessor() self.stt = stt self.realtime_stt = realtime_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() self.continuation_decider = continuation_decider or build_continuation_decider( config.continuation_decision_provider, llm, threshold=config.continuation_confidence_threshold, ) 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, audio_preprocessor=self.audio_preprocessor, stt=stt, realtime_stt=realtime_stt, llm=llm, tts=tts, ack_tts=self.ack_tts, context=context, event_bus=self.event_bus, continuation_decider=self.continuation_decider, sentence_buffer=self.sentence_buffer, ) def load(self) -> None: self.wakeword.load() self.vad_recorder.provider.load() self.audio_preprocessor.load() self.stt.load() if self.realtime_stt is not None and self.realtime_stt is not self.stt: self.realtime_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) completed += result.completed_turns if result.success: completed += 0 if result.completed_turns else 1 else: failed += result.failed_turns or 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) def _audio_chunks(segment: AudioSegment, *, chunk_ms: int) -> list[AudioSegment]: if chunk_ms <= 0 or segment.duration_ms <= chunk_ms: return [segment] bytes_per_ms = max(1, int(segment.sample_rate * segment.channels * 2 / 1000)) chunk_bytes = max(2 * segment.channels, bytes_per_ms * chunk_ms) chunk_bytes -= chunk_bytes % (2 * segment.channels) chunks: list[AudioSegment] = [] offset = 0 start_ms = segment.start_time_ms while offset < len(segment.pcm): data = segment.pcm[offset : offset + chunk_bytes] duration_ms = max(1, int(len(data) / bytes_per_ms)) chunks.append( AudioSegment( data, segment.sample_rate, segment.channels, start_ms, start_ms + duration_ms, dict(segment.metadata), ) ) offset += len(data) start_ms += duration_ms return chunks