diff --git a/python/packages/core/agent_framework/_sessions.py b/python/packages/core/agent_framework/_sessions.py index da8cfbcb318..9fcce18a071 100644 --- a/python/packages/core/agent_framework/_sessions.py +++ b/python/packages/core/agent_framework/_sessions.py @@ -2335,8 +2335,12 @@ def _append_messages() -> None: for message in new_messages: file_handle.write(f"{self._serialize_json_message(message)}\n") return + existing_messages = self._read_msgpack_messages(file_path) if file_path.exists() else [] + new_messages = filter_new_messages(existing_messages, messages) + if not new_messages: + return with file_path.open("ab") as file_handle: - for message in messages: + for message in new_messages: serialized = _DEFAULT_MSGPACK_ENCODER.encode(message.to_dict()) file_handle.write(len(serialized).to_bytes(self._MSGPACK_RECORD_HEADER_BYTES, "big")) file_handle.write(serialized) diff --git a/python/packages/core/tests/core/test_sessions.py b/python/packages/core/tests/core/test_sessions.py index 193859d3aa8..71457093d84 100644 --- a/python/packages/core/tests/core/test_sessions.py +++ b/python/packages/core/tests/core/test_sessions.py @@ -1636,6 +1636,28 @@ async def test_stores_and_loads_length_prefixed_msgpack(self, tmp_path: Path) -> assert first_record_length > 0 assert raw[4 : 4 + first_record_length] == msgspec.msgpack.encode(messages[0].to_dict()) + @pytest.mark.parametrize("serialization_format", ["json", "msgpack"]) + async def test_save_messages_deduplicates_replayed_transcript( + self, tmp_path: Path, serialization_format: Literal["json", "msgpack"] + ) -> None: + provider = FileHistoryProvider(tmp_path, serialization_format=serialization_format) + first_turn = [ + Message(role="user", contents=["hello"]), + Message(role="assistant", contents=["hi there"]), + ] + full_transcript = [ + Message(role="user", contents=["hello"]), + Message(role="assistant", contents=["hi there"]), + Message(role="user", contents=["follow-up"]), + Message(role="assistant", contents=["reply"]), + ] + + await provider.save_messages("replayed-transcript", first_turn) + await provider.save_messages("replayed-transcript", full_transcript) + + loaded = await provider.get_messages("replayed-transcript") + assert [message.text for message in loaded] == ["hello", "hi there", "follow-up", "reply"] + @pytest.mark.parametrize("serialization_format", ["json", "msgpack"]) async def test_round_trips_marked_refusal_text( self, tmp_path: Path, serialization_format: Literal["json", "msgpack"] @@ -1913,9 +1935,12 @@ def tracked_open(path: Path, *args: Any, **kwargs: Any) -> Any: 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: + @pytest.mark.parametrize("serialization_format", ["json", "msgpack"]) + async def test_save_messages_deduplicates_identical_messages( + self, tmp_path: Path, serialization_format: Literal["json", "msgpack"] + ) -> None: """Test that FileHistoryProvider does not re-append already persisted messages.""" - provider = FileHistoryProvider(tmp_path) + provider = FileHistoryProvider(tmp_path, serialization_format=serialization_format) msg1 = Message(role="user", contents=["hello"]) msg2 = Message(role="assistant", contents=["hi there"]) @@ -1928,9 +1953,12 @@ async def test_save_messages_deduplicates_identical_messages(self, tmp_path: Pat loaded = await provider.get_messages("s1") assert len(loaded) == 2 - async def test_save_messages_only_appends_new_messages(self, tmp_path: Path) -> None: + @pytest.mark.parametrize("serialization_format", ["json", "msgpack"]) + async def test_save_messages_only_appends_new_messages( + self, tmp_path: Path, serialization_format: Literal["json", "msgpack"] + ) -> None: """Test that FileHistoryProvider filters out old messages and only appends new ones""" - provider = FileHistoryProvider(tmp_path) + provider = FileHistoryProvider(tmp_path, serialization_format=serialization_format) msg1 = Message(role="user", contents=["hello"]) msg2 = Message(role="assistant", contents=["hi there"]) @@ -1945,9 +1973,12 @@ async def test_save_messages_only_appends_new_messages(self, tmp_path: Path) -> 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: + @pytest.mark.parametrize("serialization_format", ["json", "msgpack"]) + async def test_save_messages_different_roles_same_text_not_deduplicated( + self, tmp_path: Path, serialization_format: Literal["json", "msgpack"] + ) -> None: """Test that messages with the same text but different roles are kept separate.""" - provider = FileHistoryProvider(tmp_path) + provider = FileHistoryProvider(tmp_path, serialization_format=serialization_format) msg1 = Message(role="user", contents=["ping"]) msg2 = Message(role="assistant", contents=["ping"]) @@ -1974,9 +2005,12 @@ async def test_deduplication_file_integrity(self, tmp_path: Path) -> None: 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: + @pytest.mark.parametrize("serialization_format", ["json", "msgpack"]) + async def test_save_messages_preserves_duplicate_content( + self, tmp_path: Path, serialization_format: Literal["json", "msgpack"] + ) -> None: """Test that two identical user turns in the same batch are both persisted.""" - provider = FileHistoryProvider(tmp_path) + provider = FileHistoryProvider(tmp_path, serialization_format=serialization_format) yes_1 = Message(role="user", contents=["yes"]) yes_2 = Message(role="user", contents=["yes"])