Skip to content
Draft
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
9 changes: 3 additions & 6 deletions src/google/adk/flows/llm_flows/base_llm_flow.py
Original file line number Diff line number Diff line change
Expand Up @@ -678,12 +678,9 @@ async def run_live(
# the same function response. By handling agent transfer here,
# we ensure that only child agent processes its own function
# responses after the transfer.
if (
event.content
and event.content.parts
and event.content.parts[0].function_response
and event.content.parts[0].function_response.name
== 'transfer_to_agent'
if any(
function_response.name == 'transfer_to_agent'
for function_response in event.get_function_responses()
):
await asyncio.sleep(DEFAULT_TRANSFER_AGENT_DELAY)
# cancel the tasks that belongs to the closed connection.
Expand Down
52 changes: 28 additions & 24 deletions tests/unittests/flows/llm_flows/test_base_llm_flow.py
Original file line number Diff line number Diff line change
Expand Up @@ -1242,9 +1242,20 @@ async def mock_receive():
assert events[2].author == 'user'


@pytest.mark.parametrize(
('function_response_names', 'should_transfer'),
[
(('set_state', 'transfer_to_agent'), True),
(('transfer_to_agent', 'set_state'), True),
(('transfer_to_agent',), True),
(('set_state', 'other_tool'), False),
],
)
@pytest.mark.asyncio
async def test_run_live_clears_resumption_handle_on_transfer():
"""Test that run_live clears session resumption handles when transferring to another agent."""
async def test_run_live_transfers_regardless_of_function_response_order(
function_response_names: tuple[str, ...], should_transfer: bool
):
"""Live transfer finds its response regardless of parallel result order."""

agent = Agent(name='test_agent')
invocation_context = await testing_utils.create_invocation_context(
Expand All @@ -1262,32 +1273,25 @@ async def test_run_live_clears_resumption_handle_on_transfer():

flow = BaseLlmFlowForTesting()

# Mock _receive_from_model to yield an event that triggers transfer
part = types.Part(
function_response=types.FunctionResponse(name='transfer_to_agent')
)
content = types.Content(parts=[part])
parts = [
types.Part(
function_response=types.FunctionResponse(name=function_response_name)
)
for function_response_name in function_response_names
]
content = types.Content(parts=parts)
transfer_event = Event(
id=Event.new_id(),
invocation_id=invocation_context.invocation_id,
author=agent.name,
)
transfer_event.content = content
transfer_event.actions = mock.Mock()
transfer_event.actions.transfer_to_agent = 'sub_agent'

class StopTest(Exception):
pass

receive_call_count = 0
transfer_event.actions.transfer_to_agent = (
'sub_agent' if should_transfer else None
)

async def mock_receive_from_model(*args, **kwargs):
nonlocal receive_call_count
receive_call_count += 1
if receive_call_count == 1:
yield transfer_event
else:
raise StopTest()
yield transfer_event

flow._receive_from_model = mock.Mock(side_effect=mock_receive_from_model)

Expand Down Expand Up @@ -1315,12 +1319,12 @@ async def mock_run_live_sub_agent(child_ctx, *args, **kwargs):
mock_connection = mock.AsyncMock()
mock_connect.return_value.__aenter__.return_value = mock_connection

try:
async for _ in flow.run_live(invocation_context):
pass
except StopTest:
async for _ in flow.run_live(invocation_context):
pass

assert mock_sub_agent.run_live.call_count == int(should_transfer)
assert mock_connection.close.await_count == int(should_transfer)

# Verify that parent's resumption handles were not cleared
assert invocation_context.live_session_resumption_handle == 'test_handle'
assert (
Expand Down