diff --git a/src/google/adk/integrations/redis/_redis_session_service.py b/src/google/adk/integrations/redis/_redis_session_service.py index 2717b76720..041940db1f 100644 --- a/src/google/adk/integrations/redis/_redis_session_service.py +++ b/src/google/adk/integrations/redis/_redis_session_service.py @@ -382,9 +382,23 @@ async def append_event(self, session: Session, event: Event) -> Event: ) key = self._session_key(session.app_name, session.user_id, session.id) - await client.set( - key, - storage_session.model_dump_json(), - ex=self.config.ttl_seconds if self.config.ttl_seconds > 0 else None, - ) + + async def append_to_storage(pipe: Any) -> None: + raw = await pipe.get(key) + if raw: + # The caller may hold a filtered view. Reload on every transaction + # attempt so concurrent appends are preserved as well. + storage_session.events = Session.model_validate_json(raw).events + if not event.partial: + storage_session.events.append(event) + else: + storage_session.events = session.events + pipe.multi() + pipe.set( + key, + storage_session.model_dump_json(), + ex=self.config.ttl_seconds if self.config.ttl_seconds > 0 else None, + ) + + await client.transaction(append_to_storage, key) return event diff --git a/tests/unittests/integrations/redis/_fake_redis.py b/tests/unittests/integrations/redis/_fake_redis.py index 68b30befd3..3047c4fb6a 100644 --- a/tests/unittests/integrations/redis/_fake_redis.py +++ b/tests/unittests/integrations/redis/_fake_redis.py @@ -17,7 +17,11 @@ from __future__ import annotations from collections.abc import AsyncIterator +from collections.abc import Awaitable +from collections.abc import Callable import re +from typing import Any +from unittest.mock import Mock class FakeRedisAsync: @@ -29,6 +33,7 @@ def __init__(self) -> None: self._created_at: dict[str, float] = {} self._current_time: float = 0.0 self.scan_patterns: list[str] = [] + self._versions: dict[str, int] = {} def advance_time(self, seconds: float) -> None: self._current_time += seconds @@ -43,6 +48,7 @@ def _is_expired(self, key: str) -> bool: self._store.pop(key, None) self._ex_store.pop(key, None) self._created_at.pop(key, None) + self._versions[key] = self._versions.get(key, 0) + 1 return True return False @@ -63,6 +69,7 @@ async def set( self._store[key] = value self._ex_store[key] = ex self._created_at[key] = self._current_time + self._versions[key] = self._versions.get(key, 0) + 1 return True async def delete(self, key: str) -> int: @@ -70,9 +77,27 @@ async def delete(self, key: str) -> int: self._created_at.pop(key, None) if key in self._store: del self._store[key] + self._versions[key] = self._versions.get(key, 0) + 1 return 1 return 0 + async def transaction( + self, func: Callable[[Any], Awaitable[None]], key: str + ) -> list[bool | None]: + """Retries a queued write if the watched key changed or expired.""" + while True: + self._is_expired(key) + version = self._versions.get(key, 0) + pipe = Mock(get=self.get) + await func(pipe) + self._is_expired(key) + if self._versions.get(key, 0) != version: + continue + return [ + await self.set(*call.args, **call.kwargs) + for call in pipe.set.call_args_list + ] + @staticmethod def _glob_match(pattern: str, key: str) -> bool: """Matches a key the way Redis glob-style patterns do.""" diff --git a/tests/unittests/integrations/redis/test_redis_session_service.py b/tests/unittests/integrations/redis/test_redis_session_service.py index 8197dcd91b..cf8a03cee9 100644 --- a/tests/unittests/integrations/redis/test_redis_session_service.py +++ b/tests/unittests/integrations/redis/test_redis_session_service.py @@ -24,6 +24,7 @@ from google.adk.integrations.redis._config import RedisSessionServiceConfig from google.adk.integrations.redis._redis_session_service import RedisSessionService from google.adk.sessions.base_session_service import GetSessionConfig +from google.genai import types import pytest from ._fake_redis import FakeRedisAsync @@ -176,6 +177,137 @@ async def test_get_session_with_after_timestamp(session_service): assert [e.author for e in fetched.events] == ["user_3", "user_4"] +@pytest.mark.asyncio +@pytest.mark.parametrize( + "config, start", + [ + (None, 0), + (GetSessionConfig(num_recent_events=2), 2), + (GetSessionConfig(num_recent_events=0), 4), + (GetSessionConfig(after_timestamp=102), 2), + (GetSessionConfig(num_recent_events=3, after_timestamp=102), 2), + ], +) +async def test_append_event_preserves_filtered_history( + session_service, config, start +): + """A filtered view can grow without removing stored events.""" + session = await session_service.create_session( + app_name="app1", user_id="u1", state={"original": "value"} + ) + events = [ + Event( + author="user", + invocation_id=f"invocation-{i}", + timestamp=100.0 + i, + content=types.Content(parts=[types.Part(text=f"event-{i}")]), + ) + for i in range(6) + ] + for event in events[:4]: + await session_service.append_event(session, event) + view = await session_service.get_session( + app_name="app1", user_id="u1", session_id=session.id, config=config + ) + assert view.events == events[start:4] + events[4].actions.state_delta = { + "added": 1, + "app:mode": "test", + "user:theme": "dark", + "temp:scratch": "local", + } + + for event in events[4:]: + assert await session_service.append_event(view, event) is event + + stored = await session_service.get_session( + app_name="app1", user_id="u1", session_id=session.id + ) + assert stored.events == events + assert view.events == events[start:] + assert stored.state == { + "original": "value", + "added": 1, + "app:mode": "test", + "user:theme": "dark", + } + assert view.state["temp:scratch"] == "local" + assert "temp:scratch" not in stored.events[4].actions.state_delta + + +@pytest.mark.asyncio +async def test_partial_event_preserves_filtered_history(session_service): + session = await session_service.create_session(app_name="app1", user_id="u1") + old_event = Event(author="user", invocation_id="old") + await session_service.append_event(session, old_event) + view = await session_service.get_session( + app_name="app1", + user_id="u1", + session_id=session.id, + config=GetSessionConfig(num_recent_events=0), + ) + partial = Event(author="agent", invocation_id="partial", partial=True) + + assert await session_service.append_event(view, partial) is partial + + stored = await session_service.get_session( + app_name="app1", user_id="u1", session_id=session.id + ) + assert stored.events == [old_event] + assert view.events == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", ["append", "delete", "expire"]) +async def test_append_event_retries_storage_changes( + session_service, fake_redis, monkeypatch, change +): + """A concurrent append is retained; missing storage keeps recreation behavior.""" + session = await session_service.create_session(app_name="app1", user_id="u1") + old_event = Event(author="user", invocation_id="old") + await session_service.append_event(session, old_event) + view = await session_service.get_session( + app_name="app1", + user_id="u1", + session_id=session.id, + config=GetSessionConfig(num_recent_events=0), + ) + other_event = Event(author="agent", invocation_id="other") + new_event = Event(author="user", invocation_id="new") + key = session_service._session_key("app1", "u1", session.id) + original_get = fake_redis.get + changed = False + + async def get_with_concurrent_change(read_key): + nonlocal changed + raw = await original_get(read_key) + if read_key == key and not changed: + changed = True + if change == "append": + await session_service.append_event(session, other_event) + elif change == "delete": + await session_service.delete_session( + app_name="app1", user_id="u1", session_id=session.id + ) + else: + fake_redis.advance_time(3601) + return raw + + monkeypatch.setattr(fake_redis, "get", get_with_concurrent_change) + await session_service.append_event(view, new_event) + + stored = await session_service.get_session( + app_name="app1", user_id="u1", session_id=session.id + ) + expected = [old_event, other_event] if change == "append" else [] + assert stored.events == expected + [new_event] + assert view.events == [new_event] + fake_redis.advance_time(3599) + assert await original_get(key) is not None + fake_redis.advance_time(2) + assert await original_get(key) is None + + @pytest.mark.asyncio async def test_list_sessions(session_service): await session_service.create_session(