chore: Remove temp state deltas before appending an event

PiperOrigin-RevId: 816902208
This commit is contained in:
Wei Sun (Jack)
2025-10-08 16:13:52 -07:00
committed by Copybara-Service
parent cbe60c47aa
commit 24342e95f8
3 changed files with 56 additions and 2 deletions
@@ -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
@@ -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
@@ -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'