examples(websocket-server): allow interruptions

This commit is contained in:
Aleix Conchillo Flaqué
2024-11-04 13:05:02 -08:00
parent 0ac9e2dd3f
commit 602915ae18

View File

@@ -12,7 +12,7 @@ from pipecat.audio.vad.silero import SileroVADAnalyzer
from pipecat.frames.frames import LLMMessagesFrame from pipecat.frames.frames import LLMMessagesFrame
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 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.services.cartesia import CartesiaTTSService from pipecat.services.cartesia import CartesiaTTSService
from pipecat.services.deepgram import DeepgramSTTService from pipecat.services.deepgram import DeepgramSTTService
@@ -35,6 +35,7 @@ logger.add(sys.stderr, level="DEBUG")
async def main(): async def main():
transport = WebsocketServerTransport( transport = WebsocketServerTransport(
params=WebsocketServerParams( params=WebsocketServerParams(
audio_out_sample_rate=16000,
audio_out_enabled=True, audio_out_enabled=True,
add_wav_header=True, add_wav_header=True,
vad_enabled=True, vad_enabled=True,
@@ -50,6 +51,7 @@ async def main():
tts = CartesiaTTSService( tts = CartesiaTTSService(
api_key=os.getenv("CARTESIA_API_KEY"), api_key=os.getenv("CARTESIA_API_KEY"),
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
sample_rate=16000,
) )
messages = [ messages = [
@@ -74,7 +76,7 @@ async def main():
] ]
) )
task = PipelineTask(pipeline) task = PipelineTask(pipeline, params=PipelineParams(allow_interruptions=True))
@transport.event_handler("on_client_connected") @transport.event_handler("on_client_connected")
async def on_client_connected(transport, client): async def on_client_connected(transport, client):