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
Original file line number Diff line number Diff line change
Expand Up @@ -130,8 +130,20 @@ def __init__(
else:
self._redis_client = redis.from_url(redis_url, decode_responses=True) # type: ignore[no-untyped-call]

# Keys length-prefix each component ("<len>:<value>") so the join stays
# injective no matter which bytes the source/session ids carry; any fixed
# separator can be smuggled inside an opaque id and collide two sessions.
# Sessions written before source_id scoping live under
# "<key_prefix>:<session_id>"; reads merge that legacy list in place and
# writes only ever touch the scoped key, so upgrading never moves data.

def _redis_key(self, session_id: str | None) -> str:
"""Get the Redis key for a given session's messages."""
parts = (self.key_prefix, self.source_id, session_id or "default")
return "".join(f"{len(part)}:{part}" for part in parts)

def _legacy_redis_key(self, session_id: str | None) -> str:
"""Pre-scoping key layout, read only to migrate existing sessions."""
return f"{self.key_prefix}:{session_id or 'default'}"

async def get_messages(
Expand Down Expand Up @@ -159,6 +171,18 @@ async def get_messages(
# version-independent, and the outer cast pins the element type that
# ``decode_responses=True`` guarantees.
redis_messages = cast("list[str]", await _redis_result(cast("Any", self._redis_client).lrange(key, 0, -1)))
legacy_key = self._legacy_redis_key(session_id)
if legacy_key != key:
# Sessions last written before source_id scoping stay readable in
# place: the merged view is the legacy list followed by the scoped
# one, since new writes only ever land on the scoped key. The legacy
# key is never renamed or deleted on this path; an upgrade cannot
# fork history, and removing the old key is an explicit admin call.
Comment thread
he-yufeng marked this conversation as resolved.
legacy_messages = cast(
"list[str]", await _redis_result(cast("Any", self._redis_client).lrange(legacy_key, 0, -1))
)
if legacy_messages:
redis_messages = [*legacy_messages, *redis_messages]
messages: list[Message] = []
for serialized in redis_messages:
messages.append(Message.from_dict(self._deserialize_json(serialized)))
Expand Down Expand Up @@ -187,8 +211,7 @@ async def save_messages(
if self.max_messages == 0:
# Retention is disabled. Trimming cannot express this - LTRIM key 0 -1 keeps
# the whole list - so return before serializing: no payload reaches Redis, an
# AOF or a replica. Stored history is deliberately left alone. _redis_key omits
# source_id, so the list can belong to a co-located provider, and removing
# AOF, or a replica. Stored history is deliberately left alone; removing
# stored history is what clear() is for.
return

Expand Down Expand Up @@ -225,6 +248,10 @@ def _deserialize_json(data: str) -> dict[str, Any]:
async def clear(self, session_id: str | None) -> None:
"""Clear all messages for a session.

Only the scoped key is deleted. A pre-scoping legacy list belongs to
whichever sources shared it, so it is left for an explicit admin
cleanup rather than being removed by one source's clear().

Args:
session_id: The session ID to clear messages for.
"""
Expand Down
93 changes: 79 additions & 14 deletions python/packages/redis/tests/test_providers.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@ def mock_redis_client():
client.llen = AsyncMock(return_value=0)
client.ltrim = AsyncMock()
client.delete = AsyncMock()
client.exists = AsyncMock(return_value=0)

mock_pipeline = AsyncMock()
mock_pipeline.rpush = AsyncMock()
Expand Down Expand Up @@ -424,15 +425,33 @@ def test_key_format(self, mock_redis_client: MagicMock):
mock_from_url.return_value = mock_redis_client
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379", key_prefix="msgs")

assert provider._redis_key("session-123") == "msgs:session-123"
assert provider._redis_key(None) == "msgs:default"
assert provider._redis_key("session-123") == "4:msgs3:mem11:session-123"
assert provider._redis_key(None) == "4:msgs3:mem7:default"

def test_key_join_is_injective(self, mock_redis_client: MagicMock):
# moonbox3's review case: any fixed separator can be smuggled inside an
# opaque id, so the components are length-prefixed instead.
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
first = RedisHistoryProvider("audit\x1fx", redis_url="redis://localhost:6379", key_prefix="msgs")
second = RedisHistoryProvider("audit", redis_url="redis://localhost:6379", key_prefix="msgs")

assert first._redis_key("y") != second._redis_key("x\x1fy")

def test_keys_isolated_per_source_id(self, mock_redis_client: MagicMock):
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
first = RedisHistoryProvider("audit", redis_url="redis://localhost:6379", key_prefix="msgs")
second = RedisHistoryProvider("primary", redis_url="redis://localhost:6379", key_prefix="msgs")

assert first._redis_key("s1") != second._redis_key("s1")


class TestRedisHistoryProviderGetMessages:
async def test_returns_deserialized_messages(self, mock_redis_client: MagicMock):
msg1 = Message(role="user", contents=["Hello"])
msg2 = Message(role="assistant", contents=["Hi!"])
mock_redis_client.lrange = AsyncMock(return_value=[json.dumps(msg1.to_dict()), json.dumps(msg2.to_dict())])
mock_redis_client.lrange = AsyncMock(side_effect=[[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
Expand All @@ -455,10 +474,45 @@ async def test_empty_returns_empty(self, mock_redis_client: MagicMock):
messages = await provider.get_messages("s1")
assert messages == []

async def test_legacy_key_merges_in_place_on_read(self, mock_redis_client: MagicMock):
msg = Message(role="user", contents=["legacy hello"])
legacy_payload = json.dumps(msg.to_dict())
# scoped key empty, legacy key still holds the pre-scoping data
mock_redis_client.lrange = AsyncMock(side_effect=[[], [legacy_payload]])

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")

messages = await provider.get_messages("s1")
assert len(messages) == 1
assert messages[0].text == "legacy hello"
# the legacy list stays put: an upgrade must not move data out from
# under older instances that still read it
mock_redis_client.renamenx.assert_not_called()
mock_redis_client.delete.assert_not_called()

async def test_mixed_version_writes_merge_legacy_first(self, mock_redis_client: MagicMock):
old_world = Message(role="user", contents=["written by old version"])
new_world = Message(role="assistant", contents=["written by new version"])
# rolling upgrade: new writes land on the scoped key, the legacy key
# still holds what older instances wrote; both stay visible
mock_redis_client.lrange = AsyncMock(
side_effect=[[json.dumps(new_world.to_dict())], [json.dumps(old_world.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")

messages = await provider.get_messages("s1")
assert [m.text for m in messages] == ["written by old version", "written by new version"]

async def test_returns_messages_when_lrange_is_synchronous(self, mock_redis_client: MagicMock):
"""redis-py types several commands as returning a value or an awaitable; handle both."""
msg = Message(role="user", contents=["Hello"])
mock_redis_client.lrange = MagicMock(return_value=[json.dumps(msg.to_dict())])
# scoped key holds the message; the pre-scoping legacy key is empty
mock_redis_client.lrange = MagicMock(side_effect=[[json.dumps(msg.to_dict())], []])

with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
Expand Down Expand Up @@ -510,7 +564,7 @@ async def test_max_messages_trimming(self, mock_redis_client: MagicMock):

await provider.save_messages("s1", [Message(role="user", contents=["msg"])])

mock_redis_client.ltrim.assert_called_once_with("chat_messages:s1", -10, -1)
mock_redis_client.ltrim.assert_called_once_with("13:chat_messages3:mem2:s1", -10, -1)

async def test_no_trim_when_under_limit(self, mock_redis_client: MagicMock):
mock_redis_client.llen = AsyncMock(return_value=3)
Expand Down Expand Up @@ -542,13 +596,11 @@ async def test_max_messages_zero_retains_nothing(self, mock_redis_client: MagicM
mock_redis_client.ltrim.assert_not_called()

async def test_max_messages_zero_leaves_stored_history_alone(self, mock_redis_client: MagicMock):
"""Disabling retention must not delete history this provider does not own.
"""A retention setting of zero writes nothing and must not touch stored history.

``_redis_key`` omits ``source_id``, so two providers with the default prefix
share ``{key_prefix}:{session_id}``. Persisting runs in reverse provider order,
so a zero-retention provider that deleted the key would drop a co-located
provider's just-written history on every turn. Removing stored history is
``clear()``'s job, not a retention setting's.
``LTRIM key 0 -1`` is Redis's "keep the whole list", so trimming cannot
express zero, and deleting would conflate a retention setting with what
only ``clear()`` is allowed to do.
"""
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
Expand All @@ -566,15 +618,28 @@ 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("13:chat_messages3:mem9:session-1")

async def test_clear_leaves_other_source_ids_untouched(self, mock_redis_client: MagicMock):
with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
audit = RedisHistoryProvider("audit", redis_url="redis://localhost:6379")
primary = RedisHistoryProvider("primary", redis_url="redis://localhost:6379")

await audit.clear("session-1")
# the destructive case from #7471: clearing one provider must not
# delete the shared session's messages belonging to another provider,
# and the pre-scoping legacy list is never this provider's to delete
mock_redis_client.delete.assert_called_once_with("13:chat_messages5:audit9:session-1")
assert primary._redis_key("session-1") not in mock_redis_client.delete.call_args.args


class TestRedisHistoryProviderBeforeAfterRun:
"""Test before_run/after_run integration via HistoryProvider defaults."""

async def test_before_run_loads_history(self, mock_redis_client: MagicMock):
msg = Message(role="user", contents=["old msg"])
mock_redis_client.lrange = AsyncMock(return_value=[json.dumps(msg.to_dict())])
mock_redis_client.lrange = AsyncMock(side_effect=[[json.dumps(msg.to_dict())], []])

with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
Expand Down Expand Up @@ -695,7 +760,7 @@ async def test_trimmed_messages_not_reappended(self, mock_redis_client: MagicMoc
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.lrange = AsyncMock(side_effect=[[json.dumps(msg_new.to_dict())], []])

with patch("agent_framework_redis._history_provider.redis.from_url") as mock_from_url:
mock_from_url.return_value = mock_redis_client
Expand Down
Loading