From 8381ad5553911678f8f605e87041957818c20328 Mon Sep 17 00:00:00 2001 From: A Vertex SDK engineer Date: Thu, 24 Sep 2026 08:50:21 -0700 Subject: [PATCH] fix: replay session_events into a scratch session and delete it afterwards `AdkApp.async_stream_query(session_events=...)` raised `AttributeError` in the `vertexai` template and, in every template, left behind the managed session it created for the replayed events. The session is now read back before the replay and deleted when the query ends. Fixes #7118, fixes #7119 PiperOrigin-RevId: 987536032 --- agentplatform/frameworks/adk.py | 55 ++++-- .../frameworks/test_frameworks_adk.py | 63 ++++++- .../test_agent_engine_templates_adk.py | 177 ++++++++++++++++++ vertexai/agent_engines/templates/adk.py | 60 ++++-- 4 files changed, 317 insertions(+), 38 deletions(-) diff --git a/agentplatform/frameworks/adk.py b/agentplatform/frameworks/adk.py index 45d755453a..bea3fdc50f 100644 --- a/agentplatform/frameworks/adk.py +++ b/agentplatform/frameworks/adk.py @@ -1209,6 +1209,8 @@ async def async_stream_query( session_events (Optional[List[Dict[str, Any]]]): Optional. The session events to use for the query. This will be used to initialize the session if `session_id` is not provided. + That session exists only for this query and is deleted once the + stream ends. run_config (Optional[Dict[str, Any]]): Optional. The run config to use for the query. If you want to pass in a `run_config` pydantic object, you can pass in a dict @@ -1244,10 +1246,15 @@ async def async_stream_query( raise ValueError( "Only one of session_id and session_events should be specified." ) + scratch_session_id = None if not session_id: session = await self.async_create_session(user_id=user_id) session_id = session["id"] if session_events is not None: + scratch_session_id = session_id + + try: + if scratch_session_id: # We allow for session_events to be an empty list. from google.adk.events.event import Event @@ -1255,7 +1262,7 @@ async def async_stream_query( session_obj = await session_service.get_session( app_name=self._app_name(), user_id=user_id, - session_id=session_id, + session_id=scratch_session_id, ) for event in session_events: if not isinstance(event, Event): @@ -1265,28 +1272,40 @@ async def async_stream_query( event=event, ) - run_config = _validate_run_config(run_config) - if run_config: - events_async = self._tmpl_attrs.get("runner").run_async( - user_id=user_id, - session_id=session_id, - new_message=content, - run_config=run_config, - **kwargs, - ) - else: - events_async = self._tmpl_attrs.get("runner").run_async( - user_id=user_id, - session_id=session_id, - new_message=content, - **kwargs, - ) + run_config = _validate_run_config(run_config) + if run_config: + events_async = self._tmpl_attrs.get("runner").run_async( + user_id=user_id, + session_id=session_id, + new_message=content, + run_config=run_config, + **kwargs, + ) + else: + events_async = self._tmpl_attrs.get("runner").run_async( + user_id=user_id, + session_id=session_id, + new_message=content, + **kwargs, + ) - try: async for event in events_async: # Yield the event data as a dictionary yield _runtimes_utils.dump_event_for_json(event) finally: + # The caller never sees the id of a session created for + # session_events, so nothing else can delete it. + if scratch_session_id: + try: + await self._tmpl_attrs.get("session_service").delete_session( + app_name=self._app_name(), + user_id=user_id, + session_id=scratch_session_id, + ) + except Exception as e: # pylint: disable=broad-exception-caught + _warn( + f"Failed to delete scratch session {scratch_session_id}: {e}" + ) # Avoid telemetry data loss having to do with CPU throttling on instance turndown _ = await _force_flush_otel( tracing_enabled=self._tracing_enabled(), diff --git a/tests/unit/agentplatform/frameworks/test_frameworks_adk.py b/tests/unit/agentplatform/frameworks/test_frameworks_adk.py index 11932c7c1b..d646a5123c 100644 --- a/tests/unit/agentplatform/frameworks/test_frameworks_adk.py +++ b/tests/unit/agentplatform/frameworks/test_frameworks_adk.py @@ -612,7 +612,7 @@ async def test_async_stream_query_with_empty_session_events( events.append(event) assert app._tmpl_attrs.get("session_service") is not None sessions = app.list_sessions(user_id=_TEST_USER_ID) - assert len(sessions.sessions) == 1 + assert not sessions.sessions @pytest.mark.asyncio async def test_async_stream_query_with_session_events( @@ -633,7 +633,66 @@ async def test_async_stream_query_with_session_events( events.append(event) assert app._tmpl_attrs.get("session_service") is not None sessions = app.list_sessions(user_id=_TEST_USER_ID) - assert len(sessions.sessions) == 1 + assert not sessions.sessions + + @pytest.mark.asyncio + async def test_async_stream_query_deletes_session_when_replay_fails( + self, + default_instrumentor_builder_mock: mock.Mock, + get_project_id_mock: mock.Mock, + ): + app = adk_template.AdkApp(agent=Agent(name=_TEST_AGENT_NAME, model=_TEST_MODEL)) + app.set_up() + app._tmpl_attrs["runner"] = _MockRunner() + with pytest.raises(ValueError): + async for _ in app.async_stream_query( + user_id=_TEST_USER_ID, + session_events=[123], + message="test message", + ): + pass + assert not app.list_sessions(user_id=_TEST_USER_ID).sessions + + @pytest.mark.asyncio + async def test_async_stream_query_keeps_session_without_session_events( + self, + default_instrumentor_builder_mock: mock.Mock, + get_project_id_mock: mock.Mock, + ): + app = adk_template.AdkApp(agent=Agent(name=_TEST_AGENT_NAME, model=_TEST_MODEL)) + app.set_up() + app._tmpl_attrs["runner"] = _MockRunner() + async for _ in app.async_stream_query( + user_id=_TEST_USER_ID, + message="test message", + ): + pass + assert len(app.list_sessions(user_id=_TEST_USER_ID).sessions) == 1 + + @pytest.mark.asyncio + async def test_async_stream_query_warns_when_session_delete_fails( + self, + default_instrumentor_builder_mock: mock.Mock, + get_project_id_mock: mock.Mock, + ): + app = adk_template.AdkApp(agent=Agent(name=_TEST_AGENT_NAME, model=_TEST_MODEL)) + app.set_up() + app._tmpl_attrs["runner"] = _MockRunner() + with mock.patch.object( + app._tmpl_attrs["session_service"], + "delete_session", + side_effect=RuntimeError("delete failed"), + ), mock.patch.object(adk_template, "_warn") as warn_mock: + events = [ + event + async for event in app.async_stream_query( + user_id=_TEST_USER_ID, + session_events=[], + message="test message", + ) + ] + assert len(events) == 1 + warn_mock.assert_called_once() @pytest.mark.asyncio @mock.patch.dict( diff --git a/tests/unit/vertex_adk/test_agent_engine_templates_adk.py b/tests/unit/vertex_adk/test_agent_engine_templates_adk.py index a0d36dc1ac..716f84e98d 100644 --- a/tests/unit/vertex_adk/test_agent_engine_templates_adk.py +++ b/tests/unit/vertex_adk/test_agent_engine_templates_adk.py @@ -70,6 +70,83 @@ def __init__(self, name: str, model: str): _TEST_USER_ID = "test_user_id" _TEST_AGENT_NAME = "test_agent" _TEST_AGENT = Agent(name=_TEST_AGENT_NAME, model=_TEST_MODEL) +_TEST_SESSION_EVENTS = [ + { + "author": "user", + "content": { + "parts": [ + { + "text": "What is the exchange rate from US dollars to " + "Swedish krona on 2025-09-25?" + } + ], + "role": "user", + }, + "id": "8967297909049524224", + "invocationId": "e-308f65d7-a99f-41e3-b80d-40feb5f1b065", + "timestamp": 1765832134.629513, + }, + { + "author": "currency_exchange_agent", + "content": { + "parts": [ + { + "functionCall": { + "args": { + "currency_date": "2025-09-25", + "currency_from": "USD", + "currency_to": "SEK", + }, + "id": "adk-136738ad-9e57-4cfb-8e23-b0f3e50a37d7", + "name": "get_exchange_rate", + } + } + ], + "role": "model", + }, + "id": "3155402589927899136", + "invocationId": "e-308f65d7-a99f-41e3-b80d-40feb5f1b065", + "timestamp": 1765832134.723713, + }, + { + "author": "currency_exchange_agent", + "content": { + "parts": [ + { + "functionResponse": { + "id": "adk-136738ad-9e57-4cfb-8e23-b0f3e50a37d7", + "name": "get_exchange_rate", + "response": { + "amount": 1, + "base": "USD", + "date": "2025-09-25", + "rates": {"SEK": 9.4118}, + }, + } + } + ], + "role": "user", + }, + "id": "1678221912150376448", + "invocationId": "e-308f65d7-a99f-41e3-b80d-40feb5f1b065", + "timestamp": 1765832135.764961, + }, + { + "author": "currency_exchange_agent", + "content": { + "parts": [ + { + "text": "The exchange rate from US dollars to Swedish " + "krona on 2025-09-25 is 1 USD to 9.4118 SEK." + } + ], + "role": "model", + }, + "id": "2470855446567583744", + "invocationId": "e-308f65d7-a99f-41e3-b80d-40feb5f1b065", + "timestamp": 1765832135.853299, + }, +] _TEST_SESSION = { "id": "ca18c25a-644b-4e13-9b24-78c150ec3eb9", "app_name": "default_app_name", @@ -317,6 +394,23 @@ async def run_async(self, *args, **kwargs): ) +class _SessionReadingRunner: + """Records which events the session holds when the run starts.""" + + def __init__(self, session_service, app_name): + self._session_service = session_service + self._app_name = app_name + self.seen_event_ids = None + + async def run_async(self, *, user_id, session_id, **kwargs): + session = await self._session_service.get_session( + app_name=self._app_name, user_id=user_id, session_id=session_id + ) + self.seen_event_ids = [event.id for event in session.events] + async for event in _MockRunner().run_async(): + yield event + + @pytest.mark.usefixtures("google_auth_mock") class TestAdkApp: def test_adk_version(self): @@ -436,6 +530,89 @@ async def test_async_stream_query( events.append(event) assert len(events) == 1 + @pytest.mark.asyncio + async def test_async_stream_query_replays_session_events_then_deletes_session( + self, + default_instrumentor_builder_mock: mock.Mock, + get_project_id_mock: mock.Mock, + ): + app = agent_engines.AdkApp(agent=_TEST_AGENT) + app.set_up() + runner = _SessionReadingRunner( + app._tmpl_attrs["session_service"], app._app_name() + ) + app._tmpl_attrs["runner"] = runner + events = [ + event + async for event in app.async_stream_query( + user_id=_TEST_USER_ID, + session_events=_TEST_SESSION_EVENTS, + message="on the day after that?", + ) + ] + assert len(events) == 1 + assert runner.seen_event_ids == [e["id"] for e in _TEST_SESSION_EVENTS] + assert not app.list_sessions(user_id=_TEST_USER_ID).sessions + + @pytest.mark.asyncio + async def test_async_stream_query_deletes_session_when_replay_fails( + self, + default_instrumentor_builder_mock: mock.Mock, + get_project_id_mock: mock.Mock, + ): + app = agent_engines.AdkApp(agent=_TEST_AGENT) + app.set_up() + app._tmpl_attrs["runner"] = _MockRunner() + with pytest.raises(ValueError): + async for _ in app.async_stream_query( + user_id=_TEST_USER_ID, + session_events=[123], + message="test message", + ): + pass + assert not app.list_sessions(user_id=_TEST_USER_ID).sessions + + @pytest.mark.asyncio + async def test_async_stream_query_keeps_session_without_session_events( + self, + default_instrumentor_builder_mock: mock.Mock, + get_project_id_mock: mock.Mock, + ): + app = agent_engines.AdkApp(agent=_TEST_AGENT) + app.set_up() + app._tmpl_attrs["runner"] = _MockRunner() + async for _ in app.async_stream_query( + user_id=_TEST_USER_ID, + message="test message", + ): + pass + assert len(app.list_sessions(user_id=_TEST_USER_ID).sessions) == 1 + + @pytest.mark.asyncio + async def test_async_stream_query_warns_when_session_delete_fails( + self, + default_instrumentor_builder_mock: mock.Mock, + get_project_id_mock: mock.Mock, + ): + app = agent_engines.AdkApp(agent=_TEST_AGENT) + app.set_up() + app._tmpl_attrs["runner"] = _MockRunner() + with mock.patch.object( + app._tmpl_attrs["session_service"], + "delete_session", + side_effect=RuntimeError("delete failed"), + ), mock.patch.object(adk_template, "_warn") as warn_mock: + events = [ + event + async for event in app.async_stream_query( + user_id=_TEST_USER_ID, + session_events=[], + message="test message", + ) + ] + assert len(events) == 1 + warn_mock.assert_called_once() + def test_set_up_runner_auto_create_session_enabled( self, default_instrumentor_builder_mock: mock.Mock, diff --git a/vertexai/agent_engines/templates/adk.py b/vertexai/agent_engines/templates/adk.py index a2e700cdc2..fc1907e342 100644 --- a/vertexai/agent_engines/templates/adk.py +++ b/vertexai/agent_engines/templates/adk.py @@ -1180,6 +1180,8 @@ async def async_stream_query( session_events (Optional[List[Dict[str, Any]]]): Optional. The session events to use for the query. This will be used to initialize the session if `session_id` is not provided. + That session exists only for this query and is deleted once the + stream ends. run_config (Optional[Dict[str, Any]]): Optional. The run config to use for the query. If you want to pass in a `run_config` pydantic object, you can pass in a dict @@ -1215,44 +1217,66 @@ async def async_stream_query( raise ValueError( "Only one of session_id and session_events should be specified." ) + scratch_session_id = None if not session_id: session = await self.async_create_session(user_id=user_id) session_id = session["id"] if session_events is not None: + scratch_session_id = session_id + + try: + if scratch_session_id: # We allow for session_events to be an empty list. from google.adk.events.event import Event session_service = self._tmpl_attrs.get("session_service") + session_obj = await session_service.get_session( + app_name=self._app_name(), + user_id=user_id, + session_id=scratch_session_id, + ) for event in session_events: if not isinstance(event, Event): event = Event.model_validate(event) await session_service.append_event( - session=session, + session=session_obj, event=event, ) - run_config = _validate_run_config(run_config) - if run_config: - events_async = self._tmpl_attrs.get("runner").run_async( - user_id=user_id, - session_id=session_id, - new_message=content, - run_config=run_config, - **kwargs, - ) - else: - events_async = self._tmpl_attrs.get("runner").run_async( - user_id=user_id, - session_id=session_id, - new_message=content, - **kwargs, - ) + run_config = _validate_run_config(run_config) + if run_config: + events_async = self._tmpl_attrs.get("runner").run_async( + user_id=user_id, + session_id=session_id, + new_message=content, + run_config=run_config, + **kwargs, + ) + else: + events_async = self._tmpl_attrs.get("runner").run_async( + user_id=user_id, + session_id=session_id, + new_message=content, + **kwargs, + ) - try: async for event in events_async: # Yield the event data as a dictionary yield _utils.dump_event_for_json(event) finally: + # The caller never sees the id of a session created for + # session_events, so nothing else can delete it. + if scratch_session_id: + try: + await self._tmpl_attrs.get("session_service").delete_session( + app_name=self._app_name(), + user_id=user_id, + session_id=scratch_session_id, + ) + except Exception as e: # pylint: disable=broad-exception-caught + _warn( + f"Failed to delete scratch session {scratch_session_id}: {e}" + ) # Avoid telemetry data loss having to do with CPU throttling on instance turndown _ = await _force_flush_otel( tracing_enabled=self._tracing_enabled(),