Skip to content
Open
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
99 changes: 94 additions & 5 deletions python/packages/core/agent_framework/_sessions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down Expand Up @@ -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]]:
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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:
Expand Down
185 changes: 184 additions & 1 deletion python/packages/core/tests/core/test_sessions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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"
16 changes: 10 additions & 6 deletions python/packages/redis/agent_framework_redis/_history_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down
Loading
Loading