mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
chore: Remove temp state deltas before appending an event
PiperOrigin-RevId: 816902208
This commit is contained in:
committed by
Copybara-Service
parent
cbe60c47aa
commit
24342e95f8
@@ -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'
|
||||
|
||||
Reference in New Issue
Block a user