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
64 changes: 64 additions & 0 deletions src/google/adk/a2a/agent/_remote_a2a_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,8 @@
# Constants
A2A_METADATA_PREFIX = "a2a:"
DEFAULT_TIMEOUT = 600.0
# The error_code of an event for a remote task that ended in TASK_STATE_FAILED.
A2A_TASK_FAILED_ERROR_CODE = "A2A_TASK_FAILED"

_DEFAULT_PORTS = {"http": 80, "https": 443}

Expand Down Expand Up @@ -493,6 +495,61 @@ def _create_task_failure_events(
return error_event, finish_event


def _failed_task_status(
a2a_response: _compat.A2AClientEvent | A2AMessage,
) -> Any:
"""The status of the remote task a response reports as failed, if any.

A streamed failure arrives as a status update, which carries the state the
task ended in; a non-streamed one, as the task itself.
"""
if not isinstance(a2a_response, tuple):
return None
task, update = a2a_response
if update is None:
status = getattr(task, "status", None)
elif isinstance(update, A2ATaskStatusUpdateEvent):
status = update.status
else:
return None
if getattr(status, "state", None) != _compat.TS_FAILED:
return None
return status


def _mark_task_failed(
event: Optional[Event],
status: Any,
ctx: InvocationContext,
agent_name: str,
) -> Event:
"""Marks the event of a failed remote task as an error, creating one if the
response converted to none.

The text the remote agent sent with the failure is its account of why, so it
becomes the error message rather than being read as an ordinary answer. Only
the failed status's own message counts: a non-streamed task with none falls
back to its history, whose last agent message is earlier progress, not why
the task failed.
"""
if event is None:
event = Event(
author=agent_name,
invocation_id=ctx.invocation_id,
branch=ctx.branch,
)
event.error_code = event.error_code or A2A_TASK_FAILED_ERROR_CODE
if not event.error_message:
message = _compat.normalize_message(getattr(status, "message", None))
text = "\n".join(
_compat.part_text(part)
for part in (message.parts if message else [])
if _compat.is_text_part(part) and _compat.part_text(part)
)
event.error_message = text or "Remote A2A task failed"
return event


def _add_mock_function_call(event: Event, state: TaskState) -> None:
"""Generates a mock function call for input-required events if applicable."""
if event.content is None:
Expand Down Expand Up @@ -1810,6 +1867,13 @@ async def _run_async_impl(
event = await self._handle_a2a_response_v2(a2a_response, ctx)
else:
event = await self._handle_a2a_response(a2a_response, ctx)
# The response converters keep a failed task's text as content and
# drop its state, which would read as an ordinary answer. Task mode
# reports the failure with events of its own, below.
if self.mode != "task" and (
failed_status := _failed_task_status(a2a_response)
):
event = _mark_task_failed(event, failed_status, ctx, self.name)
if not event:
continue

Expand Down
134 changes: 134 additions & 0 deletions tests/unittests/a2a/agent/test_remote_a2a_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -8304,3 +8304,137 @@ async def override_metadata(ctx, response, event):
assert turn[-1].custom_metadata["a2a:task_id"] == "caller-task"
assert turn[-1].custom_metadata["a2a:context_id"] == "caller-context"
assert turn[-1].custom_metadata["a2a:response"]["id"] == "currency-task"


def _agent_message(message_id, text):
return _compat.make_message(
message_id=message_id,
role="agent",
parts=[_compat.make_text_part(text)],
)


def _failed_task_stream(*, streaming, use_v2, text, progress=None):
"""A task that works, then fails: a status update when streamed, the task
itself when not. *progress* is what the remote agent said while working: a
working status update when streamed, the task's history when not."""
message = _agent_message("failure", text) if text else None
if streaming:
stream = _bare_task_stream(_compat.TS_FAILED, use_v2=use_v2)
if message is not None:
status = _compat.make_task_status(_compat.TS_FAILED, message=message)
if _compat.IS_A2A_V1:
stream[-1].status_update.status.CopyFrom(status)
else:
stream[-1][0].status = status
stream[-1][1].status = status
if progress is not None:
working = _compat.make_task_status(
_compat.TS_WORKING, message=_agent_message("progress", progress)
)
update = _compat.make_task_status_update_event(
task_id="currency-task",
context_id="currency-context",
status=working,
final=False,
)
if _compat.IS_A2A_V1:
from a2a.types import StreamResponse

stream.insert(-1, StreamResponse(status_update=update))
else:
task = stream[0][0].model_copy(deep=True)
task.status = working
stream.insert(-1, (task, update))
return stream
task = _compat.make_task(
id="currency-task",
context_id="currency-context",
status=_compat.make_task_status(_compat.TS_FAILED, message=message),
history=[_agent_message("progress", progress)] if progress else [],
metadata=(
{remote_a2a_agent._NEW_A2A_ADK_INTEGRATION_EXTENSION: True}
if use_v2
else None
),
)
return [_make_stream_task(task)]


@pytest.mark.parametrize("use_v2", [False, True], ids=["legacy", "v2"])
@pytest.mark.parametrize("streaming", [True, False])
@pytest.mark.parametrize(
"text, error_message",
[
("claude exited with code 1", "claude exited with code 1"),
(None, "Remote A2A task failed"),
],
ids=["with-reason", "without-reason"],
)
async def test_failed_remote_task_is_an_error_event(
use_v2, streaming, text, error_message
):
"""A failed task ends the turn as an error, not as the remote's answer."""
completed = _bare_task_stream(_compat.TS_COMPLETED, use_v2=use_v2)
_, turns, session = await _run_remote_task_responses([
_failed_task_stream(streaming=streaming, use_v2=use_v2, text=text),
completed,
completed,
])

failure = turns[0][-1]
assert failure.error_code == remote_a2a_agent.A2A_TASK_FAILED_ERROR_CODE
assert failure.error_message == error_message
assert failure.is_final_response()
assert failure.custom_metadata["a2a:task_id"] == "currency-task"
# The failure is stored as an error too, not only shown as one.
assert [event.id for event in session.events if event.error_code] == [
failure.id
]
# Only the failure is marked: the working update before it and the later
# turns are not.
assert all(event.error_code is None for event in turns[0][:-1])
assert all(event.error_code is None for turn in turns[1:] for event in turn)


@pytest.mark.parametrize("use_v2", [False, True], ids=["legacy", "v2"])
@pytest.mark.parametrize("streaming", [True, False])
async def test_a_failure_without_a_reason_is_not_worded_with_earlier_progress(
use_v2, streaming
):
"""What the remote agent said while working is not why its task failed: a
failure that gives no reason gets the generic message, streamed or not."""
completed = _bare_task_stream(_compat.TS_COMPLETED, use_v2=use_v2)
_, turns, _ = await _run_remote_task_responses([
_failed_task_stream(
streaming=streaming,
use_v2=use_v2,
text=None,
progress="Reading a.txt.",
),
completed,
completed,
])

failure = turns[0][-1]
assert failure.error_code == remote_a2a_agent.A2A_TASK_FAILED_ERROR_CODE
assert failure.error_message == "Remote A2A task failed"


@pytest.mark.parametrize("use_v2", [False, True], ids=["legacy", "v2"])
@pytest.mark.parametrize(
"state",
[
_compat.TS_COMPLETED,
_compat.TS_CANCELED,
_compat.TS_REJECTED,
_compat.TS_INPUT_REQUIRED,
],
)
async def test_only_a_failed_remote_task_is_an_error_event(use_v2, state):
_, turns, _ = await _run_remote_task_responses(
[_bare_task_stream(state, use_v2=use_v2)]
+ [_bare_task_stream(_compat.TS_COMPLETED, use_v2=use_v2)] * 2
)

assert all(event.error_code is None for turn in turns for event in turn)
Loading