Merge pull request #195 from pipecat-ai/aleix/function-calling-move-to-llmservice

function calling move to LLMService
This commit is contained in:
Aleix Conchillo Flaqué
2024-06-01 02:36:29 +08:00
committed by GitHub
5 changed files with 64 additions and 62 deletions

View File

@@ -9,6 +9,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
### Added ### Added
- Added function calling (LLMService.register_function()). This will allow the
LLM to call functions you have registered when needed. For example, if you
register a function to get the weather in Los Angeles and ask the LLM about
the weather in Los Angeles, the LLM will call your function.
See https://platform.openai.com/docs/guides/function-calling
- Added Cartesia TTS support (https://cartesia.ai/) - Added Cartesia TTS support (https://cartesia.ai/)
### Fixed ### Fixed

View File

@@ -7,9 +7,9 @@
import asyncio import asyncio
import aiohttp import aiohttp
import os import os
import json
import sys import sys
from pipecat.frames.frames import TextFrame
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 PipelineTask
@@ -17,18 +17,13 @@ from pipecat.processors.aggregators.llm_response import (
LLMAssistantContextAggregator, LLMAssistantContextAggregator,
LLMUserContextAggregator, LLMUserContextAggregator,
) )
from pipecat.services.openai import OpenAILLMContext
from pipecat.processors.logger import FrameLogger from pipecat.processors.logger import FrameLogger
from pipecat.services.elevenlabs import ElevenLabsTTSService from pipecat.services.elevenlabs import ElevenLabsTTSService
from pipecat.services.openai import OpenAILLMService from pipecat.services.openai import OpenAILLMContext, OpenAILLMService
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 openai.types.chat import (
ChatCompletionToolParam, from openai.types.chat import ChatCompletionToolParam
)
from pipecat.frames.frames import (
TextFrame
)
from runner import configure from runner import configure

View File

