From be8cd0ce1cc770e23dceadc3fdff445069769a53 Mon Sep 17 00:00:00 2001 From: Venkatesh Shanbhag <91714892+theshanbhag@users.noreply.github.com> Date: Fri, 25 Sep 2026 14:15:19 +0530 Subject: [PATCH] feat(integrations): add MongoDB-backed session and memory services Add MongoDbSessionService, which persists sessions, events, and app:/user:-scoped state in MongoDB with optimistic-concurrency revisions, and MongoDbMemoryService, which ingests session events into a keyword-indexed memory collection for long-term recall. Register the MONGODB_* experimental features (enabled by default) and export both services from google.adk.integrations.mongodb. Also add a README documenting the MongoDB integration (search toolset, both services, and MongoDbToolSettings) and add mongomock to the test dependencies for the new unit tests. NO_UNIT_GUIDE=The integration README added in this commit documents all MongoDB units; per-unit guides will follow when the experimental API stabilizes. --- constraints-3.10.txt | 17 +- constraints-3.11.txt | 17 +- constraints-3.12.txt | 17 +- constraints-3.13.txt | 17 +- constraints-3.14.txt | 17 +- pyproject.toml | 1 + src/google/adk/features/_feature_registry.py | 12 +- src/google/adk/integrations/mongodb/README.md | 179 ++++++ .../adk/integrations/mongodb/__init__.py | 9 +- .../integrations/mongodb/_memory_service.py | 255 +++++++++ .../integrations/mongodb/_session_service.py | 522 ++++++++++++++++++ .../mongodb/test_memory_service.py | 148 +++++ .../mongodb/test_session_service.py | 281 ++++++++++ 13 files changed, 1459 insertions(+), 33 deletions(-) create mode 100644 src/google/adk/integrations/mongodb/README.md create mode 100644 src/google/adk/integrations/mongodb/_memory_service.py create mode 100644 src/google/adk/integrations/mongodb/_session_service.py create mode 100644 tests/unittests/integrations/mongodb/test_memory_service.py create mode 100644 tests/unittests/integrations/mongodb/test_session_service.py diff --git a/constraints-3.10.txt b/constraints-3.10.txt index 9c8a04b9b7b..fa3be57201a 100644 --- a/constraints-3.10.txt +++ b/constraints-3.10.txt @@ -1,5 +1,5 @@ # This file was autogenerated by uv via the following command: -# uv pip compile pyproject.toml --all-extras --python-version 3.10 --no-emit-package google-adk --exclude-newer 2026-09-04 --index-url https://pypi.org/simple -o constraints-3.10.txt +# uv pip compile pyproject.toml --all-extras --python-version 3.10 --no-emit-package google-adk --exclude-newer 2026-09-21 --index-url https://pypi.org/simple -o constraints-3.10.txt a2a-sdk==1.1.2 # via # -c constraints-3.10.txt.stable.tmp @@ -31,7 +31,6 @@ aiohttp==3.14.1 # daytona-analytics-api-client-async # daytona-api-client-async # daytona-toolbox-api-client-async - # google-adk # google-cloud-aiplatform # kubernetes # langchain-community @@ -989,6 +988,10 @@ mmh3==5.2.1 # via # -c constraints-3.10.txt.stable.tmp # google-cloud-spanner +mongomock==4.3.0 + # via + # -c constraints-3.10.txt.stable.tmp + # google-adk (pyproject.toml) multidict==6.7.1 # via # -c constraints-3.10.txt.stable.tmp @@ -1199,6 +1202,7 @@ packaging==26.2 # langchain-core # langsmith # marshmallow + # mongomock # opentelemetry-instrumentation # pyink # pyproject-api @@ -1505,6 +1509,7 @@ pytokens==0.4.1 pytz==2026.2 # via # -c constraints-3.10.txt.stable.tmp + # mongomock # oci # pandas pyyaml==6.0.3 @@ -1601,6 +1606,10 @@ scipy==1.15.3 # via # -c constraints-3.10.txt.stable.tmp # scikit-learn +sentinels==1.1.1 + # via + # -c constraints-3.10.txt.stable.tmp + # mongomock setuptools==83.0.0 # via # -c constraints-3.10.txt.stable.tmp @@ -1907,10 +1916,6 @@ tzdata==2026.3 # via # -c constraints-3.10.txt.stable.tmp # pandas -tzlocal==5.4.4 - # via - # -c constraints-3.10.txt.stable.tmp - # google-adk uritemplate==4.2.0 # via # -c constraints-3.10.txt.stable.tmp diff --git a/constraints-3.11.txt b/constraints-3.11.txt index 5a44ea8968c..c56e018a891 100644 --- a/constraints-3.11.txt +++ b/constraints-3.11.txt @@ -1,5 +1,5 @@ # This file was autogenerated by uv via the following command: -# uv pip compile pyproject.toml --all-extras --python-version 3.11 --no-emit-package google-adk --exclude-newer 2026-09-04 --index-url https://pypi.org/simple -o constraints-3.11.txt +# uv pip compile pyproject.toml --all-extras --python-version 3.11 --no-emit-package google-adk --exclude-newer 2026-09-21 --index-url https://pypi.org/simple -o constraints-3.11.txt a2a-sdk==1.1.1 # via # -c constraints-3.11.txt.stable.tmp @@ -32,7 +32,6 @@ aiohttp==3.14.1 # daytona-analytics-api-client-async # daytona-api-client-async # daytona-toolbox-api-client-async - # google-adk # google-cloud-aiplatform # instructor # kubernetes @@ -1113,6 +1112,10 @@ mmh3==5.2.1 # -c constraints-3.11.txt.stable.tmp # chromadb # google-cloud-spanner +mongomock==4.3.0 + # via + # -c constraints-3.11.txt.stable.tmp + # google-adk (pyproject.toml) multidict==6.7.1 # via # -c constraints-3.11.txt.stable.tmp @@ -1364,6 +1367,7 @@ packaging==26.2 # langchain-core # langsmith # marshmallow + # mongomock # onnxruntime # opentelemetry-instrumentation # pyink @@ -1740,6 +1744,7 @@ pytube==15.0.0 pytz==2026.2 # via # -c constraints-3.11.txt.stable.tmp + # mongomock # oci # pandas pyyaml==6.0.3 @@ -1859,6 +1864,10 @@ scipy==1.17.1 # via # -c constraints-3.11.txt.stable.tmp # scikit-learn +sentinels==1.1.1 + # via + # -c constraints-3.11.txt.stable.tmp + # mongomock setuptools==83.0.0 # via # -c constraints-3.11.txt.stable.tmp @@ -2171,10 +2180,6 @@ tzdata==2026.3 # -c constraints-3.11.txt.stable.tmp # pandas # pendulum -tzlocal==5.4.4 - # via - # -c constraints-3.11.txt.stable.tmp - # google-adk uc-micro-py==2.0.0 # via # -c constraints-3.11.txt.stable.tmp diff --git a/constraints-3.12.txt b/constraints-3.12.txt index 0adc156e4dd..8d2ed788e51 100644 --- a/constraints-3.12.txt +++ b/constraints-3.12.txt @@ -1,5 +1,5 @@ # This file was autogenerated by uv via the following command: -# uv pip compile pyproject.toml --all-extras --python-version 3.12 --no-emit-package google-adk --exclude-newer 2026-09-04 --index-url https://pypi.org/simple -o constraints-3.12.txt +# uv pip compile pyproject.toml --all-extras --python-version 3.12 --no-emit-package google-adk --exclude-newer 2026-09-21 --index-url https://pypi.org/simple -o constraints-3.12.txt a2a-sdk==1.1.1 # via # -c constraints-3.12.txt.stable.tmp @@ -31,7 +31,6 @@ aiohttp==3.14.1 # daytona-analytics-api-client-async # daytona-api-client-async # daytona-toolbox-api-client-async - # google-adk # google-cloud-aiplatform # kubernetes # langchain-community @@ -973,6 +972,10 @@ mmh3==5.2.1 # via # -c constraints-3.12.txt.stable.tmp # google-cloud-spanner +mongomock==4.3.0 + # via + # -c constraints-3.12.txt.stable.tmp + # google-adk (pyproject.toml) multidict==6.7.1 # via # -c constraints-3.12.txt.stable.tmp @@ -1187,6 +1190,7 @@ packaging==26.2 # langchain-core # langsmith # marshmallow + # mongomock # opentelemetry-instrumentation # pyink # pyproject-api @@ -1493,6 +1497,7 @@ pytokens==0.4.1 pytz==2026.2 # via # -c constraints-3.12.txt.stable.tmp + # mongomock # oci # pandas pyyaml==6.0.3 @@ -1597,6 +1602,10 @@ scipy==1.18.0 # via # -c constraints-3.12.txt.stable.tmp # scikit-learn +sentinels==1.1.1 + # via + # -c constraints-3.12.txt.stable.tmp + # mongomock setuptools==83.0.0 # via # -c constraints-3.12.txt.stable.tmp @@ -1874,10 +1883,6 @@ tzdata==2026.3 # via # -c constraints-3.12.txt.stable.tmp # pandas -tzlocal==5.4.4 - # via - # -c constraints-3.12.txt.stable.tmp - # google-adk uritemplate==4.2.0 # via # -c constraints-3.12.txt.stable.tmp diff --git a/constraints-3.13.txt b/constraints-3.13.txt index 6e06e27df12..0765fd018e7 100644 --- a/constraints-3.13.txt +++ b/constraints-3.13.txt @@ -1,5 +1,5 @@ # This file was autogenerated by uv via the following command: -# uv pip compile pyproject.toml --all-extras --python-version 3.13 --no-emit-package google-adk --exclude-newer 2026-09-04 --index-url https://pypi.org/simple -o constraints-3.13.txt +# uv pip compile pyproject.toml --all-extras --python-version 3.13 --no-emit-package google-adk --exclude-newer 2026-09-21 --index-url https://pypi.org/simple -o constraints-3.13.txt a2a-sdk==1.1.1 # via # -c constraints-3.13.txt.stable.tmp @@ -31,7 +31,6 @@ aiohttp==3.14.1 # daytona-analytics-api-client-async # daytona-api-client-async # daytona-toolbox-api-client-async - # google-adk # google-cloud-aiplatform # kubernetes # langchain-community @@ -965,6 +964,10 @@ mmh3==5.2.1 # via # -c constraints-3.13.txt.stable.tmp # google-cloud-spanner +mongomock==4.3.0 + # via + # -c constraints-3.13.txt.stable.tmp + # google-adk (pyproject.toml) multidict==6.7.1 # via # -c constraints-3.13.txt.stable.tmp @@ -1179,6 +1182,7 @@ packaging==26.2 # langchain-core # langsmith # marshmallow + # mongomock # opentelemetry-instrumentation # pyink # pyproject-api @@ -1485,6 +1489,7 @@ pytokens==0.4.1 pytz==2026.2 # via # -c constraints-3.13.txt.stable.tmp + # mongomock # oci # pandas pyyaml==6.0.3 @@ -1589,6 +1594,10 @@ scipy==1.18.0 # via # -c constraints-3.13.txt.stable.tmp # scikit-learn +sentinels==1.1.1 + # via + # -c constraints-3.13.txt.stable.tmp + # mongomock setuptools==83.0.0 # via # -c constraints-3.13.txt.stable.tmp @@ -1855,10 +1864,6 @@ tzdata==2026.3 # via # -c constraints-3.13.txt.stable.tmp # pandas -tzlocal==5.4.4 - # via - # -c constraints-3.13.txt.stable.tmp - # google-adk uritemplate==4.2.0 # via # -c constraints-3.13.txt.stable.tmp diff --git a/constraints-3.14.txt b/constraints-3.14.txt index f07aa331263..08c82b6e6ba 100644 --- a/constraints-3.14.txt +++ b/constraints-3.14.txt @@ -1,5 +1,5 @@ # This file was autogenerated by uv via the following command: -# uv pip compile pyproject.toml --all-extras --python-version 3.14 --no-emit-package google-adk --exclude-newer 2026-09-04 --index-url https://pypi.org/simple -o constraints-3.14.txt +# uv pip compile pyproject.toml --all-extras --python-version 3.14 --no-emit-package google-adk --exclude-newer 2026-09-21 --index-url https://pypi.org/simple -o constraints-3.14.txt a2a-sdk==1.1.1 # via # -c constraints-3.14.txt.stable.tmp @@ -31,7 +31,6 @@ aiohttp==3.14.1 # daytona-analytics-api-client-async # daytona-api-client-async # daytona-toolbox-api-client-async - # google-adk # google-cloud-aiplatform # kubernetes # langchain-community @@ -965,6 +964,10 @@ mmh3==5.2.1 # via # -c constraints-3.14.txt.stable.tmp # google-cloud-spanner +mongomock==4.3.0 + # via + # -c constraints-3.14.txt.stable.tmp + # google-adk (pyproject.toml) multidict==6.7.1 # via # -c constraints-3.14.txt.stable.tmp @@ -1179,6 +1182,7 @@ packaging==26.2 # langchain-core # langsmith # marshmallow + # mongomock # opentelemetry-instrumentation # pyink # pyproject-api @@ -1485,6 +1489,7 @@ pytokens==0.4.1 pytz==2026.2 # via # -c constraints-3.14.txt.stable.tmp + # mongomock # oci # pandas pyyaml==6.0.3 @@ -1589,6 +1594,10 @@ scipy==1.18.0 # via # -c constraints-3.14.txt.stable.tmp # scikit-learn +sentinels==1.1.1 + # via + # -c constraints-3.14.txt.stable.tmp + # mongomock setuptools==83.0.0 # via # -c constraints-3.14.txt.stable.tmp @@ -1855,10 +1864,6 @@ tzdata==2026.3 # via # -c constraints-3.14.txt.stable.tmp # pandas -tzlocal==5.4.4 - # via - # -c constraints-3.14.txt.stable.tmp - # google-adk uritemplate==4.2.0 # via # -c constraints-3.14.txt.stable.tmp diff --git a/pyproject.toml b/pyproject.toml index ef23e4dc380..bcfe9870196 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -327,6 +327,7 @@ optional-dependencies.test = [ "llama-index-readers-file>=0.4", "lxml>=5.3", "mcp>=1.24,<3", + "mongomock>=4.3,<5", # In-memory MongoDB used by the mongodb integration tests. "nltk!=3.10.1", # Transitive via rouge-score and llama-index-core; 3.10.1's import hook breaks any venv living inside the working directory (reverted upstream in nltk/nltk#3732). "openai>=2.20,<3", "openpyxl>=3.1.5,<4", diff --git a/src/google/adk/features/_feature_registry.py b/src/google/adk/features/_feature_registry.py index 87c9440f35b..a166227fecb 100644 --- a/src/google/adk/features/_feature_registry.py +++ b/src/google/adk/features/_feature_registry.py @@ -59,8 +59,10 @@ class FeatureName(str, Enum): # enum member by name. Keeping it private avoids a backward-compat # obligation for what is intended as a temporary, internal kill-switch. _MCP_GRACEFUL_ERROR_HANDLING = "MCP_GRACEFUL_ERROR_HANDLING" - MONGODB_TOOLSET = "MONGODB_TOOLSET" + MONGODB_MEMORY_SERVICE = "MONGODB_MEMORY_SERVICE" + MONGODB_SESSION_SERVICE = "MONGODB_SESSION_SERVICE" MONGODB_TOOL_SETTINGS = "MONGODB_TOOL_SETTINGS" + MONGODB_TOOLSET = "MONGODB_TOOLSET" PROGRESSIVE_SSE_STREAMING = "PROGRESSIVE_SSE_STREAMING" PUBSUB_TOOL_CONFIG = "PUBSUB_TOOL_CONFIG" PUBSUB_TOOLSET = "PUBSUB_TOOLSET" @@ -191,12 +193,18 @@ class FeatureConfig: FeatureName._MCP_GRACEFUL_ERROR_HANDLING: FeatureConfig( FeatureStage.EXPERIMENTAL, default_on=True ), - FeatureName.MONGODB_TOOLSET: FeatureConfig( + FeatureName.MONGODB_MEMORY_SERVICE: FeatureConfig( + FeatureStage.EXPERIMENTAL, default_on=True + ), + FeatureName.MONGODB_SESSION_SERVICE: FeatureConfig( FeatureStage.EXPERIMENTAL, default_on=True ), FeatureName.MONGODB_TOOL_SETTINGS: FeatureConfig( FeatureStage.EXPERIMENTAL, default_on=True ), + FeatureName.MONGODB_TOOLSET: FeatureConfig( + FeatureStage.EXPERIMENTAL, default_on=True + ), FeatureName.PROGRESSIVE_SSE_STREAMING: FeatureConfig( FeatureStage.EXPERIMENTAL, default_on=True ), diff --git a/src/google/adk/integrations/mongodb/README.md b/src/google/adk/integrations/mongodb/README.md new file mode 100644 index 00000000000..4ad95f6f9ee --- /dev/null +++ b/src/google/adk/integrations/mongodb/README.md @@ -0,0 +1,179 @@ +# MongoDB Integration for ADK + +This integration connects the Google Agent Development Kit (ADK) to MongoDB. +It provides search tools agents can call, plus MongoDB-backed session and +memory services. It works with MongoDB Atlas and self-managed MongoDB 8.0+ +deployments. + +> **Experimental:** All classes in this integration are experimental and their +> APIs may change in future releases. + +## Features + +- **Vector Search Tool (`mongodb_vector_search`):** Runs Atlas Vector Search + (`$vectorSearch`) queries against a collection, with optional pre-filtering, + configurable limits, and result projection. +- **Hybrid Search Tool (`mongodb_hybrid_search`):** Combines full-text search + and vector search with reciprocal rank fusion (`$rankFusion`), with tunable + vector/text weights. +- **Session Persistence (`MongoDbSessionService`):** Stores sessions, events, + and `app:` / `user:` / session-scoped state in MongoDB, with optimistic + concurrency control for concurrent writers. +- **Long-term Memory (`MongoDbMemoryService`):** Ingests session events into a + memory collection and recalls them by keyword matching. +- **Flexible Connection Options:** Pass a connection string and let the + integration create/own the PyMongo client, or bring your own pre-configured + `pymongo.MongoClient`. + +## Installation / Dependencies + +Install ADK with the `mongodb` extra, which pulls in `pymongo`: + +```bash +pip install google-adk[mongodb] +``` + +or if ADK is already installed: + +```bash +pip install "pymongo>=4.9,<5" +``` + +## Requirements + +| Component | Requirement | +| :--- | :--- | +| `mongodb_vector_search` | A **vector search index** on the collection (MongoDB Atlas or MongoDB 8.0+). | +| `mongodb_hybrid_search` | A **vector search index** and a **full-text search index** on the collection, and a deployment that supports `$rankFusion` (MongoDB Atlas or MongoDB 8.0+). | +| `MongoDbSessionService` | Any MongoDB deployment reachable by PyMongo (no search indexes needed). | +| `MongoDbMemoryService` | Any MongoDB deployment reachable by PyMongo (no search indexes needed). | + +## Quick Start: Search Tools + +`MongoDbToolset` exposes `mongodb_vector_search` and `mongodb_hybrid_search` +to the agent. The client, database name, and settings are bound on the +toolset and hidden from the model. + +```python +from google.adk.agents import Agent +from google.adk.integrations.mongodb import MongoDbToolset + +toolset = MongoDbToolset( + connection_string="mongodb+srv://user:pass@cluster.mongodb.net/", + database_name="products_db", +) + +agent = Agent( + model="gemini-2.5-flash", + name="product_search_agent", + instruction=( + "Search the products collection with your MongoDB tools. " + "Embed the user's query yourself before calling the search tools." + ), + tools=[toolset], +) +``` + +The embedding vector for `query_embedding` must be produced by your own +embedding model — typically exposed to the agent as an additional tool — and +must match the dimensions and content of the vectors stored in the +collection. + +To bring your own client instead of a connection string: + +```python +from pymongo import MongoClient +from google.adk.integrations.mongodb import MongoDbToolset + +client = MongoClient("mongodb+srv://user:pass@cluster.mongodb.net/") +toolset = MongoDbToolset(database_name="products_db", mongo_client=client) +``` + +## Quick Start: Session Service + +```python +from google.adk.integrations.mongodb import MongoDbSessionService +from google.adk.runners import Runner + +session_service = MongoDbSessionService( + connection_string="mongodb+srv://user:pass@cluster.mongodb.net/", + database_name="my_app", +) + +runner = Runner( + agent=agent, + app_name="my_app", + session_service=session_service, +) +``` + +### Document Layout + +`MongoDbSessionService` stores data in four collections (names configurable): + +| Collection | Document key | Contents | +| :--- | :--- | :--- | +| `sessions` | `//` | Session-scoped state (JSON-encoded), timestamps, and an optimistic-concurrency `revision`. | +| `events` | `///` | Full serialized event under `event_data`. | +| `app_states` | `` | App-scoped state (`app:` prefixed keys). | +| `user_states` | `/` | User-scoped state (`user:` prefixed keys). | + +State buckets are stored JSON-encoded so state keys containing characters +MongoDB forbids in document fields (e.g. `.`, `$`) round-trip safely. Event +appends use per-session locking plus a revision check, and raise +`StaleSessionError` if the session was modified in storage since it was +loaded. + +## Quick Start: Memory Service + +```python +from google.adk.integrations.mongodb import MongoDbMemoryService +from google.adk.runners import Runner + +memory_service = MongoDbMemoryService( + connection_string="mongodb+srv://user:pass@cluster.mongodb.net/", + database_name="my_app", +) + +runner = Runner( + agent=agent, + app_name="my_app", + memory_service=memory_service, +) +``` + +Events ingested via `add_session_to_memory` are stored as memory documents +(one per event) keyed by `///`, each +carrying the lowercase keywords extracted from its text content (a compact +English stop-word list is ignored; pass `stop_words` to customize). +`search_memory` returns entries whose keyword array intersects the query's +keywords. Ingestion is idempotent — re-adding a session overwrites its +memories rather than duplicating them. + +## Configuration: `MongoDbToolSettings` + +`MongoDbToolSettings` customizes the defaults used by the search tools: + +| Field | Type | Default | Description | +| :--- | :--- | :--- | :--- | +| `default_vector_index_name` | `str` | `"vector_index"` | Name of the vector search index to query. | +| `default_search_index_name` | `str` | `"default"` | Name of the full-text search index used by hybrid search. | +| `default_embedding_field` | `str` | `"embedding"` | Document field that stores embedding vectors. | +| `default_limit` | `int` | `4` | Default number of documents returned by a search operation. | +| `max_results` | `int` | `50` | Maximum number of documents a search operation may return. | +| `default_num_candidates` | `int` | `100` | Default number of nearest neighbors considered by vector search. | + +```python +from google.adk.integrations.mongodb import MongoDbToolset +from google.adk.integrations.mongodb import MongoDbToolSettings + +toolset = MongoDbToolset( + connection_string="mongodb+srv://user:pass@cluster.mongodb.net/", + database_name="products_db", + settings=MongoDbToolSettings(default_limit=8, max_results=20), +) +``` + +Per-call `limit`, `num_candidates`, `index_name`, `embedding_field`, and +`output_fields` arguments passed by the model override these defaults +(`limit` is still capped by `max_results`). diff --git a/src/google/adk/integrations/mongodb/__init__.py b/src/google/adk/integrations/mongodb/__init__.py index da41da744d2..cc314a6e5e2 100644 --- a/src/google/adk/integrations/mongodb/__init__.py +++ b/src/google/adk/integrations/mongodb/__init__.py @@ -15,7 +15,8 @@ """MongoDB Integration (Experimental). MongoDB tools for vector search and hybrid search against collections on -MongoDB Atlas or MongoDB 8.0+ deployments. +MongoDB Atlas or MongoDB 8.0+ deployments, plus MongoDB-backed session and +memory services. """ from __future__ import annotations @@ -23,12 +24,16 @@ import typing if typing.TYPE_CHECKING: + from ._memory_service import MongoDbMemoryService from ._mongodb_toolset import MongoDbToolset + from ._session_service import MongoDbSessionService from ._settings import MongoDbToolSettings # Map attribute names to relative module paths _lazy_imports = { + "MongoDbMemoryService": "._memory_service", "MongoDbToolset": "._mongodb_toolset", + "MongoDbSessionService": "._session_service", "MongoDbToolSettings": "._settings", } @@ -49,6 +54,8 @@ def __dir__() -> list[str]: __all__ = [ + "MongoDbMemoryService", + "MongoDbSessionService", "MongoDbToolset", "MongoDbToolSettings", ] diff --git a/src/google/adk/integrations/mongodb/_memory_service.py b/src/google/adk/integrations/mongodb/_memory_service.py new file mode 100644 index 00000000000..7fab3126ff5 --- /dev/null +++ b/src/google/adk/integrations/mongodb/_memory_service.py @@ -0,0 +1,255 @@ +# 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. + +from __future__ import annotations + +import asyncio +import logging +import re +from typing import Any +from typing import TYPE_CHECKING + +from typing_extensions import override + +from . import _client +from ...features import experimental +from ...features import FeatureName +from ...memory import _utils +from ...memory.base_memory_service import BaseMemoryService +from ...memory.base_memory_service import SearchMemoryResponse +from ...memory.memory_entry import MemoryEntry + +if TYPE_CHECKING: + from pymongo import MongoClient + + from ...sessions.session import Session + +logger = logging.getLogger("google_adk." + __name__) + +DEFAULT_MEMORIES_COLLECTION = "memories" + +# Compact English stop-word list ignored when extracting keywords. Kept +# intentionally short; callers can pass their own set via `stop_words`. +DEFAULT_STOP_WORDS = { + "a", + "an", + "and", + "are", + "as", + "at", + "be", + "but", + "by", + "for", + "from", + "has", + "have", + "how", + "i", + "if", + "in", + "into", + "is", + "it", + "its", + "me", + "my", + "of", + "on", + "or", + "our", + "so", + "that", + "the", + "their", + "this", + "to", + "was", + "we", + "what", + "when", + "where", + "which", + "who", + "will", + "with", + "you", + "your", +} + + +@experimental(FeatureName.MONGODB_MEMORY_SERVICE) +class MongoDbMemoryService(BaseMemoryService): + """Memory service that uses MongoDB as the backend. + + Events ingested from sessions are stored as memory documents (one per + event) in a single collection, each carrying the lowercase keywords + extracted from its text content. `search_memory` matches documents whose + keyword array intersects the query's keywords. + + Example: + ```python + memory_service = MongoDbMemoryService( + connection_string="mongodb+srv://user:pass@cluster.mongodb.net/", + database_name="my_app", + ) + runner = Runner( + agent=agent, app_name="my_app", memory_service=memory_service + ) + ``` + """ + + def __init__( + self, + *, + database_name: str, + connection_string: str | None = None, + mongo_client: MongoClient | None = None, + memories_collection: str = DEFAULT_MEMORIES_COLLECTION, + stop_words: set[str] | None = None, + ): + """Initializes the MongoDB memory service. + + Args: + database_name: The MongoDB database used to store memories. + connection_string: The MongoDB connection string (URI) used to create a + client owned by this service. Requires the `pymongo` package + (`pip install google-adk[mongodb]`). + mongo_client: An existing PyMongo client to use instead of creating one + from `connection_string`. The caller keeps ownership of the client. + memories_collection: Collection name for memory documents. + stop_words: Words to ignore when extracting keywords. Defaults to a + standard English stop-word list. + """ + if mongo_client is not None and connection_string is not None: + raise ValueError( + "Only one of `connection_string` and `mongo_client` may be provided." + ) + if mongo_client is not None: + self._client = mongo_client + self._owns_client = False + elif connection_string is not None: + self._client = _client.get_mongo_client(connection_string) + self._owns_client = True + else: + raise ValueError( + "Either `connection_string` or `mongo_client` must be provided." + ) + self._database_name = database_name + self.memories_collection = memories_collection + self.stop_words = ( + stop_words if stop_words is not None else DEFAULT_STOP_WORDS + ) + + def _memories(self): + return self._client[self._database_name][self.memories_collection] + + def _extract_keywords(self, text: str) -> set[str]: + """Extracts lowercase keywords from text, ignoring stop words.""" + words = re.findall(r"[a-z0-9]+", text.lower()) + return {word for word in words if word not in self.stop_words} + + @override + async def add_session_to_memory(self, session: Session) -> None: + """Ingests the session's text events into the memory collection. + + Ingestion is idempotent: memory documents are keyed by + `///`, so re-adding a session + overwrites its memories rather than duplicating them. + """ + + def _add() -> None: + for event in session.events: + if not event.content or not event.content.parts: + continue + text = " ".join( + [part.text for part in event.content.parts if part.text] + ) + if not text: + continue + keywords = self._extract_keywords(text) + if not keywords: + continue + memory_id = ( + f"{session.app_name}/{session.user_id}/{session.id}/{event.id}" + ) + self._memories().replace_one( + {"_id": memory_id}, + { + "app_name": session.app_name, + "user_id": session.user_id, + "session_id": session.id, + "author": event.author, + "keywords": sorted(keywords), + "content": event.content.model_dump( + exclude_none=True, mode="json" + ), + "timestamp": event.timestamp, + }, + upsert=True, + ) + + await asyncio.to_thread(_add) + + @override + async def search_memory( + self, *, app_name: str, user_id: str, query: str + ) -> SearchMemoryResponse: + """Searches memory for events matching the query's keywords.""" + keywords = self._extract_keywords(query) + if not keywords: + return SearchMemoryResponse() + + def _search() -> list[dict[str, Any]]: + cursor = self._memories().find({ + "app_name": app_name, + "user_id": user_id, + "keywords": {"$in": sorted(keywords)}, + }) + return list(cursor) + + docs = await asyncio.to_thread(_search) + + seen = set() + memories = [] + for doc in docs: + try: + from google.genai import types + + content = types.Content.model_validate(doc["content"]) + entry = MemoryEntry( + id=doc.get("_id"), + content=content, + author=doc.get("author"), + timestamp=_utils.format_timestamp(doc.get("timestamp", 0.0)), + ) + except Exception as exc: + logger.warning(f"Failed to parse memory entry: {exc}") + continue + content_text = ( + " ".join([part.text for part in content.parts if part.text]) + if content.parts + else "" + ) + key = (entry.author, content_text, entry.timestamp) + if key not in seen: + seen.add(key) + memories.append(entry) + + return SearchMemoryResponse(memories=memories) + + async def close(self) -> None: + """Closes the MongoDB client if it was created by this service.""" + if self._owns_client: + await asyncio.to_thread(self._client.close) diff --git a/src/google/adk/integrations/mongodb/_session_service.py b/src/google/adk/integrations/mongodb/_session_service.py new file mode 100644 index 00000000000..b6f63b95867 --- /dev/null +++ b/src/google/adk/integrations/mongodb/_session_service.py @@ -0,0 +1,522 @@ +# 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. + +from __future__ import annotations + +import asyncio +from contextlib import asynccontextmanager +import copy +import json +import logging +import time +from typing import Any +from typing import AsyncGenerator +from typing import TYPE_CHECKING + +from typing_extensions import override + +from . import _client +from ...errors._stale_session_error import StaleSessionError +from ...errors.already_exists_error import AlreadyExistsError +from ...errors.session_not_found_error import SessionNotFoundError +from ...events.event import Event +from ...features import experimental +from ...features import FeatureName +from ...platform import uuid as platform_uuid +from ...sessions import _session_util +from ...sessions.base_session_service import BaseSessionService +from ...sessions.base_session_service import GetSessionConfig +from ...sessions.base_session_service import ListSessionsResponse +from ...sessions.session import Session +from ...sessions.state import State + +if TYPE_CHECKING: + from pymongo import MongoClient + +logger = logging.getLogger("google_adk." + __name__) + +_STALE_SESSION_ERROR_MESSAGE = ( + "The session has been modified in storage since it was loaded. " + "Please reload the session before appending more events." +) + +DEFAULT_SESSIONS_COLLECTION = "sessions" +DEFAULT_EVENTS_COLLECTION = "events" +DEFAULT_APP_STATE_COLLECTION = "app_states" +DEFAULT_USER_STATE_COLLECTION = "user_states" + +_SessionLockKey = tuple[str, str, str] + + +def _dumps_state(state: dict[str, Any]) -> str: + """Serializes a state bucket, coercing non-JSON values first.""" + return json.dumps(_session_util.make_json_safe_state(state)) + + +def _loads_state(raw: Any) -> dict[str, Any]: + """Parses a stored state bucket (JSON string or legacy dict).""" + if isinstance(raw, str): + return json.loads(raw) + return dict(raw or {}) + + +@experimental(FeatureName.MONGODB_SESSION_SERVICE) +class MongoDbSessionService(BaseSessionService): + """Session service that uses MongoDB as the backend. + + Document layout within the configured database: + + - ``: one document per session, keyed by + `//`, holding the session-scoped state + (JSON-encoded), timestamps, and an optimistic-concurrency `revision`. + - ``: one document per event, keyed by + `///`, holding the full + serialized event under `event_data`. + - ``: one document per app, keyed by ``. + - ``: one document per user, keyed by + `/`. + + State buckets are stored JSON-encoded so state keys containing characters + that MongoDB forbids in document fields (e.g. `.`, `$`) round-trip safely. + + Example: + ```python + session_service = MongoDbSessionService( + connection_string="mongodb+srv://user:pass@cluster.mongodb.net/", + database_name="my_app", + ) + runner = Runner( + agent=agent, app_name="my_app", session_service=session_service + ) + ``` + """ + + def __init__( + self, + *, + database_name: str, + connection_string: str | None = None, + mongo_client: MongoClient | None = None, + sessions_collection: str = DEFAULT_SESSIONS_COLLECTION, + events_collection: str = DEFAULT_EVENTS_COLLECTION, + app_state_collection: str = DEFAULT_APP_STATE_COLLECTION, + user_state_collection: str = DEFAULT_USER_STATE_COLLECTION, + ): + """Initializes the MongoDB session service. + + Args: + database_name: The MongoDB database used to store sessions, events and + shared state. + connection_string: The MongoDB connection string (URI) used to create a + client owned by this service. Requires the `pymongo` package + (`pip install google-adk[mongodb]`). + mongo_client: An existing PyMongo client to use instead of creating one + from `connection_string`. The caller keeps ownership of the client. + sessions_collection: Collection name for session documents. + events_collection: Collection name for event documents. + app_state_collection: Collection name for app state documents. + user_state_collection: Collection name for user state documents. + """ + if mongo_client is not None and connection_string is not None: + raise ValueError( + "Only one of `connection_string` and `mongo_client` may be provided." + ) + if mongo_client is not None: + self._client = mongo_client + self._owns_client = False + elif connection_string is not None: + self._client = _client.get_mongo_client(connection_string) + self._owns_client = True + else: + raise ValueError( + "Either `connection_string` or `mongo_client` must be provided." + ) + self._database_name = database_name + self.sessions_collection = sessions_collection + self.events_collection = events_collection + self.app_state_collection = app_state_collection + self.user_state_collection = user_state_collection + + # Per-session locks used to serialize append_event calls in this process. + self._session_locks: dict[_SessionLockKey, asyncio.Lock] = {} + self._session_lock_ref_count: dict[_SessionLockKey, int] = {} + self._session_locks_guard = asyncio.Lock() + + def _sessions(self): + return self._client[self._database_name][self.sessions_collection] + + def _events(self): + return self._client[self._database_name][self.events_collection] + + def _app_states(self): + return self._client[self._database_name][self.app_state_collection] + + def _user_states(self): + return self._client[self._database_name][self.user_state_collection] + + @staticmethod + def _session_key(app_name: str, user_id: str, session_id: str) -> str: + return f"{app_name}/{user_id}/{session_id}" + + @staticmethod + def _user_key(app_name: str, user_id: str) -> str: + return f"{app_name}/{user_id}" + + @asynccontextmanager + async def _with_session_lock( + self, *, app_name: str, user_id: str, session_id: str + ) -> AsyncGenerator[None]: + """Serializes event appends for the same session within this process.""" + lock_key = (app_name, user_id, session_id) + async with self._session_locks_guard: + lock = self._session_locks.get(lock_key) + if lock is None: + lock = asyncio.Lock() + self._session_locks[lock_key] = lock + self._session_lock_ref_count[lock_key] = ( + self._session_lock_ref_count.get(lock_key, 0) + 1 + ) + + try: + async with lock: + yield + finally: + async with self._session_locks_guard: + remaining = self._session_lock_ref_count.get(lock_key, 0) - 1 + if remaining <= 0 and not lock.locked(): + self._session_lock_ref_count.pop(lock_key, None) + self._session_locks.pop(lock_key, None) + else: + self._session_lock_ref_count[lock_key] = remaining + + @staticmethod + def _merge_state( + app_state: dict[str, Any] | None, + user_state: dict[str, Any] | None, + session_state: dict[str, Any], + ) -> dict[str, Any]: + """Merges app, user, and session states into a single state dictionary.""" + merged_state = copy.deepcopy(session_state) + for key, value in (app_state or {}).items(): + merged_state[State.APP_PREFIX + key] = value + for key, value in (user_state or {}).items(): + merged_state[State.USER_PREFIX + key] = value + return merged_state + + def _read_app_state(self, app_name: str) -> dict[str, Any]: + doc = self._app_states().find_one({"_id": app_name}) + return _loads_state(doc.get("state")) if doc else {} + + def _read_user_state(self, app_name: str, user_id: str) -> dict[str, Any]: + doc = self._user_states().find_one( + {"_id": self._user_key(app_name, user_id)} + ) + return _loads_state(doc.get("state")) if doc else {} + + def _merge_state_bucket( + self, collection: Any, doc_id: str, delta: dict[str, Any] + ) -> dict[str, Any]: + """Merges delta into a stored state bucket and returns the merged state.""" + existing = collection.find_one({"_id": doc_id}) + merged = _loads_state(existing.get("state")) if existing else {} + merged.update(delta) + collection.update_one( + {"_id": doc_id}, {"$set": {"state": _dumps_state(merged)}}, upsert=True + ) + return merged + + def _to_session( + self, + doc: dict[str, Any], + merged_state: dict[str, Any], + events: list[Event], + ) -> Session: + session = Session( + id=doc["id"], + app_name=doc["app_name"], + user_id=doc["user_id"], + state=merged_state, + events=events, + last_update_time=doc.get("update_time", 0.0), + ) + session._storage_update_marker = str(doc.get("revision", 0)) + return session + + @override + async def create_session( + self, + *, + app_name: str, + user_id: str, + state: dict[str, Any] | None = None, + session_id: str | None = None, + ) -> Session: + """Creates a new session in MongoDB.""" + + def _create() -> tuple[str, dict[str, Any]]: + sid = session_id or platform_uuid.new_uuid() + state_deltas = _session_util.extract_state_delta(state or {}) + + app_state = ( + self._merge_state_bucket( + self._app_states(), app_name, state_deltas["app"] + ) + if state_deltas["app"] + else self._read_app_state(app_name) + ) + user_state = ( + self._merge_state_bucket( + self._user_states(), + self._user_key(app_name, user_id), + state_deltas["user"], + ) + if state_deltas["user"] + else self._read_user_state(app_name, user_id) + ) + + now = time.time() + doc = { + "_id": self._session_key(app_name, user_id, sid), + "id": sid, + "app_name": app_name, + "user_id": user_id, + "state": _dumps_state(state_deltas["session"]), + "create_time": now, + "update_time": now, + "revision": 0, + } + try: + self._sessions().insert_one(doc) + except Exception as exc: + if exc.__class__.__name__ == "DuplicateKeyError": + raise AlreadyExistsError(f"Session {sid} already exists.") from exc + raise + merged = self._merge_state(app_state, user_state, state_deltas["session"]) + return sid, merged + + sid, merged_state = await asyncio.to_thread(_create) + session = Session( + id=sid, + app_name=app_name, + user_id=user_id, + state=merged_state, + events=[], + last_update_time=time.time(), + ) + session._storage_update_marker = "0" + return session + + @override + async def get_session( + self, + *, + app_name: str, + user_id: str, + session_id: str, + config: GetSessionConfig | None = None, + ) -> Session | None: + """Gets a session from MongoDB.""" + + def _get() -> Session | None: + doc = self._sessions().find_one( + {"_id": self._session_key(app_name, user_id, session_id)} + ) + if not doc: + return None + + # A requested count of zero asks for no event history at all (callers + # use it to probe whether a session exists), so skip the events query. + events: list[Event] = [] + if config is None or config.num_recent_events != 0: + query: dict[str, Any] = { + "app_name": app_name, + "user_id": user_id, + "session_id": session_id, + } + if config and config.after_timestamp is not None: + query["timestamp"] = {"$gte": config.after_timestamp} + + cursor = self._events().find(query).sort("timestamp", -1) + if config and config.num_recent_events is not None: + cursor = cursor.limit(config.num_recent_events) + event_docs = list(cursor) + event_docs.reverse() # restore chronological order + events = [ + Event.model_validate(event_doc["event_data"]) + for event_doc in event_docs + ] + + merged = self._merge_state( + self._read_app_state(app_name), + self._read_user_state(app_name, user_id), + _loads_state(doc.get("state")), + ) + return self._to_session(doc, merged, events) + + return await asyncio.to_thread(_get) + + @override + async def list_sessions( + self, *, app_name: str, user_id: str | None = None + ) -> ListSessionsResponse: + """Lists sessions from MongoDB, oldest update first.""" + + def _list() -> list[Session]: + query: dict[str, Any] = {"app_name": app_name} + if user_id: + query["user_id"] = user_id + docs = list(self._sessions().find(query)) + + app_state = self._read_app_state(app_name) + user_ids = {doc["user_id"] for doc in docs} + user_states = { + uid: self._read_user_state(app_name, uid) for uid in user_ids + } + + sessions = [ + self._to_session( + doc, + self._merge_state( + app_state, + user_states.get(doc["user_id"], {}), + _loads_state(doc.get("state")), + ), + [], + ) + for doc in docs + ] + sessions.sort(key=lambda s: (s.last_update_time, s.user_id, s.id)) + return sessions + + return ListSessionsResponse(sessions=await asyncio.to_thread(_list)) + + @override + async def delete_session( + self, *, app_name: str, user_id: str, session_id: str + ) -> None: + """Deletes a session and its events from MongoDB.""" + + def _delete() -> None: + self._events().delete_many( + {"app_name": app_name, "user_id": user_id, "session_id": session_id} + ) + self._sessions().delete_one( + {"_id": self._session_key(app_name, user_id, session_id)} + ) + + await asyncio.to_thread(_delete) + + @override + async def get_user_state( + self, *, app_name: str, user_id: str + ) -> dict[str, Any]: + """Returns the user-scoped state for the given app and user.""" + return await asyncio.to_thread(self._read_user_state, app_name, user_id) + + @override + async def append_event(self, session: Session, event: Event) -> Event: + """Appends an event to a session in MongoDB.""" + if event.partial: + return event + + self._apply_temp_state(session, event) + event = self._trim_temp_delta_state(event) + + state_delta = ( + event.actions.state_delta + if event.actions and event.actions.state_delta + else {} + ) + state_deltas = _session_util.extract_state_delta(state_delta) + + async with self._with_session_lock( + app_name=session.app_name, + user_id=session.user_id, + session_id=session.id, + ): + + def _append() -> int: + if state_deltas["app"]: + self._merge_state_bucket( + self._app_states(), session.app_name, state_deltas["app"] + ) + if state_deltas["user"]: + self._merge_state_bucket( + self._user_states(), + self._user_key(session.app_name, session.user_id), + state_deltas["user"], + ) + + session_only_state = { + key: value + for key, value in session.state.items() + if not key.startswith(State.APP_PREFIX) + and not key.startswith(State.USER_PREFIX) + and not key.startswith(State.TEMP_PREFIX) + } + session_only_state.update(state_deltas["session"]) + + session_doc_id = self._session_key( + session.app_name, session.user_id, session.id + ) + current = self._sessions().find_one({"_id": session_doc_id}) + if not current: + raise SessionNotFoundError(f"Session {session.id} not found.") + current_revision = current.get("revision", 0) + if session._storage_update_marker is not None and ( + session._storage_update_marker != str(current_revision) + ): + raise StaleSessionError(_STALE_SESSION_ERROR_MESSAGE) + + # The revision filter makes the update a no-op when a concurrent + # writer bumped the revision between our read and write. + updated = self._sessions().find_one_and_update( + {"_id": session_doc_id, "revision": current_revision}, + { + "$set": { + "state": _dumps_state(session_only_state), + "update_time": event.timestamp, + }, + "$inc": {"revision": 1}, + }, + return_document=True, + ) + if updated is None: + raise StaleSessionError(_STALE_SESSION_ERROR_MESSAGE) + + # Upsert keeps event ingestion idempotent across retries. + self._events().replace_one( + {"_id": f"{session_doc_id}/{event.id}"}, + { + "app_name": session.app_name, + "user_id": session.user_id, + "session_id": session.id, + "timestamp": event.timestamp, + "event_data": event.model_dump(exclude_none=True, mode="json"), + }, + upsert=True, + ) + return int(updated.get("revision", current_revision + 1)) + + new_revision = await asyncio.to_thread(_append) + session._storage_update_marker = str(new_revision) + session.last_update_time = event.timestamp + + await super().append_event(session, event) + return event + + async def close(self) -> None: + """Closes the MongoDB client if it was created by this service.""" + if self._owns_client: + await asyncio.to_thread(self._client.close) diff --git a/tests/unittests/integrations/mongodb/test_memory_service.py b/tests/unittests/integrations/mongodb/test_memory_service.py new file mode 100644 index 00000000000..b8900884e24 --- /dev/null +++ b/tests/unittests/integrations/mongodb/test_memory_service.py @@ -0,0 +1,148 @@ +# 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. + +"""Tests for MongoDbMemoryService, backed by mongomock.""" + +from __future__ import annotations + +from google.adk.events.event import Event +from google.adk.integrations.mongodb import MongoDbMemoryService +from google.adk.sessions.session import Session +from google.genai import types +import mongomock +import pytest + +APP = "test_app" +USER = "test_user" + + +@pytest.fixture +def service(): + return MongoDbMemoryService( + mongo_client=mongomock.MongoClient(), database_name="test_db" + ) + + +def _session_with_texts(*texts: str, session_id: str = "s1") -> Session: + events = [ + Event( + invocation_id="inv", + author="user" if i % 2 == 0 else "agent", + content=types.Content( + role="user" if i % 2 == 0 else "model", + parts=[types.Part(text=text)], + ), + timestamp=float(i + 1), + ) + for i, text in enumerate(texts) + ] + return Session(id=session_id, app_name=APP, user_id=USER, events=events) + + +@pytest.mark.asyncio +async def test_add_session_then_search_by_keyword(service): + await service.add_session_to_memory( + _session_with_texts( + "I love hiking in the mountains", + "My favorite trail is the Pacific Crest Trail", + ) + ) + + response = await service.search_memory( + app_name=APP, user_id=USER, query="hiking" + ) + assert len(response.memories) == 1 + assert "hiking" in response.memories[0].content.parts[0].text + assert response.memories[0].author == "user" + + # A query can match several memories. + response = await service.search_memory( + app_name=APP, user_id=USER, query="trail mountains" + ) + assert len(response.memories) == 2 + + +@pytest.mark.asyncio +async def test_search_memory_scopes_by_app_and_user(service): + await service.add_session_to_memory(_session_with_texts("remember hiking")) + + assert ( + await service.search_memory( + app_name="other_app", user_id=USER, query="hiking" + ) + ).memories == [] + assert ( + await service.search_memory( + app_name=APP, user_id="other_user", query="hiking" + ) + ).memories == [] + + +@pytest.mark.asyncio +async def test_reingesting_session_does_not_duplicate(service): + session = _session_with_texts("I love hiking") + await service.add_session_to_memory(session) + await service.add_session_to_memory(session) + + response = await service.search_memory( + app_name=APP, user_id=USER, query="hiking" + ) + assert len(response.memories) == 1 + + +@pytest.mark.asyncio +async def test_search_memory_ignores_stop_words(service): + await service.add_session_to_memory(_session_with_texts("the cat sat")) + + # "the" is a stop word, so a stop-words-only query matches nothing even + # though the ingested text contains "the". + assert ( + await service.search_memory(app_name=APP, user_id=USER, query="the") + ).memories == [] + response = await service.search_memory( + app_name=APP, user_id=USER, query="cat" + ) + assert len(response.memories) == 1 + + +@pytest.mark.asyncio +async def test_search_memory_empty_or_stop_word_query(service): + await service.add_session_to_memory(_session_with_texts("remember hiking")) + assert ( + await service.search_memory(app_name=APP, user_id=USER, query="") + ).memories == [] + assert ( + await service.search_memory(app_name=APP, user_id=USER, query="?!") + ).memories == [] + + +@pytest.mark.asyncio +async def test_add_session_skips_events_without_text(service): + session = Session(id="s1", app_name=APP, user_id=USER) + session.events.append(Event(invocation_id="inv", author="agent")) + await service.add_session_to_memory(session) + assert ( + await service.search_memory(app_name=APP, user_id=USER, query="anything") + ).memories == [] + + +def test_constructor_validates_client_args(): + with pytest.raises(ValueError): + MongoDbMemoryService(database_name="db") + with pytest.raises(ValueError): + MongoDbMemoryService( + database_name="db", + mongo_client=mongomock.MongoClient(), + connection_string="mongodb://localhost:27017", + ) diff --git a/tests/unittests/integrations/mongodb/test_session_service.py b/tests/unittests/integrations/mongodb/test_session_service.py new file mode 100644 index 00000000000..a446ba0d271 --- /dev/null +++ b/tests/unittests/integrations/mongodb/test_session_service.py @@ -0,0 +1,281 @@ +# 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. + +"""Tests for MongoDbSessionService, backed by mongomock.""" + +from __future__ import annotations + +from google.adk.errors import StaleSessionError +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.integrations.mongodb import MongoDbSessionService +from google.adk.sessions.base_session_service import GetSessionConfig +from google.genai import types +import mongomock +import pytest + +APP = "test_app" +USER = "test_user" +USER_2 = "other_user" + + +@pytest.fixture +def client(): + return mongomock.MongoClient() + + +@pytest.fixture +def service(client): + return MongoDbSessionService(mongo_client=client, database_name="test_db") + + +def _text_event( + text: str, *, author: str = "user", timestamp: float = 1.0 +) -> Event: + return Event( + invocation_id="inv", + author=author, + content=types.Content( + role="user" if author == "user" else "model", + parts=[types.Part(text=text)], + ), + timestamp=timestamp, + ) + + +@pytest.mark.asyncio +async def test_create_and_get_session(service): + session = await service.create_session( + app_name=APP, + user_id=USER, + state={"session_key": "s1", "app:app_key": "a1", "user:user_key": "u1"}, + session_id="s1", + ) + + assert session.id == "s1" + assert session.state["session_key"] == "s1" + assert session.state["app:app_key"] == "a1" + assert session.state["user:user_key"] == "u1" + assert session.events == [] + + fetched = await service.get_session( + app_name=APP, user_id=USER, session_id="s1" + ) + assert fetched is not None + assert fetched.state == session.state + assert fetched._storage_update_marker == "0" + + # Missing session returns None. + assert ( + await service.get_session(app_name=APP, user_id=USER, session_id="nope") + is None + ) + + +@pytest.mark.asyncio +async def test_create_duplicate_session_raises(service): + await service.create_session(app_name=APP, user_id=USER, session_id="s1") + with pytest.raises(AlreadyExistsError): + await service.create_session(app_name=APP, user_id=USER, session_id="s1") + + +@pytest.mark.asyncio +async def test_constructor_validates_client_args(client): + with pytest.raises(ValueError): + MongoDbSessionService(database_name="db") + with pytest.raises(ValueError): + MongoDbSessionService( + database_name="db", + mongo_client=client, + connection_string="mongodb://localhost:27017", + ) + + +@pytest.mark.asyncio +async def test_append_event_persists_and_merges_state(service): + session = await service.create_session( + app_name=APP, user_id=USER, session_id="s1" + ) + event = _text_event("hello mongodb", timestamp=10.0) + event.actions = EventActions( + state_delta={"turn": 1, "user:theme": "dark", "temp:scratch": "x"} + ) + await service.append_event(session, event) + + fetched = await service.get_session( + app_name=APP, user_id=USER, session_id="s1" + ) + assert len(fetched.events) == 1 + assert fetched.events[0].content.parts[0].text == "hello mongodb" + assert fetched.state["turn"] == 1 + assert fetched.state["user:theme"] == "dark" + # temp state lives on the in-memory session but is never persisted. + assert "temp:scratch" in session.state + assert "temp:scratch" not in fetched.state + assert fetched._storage_update_marker == "1" + # The in-memory session advanced too. + assert session._storage_update_marker == "1" + assert len(session.events) == 1 + + # User state is shared across the user's sessions. + session2 = await service.create_session( + app_name=APP, user_id=USER, session_id="s2" + ) + assert session2.state["user:theme"] == "dark" + assert await service.get_user_state(app_name=APP, user_id=USER) == { + "theme": "dark" + } + + +@pytest.mark.asyncio +async def test_append_event_missing_session_raises(service): + session = await service.create_session( + app_name=APP, user_id=USER, session_id="s1" + ) + session.id = "deleted" + with pytest.raises(SessionNotFoundError): + await service.append_event(session, _text_event("hi")) + + +@pytest.mark.asyncio +async def test_append_event_stale_session_raises(service): + session = await service.create_session( + app_name=APP, user_id=USER, session_id="s1" + ) + stale_copy = await service.get_session( + app_name=APP, user_id=USER, session_id="s1" + ) + + await service.append_event(session, _text_event("first", timestamp=1.0)) + # stale_copy still holds revision marker "0" while storage is at "1". + with pytest.raises(StaleSessionError): + await service.append_event(stale_copy, _text_event("second", timestamp=2.0)) + + +@pytest.mark.asyncio +async def test_append_event_is_idempotent_on_retry(service): + session = await service.create_session( + app_name=APP, user_id=USER, session_id="s1" + ) + event = _text_event("hello", timestamp=1.0) + await service.append_event(session, event) + + # Simulate a storage-level retry of the same event write: the document is + # replaced in place instead of duplicated. + service._events().replace_one( + {"_id": f"{APP}/{USER}/s1/{event.id}"}, + { + "app_name": APP, + "user_id": USER, + "session_id": "s1", + "timestamp": event.timestamp, + "event_data": event.model_dump(exclude_none=True, mode="json"), + }, + upsert=True, + ) + fetched = await service.get_session( + app_name=APP, user_id=USER, session_id="s1" + ) + assert len(fetched.events) == 1 + + +@pytest.mark.asyncio +async def test_get_session_config_filters_events(service): + session = await service.create_session( + app_name=APP, user_id=USER, session_id="s1" + ) + for i in range(1, 5): + await service.append_event( + session, _text_event(f"event {i}", timestamp=float(i)) + ) + + # num_recent_events=0 -> no events loaded. + fetched = await service.get_session( + app_name=APP, + user_id=USER, + session_id="s1", + config=GetSessionConfig(num_recent_events=0), + ) + assert fetched.events == [] + + # num_recent_events=2 -> two most recent, in chronological order. + fetched = await service.get_session( + app_name=APP, + user_id=USER, + session_id="s1", + config=GetSessionConfig(num_recent_events=2), + ) + assert [e.content.parts[0].text for e in fetched.events] == [ + "event 3", + "event 4", + ] + + # after_timestamp -> only events at or after the cursor. + fetched = await service.get_session( + app_name=APP, + user_id=USER, + session_id="s1", + config=GetSessionConfig(after_timestamp=2.5), + ) + assert [e.content.parts[0].text for e in fetched.events] == [ + "event 3", + "event 4", + ] + + +@pytest.mark.asyncio +async def test_list_sessions_sorted_oldest_first(service): + await service.create_session(app_name=APP, user_id=USER, session_id="s1") + await service.create_session(app_name=APP, user_id=USER, session_id="s2") + await service.create_session(app_name=APP, user_id=USER_2, session_id="s3") + + response = await service.list_sessions(app_name=APP, user_id=USER) + assert [s.id for s in response.sessions] == ["s1", "s2"] + + response = await service.list_sessions(app_name=APP) + assert {s.id for s in response.sessions} == {"s1", "s2", "s3"} + + +@pytest.mark.asyncio +async def test_delete_session_removes_session_and_events(service): + session = await service.create_session( + app_name=APP, user_id=USER, session_id="s1" + ) + await service.append_event(session, _text_event("bye", timestamp=1.0)) + + await service.delete_session(app_name=APP, user_id=USER, session_id="s1") + + assert ( + await service.get_session(app_name=APP, user_id=USER, session_id="s1") + is None + ) + assert service._events().count_documents({}) == 0 + + +@pytest.mark.asyncio +async def test_state_keys_with_dots_round_trip(service): + """MongoDB forbids dots in document keys; JSON-encoded state does not care.""" + session = await service.create_session( + app_name=APP, user_id=USER, state={"nested.key": 1}, session_id="s1" + ) + await service.append_event( + session, + _text_event("x", timestamp=1.0), + ) + fetched = await service.get_session( + app_name=APP, user_id=USER, session_id="s1" + ) + assert fetched.state["nested.key"] == 1