This is working but RTVI still gets the transcript

FEAT: Example of muting LLMMessages to LLM.
This commit is contained in:
James Hush
2025-03-20 12:10:11 +08:00
parent f31e77c4f6
commit 1a237ddae8

View File

@@ -29,18 +29,26 @@ from runner import configure
from pipecat.audio.vad.silero import SileroVADAnalyzer from pipecat.audio.vad.silero import SileroVADAnalyzer
from pipecat.frames.frames import ( from pipecat.frames.frames import (
BotSpeakingFrame,
BotStartedSpeakingFrame, BotStartedSpeakingFrame,
BotStoppedSpeakingFrame, BotStoppedSpeakingFrame,
Frame, Frame,
InterimTranscriptionFrame,
LLMMessagesFrame,
LLMTextFrame,
OutputImageRawFrame, OutputImageRawFrame,
SpriteFrame, SpriteFrame,
STTMuteFrame,
TranscriptionFrame,
) )
from pipecat.pipeline.pipeline import Pipeline from pipecat.pipeline.pipeline import Pipeline
from pipecat.pipeline.runner import PipelineRunner from pipecat.pipeline.runner import PipelineRunner
from pipecat.pipeline.task import PipelineParams, PipelineTask from pipecat.pipeline.task import PipelineParams, PipelineTask
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext
from pipecat.processors.filters.stt_mute_filter import STTMuteConfig, STTMuteFilter, STTMuteStrategy
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
from pipecat.processors.frameworks.rtvi import RTVIConfig, RTVIObserver, RTVIProcessor from pipecat.processors.frameworks.rtvi import RTVIConfig, RTVIObserver, RTVIProcessor
from pipecat.services.deepgram import DeepgramSTTService
from pipecat.services.elevenlabs import ElevenLabsTTSService from pipecat.services.elevenlabs import ElevenLabsTTSService
from pipecat.services.openai import OpenAILLMService from pipecat.services.openai import OpenAILLMService
from pipecat.transports.services.daily import DailyParams, DailyTransport from pipecat.transports.services.daily import DailyParams, DailyTransport
@@ -49,6 +57,49 @@ load_dotenv(override=True)
logger.remove(0) logger.remove(0)
logger.add(sys.stderr, level="DEBUG") logger.add(sys.stderr, level="DEBUG")
class TranscriptionMuteProcessor(FrameProcessor):
"""Takes in STTMuteFrame and mutes TranscriptionFrame based on its content."""
def __init__(self, **kwargs):
super().__init__(**kwargs)
self._is_muted = False
async def process_frame(self, frame: Frame, direction: FrameDirection):
"""Process incoming frames and mute TranscriptionFrame based on STTMuteFrame content.
Args:
frame: The incoming frame to process
direction: The direction of frame flow in the pipeline
"""
await super().process_frame(frame, direction)
if isinstance(frame, STTMuteFrame):
logger.debug(f"TranscriptionMuteProcessor: Mute state: {frame.mute}")
self._is_muted = frame.mute
if isinstance(
frame,
(
TranscriptionFrame,
InterimTranscriptionFrame,
LLMMessagesFrame,
),
):
# Only pass frames when not muted
if not self._is_muted:
await self.push_frame(frame, direction)
else:
logger.debug(
f"{frame.__class__.__name__} suppressed - Transcription STT currently muted"
)
else:
# Pass all other frames through
if not isinstance(frame, BotSpeakingFrame):
logger.debug(f"+++ TranscriptionMuteProcessor pushing: {frame}")
await self.push_frame(frame, direction)
sprites = [] sprites = []
script_dir = os.path.dirname(__file__) script_dir = os.path.dirname(__file__)
@@ -128,7 +179,8 @@ async def main():
camera_out_height=576, camera_out_height=576,
vad_enabled=True, vad_enabled=True,
vad_analyzer=SileroVADAnalyzer(), vad_analyzer=SileroVADAnalyzer(),
transcription_enabled=True, vad_audio_passthrough=True,
# transcription_enabled=True,
# #
# Spanish # Spanish
# #
@@ -183,9 +235,20 @@ async def main():
# #
rtvi = RTVIProcessor(config=RTVIConfig(config=[])) rtvi = RTVIProcessor(config=RTVIConfig(config=[]))
stt = DeepgramSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
stt_mute_processor = STTMuteFilter(
config=STTMuteConfig(strategies={STTMuteStrategy.ALWAYS}),
)
transcription_mute_processor = TranscriptionMuteProcessor()
pipeline = Pipeline( pipeline = Pipeline(
[ [
transport.input(), transport.input(),
stt,
stt_mute_processor,
transcription_mute_processor,
rtvi, rtvi,
context_aggregator.user(), context_aggregator.user(),
llm, llm,
@@ -213,7 +276,7 @@ async def main():
@transport.event_handler("on_first_participant_joined") @transport.event_handler("on_first_participant_joined")
async def on_first_participant_joined(transport, participant): async def on_first_participant_joined(transport, participant):
await transport.capture_participant_transcription(participant["id"]) # await transport.capture_participant_transcription(participant["id"])
await task.queue_frames([context_aggregator.user().get_context_frame()]) await task.queue_frames([context_aggregator.user().get_context_frame()])
@transport.event_handler("on_participant_left") @transport.event_handler("on_participant_left")