mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
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:
committed by
Copybara-Service
parent
0758f877b1
commit
64a44c2897
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user