Add a LLMMessagesTransformFrame to facilitate programmatically editing context in a frame-based way.

The previous approach required the caller to directly grab a reference to the context object, grab a "snapshot" of its messages *at that point in time*, transform the messages, and then push an `LLMMessagesUpdateFrame` with the transformed messages. This approach can lead to problems: what if there had already been a change to the context queued in the pipeline? The transformed messages would simply overwrite it without consideration.
This commit is contained in:
Paul Kompfner
2026-02-05 16:31:57 -05:00
parent 50dace147d
commit 3aa403a16f
5 changed files with 119 additions and 3 deletions

View File

@@ -0,0 +1,3 @@
- Added `LLMMessagesTransformFrame` to facilitate programmatically editing context in a frame-based way.
The previous approach required the caller to directly grab a reference to the context object, grab a "snapshot" of its messages _at that point in time_, transform the messages, and then push an `LLMMessagesUpdateFrame` with the transformed messages. This approach can lead to problems: what if there had already been a change to the context queued in the pipeline? The transformed messages would simply overwrite it without consideration.

View File

@@ -38,7 +38,7 @@ from pipecat.utils.time import nanoseconds_to_str
from pipecat.utils.utils import obj_count, obj_id from pipecat.utils.utils import obj_count, obj_id
if TYPE_CHECKING: if TYPE_CHECKING:
from pipecat.processors.aggregators.llm_context import LLMContext, NotGiven from pipecat.processors.aggregators.llm_context import LLMContext, LLMContextMessage, NotGiven
from pipecat.processors.frame_processor import FrameProcessor from pipecat.processors.frame_processor import FrameProcessor
@@ -829,6 +829,25 @@ class LLMMessagesUpdateFrame(DataFrame):
run_llm: Optional[bool] = None run_llm: Optional[bool] = None
@dataclass
class LLMMessagesTransformFrame(DataFrame):
"""Frame containing a transform function to modify the current context's LLM messages.
A frame containing a transform function that takes the context's current list
of LLM messages and returns a modified list.
Only compatible with LLMContext and not the deprecated OpenAILLMContext.
Parameters:
transform: A function that takes a list of messages and returns a
modified list.
run_llm: Whether the context update should be sent to the LLM.
"""
transform: Callable[[List["LLMContextMessage"]], List["LLMContextMessage"]]
run_llm: Optional[bool] = None
@dataclass @dataclass
class LLMSetToolsFrame(DataFrame): class LLMSetToolsFrame(DataFrame):
"""Frame containing tools for LLM function calling. """Frame containing tools for LLM function calling.

View File

@@ -19,7 +19,7 @@ import base64
import io import io
import wave import wave
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Dict, List, Optional, TypeAlias, Union from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, TypeAlias, Union
from loguru import logger from loguru import logger
from openai._types import NOT_GIVEN as OPEN_AI_NOT_GIVEN from openai._types import NOT_GIVEN as OPEN_AI_NOT_GIVEN
@@ -374,6 +374,19 @@ class LLMContext:
""" """
self._messages[:] = messages self._messages[:] = messages
def transform_messages(
self, transform: Callable[[List[LLMContextMessage]], List[LLMContextMessage]]
):
"""Transform the current messages using the provided function.
Args:
transform: A function that takes the current list of messages and returns
a modified list of messages to set in the context.
"""
current_messages = self._messages
new_messages = transform(current_messages)
self.set_messages(new_messages)
def set_tools(self, tools: ToolsSchema | NotGiven = NOT_GIVEN): def set_tools(self, tools: ToolsSchema | NotGiven = NOT_GIVEN):
"""Set the available tools for the LLM. """Set the available tools for the LLM.

View File

