Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from typing_extensions import Self

from .._component_config import Component, ComponentModel
from ..models import ChatCompletionClient, FunctionExecutionResultMessage, LLMMessage
from ..models import AssistantMessage, ChatCompletionClient, FunctionExecutionResultMessage, LLMMessage
from ..tools import ToolSchema
from ._chat_completion_context import ChatCompletionContext

Expand All @@ -16,6 +16,27 @@ class TokenLimitedChatCompletionContextConfig(BaseModel):
initial_messages: List[LLMMessage] | None = None


def _is_function_call_message(message: LLMMessage) -> bool:
return isinstance(message, AssistantMessage) and isinstance(message.content, list)


def _remove_middle_message(messages: List[LLMMessage]) -> None:
"""Remove the middle message without splitting an adjacent tool call/result pair."""
middle_index = len(messages) // 2
message = messages[middle_index]

if _is_function_call_message(message):
if middle_index + 1 < len(messages) and isinstance(messages[middle_index + 1], FunctionExecutionResultMessage):
del messages[middle_index : middle_index + 2]
return
elif isinstance(message, FunctionExecutionResultMessage):
if middle_index > 0 and _is_function_call_message(messages[middle_index - 1]):
del messages[middle_index - 1 : middle_index + 1]
return

messages.pop(middle_index)


class TokenLimitedChatCompletionContext(ChatCompletionContext, Component[TokenLimitedChatCompletionContextConfig]):
"""(Experimental) A token based chat completion context maintains a view of the context up to a token limit.

Expand Down Expand Up @@ -61,14 +82,12 @@ async def get_messages(self) -> List[LLMMessage]:
if self._token_limit is None:
remaining_tokens = self._model_client.remaining_tokens(messages, tools=self._tool_schema)
while remaining_tokens < 0 and len(messages) > 0:
middle_index = len(messages) // 2
messages.pop(middle_index)
_remove_middle_message(messages)
remaining_tokens = self._model_client.remaining_tokens(messages, tools=self._tool_schema)
else:
token_count = self._model_client.count_tokens(messages, tools=self._tool_schema)
while token_count > self._token_limit and len(messages) > 0:
middle_index = len(messages) // 2
messages.pop(middle_index)
_remove_middle_message(messages)
token_count = self._model_client.count_tokens(messages, tools=self._tool_schema)
if messages and isinstance(messages[0], FunctionExecutionResultMessage):
# Handle the first message is a function call result message.
Expand Down
46 changes: 45 additions & 1 deletion python/packages/autogen-core/tests/test_model_context.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
from typing import List
from typing import List, Sequence
from unittest.mock import MagicMock

import pytest
from autogen_core import FunctionCall
from autogen_core.model_context import (
BufferedChatCompletionContext,
HeadAndTailChatCompletionContext,
Expand All @@ -10,6 +12,7 @@
from autogen_core.models import (
AssistantMessage,
ChatCompletionClient,
FunctionExecutionResult,
FunctionExecutionResultMessage,
LLMMessage,
UserMessage,
Expand Down Expand Up @@ -208,3 +211,44 @@ async def test_token_limited_model_context_openai_with_function_result(
assert type(retrieved[0]) == UserMessage # Function result should be removed
assert type(retrieved[1]) == AssistantMessage
assert type(retrieved[2]) == UserMessage


@pytest.mark.asyncio
@pytest.mark.parametrize("pair_start", [2, 3], ids=["result-at-middle", "call-at-middle"])
@pytest.mark.parametrize("token_limit", [6, None], ids=["count-tokens", "remaining-tokens"])
async def test_token_limited_model_context_keeps_function_call_result_pairs(
pair_start: int, token_limit: int | None
) -> None:
def count_tokens(messages: Sequence[LLMMessage], **_: object) -> int:
return len(messages)

def remaining_tokens(messages: Sequence[LLMMessage], **_: object) -> int:
return 6 - len(messages)

model_client = MagicMock(spec=ChatCompletionClient)
model_client.count_tokens.side_effect = count_tokens
model_client.remaining_tokens.side_effect = remaining_tokens
model_context = TokenLimitedChatCompletionContext(model_client=model_client, token_limit=token_limit)

function_call = AssistantMessage(
content=[FunctionCall(id="call_1", arguments="{}", name="tool")], source="assistant"
)
function_result = FunctionExecutionResultMessage(
content=[FunctionExecutionResult(content="ok", name="tool", call_id="call_1")]
)
messages: List[LLMMessage] = [
UserMessage(content="m0", source="user"),
UserMessage(content="m1", source="user"),
UserMessage(content="m2", source="user"),
UserMessage(content="m3", source="user"),
UserMessage(content="m4", source="user"),
]
messages[pair_start:pair_start] = [function_call, function_result]
for message in messages:
await model_context.add_message(message)

retrieved = await model_context.get_messages()

assert function_call not in retrieved
assert function_result not in retrieved
assert len(retrieved) == 5