From ce9c39f5a85ed12c22009693b5e6bc65f4641633 Mon Sep 17 00:00:00 2001 From: Shangjie Chen Date: Mon, 29 Sep 2025 22:45:26 -0700 Subject: [PATCH] feat: Implement checkpoint and resume logic for LoopAgent PiperOrigin-RevId: 813096880 --- src/google/adk/agents/loop_agent.py | 80 ++++++++--- tests/unittests/agents/test_loop_agent.py | 136 ++++++++++++++---- .../flows/llm_flows/test_agent_transfer.py | 27 +++- .../runners/test_pause_invocation.py | 13 ++ 4 files changed, 215 insertions(+), 41 deletions(-) diff --git a/src/google/adk/agents/loop_agent.py b/src/google/adk/agents/loop_agent.py index 617749da..9965f18f 100644 --- a/src/google/adk/agents/loop_agent.py +++ b/src/google/adk/agents/loop_agent.py @@ -16,6 +16,7 @@ from __future__ import annotations +import logging from typing import Any from typing import AsyncGenerator from typing import ClassVar @@ -24,15 +25,17 @@ from typing import Optional from typing_extensions import override -from ..agents.invocation_context import InvocationContext from ..events.event import Event from ..utils.context_utils import Aclosing from ..utils.feature_decorator import experimental from .base_agent import BaseAgent from .base_agent import BaseAgentState from .base_agent_config import BaseAgentConfig +from .invocation_context import InvocationContext from .loop_agent_config import LoopAgentConfig +logger = logging.getLogger('google_adk.' + __name__) + @experimental class LoopAgentState(BaseAgentState): @@ -69,11 +72,32 @@ class LoopAgent(BaseAgent): if not self.sub_agents: return - times_looped = 0 - while not self.max_iterations or times_looped < self.max_iterations: - for sub_agent in self.sub_agents: - should_exit = False - pause_invocation = False + agent_state = self._load_agent_state(ctx, LoopAgentState) + is_resuming_at_current_agent = agent_state is not None + times_looped, start_index = self._get_start_state(agent_state) + + should_exit = False + pause_invocation = False + while ( + not self.max_iterations or times_looped < self.max_iterations + ) and not (should_exit or pause_invocation): + for i in range(start_index, len(self.sub_agents)): + sub_agent = self.sub_agents[i] + + if ctx.is_resumable and not is_resuming_at_current_agent: + # If we are resuming from the current event, it means the same event + # has already been logged, so we should avoid yielding it again. + agent_state = LoopAgentState( + current_sub_agent=sub_agent.name, + times_looped=times_looped, + ) + yield self._create_agent_state_event(ctx, agent_state=agent_state) + + # Reset the sub-agent's state in the context to ensure that each + # sub-agent starts fresh. + if not is_resuming_at_current_agent: + ctx.reset_agent_state(sub_agent.name) + is_resuming_at_current_agent = False async with Aclosing(sub_agent.run_async(ctx)) as agen: async for event in agen: @@ -83,18 +107,42 @@ class LoopAgent(BaseAgent): if ctx.should_pause_invocation(event): pause_invocation = True - # Indicates that the loop agent should exist after running this - # sub-agent. - if should_exit: - return - - # Indicates that the invocation should be paused after running this - # sub-agent. - if pause_invocation: - return + if should_exit or pause_invocation: + break # break inner for loop + # Restart from the beginning of the loop. + start_index = 0 times_looped += 1 - return + + # If the invocation is paused, we should not yield the end of agent event. + if pause_invocation: + return + + if ctx.is_resumable: + yield self._create_agent_state_event(ctx, end_of_agent=True) + + def _get_start_state( + self, + agent_state: Optional[LoopAgentState], + ) -> tuple[int, int]: + """Computes the start state of the loop agent from the agent state.""" + if not agent_state: + return 0, 0 + + times_looped = agent_state.times_looped + start_index = 0 + if agent_state.current_sub_agent: + try: + sub_agent_names = [sub_agent.name for sub_agent in self.sub_agents] + start_index = sub_agent_names.index(agent_state.current_sub_agent) + except ValueError: + # A sub-agent was removed so the agent name is not found. + # For now, we restart from the beginning. + logger.warning( + 'Sub-agent %s was not found. Restarting from the beginning.', + agent_state.current_sub_agent, + ) + return times_looped, start_index @override async def _run_live_impl( diff --git a/tests/unittests/agents/test_loop_agent.py b/tests/unittests/agents/test_loop_agent.py index 928308a4..746135d0 100644 --- a/tests/unittests/agents/test_loop_agent.py +++ b/tests/unittests/agents/test_loop_agent.py @@ -19,6 +19,8 @@ from typing import AsyncGenerator from google.adk.agents.base_agent import BaseAgent from google.adk.agents.invocation_context import InvocationContext from google.adk.agents.loop_agent import LoopAgent +from google.adk.agents.loop_agent import LoopAgentState +from google.adk.apps import ResumabilityConfig from google.adk.events.event import Event from google.adk.events.event_actions import EventActions from google.adk.sessions.in_memory_session_service import InMemorySessionService @@ -26,6 +28,10 @@ from google.genai import types import pytest from typing_extensions import override +from .. import testing_utils + +END_OF_AGENT = testing_utils.END_OF_AGENT + class _TestingAgent(BaseAgent): @@ -72,13 +78,13 @@ class _TestingAgentWithEscalateAction(BaseAgent): author=self.name, invocation_id=ctx.invocation_id, content=types.Content( - parts=[types.Part(text=f'I have done my job after escalation!!')] + parts=[types.Part(text='I have done my job after escalation!!')] ), ) async def _create_parent_invocation_context( - test_name: str, agent: BaseAgent + test_name: str, agent: BaseAgent, resumable: bool = False ) -> InvocationContext: session_service = InMemorySessionService() session = await session_service.create_session( @@ -89,11 +95,13 @@ async def _create_parent_invocation_context( agent=agent, session=session, session_service=session_service, + resumability_config=ResumabilityConfig(is_resumable=resumable), ) @pytest.mark.asyncio -async def test_run_async(request: pytest.FixtureRequest): +@pytest.mark.parametrize('resumable', [True, False]) +async def test_run_async(request: pytest.FixtureRequest, resumable: bool): agent = _TestingAgent(name=f'{request.function.__name__}_test_agent') loop_agent = LoopAgent( name=f'{request.function.__name__}_test_loop_agent', @@ -103,15 +111,60 @@ async def test_run_async(request: pytest.FixtureRequest): ], ) parent_ctx = await _create_parent_invocation_context( - request.function.__name__, loop_agent + request.function.__name__, loop_agent, resumable=resumable ) events = [e async for e in loop_agent.run_async(parent_ctx)] - assert len(events) == 2 - assert events[0].author == agent.name - assert events[1].author == agent.name - assert events[0].content.parts[0].text == f'Hello, async {agent.name}!' - assert events[1].content.parts[0].text == f'Hello, async {agent.name}!' + simplified_events = testing_utils.simplify_resumable_app_events(events) + if resumable: + expected_events = [ + ( + loop_agent.name, + {'current_sub_agent': agent.name, 'times_looped': 0}, + ), + (agent.name, f'Hello, async {agent.name}!'), + ( + loop_agent.name, + {'current_sub_agent': agent.name, 'times_looped': 1}, + ), + (agent.name, f'Hello, async {agent.name}!'), + (loop_agent.name, END_OF_AGENT), + ] + else: + expected_events = [ + (agent.name, f'Hello, async {agent.name}!'), + (agent.name, f'Hello, async {agent.name}!'), + ] + assert simplified_events == expected_events + + +@pytest.mark.asyncio +async def test_resume_async(request: pytest.FixtureRequest): + agent_1 = _TestingAgent(name=f'{request.function.__name__}_test_agent_1') + agent_2 = _TestingAgent(name=f'{request.function.__name__}_test_agent_2') + loop_agent = LoopAgent( + name=f'{request.function.__name__}_test_loop_agent', + max_iterations=2, + sub_agents=[ + agent_1, + agent_2, + ], + ) + parent_ctx = await _create_parent_invocation_context( + request.function.__name__, loop_agent, resumable=True + ) + parent_ctx.agent_states[loop_agent.name] = LoopAgentState( + current_sub_agent=agent_2.name, times_looped=1 + ).model_dump(mode='json') + + events = [e async for e in loop_agent.run_async(parent_ctx)] + + simplified_events = testing_utils.simplify_resumable_app_events(events) + expected_events = [ + (agent_2.name, f'Hello, async {agent_2.name}!'), + (loop_agent.name, END_OF_AGENT), + ] + assert simplified_events == expected_events @pytest.mark.asyncio @@ -129,7 +182,10 @@ async def test_run_async_skip_if_no_sub_agent(request: pytest.FixtureRequest): @pytest.mark.asyncio -async def test_run_async_with_escalate_action(request: pytest.FixtureRequest): +@pytest.mark.parametrize('resumable', [True, False]) +async def test_run_async_with_escalate_action( + request: pytest.FixtureRequest, resumable: bool +): non_escalating_agent = _TestingAgent( name=f'{request.function.__name__}_test_non_escalating_agent' ) @@ -144,20 +200,52 @@ async def test_run_async_with_escalate_action(request: pytest.FixtureRequest): sub_agents=[non_escalating_agent, escalating_agent, ignored_agent], ) parent_ctx = await _create_parent_invocation_context( - request.function.__name__, loop_agent + request.function.__name__, loop_agent, resumable=resumable ) events = [e async for e in loop_agent.run_async(parent_ctx)] - # Only two events are generated because the sub escalating_agent escalates. - assert len(events) == 3 - assert events[0].author == non_escalating_agent.name - assert events[1].author == escalating_agent.name - assert events[0].content.parts[0].text == ( - f'Hello, async {non_escalating_agent.name}!' - ) - assert events[1].content.parts[0].text == ( - f'Hello, async {escalating_agent.name}!' - ) - assert ( - events[2].content.parts[0].text == 'I have done my job after escalation!!' - ) + simplified_events = testing_utils.simplify_resumable_app_events(events) + + if resumable: + expected_events = [ + ( + loop_agent.name, + { + 'current_sub_agent': non_escalating_agent.name, + 'times_looped': 0, + }, + ), + ( + non_escalating_agent.name, + f'Hello, async {non_escalating_agent.name}!', + ), + ( + loop_agent.name, + {'current_sub_agent': escalating_agent.name, 'times_looped': 0}, + ), + ( + escalating_agent.name, + f'Hello, async {escalating_agent.name}!', + ), + ( + escalating_agent.name, + 'I have done my job after escalation!!', + ), + (loop_agent.name, END_OF_AGENT), + ] + else: + expected_events = [ + ( + non_escalating_agent.name, + f'Hello, async {non_escalating_agent.name}!', + ), + ( + escalating_agent.name, + f'Hello, async {escalating_agent.name}!', + ), + ( + escalating_agent.name, + 'I have done my job after escalation!!', + ), + ] + assert simplified_events == expected_events diff --git a/tests/unittests/flows/llm_flows/test_agent_transfer.py b/tests/unittests/flows/llm_flows/test_agent_transfer.py index 4f331c41..f7d17d65 100644 --- a/tests/unittests/flows/llm_flows/test_agent_transfer.py +++ b/tests/unittests/flows/llm_flows/test_agent_transfer.py @@ -14,6 +14,7 @@ from google.adk.agents.llm_agent import Agent from google.adk.agents.loop_agent import LoopAgent +from google.adk.agents.loop_agent import LoopAgentState from google.adk.agents.sequential_agent import SequentialAgent from google.adk.agents.sequential_agent import SequentialAgentState from google.adk.apps.app import App @@ -469,12 +470,36 @@ def test_auto_to_loop(is_resumable: bool): ('root_agent', transfer_call_part('sub_agent_1')), ('root_agent', TRANSFER_RESPONSE_PART), # Loops. + ( + 'sub_agent_1', + LoopAgentState(current_sub_agent='sub_agent_1_1').model_dump( + mode='json' + ), + ), ('sub_agent_1_1', 'response1'), ('sub_agent_1_1', END_OF_AGENT), + ( + 'sub_agent_1', + LoopAgentState(current_sub_agent='sub_agent_1_2').model_dump( + mode='json' + ), + ), ('sub_agent_1_2', 'response2'), ('sub_agent_1_2', END_OF_AGENT), + ( + 'sub_agent_1', + LoopAgentState( + current_sub_agent='sub_agent_1_1', times_looped=1 + ).model_dump(mode='json'), + ), ('sub_agent_1_1', 'response3'), ('sub_agent_1_1', END_OF_AGENT), + ( + 'sub_agent_1', + LoopAgentState( + current_sub_agent='sub_agent_1_2', times_looped=1 + ).model_dump(mode='json'), + ), # Exits. ('sub_agent_1_2', Part.from_function_call(name='exit_loop', args={})), ( @@ -484,7 +509,7 @@ def test_auto_to_loop(is_resumable: bool): ), ), ('sub_agent_1_2', END_OF_AGENT), - # Later expect the loop agent to also yield agent state events. + ('sub_agent_1', END_OF_AGENT), ('root_agent', END_OF_AGENT), ] # Same session, different invocation. diff --git a/tests/unittests/runners/test_pause_invocation.py b/tests/unittests/runners/test_pause_invocation.py index 8d40961b..f8f8ffc1 100644 --- a/tests/unittests/runners/test_pause_invocation.py +++ b/tests/unittests/runners/test_pause_invocation.py @@ -21,6 +21,7 @@ from google.adk.agents.base_agent import BaseAgent from google.adk.agents.invocation_context import InvocationContext from google.adk.agents.llm_agent import LlmAgent from google.adk.agents.loop_agent import LoopAgent +from google.adk.agents.loop_agent import LoopAgentState from google.adk.agents.parallel_agent import ParallelAgent from google.adk.agents.sequential_agent import SequentialAgent from google.adk.agents.sequential_agent import SequentialAgentState @@ -368,8 +369,20 @@ class TestPauseInvocationWithLoopAgent(BasePauseInvocationTest): ): """Tests that a LoopAgent pauses on long running function call.""" assert testing_utils.simplify_resumable_app_events(runner.run("test")) == [ + ( + "root_agent", + LoopAgentState(current_sub_agent="sub_agent_1").model_dump( + mode="json" + ), + ), ("sub_agent_1", "sub agent 1 response"), ("sub_agent_1", END_OF_AGENT), + ( + "root_agent", + LoopAgentState(current_sub_agent="sub_agent_2").model_dump( + mode="json" + ), + ), ("sub_agent_2", Part.from_function_call(name="test_tool", args={})), ]