mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat: add generate/create modes for Vertex AI Memory Bank writes
add_events_to_memory now supports memory_write_mode to select generate (event-based extraction/consolidation) or create (direct raw fact writes via memory_facts). This now lets custom memory pipelines while keeping generate as the default path Co-authored-by: George Weale <gweale@google.com> PiperOrigin-RevId: 869897256
This commit is contained in:
committed by
Copybara-Service
parent
d5332f4434
commit
811e50a0cb
@@ -345,7 +345,8 @@ class Context(ReadonlyContext):
|
||||
|
||||
Args:
|
||||
events: Explicit events to add to memory.
|
||||
custom_metadata: Optional standard metadata for memory generation.
|
||||
custom_metadata: Optional metadata forwarded to the configured memory
|
||||
service. Supported keys are implementation-specific.
|
||||
|
||||
Raises:
|
||||
ValueError: If memory service is not available.
|
||||
@@ -362,6 +363,33 @@ class Context(ReadonlyContext):
|
||||
custom_metadata=custom_metadata,
|
||||
)
|
||||
|
||||
async def add_memory(
|
||||
self,
|
||||
*,
|
||||
memories: Sequence[str],
|
||||
custom_metadata: Mapping[str, object] | None = None,
|
||||
) -> None:
|
||||
"""Adds explicit memory items directly to the memory service.
|
||||
|
||||
Uses this callback's current session identifiers as memory scope.
|
||||
|
||||
Args:
|
||||
memories: Explicit memory items to add.
|
||||
custom_metadata: Optional metadata forwarded to the configured memory
|
||||
service. Supported keys are implementation-specific.
|
||||
|
||||
Raises:
|
||||
ValueError: If memory service is not available.
|
||||
"""
|
||||
if self._invocation_context.memory_service is None:
|
||||
raise ValueError("Cannot add memory: memory service is not available.")
|
||||
await self._invocation_context.memory_service.add_memory(
|
||||
app_name=self._invocation_context.session.app_name,
|
||||
user_id=self._invocation_context.session.user_id,
|
||||
memories=memories,
|
||||
custom_metadata=custom_metadata,
|
||||
)
|
||||
|
||||
async def search_memory(self, query: str) -> SearchMemoryResponse:
|
||||
"""Searches the memory of the current user.
|
||||
|
||||
|
||||
@@ -86,13 +86,40 @@ class BaseMemoryService(ABC):
|
||||
session_id: Optional session ID for memory scope/partitioning.
|
||||
custom_metadata: Optional, portable metadata for memory generation. Prefer
|
||||
this for service-specific fields (e.g., TTL) that may later become
|
||||
first-class API parameters.
|
||||
first-class API parameters. Supported keys are
|
||||
implementation-defined by each memory service.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"This memory service does not support adding event deltas. "
|
||||
"Call add_session_to_memory(session) to ingest the full session."
|
||||
)
|
||||
|
||||
async def add_memory(
|
||||
self,
|
||||
*,
|
||||
app_name: str,
|
||||
user_id: str,
|
||||
memories: Sequence[str],
|
||||
custom_metadata: Mapping[str, object] | None = None,
|
||||
) -> None:
|
||||
"""Adds explicit memory items directly to the memory service.
|
||||
|
||||
This is intended for services that support direct memory writes in addition
|
||||
to event-based memory generation.
|
||||
|
||||
Args:
|
||||
app_name: The application name for memory scope.
|
||||
user_id: The user ID for memory scope.
|
||||
memories: Explicit memory items to add.
|
||||
custom_metadata: Optional, portable metadata for memory writes. Supported
|
||||
keys are implementation-defined by each memory service.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"This memory service does not support direct memory writes. "
|
||||
"Call add_events_to_memory(...) or add_session_to_memory(session) "
|
||||
"instead."
|
||||
)
|
||||
|
||||
@abstractmethod
|
||||
async def search_memory(
|
||||
self,
|
||||
|
||||
@@ -17,6 +17,7 @@ from __future__ import annotations
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
import logging
|
||||
from typing import Optional
|
||||
from typing import TYPE_CHECKING
|
||||
@@ -37,7 +38,7 @@ if TYPE_CHECKING:
|
||||
|
||||
logger = logging.getLogger('google_adk.' + __name__)
|
||||
|
||||
_GENERATE_MEMORIES_CONFIG_KEYS = frozenset({
|
||||
_GENERATE_MEMORIES_CONFIG_FALLBACK_KEYS = frozenset({
|
||||
'disable_consolidation',
|
||||
'disable_memory_revisions',
|
||||
'http_options',
|
||||
@@ -49,6 +50,20 @@ _GENERATE_MEMORIES_CONFIG_KEYS = frozenset({
|
||||
'wait_for_completion',
|
||||
})
|
||||
|
||||
_CREATE_MEMORY_CONFIG_FALLBACK_KEYS = frozenset({
|
||||
'description',
|
||||
'disable_memory_revisions',
|
||||
'display_name',
|
||||
'expire_time',
|
||||
'http_options',
|
||||
'metadata',
|
||||
'revision_expire_time',
|
||||
'revision_ttl',
|
||||
'topics',
|
||||
'ttl',
|
||||
'wait_for_completion',
|
||||
})
|
||||
|
||||
|
||||
def _supports_generate_memories_metadata() -> bool:
|
||||
"""Returns whether installed Vertex SDK supports config.metadata."""
|
||||
@@ -62,6 +77,61 @@ def _supports_generate_memories_metadata() -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _supports_create_memory_metadata() -> bool:
|
||||
"""Returns whether installed Vertex SDK supports create config.metadata."""
|
||||
try:
|
||||
from vertexai._genai.types import common as vertex_common_types
|
||||
except ImportError:
|
||||
return False
|
||||
return 'metadata' in vertex_common_types.AgentEngineMemoryConfig.model_fields
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _get_generate_memories_config_keys() -> frozenset[str]:
|
||||
"""Returns supported config keys for memories.generate.
|
||||
|
||||
Uses SDK runtime model fields when available and falls back to a static
|
||||
allowlist to preserve compatibility when introspection is unavailable.
|
||||
"""
|
||||
try:
|
||||
from vertexai._genai.types import common as vertex_common_types
|
||||
except ImportError:
|
||||
return _GENERATE_MEMORIES_CONFIG_FALLBACK_KEYS
|
||||
|
||||
try:
|
||||
model_fields = (
|
||||
vertex_common_types.GenerateAgentEngineMemoriesConfig.model_fields
|
||||
)
|
||||
except AttributeError:
|
||||
return _GENERATE_MEMORIES_CONFIG_FALLBACK_KEYS
|
||||
|
||||
if not isinstance(model_fields, Mapping):
|
||||
return _GENERATE_MEMORIES_CONFIG_FALLBACK_KEYS
|
||||
return frozenset(model_fields.keys())
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _get_create_memory_config_keys() -> frozenset[str]:
|
||||
"""Returns supported config keys for memories.create.
|
||||
|
||||
Uses SDK runtime model fields when available and falls back to a static
|
||||
allowlist to preserve compatibility when introspection is unavailable.
|
||||
"""
|
||||
try:
|
||||
from vertexai._genai.types import common as vertex_common_types
|
||||
except ImportError:
|
||||
return _CREATE_MEMORY_CONFIG_FALLBACK_KEYS
|
||||
|
||||
try:
|
||||
model_fields = vertex_common_types.AgentEngineMemoryConfig.model_fields
|
||||
except AttributeError:
|
||||
return _CREATE_MEMORY_CONFIG_FALLBACK_KEYS
|
||||
|
||||
if not isinstance(model_fields, Mapping):
|
||||
return _CREATE_MEMORY_CONFIG_FALLBACK_KEYS
|
||||
return frozenset(model_fields.keys())
|
||||
|
||||
|
||||
class VertexAiMemoryBankService(BaseMemoryService):
|
||||
"""Implementation of the BaseMemoryService using Vertex AI Memory Bank."""
|
||||
|
||||
@@ -122,6 +192,15 @@ class VertexAiMemoryBankService(BaseMemoryService):
|
||||
session_id: str | None = None,
|
||||
custom_metadata: Mapping[str, object] | None = None,
|
||||
) -> None:
|
||||
"""Adds events to Vertex AI Memory Bank via memories.generate.
|
||||
|
||||
Args:
|
||||
app_name: The application name for memory scope.
|
||||
user_id: The user ID for memory scope.
|
||||
events: The events to process for memory generation.
|
||||
session_id: Optional session ID. Currently unused.
|
||||
custom_metadata: Optional service-specific metadata for generate config.
|
||||
"""
|
||||
_ = session_id
|
||||
await self._add_events_to_memory_from_events(
|
||||
app_name=app_name,
|
||||
@@ -130,6 +209,23 @@ class VertexAiMemoryBankService(BaseMemoryService):
|
||||
custom_metadata=custom_metadata,
|
||||
)
|
||||
|
||||
@override
|
||||
async def add_memory(
|
||||
self,
|
||||
*,
|
||||
app_name: str,
|
||||
user_id: str,
|
||||
memories: Sequence[str],
|
||||
custom_metadata: Mapping[str, object] | None = None,
|
||||
) -> None:
|
||||
"""Adds explicit memory items via Vertex memories.create."""
|
||||
await self._add_memories_via_create(
|
||||
app_name=app_name,
|
||||
user_id=user_id,
|
||||
memories=memories,
|
||||
custom_metadata=custom_metadata,
|
||||
)
|
||||
|
||||
async def _add_events_to_memory_from_events(
|
||||
self,
|
||||
*,
|
||||
@@ -166,6 +262,34 @@ class VertexAiMemoryBankService(BaseMemoryService):
|
||||
else:
|
||||
logger.info('No events to add to memory.')
|
||||
|
||||
async def _add_memories_via_create(
|
||||
self,
|
||||
*,
|
||||
app_name: str,
|
||||
user_id: str,
|
||||
memories: Sequence[str],
|
||||
custom_metadata: Mapping[str, object] | None = None,
|
||||
) -> None:
|
||||
"""Adds direct memory items without server-side extraction."""
|
||||
if not self._agent_engine_id:
|
||||
raise ValueError('Agent Engine ID is required for Memory Bank.')
|
||||
|
||||
memory_texts = _validate_memory_texts(memories)
|
||||
api_client = self._get_api_client()
|
||||
config = _build_create_memory_config(custom_metadata)
|
||||
for memory_text in memory_texts:
|
||||
operation = await api_client.agent_engines.memories.create(
|
||||
name='reasoningEngines/' + self._agent_engine_id,
|
||||
fact=memory_text,
|
||||
scope={
|
||||
'app_name': app_name,
|
||||
'user_id': user_id,
|
||||
},
|
||||
config=config,
|
||||
)
|
||||
logger.info('Create memory response received.')
|
||||
logger.debug('Create memory response: %s', operation)
|
||||
|
||||
@override
|
||||
async def search_memory(self, *, app_name: str, user_id: str, query: str):
|
||||
if not self._agent_engine_id:
|
||||
@@ -237,6 +361,7 @@ def _build_generate_memories_config(
|
||||
"""Builds a valid memories.generate config from caller metadata."""
|
||||
config: dict[str, object] = {'wait_for_completion': False}
|
||||
supports_metadata = _supports_generate_memories_metadata()
|
||||
config_keys = _get_generate_memories_config_keys()
|
||||
if not custom_metadata:
|
||||
return config
|
||||
|
||||
@@ -267,7 +392,7 @@ def _build_generate_memories_config(
|
||||
' mapping.'
|
||||
)
|
||||
continue
|
||||
if key in _GENERATE_MEMORIES_CONFIG_KEYS:
|
||||
if key in config_keys:
|
||||
if value is None:
|
||||
continue
|
||||
config[key] = value
|
||||
@@ -304,6 +429,96 @@ def _build_generate_memories_config(
|
||||
return config
|
||||
|
||||
|
||||
def _build_create_memory_config(
|
||||
custom_metadata: Mapping[str, object] | None,
|
||||
) -> dict[str, object]:
|
||||
"""Builds a valid memories.create config from caller metadata."""
|
||||
config: dict[str, object] = {'wait_for_completion': False}
|
||||
supports_metadata = _supports_create_memory_metadata()
|
||||
config_keys = _get_create_memory_config_keys()
|
||||
if not custom_metadata:
|
||||
return config
|
||||
|
||||
logger.debug('Memory creation metadata: %s', custom_metadata)
|
||||
|
||||
metadata_by_key: dict[str, object] = {}
|
||||
for key, value in custom_metadata.items():
|
||||
if key == 'metadata':
|
||||
if value is None:
|
||||
continue
|
||||
if not supports_metadata:
|
||||
logger.warning(
|
||||
'Ignoring metadata because installed Vertex SDK does not support'
|
||||
' create config.metadata.'
|
||||
)
|
||||
continue
|
||||
if isinstance(value, Mapping):
|
||||
config['metadata'] = _build_vertex_metadata(value)
|
||||
else:
|
||||
logger.warning(
|
||||
'Ignoring metadata because custom_metadata["metadata"] is not a'
|
||||
' mapping.'
|
||||
)
|
||||
continue
|
||||
if key in config_keys:
|
||||
if value is None:
|
||||
continue
|
||||
config[key] = value
|
||||
else:
|
||||
metadata_by_key[key] = value
|
||||
|
||||
if not metadata_by_key:
|
||||
return config
|
||||
|
||||
if not supports_metadata:
|
||||
logger.warning(
|
||||
'Ignoring custom metadata keys %s because installed Vertex SDK does '
|
||||
'not support create config.metadata.',
|
||||
sorted(metadata_by_key.keys()),
|
||||
)
|
||||
return config
|
||||
|
||||
existing_metadata = config.get('metadata')
|
||||
if existing_metadata is None:
|
||||
config['metadata'] = _build_vertex_metadata(metadata_by_key)
|
||||
return config
|
||||
|
||||
if isinstance(existing_metadata, Mapping):
|
||||
merged_metadata = dict(existing_metadata)
|
||||
merged_metadata.update(_build_vertex_metadata(metadata_by_key))
|
||||
config['metadata'] = merged_metadata
|
||||
return config
|
||||
|
||||
logger.warning(
|
||||
'Ignoring custom metadata keys %s because config.metadata is not a'
|
||||
' mapping.',
|
||||
sorted(metadata_by_key.keys()),
|
||||
)
|
||||
return config
|
||||
|
||||
|
||||
def _validate_memory_texts(
|
||||
memories: Sequence[str],
|
||||
) -> list[str]:
|
||||
"""Validates direct textual memory items passed to add_memory."""
|
||||
if isinstance(memories, str):
|
||||
raise TypeError('memories must be a sequence of strings.')
|
||||
if not isinstance(memories, Sequence):
|
||||
raise TypeError('memories must be a sequence of strings.')
|
||||
memory_texts: list[str] = []
|
||||
for index, raw_memory in enumerate(memories):
|
||||
if not isinstance(raw_memory, str):
|
||||
raise TypeError(f'memories[{index}] must be a string.')
|
||||
memory_text = raw_memory.strip()
|
||||
if not memory_text:
|
||||
raise ValueError(f'memories[{index}] must not be empty.')
|
||||
memory_texts.append(memory_text)
|
||||
|
||||
if not memory_texts:
|
||||
raise ValueError('memories must contain at least one entry.')
|
||||
return memory_texts
|
||||
|
||||
|
||||
def _build_vertex_metadata(
|
||||
metadata_by_key: Mapping[str, object],
|
||||
) -> dict[str, object]:
|
||||
|
||||
@@ -412,6 +412,37 @@ class TestCallbackContextAddEventsToMemory:
|
||||
):
|
||||
await context.add_events_to_memory(events=[MagicMock()])
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_memory_forwards_metadata(self, mock_invocation_context):
|
||||
"""Tests that add_memory forwards memories and metadata."""
|
||||
memory_service = AsyncMock()
|
||||
mock_invocation_context.memory_service = memory_service
|
||||
memories = ["fact one"]
|
||||
metadata = {"ttl": "6000s"}
|
||||
|
||||
context = CallbackContext(mock_invocation_context)
|
||||
await context.add_memory(memories=memories, custom_metadata=metadata)
|
||||
|
||||
memory_service.add_memory.assert_called_once_with(
|
||||
app_name=mock_invocation_context.session.app_name,
|
||||
user_id=mock_invocation_context.session.user_id,
|
||||
memories=memories,
|
||||
custom_metadata=metadata,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_memory_no_service_raises(self, mock_invocation_context):
|
||||
"""Tests that add_memory raises ValueError with no service."""
|
||||
mock_invocation_context.memory_service = None
|
||||
|
||||
context = CallbackContext(mock_invocation_context)
|
||||
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match=r"Cannot add memory: memory service is not available\.",
|
||||
):
|
||||
await context.add_memory(memories=["fact one"])
|
||||
|
||||
|
||||
class TestToolContextAddSessionToMemory:
|
||||
"""Test the add_session_to_memory method in ToolContext."""
|
||||
|
||||
@@ -486,3 +486,33 @@ class TestContextMemoryMethods:
|
||||
match=r"Cannot add events to memory: memory service is not available\.",
|
||||
):
|
||||
await context.add_events_to_memory(events=[MagicMock()])
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_memory_forwards_metadata(self, mock_invocation_context):
|
||||
"""Tests that add_memory forwards memories and metadata."""
|
||||
memory_service = AsyncMock()
|
||||
mock_invocation_context.memory_service = memory_service
|
||||
memories = ["fact one"]
|
||||
metadata = {"ttl": "6000s"}
|
||||
|
||||
context = Context(mock_invocation_context)
|
||||
await context.add_memory(memories=memories, custom_metadata=metadata)
|
||||
|
||||
memory_service.add_memory.assert_called_once_with(
|
||||
app_name=mock_invocation_context.session.app_name,
|
||||
user_id=mock_invocation_context.session.user_id,
|
||||
memories=memories,
|
||||
custom_metadata=metadata,
|
||||
)
|
||||
|
||||
async def test_add_memory_no_service_raises(self, mock_invocation_context):
|
||||
"""Test that add_memory raises ValueError when no service."""
|
||||
mock_invocation_context.memory_service = None
|
||||
|
||||
context = Context(mock_invocation_context)
|
||||
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match=r"Cannot add memory: memory service is not available\.",
|
||||
):
|
||||
await context.add_memory(memories=["fact one"])
|
||||
|
||||
@@ -19,6 +19,7 @@ from typing import Optional
|
||||
from unittest import mock
|
||||
|
||||
from google.adk.events.event import Event
|
||||
from google.adk.memory import vertex_ai_memory_bank_service as memory_service_module
|
||||
from google.adk.memory.vertex_ai_memory_bank_service import VertexAiMemoryBankService
|
||||
from google.adk.sessions.session import Session
|
||||
from google.genai import types
|
||||
@@ -36,6 +37,10 @@ def _supports_generate_memories_metadata() -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _supports_create_memory_metadata() -> bool:
|
||||
return 'metadata' in vertex_common_types.AgentEngineMemoryConfig.model_fields
|
||||
|
||||
|
||||
class _AsyncListIterator:
|
||||
"""Minimal async iterator wrapper for list-like results."""
|
||||
|
||||
@@ -114,11 +119,58 @@ def mock_vertex_ai_memory_bank_service(
|
||||
)
|
||||
|
||||
|
||||
def test_build_generate_memories_config_uses_runtime_config_keys():
|
||||
with (
|
||||
mock.patch.object(
|
||||
memory_service_module,
|
||||
'_get_generate_memories_config_keys',
|
||||
return_value=frozenset({'wait_for_completion', 'new_generate_key'}),
|
||||
),
|
||||
mock.patch.object(
|
||||
memory_service_module,
|
||||
'_supports_generate_memories_metadata',
|
||||
return_value=False,
|
||||
),
|
||||
):
|
||||
config = memory_service_module._build_generate_memories_config(
|
||||
{'new_generate_key': 'value'}
|
||||
)
|
||||
|
||||
assert config == {
|
||||
'wait_for_completion': False,
|
||||
'new_generate_key': 'value',
|
||||
}
|
||||
|
||||
|
||||
def test_build_create_memory_config_uses_runtime_config_keys():
|
||||
with (
|
||||
mock.patch.object(
|
||||
memory_service_module,
|
||||
'_get_create_memory_config_keys',
|
||||
return_value=frozenset({'wait_for_completion', 'new_create_key'}),
|
||||
),
|
||||
mock.patch.object(
|
||||
memory_service_module,
|
||||
'_supports_create_memory_metadata',
|
||||
return_value=False,
|
||||
),
|
||||
):
|
||||
config = memory_service_module._build_create_memory_config(
|
||||
{'new_create_key': 'value'}
|
||||
)
|
||||
|
||||
assert config == {
|
||||
'wait_for_completion': False,
|
||||
'new_create_key': 'value',
|
||||
}
|
||||
|
||||
|
||||
@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.create = mock.AsyncMock()
|
||||
mock_async_client.agent_engines.memories.retrieve = mock.AsyncMock()
|
||||
|
||||
mock_client = mock.MagicMock()
|
||||
@@ -237,6 +289,7 @@ async def test_add_events_to_memory_without_session_id(
|
||||
]
|
||||
)
|
||||
vertex_common_types.GenerateAgentEngineMemoriesConfig(**generate_config)
|
||||
mock_vertexai_client.agent_engines.memories.create.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -373,6 +426,86 @@ async def test_add_events_to_memory_with_filtered_events_skips_rpc(
|
||||
)
|
||||
|
||||
mock_vertexai_client.agent_engines.memories.generate.assert_not_called()
|
||||
mock_vertexai_client.agent_engines.memories.create.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_memory_calls_create(
|
||||
mock_vertexai_client,
|
||||
):
|
||||
memory_service = mock_vertex_ai_memory_bank_service()
|
||||
await memory_service.add_memory(
|
||||
app_name=MOCK_SESSION.app_name,
|
||||
user_id=MOCK_SESSION.user_id,
|
||||
memories=['fact one', 'fact two'],
|
||||
custom_metadata={
|
||||
'ttl': '6000s',
|
||||
'source': 'agent',
|
||||
},
|
||||
)
|
||||
|
||||
expected_config = {
|
||||
'wait_for_completion': False,
|
||||
'ttl': '6000s',
|
||||
}
|
||||
if _supports_create_memory_metadata():
|
||||
expected_config['metadata'] = {'source': {'string_value': 'agent'}}
|
||||
|
||||
mock_vertexai_client.agent_engines.memories.generate.assert_not_called()
|
||||
mock_vertexai_client.agent_engines.memories.create.assert_has_awaits([
|
||||
mock.call(
|
||||
name='reasoningEngines/123',
|
||||
fact='fact one',
|
||||
scope={'app_name': MOCK_APP_NAME, 'user_id': MOCK_USER_ID},
|
||||
config=expected_config,
|
||||
),
|
||||
mock.call(
|
||||
name='reasoningEngines/123',
|
||||
fact='fact two',
|
||||
scope={'app_name': MOCK_APP_NAME, 'user_id': MOCK_USER_ID},
|
||||
config=expected_config,
|
||||
),
|
||||
])
|
||||
assert mock_vertexai_client.agent_engines.memories.create.await_count == 2
|
||||
|
||||
create_config = (
|
||||
mock_vertexai_client.agent_engines.memories.create.call_args.kwargs[
|
||||
'config'
|
||||
]
|
||||
)
|
||||
vertex_common_types.AgentEngineMemoryConfig(**create_config)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_memory_missing_memories_raises(
|
||||
mock_vertexai_client,
|
||||
):
|
||||
memory_service = mock_vertex_ai_memory_bank_service()
|
||||
with pytest.raises(
|
||||
ValueError, match=r'memories must contain at least one entry'
|
||||
):
|
||||
await memory_service.add_memory(
|
||||
app_name=MOCK_SESSION.app_name,
|
||||
user_id=MOCK_SESSION.user_id,
|
||||
memories=[],
|
||||
)
|
||||
mock_vertexai_client.agent_engines.memories.generate.assert_not_called()
|
||||
mock_vertexai_client.agent_engines.memories.create.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_memory_with_invalid_memory_type_raises(
|
||||
mock_vertexai_client,
|
||||
):
|
||||
memory_service = mock_vertex_ai_memory_bank_service()
|
||||
with pytest.raises(TypeError, match=r'memories\[0\] must be a string'):
|
||||
await memory_service.add_memory(
|
||||
app_name=MOCK_SESSION.app_name,
|
||||
user_id=MOCK_SESSION.user_id,
|
||||
memories=[123],
|
||||
)
|
||||
mock_vertexai_client.agent_engines.memories.generate.assert_not_called()
|
||||
mock_vertexai_client.agent_engines.memories.create.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
Reference in New Issue
Block a user