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(),