455 lines
18 KiB
Python
455 lines
18 KiB
Python
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(),
|
||
)
|