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:
Paul Kompfner
2025-09-24 15:54:01 -04:00
parent ceba27e696
commit 236ac93ac6
5 changed files with 53 additions and 45 deletions

View File

@@ -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()

View File

@@ -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:

View File

@@ -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.

View File

@@ -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")

View File

@@ -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
)