Skip to content

Commit 0a70337

Browse files
GWealecopybara-github
authored andcommitted
fix: raise SessionNotFoundError when appending to a missing session
Co-authored-by: George Weale <gweale@google.com> PiperOrigin-RevId: 954860721
1 parent 9adf011 commit 0a70337

6 files changed

Lines changed: 57 additions & 4 deletions

File tree

src/google/adk/errors/session_not_found_error.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -21,5 +21,5 @@ class SessionNotFoundError(ValueError):
2121
Inherits from ValueError (for backward compatibility).
2222
"""
2323

24-
def __init__(self, message="Session not found."):
24+
def __init__(self, message: str = "Session not found.") -> None:
2525
super().__init__(message)

src/google/adk/integrations/firestore/firestore_session_service.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,7 @@
2929
from typing import Optional
3030

3131
from ...errors.already_exists_error import AlreadyExistsError
32+
from ...errors.session_not_found_error import SessionNotFoundError
3233
from ...events.event import Event
3334
from ...platform import uuid as platform_uuid
3435
from ...sessions import _session_util
@@ -504,7 +505,7 @@ async def _append_txn(transaction: firestore.AsyncTransaction) -> int:
504505
# 1. Reads
505506
session_snap = await session_ref.get(transaction=transaction)
506507
if not session_snap.exists:
507-
raise ValueError(f"Session {session.id} not found.")
508+
raise SessionNotFoundError(f"Session {session.id} not found.")
508509

509510
session_doc = session_snap.to_dict() or {}
510511
if session_doc.get("status") == "DELETING":

src/google/adk/sessions/database_session_service.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -49,6 +49,7 @@
4949

5050
from . import _session_util
5151
from ..errors.already_exists_error import AlreadyExistsError
52+
from ..errors.session_not_found_error import SessionNotFoundError
5253
from ..events.event import Event
5354
from .base_session_service import BaseSessionService
5455
from .base_session_service import GetSessionConfig
@@ -779,7 +780,7 @@ async def append_event(self, session: Session, event: Event) -> Event:
779780
storage_session_result = await sql_session.execute(storage_session_stmt)
780781
storage_session = storage_session_result.scalars().one_or_none()
781782
if storage_session is None:
782-
raise ValueError(f"Session {session.id} not found.")
783+
raise SessionNotFoundError(f"Session {session.id} not found.")
783784
storage_update_time = storage_session.get_update_timestamp(
784785
is_sqlite=is_sqlite, is_postgresql=is_postgresql
785786
)

src/google/adk/sessions/sqlite_session_service.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@
3131

3232
from . import _session_util
3333
from ..errors.already_exists_error import AlreadyExistsError
34+
from ..errors.session_not_found_error import SessionNotFoundError
3435
from ..events.event import Event
3536
from .base_session_service import BaseSessionService
3637
from .base_session_service import GetSessionConfig
@@ -388,7 +389,7 @@ async def append_event(self, session: Session, event: Event) -> Event:
388389
) as cursor:
389390
row = await cursor.fetchone()
390391
if row is None:
391-
raise ValueError(f"Session {session.id} not found.")
392+
raise SessionNotFoundError(f"Session {session.id} not found.")
392393
storage_update_time = row["update_time"]
393394
if storage_update_time > session.last_update_time:
394395
raise ValueError(

tests/unittests/integrations/firestore/test_firestore_session_service.py

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
from unittest import mock
2121

2222
from google.adk.errors.already_exists_error import AlreadyExistsError
23+
from google.adk.errors.session_not_found_error import SessionNotFoundError
2324
from google.adk.events.event import Event
2425
from google.adk.events.event import EventActions
2526
from google.adk.integrations.firestore.firestore_session_service import FirestoreSessionService
@@ -268,6 +269,28 @@ async def test_append_event(mock_firestore_client):
268269
assert session.last_update_time == event.timestamp
269270

270271

272+
@pytest.mark.asyncio
273+
async def test_append_event_session_not_found(mock_firestore_client):
274+
service = FirestoreSessionService(client=mock_firestore_client)
275+
session = Session(id="test_session", app_name="test_app", user_id="test_user")
276+
event = Event(invocation_id="test_inv", author="user")
277+
278+
session_doc_snapshot = mock.MagicMock()
279+
session_doc_snapshot.exists = False
280+
281+
root_coll = mock_firestore_client.collection.return_value
282+
app_ref = root_coll.document.return_value
283+
users_coll = app_ref.collection.return_value
284+
user_ref = users_coll.document.return_value
285+
sessions_ref = user_ref.collection.return_value
286+
session_doc_ref = sessions_ref.document.return_value
287+
session_doc_ref.get = mock.AsyncMock(return_value=session_doc_snapshot)
288+
289+
with mock.patch("google.cloud.firestore.async_transactional", lambda x: x):
290+
with pytest.raises(SessionNotFoundError):
291+
await service.append_event(session, event)
292+
293+
271294
@pytest.mark.asyncio
272295
async def test_append_event_with_state_delta(mock_firestore_client):
273296
service = FirestoreSessionService(client=mock_firestore_client)

tests/unittests/sessions/test_session_service.py

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
from unittest import mock
2323

2424
from google.adk.errors.already_exists_error import AlreadyExistsError
25+
from google.adk.errors.session_not_found_error import SessionNotFoundError
2526
from google.adk.events.event import Event
2627
from google.adk.events.event_actions import EventActions
2728
from google.adk.features import FeatureName
@@ -802,6 +803,32 @@ async def test_session_last_update_time_updates_on_event(session_service):
802803
assert refreshed_session.last_update_time > original_update_time
803804

804805

806+
@pytest.mark.asyncio
807+
@pytest.mark.parametrize(
808+
'service_type', [SessionServiceType.DATABASE, SessionServiceType.SQLITE]
809+
)
810+
async def test_append_event_to_deleted_session_raises_session_not_found(
811+
service_type, tmp_path
812+
):
813+
session_service = get_session_service(service_type, tmp_path)
814+
try:
815+
app_name = 'my_app'
816+
user_id = 'user'
817+
session = await session_service.create_session(
818+
app_name=app_name, user_id=user_id
819+
)
820+
await session_service.delete_session(
821+
app_name=app_name, user_id=user_id, session_id=session.id
822+
)
823+
824+
event = Event(invocation_id='inv1', author='user')
825+
with pytest.raises(SessionNotFoundError):
826+
await session_service.append_event(session, event)
827+
finally:
828+
if isinstance(session_service, DatabaseSessionService):
829+
await session_service.close()
830+
831+
805832
@pytest.mark.asyncio
806833
async def test_append_event_to_stale_session():
807834
session_service = get_session_service(

0 commit comments

Comments
 (0)