@@ -1,11 +1,15 @@
#
# Copyright (c) 2024, Daily
#
# SPDX-License-Identifier: BSD 2-Clause License
#
import asyncio import asyncio
import aiohttp import aiohttp
import copy
import json
import os import os
import re
import sys import sys
import wave import wave
from typing import List from typing import List
from openai._types import NotGiven, NOT_GIVEN from openai._types import NotGiven, NOT_GIVEN
@@ -14,23 +18,18 @@ from openai.types.chat import (
ChatCompletionToolParam, ChatCompletionToolParam,
) )
from pipecat.frames.frames import AudioRawFrame
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.llm_response import LLMUserContextAggregator, LLMAssistantContextAggregator from pipecat.processors.aggregators.llm_response import LLMUserContextAggregator, LLMAssistantContextAggregator
from pipecat.processors.logger import FrameLogger from pipecat.processors.logger import FrameLogger
from pipecat.frames.frames import ( from pipecat.processors.frame_processor import FrameDirection
Frame,
LLMMessagesFrame,
AudioRawFrame,
)
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
from pipecat.services.elevenlabs import ElevenLabsTTSService from pipecat.services.elevenlabs import ElevenLabsTTSService
from pipecat.services.openai import OpenAILLMService from pipecat.services.openai import OpenAILLMContext, OpenAILLMContextFrame, OpenAILLMService
from pipecat.services.ai_services import AIService from pipecat.services.ai_services import AIService
from pipecat.transports.services.daily import DailyParams, DailyTranscriptionSettings, DailyTransport from pipecat.transports.services.daily import DailyParams, DailyTransport
from pipecat.vad.silero import SileroVADAnalyzer from pipecat.vad.silero import SileroVADAnalyzer
from pipecat.services.openai import OpenAILLMContext, OpenAILLMContextFrame
from runner import configure from runner import configure
@@ -242,7 +241,6 @@ class IntakeProcessor:
self._context.add_message( self._context.add_message(
{"role": "system", "content": "Finally, ask the user the reason for their doctor visit today. Once they answer, call the list_visit_reasons function."}) {"role": "system", "content": "Finally, ask the user the reason for their doctor visit today. Once they answer, call the list_visit_reasons function."})
await llm.process_frame(OpenAILLMContextFrame(self._context), FrameDirection.DOWNSTREAM) await llm.process_frame(OpenAILLMContextFrame(self._context), FrameDirection.DOWNSTREAM)
pass
async def start_visit_reasons(self, llm): async def start_visit_reasons(self, llm):
print("!!! doing start visit reasons") print("!!! doing start visit reasons")
@@ -251,7 +249,6 @@ class IntakeProcessor:
self._context.add_message({"role": "system", self._context.add_message({"role": "system",
"content": "Now, thank the user and end the conversation."}) "content": "Now, thank the user and end the conversation."})
await llm.process_frame(OpenAILLMContextFrame(self._context), FrameDirection.DOWNSTREAM) await llm.process_frame(OpenAILLMContextFrame(self._context), FrameDirection.DOWNSTREAM)
pass
async def save_data(self, llm, args): async def save_data(self, llm, args):
logger.info(f"!!! Saving data: {args}") logger.info(f"!!! Saving data: {args}")
@@ -305,12 +302,10 @@ async def main(room_url: str, token):
model="gpt-4o") model="gpt-4o")
messages = [] messages = []
context = OpenAILLMContext( context = OpenAILLMContext(messages=messages)
messages=messages,
)
user_context = LLMUserContextAggregator(context) user_context = LLMUserContextAggregator(context)
assistant_context = LLMAssistantContextAggregator(context) assistant_context = LLMAssistantContextAggregator(context)
# checklist = ChecklistProcessor(context, llm)
intake = IntakeProcessor(context, llm) intake = IntakeProcessor(context, llm)
llm.register_function("verify_birthday", intake.verify_birthday) llm.register_function("verify_birthday", intake.verify_birthday)
llm.register_function( llm.register_function(
@@ -329,19 +324,20 @@ async def main(room_url: str, token):
"list_visit_reasons", "list_visit_reasons",
intake.save_data, intake.save_data,
start_callback=intake.start_visit_reasons) start_callback=intake.start_visit_reasons)
fl = FrameLogger("LLM Output") fl = FrameLogger("LLM Output")
pipeline = Pipeline([ pipeline = Pipeline([
transport.input(), transport.input(), # Transport input
user_context, user_context, # User responses
llm, llm, # LLM
fl, fl, # Frame logger
tts, tts, # TTS
transport.output(), transport.output(), # Transport output
assistant_context, assistant_context, # Assistant responses
]) ])
task = PipelineTask(pipeline, allow_interruptions=False) task = PipelineTask(pipeline, PipelineParams(allow_interruptions=False))
@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):

View File

@@ -43,6 +43,31 @@ class LLMService(AIService):
def __init__(self): def __init__(self):
super().__init__() super().__init__()
self._callbacks = {}
self._start_callbacks = {}
# TODO-CB: callback function type
def register_function(self, function_name: str, callback, start_callback=None):
self._callbacks[function_name] = callback
if start_callback:
self._start_callbacks[function_name] = start_callback
def unregister_function(self, function_name: str):
del self._callbacks[function_name]
if self._start_callbacks[function_name]:
del self._start_callbacks[function_name]
def has_function(self, function_name: str):
return function_name in self._callbacks.keys()
async def call_function(self, function_name: str, args):
if function_name in self._callbacks.keys():
return await self._callbacks[function_name](self, args)
return None
async def call_start_function(self, function_name: str):
if function_name in self._start_callbacks.keys():
await self._start_callbacks[function_name](self)
class TTSService(AIService): class TTSService(AIService):

View File

