From 24342e95f868ab6511fa4472bffcb1356b3d952c Mon Sep 17 00:00:00 2001 From: "Wei Sun (Jack)" Date: Wed, 8 Oct 2025 16:13:08 -0700 Subject: [PATCH] chore: Remove temp state deltas before appending an event PiperOrigin-RevId: 816902208 --- .../adk/sessions/base_session_service.py | 19 ++++++++-- .../adk/sessions/database_session_service.py | 3 ++ .../sessions/test_session_service.py | 36 +++++++++++++++++++ 3 files changed, 56 insertions(+), 2 deletions(-) diff --git a/src/google/adk/sessions/base_session_service.py b/src/google/adk/sessions/base_session_service.py index 25e46ba1..a76a6b9d 100644 --- a/src/google/adk/sessions/base_session_service.py +++ b/src/google/adk/sessions/base_session_service.py @@ -12,6 +12,8 @@ # See the License for the specific language governing permissions and # limitations under the License. +from __future__ import annotations + import abc from typing import Any from typing import Optional @@ -95,11 +97,24 @@ class BaseSessionService(abc.ABC): """Appends an event to a session object.""" if event.partial: return event - self.__update_session_state(session, event) + event = self._trim_temp_delta_state(event) + self._update_session_state(session, event) session.events.append(event) return event - def __update_session_state(self, session: Session, event: Event) -> None: + def _trim_temp_delta_state(self, event: Event) -> Event: + """Removes temporary state delta keys from the event.""" + if not event.actions or not event.actions.state_delta: + return event + + event.actions.state_delta = { + key: value + for key, value in event.actions.state_delta.items() + if not key.startswith(State.TEMP_PREFIX) + } + return event + + def _update_session_state(self, session: Session, event: Event) -> None: """Updates the session state based on the event.""" if not event.actions or not event.actions.state_delta: return diff --git a/src/google/adk/sessions/database_session_service.py b/src/google/adk/sessions/database_session_service.py index 959524c6..3ed87847 100644 --- a/src/google/adk/sessions/database_session_service.py +++ b/src/google/adk/sessions/database_session_service.py @@ -599,6 +599,9 @@ class DatabaseSessionService(BaseSessionService): if event.partial: return event + # Trim temp state before persisting + event = self._trim_temp_delta_state(event) + # 1. Check if timestamp is stale # 2. Update session attributes based on event config # 3. Store event to table diff --git a/tests/unittests/sessions/test_session_service.py b/tests/unittests/sessions/test_session_service.py index c2a3a1d9..2ca265cb 100644 --- a/tests/unittests/sessions/test_session_service.py +++ b/tests/unittests/sessions/test_session_service.py @@ -441,3 +441,39 @@ async def test_append_event_with_fields(service_type): retrieved_event = retrieved_session.events[0] assert retrieved_event == event + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + 'service_type', [SessionServiceType.IN_MEMORY, SessionServiceType.DATABASE] +) +async def test_append_event_should_trim_temp_delta_state(service_type): + session_service = get_session_service(service_type) + app_name = 'my_app' + user_id = 'user' + + session = await session_service.create_session( + app_name=app_name, user_id=user_id + ) + + event = Event( + invocation_id='invocation', + author='user', + content=types.Content(role='user', parts=[types.Part(text='text')]), + actions=EventActions( + state_delta={ + 'app:key': 'app_value', + 'temp:key': 'temp_value', + } + ), + ) + + await session_service.append_event(session, event) + + updated_session = await session_service.get_session( + app_name=app_name, user_id=user_id, session_id=session.id + ) + + last_event = updated_session.events[-1] + assert 'temp:key' not in last_event.actions.state_delta + assert last_event.actions.state_delta['app:key'] == 'app_value'