[Pipeline 状态机]:完成TurnController重构,包含wake到standby闭环和错误恢复

This commit is contained in:
mkbk
2026-06-17 21:29:28 +08:00
parent a251ff379d
commit f9da304568
5 changed files with 364 additions and 7 deletions
@@ -50,7 +50,7 @@
- [x] 7.1 更新 OpenSpec 以描述 stage 化 pipeline、事件总线、TurnController 和主说话人端点;前置条件:公开参考已确认;验收标准:proposal/design/spec/tasks 覆盖新架构和任务;测试要点:OpenSpec strict;优先级:P0;预计:45 分钟。
- [x] 7.2 实现 pipeline event bus 和终端事件映射;前置条件:7.1 完成;验收标准:所有 live 用户可见状态由事件产生;测试要点:事件顺序和终端文案测试;优先级:P0;预计:60 分钟。
- [ ] 7.3 实现 `TurnController``VoiceAssistantPipeline`;前置条件:7.2 完成;验收标准:`run-live` 使用统一 pipeline,成功/失败 turn 均恢复待机;测试要点:两轮 fake runtime、错误恢复、上下文回归;优先级:P0;预计:60 分钟。
- [x] 7.3 实现 `TurnController``VoiceAssistantPipeline`;前置条件:7.2 完成;验收标准:`run-live` 使用统一 pipeline,成功/失败 turn 均恢复待机;测试要点:两轮 fake runtime、错误恢复、上下文回归;优先级:P0;预计:60 分钟。
- [ ] 7.4 实现本轮主说话人端点;前置条件:7.3 完成;验收标准:主说话人音色消失约 300 ms 后结束采集;测试要点:一次提问后背景噪声不拖尾、短暂停顿不断句、画像不足回退;优先级:P0;预计:60 分钟。
- [ ] 7.5 更新 README、`.env.example`、本地 `.env` 非密钥配置;前置条件:7.2 至 7.4 完成;验收标准:运行说明匹配新 pipeline;测试要点:`--show-config` 不泄露 key;优先级:P0;预计:30 分钟。
- [ ] 7.6 验证并提交“Pipeline 文档验收”模块;前置条件:7.1 至 7.5 完成;验收标准:compileall、unittest、security-check、model-check、device-check、OpenSpec strict 全通过;优先级:P0;预计:30 分钟。
+3
View File
@@ -1,6 +1,7 @@
"""Owner voice pet pipeline package."""
from .config import AppConfig
from .assistant_pipeline import TurnController, VoiceAssistantPipeline
from .events import PipelineEvent, PipelineEventBus
from .models import (
AudioFrame,
@@ -30,6 +31,8 @@ from .ui import ConsolePetWindow, PetStateController, PetVisualState
__all__ = [
"AppConfig",
"TurnController",
"VoiceAssistantPipeline",
"PipelineEvent",
"PipelineEventBus",
"AudioFrame",
+351
View File
@@ -0,0 +1,351 @@
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Protocol
from .config import AppConfig
from .conversation import ConversationContext
from .events import (
ACK_STARTED,
CAPTURE_STARTED,
LLM_STARTED,
PLAYBACK_FINISHED,
QUESTION_PROMPT,
RECOVERING,
SPEECH_ENDED,
SPEECH_STARTED,
STAGE_ERROR,
STANDBY_RESUMED,
STT_STARTED,
TRANSCRIPT_FINAL,
TTS_STARTED,
WAKE_DETECTED,
WAKE_LISTENING,
PipelineEvent,
PipelineEventBus,
dispatch_pipeline_event,
)
from .models import AudioSegment, ErrorCode, PipelineState, ProviderError
from .protocols import AudioTransport, LlmProvider, 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)
@dataclass(slots=True)
class RuntimeSummary:
completed_turns: int
failed_turns: int
interrupted: bool = False
last_error: ProviderError | None = None
class TurnController:
def __init__(
self,
*,
config: AppConfig,
transport: AudioTransport,
wakeword: WakeWordProvider,
vad_recorder: VadRecorder,
stt: SttProvider,
llm: LlmProvider,
tts: TtsProvider,
ack_tts: TtsProvider,
context: ConversationContext,
event_bus: PipelineEventBus,
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
self.context = context
self.event_bus = event_bus
self.sentence_buffer = sentence_buffer or SentenceBuffer()
self._states: list[PipelineState] = []
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)
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,
"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) -> 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:
self._drain_input_after_playback()
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:
self.context.append_user(user_text)
self._event(LLM_STARTED, PipelineState.THINKING, "思考中:正在生成回复", turn_id=turn_id)
assistant_text = ""
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)):
self._speak(sentence, turn_id)
for sentence in self.sentence_buffer.flush():
self._speak(sentence, turn_id)
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,
"voice-assistant-pipeline",
"llm",
),
turn_id,
)
self.context.append_assistant(assistant_text)
self._event(STANDBY_RESUMED, PipelineState.WAKE_LISTENING, "恢复待机:可继续唤醒", turn_id=turn_id)
return TurnResult(True, user_text, assistant_text, states=list(self._states))
def _speak(self, sentence: str, turn_id: int) -> None:
self._event(TTS_STARTED, PipelineState.SPEAKING, "播放中:正在播报回复", turn_id=turn_id)
segment = self.tts.synthesize(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()
def _drain_input_after_playback(self) -> None:
if self.config.post_playback_drain_ms <= 0:
return
self.transport.flush_input()
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) -> 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)
class VoiceAssistantPipeline:
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
self.event_bus = event_bus or PipelineEventBus()
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,
stt=stt,
llm=llm,
tts=tts,
ack_tts=self.ack_tts,
context=context,
event_bus=self.event_bus,
sentence_buffer=self.sentence_buffer,
)
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()
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:
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)
+4 -3
View File
@@ -5,6 +5,7 @@ from dataclasses import dataclass, field
from typing import Protocol
from .config import AppConfig
from .assistant_pipeline import VoiceAssistantPipeline
from .conversation import ConversationContext
from .events import (
ACK_STARTED,
@@ -325,7 +326,7 @@ class LiveVoiceRuntime:
dispatch_pipeline_event(self.reporter, event)
def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = None) -> LiveVoiceRuntime:
def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = None) -> VoiceAssistantPipeline:
errors = config.validate_basic()
if errors:
raise errors[0]
@@ -344,7 +345,7 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
vad_provider = SherpaOnnxVadProvider(config.speech_models_dir, threshold=config.vad_threshold)
else:
vad_provider = EnergyVadProvider()
return LiveVoiceRuntime(
return VoiceAssistantPipeline(
config=config,
transport=SoundDeviceAudioTransport(output_device=config.audio_output_device),
wakeword=SherpaOnnxKeywordWakeWordProvider(
@@ -369,5 +370,5 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
max_messages=config.context_max_messages,
max_chars=config.context_max_chars,
),
reporter=reporter,
reporter=reporter or TerminalRuntimeReporter(),
)
+5 -3
View File
@@ -21,7 +21,7 @@ from owner_voice_pet.events import (
)
from owner_voice_pet.llm import MockLlmProvider
from owner_voice_pet.models import AudioFrame, AudioSegment, Transcript
from owner_voice_pet.runtime import LiveVoiceRuntime
from owner_voice_pet.assistant_pipeline import VoiceAssistantPipeline
from owner_voice_pet.transport import MemoryAudioTransport
from owner_voice_pet.tts import SineTtsProvider
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder
@@ -82,7 +82,7 @@ class RecordingReporter:
self.errors.append(f"{stage}:{code}:{message}")
def make_runtime(texts: list[str], context: ConversationContext | None = None) -> tuple[LiveVoiceRuntime, QueueSttProvider, MockLlmProvider, MemoryAudioTransport, RecordingReporter]:
def make_runtime(texts: list[str], context: ConversationContext | None = None) -> tuple[VoiceAssistantPipeline, QueueSttProvider, MockLlmProvider, MemoryAudioTransport, RecordingReporter]:
frames = []
for idx, _text in enumerate(texts):
base_id = idx * 5
@@ -95,7 +95,7 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) -
tts = SineTtsProvider()
reporter = RecordingReporter()
event_bus = PipelineEventBus()
runtime = LiveVoiceRuntime(
runtime = VoiceAssistantPipeline(
config=AppConfig(llm_api_key="secret", speech_provider="cloud", post_playback_drain_ms=0),
transport=transport,
wakeword=KeywordWakeWordProvider(),
@@ -113,6 +113,8 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) -
class LiveRuntimeTests(unittest.TestCase):
def test_repeated_runtime_runs_two_turns_and_returns_to_standby(self) -> None:
runtime, stt, llm, transport, reporter = make_runtime(["第一问", "第二问"])
self.assertIsInstance(runtime, VoiceAssistantPipeline)
self.assertIsNotNone(runtime.controller)
summary = runtime.run(max_turns=2)
self.assertEqual(summary.completed_turns, 2)