Remove remaining usage of OpenAILLMContext throughout the codebase in favor of LLMContext, except for:
- Usage in classes that are already deprecated - Usage related to realtime LLMs, which don't yet support `LLMContext` - Usage in (soon-to-be-deprecated) code paths related to `OpenAILLMContext` itself and associated machinery
This commit is contained in:
@@ -34,7 +34,8 @@ from pipecat.frames.frames import EndTaskFrame, LLMRunFrame, OutputImageRawFrame
|
|||||||
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.llm_context import LLMContext
|
||||||
|
from pipecat.processors.aggregators.llm_response_universal import LLMContextAggregatorPair
|
||||||
from pipecat.processors.audio.audio_buffer_processor import AudioBufferProcessor
|
from pipecat.processors.audio.audio_buffer_processor import AudioBufferProcessor
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.runner.types import RunnerArguments
|
from pipecat.runner.types import RunnerArguments
|
||||||
@@ -283,8 +284,8 @@ async def run_eval_pipeline(
|
|||||||
},
|
},
|
||||||
]
|
]
|
||||||
|
|
||||||
context = OpenAILLMContext(messages, tools)
|
context = LLMContext(messages, tools)
|
||||||
context_aggregator = llm.create_context_aggregator(context)
|
context_aggregator = LLMContextAggregatorPair(context)
|
||||||
|
|
||||||
audio_buffer = AudioBufferProcessor()
|
audio_buffer = AudioBufferProcessor()
|
||||||
|
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ LLM processing, and text-to-speech components in conversational AI pipelines.
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
|
from abc import abstractmethod
|
||||||
from typing import Any, Dict, List, Literal, Optional, Set
|
from typing import Any, Dict, List, Literal, Optional, Set
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -169,6 +170,11 @@ class LLMContextAggregator(FrameProcessor):
|
|||||||
"""Reset the aggregation state."""
|
"""Reset the aggregation state."""
|
||||||
self._aggregation = ""
|
self._aggregation = ""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
async def push_aggregation(self):
|
||||||
|
"""Push the current aggregation downstream."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
class LLMUserAggregator(LLMContextAggregator):
|
class LLMUserAggregator(LLMContextAggregator):
|
||||||
"""User LLM aggregator that processes speech-to-text transcriptions.
|
"""User LLM aggregator that processes speech-to-text transcriptions.
|
||||||
@@ -301,7 +307,7 @@ class LLMUserAggregator(LLMContextAggregator):
|
|||||||
frame = LLMContextFrame(self._context)
|
frame = LLMContextFrame(self._context)
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
async def _push_aggregation(self):
|
async def push_aggregation(self):
|
||||||
"""Push the current aggregation based on interruption strategies and conditions."""
|
"""Push the current aggregation based on interruption strategies and conditions."""
|
||||||
if len(self._aggregation) > 0:
|
if len(self._aggregation) > 0:
|
||||||
if self.interruption_strategies and self._bot_speaking:
|
if self.interruption_strategies and self._bot_speaking:
|
||||||
@@ -392,7 +398,7 @@ class LLMUserAggregator(LLMContextAggregator):
|
|||||||
# pushing the aggregation as we will probably get a final transcription.
|
# pushing the aggregation as we will probably get a final transcription.
|
||||||
if len(self._aggregation) > 0:
|
if len(self._aggregation) > 0:
|
||||||
if not self._seen_interim_results:
|
if not self._seen_interim_results:
|
||||||
await self._push_aggregation()
|
await self.push_aggregation()
|
||||||
# Handles the case where both the user and the bot are not speaking,
|
# Handles the case where both the user and the bot are not speaking,
|
||||||
# and the bot was previously speaking before the user interruption.
|
# and the bot was previously speaking before the user interruption.
|
||||||
# So in this case we are resetting the aggregation timer
|
# So in this case we are resetting the aggregation timer
|
||||||
@@ -471,7 +477,7 @@ class LLMUserAggregator(LLMContextAggregator):
|
|||||||
await self._maybe_emulate_user_speaking()
|
await self._maybe_emulate_user_speaking()
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
if not self._user_speaking:
|
if not self._user_speaking:
|
||||||
await self._push_aggregation()
|
await self.push_aggregation()
|
||||||
|
|
||||||
# If we are emulating VAD we still need to send the user stopped
|
# If we are emulating VAD we still need to send the user stopped
|
||||||
# speaking frame.
|
# speaking frame.
|
||||||
@@ -607,12 +613,12 @@ class LLMAssistantAggregator(LLMContextAggregator):
|
|||||||
elif isinstance(frame, UserImageRawFrame) and frame.request and frame.request.tool_call_id:
|
elif isinstance(frame, UserImageRawFrame) and frame.request and frame.request.tool_call_id:
|
||||||
await self._handle_user_image_frame(frame)
|
await self._handle_user_image_frame(frame)
|
||||||
elif isinstance(frame, BotStoppedSpeakingFrame):
|
elif isinstance(frame, BotStoppedSpeakingFrame):
|
||||||
await self._push_aggregation()
|
await self.push_aggregation()
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
else:
|
else:
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
async def _push_aggregation(self):
|
async def push_aggregation(self):
|
||||||
"""Push the current assistant aggregation with timestamp."""
|
"""Push the current assistant aggregation with timestamp."""
|
||||||
if not self._aggregation:
|
if not self._aggregation:
|
||||||
return
|
return
|
||||||
@@ -644,7 +650,7 @@ class LLMAssistantAggregator(LLMContextAggregator):
|
|||||||
await self.push_context_frame(FrameDirection.UPSTREAM)
|
await self.push_context_frame(FrameDirection.UPSTREAM)
|
||||||
|
|
||||||
async def _handle_interruptions(self, frame: InterruptionFrame):
|
async def _handle_interruptions(self, frame: InterruptionFrame):
|
||||||
await self._push_aggregation()
|
await self.push_aggregation()
|
||||||
self._started = 0
|
self._started = 0
|
||||||
await self.reset()
|
await self.reset()
|
||||||
|
|
||||||
@@ -778,7 +784,7 @@ class LLMAssistantAggregator(LLMContextAggregator):
|
|||||||
text=frame.request.context,
|
text=frame.request.context,
|
||||||
)
|
)
|
||||||
|
|
||||||
await self._push_aggregation()
|
await self.push_aggregation()
|
||||||
await self.push_context_frame(FrameDirection.UPSTREAM)
|
await self.push_context_frame(FrameDirection.UPSTREAM)
|
||||||
|
|
||||||
async def _handle_llm_start(self, _: LLMFullResponseStartFrame):
|
async def _handle_llm_start(self, _: LLMFullResponseStartFrame):
|
||||||
@@ -786,7 +792,7 @@ class LLMAssistantAggregator(LLMContextAggregator):
|
|||||||
|
|
||||||
async def _handle_llm_end(self, _: LLMFullResponseEndFrame):
|
async def _handle_llm_end(self, _: LLMFullResponseEndFrame):
|
||||||
self._started -= 1
|
self._started -= 1
|
||||||
await self._push_aggregation()
|
await self.push_aggregation()
|
||||||
|
|
||||||
async def _handle_text(self, frame: TextFrame):
|
async def _handle_text(self, frame: TextFrame):
|
||||||
if not self._started:
|
if not self._started:
|
||||||
|
|||||||
@@ -12,14 +12,14 @@ in conversational pipelines.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from pipecat.frames.frames import TextFrame
|
from pipecat.frames.frames import TextFrame
|
||||||
from pipecat.processors.aggregators.llm_response import LLMUserContextAggregator
|
from pipecat.processors.aggregators.llm_context import LLMContext
|
||||||
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext
|
from pipecat.processors.aggregators.llm_response_universal import LLMUserAggregator
|
||||||
|
|
||||||
|
|
||||||
class UserResponseAggregator(LLMUserContextAggregator):
|
class UserResponseAggregator(LLMUserAggregator):
|
||||||
"""Aggregates user responses into TextFrame objects.
|
"""Aggregates user responses into TextFrame objects.
|
||||||
|
|
||||||
This aggregator extends LLMUserContextAggregator to specifically handle
|
This aggregator extends LLMUserAggregator to specifically handle
|
||||||
user input by collecting text responses and outputting them as TextFrame
|
user input by collecting text responses and outputting them as TextFrame
|
||||||
objects when the aggregation is complete.
|
objects when the aggregation is complete.
|
||||||
"""
|
"""
|
||||||
@@ -28,9 +28,9 @@ class UserResponseAggregator(LLMUserContextAggregator):
|
|||||||
"""Initialize the user response aggregator.
|
"""Initialize the user response aggregator.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
**kwargs: Additional arguments passed to parent LLMUserContextAggregator.
|
**kwargs: Additional arguments passed to parent LLMUserAggregator.
|
||||||
"""
|
"""
|
||||||
super().__init__(context=OpenAILLMContext(), **kwargs)
|
super().__init__(context=LLMContext(), **kwargs)
|
||||||
|
|
||||||
async def push_aggregation(self):
|
async def push_aggregation(self):
|
||||||
"""Push the aggregated user response as a TextFrame.
|
"""Push the aggregated user response as a TextFrame.
|
||||||
|
|||||||
@@ -12,14 +12,12 @@ from dotenv import load_dotenv
|
|||||||
|
|
||||||
from pipecat.adapters.schemas.function_schema import FunctionSchema
|
from pipecat.adapters.schemas.function_schema import FunctionSchema
|
||||||
from pipecat.adapters.schemas.tools_schema import ToolsSchema
|
from pipecat.adapters.schemas.tools_schema import ToolsSchema
|
||||||
|
from pipecat.frames.frames import LLMContextFrame
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
from pipecat.processors.aggregators.openai_llm_context import (
|
from pipecat.processors.aggregators.llm_context import LLMContext
|
||||||
OpenAILLMContext,
|
|
||||||
OpenAILLMContextFrame,
|
|
||||||
)
|
|
||||||
from pipecat.services.anthropic.llm import AnthropicLLMService
|
from pipecat.services.anthropic.llm import AnthropicLLMService
|
||||||
from pipecat.services.google.llm import GoogleLLMService
|
from pipecat.services.google.llm import GoogleLLMService
|
||||||
from pipecat.services.llm_service import LLMService
|
from pipecat.services.llm_service import FunctionCallParams, LLMService
|
||||||
from pipecat.services.openai.llm import OpenAILLMService
|
from pipecat.services.openai.llm import OpenAILLMService
|
||||||
from pipecat.tests.utils import run_test
|
from pipecat.tests.utils import run_test
|
||||||
|
|
||||||
@@ -48,8 +46,13 @@ def standard_tools() -> ToolsSchema:
|
|||||||
|
|
||||||
|
|
||||||
async def _test_llm_function_calling(llm: LLMService):
|
async def _test_llm_function_calling(llm: LLMService):
|
||||||
# Create an AsyncMock for the function
|
# Create a mock weather function
|
||||||
mock_fetch_weather = AsyncMock()
|
call_count = 0
|
||||||
|
|
||||||
|
async def mock_fetch_weather(params: FunctionCallParams):
|
||||||
|
nonlocal call_count
|
||||||
|
call_count += 1
|
||||||
|
pass
|
||||||
|
|
||||||
llm.register_function(None, mock_fetch_weather)
|
llm.register_function(None, mock_fetch_weather)
|
||||||
|
|
||||||
@@ -60,21 +63,19 @@ async def _test_llm_function_calling(llm: LLMService):
|
|||||||
},
|
},
|
||||||
{"role": "user", "content": " How is the weather today in San Francisco, California?"},
|
{"role": "user", "content": " How is the weather today in San Francisco, California?"},
|
||||||
]
|
]
|
||||||
context = OpenAILLMContext(messages, standard_tools())
|
context = LLMContext(messages, standard_tools())
|
||||||
# This is done by default inside the create_context_aggregator
|
|
||||||
context.set_llm_adapter(llm.get_llm_adapter())
|
|
||||||
|
|
||||||
pipeline = Pipeline([llm])
|
pipeline = Pipeline([llm])
|
||||||
|
|
||||||
frames_to_send = [OpenAILLMContextFrame(context)]
|
frames_to_send = [LLMContextFrame(context)]
|
||||||
await run_test(
|
await run_test(
|
||||||
pipeline,
|
pipeline,
|
||||||
frames_to_send=frames_to_send,
|
frames_to_send=frames_to_send,
|
||||||
expected_down_frames=None,
|
expected_down_frames=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Assert that the mock function was called
|
# Assert that the weather function was called once
|
||||||
mock_fetch_weather.assert_called_once()
|
assert call_count == 1
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.skipif(os.getenv("OPENAI_API_KEY") is None, reason="OPENAI_API_KEY is not set")
|
@pytest.mark.skipif(os.getenv("OPENAI_API_KEY") is None, reason="OPENAI_API_KEY is not set")
|
||||||
|
|||||||
@@ -10,24 +10,21 @@ from langchain.prompts import ChatPromptTemplate
|
|||||||
from langchain_core.language_models import FakeStreamingListLLM
|
from langchain_core.language_models import FakeStreamingListLLM
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
|
LLMContextAssistantTimestampFrame,
|
||||||
|
LLMContextFrame,
|
||||||
LLMFullResponseEndFrame,
|
LLMFullResponseEndFrame,
|
||||||
LLMFullResponseStartFrame,
|
LLMFullResponseStartFrame,
|
||||||
OpenAILLMContextAssistantTimestampFrame,
|
|
||||||
TextFrame,
|
TextFrame,
|
||||||
TranscriptionFrame,
|
TranscriptionFrame,
|
||||||
UserStartedSpeakingFrame,
|
UserStartedSpeakingFrame,
|
||||||
UserStoppedSpeakingFrame,
|
UserStoppedSpeakingFrame,
|
||||||
)
|
)
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
|
from pipecat.processors.aggregators.llm_context import LLMContext
|
||||||
from pipecat.processors.aggregators.llm_response import (
|
from pipecat.processors.aggregators.llm_response import (
|
||||||
LLMAssistantAggregatorParams,
|
LLMAssistantAggregatorParams,
|
||||||
LLMAssistantContextAggregator,
|
|
||||||
LLMUserContextAggregator,
|
|
||||||
)
|
|
||||||
from pipecat.processors.aggregators.openai_llm_context import (
|
|
||||||
OpenAILLMContext,
|
|
||||||
OpenAILLMContextFrame,
|
|
||||||
)
|
)
|
||||||
|
from pipecat.processors.aggregators.llm_response_universal import LLMContextAggregatorPair
|
||||||
from pipecat.processors.frame_processor import FrameProcessor
|
from pipecat.processors.frame_processor import FrameProcessor
|
||||||
from pipecat.processors.frameworks.langchain import LangchainProcessor
|
from pipecat.processors.frameworks.langchain import LangchainProcessor
|
||||||
from pipecat.tests.utils import SleepFrame, run_test
|
from pipecat.tests.utils import SleepFrame, run_test
|
||||||
@@ -67,13 +64,14 @@ class TestLangchain(unittest.IsolatedAsyncioTestCase):
|
|||||||
proc = LangchainProcessor(chain=chain)
|
proc = LangchainProcessor(chain=chain)
|
||||||
self.mock_proc = self.MockProcessor("token_collector")
|
self.mock_proc = self.MockProcessor("token_collector")
|
||||||
|
|
||||||
context = OpenAILLMContext()
|
context = LLMContext()
|
||||||
tma_in = LLMUserContextAggregator(context)
|
context_aggregator = LLMContextAggregatorPair(
|
||||||
tma_out = LLMAssistantContextAggregator(
|
context, assistant_params=LLMAssistantAggregatorParams(expect_stripped_words=False)
|
||||||
context, params=LLMAssistantAggregatorParams(expect_stripped_words=False)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
pipeline = Pipeline([tma_in, proc, self.mock_proc, tma_out])
|
pipeline = Pipeline(
|
||||||
|
[context_aggregator.user(), proc, self.mock_proc, context_aggregator.assistant()]
|
||||||
|
)
|
||||||
|
|
||||||
frames_to_send = [
|
frames_to_send = [
|
||||||
UserStartedSpeakingFrame(),
|
UserStartedSpeakingFrame(),
|
||||||
@@ -84,8 +82,8 @@ class TestLangchain(unittest.IsolatedAsyncioTestCase):
|
|||||||
expected_down_frames = [
|
expected_down_frames = [
|
||||||
UserStartedSpeakingFrame,
|
UserStartedSpeakingFrame,
|
||||||
UserStoppedSpeakingFrame,
|
UserStoppedSpeakingFrame,
|
||||||
OpenAILLMContextFrame,
|
LLMContextFrame,
|
||||||
OpenAILLMContextAssistantTimestampFrame,
|
LLMContextAssistantTimestampFrame,
|
||||||
]
|
]
|
||||||
await run_test(
|
await run_test(
|
||||||
pipeline,
|
pipeline,
|
||||||
@@ -94,4 +92,6 @@ class TestLangchain(unittest.IsolatedAsyncioTestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual("".join(self.mock_proc.token), self.expected_response)
|
self.assertEqual("".join(self.mock_proc.token), self.expected_response)
|
||||||
self.assertEqual(tma_out.messages[-1]["content"], self.expected_response)
|
self.assertEqual(
|
||||||
|
context_aggregator.assistant().messages[-1]["content"], self.expected_response
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user