Skip to content
Original file line number Diff line number Diff line change
Expand Up @@ -112,8 +112,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 All @@ -136,6 +148,16 @@ async def get_messages(
mark_feature_used(FeatureIndex.REDIS)
key = self._redis_key(session_id)
redis_messages: list[str] = await self._redis_client.lrange(key, 0, -1) # type: ignore[misc]
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: list[str] = await self._redis_client.lrange(legacy_key, 0, -1) # type: ignore[misc]
if legacy_messages:
redis_messages = [*legacy_messages, *redis_messages]
messages: list[Message] = []
if redis_messages:
for serialized in redis_messages: # type: ignore[union-attr]
Expand Down Expand Up @@ -165,8 +187,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 @@ -203,6 +224,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
88 changes: 76 additions & 12 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,6 +474,40 @@ 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"]


class TestRedisHistoryProviderSaveMessages:
async def test_saves_serialized_messages(self, mock_redis_client: MagicMock):
Expand Down Expand Up @@ -486,7 +539,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 @@ -518,13 +571,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 @@ -542,15 +593,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
Loading