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
24 changes: 19 additions & 5 deletions src/google/adk/integrations/redis/_redis_session_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
25 changes: 25 additions & 0 deletions tests/unittests/integrations/redis/_fake_redis.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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
Expand All @@ -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

Expand All @@ -63,16 +69,35 @@ 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:
self._ex_store.pop(key, None)
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."""
Expand Down
132 changes: 132 additions & 0 deletions tests/unittests/integrations/redis/test_redis_session_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down