@@ -16,7 +16,7 @@ import json
import warnings import warnings
from abc import abstractmethod from abc import abstractmethod
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any, Dict, List, Literal, Optional, Set, Type from typing import Any, Callable, Dict, List, Literal, Optional, Set, Type
from loguru import logger from loguru import logger
@@ -40,6 +40,7 @@ from pipecat.frames.frames import (
LLMFullResponseEndFrame, LLMFullResponseEndFrame,
LLMFullResponseStartFrame, LLMFullResponseStartFrame,
LLMMessagesAppendFrame, LLMMessagesAppendFrame,
LLMMessagesTransformFrame,
LLMMessagesUpdateFrame, LLMMessagesUpdateFrame,
LLMRunFrame, LLMRunFrame,
LLMSetToolChoiceFrame, LLMSetToolChoiceFrame,
@@ -270,6 +271,17 @@ class LLMContextAggregator(FrameProcessor):
""" """
self._context.set_messages(messages) self._context.set_messages(messages)
def transform_messages(
self, transform: Callable[[List[LLMContextMessage]], List[LLMContextMessage]]
):
"""Transform the context messages using a provided function.
Args:
transform: A function that takes the current list of messages and returns
a modified list of messages to set in the context.
"""
self._context.transform_messages(transform)
def set_tools(self, tools: ToolsSchema | NotGiven): def set_tools(self, tools: ToolsSchema | NotGiven):
"""Set tools in the context. """Set tools in the context.
@@ -470,6 +482,8 @@ class LLMUserAggregator(LLMContextAggregator):
await self._handle_llm_messages_append(frame) await self._handle_llm_messages_append(frame)
elif isinstance(frame, LLMMessagesUpdateFrame): elif isinstance(frame, LLMMessagesUpdateFrame):
await self._handle_llm_messages_update(frame) await self._handle_llm_messages_update(frame)
elif isinstance(frame, LLMMessagesTransformFrame):
await self._handle_llm_messages_transform(frame)
elif isinstance(frame, LLMSetToolsFrame): elif isinstance(frame, LLMSetToolsFrame):
self.set_tools(frame.tools) self.set_tools(frame.tools)
# Push the LLMSetToolsFrame as well, since some realtime (aka # Push the LLMSetToolsFrame as well, since some realtime (aka
@@ -603,6 +617,15 @@ class LLMUserAggregator(LLMContextAggregator):
if frame.run_llm: if frame.run_llm:
await self.push_context_frame() await self.push_context_frame()
async def _handle_llm_messages_transform(self, frame: LLMMessagesTransformFrame):
self.transform_messages(frame.transform)
# Mark the context as programmatically edited. This flag is stored as a
# runtime attribute on the shared context object so that both user and
# assistant aggregators can see it.
self._context._pipecat_messages_programmatically_edited = True
if frame.run_llm:
await self.push_context_frame()
async def _handle_speech_control_params(self, frame: SpeechControlParamsFrame): async def _handle_speech_control_params(self, frame: SpeechControlParamsFrame):
if frame.id in self._self_queued_frames: if frame.id in self._self_queued_frames:
return return
@@ -888,6 +911,8 @@ class LLMAssistantAggregator(LLMContextAggregator):
await self._handle_llm_messages_append(frame) await self._handle_llm_messages_append(frame)
elif isinstance(frame, LLMMessagesUpdateFrame): elif isinstance(frame, LLMMessagesUpdateFrame):
await self._handle_llm_messages_update(frame) await self._handle_llm_messages_update(frame)
elif isinstance(frame, LLMMessagesTransformFrame):
await self._handle_llm_messages_transform(frame)
elif isinstance(frame, LLMSetToolsFrame): elif isinstance(frame, LLMSetToolsFrame):
self.set_tools(frame.tools) self.set_tools(frame.tools)
elif isinstance(frame, LLMSetToolChoiceFrame): elif isinstance(frame, LLMSetToolChoiceFrame):
@@ -947,6 +972,15 @@ class LLMAssistantAggregator(LLMContextAggregator):
if frame.run_llm: if frame.run_llm:
await self.push_context_frame(FrameDirection.UPSTREAM) await self.push_context_frame(FrameDirection.UPSTREAM)
async def _handle_llm_messages_transform(self, frame: LLMMessagesTransformFrame):
self.transform_messages(frame.transform)
# Mark the context as programmatically edited. This flag is stored as a
# runtime attribute on the shared context object so that both user and
# assistant aggregators can see it.
self._context._pipecat_messages_programmatically_edited = True
if frame.run_llm:
await self.push_context_frame(FrameDirection.UPSTREAM)
async def _handle_interruptions(self, frame: InterruptionFrame): async def _handle_interruptions(self, frame: InterruptionFrame):
await self._trigger_assistant_turn_stopped() await self._trigger_assistant_turn_stopped()
self._started = 0 self._started = 0

View File

@@ -18,6 +18,7 @@ from pipecat.frames.frames import (
LLMFullResponseEndFrame, LLMFullResponseEndFrame,
LLMFullResponseStartFrame, LLMFullResponseStartFrame,
LLMMessagesAppendFrame, LLMMessagesAppendFrame,
LLMMessagesTransformFrame,
LLMMessagesUpdateFrame, LLMMessagesUpdateFrame,
LLMRunFrame, LLMRunFrame,
LLMTextFrame, LLMTextFrame,
@@ -147,6 +148,52 @@ class TestLLMUserAggregator(unittest.IsolatedAsyncioTestCase):
) )
assert context.messages[0]["content"] == "Hi there!" assert context.messages[0]["content"] == "Hi there!"
async def test_llm_messages_transform(self):
context = LLMContext()
# Set up initial messages
context.set_messages(
[
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there!"},
{"role": "user", "content": "How are you?"},
]
)
pipeline = Pipeline([LLMUserAggregator(context)])
# Transform that keeps only user messages
def keep_user_messages(messages):
return [m for m in messages if m["role"] == "user"]
frames_to_send = [LLMMessagesTransformFrame(transform=keep_user_messages)]
await run_test(
pipeline,
frames_to_send=frames_to_send,
)
assert len(context.messages) == 2
assert context.messages[0]["content"] == "Hello"
assert context.messages[1]["content"] == "How are you?"
async def test_llm_messages_transform_run(self):
context = LLMContext()
# Set up initial messages
context.set_messages([{"role": "user", "content": "Hello"}])
pipeline = Pipeline([LLMUserAggregator(context)])
# Transform that modifies the content
def uppercase_content(messages):
return [{"role": m["role"], "content": m["content"].upper()} for m in messages]
frames_to_send = [LLMMessagesTransformFrame(transform=uppercase_content, run_llm=True)]
expected_down_frames = [LLMContextFrame]
await run_test(
pipeline,
frames_to_send=frames_to_send,
expected_down_frames=expected_down_frames,
)
assert context.messages[0]["content"] == "HELLO"
async def test_default_user_turn_strategies(self): async def test_default_user_turn_strategies(self):
context = LLMContext() context = LLMContext()
user_aggregator = LLMUserAggregator(context) user_aggregator = LLMUserAggregator(context)