Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 10 additions & 3 deletions src/google/adk/memory/in_memory_memory_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@

from collections.abc import Mapping
from collections.abc import Sequence
from copy import deepcopy
import itertools
import re
import threading
Expand Down Expand Up @@ -98,7 +99,7 @@ async def add_session_to_memory(self, session: Session) -> None:
with self._lock:
self._session_events[user_key] = self._session_events.get(user_key, {})
self._session_events[user_key][session.id] = [
event
event.model_copy(update={'content': deepcopy(event.content)})
for event in session.events
if event.content and event.content.parts
]
Expand Down Expand Up @@ -128,7 +129,10 @@ async def add_events_to_memory(
existing_ids = {event.id for event in existing_events}
for event in events_to_add:
if event.id not in existing_ids:
existing_events.append(event)
# Snapshot the fields used by memory without copying opaque outputs.
existing_events.append(
event.model_copy(update={'content': deepcopy(event.content)})
)
existing_ids.add(event.id)
self._session_events[user_key][scoped_session_id] = existing_events

Expand Down Expand Up @@ -186,5 +190,8 @@ async def search_memory(
# so it is stable and events matching equally stay in insertion order.
scored_memories.sort(key=lambda scored_memory: -scored_memory[0])
return SearchMemoryResponse(
memories=[memory for _, memory in scored_memories[:_MAX_SEARCH_RESULTS]]
memories=[
memory.model_copy(update={'content': deepcopy(memory.content)})
for _, memory in scored_memories[:_MAX_SEARCH_RESULTS]
]
)
85 changes: 85 additions & 0 deletions tests/unittests/memory/test_in_memory_memory_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -602,3 +602,88 @@ def reader():
assert (
not errors
), f'search_memory raced with concurrent writes: {errors[0]!r}'


@pytest.mark.parametrize('ingest_session', [True, False])
async def test_ingested_memories_do_not_follow_caller_event_mutations(
ingest_session: bool,
) -> None:
"""Changing an ingested event does not rewrite the stored memory."""
service = InMemoryMemoryService()
event = Event(
id='event',
author='user',
timestamp=12345,
content=types.Content(parts=[types.Part(text='I prefer jasmine tea.')]),
)
session = Session(
app_name=MOCK_APP_NAME,
user_id=MOCK_USER_ID,
id='session',
events=[event],
)
if ingest_session:
await service.add_session_to_memory(session)
else:
await service.add_events_to_memory(
app_name=session.app_name,
user_id=session.user_id,
session_id=session.id,
events=session.events,
)

event.author = 'changed'
event.content.parts[0].text = 'I prefer coffee.'
result = await service.search_memory(
app_name=MOCK_APP_NAME, user_id=MOCK_USER_ID, query='jasmine'
)

assert len(result.memories) == 1
assert result.memories[0].author == 'user'
assert result.memories[0].content.parts[0].text == 'I prefer jasmine tea.'


@pytest.mark.parametrize('ingest_session', [True, False])
async def test_retrieved_memories_do_not_mutate_storage_or_session_history(
ingest_session: bool,
) -> None:
"""Editing a search result leaves subsequent recall and its source intact."""
service = InMemoryMemoryService()
session = Session(
app_name=MOCK_APP_NAME,
user_id=MOCK_USER_ID,
id='session',
events=[
Event(
id='event',
author='user',
timestamp=12345,
content=types.Content(
parts=[types.Part(text='I prefer jasmine tea.')]
),
),
],
)
if ingest_session:
await service.add_session_to_memory(session)
else:
await service.add_events_to_memory(
app_name=session.app_name,
user_id=session.user_id,
session_id=session.id,
events=session.events,
)
first_result = await service.search_memory(
app_name=MOCK_APP_NAME, user_id=MOCK_USER_ID, query='jasmine'
)

first_result.memories[0].content.parts.clear()
second_result = await service.search_memory(
app_name=MOCK_APP_NAME, user_id=MOCK_USER_ID, query='jasmine'
)

assert len(second_result.memories) == 1
assert (
second_result.memories[0].content.parts[0].text == 'I prefer jasmine tea.'
)
assert session.events[0].content.parts[0].text == 'I prefer jasmine tea.'
Loading