@@ -29,12 +29,7 @@ from pipecat.frames.frames import (
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext, OpenAILLMContextFrame from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext, OpenAILLMContextFrame
from pipecat.processors.frame_processor import FrameDirection from pipecat.processors.frame_processor import FrameDirection
from pipecat.services.ai_services import LLMService, ImageGenService from pipecat.services.ai_services import LLMService, ImageGenService
from openai.types.chat import (
ChatCompletionSystemMessageParam,
ChatCompletionFunctionMessageParam,
ChatCompletionToolParam,
ChatCompletionUserMessageParam,
)
from loguru import logger from loguru import logger
try: try:
@@ -43,7 +38,9 @@ try:
from openai.types.chat import ( from openai.types.chat import (
ChatCompletion, ChatCompletion,
ChatCompletionChunk, ChatCompletionChunk,
ChatCompletionFunctionMessageParam,
ChatCompletionMessageParam, ChatCompletionMessageParam,
ChatCompletionToolParam
) )
except ModuleNotFoundError as e: except ModuleNotFoundError as e:
logger.error(f"Exception: {e}") logger.error(f"Exception: {e}")
@@ -70,23 +67,10 @@ class BaseOpenAILLMService(LLMService):
super().__init__() super().__init__()
self._model: str = model self._model: str = model
self._client = self.create_client(api_key=api_key, base_url=base_url) self._client = self.create_client(api_key=api_key, base_url=base_url)
self._callbacks = {}
self._start_callbacks = {}
def create_client(self, api_key=None, base_url=None): def create_client(self, api_key=None, base_url=None):
return AsyncOpenAI(api_key=api_key, base_url=base_url) return AsyncOpenAI(api_key=api_key, base_url=base_url)
# TODO-CB: callback function type
def register_function(self, function_name, callback, start_callback=None):
self._callbacks[function_name] = callback
if start_callback:
self._start_callbacks[function_name] = start_callback
def unregister_function(self, function_name):
del self._callbacks[function_name]
if self._start_callbacks[function_name]:
del self._start_callbacks[function_name]
async def _stream_chat_completions( async def _stream_chat_completions(
self, context: OpenAILLMContext self, context: OpenAILLMContext
) -> AsyncStream[ChatCompletionChunk]: ) -> AsyncStream[ChatCompletionChunk]:
@@ -159,10 +143,7 @@ class BaseOpenAILLMService(LLMService):
if tool_call.function and tool_call.function.name: if tool_call.function and tool_call.function.name:
function_name += tool_call.function.name function_name += tool_call.function.name
tool_call_id = tool_call.id tool_call_id = tool_call.id
# only send a function start frame if we're not handling the function call await self.call_start_function(function_name)
if function_name in self._callbacks.keys():
if function_name in self._start_callbacks.keys():
await self._start_callbacks[function_name](self)
if tool_call.function and tool_call.function.arguments: if tool_call.function and tool_call.function.arguments:
# Keep iterating through the response to collect all the argument fragments # Keep iterating through the response to collect all the argument fragments
arguments += tool_call.function.arguments arguments += tool_call.function.arguments
@@ -176,9 +157,8 @@ class BaseOpenAILLMService(LLMService):
# the context, and re-prompt to get a chat answer. If we don't have a registered # the context, and re-prompt to get a chat answer. If we don't have a registered
# handler, raise an exception. # handler, raise an exception.
if function_name and arguments: if function_name and arguments:
if function_name in self._callbacks.keys(): if self.has_function(function_name):
await self._handle_function_call(context, tool_call_id, function_name, arguments) await self._handle_function_call(context, tool_call_id, function_name, arguments)
else: else:
raise OpenAIUnhandledFunctionException( raise OpenAIUnhandledFunctionException(
f"The LLM tried to call a function named '{function_name}', but there isn't a callback registered for that function.") f"The LLM tried to call a function named '{function_name}', but there isn't a callback registered for that function.")
@@ -191,7 +171,7 @@ class BaseOpenAILLMService(LLMService):
arguments arguments
): ):
arguments = json.loads(arguments) arguments = json.loads(arguments)
result = await self._callbacks[function_name](self, arguments) result = await self.call_function(function_name, arguments)
arguments = json.dumps(arguments) arguments = json.dumps(arguments)
if isinstance(result, (str, dict)): if isinstance(result, (str, dict)):
# Handle it in "full magic mode" # Handle it in "full magic mode"