diff --git a/src/google/adk/memory/in_memory_memory_service.py b/src/google/adk/memory/in_memory_memory_service.py index 124915314d..e85496f757 100644 --- a/src/google/adk/memory/in_memory_memory_service.py +++ b/src/google/adk/memory/in_memory_memory_service.py @@ -15,6 +15,7 @@ from collections.abc import Mapping from collections.abc import Sequence +from copy import deepcopy import itertools import re import threading @@ -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 ] @@ -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 @@ -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] + ] ) diff --git a/tests/unittests/memory/test_in_memory_memory_service.py b/tests/unittests/memory/test_in_memory_memory_service.py index 1e95bfceb1..93dea40e34 100644 --- a/tests/unittests/memory/test_in_memory_memory_service.py +++ b/tests/unittests/memory/test_in_memory_memory_service.py @@ -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.'