Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 37 additions & 18 deletions agentplatform/frameworks/adk.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -1244,18 +1246,23 @@ 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=session_id,
session_id=scratch_session_id,
)
for event in session_events:
if not isinstance(event, Event):
Expand All @@ -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(),
Expand Down
63 changes: 61 additions & 2 deletions tests/unit/agentplatform/frameworks/test_frameworks_adk.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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(
Expand Down
177 changes: 177 additions & 0 deletions tests/unit/vertex_adk/test_agent_engine_templates_adk.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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,
Expand Down
Loading
Loading