[Wake/VAD/STT 与 Live runtime]:完成真实重复语音运行,包含云端ASR/TTS开关、run-live和临时上下文测试
This commit is contained in:
@@ -0,0 +1,267 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Protocol
|
||||
|
||||
from .config import AppConfig
|
||||
from .conversation import ConversationContext
|
||||
from .llm import OpenAICompatibleLlmProvider
|
||||
from .models import AudioSegment, ErrorCode, PipelineState, ProviderError
|
||||
from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider
|
||||
from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
|
||||
from .transport import SoundDeviceAudioTransport
|
||||
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer
|
||||
from .vad import EnergyVadProvider, VadRecorder
|
||||
|
||||
|
||||
class RuntimeReporter(Protocol):
|
||||
def status(self, state: str, message: str, *, turn_id: int | None = None) -> None:
|
||||
...
|
||||
|
||||
def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None:
|
||||
...
|
||||
|
||||
|
||||
class TerminalRuntimeReporter:
|
||||
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 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,
|
||||
vad_recorder: VadRecorder,
|
||||
stt: SttProvider,
|
||||
llm: LlmProvider,
|
||||
tts: TtsProvider,
|
||||
context: ConversationContext,
|
||||
reporter: RuntimeReporter | None = None,
|
||||
sentence_buffer: SentenceBuffer | None = None,
|
||||
) -> None:
|
||||
self.config = config
|
||||
self.transport = transport
|
||||
self.vad_recorder = vad_recorder
|
||||
self.stt = stt
|
||||
self.llm = llm
|
||||
self.tts = tts
|
||||
self.context = context
|
||||
self.reporter = reporter or TerminalRuntimeReporter()
|
||||
self.sentence_buffer = sentence_buffer or SentenceBuffer()
|
||||
self._states: list[PipelineState] = []
|
||||
|
||||
def load(self) -> None:
|
||||
self.vad_recorder.provider.load()
|
||||
self.stt.load()
|
||||
self.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:
|
||||
self._states = []
|
||||
try:
|
||||
self._state(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:
|
||||
while True:
|
||||
wake_segment = self._capture_segment(turn_id, state_message="待机:检测到语音,正在判断唤醒词")
|
||||
if isinstance(wake_segment, ProviderError):
|
||||
if wake_segment.code == ErrorCode.VAD_TIMEOUT_NO_SPEECH:
|
||||
self._state(PipelineState.WAKE_LISTENING, "待机:继续等待唤醒词“小杰小杰”", turn_id=turn_id)
|
||||
continue
|
||||
return wake_segment
|
||||
try:
|
||||
wake_transcript = self.stt.transcribe(wake_segment).normalized_text
|
||||
except ProviderError as exc:
|
||||
if exc.code == ErrorCode.STT_EMPTY_TRANSCRIPT:
|
||||
self._state(PipelineState.WAKE_LISTENING, "恢复待机:未听清唤醒词", turn_id=turn_id)
|
||||
continue
|
||||
raise
|
||||
remainder = _text_after_wake_word(wake_transcript, self.config.wake_word)
|
||||
if remainder is None:
|
||||
self._state(PipelineState.WAKE_LISTENING, "恢复待机:未命中唤醒词", turn_id=turn_id)
|
||||
continue
|
||||
self._state(PipelineState.SPEECH_DETECTING, "唤醒命中:请说出问题", turn_id=turn_id)
|
||||
if is_valid_transcript_text(remainder):
|
||||
return remainder
|
||||
user_segment = self._capture_segment(turn_id, state_message="录音中:正在听取问题")
|
||||
if isinstance(user_segment, ProviderError):
|
||||
return user_segment
|
||||
self._state(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",
|
||||
)
|
||||
return user_text
|
||||
|
||||
def _capture_segment(self, turn_id: int, *, state_message: str) -> AudioSegment | ProviderError:
|
||||
self.vad_recorder.reset()
|
||||
self.vad_recorder.provider.reset()
|
||||
self._state(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:
|
||||
result = self.vad_recorder.feed(frame)
|
||||
if isinstance(result, ProviderError):
|
||||
return result
|
||||
if isinstance(result, AudioSegment):
|
||||
return result
|
||||
|
||||
def _reply_to_user(self, user_text: str, turn_id: int) -> TurnResult:
|
||||
self.context.append_user(user_text)
|
||||
self._state(PipelineState.THINKING, f"思考中:{user_text}", 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,
|
||||
"live-runtime",
|
||||
"llm",
|
||||
),
|
||||
turn_id,
|
||||
)
|
||||
self.context.append_assistant(assistant_text)
|
||||
self._state(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._state(PipelineState.SPEAKING, "播放中:正在播报回复", turn_id=turn_id)
|
||||
segment = self.tts.synthesize(sentence)
|
||||
playback = self.transport.play_pcm(segment)
|
||||
if playback.error:
|
||||
raise playback.error
|
||||
|
||||
def _recover(self, error: ProviderError, turn_id: int) -> TurnResult:
|
||||
self.reporter.error(error.stage, error.code.value, error.message, turn_id=turn_id)
|
||||
self._state(PipelineState.ERROR_RECOVERING, "恢复待机:本轮已结束", turn_id=turn_id)
|
||||
self._state(PipelineState.WAKE_LISTENING, "恢复待机:可继续唤醒", turn_id=turn_id)
|
||||
return TurnResult(False, error=error, states=list(self._states))
|
||||
|
||||
def _state(self, state: PipelineState, message: str, *, turn_id: int) -> None:
|
||||
self._states.append(state)
|
||||
self.reporter.status(state.value, message, turn_id=turn_id)
|
||||
|
||||
|
||||
def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = None) -> LiveVoiceRuntime:
|
||||
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()
|
||||
return LiveVoiceRuntime(
|
||||
config=config,
|
||||
transport=SoundDeviceAudioTransport(output_device=config.audio_output_device),
|
||||
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=300, end_silence_ms=500, no_speech_timeout_ms=8000),
|
||||
stt=stt,
|
||||
llm=OpenAICompatibleLlmProvider(config),
|
||||
tts=tts,
|
||||
context=ConversationContext(
|
||||
max_messages=config.context_max_messages,
|
||||
max_chars=config.context_max_chars,
|
||||
),
|
||||
reporter=reporter,
|
||||
)
|
||||
|
||||
|
||||
def _text_after_wake_word(text: str, wake_word: str) -> str | None:
|
||||
compact_text = _compact(text)
|
||||
compact_wake = _compact(wake_word)
|
||||
index = compact_text.find(compact_wake)
|
||||
if index < 0:
|
||||
return None
|
||||
end = index + len(compact_wake)
|
||||
compact_remainder = compact_text[end:].strip(",,。.!!?? ")
|
||||
if not compact_remainder:
|
||||
return ""
|
||||
original = text.replace(" ", "")
|
||||
return original[-len(compact_remainder) :]
|
||||
|
||||
|
||||
def _compact(text: str) -> str:
|
||||
return "".join(ch for ch in text.strip() if not ch.isspace())
|
||||
Reference in New Issue
Block a user