Skip to content

Commit f343270

Browse files
vertex-sdk-botcopybara-github
authored andcommitted
fix: Enable automatic session creation in ADK template.
PiperOrigin-RevId: 955725374
1 parent f8eb68c commit f343270

6 files changed

Lines changed: 182 additions & 0 deletions

File tree

agentplatform/agent_engines/templates/adk.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1127,6 +1127,7 @@ def set_up(self):
11271127
artifact_service=self._tmpl_attrs.get("artifact_service"),
11281128
memory_service=self._tmpl_attrs.get("memory_service"),
11291129
credential_service=self._tmpl_attrs.get("credential_service"),
1130+
auto_create_session=True,
11301131
)
11311132
self._tmpl_attrs["in_memory_session_service"] = InMemorySessionService()
11321133
self._tmpl_attrs["in_memory_artifact_service"] = InMemoryArtifactService()

tests/unit/agentplatform/frameworks/test_frameworks_adk.py

Lines changed: 61 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -531,6 +531,67 @@ async def test_async_stream_query(
531531
events.append(event)
532532
assert len(events) == 1
533533

534+
def test_set_up_runner_auto_create_session_enabled(
535+
self,
536+
default_instrumentor_builder_mock: mock.Mock,
537+
get_project_id_mock: mock.Mock,
538+
):
539+
"""The main runner opts into auto-creating missing sessions (b/537477760)."""
540+
app = adk_template.AdkApp(agent=_TEST_AGENT)
541+
app.set_up()
542+
assert app._tmpl_attrs.get("runner").auto_create_session is True
543+
544+
@pytest.mark.asyncio
545+
async def test_runner_auto_creates_missing_session(
546+
self,
547+
default_instrumentor_builder_mock: mock.Mock,
548+
get_project_id_mock: mock.Mock,
549+
):
550+
from google.adk.runners import Runner
551+
from google.adk.sessions.in_memory_session_service import (
552+
InMemorySessionService,
553+
)
554+
555+
app = adk_template.AdkApp(agent=_TEST_AGENT)
556+
app.set_up()
557+
assert app._tmpl_attrs.get("runner").auto_create_session is True
558+
559+
app_name = app._tmpl_attrs.get("app_name")
560+
session_service = InMemorySessionService()
561+
runner = Runner(
562+
agent=_TEST_AGENT,
563+
app_name=app_name,
564+
session_service=session_service,
565+
auto_create_session=app._tmpl_attrs.get("runner").auto_create_session,
566+
)
567+
missing_session_id = "0000000000000000000"
568+
569+
# Sanity: the session does not exist yet.
570+
assert (
571+
await session_service.get_session(
572+
app_name=app_name,
573+
user_id=_TEST_USER_ID,
574+
session_id=missing_session_id,
575+
)
576+
is None
577+
)
578+
579+
session = await runner._get_or_create_session(
580+
user_id=_TEST_USER_ID,
581+
session_id=missing_session_id,
582+
)
583+
584+
assert session is not None
585+
assert session.id == missing_session_id
586+
assert (
587+
await session_service.get_session(
588+
app_name=app_name,
589+
user_id=_TEST_USER_ID,
590+
session_id=missing_session_id,
591+
)
592+
is not None
593+
)
594+
534595
@pytest.mark.asyncio
535596
async def test_async_stream_query_with_empty_session_events(
536597
self,

tests/unit/vertex_adk/test_agent_engine_templates_adk.py

Lines changed: 61 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -429,6 +429,67 @@ async def test_async_stream_query(
429429
events.append(event)
430430
assert len(events) == 1
431431

432+
def test_set_up_runner_auto_create_session_enabled(
433+
self,
434+
default_instrumentor_builder_mock: mock.Mock,
435+
get_project_id_mock: mock.Mock,
436+
):
437+
"""The main runner opts into auto-creating missing sessions (b/537477760)."""
438+
app = agent_engines.AdkApp(agent=_TEST_AGENT)
439+
app.set_up()
440+
assert app._tmpl_attrs.get("runner").auto_create_session is True
441+
442+
@pytest.mark.asyncio
443+
async def test_runner_auto_creates_missing_session(
444+
self,
445+
default_instrumentor_builder_mock: mock.Mock,
446+
get_project_id_mock: mock.Mock,
447+
):
448+
from google.adk.runners import Runner
449+
from google.adk.sessions.in_memory_session_service import (
450+
InMemorySessionService,
451+
)
452+
453+
app = agent_engines.AdkApp(agent=_TEST_AGENT)
454+
app.set_up()
455+
assert app._tmpl_attrs.get("runner").auto_create_session is True
456+
457+
app_name = app._tmpl_attrs.get("app_name")
458+
session_service = InMemorySessionService()
459+
runner = Runner(
460+
agent=_TEST_AGENT,
461+
app_name=app_name,
462+
session_service=session_service,
463+
auto_create_session=app._tmpl_attrs.get("runner").auto_create_session,
464+
)
465+
missing_session_id = "0000000000000000000"
466+
467+
# Sanity: the session does not exist yet.
468+
assert (
469+
await session_service.get_session(
470+
app_name=app_name,
471+
user_id=_TEST_USER_ID,
472+
session_id=missing_session_id,
473+
)
474+
is None
475+
)
476+
477+
session = await runner._get_or_create_session(
478+
user_id=_TEST_USER_ID,
479+
session_id=missing_session_id,
480+
)
481+
482+
assert session is not None
483+
assert session.id == missing_session_id
484+
assert (
485+
await session_service.get_session(
486+
app_name=app_name,
487+
user_id=_TEST_USER_ID,
488+
session_id=missing_session_id,
489+
)
490+
is not None
491+
)
492+
432493
@pytest.mark.asyncio
433494
@mock.patch.dict(
434495
os.environ,

tests/unit/vertex_adk/test_reasoning_engine_templates_adk.py

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -475,6 +475,63 @@ async def test_async_stream_query(self):
475475
events.append(event)
476476
assert len(events) == 1
477477

478+
def test_set_up_runner_auto_create_session_enabled(self):
479+
"""The main runner opts into auto-creating missing sessions (b/537477760)."""
480+
app = reasoning_engines.AdkApp(
481+
agent=Agent(name=_TEST_AGENT_NAME, model=_TEST_MODEL)
482+
)
483+
app.set_up()
484+
assert app._tmpl_attrs.get("runner").auto_create_session is True
485+
486+
@pytest.mark.asyncio
487+
async def test_runner_auto_creates_missing_session(self):
488+
from google.adk.runners import Runner
489+
from google.adk.sessions.in_memory_session_service import (
490+
InMemorySessionService,
491+
)
492+
493+
app = reasoning_engines.AdkApp(
494+
agent=Agent(name=_TEST_AGENT_NAME, model=_TEST_MODEL)
495+
)
496+
app.set_up()
497+
assert app._tmpl_attrs.get("runner").auto_create_session is True
498+
499+
app_name = app._tmpl_attrs.get("app_name")
500+
session_service = InMemorySessionService()
501+
runner = Runner(
502+
agent=Agent(name=_TEST_AGENT_NAME, model=_TEST_MODEL),
503+
app_name=app_name,
504+
session_service=session_service,
505+
auto_create_session=app._tmpl_attrs.get("runner").auto_create_session,
506+
)
507+
missing_session_id = "0000000000000000000"
508+
509+
# Sanity: the session does not exist yet.
510+
assert (
511+
await session_service.get_session(
512+
app_name=app_name,
513+
user_id=_TEST_USER_ID,
514+
session_id=missing_session_id,
515+
)
516+
is None
517+
)
518+
519+
session = await runner._get_or_create_session(
520+
user_id=_TEST_USER_ID,
521+
session_id=missing_session_id,
522+
)
523+
524+
assert session is not None
525+
assert session.id == missing_session_id
526+
assert (
527+
await session_service.get_session(
528+
app_name=app_name,
529+
user_id=_TEST_USER_ID,
530+
session_id=missing_session_id,
531+
)
532+
is not None
533+
)
534+
478535
@pytest.mark.asyncio
479536
async def test_async_stream_query_with_empty_session_events(self):
480537
app = reasoning_engines.AdkApp(

vertexai/agent_engines/templates/adk.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1092,6 +1092,7 @@ def set_up(self):
10921092
session_service=self._tmpl_attrs.get("session_service"),
10931093
artifact_service=self._tmpl_attrs.get("artifact_service"),
10941094
memory_service=self._tmpl_attrs.get("memory_service"),
1095+
auto_create_session=True,
10951096
)
10961097
self._tmpl_attrs["in_memory_session_service"] = InMemorySessionService()
10971098
self._tmpl_attrs["in_memory_artifact_service"] = InMemoryArtifactService()

vertexai/preview/reasoning_engines/templates/adk.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -977,6 +977,7 @@ def set_up(self):
977977
artifact_service=self._tmpl_attrs.get("artifact_service"),
978978
memory_service=self._tmpl_attrs.get("memory_service"),
979979
app_name=self._tmpl_attrs.get("app_name"),
980+
auto_create_session=True,
980981
)
981982
self._tmpl_attrs["in_memory_session_service"] = InMemorySessionService()
982983
self._tmpl_attrs["in_memory_artifact_service"] = InMemoryArtifactService()

0 commit comments

Comments
 (0)