Skip to content

Commit f0ec997

Browse files
DeanChensjcopybara-github
authored andcommitted
fix(sessions): Further fixes for DatabaseSessionService
- Fix timezone inconsistency in append_event where it used local naive time for Postgres (now uses UTC naive). - Fix potential MissingGreenlet in create_session by generating UUID in Python and calling to_session before commit. - Add regression test for create_session. Co-authored-by: Shangjie Chen <deanchen@google.com> PiperOrigin-RevId: 933531685
1 parent 63841c3 commit f0ec997

4 files changed

Lines changed: 115 additions & 18 deletions

File tree

src/google/adk/sessions/database_session_service.py

Lines changed: 29 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
from typing import TypeVar
2727

2828
from google.adk.platform import time as platform_time
29+
from google.adk.platform import uuid as platform_uuid
2930

3031
try:
3132
from sqlalchemy import delete
@@ -434,9 +435,12 @@ async def create_session(
434435
# 4. Build the session object with generated id
435436
# 5. Return the session
436437
await self._prepare_tables()
438+
has_user_provided_id = session_id is not None
439+
if session_id is None:
440+
session_id = platform_uuid.new_uuid()
437441
schema = self._get_schema_classes()
438442
async with self._rollback_on_exception_session() as sql_session:
439-
if session_id and await sql_session.get(
443+
if has_user_provided_id and await sql_session.get(
440444
schema.StorageSession, (app_name, user_id, session_id)
441445
):
442446
raise AlreadyExistsError(
@@ -484,15 +488,17 @@ async def create_session(
484488
update_time=now,
485489
)
486490
sql_session.add(storage_session)
487-
await sql_session.commit()
488491

489492
# Merge states for response
490493
merged_state = _merge_state(
491494
storage_app_state.state, storage_user_state.state, session_state
492495
)
496+
# Call to_session before commit to avoid post-commit lazy-load.
497+
await sql_session.flush()
493498
session = storage_session.to_session(
494-
state=merged_state, is_sqlite=is_sqlite
499+
state=merged_state, is_sqlite=is_sqlite, is_postgresql=is_postgresql
495500
)
501+
await sql_session.commit()
496502
return session
497503

498504
@override
@@ -555,8 +561,12 @@ async def get_session(
555561
# Convert storage session to session
556562
events = [e.to_event() for e in reversed(storage_events)]
557563
is_sqlite = self.db_engine.dialect.name == _SQLITE_DIALECT
564+
is_postgresql = self.db_engine.dialect.name == _POSTGRESQL_DIALECT
558565
session = storage_session.to_session(
559-
state=merged_state, events=events, is_sqlite=is_sqlite
566+
state=merged_state,
567+
events=events,
568+
is_sqlite=is_sqlite,
569+
is_postgresql=is_postgresql,
560570
)
561571
return session
562572

@@ -603,12 +613,17 @@ async def list_sessions(
603613

604614
sessions = []
605615
is_sqlite = self.db_engine.dialect.name == _SQLITE_DIALECT
616+
is_postgresql = self.db_engine.dialect.name == _POSTGRESQL_DIALECT
606617
for storage_session in results:
607618
session_state = storage_session.state
608619
user_state = user_states_map.get(storage_session.user_id, {})
609620
merged_state = _merge_state(app_state, user_state, session_state)
610621
sessions.append(
611-
storage_session.to_session(state=merged_state, is_sqlite=is_sqlite)
622+
storage_session.to_session(
623+
state=merged_state,
624+
is_sqlite=is_sqlite,
625+
is_postgresql=is_postgresql,
626+
)
612627
)
613628
return ListSessionsResponse(sessions=sessions)
614629

@@ -660,6 +675,7 @@ async def append_event(self, session: Session, event: Event) -> Event:
660675
# 3. Store the new event.
661676
schema = self._get_schema_classes()
662677
is_sqlite = self.db_engine.dialect.name == _SQLITE_DIALECT
678+
is_postgresql = self.db_engine.dialect.name == _POSTGRESQL_DIALECT
663679
use_row_level_locking = self._supports_row_level_locking()
664680

665681
state_delta = event.actions.state_delta if event.actions.state_delta else {}
@@ -685,7 +701,9 @@ async def append_event(self, session: Session, event: Event) -> Event:
685701
storage_session = storage_session_result.scalars().one_or_none()
686702
if storage_session is None:
687703
raise ValueError(f"Session {session.id} not found.")
688-
storage_update_time = storage_session.get_update_timestamp(is_sqlite)
704+
storage_update_time = storage_session.get_update_timestamp(
705+
is_sqlite=is_sqlite, is_postgresql=is_postgresql
706+
)
689707
storage_update_marker = storage_session.get_update_marker()
690708

691709
storage_app_state = await _select_required_state(
@@ -751,7 +769,8 @@ async def append_event(self, session: Session, event: Event) -> Event:
751769
storage_session.state | state_deltas["session"]
752770
)
753771

754-
if is_sqlite:
772+
is_postgresql = self.db_engine.dialect.name == _POSTGRESQL_DIALECT
773+
if is_sqlite or is_postgresql:
755774
update_time = datetime.fromtimestamp(
756775
event.timestamp, timezone.utc
757776
).replace(tzinfo=None)
@@ -763,7 +782,9 @@ async def append_event(self, session: Session, event: Event) -> Event:
763782
# Read revision fields before commit. Post-commit ORM attribute access
764783
# can lazy-load expired columns and trigger MissingGreenlet with asyncpg
765784
# when pool_pre_ping is enabled.
766-
last_update_time = storage_session.get_update_timestamp(is_sqlite)
785+
last_update_time = storage_session.get_update_timestamp(
786+
is_sqlite=is_sqlite, is_postgresql=is_postgresql
787+
)
767788
storage_update_marker = storage_session.get_update_marker()
768789
await sql_session.commit()
769790

src/google/adk/sessions/schemas/v0.py

Lines changed: 18 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -172,12 +172,22 @@ def update_timestamp_tz(self) -> float:
172172
and sqlalchemy_session.bind
173173
and sqlalchemy_session.bind.dialect.name == "sqlite"
174174
)
175-
return self.get_update_timestamp(is_sqlite=is_sqlite)
175+
is_postgresql = bool(
176+
sqlalchemy_session
177+
and sqlalchemy_session.bind
178+
and sqlalchemy_session.bind.dialect.name == "postgresql"
179+
)
180+
return self.get_update_timestamp(
181+
is_sqlite=is_sqlite, is_postgresql=is_postgresql
182+
)
176183

177-
def get_update_timestamp(self, is_sqlite: bool) -> float:
184+
def get_update_timestamp(
185+
self, is_sqlite: bool = False, is_postgresql: bool = False
186+
) -> float:
178187
"""Returns the time zone aware update timestamp."""
179-
if is_sqlite:
180-
# SQLite does not support timezone. SQLAlchemy returns a naive datetime
188+
del is_sqlite, is_postgresql # Unused.
189+
if self.update_time.tzinfo is None:
190+
# SQLite and PostgreSQL do not support timezone. SQLAlchemy returns a naive datetime
181191
# object without timezone information. We need to convert it to UTC
182192
# manually.
183193
return self.update_time.replace(tzinfo=timezone.utc).timestamp()
@@ -195,6 +205,7 @@ def to_session(
195205
state: dict[str, Any] | None = None,
196206
events: list[Event] | None = None,
197207
is_sqlite: bool = False,
208+
is_postgresql: bool = False,
198209
) -> Session:
199210
"""Converts the storage session to a session object."""
200211
if state is None:
@@ -208,7 +219,9 @@ def to_session(
208219
id=self.id,
209220
state=state,
210221
events=events,
211-
last_update_time=self.get_update_timestamp(is_sqlite=is_sqlite),
222+
last_update_time=self.get_update_timestamp(
223+
is_sqlite=is_sqlite, is_postgresql=is_postgresql
224+
),
212225
)
213226
session._storage_update_marker = self.get_update_marker()
214227
return session

src/google/adk/sessions/schemas/v1.py

Lines changed: 18 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -119,12 +119,22 @@ def update_timestamp_tz(self) -> float:
119119
and sqlalchemy_session.bind
120120
and sqlalchemy_session.bind.dialect.name == "sqlite"
121121
)
122-
return self.get_update_timestamp(is_sqlite=is_sqlite)
122+
is_postgresql = bool(
123+
sqlalchemy_session
124+
and sqlalchemy_session.bind
125+
and sqlalchemy_session.bind.dialect.name == "postgresql"
126+
)
127+
return self.get_update_timestamp(
128+
is_sqlite=is_sqlite, is_postgresql=is_postgresql
129+
)
123130

124-
def get_update_timestamp(self, is_sqlite: bool) -> float:
131+
def get_update_timestamp(
132+
self, is_sqlite: bool = False, is_postgresql: bool = False
133+
) -> float:
125134
"""Returns the time zone aware update timestamp."""
126-
if is_sqlite:
127-
# SQLite does not support timezone. SQLAlchemy returns a naive datetime
135+
del is_sqlite, is_postgresql # Unused.
136+
if self.update_time.tzinfo is None:
137+
# SQLite and PostgreSQL do not support timezone. SQLAlchemy returns a naive datetime
128138
# object without timezone information. We need to convert it to UTC
129139
# manually.
130140
return self.update_time.replace(tzinfo=timezone.utc).timestamp()
@@ -142,6 +152,7 @@ def to_session(
142152
state: dict[str, Any] | None = None,
143153
events: list[Event] | None = None,
144154
is_sqlite: bool = False,
155+
is_postgresql: bool = False,
145156
) -> Session:
146157
"""Converts the storage session to a session object."""
147158
if state is None:
@@ -155,7 +166,9 @@ def to_session(
155166
id=self.id,
156167
state=state,
157168
events=events,
158-
last_update_time=self.get_update_timestamp(is_sqlite=is_sqlite),
169+
last_update_time=self.get_update_timestamp(
170+
is_sqlite=is_sqlite, is_postgresql=is_postgresql
171+
),
159172
)
160173
session._storage_update_marker = self.get_update_marker()
161174
return session

tests/unittests/sessions/test_session_service.py

Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1335,6 +1335,56 @@ def _spy_factory():
13351335
await service.close()
13361336

13371337

1338+
@pytest.mark.asyncio
1339+
async def test_create_session_reads_storage_revision_before_commit():
1340+
"""create_session captures session revision before commit completes."""
1341+
service = DatabaseSessionService('sqlite+aiosqlite:///:memory:')
1342+
await service._prepare_tables()
1343+
schema = service._get_schema_classes()
1344+
original_get_update_timestamp = schema.StorageSession.get_update_timestamp
1345+
original_get_update_marker = schema.StorageSession.get_update_marker
1346+
revision_read_state = {'committed': False, 'post_commit_reads': 0}
1347+
1348+
def _track_revision_read(original):
1349+
def wrapper(self, *args, **kwargs):
1350+
if revision_read_state['committed']:
1351+
revision_read_state['post_commit_reads'] += 1
1352+
return original(self, *args, **kwargs)
1353+
1354+
return wrapper
1355+
1356+
schema.StorageSession.get_update_timestamp = _track_revision_read(
1357+
original_get_update_timestamp
1358+
)
1359+
schema.StorageSession.get_update_marker = _track_revision_read(
1360+
original_get_update_marker
1361+
)
1362+
1363+
try:
1364+
original_factory = service.database_session_factory
1365+
1366+
def _spy_factory():
1367+
return _CommitOrderSpySession(
1368+
original_factory(),
1369+
on_committed=lambda: revision_read_state.update({'committed': True}),
1370+
)
1371+
1372+
service.database_session_factory = _spy_factory
1373+
1374+
session = await service.create_session(
1375+
app_name='app', user_id='user', session_id='s1'
1376+
)
1377+
1378+
assert revision_read_state['post_commit_reads'] == 0
1379+
assert session.last_update_time is not None
1380+
assert session._storage_update_marker is not None
1381+
finally:
1382+
schema.StorageSession.get_update_timestamp = original_get_update_timestamp
1383+
schema.StorageSession.get_update_marker = original_get_update_marker
1384+
1385+
await service.close()
1386+
1387+
13381388
@pytest.mark.asyncio
13391389
async def test_delete_session_calls_rollback_on_commit_failure():
13401390
"""Verifies that a commit failure during delete_session triggers an explicit

0 commit comments

Comments
 (0)