fix: Migrate VertexAiMemoryBankService to use the async Vertex AI client

This change updates the VertexAiMemoryBankService to utilize the asynchronous interface provided by `vertexai.Client().aio`. This involves:
-   Retrieving the async client via `_get_api_client().aio`.
-   Awaiting calls to `generate` and `retrieve`.
-   Using `async for` to iterate over the results of the `retrieve` method, as it now returns an async iterator

Close #4386

Co-authored-by: George Weale <gweale@google.com>
PiperOrigin-RevId: 867675311
This commit is contained in:
George Weale
2026-02-09 10:45:15 -08:00
committed by Copybara-Service
parent 0758f877b1
commit 64a44c2897
2 changed files with 92 additions and 28 deletions
@@ -27,6 +27,8 @@ from .base_memory_service import SearchMemoryResponse
from .memory_entry import MemoryEntry
if TYPE_CHECKING:
import vertexai
from ..sessions.session import Session
logger = logging.getLogger('google_adk.' + __name__)
@@ -88,8 +90,8 @@ class VertexAiMemoryBankService(BaseMemoryService):
'content': event.content.model_dump(exclude_none=True, mode='json')
})
if events:
client = self._get_api_client()
operation = client.agent_engines.memories.generate(
api_client = self._get_api_client()
operation = await api_client.agent_engines.memories.generate(
name='reasoningEngines/' + self._agent_engine_id,
direct_contents_source={'events': events},
scope={
@@ -108,22 +110,24 @@ class VertexAiMemoryBankService(BaseMemoryService):
if not self._agent_engine_id:
raise ValueError('Agent Engine ID is required for Memory Bank.')
client = self._get_api_client()
retrieved_memories_iterator = client.agent_engines.memories.retrieve(
name='reasoningEngines/' + self._agent_engine_id,
scope={
'app_name': app_name,
'user_id': user_id,
},
similarity_search_params={
'search_query': query,
},
api_client = self._get_api_client()
retrieved_memories_iterator = (
await api_client.agent_engines.memories.retrieve(
name='reasoningEngines/' + self._agent_engine_id,
scope={
'app_name': app_name,
'user_id': user_id,
},
similarity_search_params={
'search_query': query,
},
)
)
logger.info('Search memory response received.')
memory_events = []
for retrieved_memory in retrieved_memories_iterator:
memory_events: list[MemoryEntry] = []
async for retrieved_memory in retrieved_memories_iterator:
# TODO: add more complex error handling
logger.debug('Retrieved memory: %s', retrieved_memory)
memory_events.append(
@@ -138,13 +142,14 @@ class VertexAiMemoryBankService(BaseMemoryService):
)
return SearchMemoryResponse(memories=memory_events)
def _get_api_client(self):
def _get_api_client(self) -> vertexai.AsyncClient:
"""Instantiates an API client for the given project and location.
It needs to be instantiated inside each request so that the event loop
management can be properly propagated.
Returns:
An API client for the given project and location or express mode api key.
An async API client for the given project and location or express mode api
key.
"""
import vertexai
@@ -152,7 +157,7 @@ class VertexAiMemoryBankService(BaseMemoryService):
project=self._project,
location=self._location,
api_key=self._express_mode_api_key,
)
).aio
def _should_filter_out_event(content: types.Content) -> bool:
@@ -13,6 +13,8 @@
# limitations under the License.
from datetime import datetime
from typing import Any
from typing import Iterable
from typing import Optional
from unittest import mock
@@ -25,6 +27,25 @@ import pytest
MOCK_APP_NAME = 'test-app'
MOCK_USER_ID = 'test-user'
class _AsyncListIterator:
"""Minimal async iterator wrapper for list-like results."""
def __init__(self, items: Iterable[Any]):
self._items = list(items)
self._index = 0
def __aiter__(self) -> '_AsyncListIterator':
return self
async def __anext__(self) -> Any:
if self._index >= len(self._items):
raise StopAsyncIteration
item = self._items[self._index]
self._index += 1
return item
MOCK_SESSION = Session(
app_name=MOCK_APP_NAME,
user_id=MOCK_USER_ID,
@@ -88,11 +109,15 @@ def mock_vertex_ai_memory_bank_service(
@pytest.fixture
def mock_vertexai_client():
with mock.patch('vertexai.Client') as mock_client_constructor:
mock_async_client = mock.MagicMock()
mock_async_client.agent_engines.memories.generate = mock.AsyncMock()
mock_async_client.agent_engines.memories.retrieve = mock.AsyncMock()
mock_client = mock.MagicMock()
mock_client.agent_engines.memories.generate = mock.MagicMock()
mock_client.agent_engines.memories.retrieve = mock.MagicMock()
mock_client.aio = mock_async_client
mock_client_constructor.return_value = mock_client
yield mock_client
yield mock_async_client
@pytest.mark.asyncio
@@ -115,7 +140,7 @@ async def test_add_session_to_memory(mock_vertexai_client):
memory_service = mock_vertex_ai_memory_bank_service()
await memory_service.add_session_to_memory(MOCK_SESSION)
mock_vertexai_client.agent_engines.memories.generate.assert_called_once_with(
mock_vertexai_client.agent_engines.memories.generate.assert_awaited_once_with(
name='reasoningEngines/123',
direct_contents_source={
'events': [
@@ -136,7 +161,7 @@ async def test_add_empty_session_to_memory(mock_vertexai_client):
memory_service = mock_vertex_ai_memory_bank_service()
await memory_service.add_session_to_memory(MOCK_SESSION_WITH_EMPTY_EVENTS)
mock_vertexai_client.agent_engines.memories.generate.assert_not_called()
mock_vertexai_client.agent_engines.memories.generate.assert_not_awaited()
@pytest.mark.asyncio
@@ -147,16 +172,16 @@ async def test_search_memory(mock_vertexai_client):
2024, 12, 12, 12, 12, 12, 123456
)
mock_vertexai_client.agent_engines.memories.retrieve.return_value = [
retrieved_memory
]
mock_vertexai_client.agent_engines.memories.retrieve.return_value = (
_AsyncListIterator([retrieved_memory])
)
memory_service = mock_vertex_ai_memory_bank_service()
result = await memory_service.search_memory(
app_name=MOCK_APP_NAME, user_id=MOCK_USER_ID, query='query'
)
mock_vertexai_client.agent_engines.memories.retrieve.assert_called_once_with(
mock_vertexai_client.agent_engines.memories.retrieve.assert_awaited_once_with(
name='reasoningEngines/123',
scope={'app_name': MOCK_APP_NAME, 'user_id': MOCK_USER_ID},
similarity_search_params={'search_query': 'query'},
@@ -168,17 +193,51 @@ async def test_search_memory(mock_vertexai_client):
@pytest.mark.asyncio
async def test_search_memory_empty_results(mock_vertexai_client):
mock_vertexai_client.agent_engines.memories.retrieve.return_value = []
mock_vertexai_client.agent_engines.memories.retrieve.return_value = (
_AsyncListIterator([])
)
memory_service = mock_vertex_ai_memory_bank_service()
result = await memory_service.search_memory(
app_name=MOCK_APP_NAME, user_id=MOCK_USER_ID, query='query'
)
mock_vertexai_client.agent_engines.memories.retrieve.assert_called_once_with(
mock_vertexai_client.agent_engines.memories.retrieve.assert_awaited_once_with(
name='reasoningEngines/123',
scope={'app_name': MOCK_APP_NAME, 'user_id': MOCK_USER_ID},
similarity_search_params={'search_query': 'query'},
)
assert len(result.memories) == 0
@pytest.mark.asyncio
async def test_search_memory_uses_async_client_path():
sync_client = mock.MagicMock()
sync_client.agent_engines.memories.retrieve.side_effect = AssertionError(
'sync retrieve should not be called'
)
async_client = mock.MagicMock()
async_client.agent_engines.memories.retrieve = mock.AsyncMock(
return_value=_AsyncListIterator([])
)
with mock.patch('vertexai.Client') as mock_client_constructor:
mock_client_constructor.return_value = mock.MagicMock(
aio=async_client,
agent_engines=sync_client.agent_engines,
)
memory_service = mock_vertex_ai_memory_bank_service()
await memory_service.search_memory(
app_name=MOCK_APP_NAME,
user_id=MOCK_USER_ID,
query='query',
)
async_client.agent_engines.memories.retrieve.assert_awaited_once_with(
name='reasoningEngines/123',
scope={'app_name': MOCK_APP_NAME, 'user_id': MOCK_USER_ID},
similarity_search_params={'search_query': 'query'},
)
sync_client.agent_engines.memories.retrieve.assert_not_called()