diff --git a/python/packages/redis/agent_framework_redis/_history_provider.py b/python/packages/redis/agent_framework_redis/_history_provider.py index 8ec968a66e..660d5f079b 100644 --- a/python/packages/redis/agent_framework_redis/_history_provider.py +++ b/python/packages/redis/agent_framework_redis/_history_provider.py @@ -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 (":") 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 + # ":"; 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( @@ -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. + 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))) @@ -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 @@ -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. """ diff --git a/python/packages/redis/tests/test_providers.py b/python/packages/redis/tests/test_providers.py index c6ba48a0af..d885f4d25f 100644 --- a/python/packages/redis/tests/test_providers.py +++ b/python/packages/redis/tests/test_providers.py @@ -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() @@ -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 @@ -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 @@ -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) @@ -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 @@ -566,7 +618,20 @@ 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: @@ -574,7 +639,7 @@ class TestRedisHistoryProviderBeforeAfterRun: 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 @@ -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