Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion python/packages/core/agent_framework/_sessions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
50 changes: 42 additions & 8 deletions python/packages/core/tests/core/test_sessions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down Expand Up @@ -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"])
Expand All @@ -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"])
Expand All @@ -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"])
Expand All @@ -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"])
Expand Down
Loading