diff --git a/python/packages/redis/agent_framework_redis/_history_provider.py b/python/packages/redis/agent_framework_redis/_history_provider.py index a7703db8a2..f260e3f2da 100644 --- a/python/packages/redis/agent_framework_redis/_history_provider.py +++ b/python/packages/redis/agent_framework_redis/_history_provider.py @@ -106,9 +106,13 @@ def __init__( else: self._redis_client = redis.from_url(redis_url, decode_responses=True) # type: ignore[no-untyped-call] + # Unit separator: source ids and session ids are opaque strings and can + # legitimately contain ':', which would make colon-joined keys ambiguous. + _KEY_SEP = "\x1f" + def _redis_key(self, session_id: str | None) -> str: """Get the Redis key for a given session's messages.""" - return f"{self.key_prefix}:{session_id or 'default'}" + return self._KEY_SEP.join([self.key_prefix, self.source_id, session_id or "default"]) async def get_messages( self, diff --git a/python/packages/redis/tests/test_providers.py b/python/packages/redis/tests/test_providers.py index 55aee29662..a73410182d 100644 --- a/python/packages/redis/tests/test_providers.py +++ b/python/packages/redis/tests/test_providers.py @@ -420,8 +420,16 @@ 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") == "msgs\x1fmem\x1fsession-123" + assert provider._redis_key(None) == "msgs\x1fmem\x1fdefault" + + 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: @@ -482,7 +490,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("chat_messages\x1fmem\x1fs1", -10, -1) async def test_no_trim_when_under_limit(self, mock_redis_client: MagicMock): mock_redis_client.llen = AsyncMock(return_value=3) @@ -503,7 +511,19 @@ 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("chat_messages\x1fmem\x1fsession-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("chat_messages\x1faudit\x1fsession-1") + assert primary._redis_key("session-1") not in mock_redis_client.delete.call_args.args class TestRedisHistoryProviderBeforeAfterRun: