[实时转写]:完成录音期间转写显示,包含partial事件、本地streaming STT和终端实时反馈
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user