[实时转写]:完成录音期间转写显示,包含partial事件、本地streaming STT和终端实时反馈

This commit is contained in:
mkbk
2026-06-17 22:16:56 +08:00
parent a64bb86da4
commit a77a172412
16 changed files with 331 additions and 17 deletions
+54 -2
View File
@@ -18,6 +18,7 @@ from .events import (
STANDBY_RESUMED,
STT_STARTED,
TRANSCRIPT_FINAL,
TRANSCRIPT_PARTIAL,
TTS_STARTED,
WAKE_DETECTED,
WAKE_LISTENING,
@@ -25,8 +26,16 @@ from .events import (
PipelineEventBus,
dispatch_pipeline_event,
)
from .models import AudioSegment, ErrorCode, PipelineState, ProviderError
from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, WakeWordProvider
from .models import AudioFrame, AudioSegment, ErrorCode, PipelineState, ProviderError
from .protocols import (
AudioTransport,
LlmProvider,
RealtimeSttProvider,
RealtimeTranscriptSession,
SttProvider,
TtsProvider,
WakeWordProvider,
)
from .stt import is_valid_transcript_text
from .tts import SentenceBuffer
from .vad import VadRecorder
@@ -69,6 +78,7 @@ class TurnController:
wakeword: WakeWordProvider,
vad_recorder: VadRecorder,
stt: SttProvider,
realtime_stt: RealtimeSttProvider | None,
llm: LlmProvider,
tts: TtsProvider,
ack_tts: TtsProvider,
@@ -81,6 +91,7 @@ class TurnController:
self.wakeword = wakeword
self.vad_recorder = vad_recorder
self.stt = stt
self.realtime_stt = realtime_stt
self.llm = llm
self.tts = tts
self.ack_tts = ack_tts
@@ -146,6 +157,7 @@ class TurnController:
def _capture_segment(self, turn_id: int, *, state_message: str) -> AudioSegment | ProviderError:
self.vad_recorder.reset()
self.vad_recorder.provider.reset()
realtime_session = self._start_realtime_transcript()
self._event(CAPTURE_STARTED, PipelineState.RECORDING, state_message, turn_id=turn_id)
while True:
frames = self.transport.read_frames(timeout_ms=100)
@@ -158,7 +170,11 @@ class TurnController:
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:
self._emit_realtime_transcript(realtime_session, frame, turn_id)
if isinstance(result, AudioSegment):
if realtime_session is not None:
self._finish_realtime_transcript(realtime_session, turn_id)
self._event(
SPEECH_ENDED,
PipelineState.RECORDING,
@@ -168,6 +184,37 @@ class TurnController:
)
return result
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,
) -> None:
transcript = realtime_session.accept_frame(frame)
if transcript is None:
return
text = transcript.normalized_text
if not is_valid_transcript_text(text):
return
self._event(TRANSCRIPT_PARTIAL, PipelineState.RECORDING, "", turn_id=turn_id, payload={"text": text})
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:
@@ -263,6 +310,7 @@ class VoiceAssistantPipeline:
llm: LlmProvider,
tts: TtsProvider,
context: ConversationContext,
realtime_stt: RealtimeSttProvider | None = None,
ack_tts: TtsProvider | None = None,
reporter: RuntimeReporter | None = None,
event_bus: PipelineEventBus | None = None,
@@ -273,6 +321,7 @@ class VoiceAssistantPipeline:
self.wakeword = wakeword
self.vad_recorder = vad_recorder
self.stt = stt
self.realtime_stt = realtime_stt
self.llm = llm
self.tts = tts
self.ack_tts = ack_tts or tts
@@ -288,6 +337,7 @@ class VoiceAssistantPipeline:
wakeword=wakeword,
vad_recorder=vad_recorder,
stt=stt,
realtime_stt=realtime_stt,
llm=llm,
tts=tts,
ack_tts=self.ack_tts,
@@ -300,6 +350,8 @@ class VoiceAssistantPipeline:
self.wakeword.load()
self.vad_recorder.provider.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()