diff --git a/src/google/adk/sessions/database_session_service.py b/src/google/adk/sessions/database_session_service.py index 24f525ba..6b19464e 100644 --- a/src/google/adk/sessions/database_session_service.py +++ b/src/google/adk/sessions/database_session_service.py @@ -531,6 +531,16 @@ class DatabaseSessionService(BaseSessionService): schema = self._get_schema_classes() is_sqlite = self.db_engine.dialect.name == _SQLITE_DIALECT use_row_level_locking = self._supports_row_level_locking() + + state_delta = ( + event.actions.state_delta + if event.actions and event.actions.state_delta + else {} + ) + state_deltas = _session_util.extract_state_delta(state_delta) + has_app_delta = bool(state_deltas["app"]) + has_user_delta = bool(state_deltas["user"]) + async with self._with_session_lock( app_name=session.app_name, user_id=session.user_id, @@ -554,7 +564,7 @@ class DatabaseSessionService(BaseSessionService): sql_session=sql_session, state_model=schema.StorageAppState, predicates=(schema.StorageAppState.app_name == session.app_name,), - use_row_level_locking=use_row_level_locking, + use_row_level_locking=use_row_level_locking and has_app_delta, missing_message=( "App state missing for app_name=" f"{session.app_name!r}. Session state tables should be " @@ -568,7 +578,7 @@ class DatabaseSessionService(BaseSessionService): schema.StorageUserState.app_name == session.app_name, schema.StorageUserState.user_id == session.user_id, ), - use_row_level_locking=use_row_level_locking, + use_row_level_locking=use_row_level_locking and has_user_delta, missing_message=( "User state missing for app_name=" f"{session.app_name!r}, user_id={session.user_id!r}. " @@ -599,23 +609,19 @@ class DatabaseSessionService(BaseSessionService): storage_events = [e async for e in result] session.events = [e.to_event() for e in storage_events] - # Extract state delta - if event.actions and event.actions.state_delta: - state_deltas = _session_util.extract_state_delta( - event.actions.state_delta + # Merge pre-extracted state deltas into storage. + if has_app_delta: + storage_app_state.state = ( + storage_app_state.state | state_deltas["app"] + ) + if has_user_delta: + storage_user_state.state = ( + storage_user_state.state | state_deltas["user"] + ) + if state_deltas["session"]: + storage_session.state = ( + storage_session.state | state_deltas["session"] ) - app_state_delta = state_deltas["app"] - user_state_delta = state_deltas["user"] - session_state_delta = state_deltas["session"] - # Merge state and update storage - if app_state_delta: - storage_app_state.state = storage_app_state.state | app_state_delta - if user_state_delta: - storage_user_state.state = ( - storage_user_state.state | user_state_delta - ) - if session_state_delta: - storage_session.state = storage_session.state | session_state_delta if is_sqlite: update_time = datetime.fromtimestamp( diff --git a/tests/unittests/sessions/test_session_service.py b/tests/unittests/sessions/test_session_service.py index 25530bed..4e277195 100644 --- a/tests/unittests/sessions/test_session_service.py +++ b/tests/unittests/sessions/test_session_service.py @@ -1153,3 +1153,92 @@ async def test_prepare_tables_idempotent_after_creation(): assert session.id == 's1' finally: await service.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + 'state_delta, expect_app_lock, expect_user_lock', + [ + pytest.param( + None, + False, + False, + id='no_state_delta', + ), + pytest.param( + {'session_key': 'v'}, + False, + False, + id='session_only_delta', + ), + pytest.param( + {'app:key': 'v'}, + True, + False, + id='app_delta_only', + ), + pytest.param( + {'user:key': 'v'}, + False, + True, + id='user_delta_only', + ), + pytest.param( + {'app:a': '1', 'user:b': '2', 'sk': '3'}, + True, + True, + id='all_scopes', + ), + ], +) +async def test_append_event_locks_only_scopes_with_deltas( + state_delta, expect_app_lock, expect_user_lock +): + """FOR UPDATE should only be requested for state scopes that have deltas.""" + service = DatabaseSessionService('sqlite+aiosqlite:///:memory:') + + lock_requests = [] + original_fn = database_session_service._select_required_state + + async def tracking_fn(**kwargs): + lock_requests.append({ + 'model': kwargs['state_model'].__tablename__, + 'use_row_level_locking': kwargs['use_row_level_locking'], + }) + return await original_fn(**kwargs) + + try: + session = await service.create_session( + app_name='app', user_id='user', session_id='s1' + ) + + database_session_service._select_required_state = tracking_fn + lock_requests.clear() + + event_kwargs = {'invocation_id': 'inv', 'author': 'user'} + if state_delta is not None: + event_kwargs['actions'] = EventActions(state_delta=state_delta) + event = Event(**event_kwargs) + await service.append_event(session, event) + + app_req = next( + (r for r in lock_requests if r['model'] == 'app_states'), None + ) + user_req = next( + (r for r in lock_requests if r['model'] == 'user_states'), None + ) + + # SQLite doesn't support row-level locking so use_row_level_locking is + # always False. The important check is that locking is not requested + # when there is no delta (it must never be True without a delta). + if not expect_app_lock: + assert ( + app_req is None or not app_req['use_row_level_locking'] + ), 'app_states should not be locked without an app: delta' + if not expect_user_lock: + assert ( + user_req is None or not user_req['use_row_level_locking'] + ), 'user_states should not be locked without a user: delta' + finally: + database_session_service._select_required_state = original_fn + await service.close()