diff --git a/python/packages/autogen-core/src/autogen_core/model_context/_token_limited_chat_completion_context.py b/python/packages/autogen-core/src/autogen_core/model_context/_token_limited_chat_completion_context.py index b8a0258a4d62..6336458f5241 100644 --- a/python/packages/autogen-core/src/autogen_core/model_context/_token_limited_chat_completion_context.py +++ b/python/packages/autogen-core/src/autogen_core/model_context/_token_limited_chat_completion_context.py @@ -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 @@ -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. @@ -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. diff --git a/python/packages/autogen-core/tests/test_model_context.py b/python/packages/autogen-core/tests/test_model_context.py index bcd2a9c87d9e..9d9c678f865e 100644 --- a/python/packages/autogen-core/tests/test_model_context.py +++ b/python/packages/autogen-core/tests/test_model_context.py @@ -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, @@ -10,6 +12,7 @@ from autogen_core.models import ( AssistantMessage, ChatCompletionClient, + FunctionExecutionResult, FunctionExecutionResultMessage, LLMMessage, UserMessage, @@ -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