fix: make SqliteSessionService state merges use dict.update() semantics
Merge https://github.com/google/adk-python/pull/6729 Fixes #6728 PiperOrigin-RevId: 967540198
This commit is contained in:
committed by
Copybara-Service
parent
c986ff0fce
commit
e4ba7040fb
@@ -48,6 +48,25 @@ logger = logging.getLogger("google_adk." + __name__)
|
||||
|
||||
PRAGMA_FOREIGN_KEYS = "PRAGMA foreign_keys = ON"
|
||||
|
||||
# Merges {delta} into {state} with dict.update() semantics: keys in the delta
|
||||
# always win with their delta value (including SQL NULL / JSON null), unlike
|
||||
# json_patch() which deep-merges dict values and treats null as "delete key".
|
||||
_MERGE_STATE_SQL = """
|
||||
SELECT json_group_object(
|
||||
key,
|
||||
CASE
|
||||
WHEN type IN ('object','array') THEN json(value)
|
||||
WHEN type IN ('true','false') THEN json(type)
|
||||
ELSE value
|
||||
END)
|
||||
FROM (
|
||||
SELECT key, value, type FROM json_each({delta})
|
||||
UNION ALL
|
||||
SELECT key, value, type FROM json_each({state})
|
||||
WHERE key NOT IN (SELECT key FROM json_each({delta}))
|
||||
)
|
||||
"""
|
||||
|
||||
APP_STATES_TABLE_SCHEMA = """
|
||||
CREATE TABLE IF NOT EXISTS app_states (
|
||||
app_name TEXT PRIMARY KEY,
|
||||
@@ -557,11 +576,11 @@ class SqliteSessionService(BaseSessionService):
|
||||
delta: dict[str, Any],
|
||||
now: float,
|
||||
) -> None:
|
||||
"""Atomically inserts or updates app state using json_patch."""
|
||||
"""Atomically inserts or updates app state with dict.update() semantics."""
|
||||
await db.execute(
|
||||
"""
|
||||
f"""
|
||||
INSERT INTO app_states (app_name, state, update_time) VALUES (?, ?, ?)
|
||||
ON CONFLICT(app_name) DO UPDATE SET state=json_patch(state, excluded.state), update_time=excluded.update_time
|
||||
ON CONFLICT(app_name) DO UPDATE SET state=({_MERGE_STATE_SQL.format(delta='excluded.state', state='state')}), update_time=excluded.update_time
|
||||
""",
|
||||
(app_name, json.dumps(delta), now),
|
||||
)
|
||||
@@ -574,11 +593,11 @@ class SqliteSessionService(BaseSessionService):
|
||||
delta: dict[str, Any],
|
||||
now: float,
|
||||
) -> None:
|
||||
"""Atomically inserts or updates user state using json_patch."""
|
||||
"""Atomically inserts or updates user state with dict.update() semantics."""
|
||||
await db.execute(
|
||||
"""
|
||||
f"""
|
||||
INSERT INTO user_states (app_name, user_id, state, update_time) VALUES (?, ?, ?, ?)
|
||||
ON CONFLICT(app_name, user_id) DO UPDATE SET state=json_patch(state, excluded.state), update_time=excluded.update_time
|
||||
ON CONFLICT(app_name, user_id) DO UPDATE SET state=({_MERGE_STATE_SQL.format(delta='excluded.state', state='state')}), update_time=excluded.update_time
|
||||
""",
|
||||
(app_name, user_id, json.dumps(delta), now),
|
||||
)
|
||||
@@ -592,11 +611,13 @@ class SqliteSessionService(BaseSessionService):
|
||||
delta: dict[str, Any],
|
||||
now: float,
|
||||
) -> None:
|
||||
"""Atomically updates session state using json_patch."""
|
||||
"""Atomically updates session state with dict.update() semantics."""
|
||||
await db.execute(
|
||||
"UPDATE sessions SET state=json_patch(state, ?), update_time=? WHERE"
|
||||
" app_name=? AND user_id=? AND id=?",
|
||||
"UPDATE sessions SET"
|
||||
f" state=({_MERGE_STATE_SQL.format(delta='?', state='state')}),"
|
||||
" update_time=? WHERE app_name=? AND user_id=? AND id=?",
|
||||
(
|
||||
json.dumps(delta),
|
||||
json.dumps(delta),
|
||||
now,
|
||||
app_name,
|
||||
|
||||
@@ -647,6 +647,454 @@ async def test_session_state_is_not_shared(session_service):
|
||||
assert session1b.state == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dict_valued_state_delta_replaces_stored_value(session_service):
|
||||
"""A dict-valued delta replaces the stored value, it is not deep-merged."""
|
||||
app_name = 'my_app'
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name,
|
||||
user_id='u1',
|
||||
session_id='s1',
|
||||
state={'profile': {'name': 'ada', 'role': 'admin'}},
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'profile': {'name': 'bob'}}),
|
||||
)
|
||||
await session_service.append_event(session=session, event=event)
|
||||
|
||||
reloaded = await session_service.get_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1'
|
||||
)
|
||||
assert reloaded.state.get('profile') == {'name': 'bob'}
|
||||
assert session.state.get('profile') == {'name': 'bob'}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_none_valued_state_delta_is_stored_not_dropped(session_service):
|
||||
"""A None-valued delta stores null, it does not delete the key."""
|
||||
app_name = 'my_app'
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1', state={'flag': True}
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'flag': None}),
|
||||
)
|
||||
await session_service.append_event(session=session, event=event)
|
||||
|
||||
reloaded = await session_service.get_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1'
|
||||
)
|
||||
assert 'flag' in reloaded.state
|
||||
assert reloaded.state.get('flag') is None
|
||||
assert 'flag' in session.state
|
||||
assert session.state.get('flag') is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_boolean_state_survives_unrelated_state_delta(session_service):
|
||||
"""An unrelated state_delta must not corrupt a stored boolean's type."""
|
||||
app_name = 'my_app'
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1', state={'flag': False}
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'new_flag': True}),
|
||||
)
|
||||
await session_service.append_event(session=session, event=event)
|
||||
|
||||
reloaded = await session_service.get_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1'
|
||||
)
|
||||
assert reloaded.state.get('new_flag') is True
|
||||
assert reloaded.state.get('flag') is False
|
||||
assert session.state.get('new_flag') is True
|
||||
assert session.state.get('flag') is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dict_valued_state_delta_replaces_stored_value(session_service):
|
||||
"""A dict-valued delta replaces the stored value, it is not deep-merged."""
|
||||
app_name = 'my_app'
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name,
|
||||
user_id='u1',
|
||||
session_id='s1',
|
||||
state={'profile': {'name': 'ada', 'role': 'admin'}},
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'profile': {'name': 'bob'}}),
|
||||
)
|
||||
await session_service.append_event(session=session, event=event)
|
||||
|
||||
reloaded = await session_service.get_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1'
|
||||
)
|
||||
assert reloaded.state.get('profile') == {'name': 'bob'}
|
||||
assert session.state.get('profile') == {'name': 'bob'}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_none_valued_state_delta_is_stored_not_dropped(session_service):
|
||||
"""A None-valued delta stores null, it does not delete the key."""
|
||||
app_name = 'my_app'
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1', state={'flag': True}
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'flag': None}),
|
||||
)
|
||||
await session_service.append_event(session=session, event=event)
|
||||
|
||||
reloaded = await session_service.get_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1'
|
||||
)
|
||||
assert 'flag' in reloaded.state
|
||||
assert reloaded.state.get('flag') is None
|
||||
assert 'flag' in session.state
|
||||
assert session.state.get('flag') is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_boolean_state_survives_unrelated_state_delta(session_service):
|
||||
"""An unrelated state_delta must not corrupt a stored boolean's type."""
|
||||
app_name = 'my_app'
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1', state={'flag': False}
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'new_flag': True}),
|
||||
)
|
||||
await session_service.append_event(session=session, event=event)
|
||||
|
||||
reloaded = await session_service.get_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1'
|
||||
)
|
||||
assert reloaded.state.get('new_flag') is True
|
||||
assert reloaded.state.get('flag') is False
|
||||
assert session.state.get('new_flag') is True
|
||||
assert session.state.get('flag') is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_app_state_dict_valued_delta_replaces_stored_value(
|
||||
session_service,
|
||||
):
|
||||
"""A dict-valued delta to app: state replaces it, it is not deep-merged."""
|
||||
app_name = 'my_app'
|
||||
session1 = await session_service.create_session(
|
||||
app_name=app_name,
|
||||
user_id='u1',
|
||||
session_id='s1',
|
||||
state={'app:cfg': {'name': 'ada', 'role': 'admin'}},
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'app:cfg': {'name': 'bob'}}),
|
||||
)
|
||||
await session_service.append_event(session=session1, event=event)
|
||||
|
||||
# A different user's session should see the replaced, not merged, value.
|
||||
session2 = await session_service.create_session(
|
||||
app_name=app_name, user_id='u2', session_id='s2'
|
||||
)
|
||||
assert session2.state.get('app:cfg') == {'name': 'bob'}
|
||||
assert session1.state.get('app:cfg') == {'name': 'bob'}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_state_none_valued_delta_is_stored_not_dropped(
|
||||
session_service,
|
||||
):
|
||||
"""A None-valued delta to user: state stores null, it does not delete it."""
|
||||
app_name = 'my_app'
|
||||
session1 = await session_service.create_session(
|
||||
app_name=app_name,
|
||||
user_id='u1',
|
||||
session_id='s1',
|
||||
state={'user:pref': 'dark_mode'},
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'user:pref': None}),
|
||||
)
|
||||
await session_service.append_event(session=session1, event=event)
|
||||
|
||||
# Another session for the same user should see the null, not a dropped key.
|
||||
session1b = await session_service.create_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1b'
|
||||
)
|
||||
assert 'user:pref' in session1b.state
|
||||
assert session1b.state.get('user:pref') is None
|
||||
assert 'user:pref' in session1.state
|
||||
assert session1.state.get('user:pref') is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dict_valued_state_delta_replaces_stored_value(session_service):
|
||||
"""A dict-valued delta replaces the stored value, it is not deep-merged."""
|
||||
app_name = 'my_app'
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name,
|
||||
user_id='u1',
|
||||
session_id='s1',
|
||||
state={'profile': {'name': 'ada', 'role': 'admin'}},
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'profile': {'name': 'bob'}}),
|
||||
)
|
||||
await session_service.append_event(session=session, event=event)
|
||||
|
||||
reloaded = await session_service.get_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1'
|
||||
)
|
||||
assert reloaded.state.get('profile') == {'name': 'bob'}
|
||||
assert session.state.get('profile') == {'name': 'bob'}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_none_valued_state_delta_is_stored_not_dropped(session_service):
|
||||
"""A None-valued delta stores null, it does not delete the key."""
|
||||
app_name = 'my_app'
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1', state={'flag': True}
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'flag': None}),
|
||||
)
|
||||
await session_service.append_event(session=session, event=event)
|
||||
|
||||
reloaded = await session_service.get_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1'
|
||||
)
|
||||
assert 'flag' in reloaded.state
|
||||
assert reloaded.state.get('flag') is None
|
||||
assert 'flag' in session.state
|
||||
assert session.state.get('flag') is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_boolean_state_survives_unrelated_state_delta(session_service):
|
||||
"""An unrelated state_delta must not corrupt a stored boolean's type."""
|
||||
app_name = 'my_app'
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1', state={'flag': False}
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'new_flag': True}),
|
||||
)
|
||||
await session_service.append_event(session=session, event=event)
|
||||
|
||||
reloaded = await session_service.get_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1'
|
||||
)
|
||||
assert reloaded.state.get('new_flag') is True
|
||||
assert reloaded.state.get('flag') is False
|
||||
assert session.state.get('new_flag') is True
|
||||
assert session.state.get('flag') is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_app_state_dict_valued_delta_replaces_stored_value(
|
||||
session_service,
|
||||
):
|
||||
"""A dict-valued delta to app: state replaces it, it is not deep-merged."""
|
||||
app_name = 'my_app'
|
||||
session1 = await session_service.create_session(
|
||||
app_name=app_name,
|
||||
user_id='u1',
|
||||
session_id='s1',
|
||||
state={'app:cfg': {'name': 'ada', 'role': 'admin'}},
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'app:cfg': {'name': 'bob'}}),
|
||||
)
|
||||
await session_service.append_event(session=session1, event=event)
|
||||
|
||||
# A different user's session should see the replaced, not merged, value.
|
||||
session2 = await session_service.create_session(
|
||||
app_name=app_name, user_id='u2', session_id='s2'
|
||||
)
|
||||
assert session2.state.get('app:cfg') == {'name': 'bob'}
|
||||
assert session1.state.get('app:cfg') == {'name': 'bob'}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_state_none_valued_delta_is_stored_not_dropped(
|
||||
session_service,
|
||||
):
|
||||
"""A None-valued delta to user: state stores null, it does not delete it."""
|
||||
app_name = 'my_app'
|
||||
session1 = await session_service.create_session(
|
||||
app_name=app_name,
|
||||
user_id='u1',
|
||||
session_id='s1',
|
||||
state={'user:pref': 'dark_mode'},
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'user:pref': None}),
|
||||
)
|
||||
await session_service.append_event(session=session1, event=event)
|
||||
|
||||
# Another session for the same user should see the null, not a dropped key.
|
||||
session1b = await session_service.create_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1b'
|
||||
)
|
||||
assert 'user:pref' in session1b.state
|
||||
assert session1b.state.get('user:pref') is None
|
||||
assert 'user:pref' in session1.state
|
||||
assert session1.state.get('user:pref') is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dict_valued_state_delta_replaces_stored_value(session_service):
|
||||
"""A dict-valued delta replaces the stored value, it is not deep-merged."""
|
||||
app_name = 'my_app'
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name,
|
||||
user_id='u1',
|
||||
session_id='s1',
|
||||
state={'profile': {'name': 'ada', 'role': 'admin'}},
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'profile': {'name': 'bob'}}),
|
||||
)
|
||||
await session_service.append_event(session=session, event=event)
|
||||
|
||||
reloaded = await session_service.get_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1'
|
||||
)
|
||||
assert reloaded.state.get('profile') == {'name': 'bob'}
|
||||
assert session.state.get('profile') == {'name': 'bob'}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_none_valued_state_delta_is_stored_not_dropped(session_service):
|
||||
"""A None-valued delta stores null, it does not delete the key."""
|
||||
app_name = 'my_app'
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1', state={'flag': True}
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'flag': None}),
|
||||
)
|
||||
await session_service.append_event(session=session, event=event)
|
||||
|
||||
reloaded = await session_service.get_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1'
|
||||
)
|
||||
assert 'flag' in reloaded.state
|
||||
assert reloaded.state.get('flag') is None
|
||||
assert 'flag' in session.state
|
||||
assert session.state.get('flag') is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_boolean_state_survives_unrelated_state_delta(session_service):
|
||||
"""An unrelated state_delta must not corrupt a stored boolean's type."""
|
||||
app_name = 'my_app'
|
||||
session = await session_service.create_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1', state={'flag': False}
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'new_flag': True}),
|
||||
)
|
||||
await session_service.append_event(session=session, event=event)
|
||||
|
||||
reloaded = await session_service.get_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1'
|
||||
)
|
||||
assert reloaded.state.get('new_flag') is True
|
||||
assert reloaded.state.get('flag') is False
|
||||
assert session.state.get('new_flag') is True
|
||||
assert session.state.get('flag') is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_app_state_dict_valued_delta_replaces_stored_value(
|
||||
session_service,
|
||||
):
|
||||
"""A dict-valued delta to app: state replaces it, it is not deep-merged."""
|
||||
app_name = 'my_app'
|
||||
session1 = await session_service.create_session(
|
||||
app_name=app_name,
|
||||
user_id='u1',
|
||||
session_id='s1',
|
||||
state={'app:cfg': {'name': 'ada', 'role': 'admin'}},
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'app:cfg': {'name': 'bob'}}),
|
||||
)
|
||||
await session_service.append_event(session=session1, event=event)
|
||||
|
||||
# A different user's session should see the replaced, not merged, value.
|
||||
session2 = await session_service.create_session(
|
||||
app_name=app_name, user_id='u2', session_id='s2'
|
||||
)
|
||||
assert session2.state.get('app:cfg') == {'name': 'bob'}
|
||||
assert session1.state.get('app:cfg') == {'name': 'bob'}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_state_none_valued_delta_is_stored_not_dropped(
|
||||
session_service,
|
||||
):
|
||||
"""A None-valued delta to user: state stores null, it does not delete it."""
|
||||
app_name = 'my_app'
|
||||
session1 = await session_service.create_session(
|
||||
app_name=app_name,
|
||||
user_id='u1',
|
||||
session_id='s1',
|
||||
state={'user:pref': 'dark_mode'},
|
||||
)
|
||||
event = Event(
|
||||
invocation_id='inv1',
|
||||
author='user',
|
||||
actions=EventActions(state_delta={'user:pref': None}),
|
||||
)
|
||||
await session_service.append_event(session=session1, event=event)
|
||||
|
||||
# Another session for the same user should see the null, not a dropped key.
|
||||
session1b = await session_service.create_session(
|
||||
app_name=app_name, user_id='u1', session_id='s1b'
|
||||
)
|
||||
assert 'user:pref' in session1b.state
|
||||
assert session1b.state.get('user:pref') is None
|
||||
assert 'user:pref' in session1.state
|
||||
assert session1.state.get('user:pref') is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_temp_state_is_not_persisted_in_state_or_events(session_service):
|
||||
app_name = 'my_app'
|
||||
|
||||
Reference in New Issue
Block a user