diff --git a/python/packages/core/agent_framework/_sessions.py b/python/packages/core/agent_framework/_sessions.py index fc6e395a20..ae218f8e94 100644 --- a/python/packages/core/agent_framework/_sessions.py +++ b/python/packages/core/agent_framework/_sessions.py @@ -65,6 +65,7 @@ JsonDumps: TypeAlias = Callable[[Any], str | bytes] JsonLoads: TypeAlias = Callable[[str | bytes], Any] ServiceSessionId: TypeAlias = Mapping[str, Any] +MessageIdentity: TypeAlias = tuple[str, ...] StateT = TypeVar("StateT") StateEncoder: TypeAlias = Callable[[Any], Mapping[str, Any]] StateDecoder: TypeAlias = Callable[[Mapping[str, Any]], Any] @@ -184,6 +185,74 @@ def _deduplicate_origin_session_ids(origin_session_ids: Iterable[str]) -> list[s return unique_origin_session_ids +def get_message_identity(message: Message) -> MessageIdentity: + """Return a stable identity for a message for deduplication. + + Uses the message's ID if available, otherwise falls back to a hash of + its role and serialized contents to prevent duplicate persistence. + """ + msg_id = getattr(message, "message_id", None) + if msg_id is None: + msg_id = getattr(message, "id", None) + if msg_id is not None: + return ("id", str(msg_id)) + + try: + contents_data = [c.to_dict() for c in message.contents] if message.contents else [] + serialized = json.dumps(contents_data, sort_keys=True, ensure_ascii=False) + return ("content", str(message.role), serialized) + except Exception: + return ("content", str(message.role), str(message.contents)) + + +def _get_message_hash(message: Message) -> tuple: + """Stable hash for sequence matching.""" + return get_message_identity(message) + + +def filter_new_messages(existing: Sequence[Message], incoming: Sequence[Message]) -> list[Message]: + """Filters incoming messages to only those that are truly new. + + Handles both 'append-only' and 'full transcript replay' scenarios. + Prevents superlinear growth and preserves legitimate duplicate turns. + """ + if not existing: + return list(incoming) + + 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) :]) + + last_existing_hash = existing_hashes[-1] + try: + split_idx = -1 + for i in range(len(incoming_hashes) - 1, -1, -1): + if incoming_hashes[i] == last_existing_hash: + match = True + for j in range(1, min(i + 1, len(existing_hashes))): + if incoming_hashes[i - j] != existing_hashes[-(j + 1)]: + match = False + break + if match: + split_idx = i + break + + if split_idx != -1: + return list(incoming[split_idx + 1 :]) + except Exception: + logger.debug("sequence alignment check failed, falling back to set-based deduplication") + + existing_set = set(existing_hashes) + new_msgs = [] + 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 + + def _is_middleware_sequence( middleware: MiddlewareTypes | Sequence[MiddlewareTypes], ) -> TypeGuard[Sequence[MiddlewareTypes]]: @@ -1833,10 +1902,13 @@ async def save_messages( ) -> None: """Persist messages to session state.""" mark_feature_used(FeatureIndex.CORE_IN_MEMORY_HISTORY_PROVIDER) - if state is None: + if state is None or not messages: return existing = state.get("messages", []) - state["messages"] = [*existing, *messages] + new_messages = filter_new_messages(existing, messages) + + if new_messages: + state["messages"] = [*existing, *new_messages] @experimental(feature_id=ExperimentalFeature.FILE_HISTORY) @@ -1995,9 +2067,26 @@ async def save_messages( def _append_messages() -> None: with file_lock: if self.serialization_format == "json": - with file_path.open("a", encoding="utf-8") as file_handle: - for message in messages: - file_handle.write(f"{self._serialize_json_message(message)}\n") + existing_messages: list[Message] = [] + if file_path.exists(): + with file_path.open("r", encoding="utf-8") as f: + for line in f: + line = line.strip() + if not line: + continue + try: + payload = self.loads(line) + msg = Message.from_dict(dict(cast(Mapping[str, Any], payload))) + existing_messages.append(msg) + except Exception: + logger.debug("failed to parse history line for deduplication") + continue + + new_messages = filter_new_messages(existing_messages, messages) + if new_messages: + with file_path.open("a", encoding="utf-8") as file_handle: + for message in new_messages: + file_handle.write(f"{self._serialize_json_message(message)}\n") return with file_path.open("ab") as file_handle: for message in messages: diff --git a/python/packages/core/tests/core/test_sessions.py b/python/packages/core/tests/core/test_sessions.py index 5a229b2a41..38d634772a 100644 --- a/python/packages/core/tests/core/test_sessions.py +++ b/python/packages/core/tests/core/test_sessions.py @@ -8,7 +8,7 @@ from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass from pathlib import Path -from typing import Any +from typing import TYPE_CHECKING, Any, cast from unittest.mock import patch import msgspec @@ -41,6 +41,9 @@ from agent_framework._telemetry import FeatureIndex from agent_framework.exceptions import MiddlewareException +if TYPE_CHECKING: + from agent_framework._agents import SupportsAgentRun + # --------------------------------------------------------------------------- # SessionContext tests # --------------------------------------------------------------------------- @@ -1375,6 +1378,111 @@ async def test_source_id_attribution(self) -> None: ctx.extend_messages("custom-source", [Message(role="user", contents=["test"])]) assert "custom-source" in ctx.context_messages + async def test_save_messages_deduplicates_identical_messages(self) -> None: + """Test that save_messages does not re-append messages already in the store.""" + provider = InMemoryHistoryProvider() + state: dict[str, Any] = {} + + msg1 = Message(role="user", contents=["hello"]) + msg2 = Message(role="assistant", contents=["hi there"]) + + await provider.save_messages("s1", [msg1, msg2], state=state) + assert len(state["messages"]) == 2 + + await provider.save_messages("s1", [msg1, msg2], state=state) + assert len(state["messages"]) == 2 + + async def test_save_messages_only_appends_new_messages(self) -> None: + """Test that save_messages filters out old messages and only appends new ones""" + provider = InMemoryHistoryProvider() + state: dict[str, Any] = {} + + msg1 = Message(role="user", contents=["hello"]) + msg2 = Message(role="assistant", contents=["hi there"]) + msg3 = Message(role="user", contents=["how are you?"]) + + await provider.save_messages("s1", [msg1, msg2], state=state) + assert len(state["messages"]) == 2 + + await provider.save_messages("s1", [msg1, msg2, msg3], state=state) + assert len(state["messages"]) == 3 + assert state["messages"][2].text == "how are you?" + + async def test_save_messages_different_roles_same_text_not_deduplicated(self) -> None: + """Test that messages with the same text but different roles are kept separate.""" + provider = InMemoryHistoryProvider() + state: dict[str, Any] = {} + + msg1 = Message(role="user", contents=["ping"]) + msg2 = Message(role="assistant", contents=["ping"]) + + await provider.save_messages("s1", [msg1, msg2], state=state) + assert len(state["messages"]) == 2 + + async def test_save_messages_deduplication_with_none_state(self) -> None: + """Test that save_messages with None state does not raise.""" + provider = InMemoryHistoryProvider() + msg = Message(role="user", contents=["hello"]) + await provider.save_messages("s1", [msg], state=None) + + async def test_full_loop_does_not_grow_superlinearly(self) -> None: + """Regression test: a multi-round looped run must not re-persist the whole + conversation on every round.""" + from agent_framework import AgentResponse + + provider = InMemoryHistoryProvider() + session = AgentSession() + provider_state = session.state.setdefault(provider.source_id, {}) + ctx1 = SessionContext(session_id="s1", input_messages=[Message(role="user", contents=["turn 1"])]) + + await provider.before_run( + agent=cast("SupportsAgentRun", None), + session=session, + context=ctx1, + state=provider_state, + ) + ctx1._response = AgentResponse(messages=[Message(role="assistant", contents=["reply 1"])]) + await provider.after_run( + agent=cast("SupportsAgentRun", None), + session=session, + context=ctx1, + state=provider_state, + ) + + ctx2 = SessionContext(session_id="s1", input_messages=[Message(role="user", contents=["turn 2"])]) + await provider.before_run( + agent=cast("SupportsAgentRun", None), + session=session, + context=ctx2, + state=provider_state, + ) + ctx2._response = AgentResponse(messages=[Message(role="assistant", contents=["reply 2"])]) + await provider.after_run( + agent=cast("SupportsAgentRun", None), + session=session, + context=ctx2, + state=provider_state, + ) + + stored = session.state[provider.source_id]["messages"] + assert len(stored) == 4 + texts = [m.text for m in stored] + assert texts == ["turn 1", "reply 1", "turn 2", "reply 2"] + + async def test_save_messages_preserves_duplicate_content(self) -> None: + """Two separate user 'yes' replies in the same batch must both be persisted.""" + provider = InMemoryHistoryProvider() + state: dict[str, Any] = {} + + yes_1 = Message(role="user", contents=["yes"]) + yes_2 = Message(role="user", contents=["yes"]) + + await provider.save_messages("s1", [yes_1, yes_2], state=state) + + assert len(state["messages"]) == 2 + assert state["messages"][0].text == "yes" + assert state["messages"][1].text == "yes" + class TestFileHistoryProvider: def test_is_marked_experimental(self) -> None: @@ -1658,3 +1766,78 @@ def tracked_open(path: Path, *args: Any, **kwargs: Any) -> Any: assert not overlap_detected loaded = await provider.get_messages(session_id) assert [message.text for message in loaded] == ["first", "second"] + + async def test_save_messages_deduplicates_identical_messages(self, tmp_path: Path) -> None: + """Test that FileHistoryProvider does not re-append already persisted messages.""" + provider = FileHistoryProvider(tmp_path) + + msg1 = Message(role="user", contents=["hello"]) + msg2 = Message(role="assistant", contents=["hi there"]) + + await provider.save_messages("s1", [msg1, msg2]) + loaded = await provider.get_messages("s1") + assert len(loaded) == 2 + + await provider.save_messages("s1", [msg1, msg2]) + loaded = await provider.get_messages("s1") + assert len(loaded) == 2 + + async def test_save_messages_only_appends_new_messages(self, tmp_path: Path) -> None: + """Test that FileHistoryProvider filters out old messages and only appends new ones""" + provider = FileHistoryProvider(tmp_path) + + msg1 = Message(role="user", contents=["hello"]) + msg2 = Message(role="assistant", contents=["hi there"]) + msg3 = Message(role="user", contents=["how are you?"]) + + await provider.save_messages("s1", [msg1, msg2]) + loaded = await provider.get_messages("s1") + assert len(loaded) == 2 + + await provider.save_messages("s1", [msg1, msg2, msg3]) + loaded = await provider.get_messages("s1") + assert len(loaded) == 3 + assert loaded[2].text == "how are you?" + + async def test_save_messages_different_roles_same_text_not_deduplicated(self, tmp_path: Path) -> None: + """Test that messages with the same text but different roles are kept separate.""" + provider = FileHistoryProvider(tmp_path) + + msg1 = Message(role="user", contents=["ping"]) + msg2 = Message(role="assistant", contents=["ping"]) + + await provider.save_messages("s1", [msg1, msg2]) + loaded = await provider.get_messages("s1") + assert len(loaded) == 2 + + async def test_deduplication_file_integrity(self, tmp_path: Path) -> None: + """Test that deduplication writes the correct number of lines to the JSONL file.""" + provider = FileHistoryProvider(tmp_path) + + msg1 = Message(role="user", contents=["hello"]) + msg2 = Message(role="assistant", contents=["hi there"]) + msg3 = Message(role="user", contents=["follow-up"]) + + await provider.save_messages("s1", [msg1, msg2]) + + session_file = provider._session_file_path("s1") + raw_lines = (await asyncio.to_thread(session_file.read_text, encoding="utf-8")).splitlines() + assert len(raw_lines) == 2 + + await provider.save_messages("s1", [msg1, msg2, msg3]) + raw_lines = (await asyncio.to_thread(session_file.read_text, encoding="utf-8")).splitlines() + assert len(raw_lines) == 3 + + async def test_save_messages_preserves_duplicate_content(self, tmp_path: Path) -> None: + """Test that two identical user turns in the same batch are both persisted.""" + provider = FileHistoryProvider(tmp_path) + + yes_1 = Message(role="user", contents=["yes"]) + yes_2 = Message(role="user", contents=["yes"]) + + await provider.save_messages("s1", [yes_1, yes_2]) + loaded = await provider.get_messages("s1") + + assert len(loaded) == 2 + assert loaded[0].text == "yes" + assert loaded[1].text == "yes" diff --git a/python/packages/redis/agent_framework_redis/_history_provider.py b/python/packages/redis/agent_framework_redis/_history_provider.py index a7703db8a2..8805b0b233 100644 --- a/python/packages/redis/agent_framework_redis/_history_provider.py +++ b/python/packages/redis/agent_framework_redis/_history_provider.py @@ -8,12 +8,13 @@ from __future__ import annotations +import json from collections.abc import Sequence from typing import Any, ClassVar import redis.asyncio as redis from agent_framework import Message -from agent_framework._sessions import HistoryProvider +from agent_framework._sessions import HistoryProvider, filter_new_messages from agent_framework._telemetry import mark_feature_used from redis.credentials import CredentialProvider @@ -157,7 +158,14 @@ async def save_messages( return key = self._redis_key(session_id) - serialized_messages = [self._serialize_json(msg) for msg in messages] + + existing_messages = await self.get_messages(session_id, state=state, **kwargs) + new_messages = filter_new_messages(existing_messages, messages) + + if not new_messages: + return + + serialized_messages = [self._serialize_json(msg) for msg in new_messages] async with self._redis_client.pipeline(transaction=True) as pipe: for serialized in serialized_messages: @@ -172,15 +180,11 @@ async def save_messages( @staticmethod def _serialize_json(message: Message) -> str: """Serialize a Message to a JSON string for Redis storage.""" - import json - return json.dumps(message.to_dict()) @staticmethod def _deserialize_json(data: str) -> dict[str, Any]: """Deserialize a JSON string from Redis to a dict.""" - import json - return json.loads(data) async def clear(self, session_id: str | None) -> None: diff --git a/python/packages/redis/tests/test_providers.py b/python/packages/redis/tests/test_providers.py index 55aee29662..f5a004d15f 100644 --- a/python/packages/redis/tests/test_providers.py +++ b/python/packages/redis/tests/test_providers.py @@ -569,3 +569,93 @@ async def test_after_run_skips_when_no_messages(self, mock_redis_client: MagicMo ) # type: ignore[arg-type] mock_redis_client.pipeline.assert_not_called() + + +class TestRedisHistoryProviderDeduplication: + """Tests for Redis save_messages deduplication and trimming behavior.""" + + async def test_deduplicates_identical_messages(self, mock_redis_client: MagicMock): + msg1 = Message(role="user", contents=["hello"]) + msg2 = Message(role="assistant", contents=["hi there"]) + + mock_redis_client.lrange = AsyncMock(return_value=[json.dumps(msg1.to_dict()), json.dumps(msg2.to_dict())]) + + 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") + + await provider.save_messages("s1", [msg1, msg2]) + + pipeline = mock_redis_client.pipeline.return_value.__aenter__.return_value + pipeline.rpush.assert_not_called() + pipeline.execute.assert_not_called() + + async def test_only_appends_new_messages(self, mock_redis_client: MagicMock): + msg1 = Message(role="user", contents=["hello"]) + msg2 = Message(role="assistant", contents=["hi there"]) + msg3 = Message(role="user", contents=["how are you?"]) + + mock_redis_client.lrange = AsyncMock(return_value=[json.dumps(msg1.to_dict()), json.dumps(msg2.to_dict())]) + + 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") + + await provider.save_messages("s1", [msg1, msg2, msg3]) + + pipeline = mock_redis_client.pipeline.return_value.__aenter__.return_value + assert pipeline.rpush.call_count == 1 + + call_args = pipeline.rpush.call_args[0] + pushed_msg_dict = json.loads(call_args[1]) + 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"]) + + mock_redis_client.lrange = AsyncMock(return_value=[json.dumps(msg1.to_dict())]) + + 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") + + msg2 = Message(role="assistant", contents=["ping"]) + await provider.save_messages("s1", [msg1, msg2]) + + pipeline = mock_redis_client.pipeline.return_value.__aenter__.return_value + assert pipeline.rpush.call_count == 1 + + async def test_trimmed_messages_not_reappended(self, mock_redis_client: MagicMock): + """Messages trimmed by max_messages should not be re-appended + when the caller resends the full transcript. Sequence matching + handles this without needing a :seen set.""" + msg_old = Message(role="user", contents=["old"]) + msg_new = Message(role="assistant", contents=["new"]) + + mock_redis_client.lrange = AsyncMock(return_value=[json.dumps(msg_new.to_dict())]) + + 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") + + await provider.save_messages("s1", [msg_old, msg_new]) + + pipeline = mock_redis_client.pipeline.return_value.__aenter__.return_value + + pipeline.rpush.assert_not_called() + + async def test_preserves_duplicate_content(self, mock_redis_client: MagicMock): + """Two separate user 'yes' replies must both be persisted.""" + yes_1 = Message(role="user", contents=["yes"]) + yes_2 = Message(role="user", contents=["yes"]) + + mock_redis_client.lrange = AsyncMock(return_value=[]) + + 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") + + await provider.save_messages("s1", [yes_1, yes_2]) + + pipeline = mock_redis_client.pipeline.return_value.__aenter__.return_value + assert pipeline.rpush.call_count == 2