Skip to content
Merged
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
31 changes: 14 additions & 17 deletions python/packages/core/agent_framework/_sessions.py
Original file line number Diff line number Diff line change
Expand Up @@ -175,7 +175,7 @@ def _get_message_hash(message: Message) -> MessageIdentity:


def filter_new_messages(existing: Sequence[Message], incoming: Sequence[Message]) -> list[Message]:
"""Filters incoming messages to only those that are truly new.
"""Return messages after the ordered overlap with persisted history.

Handles both 'append-only' and 'full transcript replay' scenarios.
Prevents superlinear growth and preserves legitimate duplicate turns.
Expand All @@ -186,23 +186,20 @@ def filter_new_messages(existing: Sequence[Message], incoming: Sequence[Message]
existing_hashes = [_get_message_hash(m) for m in existing]
incoming_hashes = [_get_message_hash(m) for m in incoming]

if len(incoming) >= len(existing) and incoming_hashes[: len(existing_hashes)] == existing_hashes:
return list(incoming[len(existing) :])
for i in range(len(incoming_hashes) - len(existing_hashes) + 1):
if incoming_hashes[i : i + len(existing_hashes)] == existing_hashes:
if i == 0 and len(existing) == 1 and existing[-1].role == "user" and existing_hashes[-1][0] != "id":
break # A repeated input without an ID is more important to retain than a possible replay.
return list(incoming[i + len(existing_hashes) :])

try:
for i in range(len(incoming_hashes) - len(existing_hashes) + 1):
if incoming_hashes[i : i + len(existing_hashes)] == existing_hashes:
return list(incoming[i + len(existing_hashes) :])
except Exception:
logger.debug("sequence alignment check failed, falling back to set-based deduplication")

existing_set = set(existing_hashes)
new_msgs: list[Message] = []
for m, h in zip(incoming, incoming_hashes):
if h not in existing_set:
new_msgs.append(m)
existing_set.add(h)
return new_msgs
for overlap in range(min(len(existing_hashes), len(incoming_hashes)), 0, -1):
if existing_hashes[-overlap:] != incoming_hashes[:overlap]:
continue
if overlap == 1 and existing[-1].role == "user" and existing_hashes[-1][0] != "id":
continue
return list(incoming[overlap:])

return list(incoming)


@dataclass(frozen=True, slots=True)
Expand Down
15 changes: 8 additions & 7 deletions python/packages/core/tests/core/test_middleware_with_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -1772,14 +1772,15 @@ async def process(self, context: AgentContext, call_next: Callable[[], Awaitable
second_after = thread_states[3]
assert second_after["before_next"] is False
assert second_after["messages_count"] == 1 # Input messages unchanged
assert second_after["thread_count"] == 3 # Previous history (2) + current input (1)
assert second_after["thread_count"] == 4 # Both runs persist their input and response.
assert second_after["messages_text"] == ["second message"]
# Thread should contain: first input + first response + second input
assert "first message" in second_after["thread_messages_text"]
assert "second message" in second_after["thread_messages_text"]
# "test response" should only appear once since the duplicate was correctly filtered
response_count = sum(1 for text in second_after["thread_messages_text"] if "test response" in text)
assert response_count == 1
# Repeated response text remains attached to each run, in conversation order.
assert second_after["thread_messages_text"] == [
"first message",
"test response",
"second message",
"test response",
]


class TestChatAgentChatMiddleware:
Expand Down
79 changes: 79 additions & 0 deletions python/packages/core/tests/core/test_sessions.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@
_run_identity_scope,
_RunPersistenceGate,
_suspend_run_persistence_gate,
filter_new_messages,
is_local_history_conversation_id,
)
from agent_framework._telemetry import FeatureIndex
Expand All @@ -53,6 +54,51 @@
if TYPE_CHECKING:
from agent_framework._agents import SupportsAgentRun


def test_filter_new_messages_preserves_repeated_user_turn() -> None:
existing = [Message(role="user", contents=["yes"]), Message(role="assistant", contents=["first reply"])]
incoming = [Message(role="user", contents=["yes"]), Message(role="assistant", contents=["second reply"])]

assert filter_new_messages(existing, incoming) == incoming


def test_filter_new_messages_preserves_repeated_user_after_unanswered_input() -> None:
previous_input = Message(role="user", contents=["yes"])
repeated_input = Message(role="user", contents=["yes"])
reply = Message(role="assistant", contents=["second reply"])

assert filter_new_messages([previous_input], [repeated_input]) == [repeated_input]
assert filter_new_messages([previous_input], [repeated_input, reply]) == [repeated_input, reply]

identified_input = Message(role="user", contents=["yes"], message_id="turn-1")
assert filter_new_messages([identified_input], [identified_input]) == []


def test_filter_new_messages_aligns_partial_replay_with_repeated_content() -> None:
existing = [
Message(role="user", contents=["yes"]),
Message(role="assistant", contents=["first reply"]),
Message(role="user", contents=["yes"]),
Message(role="assistant", contents=["second reply"]),
]
new_messages = [Message(role="user", contents=["yes"]), Message(role="assistant", contents=["third reply"])]
incoming = [*existing[-2:], *new_messages]

assert filter_new_messages(existing, incoming) == new_messages


def test_filter_new_messages_keeps_tool_call_result_pairs_on_replay() -> None:
call = Message(
role="assistant", contents=[Content.from_function_call(call_id="call-1", name="lookup", arguments="{}")]
)
result = Message(role="tool", contents=[Content.from_function_result(call_id="call-1", result="found")])
existing = [Message(role="user", contents=["lookup"]), call, result]
new_messages = [Message(role="user", contents=["lookup"]), Message(role="assistant", contents=["again"])]

assert filter_new_messages(existing, [*existing, *new_messages]) == new_messages
assert filter_new_messages(existing, [result, *new_messages]) == new_messages


# ---------------------------------------------------------------------------
# SessionContext tests
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -1624,6 +1670,23 @@ async def test_save_messages_preserves_duplicate_content(self) -> None:
assert state["messages"][0].text == "yes"
assert state["messages"][1].text == "yes"

async def test_save_messages_preserves_repeated_user_turn_across_saves(self) -> None:
provider = InMemoryHistoryProvider()
state: dict[str, Any] = {}

await provider.save_messages(
"s1",
[Message(role="user", contents=["yes"]), Message(role="assistant", contents=["first reply"])],
state=state,
)
await provider.save_messages(
"s1",
[Message(role="user", contents=["yes"]), Message(role="assistant", contents=["second reply"])],
state=state,
)

assert [message.text for message in state["messages"]] == ["yes", "first reply", "yes", "second reply"]

async def test_save_messages_handles_replayed_transcript_with_duplicates(self) -> None:
provider = InMemoryHistoryProvider()
state: dict[str, Any] = {}
Expand Down Expand Up @@ -2111,6 +2174,22 @@ async def test_save_messages_preserves_duplicate_content(
assert loaded[0].text == "yes"
assert loaded[1].text == "yes"

@pytest.mark.parametrize("serialization_format", ["json", "msgpack"])
async def test_save_messages_preserves_repeated_user_turn_across_saves(
self, tmp_path: Path, serialization_format: Literal["json", "msgpack"]
) -> None:
provider = FileHistoryProvider(tmp_path, serialization_format=serialization_format)

await provider.save_messages(
"s1", [Message(role="user", contents=["yes"]), Message(role="assistant", contents=["first reply"])]
)
await provider.save_messages(
"s1", [Message(role="user", contents=["yes"]), Message(role="assistant", contents=["second reply"])]
)

loaded = await provider.get_messages("s1")
assert [message.text for message in loaded] == ["yes", "first reply", "yes", "second reply"]


# ---------------------------------------------------------------------------
# Run-persistence gate tests
Expand Down
22 changes: 21 additions & 1 deletion python/packages/redis/tests/test_providers.py
Original file line number Diff line number Diff line change
Expand Up @@ -833,7 +833,7 @@ async def test_only_appends_new_messages(self, mock_redis_client: MagicMock):
assert pushed_msg_dict["contents"][0]["text"] == "how are you?"

async def test_different_roles_same_text_not_deduplicated(self, mock_redis_client: MagicMock):
msg1 = Message(role="user", contents=["ping"])
msg1 = Message(role="user", contents=["ping"], message_id="original-ping")

mock_redis_client.lrange = AsyncMock(return_value=[json.dumps(msg1.to_dict())])

Expand All @@ -846,6 +846,26 @@ async def test_different_roles_same_text_not_deduplicated(self, mock_redis_clien

pipeline = mock_redis_client.pipeline.return_value.__aenter__.return_value
assert pipeline.rpush.call_count == 1
pushed_msg_dict = json.loads(pipeline.rpush.call_args.args[1])
assert pushed_msg_dict["role"] == "assistant"

async def test_repeated_user_turn_is_not_deduplicated(self, mock_redis_client: MagicMock):
previous = [Message(role="user", contents=["yes"]), Message(role="assistant", contents=["first reply"])]
mock_redis_client.lrange = AsyncMock(return_value=[json.dumps(msg.to_dict()) for msg in previous])

with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379", application_id="test-app")

await provider.save_messages(
"s1", [Message(role="user", contents=["yes"]), Message(role="assistant", contents=["second reply"])]
)

pipeline = mock_redis_client.pipeline.return_value.__aenter__.return_value
assert [json.loads(call.args[1])["contents"][0]["text"] for call in pipeline.rpush.await_args_list] == [
"yes",
"second reply",
]

async def test_trimmed_messages_not_reappended(self, mock_redis_client: MagicMock):
"""Messages trimmed by max_messages should not be re-appended
Expand Down
Loading