From e6087fb139ed0ecf5997534b9634ca0fa8b2de10 Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Wed, 22 Jul 2026 00:15:21 +0530 Subject: [PATCH 1/3] 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/3] 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/3] 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)