fix(flows): inject transfer_to_agent tool for HITL confirmation resume

Merge https://github.com/google/adk-python/pull/5669

Closes #5633

Co-authored-by: George Weale <gweale@google.com>
COPYBARA_INTEGRATE_REVIEW=https://github.com/google/adk-python/pull/5669 from settler-av:fix/transfer-to-agent-confirmation 91a210dd7cd5b1dcbc5a98c7c5068c73701d6325
PiperOrigin-RevId: 963694454
This commit is contained in:
Adnan Vahora
2026-08-12 15:31:23 -07:00
committed by Copybara-Service
parent 374aab372a
commit 0897bee6a0
3 changed files with 332 additions and 3 deletions
@@ -50,9 +50,7 @@ class _AgentTransferLlmRequestProcessor(BaseLlmRequestProcessor):
if not transfer_targets:
return
transfer_to_agent_tool = TransferToAgentTool(
agent_names=[agent.name for agent in transfer_targets]
)
transfer_to_agent_tool = _build_transfer_tool(transfer_targets)
llm_request.append_instructions([
_build_transfer_instructions(
@@ -205,3 +203,19 @@ def _get_transfer_targets(agent: LlmAgent) -> list[BaseAgent]:
])
return result
def _build_transfer_tool(
transfer_targets: Sequence[BaseAgent],
) -> TransferToAgentTool:
"""Builds the transfer tool offering the given agents as targets.
Args:
transfer_targets: The agents that can be transferred to.
Returns:
A TransferToAgentTool for the given targets.
"""
return TransferToAgentTool(
agent_names=[target.name for target in transfer_targets]
)
@@ -31,6 +31,8 @@ from ...tools.base_tool import BaseTool
from ...tools.tool_confirmation import ToolConfirmation
from ...tools.tool_context import ToolContext
from ._base_llm_processor import BaseLlmRequestProcessor
from .agent_transfer import _build_transfer_tool
from .agent_transfer import _get_transfer_targets
from .functions import REQUEST_CONFIRMATION_FUNCTION_CALL_NAME
if TYPE_CHECKING:
@@ -329,6 +331,14 @@ class _RequestConfirmationLlmRequestProcessor(BaseLlmRequestProcessor):
)
}
from ...agents.llm_agent import LlmAgent
if isinstance(agent, LlmAgent):
transfer_targets = _get_transfer_targets(agent)
if transfer_targets:
transfer_tool = _build_transfer_tool(transfer_targets)
tools_dict[transfer_tool.name] = transfer_tool
# Step 3: Resolve confirmation targets using extracted helper.
confirmation_fc_ids = set(confirmations_by_fc_id.keys())
tools_to_resume_with_confirmation, tools_to_resume_with_args = (
@@ -331,6 +331,311 @@ async def test_request_confirmation_processor_tool_not_confirmed():
) # tool_confirmation_dict
TRANSFER_TOOL_NAME = "transfer_to_agent"
TRANSFER_FC_ID = "transfer_fc_id"
TRANSFER_CONFIRMATION_FC_ID = "transfer_confirmation_fc_id"
def _build_transfer_confirmation_events(
confirmed: bool,
agent_name: str,
) -> list[Event]:
"""Helper to build the agent + user events for a transfer_to_agent confirmation."""
original_fc = types.FunctionCall(
name=TRANSFER_TOOL_NAME,
args={"agent_name": "sub_agent"},
id=TRANSFER_FC_ID,
)
tool_confirmation = ToolConfirmation(
confirmed=False, hint="Approve transfer?"
)
original_fc_event = Event(
author=agent_name,
content=types.Content(parts=[types.Part(function_call=original_fc)]),
)
confirmation_requested_event = Event(
author="user",
content=types.Content(
parts=[
types.Part(
function_response=types.FunctionResponse(
name=TRANSFER_TOOL_NAME,
id=TRANSFER_FC_ID,
response={"status": "waiting_for_confirm"},
)
)
]
),
actions=EventActions(
requested_tool_confirmations={TRANSFER_FC_ID: tool_confirmation}
),
)
agent_event = Event(
author=agent_name,
content=types.Content(
parts=[
types.Part(
function_call=types.FunctionCall(
name=functions.REQUEST_CONFIRMATION_FUNCTION_CALL_NAME,
args={
"originalFunctionCall": original_fc.model_dump(
exclude_none=True, by_alias=True
),
"toolConfirmation": tool_confirmation.model_dump(
by_alias=True, exclude_none=True
),
},
id=TRANSFER_CONFIRMATION_FC_ID,
)
)
]
),
)
user_confirmation = ToolConfirmation(confirmed=confirmed)
user_event = Event(
author="user",
content=types.Content(
parts=[
types.Part(
function_response=types.FunctionResponse(
name=functions.REQUEST_CONFIRMATION_FUNCTION_CALL_NAME,
id=TRANSFER_CONFIRMATION_FC_ID,
response={
"response": user_confirmation.model_dump_json()
},
)
)
]
),
)
return [
original_fc_event,
confirmation_requested_event,
agent_event,
user_event,
]
@pytest.mark.asyncio
async def test_request_confirmation_transfer_to_agent_approved():
"""Test that transfer_to_agent is injected into tools_dict when confirmed."""
sub_agent = LlmAgent(name="sub_agent", model="gemini-2.0-flash")
agent = LlmAgent(
name="orchestrator", model="gemini-2.0-flash", sub_agents=[sub_agent]
)
invocation_context = await testing_utils.create_invocation_context(
agent=agent
)
llm_request = LlmRequest()
invocation_context.session.events.extend(
_build_transfer_confirmation_events(confirmed=True, agent_name=agent.name)
)
expected_event = Event(
author="agent",
content=types.Content(
parts=[
types.Part(
function_response=types.FunctionResponse(
name=TRANSFER_TOOL_NAME,
id=TRANSFER_FC_ID,
response={},
)
)
]
),
)
with patch(
"google.adk.flows.llm_flows.functions.handle_function_call_list_async"
) as mock_handle:
mock_handle.return_value = expected_event
events = []
async for event in request_processor.run_async(
invocation_context, llm_request
):
events.append(event)
assert len(events) == 1
mock_handle.assert_called_once()
args, _ = mock_handle.call_args
tools_dict = args[2]
assert TRANSFER_TOOL_NAME in tools_dict
@pytest.mark.asyncio
async def test_request_confirmation_transfer_to_agent_rejected():
"""Test that transfer_to_agent is injected even when rejected."""
sub_agent = LlmAgent(name="sub_agent", model="gemini-2.0-flash")
agent = LlmAgent(
name="orchestrator", model="gemini-2.0-flash", sub_agents=[sub_agent]
)
invocation_context = await testing_utils.create_invocation_context(
agent=agent
)
llm_request = LlmRequest()
invocation_context.session.events.extend(
_build_transfer_confirmation_events(
confirmed=False, agent_name=agent.name
)
)
expected_event = Event(
author="agent",
content=types.Content(
parts=[
types.Part(
function_response=types.FunctionResponse(
name=TRANSFER_TOOL_NAME,
id=TRANSFER_FC_ID,
response={"error": "Tool execution not confirmed"},
)
)
]
),
)
with patch(
"google.adk.flows.llm_flows.functions.handle_function_call_list_async"
) as mock_handle:
mock_handle.return_value = expected_event
events = []
async for event in request_processor.run_async(
invocation_context, llm_request
):
events.append(event)
assert len(events) == 1
mock_handle.assert_called_once()
args, _ = mock_handle.call_args
tools_dict = args[2]
assert TRANSFER_TOOL_NAME in tools_dict
@pytest.mark.asyncio
async def test_request_confirmation_no_sub_agents_no_transfer_tool():
"""Test that transfer_to_agent is NOT injected when agent has no sub_agents."""
agent = LlmAgent(
name="test_agent", model="gemini-2.0-flash", tools=[mock_tool]
)
invocation_context = await testing_utils.create_invocation_context(
agent=agent
)
llm_request = LlmRequest()
original_fc = types.FunctionCall(
name=MOCK_TOOL_NAME, args={"param1": "test"}, id=MOCK_FUNCTION_CALL_ID
)
tool_confirmation = ToolConfirmation(confirmed=False, hint="test hint")
invocation_context.session.events.append(
Event(
author=agent.name,
content=types.Content(parts=[types.Part(function_call=original_fc)]),
)
)
invocation_context.session.events.append(
Event(
author="user",
content=types.Content(
parts=[
types.Part(
function_response=types.FunctionResponse(
name=MOCK_TOOL_NAME,
id=MOCK_FUNCTION_CALL_ID,
response={"status": "waiting_for_confirm"},
)
)
]
),
actions=EventActions(
requested_tool_confirmations={
MOCK_FUNCTION_CALL_ID: tool_confirmation
}
),
)
)
invocation_context.session.events.append(
Event(
author=agent.name,
content=types.Content(
parts=[
types.Part(
function_call=types.FunctionCall(
name=functions.REQUEST_CONFIRMATION_FUNCTION_CALL_NAME,
args={
"originalFunctionCall": original_fc.model_dump(
exclude_none=True, by_alias=True
),
"toolConfirmation": tool_confirmation.model_dump(
by_alias=True, exclude_none=True
),
},
id=MOCK_CONFIRMATION_FUNCTION_CALL_ID,
)
)
]
),
)
)
user_confirmation = ToolConfirmation(confirmed=True)
invocation_context.session.events.append(
Event(
author="user",
content=types.Content(
parts=[
types.Part(
function_response=types.FunctionResponse(
name=functions.REQUEST_CONFIRMATION_FUNCTION_CALL_NAME,
id=MOCK_CONFIRMATION_FUNCTION_CALL_ID,
response={
"response": user_confirmation.model_dump_json()
},
)
)
]
),
)
)
expected_event = Event(
author="agent",
content=types.Content(
parts=[
types.Part(
function_response=types.FunctionResponse(
name=MOCK_TOOL_NAME,
id=MOCK_FUNCTION_CALL_ID,
response={"result": "Mock tool result with test"},
)
)
]
),
)
with patch(
"google.adk.flows.llm_flows.functions.handle_function_call_list_async"
) as mock_handle:
mock_handle.return_value = expected_event
events = []
async for event in request_processor.run_async(
invocation_context, llm_request
):
events.append(event)
assert len(events) == 1
mock_handle.assert_called_once()
args, _ = mock_handle.call_args
tools_dict = args[2]
assert TRANSFER_TOOL_NAME not in tools_dict
assert MOCK_TOOL_NAME in tools_dict
@pytest.mark.asyncio
async def test_request_confirmation_processor_finds_user_confirmation_in_default_branch():
"""Processor finds user confirmation in default branch when agent is in child branch.