from __future__ import annotations import sys from dataclasses import dataclass, field from typing import Protocol from .audio_preprocess import NoopAudioPreprocessor, SherpaOnnxDenoiserPreprocessor from .config import AppConfig from .assistant_pipeline import VoiceAssistantPipeline from .conversation import ConversationContext from .events import ( ACK_STARTED, CAPTURE_STARTED, LLM_STARTED, PipelineEvent, PipelineEventBus, PLAYBACK_FINISHED, QUESTION_PROMPT, RECOVERING, SPEECH_ENDED, SPEECH_STARTED, STAGE_ERROR, STANDBY_RESUMED, STT_STARTED, TRANSCRIPT_FINAL, TTS_STARTED, WAKE_DETECTED, WAKE_LISTENING, dispatch_pipeline_event, ) from .llm import OpenAICompatibleLlmProvider from .models import AudioSegment, ErrorCode, PipelineState, ProviderError from .protocols import AudioTransport, LlmProvider, RealtimeSttProvider, SttProvider, TtsProvider, WakeWordProvider from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text from .transport import SoundDeviceAudioTransport from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer, make_end_chime, sanitize_tts_text from .vad import EnergyVadProvider, HybridVadProvider, PrimarySpeakerVadRecorder, SherpaOnnxVadProvider, VadRecorder from .wakeword import SherpaOnnxKeywordWakeWordProvider 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: ... class TerminalRuntimeReporter: def handle_event(self, event: PipelineEvent) -> None: dispatch_pipeline_event(self, event) def status(self, state: str, message: str, *, turn_id: int | None = None) -> None: prefix = f"[第{turn_id}轮] " if turn_id is not None else "" print(f"{prefix}{message}", flush=True) def transcript(self, text: str, *, final: bool, turn_id: int | None = None) -> None: prefix = f"[第{turn_id}轮] " if turn_id is not None else "" label = "转写结果" if final else "实时转写" print(f"{prefix}{label}:{text}", flush=True) def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None: prefix = f"[第{turn_id}轮] " if turn_id is not None else "" print(f"{prefix}{stage}失败:{code} {message}", file=sys.stderr, flush=True) @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 LiveVoiceRuntime: 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 or TerminalRuntimeReporter() self.event_bus = event_bus or PipelineEventBus() self.event_bus.subscribe(self._report_event) self.sentence_buffer = sentence_buffer or SentenceBuffer() self._states: list[PipelineState] = [] self._cached_ack_text: str | None = None self._cached_ack_segment: AudioSegment | None = None 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() self.prepare_ack_audio() def prepare_ack_audio(self) -> None: text = self.config.wake_ack_text.strip() if not text: self._cached_ack_text = None self._cached_ack_segment = None return self._cached_ack_text = text self._cached_ack_segment = self.ack_tts.synthesize(text) 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: 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 shutdown(self) -> None: self.transport.stop() def _wait_for_wake_and_user_text(self, turn_id: int) -> str | ProviderError: wake_error = self._wait_for_local_wake(turn_id) 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, "live-runtime", "stt", ) self._event(TRANSCRIPT_FINAL, PipelineState.TRANSCRIBING, "", turn_id=turn_id, payload={"text": user_text}) return user_text def _wait_for_local_wake(self, turn_id: int) -> 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: return None try: self._event(ACK_STARTED, PipelineState.SPEAKING, f"应答中:{text}", turn_id=turn_id) segment = self._ack_segment(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 _ack_segment(self, text: str) -> AudioSegment: if self._cached_ack_text != text or self._cached_ack_segment is None: self._cached_ack_text = text self._cached_ack_segment = self.ack_tts.synthesize(text) return self._cached_ack_segment 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 _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 = "" spoken_parts: list[str] = [] 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)): spoken = self._speak(sentence, turn_id) if spoken: spoken_parts.append(spoken) for sentence in self.sentence_buffer.flush(): spoken = self._speak(sentence, turn_id) if spoken: spoken_parts.append(spoken) 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, "live-runtime", "llm", ), turn_id, ) spoken_text = "".join(spoken_parts) if not spoken_text.strip(): return self._recover( ProviderError( ErrorCode.TTS_EMPTY_AUDIO, "LLM reply contained no speakable text after TTS sanitization", True, "live-runtime", "tts", ), turn_id, ) self.context.append_assistant(spoken_text) self._play_end_chime() self._event(STANDBY_RESUMED, PipelineState.WAKE_LISTENING, "恢复待机:可继续唤醒", turn_id=turn_id) return TurnResult(True, user_text, spoken_text, states=list(self._states)) def _speak(self, sentence: str, turn_id: int) -> str: spoken_sentence = sanitize_tts_text(sentence) if not spoken_sentence: return "" self._event(TTS_STARTED, PipelineState.SPEAKING, "播放中:正在播报回复", turn_id=turn_id) segment = self.tts.synthesize(spoken_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() return spoken_sentence def _play_end_chime(self) -> None: if not self.config.end_chime_enabled: return segment = make_end_chime( file_path=self.config.end_chime_file, frequency_hz=self.config.end_chime_frequency_hz, duration_ms=self.config.end_chime_duration_ms, sample_rate=self.config.sample_rate, channels=self.config.channels, ) playback = self.transport.play_pcm(segment) if playback.error: return self._drain_input_after_playback() 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) def _report_event(self, event: PipelineEvent) -> None: handler = getattr(self.reporter, "handle_event", None) if callable(handler): handler(event) return dispatch_pipeline_event(self.reporter, event) def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = None) -> VoiceAssistantPipeline: errors = config.validate_basic() if errors: raise errors[0] if config.speech_provider == "cloud": stt: SttProvider = CloudAsrSttProvider(config) tts: TtsProvider = CloudTtsProvider(config) else: stt = SherpaOnnxSttProvider(str(config.speech_models_dir)) tts = MacSayTtsProvider() realtime_stt: RealtimeSttProvider | None = None if config.realtime_transcript_enabled: if isinstance(stt, SherpaOnnxSttProvider): realtime_stt = stt else: realtime_stt = SherpaOnnxSttProvider(str(config.speech_models_dir)) if config.vad_provider == "hybrid": vad_provider = HybridVadProvider( SherpaOnnxVadProvider(config.speech_models_dir, threshold=config.vad_threshold), EnergyVadProvider(), ) elif config.vad_provider == "local": vad_provider = SherpaOnnxVadProvider(config.speech_models_dir, threshold=config.vad_threshold) else: vad_provider = EnergyVadProvider() recorder_cls = PrimarySpeakerVadRecorder if config.endpoint_mode == "primary_speaker" else VadRecorder recorder_kwargs = { "provider": vad_provider, "min_duration_ms": config.vad_min_duration_ms, "end_silence_ms": config.vad_end_silence_ms, "no_speech_timeout_ms": config.vad_no_speech_timeout_ms, "max_recording_ms": config.vad_max_recording_ms, } if recorder_cls is PrimarySpeakerVadRecorder: recorder_kwargs.update( { "speaker_profile_ms": config.speaker_profile_ms, "speaker_profile_min_ms": config.speaker_profile_min_ms, "speaker_absent_ms": config.speaker_absent_ms, "similarity_threshold": config.speaker_similarity_threshold, "min_rms": config.speaker_min_rms, } ) audio_preprocessor = ( SherpaOnnxDenoiserPreprocessor(config.speech_models_dir) if config.noise_filter_enabled else NoopAudioPreprocessor() ) return VoiceAssistantPipeline( config=config, transport=SoundDeviceAudioTransport(output_device=config.audio_output_device), wakeword=SherpaOnnxKeywordWakeWordProvider( config.speech_models_dir, keyword=config.wake_word, keywords_file=config.wake_keywords_file, threshold=config.wake_kws_threshold, score=config.wake_kws_score, ), vad_recorder=recorder_cls(**recorder_kwargs), audio_preprocessor=audio_preprocessor, stt=stt, realtime_stt=realtime_stt, llm=OpenAICompatibleLlmProvider(config), tts=tts, ack_tts=MacSayTtsProvider(), context=ConversationContext( max_messages=config.context_max_messages, max_chars=config.context_max_chars, ), reporter=reporter or TerminalRuntimeReporter(), )