diff --git a/python/packages/redis/agent_framework_redis/_history_provider.py b/python/packages/redis/agent_framework_redis/_history_provider.py index 3443c3b9a2c..002ba486241 100644 --- a/python/packages/redis/agent_framework_redis/_history_provider.py +++ b/python/packages/redis/agent_framework_redis/_history_provider.py @@ -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 @@ -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 (":") 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 + # ":" 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( @@ -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): + with suppress(Exception): # a legacy key that vanished mid-read is a no-op + await self._redis_client.renamenx(legacy_key, key) + 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] @@ -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.""" diff --git a/python/packages/redis/tests/test_providers.py b/python/packages/redis/tests/test_providers.py index eb892095332..a893cabbf3a 100644 --- a/python/packages/redis/tests/test_providers.py +++ b/python/packages/redis/tests/test_providers.py @@ -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() @@ -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: @@ -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): @@ -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) @@ -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: