feat: Make resumable llm agents yield checkpoint events

PiperOrigin-RevId: 813001108
This commit is contained in:
Xinran (Sherry) Tang
2025-09-29 17:08:58 -07:00
committed by Copybara-Service
parent 609a2358eb
commit f005414895
10 changed files with 1133 additions and 206 deletions
+22 -1
View File
@@ -242,8 +242,10 @@ class InvocationContext(BaseModel):
def user_id(self) -> str:
return self.session.user_id
def get_events(
# TODO: Move this method from invocation_context to a dedicated module.
def _get_events(
self,
*,
current_invocation: bool = False,
current_branch: bool = False,
) -> list[Event]:
@@ -304,6 +306,25 @@ class InvocationContext(BaseModel):
return False
# TODO: Move this method from invocation_context to a dedicated module.
# TODO: Converge this method with find_matching_function_call in llm_flows.
def _find_matching_function_call(
self, function_response_event: Event
) -> Optional[Event]:
"""Finds the function call event in the current invocation that matches the function response id."""
function_responses = function_response_event.get_function_responses()
if not function_responses:
return None
function_call_id = function_responses[0].id
events = self._get_events(current_invocation=True)
# The last event is function_response_event, so we search backwards from the
# one before it.
for event in reversed(events[:-1]):
if any(fc.id == function_call_id for fc in event.get_function_calls()):
return event
return None
def new_invocation_context_id() -> str:
return "e-" + str(uuid.uuid4())
+86
View File
@@ -54,6 +54,7 @@ from ..tools.tool_context import ToolContext
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 .callback_context import CallbackContext
from .invocation_context import InvocationContext
@@ -337,6 +338,20 @@ class LlmAgent(BaseAgent):
async def _run_async_impl(
self, ctx: InvocationContext
) -> AsyncGenerator[Event, None]:
agent_state = self._load_agent_state(ctx, BaseAgentState)
# If there is an sub-agent to resume, run it and then end the current
# agent.
if agent_state is not None and (
agent_to_transfer := self._get_subagent_to_resume(ctx)
):
async with Aclosing(agent_to_transfer.run_async(ctx)) as agen:
async for event in agen:
yield event
yield self._create_agent_state_event(ctx, end_of_agent=True)
return
async with Aclosing(self._llm_flow.run_async(ctx)) as agen:
async for event in agen:
self.__maybe_save_output_to_state(event)
@@ -344,6 +359,9 @@ class LlmAgent(BaseAgent):
if ctx.should_pause_invocation(event):
return
if ctx.is_resumable:
yield self._create_agent_state_event(ctx, end_of_agent=True)
@override
async def _run_live_impl(
self, ctx: InvocationContext
@@ -498,6 +516,74 @@ class LlmAgent(BaseAgent):
else:
return AutoFlow()
def _get_subagent_to_resume(
self, ctx: InvocationContext
) -> Optional[BaseAgent]:
"""Returns the sub-agent in the llm tree to resume if it exists.
There are 2 cases where we need to transfer to and resume a sub-agent:
1. The last event is a transfer to agent response from the current agent.
In this case, we need to return the agent specified in the response.
2. The last event's author isn't the current agent, or the user is
responding to another agent's tool call.
In this case, we need to return the LAST agent being transferred to
from the current agent.
"""
events = ctx._get_events(current_invocation=True, current_branch=True)
if not events:
return None
last_event = events[-1]
if last_event.author == self.name:
# Last event is from current agent. Return transfer_to_agent in the event
# if it exists, or None.
return self.__get_transfer_to_agent_or_none(last_event, self.name)
# Last event is from user or another agent.
if last_event.author == 'user':
function_call_event = ctx._find_matching_function_call(last_event)
if not function_call_event:
raise ValueError(
'No agent to transfer to for resuming agent from function response'
f' {self.name}'
)
if function_call_event.author == self.name:
# User is responding to a tool call from the current agent.
# Current agent should continue, so no sub-agent to resume.
return None
# Last event is from another agent, or from user for another agent's tool
# call. We need to find the last agent we transferred to.
for event in reversed(events):
if agent := self.__get_transfer_to_agent_or_none(event, self.name):
return agent
return None
def __get_agent_to_run(self, agent_name: str) -> BaseAgent:
"""Find the agent to run under the root agent by name."""
agent_to_run = self.root_agent.find_agent(agent_name)
if not agent_to_run:
raise ValueError(f'Agent {agent_name} not found in the agent tree.')
return agent_to_run
def __get_transfer_to_agent_or_none(
self, event: Event, from_agent: str
) -> Optional[BaseAgent]:
"""Returns the agent to run if the event is a transfer to agent response."""
function_responses = event.get_function_responses()
if not function_responses:
return None
for function_response in function_responses:
if (
function_response.name == 'transfer_to_agent'
and event.author == from_agent
and event.actions.transfer_to_agent != from_agent
):
return self.__get_agent_to_run(event.actions.transfer_to_agent)
return None
def __maybe_save_output_to_state(self, event: Event):
"""Saves the model output to state if needed."""
# skip if the event was authored by some other agent (e.g. current agent
@@ -376,6 +376,28 @@ class BaseLlmFlow(ABC):
if invocation_context.end_invocation:
return
# Resume the LLM agent based on the last event from the current branch.
# 1. User content: continue the normal flow
# 2. Function call: call the tool and get the response event.
events = invocation_context._get_events(
current_invocation=True, current_branch=True
)
if (
invocation_context.is_resumable
and events
and events[-1].get_function_calls()
):
model_response_event = events[-1]
async with Aclosing(
self._postprocess_handle_function_calls_async(
invocation_context, model_response_event, llm_request
)
) as agen:
async for event in agen:
event.id = Event.new_id()
yield event
return
# Calls the LLM.
model_response_event = Event(
id=Event.new_id(),
+2 -1
View File
@@ -135,7 +135,8 @@ def _rearrange_events_for_latest_function_response(
Returns:
A list of events with the latest function_response rearranged.
"""
if not events:
if len(events) < 2:
# No need to process, since there is no function_call.
return events
function_responses = events[-1].get_function_responses()
+10 -1
View File
@@ -606,7 +606,16 @@ class Runner:
event = find_matching_function_call(session.events)
if event and event.author:
return root_agent.find_agent(event.author)
for event in filter(lambda e: e.author != 'user', reversed(session.events)):
def _event_filter(event: Event) -> bool:
"""Filters out user-authored events and agent state change events."""
if event.author == 'user':
return False
if event.actions.agent_state is not None or event.actions.end_of_agent:
return False
return True
for event in filter(_event_filter, reversed(session.events)):
if event.author == root_agent.name:
# Found root agent.
return root_agent