722 lines
28 KiB
Python
722 lines
28 KiB
Python
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
|