test(sessions): run the shared contract tests against more backends

Co-authored-by: George Weale <gweale@google.com>
PiperOrigin-RevId: 967478023
This commit is contained in:
George Weale
2026-08-19 16:48:30 -07:00
committed by Copybara-Service
parent 4e68bad199
commit b4a9acbcb9
5 changed files with 287 additions and 85 deletions
+15
View File
@@ -0,0 +1,15 @@
# Copyright 2026 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Unit tests for the integrations package."""
@@ -0,0 +1,78 @@
# Copyright 2026 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""In-memory stand-in for the part of redis.asyncio that ADK calls."""
from __future__ import annotations
from collections.abc import AsyncIterator
class FakeRedisAsync:
"""In-memory asynchronous Redis mock for testing."""
def __init__(self) -> None:
self._store: dict[str, str] = {}
self._ex_store: dict[str, int | None] = {}
self._created_at: dict[str, float] = {}
self._current_time: float = 0.0
def advance_time(self, seconds: float) -> None:
self._current_time += seconds
def _is_expired(self, key: str) -> bool:
if key not in self._store:
return True
ttl = self._ex_store.get(key)
if ttl is not None and ttl > 0:
created = self._created_at.get(key, 0.0)
if self._current_time - created >= ttl:
self._store.pop(key, None)
self._ex_store.pop(key, None)
self._created_at.pop(key, None)
return True
return False
async def get(self, key: str) -> str | None:
if self._is_expired(key):
return None
return self._store.get(key)
async def set(
self,
key: str,
value: str,
ex: int | None = None,
nx: bool = False,
) -> bool | None:
if nx and not self._is_expired(key):
return None
self._store[key] = value
self._ex_store[key] = ex
self._created_at[key] = self._current_time
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]
return 1
return 0
async def scan_iter(self, match: str) -> AsyncIterator[str]:
prefix = match.rstrip("*")
for k in list(self._store):
if not self._is_expired(k) and k.startswith(prefix):
yield k
@@ -26,64 +26,7 @@ from google.adk.integrations.redis._redis_session_service import RedisSessionSer
from google.adk.sessions.base_session_service import GetSessionConfig
import pytest
class FakeRedisAsync:
"""In-memory asynchronous Redis mock for testing."""
def __init__(self):
self._store: dict[str, str] = {}
self._ex_store: dict[str, int | None] = {}
self._created_at: dict[str, float] = {}
self._current_time: float = 0.0
def advance_time(self, seconds: float) -> None:
self._current_time += seconds
def _is_expired(self, key: str) -> bool:
if key not in self._store:
return True
ttl = self._ex_store.get(key)
if ttl is not None and ttl > 0:
created = self._created_at.get(key, 0.0)
if self._current_time - created >= ttl:
self._store.pop(key, None)
self._ex_store.pop(key, None)
self._created_at.pop(key, None)
return True
return False
async def get(self, key: str) -> str | None:
if self._is_expired(key):
return None
return self._store.get(key)
async def set(
self,
key: str,
value: str,
ex: int | None = None,
nx: bool = False,
) -> bool | None:
if nx and not self._is_expired(key):
return None
self._store[key] = value
self._ex_store[key] = ex
self._created_at[key] = self._current_time
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]
return 1
return 0
async def scan_iter(self, match: str):
prefix = match.rstrip("*")
for k in list(self._store):
if not self._is_expired(k) and k.startswith(prefix):
yield k
from ._fake_redis import FakeRedisAsync
@pytest.fixture
+173
View File
@@ -0,0 +1,173 @@
# Copyright 2026 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Backend registry behind the shared session service contract tests.
Every test that takes the ``session_service`` fixture states a behavior all
``BaseSessionService`` implementations owe their callers. A backend is only
held to those behaviors once it is registered here, so one left out of this
list can drift from the contract with no test disagreeing.
A backend that fails a contract test has to record it in ``divergences`` with
a written reason. The test is then marked ``xfail(strict=True)``, so the entry
becomes a defect anyone can pick up, and whoever fixes the backend has to
delete the entry in the same change.
The Vertex AI and Firestore backends are still missing from the list. Each
needs a stateful in-memory stand-in for its storage API first: the Firestore
tests drive a call-by-call mock that holds no state, and the Agent Engine fake
keys sessions by id alone rather than by app and user.
"""
from __future__ import annotations
from collections.abc import AsyncIterator
from collections.abc import Callable
from collections.abc import Mapping
import contextlib
import dataclasses
import pathlib
from google.adk.cli.utils.local_storage import PerAgentDatabaseSessionService
from google.adk.features import FeatureName
from google.adk.features import override_feature_enabled
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 BaseSessionService
from google.adk.sessions.database_session_service import DatabaseSessionService
from google.adk.sessions.in_memory_session_service import InMemorySessionService
from google.adk.sessions.sqlite_session_service import SqliteSessionService
import pytest
from ..integrations.redis._fake_redis import FakeRedisAsync
_MakeService = Callable[
[pathlib.Path], contextlib.AbstractAsyncContextManager[BaseSessionService]
]
@dataclasses.dataclass(frozen=True)
class _Backend:
"""A session service implementation held to the shared contract."""
name: str
make: _MakeService
divergences: Mapping[str, str] = dataclasses.field(default_factory=dict)
"""Contract test name -> the written reason this backend fails it today."""
@contextlib.asynccontextmanager
async def _make_in_memory(
tmp_path: pathlib.Path,
) -> AsyncIterator[BaseSessionService]:
del tmp_path
yield InMemorySessionService()
@contextlib.asynccontextmanager
async def _make_in_memory_light_copy(
tmp_path: pathlib.Path,
) -> AsyncIterator[BaseSessionService]:
del tmp_path
override_feature_enabled(
FeatureName.IN_MEMORY_SESSION_SERVICE_LIGHT_COPY, True
)
try:
yield InMemorySessionService()
finally:
override_feature_enabled(
FeatureName.IN_MEMORY_SESSION_SERVICE_LIGHT_COPY, False
)
@contextlib.asynccontextmanager
async def _make_database(
tmp_path: pathlib.Path,
) -> AsyncIterator[BaseSessionService]:
del tmp_path
service = DatabaseSessionService('sqlite+aiosqlite:///:memory:')
try:
yield service
finally:
await service.close()
@contextlib.asynccontextmanager
async def _make_sqlite(
tmp_path: pathlib.Path,
) -> AsyncIterator[BaseSessionService]:
yield SqliteSessionService(str(tmp_path / 'sqlite.db'))
@contextlib.asynccontextmanager
async def _make_redis(
tmp_path: pathlib.Path,
) -> AsyncIterator[BaseSessionService]:
del tmp_path
yield RedisSessionService(
config=RedisSessionServiceConfig(key_prefix='conformance:session:'),
redis_client=FakeRedisAsync(),
)
@contextlib.asynccontextmanager
async def _make_per_agent_database(
tmp_path: pathlib.Path,
) -> AsyncIterator[BaseSessionService]:
service = PerAgentDatabaseSessionService(agents_root=tmp_path)
try:
yield service
finally:
await service.close()
BACKENDS = [
_Backend('in_memory', _make_in_memory),
_Backend('in_memory_light_copy', _make_in_memory_light_copy),
_Backend('database', _make_database),
_Backend('sqlite', _make_sqlite),
# Two more Redis divergences have no contract test to hang an xfail on
# yet: it builds its key scan pattern from a truthiness check on the user
# id, so an empty one lists every user's sessions, and it writes the
# session key unconditionally on append, so appending to a session it has
# never stored creates one instead of raising.
_Backend(
'redis',
_make_redis,
divergences={
'test_list_sessions_ordered_by_last_update_time': (
'Redis sorts sessions newest first, while the base class'
' documents oldest first.'
),
'test_session_last_update_time_updates_on_event': (
'Redis stamps the session with the wall clock instead of the'
" appended event's timestamp."
),
},
),
_Backend('per_agent_database', _make_per_agent_database),
]
@pytest.fixture(params=BACKENDS, ids=lambda backend: backend.name)
async def session_service(
request: pytest.FixtureRequest, tmp_path: pathlib.Path
) -> AsyncIterator[BaseSessionService]:
"""Yields each registered backend in turn, xfailing its known divergences."""
backend: _Backend = request.param
divergence = backend.divergences.get(request.node.originalname)
if divergence is not None:
request.node.add_marker(pytest.mark.xfail(strict=True, reason=divergence))
async with backend.make(tmp_path) as service:
yield service
@@ -17,6 +17,7 @@ from contextlib import asynccontextmanager
from datetime import datetime
from datetime import timezone
import enum
import inspect
import os
import sqlite3
import time
@@ -28,8 +29,6 @@ from google.adk.errors.already_exists_error import AlreadyExistsError
from google.adk.errors.session_not_found_error import SessionNotFoundError
from google.adk.events.event import Event
from google.adk.events.event_actions import EventActions
from google.adk.features import FeatureName
from google.adk.features import override_feature_enabled
from google.adk.sessions import database_session_service
from google.adk.sessions.base_session_service import GetSessionConfig
from google.adk.sessions.database_session_service import DatabaseSessionService
@@ -50,10 +49,14 @@ from sqlalchemy.exc import ArgumentError
from sqlalchemy.ext.asyncio import create_async_engine
from sqlalchemy.pool import StaticPool
# Tests below that take `session_service` run once per backend registered in
# _conformance; each states a behavior every backend owes its callers.
from . import _conformance
from ._conformance import session_service # noqa: F401
class SessionServiceType(enum.Enum):
IN_MEMORY = 'IN_MEMORY'
IN_MEMORY_WITH_LIGHT_COPY_ENABLED = 'IN_MEMORY_WITH_LIGHT_COPY_ENABLED'
DATABASE = 'DATABASE'
SQLITE = 'SQLITE'
@@ -67,33 +70,23 @@ def get_session_service(
return DatabaseSessionService('sqlite+aiosqlite:///:memory:')
if service_type == SessionServiceType.SQLITE:
return SqliteSessionService(str(tmp_path / 'sqlite.db'))
if service_type == SessionServiceType.IN_MEMORY_WITH_LIGHT_COPY_ENABLED:
return InMemorySessionService()
return InMemorySessionService()
@pytest.fixture(
params=[
SessionServiceType.IN_MEMORY,
SessionServiceType.IN_MEMORY_WITH_LIGHT_COPY_ENABLED,
SessionServiceType.DATABASE,
SessionServiceType.SQLITE,
]
)
async def session_service(request, tmp_path):
"""Provides a session service and closes database backends on teardown."""
if request.param == SessionServiceType.IN_MEMORY_WITH_LIGHT_COPY_ENABLED:
override_feature_enabled(
FeatureName.IN_MEMORY_SESSION_SERVICE_LIGHT_COPY, True
)
service = get_session_service(request.param, tmp_path)
yield service
if isinstance(service, DatabaseSessionService):
await service.close()
if request.param == SessionServiceType.IN_MEMORY_WITH_LIGHT_COPY_ENABLED:
override_feature_enabled(
FeatureName.IN_MEMORY_SESSION_SERVICE_LIGHT_COPY, False
)
def test_recorded_divergences_name_a_contract_test():
"""A divergence keyed on anything else silently excuses no backend."""
for backend in _conformance.BACKENDS:
for test_name in backend.divergences:
test_function = globals().get(test_name)
assert test_function is not None, (
f'{backend.name} records a divergence for {test_name}, which is not'
' a test in this module'
)
parameters = inspect.signature(test_function).parameters
assert 'session_service' in parameters, (
f'{backend.name} records a divergence for {test_name}, which does'
' not take the shared contract fixture'
)
def test_database_session_service_enables_pool_pre_ping_by_default():