From e6087fb139ed0ecf5997534b9634ca0fa8b2de10 Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Wed, 22 Jul 2026 00:15:21 +0530 Subject: [PATCH 1/6] fix: prevent superlinear history growth by deduplicating messages in save_messages --- .../core/agent_framework/_sessions.py | 49 ++++++- .../packages/core/tests/core/test_sessions.py | 132 ++++++++++++++++++ .../_history_provider.py | 35 ++++- python/packages/redis/tests/test_providers.py | 55 ++++++++ 4 files changed, 268 insertions(+), 3 deletions(-) diff --git a/python/packages/core/agent_framework/_sessions.py b/python/packages/core/agent_framework/_sessions.py index 9a6dab9a7c1..ae388b1053c 100644 --- a/python/packages/core/agent_framework/_sessions.py +++ b/python/packages/core/agent_framework/_sessions.py @@ -77,6 +77,25 @@ def _deduplicate_origin_session_ids(origin_session_ids: Iterable[str]) -> list[s return unique_origin_session_ids +def _get_message_identity(message: Message) -> tuple: + """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, "id", None) + if msg_id is not None: + return ("id", msg_id) + + try: + # Use to_dict() for a stable, deterministic representation + serialized = json.dumps(message.to_dict(), sort_keys=True, ensure_ascii=False) + return ("content", message.role, serialized) + except Exception: + # Fallback if serialization fails for any reason + return ("content", message.role, str(message.contents)) + + def _is_middleware_sequence( middleware: MiddlewareTypes | Sequence[MiddlewareTypes], ) -> TypeGuard[Sequence[MiddlewareTypes]]: @@ -1115,7 +1134,15 @@ async def save_messages( if state is None: return existing = state.get("messages", []) - state["messages"] = [*existing, *messages] + existing_identities = {_get_message_identity(m) for m in existing} + new_messages = [] + for msg in messages: + identity = _get_message_identity(msg) + if identity not in existing_identities: + existing_identities.add(identity) + new_messages.append(msg) + if new_messages: + state["messages"] = [*existing, *new_messages] @experimental(feature_id=ExperimentalFeature.FILE_HISTORY) @@ -1291,9 +1318,27 @@ async def save_messages( file_lock = self._session_write_lock(file_path) def _append_messages() -> None: + existing_identities: set[tuple] = set() + 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_identities.add(_get_message_identity(msg)) + except Exception: + logger.debug("Failed to parse history line for deduplication", exc_info=True) + continue + with file_lock, file_path.open("a", encoding="utf-8") as file_handle: for message in messages: - file_handle.write(f"{self._serialize_message(message)}\n") + identity = _get_message_identity(message) + if identity not in existing_identities: + existing_identities.add(identity) + file_handle.write(f"{self._serialize_message(message)}\n") async with async_lock: await asyncio.to_thread(_append_messages) diff --git a/python/packages/core/tests/core/test_sessions.py b/python/packages/core/tests/core/test_sessions.py index c4f4061d42c..d9135ba9ac1 100644 --- a/python/packages/core/tests/core/test_sessions.py +++ b/python/packages/core/tests/core/test_sessions.py @@ -669,6 +669,77 @@ 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=None, session=session, context=ctx1, state=provider_state) # type: ignore[arg-type] + ctx1._response = AgentResponse(messages=[Message(role="assistant", contents=["reply 1"])]) + await provider.after_run(agent=None, session=session, context=ctx1, state=provider_state) # type: ignore[arg-type] + + ctx2 = SessionContext(session_id="s1", input_messages=[Message(role="user", contents=["turn 2"])]) + await provider.before_run(agent=None, session=session, context=ctx2, state=provider_state) # type: ignore[arg-type] + ctx2._response = AgentResponse(messages=[Message(role="assistant", contents=["reply 2"])]) + await provider.after_run(agent=None, session=session, context=ctx2, state=provider_state) # type: ignore[arg-type] + + 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"] + class TestFileHistoryProvider: def test_is_marked_experimental(self) -> None: @@ -882,3 +953,64 @@ 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 diff --git a/python/packages/redis/agent_framework_redis/_history_provider.py b/python/packages/redis/agent_framework_redis/_history_provider.py index ef22511faa8..3b305ca204b 100644 --- a/python/packages/redis/agent_framework_redis/_history_provider.py +++ b/python/packages/redis/agent_framework_redis/_history_provider.py @@ -8,6 +8,7 @@ from __future__ import annotations +import json from collections.abc import Sequence from typing import Any, ClassVar @@ -17,6 +18,26 @@ from redis.credentials import CredentialProvider +def _get_message_identity(message: Message) -> tuple: + """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", 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", message.role, serialized) + except Exception: + return ("content", message.role, str(message.contents)) + + class RedisHistoryProvider(HistoryProvider): """Redis-backed history provider using the new HistoryProvider hooks pattern. @@ -152,7 +173,19 @@ async def save_messages( return key = self._redis_key(session_id) - serialized_messages = [self._serialize_json(msg) for msg in messages] + existing_message = await self.get_messages(session_id, state=state, **kwargs) + existing_identities = {_get_message_identity(m) for m in existing_message} + new_messages = [] + for msg in messages: + identity = _get_message_identity(msg) + if identity not in existing_identities: + existing_identities.add(identity) + new_messages.append(msg) + + 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: diff --git a/python/packages/redis/tests/test_providers.py b/python/packages/redis/tests/test_providers.py index 74ed328b404..19a460ee33d 100644 --- a/python/packages/redis/tests/test_providers.py +++ b/python/packages/redis/tests/test_providers.py @@ -559,3 +559,58 @@ 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.""" + + 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 From c42ebe4a71644a99904a5c1179b403fec00a9dbc Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Wed, 22 Jul 2026 23:13:50 +0530 Subject: [PATCH 2/6] fix: address review feedback for history deduplication --- .../core/agent_framework/_sessions.py | 51 ++++++++++--------- .../_history_provider.py | 26 +--------- 2 files changed, 27 insertions(+), 50 deletions(-) diff --git a/python/packages/core/agent_framework/_sessions.py b/python/packages/core/agent_framework/_sessions.py index ae388b1053c..e56f5c24b7e 100644 --- a/python/packages/core/agent_framework/_sessions.py +++ b/python/packages/core/agent_framework/_sessions.py @@ -83,16 +83,17 @@ def _get_message_identity(message: Message) -> tuple: 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, "id", None) + 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", msg_id) try: - # Use to_dict() for a stable, deterministic representation - serialized = json.dumps(message.to_dict(), sort_keys=True, ensure_ascii=False) + 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", message.role, serialized) except Exception: - # Fallback if serialization fails for any reason return ("content", message.role, str(message.contents)) @@ -1318,27 +1319,27 @@ async def save_messages( file_lock = self._session_write_lock(file_path) def _append_messages() -> None: - existing_identities: set[tuple] = set() - 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_identities.add(_get_message_identity(msg)) - except Exception: - logger.debug("Failed to parse history line for deduplication", exc_info=True) - continue - - with file_lock, file_path.open("a", encoding="utf-8") as file_handle: - for message in messages: - identity = _get_message_identity(message) - if identity not in existing_identities: - existing_identities.add(identity) - file_handle.write(f"{self._serialize_message(message)}\n") + with file_lock: + existing_identities: set[tuple] = set() + 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_identities.add(_get_message_identity(msg)) + except Exception: + logger.debug("failed to parse history line for deduplication") + continue + with file_path.open("a", encoding="utf-8") as file_handle: + for message in messages: + identity = _get_message_identity(message) + if identity not in existing_identities: + existing_identities.add(identity) + file_handle.write(f"{self._serialize_message(message)}\n") async with async_lock: await asyncio.to_thread(_append_messages) diff --git a/python/packages/redis/agent_framework_redis/_history_provider.py b/python/packages/redis/agent_framework_redis/_history_provider.py index 3b305ca204b..e8dbb5cd5e6 100644 --- a/python/packages/redis/agent_framework_redis/_history_provider.py +++ b/python/packages/redis/agent_framework_redis/_history_provider.py @@ -14,30 +14,10 @@ import redis.asyncio as redis from agent_framework import Message -from agent_framework._sessions import HistoryProvider +from agent_framework._sessions import HistoryProvider, _get_message_identity from redis.credentials import CredentialProvider -def _get_message_identity(message: Message) -> tuple: - """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", 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", message.role, serialized) - except Exception: - return ("content", message.role, str(message.contents)) - - class RedisHistoryProvider(HistoryProvider): """Redis-backed history provider using the new HistoryProvider hooks pattern. @@ -200,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: From 489f53656f97460c8c172c4691c3418d9f346de3 Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Fri, 24 Jul 2026 14:58:52 +0530 Subject: [PATCH 3/6] fix: Prevent superlinear history growth by deduplicating messages --- .../core/agent_framework/_sessions.py | 22 ++++++------ .../packages/core/tests/core/test_sessions.py | 35 +++++++++++++++---- .../_history_provider.py | 6 ++-- 3 files changed, 43 insertions(+), 20 deletions(-) diff --git a/python/packages/core/agent_framework/_sessions.py b/python/packages/core/agent_framework/_sessions.py index e56f5c24b7e..51e4870d502 100644 --- a/python/packages/core/agent_framework/_sessions.py +++ b/python/packages/core/agent_framework/_sessions.py @@ -56,6 +56,7 @@ JsonDumps: TypeAlias = Callable[[Any], str | bytes] JsonLoads: TypeAlias = Callable[[str | bytes], Any] ServiceSessionId: TypeAlias = Mapping[str, Any] +MessageIdentity: TypeAlias = tuple[str, ...] def _default_json_dumps(value: Any) -> str: @@ -77,7 +78,7 @@ def _deduplicate_origin_session_ids(origin_session_ids: Iterable[str]) -> list[s return unique_origin_session_ids -def _get_message_identity(message: Message) -> tuple: +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 @@ -87,14 +88,14 @@ def _get_message_identity(message: Message) -> tuple: if msg_id is None: msg_id = getattr(message, "id", None) if msg_id is not None: - return ("id", msg_id) + 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", message.role, serialized) + return ("content", str(message.role), serialized) except Exception: - return ("content", message.role, str(message.contents)) + return ("content", str(message.role), str(message.contents)) def _is_middleware_sequence( @@ -1135,12 +1136,11 @@ async def save_messages( if state is None: return existing = state.get("messages", []) - existing_identities = {_get_message_identity(m) for m in existing} + existing_id = {id(m) for m in existing} new_messages = [] for msg in messages: - identity = _get_message_identity(msg) - if identity not in existing_identities: - existing_identities.add(identity) + if id(msg) not in existing_id: + existing_id.add(id(msg)) new_messages.append(msg) if new_messages: state["messages"] = [*existing, *new_messages] @@ -1320,7 +1320,7 @@ async def save_messages( def _append_messages() -> None: with file_lock: - existing_identities: set[tuple] = set() + existing_identities: set[MessageIdentity] = set() if file_path.exists(): with file_path.open("r", encoding="utf-8") as f: for line in f: @@ -1330,13 +1330,13 @@ def _append_messages() -> None: try: payload = self.loads(line) msg = Message.from_dict(dict(cast(Mapping[str, Any], payload))) - existing_identities.add(_get_message_identity(msg)) + existing_identities.add(get_message_identity(msg)) except Exception: logger.debug("failed to parse history line for deduplication") continue with file_path.open("a", encoding="utf-8") as file_handle: for message in messages: - identity = _get_message_identity(message) + identity = get_message_identity(message) if identity not in existing_identities: existing_identities.add(identity) file_handle.write(f"{self._serialize_message(message)}\n") diff --git a/python/packages/core/tests/core/test_sessions.py b/python/packages/core/tests/core/test_sessions.py index d9135ba9ac1..93ce31333fc 100644 --- a/python/packages/core/tests/core/test_sessions.py +++ b/python/packages/core/tests/core/test_sessions.py @@ -6,7 +6,7 @@ import time from collections.abc import Awaitable, Callable, Sequence from pathlib import Path -from typing import Any +from typing import TYPE_CHECKING, Any, cast import pytest @@ -27,6 +27,9 @@ from agent_framework._sessions import LOCAL_HISTORY_CONVERSATION_ID, is_local_history_conversation_id from agent_framework.exceptions import MiddlewareException +if TYPE_CHECKING: + from agent_framework._agents import SupportsAgentRun + # --------------------------------------------------------------------------- # SessionContext tests # --------------------------------------------------------------------------- @@ -724,16 +727,36 @@ async def test_full_loop_does_not_grow_superlinearly(self) -> None: 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=None, session=session, context=ctx1, state=provider_state) # type: ignore[arg-type] + + 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=None, session=session, context=ctx1, state=provider_state) # type: ignore[arg-type] + 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=None, session=session, context=ctx2, state=provider_state) # type: ignore[arg-type] + 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=None, session=session, context=ctx2, state=provider_state) # type: ignore[arg-type] + 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 diff --git a/python/packages/redis/agent_framework_redis/_history_provider.py b/python/packages/redis/agent_framework_redis/_history_provider.py index e8dbb5cd5e6..b242f507c64 100644 --- a/python/packages/redis/agent_framework_redis/_history_provider.py +++ b/python/packages/redis/agent_framework_redis/_history_provider.py @@ -14,7 +14,7 @@ import redis.asyncio as redis from agent_framework import Message -from agent_framework._sessions import HistoryProvider, _get_message_identity +from agent_framework._sessions import HistoryProvider, MessageIdentity, get_message_identity from redis.credentials import CredentialProvider @@ -154,10 +154,10 @@ async def save_messages( key = self._redis_key(session_id) existing_message = await self.get_messages(session_id, state=state, **kwargs) - existing_identities = {_get_message_identity(m) for m in existing_message} + existing_identities: set[MessageIdentity] = {get_message_identity(m) for m in existing_message} new_messages = [] for msg in messages: - identity = _get_message_identity(msg) + identity = get_message_identity(msg) if identity not in existing_identities: existing_identities.add(identity) new_messages.append(msg) From 04d0e003877b843e32511a00064c4b07ae0cbd6a Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Thu, 30 Jul 2026 15:54:30 +0530 Subject: [PATCH 4/6] fix: add list[Message] type hints --- python/packages/core/agent_framework/_sessions.py | 2 +- .../packages/redis/agent_framework_redis/_history_provider.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/python/packages/core/agent_framework/_sessions.py b/python/packages/core/agent_framework/_sessions.py index 07850d07327..62520a0148c 100644 --- a/python/packages/core/agent_framework/_sessions.py +++ b/python/packages/core/agent_framework/_sessions.py @@ -1209,7 +1209,7 @@ async def save_messages( return existing = state.get("messages", []) existing_id = {id(m) for m in existing} - new_messages = [] + new_messages: list[Message] = [] for msg in messages: if id(msg) not in existing_id: existing_id.add(id(msg)) diff --git a/python/packages/redis/agent_framework_redis/_history_provider.py b/python/packages/redis/agent_framework_redis/_history_provider.py index b242f507c64..a9f999fd6d2 100644 --- a/python/packages/redis/agent_framework_redis/_history_provider.py +++ b/python/packages/redis/agent_framework_redis/_history_provider.py @@ -155,7 +155,7 @@ async def save_messages( key = self._redis_key(session_id) existing_message = await self.get_messages(session_id, state=state, **kwargs) existing_identities: set[MessageIdentity] = {get_message_identity(m) for m in existing_message} - new_messages = [] + new_messages: list[Message] = [] for msg in messages: identity = get_message_identity(msg) if identity not in existing_identities: From 5b946bd6ac081106d6e64d18659206d5f176d9e7 Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Mon, 3 Aug 2026 11:46:48 +0530 Subject: [PATCH 5/6] fix(sessions): resolve deduplication churn and collapsing of identical message --- .../core/agent_framework/_sessions.py | 16 +++--- .../packages/core/tests/core/test_sessions.py | 30 +++++++++++ .../_history_provider.py | 19 ++++++- python/packages/redis/tests/test_providers.py | 52 ++++++++++++++++++- 4 files changed, 107 insertions(+), 10 deletions(-) diff --git a/python/packages/core/agent_framework/_sessions.py b/python/packages/core/agent_framework/_sessions.py index e9a55771fbd..48955b56457 100644 --- a/python/packages/core/agent_framework/_sessions.py +++ b/python/packages/core/agent_framework/_sessions.py @@ -187,7 +187,6 @@ def _deduplicate_origin_session_ids(origin_session_ids: Iterable[str]) -> list[s 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. """ @@ -196,13 +195,18 @@ def get_message_identity(message: Message) -> MessageIdentity: msg_id = getattr(message, "id", None) if msg_id is not None: return ("id", str(msg_id)) - + + new_id = str(uuid.uuid4()) 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) + message.message_id = new_id + return ("id", new_id) except Exception: - return ("content", str(message.role), str(message.contents)) + 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 _is_middleware_sequence( diff --git a/python/packages/core/tests/core/test_sessions.py b/python/packages/core/tests/core/test_sessions.py index 1aca51a6268..b1bc270458c 100644 --- a/python/packages/core/tests/core/test_sessions.py +++ b/python/packages/core/tests/core/test_sessions.py @@ -1469,6 +1469,20 @@ async def test_full_loop_does_not_grow_superlinearly(self) -> None: 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 must both be persisted in memory.""" + 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], state=state) + assert len(state["messages"]) == 1 + + await provider.save_messages("s1", [yes_2], state=state) + assert len(state["messages"]) == 2 + class TestFileHistoryProvider: def test_is_marked_experimental(self) -> None: @@ -1813,3 +1827,19 @@ async def test_deduplication_file_integrity(self, tmp_path: Path) -> None: 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 separate identical user turns 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]) + loaded = await provider.get_messages("s1") + assert len(loaded) == 1 + + await provider.save_messages("s1", [yes_2]) + loaded = await provider.get_messages("s1") + + assert len(loaded) == 2 \ No newline at end of file diff --git a/python/packages/redis/agent_framework_redis/_history_provider.py b/python/packages/redis/agent_framework_redis/_history_provider.py index aa97f334ffd..ff143150fff 100644 --- a/python/packages/redis/agent_framework_redis/_history_provider.py +++ b/python/packages/redis/agent_framework_redis/_history_provider.py @@ -158,12 +158,23 @@ async def save_messages( return key = self._redis_key(session_id) + seen_key = f"{key}:seen" + existing_message = await self.get_messages(session_id, state=state, **kwargs) existing_identities: set[MessageIdentity] = {get_message_identity(m) for m in existing_message} + + seen_members = await self._redis_client.smembers(seen_key) + seen_identities: set[MessageIdentity] = set() + for raw in seen_members: + try: + seen_identities.add(tuple(json.loads(raw))) + except Exception: + continue + new_messages: list[Message] = [] for msg in messages: identity = get_message_identity(msg) - if identity not in existing_identities: + if identity not in existing_identities and identity not in seen_identities: existing_identities.add(identity) new_messages.append(msg) @@ -175,6 +186,9 @@ async def save_messages( async with self._redis_client.pipeline(transaction=True) as pipe: for serialized in serialized_messages: await pipe.rpush(key, serialized) # type: ignore[misc] + for msg in new_messages: + identity = get_message_identity(msg) + await pipe.sadd(seen_key, json.dumps(list(identity))) # type: ignore[misc] await pipe.execute() if self.max_messages is not None: @@ -198,7 +212,8 @@ async def clear(self, session_id: str | None) -> None: Args: session_id: The session ID to clear messages for. """ - await self._redis_client.delete(self._redis_key(session_id)) + key = self._redis_key(session_id) + await self._redis_client.delete(key, f"{key}:seen") async def aclose(self) -> None: """Close the Redis connection.""" diff --git a/python/packages/redis/tests/test_providers.py b/python/packages/redis/tests/test_providers.py index edcc953ce8f..5a66d615b79 100644 --- a/python/packages/redis/tests/test_providers.py +++ b/python/packages/redis/tests/test_providers.py @@ -10,7 +10,7 @@ import pytest from agent_framework import AgentResponse, Message -from agent_framework._sessions import AgentSession, SessionContext +from agent_framework._sessions import AgentSession, SessionContext, get_message_identity from agent_framework_redis._context_provider import RedisContextProvider from agent_framework_redis._feature_usage import FeatureIndex @@ -63,9 +63,12 @@ def mock_redis_client(): client.llen = AsyncMock(return_value=0) client.ltrim = AsyncMock() client.delete = AsyncMock() + client.smembers = AsyncMock(return_value=set()) + mock_pipeline = AsyncMock() mock_pipeline.rpush = AsyncMock() + mock_pipeline.sadd = AsyncMock() mock_pipeline.execute = AsyncMock() client.pipeline.return_value.__aenter__.return_value = mock_pipeline @@ -503,7 +506,7 @@ async def test_clear_calls_delete(self, mock_redis_client: MagicMock): provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379") await provider.clear("session-1") - mock_redis_client.delete.assert_called_once_with("chat_messages:session-1") + mock_redis_client.delete.assert_called_once_with("chat_messages:session-1", "chat_messages:session-1:seen") class TestRedisHistoryProviderBeforeAfterRun: @@ -624,3 +627,48 @@ 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 + +class TestRedisHistoryProviderDeduplication: + """Tests for Redis save_messages deduplication and trimming behavior.""" + + 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.""" + 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())]) + + mock_redis_client.smembers = AsyncMock( + return_value=[ + json.dumps(list(get_message_identity(msg_old))), + json.dumps(list(get_message_identity(msg_new))), + ] + ) + + 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=[]) + mock_redis_client.smembers = AsyncMock(return_value=set()) + + 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 \ No newline at end of file From 2b499a7b66d9133495b3d150b396387ded12bdf8 Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Mon, 3 Aug 2026 15:08:54 +0530 Subject: [PATCH 6/6] fix(sessions): replace uuid/seen-set dedup with sequence aware filtering --- .../core/agent_framework/_sessions.py | 88 ++++++++++++++----- .../packages/core/tests/core/test_sessions.py | 26 +++--- .../_history_provider.py | 30 ++----- python/packages/redis/tests/test_providers.py | 37 +++----- 4 files changed, 93 insertions(+), 88 deletions(-) diff --git a/python/packages/core/agent_framework/_sessions.py b/python/packages/core/agent_framework/_sessions.py index 48955b56457..ae218f8e94f 100644 --- a/python/packages/core/agent_framework/_sessions.py +++ b/python/packages/core/agent_framework/_sessions.py @@ -187,6 +187,7 @@ def _deduplicate_origin_session_ids(origin_session_ids: Iterable[str]) -> list[s 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. """ @@ -195,18 +196,61 @@ def get_message_identity(message: Message) -> MessageIdentity: msg_id = getattr(message, "id", None) if msg_id is not None: return ("id", str(msg_id)) - - new_id = str(uuid.uuid4()) + try: - message.message_id = new_id - return ("id", new_id) + 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: - 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)) + 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( @@ -1858,15 +1902,11 @@ 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", []) - existing_id = {id(m) for m in existing} - new_messages: list[Message] = [] - for msg in messages: - if id(msg) not in existing_id: - existing_id.add(id(msg)) - new_messages.append(msg) + new_messages = filter_new_messages(existing, messages) + if new_messages: state["messages"] = [*existing, *new_messages] @@ -2027,7 +2067,7 @@ async def save_messages( def _append_messages() -> None: with file_lock: if self.serialization_format == "json": - existing_identities: set[MessageIdentity] = set() + existing_messages: list[Message] = [] if file_path.exists(): with file_path.open("r", encoding="utf-8") as f: for line in f: @@ -2037,15 +2077,15 @@ def _append_messages() -> None: try: payload = self.loads(line) msg = Message.from_dict(dict(cast(Mapping[str, Any], payload))) - existing_identities.add(get_message_identity(msg)) + existing_messages.append(msg) except Exception: logger.debug("failed to parse history line for deduplication") continue - with file_path.open("a", encoding="utf-8") as file_handle: - for message in messages: - identity = get_message_identity(message) - if identity not in existing_identities: - existing_identities.add(identity) + + 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: diff --git a/python/packages/core/tests/core/test_sessions.py b/python/packages/core/tests/core/test_sessions.py index b1bc270458c..38d634772a9 100644 --- a/python/packages/core/tests/core/test_sessions.py +++ b/python/packages/core/tests/core/test_sessions.py @@ -1470,18 +1470,18 @@ async def test_full_loop_does_not_grow_superlinearly(self) -> None: 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 must both be persisted in memory.""" + """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], state=state) - assert len(state["messages"]) == 1 - - await provider.save_messages("s1", [yes_2], state=state) + + 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: @@ -1829,17 +1829,15 @@ async def test_deduplication_file_integrity(self, tmp_path: Path) -> None: assert len(raw_lines) == 3 async def test_save_messages_preserves_duplicate_content(self, tmp_path: Path) -> None: - """Test that two separate identical user turns are both persisted.""" + """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]) + await provider.save_messages("s1", [yes_1, yes_2]) loaded = await provider.get_messages("s1") - assert len(loaded) == 1 - await provider.save_messages("s1", [yes_2]) - loaded = await provider.get_messages("s1") - - assert len(loaded) == 2 \ No newline at end of file + 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 ff143150fff..8805b0b2330 100644 --- a/python/packages/redis/agent_framework_redis/_history_provider.py +++ b/python/packages/redis/agent_framework_redis/_history_provider.py @@ -14,7 +14,7 @@ import redis.asyncio as redis from agent_framework import Message -from agent_framework._sessions import HistoryProvider, MessageIdentity, get_message_identity +from agent_framework._sessions import HistoryProvider, filter_new_messages from agent_framework._telemetry import mark_feature_used from redis.credentials import CredentialProvider @@ -158,25 +158,9 @@ async def save_messages( return key = self._redis_key(session_id) - seen_key = f"{key}:seen" - - existing_message = await self.get_messages(session_id, state=state, **kwargs) - existing_identities: set[MessageIdentity] = {get_message_identity(m) for m in existing_message} - - seen_members = await self._redis_client.smembers(seen_key) - seen_identities: set[MessageIdentity] = set() - for raw in seen_members: - try: - seen_identities.add(tuple(json.loads(raw))) - except Exception: - continue - - new_messages: list[Message] = [] - for msg in messages: - identity = get_message_identity(msg) - if identity not in existing_identities and identity not in seen_identities: - existing_identities.add(identity) - new_messages.append(msg) + + existing_messages = await self.get_messages(session_id, state=state, **kwargs) + new_messages = filter_new_messages(existing_messages, messages) if not new_messages: return @@ -186,9 +170,6 @@ async def save_messages( async with self._redis_client.pipeline(transaction=True) as pipe: for serialized in serialized_messages: await pipe.rpush(key, serialized) # type: ignore[misc] - for msg in new_messages: - identity = get_message_identity(msg) - await pipe.sadd(seen_key, json.dumps(list(identity))) # type: ignore[misc] await pipe.execute() if self.max_messages is not None: @@ -212,8 +193,7 @@ async def clear(self, session_id: str | None) -> None: Args: session_id: The session ID to clear messages for. """ - key = self._redis_key(session_id) - await self._redis_client.delete(key, f"{key}:seen") + await self._redis_client.delete(self._redis_key(session_id)) async def aclose(self) -> None: """Close the Redis connection.""" diff --git a/python/packages/redis/tests/test_providers.py b/python/packages/redis/tests/test_providers.py index 5a66d615b79..f5a004d15f8 100644 --- a/python/packages/redis/tests/test_providers.py +++ b/python/packages/redis/tests/test_providers.py @@ -10,7 +10,7 @@ import pytest from agent_framework import AgentResponse, Message -from agent_framework._sessions import AgentSession, SessionContext, get_message_identity +from agent_framework._sessions import AgentSession, SessionContext from agent_framework_redis._context_provider import RedisContextProvider from agent_framework_redis._feature_usage import FeatureIndex @@ -63,12 +63,9 @@ def mock_redis_client(): client.llen = AsyncMock(return_value=0) client.ltrim = AsyncMock() client.delete = AsyncMock() - client.smembers = AsyncMock(return_value=set()) - mock_pipeline = AsyncMock() mock_pipeline.rpush = AsyncMock() - mock_pipeline.sadd = AsyncMock() mock_pipeline.execute = AsyncMock() client.pipeline.return_value.__aenter__.return_value = mock_pipeline @@ -506,7 +503,7 @@ async def test_clear_calls_delete(self, mock_redis_client: MagicMock): provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379") await provider.clear("session-1") - mock_redis_client.delete.assert_called_once_with("chat_messages:session-1", "chat_messages:session-1:seen") + mock_redis_client.delete.assert_called_once_with("chat_messages:session-1") class TestRedisHistoryProviderBeforeAfterRun: @@ -575,7 +572,7 @@ async def test_after_run_skips_when_no_messages(self, mock_redis_client: MagicMo class TestRedisHistoryProviderDeduplication: - """Tests for Redis save_messages deduplication.""" + """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"]) @@ -628,24 +625,15 @@ 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 -class TestRedisHistoryProviderDeduplication: - """Tests for Redis save_messages deduplication and trimming behavior.""" - 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.""" + 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())]) - - mock_redis_client.smembers = AsyncMock( - return_value=[ - json.dumps(list(get_message_identity(msg_old))), - json.dumps(list(get_message_identity(msg_new))), - ] - ) - + 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") @@ -653,22 +641,21 @@ async def test_trimmed_messages_not_reappended(self, mock_redis_client: MagicMoc 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=[]) - mock_redis_client.smembers = AsyncMock(return_value=set()) - + 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 \ No newline at end of file + assert pipeline.rpush.call_count == 2