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:
committed by
Copybara-Service
parent
374aab372a
commit
0897bee6a0
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user