diff --git a/agentplatform/agent_engines/templates/adk.py b/agentplatform/agent_engines/templates/adk.py index 0334091544..ff3aecbb90 100644 --- a/agentplatform/agent_engines/templates/adk.py +++ b/agentplatform/agent_engines/templates/adk.py @@ -1127,6 +1127,7 @@ def set_up(self): artifact_service=self._tmpl_attrs.get("artifact_service"), memory_service=self._tmpl_attrs.get("memory_service"), credential_service=self._tmpl_attrs.get("credential_service"), + auto_create_session=True, ) self._tmpl_attrs["in_memory_session_service"] = InMemorySessionService() self._tmpl_attrs["in_memory_artifact_service"] = InMemoryArtifactService() diff --git a/tests/unit/agentplatform/frameworks/test_frameworks_adk.py b/tests/unit/agentplatform/frameworks/test_frameworks_adk.py index ad9caf5c42..d88cd56074 100644 --- a/tests/unit/agentplatform/frameworks/test_frameworks_adk.py +++ b/tests/unit/agentplatform/frameworks/test_frameworks_adk.py @@ -531,6 +531,67 @@ async def test_async_stream_query( events.append(event) assert len(events) == 1 + def test_set_up_runner_auto_create_session_enabled( + self, + default_instrumentor_builder_mock: mock.Mock, + get_project_id_mock: mock.Mock, + ): + """The main runner opts into auto-creating missing sessions (b/537477760).""" + app = adk_template.AdkApp(agent=_TEST_AGENT) + app.set_up() + assert app._tmpl_attrs.get("runner").auto_create_session is True + + @pytest.mark.asyncio + async def test_runner_auto_creates_missing_session( + self, + default_instrumentor_builder_mock: mock.Mock, + get_project_id_mock: mock.Mock, + ): + from google.adk.runners import Runner + from google.adk.sessions.in_memory_session_service import ( + InMemorySessionService, + ) + + app = adk_template.AdkApp(agent=_TEST_AGENT) + app.set_up() + assert app._tmpl_attrs.get("runner").auto_create_session is True + + app_name = app._tmpl_attrs.get("app_name") + session_service = InMemorySessionService() + runner = Runner( + agent=_TEST_AGENT, + app_name=app_name, + session_service=session_service, + auto_create_session=app._tmpl_attrs.get("runner").auto_create_session, + ) + missing_session_id = "0000000000000000000" + + # Sanity: the session does not exist yet. + assert ( + await session_service.get_session( + app_name=app_name, + user_id=_TEST_USER_ID, + session_id=missing_session_id, + ) + is None + ) + + session = await runner._get_or_create_session( + user_id=_TEST_USER_ID, + session_id=missing_session_id, + ) + + assert session is not None + assert session.id == missing_session_id + assert ( + await session_service.get_session( + app_name=app_name, + user_id=_TEST_USER_ID, + session_id=missing_session_id, + ) + is not None + ) + @pytest.mark.asyncio async def test_async_stream_query_with_empty_session_events( self, 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 b5d1185dfd..c80e96a5b5 100644 --- a/tests/unit/vertex_adk/test_agent_engine_templates_adk.py +++ b/tests/unit/vertex_adk/test_agent_engine_templates_adk.py @@ -429,6 +429,67 @@ async def test_async_stream_query( events.append(event) assert len(events) == 1 + def test_set_up_runner_auto_create_session_enabled( + self, + default_instrumentor_builder_mock: mock.Mock, + get_project_id_mock: mock.Mock, + ): + """The main runner opts into auto-creating missing sessions (b/537477760).""" + app = agent_engines.AdkApp(agent=_TEST_AGENT) + app.set_up() + assert app._tmpl_attrs.get("runner").auto_create_session is True + + @pytest.mark.asyncio + async def test_runner_auto_creates_missing_session( + self, + default_instrumentor_builder_mock: mock.Mock, + get_project_id_mock: mock.Mock, + ): + from google.adk.runners import Runner + from google.adk.sessions.in_memory_session_service import ( + InMemorySessionService, + ) + + app = agent_engines.AdkApp(agent=_TEST_AGENT) + app.set_up() + assert app._tmpl_attrs.get("runner").auto_create_session is True + + app_name = app._tmpl_attrs.get("app_name") + session_service = InMemorySessionService() + runner = Runner( + agent=_TEST_AGENT, + app_name=app_name, + session_service=session_service, + auto_create_session=app._tmpl_attrs.get("runner").auto_create_session, + ) + missing_session_id = "0000000000000000000" + + # Sanity: the session does not exist yet. + assert ( + await session_service.get_session( + app_name=app_name, + user_id=_TEST_USER_ID, + session_id=missing_session_id, + ) + is None + ) + + session = await runner._get_or_create_session( + user_id=_TEST_USER_ID, + session_id=missing_session_id, + ) + + assert session is not None + assert session.id == missing_session_id + assert ( + await session_service.get_session( + app_name=app_name, + user_id=_TEST_USER_ID, + session_id=missing_session_id, + ) + is not None + ) + @pytest.mark.asyncio @mock.patch.dict( os.environ, diff --git a/tests/unit/vertex_adk/test_reasoning_engine_templates_adk.py b/tests/unit/vertex_adk/test_reasoning_engine_templates_adk.py index 598ec89fd8..e956f06e1a 100644 --- a/tests/unit/vertex_adk/test_reasoning_engine_templates_adk.py +++ b/tests/unit/vertex_adk/test_reasoning_engine_templates_adk.py @@ -475,6 +475,63 @@ async def test_async_stream_query(self): events.append(event) assert len(events) == 1 + def test_set_up_runner_auto_create_session_enabled(self): + """The main runner opts into auto-creating missing sessions (b/537477760).""" + app = reasoning_engines.AdkApp( + agent=Agent(name=_TEST_AGENT_NAME, model=_TEST_MODEL) + ) + app.set_up() + assert app._tmpl_attrs.get("runner").auto_create_session is True + + @pytest.mark.asyncio + async def test_runner_auto_creates_missing_session(self): + from google.adk.runners import Runner + from google.adk.sessions.in_memory_session_service import ( + InMemorySessionService, + ) + + app = reasoning_engines.AdkApp( + agent=Agent(name=_TEST_AGENT_NAME, model=_TEST_MODEL) + ) + app.set_up() + assert app._tmpl_attrs.get("runner").auto_create_session is True + + app_name = app._tmpl_attrs.get("app_name") + session_service = InMemorySessionService() + runner = Runner( + agent=Agent(name=_TEST_AGENT_NAME, model=_TEST_MODEL), + app_name=app_name, + session_service=session_service, + auto_create_session=app._tmpl_attrs.get("runner").auto_create_session, + ) + missing_session_id = "0000000000000000000" + + # Sanity: the session does not exist yet. + assert ( + await session_service.get_session( + app_name=app_name, + user_id=_TEST_USER_ID, + session_id=missing_session_id, + ) + is None + ) + + session = await runner._get_or_create_session( + user_id=_TEST_USER_ID, + session_id=missing_session_id, + ) + + assert session is not None + assert session.id == missing_session_id + assert ( + await session_service.get_session( + app_name=app_name, + user_id=_TEST_USER_ID, + session_id=missing_session_id, + ) + is not None + ) + @pytest.mark.asyncio async def test_async_stream_query_with_empty_session_events(self): app = reasoning_engines.AdkApp( diff --git a/vertexai/agent_engines/templates/adk.py b/vertexai/agent_engines/templates/adk.py index 58cc8a719f..25af45c2dd 100644 --- a/vertexai/agent_engines/templates/adk.py +++ b/vertexai/agent_engines/templates/adk.py @@ -1092,6 +1092,7 @@ def set_up(self): session_service=self._tmpl_attrs.get("session_service"), artifact_service=self._tmpl_attrs.get("artifact_service"), memory_service=self._tmpl_attrs.get("memory_service"), + auto_create_session=True, ) self._tmpl_attrs["in_memory_session_service"] = InMemorySessionService() self._tmpl_attrs["in_memory_artifact_service"] = InMemoryArtifactService() diff --git a/vertexai/preview/reasoning_engines/templates/adk.py b/vertexai/preview/reasoning_engines/templates/adk.py index ec9fea0d78..4138a69ecf 100644 --- a/vertexai/preview/reasoning_engines/templates/adk.py +++ b/vertexai/preview/reasoning_engines/templates/adk.py @@ -977,6 +977,7 @@ def set_up(self): artifact_service=self._tmpl_attrs.get("artifact_service"), memory_service=self._tmpl_attrs.get("memory_service"), app_name=self._tmpl_attrs.get("app_name"), + auto_create_session=True, ) self._tmpl_attrs["in_memory_session_service"] = InMemorySessionService() self._tmpl_attrs["in_memory_artifact_service"] = InMemoryArtifactService()