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 @@ -9,6 +9,7 @@
from __future__ import annotations

from collections.abc import Sequence
from contextlib import suppress
from typing import Any, ClassVar

import redis.asyncio as redis
Expand Down Expand Up @@ -111,8 +112,19 @@ 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>" and migrate lazily on first read.

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 @@ -135,6 +147,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]
if not redis_messages:
# Lazy migration: a session last written with the pre-scoping key
# layout moves under its new key on first read. renamenx keeps this
# atomic and no-ops if a concurrent write already landed there; the
# legacy list then stays put and remains clearable via clear().
legacy_key = self._legacy_redis_key(session_id)
if legacy_key != key and await self._redis_client.exists(legacy_key): # type: ignore[misc]
with suppress(Exception): # a legacy key that vanished mid-read is a no-op
await self._redis_client.renamenx(legacy_key, key) # type: ignore[misc]
redis_messages = await self._redis_client.lrange(key, 0, -1) # type: ignore[misc]
messages: list[Message] = []
if redis_messages:
for serialized in redis_messages: # type: ignore[union-attr]
Expand Down Expand Up @@ -202,7 +224,7 @@ async def clear(self, session_id: str | None) -> None:
Args:
session_id: The session ID to clear messages for.
"""
await self._redis_client.delete(self._redis_key(session_id))
await self._redis_client.delete(self._redis_key(session_id), self._legacy_redis_key(session_id))

async def aclose(self) -> None:
"""Close the Redis connection."""
Expand Down
71 changes: 67 additions & 4 deletions python/packages/redis/tests/test_providers.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,8 @@ def mock_redis_client():
client.llen = AsyncMock(return_value=0)
client.ltrim = AsyncMock()
client.delete = AsyncMock()
client.exists = AsyncMock(return_value=0)
client.renamenx = AsyncMock(return_value=0)

mock_pipeline = AsyncMock()
mock_pipeline.rpush = AsyncMock()
Expand Down Expand Up @@ -424,8 +426,26 @@ 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:
Expand Down Expand Up @@ -455,6 +475,35 @@ async def test_empty_returns_empty(self, mock_redis_client: MagicMock):
messages = await provider.get_messages("s1")
assert messages == []

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

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")
mock_redis_client.renamenx.assert_called_once_with("chat_messages:s1", "13:chat_messages3:mem2:s1")
assert len(messages) == 1
assert messages[0].text == "legacy hello"

async def test_no_migration_when_new_key_has_data(self, mock_redis_client: MagicMock):
msg = Message(role="user", contents=["current"])
mock_redis_client.lrange = AsyncMock(return_value=[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
provider = RedisHistoryProvider("mem", redis_url="redis://localhost:6379")

messages = await provider.get_messages("s1")
mock_redis_client.exists.assert_not_called()
assert len(messages) == 1


class TestRedisHistoryProviderSaveMessages:
async def test_saves_serialized_messages(self, mock_redis_client: MagicMock):
Expand Down Expand Up @@ -486,7 +535,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,7 +591,21 @@ 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", "chat_messages: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
mock_redis_client.delete.assert_called_once_with(
"13:chat_messages5:audit9:session-1", "chat_messages:session-1"
)
assert primary._redis_key("session-1") not in mock_redis_client.delete.call_args.args


class TestRedisHistoryProviderBeforeAfterRun:
Expand Down
Loading