Skip to content
Open
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
15 changes: 15 additions & 0 deletions src/google/adk/flows/llm_flows/context/_contents.py
Original file line number Diff line number Diff line change
Expand Up @@ -300,6 +300,21 @@ def _build_task_input_user_content(
parts.append(types.Part(text=_SINGLE_TURN_NUDGE))
return types.Content(role='user', parts=parts)

if '@' in isolation_scope:
invocation_id = isolation_scope.rpartition('@')[2]
for event in all_events:
if (
event.invocation_id == invocation_id
and event.author == 'user'
and event.content
and event.content.parts
and not any(p.function_response for p in event.content.parts)
):
parts = list(event.content.parts)
if is_single_turn:
parts.append(types.Part(text=_SINGLE_TURN_NUDGE))
return types.Content(role='user', parts=parts)

# Fallback: workflow-node task with no originating FC. Use the
# node_input that the wrapper stamped onto ``ic.user_content``.
if user_content and user_content.parts:
Expand Down
5 changes: 4 additions & 1 deletion src/google/adk/workflow/_dynamic_node_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -336,7 +336,10 @@ async def _execute_step(
override_isolation_scope is None
and getattr(node, 'mode', None) == 'task'
):
override_isolation_scope = node_path
if not curr_parent_path:
override_isolation_scope = f"{target_node_name}@{ctx.invocation_id}"
else:
override_isolation_scope = node_path

# Rehydration chronological sequence barrier setup for the parent path
if self._enable_replay and curr_parent_path:
Expand Down
88 changes: 88 additions & 0 deletions tests/unittests/a2a/integration/test_client_server_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -787,3 +787,91 @@ def _status_update(text: str, state: int) -> pb.StreamResponse:
# the latest status is applied to the running task.
final_task, _ = results[2]
assert final_task.status.state == pb.TASK_STATE_COMPLETED


@pytest.mark.asyncio
async def test_a2a_task_clarification_and_resumption():
"""Tests that root task input preservation works across the A2A boundary.

When a RemoteA2aAgent operates in mode="task", it acts as a sub-agent on the
client, but runs as a root agent on the server. The server-side Runner must
successfully associate the clarification turn with the paused task's scope
and reconstruct the strictly-alternating task history correctly.
"""
from google.adk.agents.llm_agent import LlmAgent
from tests.unittests import testing_utils

model = testing_utils.MockModel.create(
responses=[
"What is your specific question?",
"Task completed.",
]
)
server_agent = LlmAgent(name="my_task", model=model, mode="task")

from google.adk.runners import Runner
from google.adk.a2a.executor.a2a_agent_executor import A2aAgentExecutor
from google.adk.a2a import _compat
from starlette.applications import Starlette
from a2a.server.tasks import InMemoryTaskStore as TaskStore

server_runner = Runner(
app_name="ServerApp",
agent=server_agent,
session_service=InMemorySessionService(),
)
executor = A2aAgentExecutor(runner=server_runner)
app = Starlette()
_compat.attach_a2a_routes_to_app(
app,
agent_card=agent_card,
agent_executor=executor,
task_store=TaskStore(),
)

async with app.router.lifespan_context(app):
client_agent = create_client(app, streaming=False, mode="task")

session_service = InMemorySessionService()
await session_service.create_session(
app_name="ClientApp", user_id="test_user", session_id="test_session"
)
client_runner = Runner(
app_name="ClientApp",
agent=client_agent,
session_service=session_service,
)

# Turn 1: Initial task request
new_message = types.Content(
parts=[types.Part(text="Do a task")], role="user"
)
async for _ in client_runner.run_async(
user_id="test_user", session_id="test_session", new_message=new_message
):
pass

# Turn 2: Clarification response
new_message2 = types.Content(
parts=[types.Part(text="Here is the clarification")], role="user"
)
async for _ in client_runner.run_async(
user_id="test_user", session_id="test_session", new_message=new_message2
):
pass

# Verify server-side model requests to ensure correct history reconstruction
assert len(model.requests) == 2

req1 = model.requests[0]
parts1 = [p.text for c in req1.contents if c.role == "user" for p in c.parts]
assert "Do a task" in parts1[0]
# Check strict alternation on Turn 1
assert [c.role for c in req1.contents] == ["user"]

req2 = model.requests[1]
parts2 = [p.text for c in req2.contents if c.role == "user" for p in c.parts]
assert "Do a task" in parts2[0]
assert "Here is the clarification" in parts2[1]
# Check strict alternation on Turn 2
assert [c.role for c in req2.contents] == ["user", "model", "user"]
52 changes: 52 additions & 0 deletions tests/unittests/workflow/test_task_api_e2e.py
Original file line number Diff line number Diff line change
Expand Up @@ -873,3 +873,55 @@ async def driver(ctx):
p.text or "" for c in beta_request.contents or [] for p in c.parts or []
)
assert "ALPHA_SECRET_CONVERSATION" not in rendered_beta_context


@pytest.mark.asyncio
async def test_root_task_survives_clarification(
request: pytest.FixtureRequest,
):
model = testing_utils.MockModel.create(
responses=[
"What is your specific question?",
"Got it."
]
)
agent = LlmAgent(name="my_task", model=model, mode="task")
app = App(name=request.function.__name__, root_agent=agent)
runner = testing_utils.InMemoryRunner(app=app)

await runner.run_async("Do a task")
await runner.run_async("Here is the clarification")

assert len(model.requests) == 2
req = model.requests[1]
parts = [p.text for c in req.contents if c.role == "user" for p in c.parts]
assert "Do a task" in parts[0]
assert "Here is the clarification" in parts[1]


@pytest.mark.asyncio
async def test_separate_root_tasks_same_session_get_separate_scopes(
request: pytest.FixtureRequest,
):
model = testing_utils.MockModel.create(
responses=[
_finish_part({"result": "Done A"}),
_finish_part({"result": "Done B"}),
]
)
agent = LlmAgent(name="my_task", model=model, mode="task")
app = App(name=request.function.__name__, root_agent=agent)
runner = testing_utils.InMemoryRunner(app=app)

await runner.run_async("Run A")
await runner.run_async("Run B")

events = runner.session.events
task_events = [e for e in events if e.author == "my_task" and e.output]
assert len(task_events) == 2
# scopes should be different
scope_a = task_events[0].isolation_scope
scope_b = task_events[1].isolation_scope
assert scope_a != scope_b
assert scope_a is not None
assert scope_b is not None
Loading