function call fix and user transcription frames

This commit is contained in:
Kwindla Hultman Kramer
2024-10-07 16:52:26 -07:00
parent 5426891feb
commit 40dc546b81
3 changed files with 30 additions and 12 deletions

View File

@@ -5,10 +5,14 @@
# #
import asyncio import asyncio
import aiohttp
import os import os
import sys import sys
import aiohttp
from dotenv import load_dotenv
from loguru import logger
from runner import configure
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
@@ -16,19 +20,14 @@ from pipecat.processors.aggregators.openai_llm_context import (
OpenAILLMContext, OpenAILLMContext,
) )
from pipecat.services.openai_realtime_beta import ( from pipecat.services.openai_realtime_beta import (
InputAudioTranscription,
OpenAILLMServiceRealtimeBeta, OpenAILLMServiceRealtimeBeta,
TurnDetection,
SessionProperties, SessionProperties,
TurnDetection,
) )
from pipecat.transports.services.daily import DailyParams, DailyTransport from pipecat.transports.services.daily import DailyParams, DailyTransport
from pipecat.vad.silero import SileroVADAnalyzer from pipecat.vad.silero import SileroVADAnalyzer
from runner import configure
from loguru import logger
from dotenv import load_dotenv
load_dotenv(override=True) load_dotenv(override=True)
logger.remove(0) logger.remove(0)
@@ -76,7 +75,7 @@ async def main():
audio_in_sample_rate=24000, audio_in_sample_rate=24000,
audio_out_enabled=True, audio_out_enabled=True,
audio_out_sample_rate=24000, audio_out_sample_rate=24000,
transcription_enabled=True, transcription_enabled=False,
vad_enabled=True, vad_enabled=True,
vad_analyzer=SileroVADAnalyzer(), vad_analyzer=SileroVADAnalyzer(),
vad_audio_passthrough=True, vad_audio_passthrough=True,
@@ -84,6 +83,7 @@ async def main():
) )
session_properties = SessionProperties( session_properties = SessionProperties(
input_audio_transcription=InputAudioTranscription(),
turn_detection=TurnDetection(silence_duration_ms=1000), turn_detection=TurnDetection(silence_duration_ms=1000),
tools=tools, tools=tools,
instructions=""" instructions="""
@@ -107,7 +107,11 @@ Remember, your responses should be short. Just one or two sentences, usually.
llm = OpenAILLMServiceRealtimeBeta( llm = OpenAILLMServiceRealtimeBeta(
api_key=os.getenv("OPENAI_API_KEY"), session_properties=session_properties api_key=os.getenv("OPENAI_API_KEY"), session_properties=session_properties
) )
llm.register_function(None, fetch_weather_from_api)
# you can either register a single function for all function calls, or specific functions
# llm.register_function(None, fetch_weather_from_api)
llm.register_function("get_current_weather", fetch_weather_from_api)
context = OpenAILLMContext( context = OpenAILLMContext(
# [{"role": "user", "content": "What's the weather right now in San Francisco?"}], tools # [{"role": "user", "content": "What's the weather right now in San Francisco?"}], tools
[{"role": "user", "content": "Say 'hello'."}], [{"role": "user", "content": "Say 'hello'."}],

View File

@@ -1,2 +1,2 @@
from .events import InputAudioTranscription, SessionProperties, TurnDetection
from .llm_and_context import OpenAILLMServiceRealtimeBeta from .llm_and_context import OpenAILLMServiceRealtimeBeta
from .events import SessionProperties, TurnDetection

View File

@@ -40,6 +40,7 @@ from pipecat.services.openai import (
OpenAIContextAggregatorPair, OpenAIContextAggregatorPair,
OpenAIUserContextAggregator, OpenAIUserContextAggregator,
) )
from pipecat.utils.time import time_now_iso8601
from . import events from . import events
@@ -118,6 +119,7 @@ class OpenAILLMServiceRealtimeBeta(LLMService):
base_url="wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview-2024-10-01", base_url="wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview-2024-10-01",
session_properties: events.SessionProperties = events.SessionProperties(), session_properties: events.SessionProperties = events.SessionProperties(),
start_audio_paused: bool = False, start_audio_paused: bool = False,
send_transcription_frames: bool = True,
**kwargs, **kwargs,
): ):
super().__init__(base_url=base_url, **kwargs) super().__init__(base_url=base_url, **kwargs)
@@ -126,6 +128,7 @@ class OpenAILLMServiceRealtimeBeta(LLMService):
self._session_properties = session_properties self._session_properties = session_properties
self._audio_input_paused = start_audio_paused self._audio_input_paused = start_audio_paused
self._send_transcription_frames = send_transcription_frames
self._websocket = None self._websocket = None
self._receive_task = None self._receive_task = None
self._context = None self._context = None
@@ -237,6 +240,11 @@ class OpenAILLMServiceRealtimeBeta(LLMService):
self._context.add_message({"role": "user", "content": evt.transcript}) self._context.add_message({"role": "user", "content": evt.transcript})
else: else:
logger.error("Context is None, cannot add message") logger.error("Context is None, cannot add message")
if self._send_transcription_frames:
await self.push_frame(
# no way to get a language code?
TranscriptionFrame(evt.transcript, "", time_now_iso8601())
)
elif evt.type == "response.output_item.added": elif evt.type == "response.output_item.added":
# todo: think about adding a frame for this (generally, in Pipecat/RTVI), as # todo: think about adding a frame for this (generally, in Pipecat/RTVI), as
# it could be useful for managing UI state # it could be useful for managing UI state
@@ -306,7 +314,13 @@ class OpenAILLMServiceRealtimeBeta(LLMService):
if self.has_function(function_name): if self.has_function(function_name):
run_llm = index == total_items - 1 run_llm = index == total_items - 1
if function_name in self._callbacks.keys(): if function_name in self._callbacks.keys():
f = self._callbacks[function_name] await self.call_function(
context=self._context,
tool_call_id=tool_id,
function_name=function_name,
arguments=arguments,
run_llm=run_llm,
)
elif None in self._callbacks.keys(): elif None in self._callbacks.keys():
await self.call_function( await self.call_function(
context=self._context, context=self._context,