Files
Owner/src/owner_voice_pet/runtime.py
T

455 lines
18 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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(),
)