diff --git a/site/src/content/docs/user-guide/concepts/multi-agent/agent-to-agent.mdx b/site/src/content/docs/user-guide/concepts/multi-agent/agent-to-agent.mdx index 2264340f62..4084a4fbd2 100644 --- a/site/src/content/docs/user-guide/concepts/multi-agent/agent-to-agent.mdx +++ b/site/src/content/docs/user-guide/concepts/multi-agent/agent-to-agent.mdx @@ -387,9 +387,9 @@ This feature is currently available in the Python SDK only. When a tool or hook on the served agent raises an [interrupt](../interrupts.md), the task moves to the A2A `input_required` state and waits. The client answers it, and the task resumes the paused tool exactly where it stopped. -Answering an interrupt is the one flow on this page that needs the raw [`a2a-sdk`](https://github.com/a2aproject/a2a-python) client rather than `A2AAgent`. `A2AAgent` speaks in text, so it raises `ValueError` if you pass it interrupt responses, and it drops the `DataPart` carrying the interrupt ids when it reads the reply. The examples below build A2A messages directly. +Answering an interrupt is the one flow on this page that needs the raw [`a2a-sdk`](https://github.com/a2aproject/a2a-python) client rather than `A2AAgent`. `A2AAgent` speaks in text, so it raises `ValueError` if you pass it interrupt responses, and it drops the data part carrying the interrupt ids when it reads the reply. The examples below build A2A messages directly. -Each interrupt has a server-generated id, and an answer is bound to the id of the interrupt that raised it. The server advertises the pending interrupts on the `input_required` status message as a `DataPart`, alongside the human-readable `TextPart`: +Each interrupt has a server-generated id, and an answer is bound to the id of the interrupt that raised it. The server advertises the pending interrupts on the `input_required` status message as a data part, alongside the human-readable text part: ```json { @@ -406,7 +406,7 @@ Each interrupt has a server-generated id, and an answer is bound to the id of th } ``` -To answer, send a new message on the same `taskId` containing a `DataPart` that echoes the `interruptId` back with the response: +To answer, send a new message on the same `taskId` containing a data part that echoes the `interruptId` back with the response: ```json { @@ -420,25 +420,31 @@ To answer, send a new message on the same `taskId` containing a `DataPart` that } ``` -The `response` becomes the return value of the `interrupt()` call that paused the tool. It is any JSON value except `null`, which the server refuses — a null answer would leave the interrupt unsatisfied and re-raise it. `false` and `0` are fine. Answer several interrupts in one message by sending one `DataPart` for each. +The `response` becomes the return value of the `interrupt()` call that paused the tool. It is any JSON value except `null`, which the server refuses — a null answer would leave the interrupt unsatisfied and re-raise it. `false` and `0` are fine. Answer several interrupts in one message by sending one data part for each. + +A2A data parts carry numbers as protobuf `Value`, which has no integer type — a whole-number +float like `1.0` is indistinguishable from the int `1` on the wire, so the server normalizes any +whole float back to an int before it reaches `response`. A non-whole float such as `1.5` is +unaffected. Reading the ids off the status message and answering them: ```python -from a2a.types import DataPart, Part +from a2a.helpers import new_data_part +from google.protobuf.json_format import MessageToDict # The task parked in input_required; read the interrupts it is waiting on. pending = next( - part.root.data["interrupts"] + MessageToDict(part.data)["interrupts"] for part in task.status.message.parts - if isinstance(part.root, DataPart) and "interrupts" in part.root.data + if part.HasField("data") and "interrupts" in MessageToDict(part.data) ) # Answer each one on the same taskId. answers = [ - Part(root=DataPart(data={ + new_data_part({ "interruptResponse": {"interruptId": item["interruptId"], "response": {"approved": True}} - })) + }) for item in pending ] ``` @@ -453,7 +459,7 @@ The server rejects an answer it cannot bind, before the agent runs, so a refused A task with a pending interrupt also rejects an ordinary conversational message — answer the interrupt, or cancel the task. -A `DataPart` without an `interruptResponse` key is unaffected and continues to reach the agent as structured data. +A data part without an `interruptResponse` key is unaffected and continues to reach the agent as structured data. ### Server Configuration Options diff --git a/strands-py/pyproject.toml b/strands-py/pyproject.toml index f93ef9ae3b..801c9500b5 100644 --- a/strands-py/pyproject.toml +++ b/strands-py/pyproject.toml @@ -68,8 +68,8 @@ docs = [ ] a2a = [ - "a2a-sdk>=0.3.0,<0.4.0", - "a2a-sdk[sql]>=0.3.0,<0.4.0", + "a2a-sdk>=1.1.0,<2.0.0", + "a2a-sdk[sql]>=1.1.0,<2.0.0", "uvicorn>=0.34.2,<1.0.0", "httpx>=0.28.1,<1.0.0", "fastapi>=0.133.0,<1.0.0", diff --git a/strands-py/src/strands/agent/a2a_agent.py b/strands-py/src/strands/agent/a2a_agent.py index eeb96f7a24..439fc28db2 100644 --- a/strands-py/src/strands/agent/a2a_agent.py +++ b/strands-py/src/strands/agent/a2a_agent.py @@ -15,14 +15,10 @@ import httpx from a2a.client import A2ACardResolver, ClientConfig, ClientFactory -from a2a.types import AgentCard, Message, TaskArtifactUpdateEvent, TaskStatusUpdateEvent +from a2a.types import AgentCard, SendMessageRequest from .._async import run_async -from ..multiagent.a2a._converters import ( - _STATE_TO_STOP_REASON, - convert_input_to_message, - convert_response_to_agent_result, -) +from ..multiagent.a2a._converters import convert_input_to_message, convert_responses_to_agent_result from ..types._events import AgentResultEvent from ..types.a2a import A2AResponse, A2AStreamEvent from ..types.agent import AgentInput @@ -33,13 +29,6 @@ _DEFAULT_TIMEOUT = 300 -# A2A task states that indicate the response stream is complete. -# Derived from the canonical _STATE_TO_STOP_REASON mapping in _converters. -# Terminal states (end_turn) mean no more events; input states (interrupt) mean execution is paused. -_TERMINAL_STATES = {state for state, reason in _STATE_TO_STOP_REASON.items() if reason == "end_turn"} -_INPUT_STATES = {state for state, reason in _STATE_TO_STOP_REASON.items() if reason == "interrupt"} -_COMPLETE_STATES = _TERMINAL_STATES | _INPUT_STATES - class A2AAgent(AgentBase): """Client wrapper for remote A2A agents.""" @@ -157,9 +146,9 @@ async def stream_async( Yields: An async iterator that yields events. Each event is a dictionary: - - A2AStreamEvent: {"type": "a2a_stream", "event": } - where the A2A object can be a Message, or a tuple of - (Task, TaskStatusUpdateEvent) or (Task, TaskArtifactUpdateEvent). + - A2AStreamEvent: {"type": "a2a_stream", "event": } + where the StreamResponse carries exactly one of a task, message, + status_update, or artifact_update. - AgentResultEvent: {"result": AgentResult} - always emitted last. Raises: @@ -174,20 +163,14 @@ async def stream_async( print(f"Final result: {event['result'].message}") ``` """ - last_event = None - last_complete_event = None - - async for event in self._send_message(prompt): - last_event = event - if self._is_complete_event(event): - last_complete_event = event - yield A2AStreamEvent(event) + responses: list[A2AResponse] = [] - # Use the last complete event if available, otherwise fall back to last event - final_event = last_complete_event or last_event + async for response in self._send_message(prompt): + responses.append(response) + yield A2AStreamEvent(response) - if final_event is not None: - result = convert_response_to_agent_result(final_event) + if responses: + result = convert_responses_to_agent_result(responses) yield AgentResultEvent(result) async def get_agent_card(self) -> AgentCard: @@ -214,11 +197,11 @@ async def get_agent_card(self) -> AgentCard: self._agent_card = await resolver.get_agent_card() # Populate name from card if not set - if self.name is None and self._agent_card.name is not None: + if self.name is None and self._agent_card.name: self.name = self._agent_card.name # Populate description from card if not set - if self.description is None and self._agent_card.description is not None: + if self.description is None and self._agent_card.description: self.description = self._agent_card.description logger.debug("agent=<%s>, endpoint=<%s> | discovered agent card", self.name, self.endpoint) @@ -258,7 +241,7 @@ async def _send_message(self, prompt: AgentInput) -> AsyncIterator[A2AResponse]: prompt: Input to send to the agent. Yields: - A2A response events. + A2A StreamResponse events. Raises: ValueError: If prompt is None. @@ -267,46 +250,9 @@ async def _send_message(self, prompt: AgentInput) -> AsyncIterator[A2AResponse]: raise ValueError("prompt is required for A2AAgent") message = convert_input_to_message(prompt) + request = SendMessageRequest(message=message) logger.debug("agent=<%s>, endpoint=<%s> | sending message", self.name, self.endpoint) async with self._get_a2a_client() as client: - async for event in client.send_message(message): - yield event - - def _is_complete_event(self, event: A2AResponse) -> bool: - """Check if an A2A event represents a complete response. - - Recognizes all terminal states (completed, failed, canceled, rejected) - and pausing states (input_required, auth_required) as complete events. - - Args: - event: A2A event. - - Returns: - True if the event represents a complete response. - """ - # Direct Message is always complete - if isinstance(event, Message): - return True - - # Handle tuple responses (Task, UpdateEvent | None) - if isinstance(event, tuple) and len(event) == 2: - task, update_event = event - - # Initial task response (no update event) - if update_event is None: - return True - - # Artifact update with last_chunk flag - if isinstance(update_event, TaskArtifactUpdateEvent): - if hasattr(update_event, "last_chunk") and update_event.last_chunk is not None: - return update_event.last_chunk - return False - - # Status update - check for terminal or pausing states - if isinstance(update_event, TaskStatusUpdateEvent): - if update_event.status and hasattr(update_event.status, "state"): - state = update_event.status.state - return state in _COMPLETE_STATES - - return False + async for response in client.send_message(request): + yield response diff --git a/strands-py/src/strands/multiagent/a2a/_converters.py b/strands-py/src/strands/multiagent/a2a/_converters.py index 7808ae3256..b04cc75c7a 100644 --- a/strands-py/src/strands/multiagent/a2a/_converters.py +++ b/strands-py/src/strands/multiagent/a2a/_converters.py @@ -1,10 +1,12 @@ """Conversion functions between Strands and A2A types.""" +from collections.abc import Sequence +from dataclasses import dataclass, field from typing import cast from uuid import uuid4 from a2a.types import Message as A2AMessage -from a2a.types import Part, Role, TaskArtifactUpdateEvent, TaskState, TaskStatusUpdateEvent, TextPart +from a2a.types import Part, Role, TaskArtifactUpdateEvent, TaskState, TaskStatus from ...agent.agent_result import AgentResult from ...telemetry.metrics import EventLoopMetrics @@ -15,15 +17,27 @@ # Mapping from A2A TaskState to Strands stop_reason _STATE_TO_STOP_REASON: dict[TaskState, StopReason] = { - TaskState.completed: "end_turn", - TaskState.failed: "end_turn", - TaskState.canceled: "end_turn", - TaskState.rejected: "end_turn", - TaskState.input_required: "interrupt", - TaskState.auth_required: "interrupt", + TaskState.TASK_STATE_COMPLETED: "end_turn", + TaskState.TASK_STATE_FAILED: "end_turn", + TaskState.TASK_STATE_CANCELED: "end_turn", + TaskState.TASK_STATE_REJECTED: "end_turn", + TaskState.TASK_STATE_INPUT_REQUIRED: "interrupt", + TaskState.TASK_STATE_AUTH_REQUIRED: "interrupt", } +def _task_state_to_str(task_state: TaskState) -> str: + """Render a TaskState as the kebab-case string stored in ``AgentResult.state["a2a_task_state"]``. + + TaskState's protobuf enum names are SCREAMING_SNAKE_CASE (e.g. ``TASK_STATE_INPUT_REQUIRED``); + this renders the kebab-case form (e.g. ``"input-required"``) that existing Strands callers expect. + """ + if task_state == TaskState.TASK_STATE_UNSPECIFIED: + return "unknown" + name: str = TaskState.Name(task_state) # type: ignore[attr-defined] + return name.removeprefix("TASK_STATE_").lower().replace("_", "-") + + def convert_input_to_message(prompt: AgentInput) -> A2AMessage: """Convert AgentInput to A2A Message. @@ -40,9 +54,8 @@ def convert_input_to_message(prompt: AgentInput) -> A2AMessage: if isinstance(prompt, str): return A2AMessage( - kind="message", - role=Role.user, - parts=[Part(TextPart(kind="text", text=prompt))], + role=Role.ROLE_USER, + parts=[Part(text=prompt)], message_id=message_id, ) @@ -57,16 +70,14 @@ def convert_input_to_message(prompt: AgentInput) -> A2AMessage: content = cast(list[ContentBlock], msg.get("content", [])) parts = convert_content_blocks_to_parts(content) return A2AMessage( - kind="message", - role=Role.user, + role=Role.ROLE_USER, parts=parts, message_id=message_id, ) else: parts = convert_content_blocks_to_parts(cast(list[ContentBlock], prompt)) return A2AMessage( - kind="message", - role=Role.user, + role=Role.ROLE_USER, parts=parts, message_id=message_id, ) @@ -86,29 +97,121 @@ def convert_content_blocks_to_parts(content_blocks: list[ContentBlock]) -> list[ parts = [] for block in content_blocks: if "text" in block: - parts.append(Part(TextPart(kind="text", text=block["text"]))) + parts.append(Part(text=block["text"])) return parts def _extract_task_state(response: A2AResponse) -> TaskState | None: - """Extract the task state from an A2A response. + """Extract the task state carried by a single A2A StreamResponse, if any. Args: - response: A2A response (either A2AMessage or tuple of task and update event). + response: A single StreamResponse from the A2A event stream. Returns: - The TaskState if available, None otherwise. + The TaskState if this response carries one (a ``task`` or ``status_update``), else None. """ - if isinstance(response, tuple) and len(response) == 2: - _task, update_event = response - if isinstance(update_event, TaskStatusUpdateEvent): - if update_event.status and hasattr(update_event.status, "state"): - return update_event.status.state + if response.HasField("status_update"): + return response.status_update.status.state + if response.HasField("task") and response.task.HasField("status"): + return response.task.status.state return None -def convert_response_to_agent_result(response: A2AResponse) -> AgentResult: - """Convert A2A response to AgentResult. +def _parts_to_content(parts: Sequence[Part]) -> list[ContentBlock]: + """Convert a sequence of A2A text Parts into Strands ContentBlocks. + + Drops non-text parts and empty-text parts (the latter appear as a content-less + ``last_chunk`` marker on compliant-streaming artifact updates). + """ + return [{"text": part.text} for part in parts if part.HasField("text") and part.text] + + +@dataclass +class _ResponseAccumulator: + """Accumulates content and task state across a full A2A StreamResponse sequence. + + See ``convert_responses_to_agent_result`` for the content precedence this implements. + """ + + artifact_parts: dict[str, list[ContentBlock]] = field(default_factory=dict) + artifact_order: list[str] = field(default_factory=list) + terminal_message_content: list[ContentBlock] = field(default_factory=list) + narration_content: list[ContentBlock] = field(default_factory=list) + task_content: list[ContentBlock] = field(default_factory=list) + message_content: list[ContentBlock] = field(default_factory=list) + task_state: TaskState | None = None + + def ingest(self, response: A2AResponse) -> None: + """Fold one StreamResponse event into the accumulated content and task state.""" + state = _extract_task_state(response) + if state is not None: + self.task_state = state + + if response.HasField("artifact_update"): + self._ingest_artifact_update(response.artifact_update) + elif response.HasField("status_update"): + self._ingest_status_update(response.status_update.status) + elif response.HasField("task"): + self.task_content = [ + content for artifact in response.task.artifacts for content in _parts_to_content(artifact.parts) + ] + if not self.task_content and response.task.HasField("status") and response.task.status.HasField("message"): + self.task_content = _parts_to_content(response.task.status.message.parts) + elif response.HasField("message"): + self.message_content = _parts_to_content(response.message.parts) + + def _ingest_artifact_update(self, update: TaskArtifactUpdateEvent) -> None: + """Fold one artifact_update event, honoring ``append`` (A2A: false/unset replaces).""" + artifact_id = update.artifact.artifact_id + parts_content = _parts_to_content(update.artifact.parts) + if artifact_id not in self.artifact_parts: + self.artifact_parts[artifact_id] = [] + self.artifact_order.append(artifact_id) + if update.append: + self.artifact_parts[artifact_id].extend(parts_content) + else: + self.artifact_parts[artifact_id] = parts_content + + def _ingest_status_update(self, status: TaskStatus) -> None: + """Route a status_update's message: a terminal state carries actionable text, else narration.""" + if not status.HasField("message"): + return + parts_content = _parts_to_content(status.message.parts) + if status.state in _STATE_TO_STOP_REASON: + self.terminal_message_content = parts_content + else: + self.narration_content = parts_content + + @property + def artifact_content(self) -> list[ContentBlock]: + """Accumulated artifact content across all artifact ids, in first-seen order.""" + return [content for artifact_id in self.artifact_order for content in self.artifact_parts[artifact_id]] + + @property + def content(self) -> list[ContentBlock]: + """The final content for the AgentResult, per the precedence documented on the caller.""" + artifact_content = self.artifact_content + if artifact_content or self.terminal_message_content: + return artifact_content + self.terminal_message_content + return self.narration_content or self.task_content or self.message_content + + +def convert_responses_to_agent_result(responses: Sequence[A2AResponse]) -> AgentResult: + """Convert the full sequence of A2A StreamResponse events from one call into an AgentResult. + + Each StreamResponse carries at most one of ``task`` | ``message`` | ``status_update`` | + ``artifact_update``, and no single event is guaranteed to carry the final content by itself, so + content is reconstructed across the whole stream: + - ``artifact_update`` parts accumulate per ``artifact_id``, honoring the event's ``append`` + flag (A2A schema: unset/false replaces that artifact's accumulated parts, true appends to + them), so a peer that re-sends a cumulative artifact each turn doesn't duplicate content. + - a terminal ``status_update`` message (one whose state is in ``_STATE_TO_STOP_REASON``, + e.g. input_required or failed) is appended after any artifact content, since it carries + the actionable text (an approval prompt, a failure reason) rather than duplicating it. + - a non-terminal ``status_update`` message is progress narration and is only used as a + fallback when no artifact or terminal-status content was found. + - a bare ``task`` or ``message`` response (no separate update events) supplies content + directly. Maps A2A task lifecycle states to appropriate Strands stop_reasons: - completed → end_turn @@ -119,61 +222,33 @@ def convert_response_to_agent_result(response: A2AResponse) -> AgentResult: - auth_required → interrupt (agent needs authentication) Args: - response: A2A response (either A2AMessage or tuple of task and update event). + responses: All StreamResponse events observed for one ``send_message`` call, in order. Returns: AgentResult with extracted content and metadata. """ - content: list[ContentBlock] = [] - task_state = _extract_task_state(response) - stop_reason: StopReason = _STATE_TO_STOP_REASON.get(task_state, "end_turn") if task_state else "end_turn" - - if isinstance(response, tuple) and len(response) == 2: - task, update_event = response - - # Handle artifact updates - if isinstance(update_event, TaskArtifactUpdateEvent): - if update_event.artifact and hasattr(update_event.artifact, "parts") and update_event.artifact.parts: - for part in update_event.artifact.parts: - if hasattr(part, "root") and hasattr(part.root, "text"): - content.append({"text": part.root.text}) - # Handle status updates with messages - elif isinstance(update_event, TaskStatusUpdateEvent): - if ( - update_event.status - and hasattr(update_event.status, "message") - and update_event.status.message - and update_event.status.message.parts - ): - for part in update_event.status.message.parts: - if hasattr(part, "root") and hasattr(part.root, "text"): - content.append({"text": part.root.text}) - - # Use task.artifacts when no content was extracted from the event - if not content and task and hasattr(task, "artifacts") and task.artifacts is not None: - for artifact in task.artifacts: - if hasattr(artifact, "parts") and artifact.parts: - for part in artifact.parts: - if hasattr(part, "root") and hasattr(part.root, "text"): - content.append({"text": part.root.text}) - elif isinstance(response, A2AMessage): - for part in response.parts: - if hasattr(part, "root") and hasattr(part.root, "text"): - content.append({"text": part.root.text}) + accumulator = _ResponseAccumulator() + for response in responses: + accumulator.ingest(response) + + task_state = accumulator.task_state + stop_reason: StopReason = ( + _STATE_TO_STOP_REASON.get(task_state, "end_turn") if task_state is not None else "end_turn" + ) message: Message = { "role": "assistant", - "content": content, + "content": accumulator.content, } # Build state dict with A2A metadata - state: dict[str, str] = {} + state_dict: dict[str, str] = {} if task_state is not None: - state["a2a_task_state"] = task_state.value + state_dict["a2a_task_state"] = _task_state_to_str(task_state) return AgentResult( stop_reason=stop_reason, message=message, metrics=EventLoopMetrics(), - state=state, + state=state_dict, ) diff --git a/strands-py/src/strands/multiagent/a2a/executor.py b/strands-py/src/strands/multiagent/a2a/executor.py index 78493d7a03..472e0818a7 100644 --- a/strands-py/src/strands/multiagent/a2a/executor.py +++ b/strands-py/src/strands/multiagent/a2a/executor.py @@ -9,35 +9,26 @@ """ import asyncio -import base64 import json import logging import mimetypes import uuid import warnings from collections import OrderedDict -from collections.abc import Callable +from collections.abc import Callable, Sequence from dataclasses import dataclass from typing import Any, Literal, cast +from a2a.helpers import new_data_part, new_task_from_user_message, new_text_message, new_text_part from a2a.server.agent_execution import AgentExecutor, RequestContext from a2a.server.events import EventQueue from a2a.server.tasks import TaskUpdater -from a2a.types import ( - DataPart, - FilePart, - InternalError, - InvalidParamsError, - Part, - TaskState, - TextPart, - UnsupportedOperationError, -) -from a2a.utils import new_agent_text_message, new_task -from a2a.utils.errors import ServerError +from a2a.types import Part, TaskState +from a2a.utils.errors import InternalError, InvalidParamsError, UnsupportedOperationError +from google.protobuf.json_format import MessageToDict from ...agent.agent import Agent as SAAgent -from ...agent.agent import AgentResult as SAAgentResult +from ...agent.agent_result import AgentResult as SAAgentResult from ...session.session_manager import SessionManager from ...types._snapshot import Snapshot from ...types.content import ContentBlock @@ -56,7 +47,7 @@ # A factory that builds a fresh Agent for a given A2A context_id. AgentFactory = Callable[[str], SAAgent] -# Key identifying a DataPart that carries a Strands interrupt response. The A2A payload mirrors the +# Key identifying a data Part that carries a Strands interrupt response. The A2A payload mirrors the # `InterruptResponseContent` type verbatim so the wire contract and the SDK type cannot drift. INTERRUPT_RESPONSE_KEY = "interruptResponse" @@ -278,7 +269,7 @@ async def _stream_agent( """Stream one agent invocation and translate its events to A2A updates. Raises: - ServerError: If the input does not match the agent's interrupt state — interrupt + InvalidParamsError: If the input does not match the agent's interrupt state — interrupt responses that name no parked interrupt, or fresh content for a parked task. Both fail before the agent runs, so a parked interrupt survives a rejected resume. """ @@ -287,10 +278,8 @@ async def _stream_agent( if _is_interrupt_resume(prompt): self._validate_interrupt_resume(agent, cast(list[InterruptResponseContent], prompt)) elif agent._interrupt_state.activated: - raise ServerError( - error=InvalidParamsError( - message="Task is awaiting an interrupt response and cannot accept a new message" - ) + raise InvalidParamsError( + message="Task is awaiting an interrupt response and cannot accept a new message" ) from None try: @@ -329,21 +318,21 @@ async def execute( event_queue: The A2A event queue used to send response events back to the client. Raises: - ServerError: If an unrecoverable error occurs during agent execution setup - (e.g., missing input). Agent execution errors are handled gracefully - by transitioning the task to the failed state. + InvalidParamsError: If the request carries malformed or unresolvable parameters. + InternalError: If the request is missing required content. + UnsupportedOperationError: If the operation is not supported in the current state. """ task = context.current_task if not task: - task = new_task(context.message) # type: ignore + task = new_task_from_user_message(context.message) # type: ignore await event_queue.enqueue_event(task) updater = TaskUpdater(event_queue, task.id, task.context_id) try: await self._execute_streaming(context, updater) - except ServerError: - # Re-raise ServerErrors (setup failures like missing input) + except (InvalidParamsError, InternalError, UnsupportedOperationError): + # Re-raise server-side setup failures (missing input, bad params, unsupported op). raise except asyncio.CancelledError: # asyncio.CancelledError is a BaseException (not Exception) — raised when @@ -353,7 +342,7 @@ async def execute( try: await updater.cancel( message=updater.new_agent_message( - parts=[Part(root=TextPart(text="Task cancelled due to connection termination"))] + parts=[new_text_part("Task cancelled due to connection termination")] ) ) except RuntimeError: @@ -364,9 +353,7 @@ async def execute( # Agent execution failures transition to failed state logger.exception("task_id=<%s> | agent execution failed, transitioning to failed state", task.id) try: - await updater.failed( - message=updater.new_agent_message(parts=[Part(root=TextPart(text="Agent execution failed"))]) - ) + await updater.failed(message=updater.new_agent_message(parts=[new_text_part("Agent execution failed")])) except RuntimeError: # Task already in terminal state (e.g., completed before error in cleanup) logger.debug("task_id=<%s> | task already in terminal state, cannot transition to failed", task.id) @@ -382,11 +369,11 @@ async def _execute_streaming(self, context: RequestContext, updater: TaskUpdater updater: The task updater for managing task state and sending updates. Raises: - ServerError: If input conversion fails (missing or empty content), or if the message - carries malformed interrupt responses. + InternalError: If input conversion fails (missing or empty content). + InvalidParamsError: If the message carries malformed interrupt responses. """ if not (context.message and hasattr(context.message, "parts")): - raise ServerError(error=InternalError(message="Request message is missing or has no parts")) from None + raise InternalError(message="Request message is missing or has no parts") from None # Interrupt responses resume a parked task, so they are recognized before the generic # conversion below would flatten them into text. @@ -394,9 +381,7 @@ async def _execute_streaming(self, context: RequestContext, updater: TaskUpdater if prompt is None: prompt = self._convert_a2a_parts_to_content_blocks(context.message.parts) if not prompt: - raise ServerError( - error=InternalError(message="No valid content found in request message parts") - ) from None + raise InternalError(message="No valid content found in request message parts") from None if not self.enable_a2a_compliant_streaming: warnings.warn( @@ -417,7 +402,7 @@ async def _execute_streaming(self, context: RequestContext, updater: TaskUpdater # The framework always populates context_id before execute() runs; isolation is keyed on it. context_id = context.context_id if not context_id: - raise ServerError(error=InternalError(message="Request is missing a context_id")) from None + raise InternalError(message="Request is missing a context_id") from None if self._agent_factory is not None: await self._run_with_context_agent(context_id, prompt, invocation_state, updater, stream_state) @@ -431,7 +416,7 @@ async def _handle_interrupt_result(self, result: SAAgentResult, updater: TaskUpd the A2A `input_required` state. The interrupt details are communicated to the client via the status message. - The details are carried twice: a TextPart describing what is needed, and a DataPart holding + The details are carried twice: a text Part describing what is needed, and a data Part holding each interrupt's id. Only the id lets a client address its response back to the interrupt that raised it, and an id is generated server-side so it cannot be inferred from the prose. @@ -458,10 +443,10 @@ async def _handle_interrupt_result(self, result: SAAgentResult, updater: TaskUpd # Still transition to input_required — the agent signaled it needs input. input_message = "Agent requires additional input to continue" - # The TextPart stays first so clients that only read prose are unaffected. - parts = [Part(root=TextPart(text=input_message))] + # The text Part stays first so clients that only read prose are unaffected. + parts = [new_text_part(input_message)] if pending_interrupts: - parts.append(Part(root=DataPart(data={INTERRUPTS_KEY: pending_interrupts}))) + parts.append(new_data_part({INTERRUPTS_KEY: pending_interrupts})) await updater.requires_input(message=updater.new_agent_message(parts=parts)) @@ -485,7 +470,7 @@ async def _handle_streaming_event( if text_content := event["data"]: if stream_state is not None: await updater.add_artifact( - [Part(root=TextPart(text=text_content))], + [new_text_part(text_content)], artifact_id=stream_state.artifact_id, name="agent_response", append=not stream_state.is_first_chunk, @@ -494,11 +479,11 @@ async def _handle_streaming_event( else: # Legacy use update_status with agent message await updater.update_status( - TaskState.working, - new_agent_text_message( + TaskState.TASK_STATE_WORKING, + new_text_message( text_content, - updater.context_id, - updater.task_id, + context_id=updater.context_id, + task_id=updater.task_id, ), ) @@ -524,14 +509,14 @@ async def _handle_agent_result( if stream_state.is_first_chunk: final_content = str(result) if result else "" await updater.add_artifact( - [Part(root=TextPart(text=final_content))], + [new_text_part(final_content)], artifact_id=stream_state.artifact_id, name="agent_response", last_chunk=True, ) else: await updater.add_artifact( - [Part(root=TextPart(text=""))], + [new_text_part("")], artifact_id=stream_state.artifact_id, name="agent_response", append=True, @@ -539,7 +524,7 @@ async def _handle_agent_result( ) elif final_content := str(result): await updater.add_artifact( - [Part(root=TextPart(text=final_content))], + [new_text_part(final_content)], name="agent_response", ) await updater.complete() @@ -559,12 +544,13 @@ async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None event_queue: The A2A event queue. Raises: - ServerError: If no current task exists or the task is already in a terminal state. + UnsupportedOperationError: If no current task exists or the task is already in a + terminal state. """ task = context.current_task if not task: logger.warning("context_id=<%s> | cancel requested but no current task found", context.context_id) - raise ServerError(error=UnsupportedOperationError()) from None + raise UnsupportedOperationError() from None # Cooperatively cancel the agent's execution (best-effort). In factory mode, resolve the # agent for this context; in single-agent mode, the shared agent. @@ -582,12 +568,12 @@ async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None try: await updater.cancel( - message=updater.new_agent_message(parts=[Part(root=TextPart(text="Task cancelled by client request"))]) + message=updater.new_agent_message(parts=[new_text_part("Task cancelled by client request")]) ) except RuntimeError: # TaskUpdater raises RuntimeError when task is already in a terminal state logger.warning("task_id=<%s> | cannot cancel, already in terminal state", task.id) - raise ServerError(error=UnsupportedOperationError()) from None + raise UnsupportedOperationError() from None def _get_file_type_from_mime_type(self, mime_type: str | None) -> Literal["document", "image", "video", "unknown"]: """Classify file type based on MIME type. @@ -660,16 +646,23 @@ def _strip_file_extension(self, file_name: str) -> str: return file_name.rsplit(".", 1)[0] return file_name - def _extract_interrupt_responses(self, parts: list[Part]) -> list[InterruptResponseContent] | None: + def _extract_interrupt_responses(self, parts: Sequence[Part]) -> list[InterruptResponseContent] | None: """Extract Strands interrupt responses from inbound A2A message parts. - A client resumes a task parked in ``input_required`` by sending a DataPart shaped like the + A client resumes a task parked in ``input_required`` by sending a data Part shaped like the Strands ``InterruptResponseContent`` type:: - {"kind": "data", "data": {"interruptResponse": {"interruptId": "", "response": }}} + Part(data={"interruptResponse": {"interruptId": "", "response": }}) Recognition is deliberately narrow: only the explicit shape above is treated as a resume, so - an ordinary DataPart still reaches the generic content-block path unchanged. + an ordinary data Part still reaches the generic content-block path unchanged. + + Note: + A2A data parts use protobuf ``Value``, which has no integer type — all numbers are + IEEE 754 doubles. ``MessageToDict`` renders every number as a Python ``float``, so an + integer a peer sent (e.g. ``3``) arrives as ``3.0``. This is a v1 wire-format + limitation, not a Strands bug. Downstream code resuming an interrupt should handle + numeric ``response`` values as floats. Args: parts: List of A2A Part objects from the inbound message. @@ -679,49 +672,40 @@ def _extract_interrupt_responses(self, parts: list[Part]) -> list[InterruptRespo caller should fall back to generic content-block conversion. Raises: - ServerError: If an interrupt response is malformed, carries a null response, repeats an - interrupt id, or is accompanied by unrelated content in the same message. + InvalidParamsError: If an interrupt response is malformed, carries a null response, + repeats an interrupt id, or is accompanied by unrelated content in the same message. """ responses: list[InterruptResponseContent] = [] seen_ids: set[str] = set() unrelated_parts = 0 for part in parts: - part_root = part.root - data = part_root.data if isinstance(part_root, DataPart) else None + data = MessageToDict(part.data) if part.HasField("data") else None if not isinstance(data, dict) or INTERRUPT_RESPONSE_KEY not in data: unrelated_parts += 1 continue response = data[INTERRUPT_RESPONSE_KEY] if not isinstance(response, dict): - raise ServerError( - error=InvalidParamsError( - message=f"'{INTERRUPT_RESPONSE_KEY}' must be an object with 'interruptId' and 'response'" - ) + raise InvalidParamsError( + message=f"'{INTERRUPT_RESPONSE_KEY}' must be an object with 'interruptId' and 'response'" ) from None interrupt_id = response.get("interruptId") if not isinstance(interrupt_id, str) or not interrupt_id: - raise ServerError( - error=InvalidParamsError(message="Interrupt response is missing a non-empty 'interruptId'") - ) from None + raise InvalidParamsError(message="Interrupt response is missing a non-empty 'interruptId'") from None # `Interrupt.response` of None means "not yet answered", so a null answer would leave # the interrupt unsatisfied and re-raise it — the client would see an identical # input_required and no error. Falsy answers such as False are fine. if response.get("response") is None: - raise ServerError( - error=InvalidParamsError( - message=f"Interrupt response for '{interrupt_id}' must provide a non-null 'response'" - ) + raise InvalidParamsError( + message=f"Interrupt response for '{interrupt_id}' must provide a non-null 'response'" ) from None # Two answers for one interrupt are ambiguous; reject rather than silently choosing one. if interrupt_id in seen_ids: - raise ServerError( - error=InvalidParamsError(message=f"Duplicate interrupt response for '{interrupt_id}'") - ) from None + raise InvalidParamsError(message=f"Duplicate interrupt response for '{interrupt_id}'") from None seen_ids.add(interrupt_id) responses.append( @@ -736,10 +720,8 @@ def _extract_interrupt_responses(self, parts: list[Part]) -> list[InterruptRespo # The agent resumes from interrupt responses alone; delivering both would mean dropping the # conversational content, which is exactly the silent behavior the resume path must avoid. if unrelated_parts: - raise ServerError( - error=InvalidParamsError( - message="A message carrying interrupt responses must not contain other content parts" - ) + raise InvalidParamsError( + message="A message carrying interrupt responses must not contain other content parts" ) from None logger.debug("interrupt_ids=<%s> | extracted interrupt responses from request", sorted(seen_ids)) @@ -753,15 +735,13 @@ def _validate_interrupt_resume(self, agent: SAAgent, responses: list[InterruptRe responses: Interrupt responses extracted from the inbound message. Raises: - ServerError: If the agent holds no parked interrupts, or a response names an interrupt - the agent is not waiting on. + InvalidParamsError: If the agent holds no parked interrupts, or a response names an + interrupt the agent is not waiting on. """ interrupt_state = agent._interrupt_state if not interrupt_state.activated: - raise ServerError( - error=InvalidParamsError(message="Received interrupt responses but no interrupt is pending") - ) from None + raise InvalidParamsError(message="Received interrupt responses but no interrupt is pending") from None unknown_ids = sorted( content["interruptResponse"]["interruptId"] @@ -769,11 +749,9 @@ def _validate_interrupt_resume(self, agent: SAAgent, responses: list[InterruptRe if content["interruptResponse"]["interruptId"] not in interrupt_state.interrupts ) if unknown_ids: - raise ServerError( - error=InvalidParamsError(message=f"No pending interrupt matches id(s): {', '.join(unknown_ids)}") - ) from None + raise InvalidParamsError(message=f"No pending interrupt matches id(s): {', '.join(unknown_ids)}") from None - def _convert_a2a_parts_to_content_blocks(self, parts: list[Part]) -> list[ContentBlock]: + def _convert_a2a_parts_to_content_blocks(self, parts: Sequence[Part]) -> list[ContentBlock]: """Convert A2A message parts to Strands ContentBlocks. Args: @@ -786,70 +764,59 @@ def _convert_a2a_parts_to_content_blocks(self, parts: list[Part]) -> list[Conten for part in parts: try: - part_root = part.root - - if isinstance(part_root, TextPart): - # Handle TextPart - content_blocks.append(ContentBlock(text=part_root.text)) + if part.HasField("text"): + content_blocks.append(ContentBlock(text=part.text)) - elif isinstance(part_root, FilePart): - # Handle FilePart - file_obj = part_root.file - mime_type = getattr(file_obj, "mime_type", None) - raw_file_name = getattr(file_obj, "name", "FileNameNotProvided") + elif part.HasField("raw"): + mime_type = part.media_type or None + raw_file_name = part.filename or "FileNameNotProvided" file_name = self._strip_file_extension(raw_file_name) file_type = self._get_file_type_from_mime_type(mime_type) file_format = self._get_file_format_from_mime_type(mime_type, file_type) + raw_bytes = part.raw - # Handle FileWithBytes vs FileWithUri - bytes_data = getattr(file_obj, "bytes", None) - uri_data = getattr(file_obj, "uri", None) - - if bytes_data: - try: - # A2A bytes are always base64-encoded strings - decoded_bytes = base64.b64decode(bytes_data) - except Exception as e: - raise ValueError(f"Failed to decode base64 data for file '{raw_file_name}': {e}") from e - - if file_type == "image": - content_blocks.append( - ContentBlock( - image=ImageContent( - format=file_format, # type: ignore - source=ImageSource(bytes=decoded_bytes), - ) + if file_type == "image": + content_blocks.append( + ContentBlock( + image=ImageContent( + format=file_format, # type: ignore + source=ImageSource(bytes=raw_bytes), ) ) - elif file_type == "video": - content_blocks.append( - ContentBlock( - video=VideoContent( - format=file_format, # type: ignore - source=VideoSource(bytes=decoded_bytes), - ) + ) + elif file_type == "video": + content_blocks.append( + ContentBlock( + video=VideoContent( + format=file_format, # type: ignore + source=VideoSource(bytes=raw_bytes), ) ) - else: # document or unknown - content_blocks.append( - ContentBlock( - document=DocumentContent( - format=file_format, # type: ignore - name=file_name, - source=DocumentSource(bytes=decoded_bytes), - ) + ) + else: # document or unknown + content_blocks.append( + ContentBlock( + document=DocumentContent( + format=file_format, # type: ignore + name=file_name, + source=DocumentSource(bytes=raw_bytes), ) ) - # Handle FileWithUri - elif uri_data: - # For URI files, create a text representation since Strands ContentBlocks expect bytes - content_blocks.append( - ContentBlock(text=f"[File: {file_name} ({mime_type})] - Referenced file at: {uri_data}") ) - elif isinstance(part_root, DataPart): - # Handle DataPart - convert structured data to JSON text + + elif part.HasField("url"): + mime_type = part.media_type or None + raw_file_name = part.filename or "FileNameNotProvided" + file_name = self._strip_file_extension(raw_file_name) + # For URL files, create a text representation since Strands ContentBlocks expect bytes + content_blocks.append( + ContentBlock(text=f"[File: {file_name} ({mime_type})] - Referenced file at: {part.url}") + ) + + elif part.HasField("data"): + # Handle a data Part - convert structured data to JSON text try: - data_text = json.dumps(part_root.data, indent=2) + data_text = json.dumps(MessageToDict(part.data), indent=2) content_blocks.append(ContentBlock(text=f"[Structured Data]\n{data_text}")) except Exception: logger.exception("Failed to serialize data part") diff --git a/strands-py/src/strands/multiagent/a2a/server.py b/strands-py/src/strands/multiagent/a2a/server.py index db2bf89916..f4ccdc3f90 100644 --- a/strands-py/src/strands/multiagent/a2a/server.py +++ b/strands-py/src/strands/multiagent/a2a/server.py @@ -9,11 +9,11 @@ from urllib.parse import urlparse import uvicorn -from a2a.server.apps import A2AFastAPIApplication, A2AStarletteApplication from a2a.server.events import QueueManager from a2a.server.request_handlers import DefaultRequestHandler +from a2a.server.routes import add_a2a_routes_to_fastapi, create_agent_card_routes, create_jsonrpc_routes from a2a.server.tasks import InMemoryTaskStore, PushNotificationConfigStore, PushNotificationSender, TaskStore -from a2a.types import AgentCapabilities, AgentCard, AgentSkill +from a2a.types import AgentCapabilities, AgentCard, AgentInterface, AgentSkill from fastapi import FastAPI from starlette.applications import Starlette @@ -92,8 +92,9 @@ def __init__( for backwards compatibility. Defaults to False. Raises: - ValueError: If neither or both of ``agent``/``agent_factory`` are provided, or if - ``max_contexts`` is less than 1. + ValueError: If neither or both of ``agent``/``agent_factory`` are provided, if + ``max_contexts`` is less than 1, or if the agent's ``name`` or ``description`` + is None or empty (the AgentCard is validated at construction time). """ if (agent is None) == (agent_factory is None): raise ValueError("Provide exactly one of 'agent' or 'agent_factory'.") @@ -124,6 +125,8 @@ def __init__( self.name = self.strands_agent.name self.description = self.strands_agent.description self.capabilities = AgentCapabilities(streaming=True) + self._agent_skills = skills + self._agent_card_url: str | None = None self.request_handler = DefaultRequestHandler( agent_executor=StrandsA2AExecutor( agent, @@ -135,9 +138,8 @@ def __init__( queue_manager=queue_manager, push_config_store=push_config_store, push_sender=push_sender, + agent_card=self.public_agent_card, ) - self._agent_skills = skills - self._agent_card_url: str | None = None logger.info("Strands' integration with A2A is experimental. Be aware of frequent breaking changes.") def _parse_public_url(self, url: str) -> tuple[str, str]: @@ -181,7 +183,7 @@ def public_agent_card(self) -> AgentCard: return AgentCard( name=self.name, description=self.description, - url=self.agent_card_url, + supported_interfaces=[AgentInterface(protocol_binding="JSONRPC", url=self.agent_card_url)], version=self.version, skills=self.agent_skills, default_input_modes=["text"], @@ -216,6 +218,8 @@ def agent_card_url(self) -> str: def agent_card_url(self, url: str) -> None: """Override the URL advertised in the AgentCard. + Set this before calling ``to_starlette_app``/``to_fastapi_app``. + Args: url: The URL to advertise in the AgentCard. """ @@ -230,11 +234,27 @@ def agent_skills(self) -> list[AgentSkill]: def agent_skills(self, skills: list[AgentSkill]) -> None: """Set the list of skills this agent provides. + Set this before calling ``to_starlette_app``/``to_fastapi_app``. + Args: skills: A list of AgentSkill objects to set for this agent. """ self._agent_skills = skills + _V0_3_BUILD_KEYS = frozenset({"routes", "rpc_url", "agent_card_url", "extended_agent_card_url"}) + + def _validate_app_kwargs(self, app_kwargs: dict[str, Any] | None) -> None: + """Reject v0.3 ``build()`` keys that are now managed internally.""" + if not app_kwargs: + return + invalid = self._V0_3_BUILD_KEYS & app_kwargs.keys() + if invalid: + raise ValueError( + f"app_kwargs contains keys no longer accepted: {sorted(invalid)}. " + f"Use 'http_url' for path mounting (replaces rpc_url), " + f"the 'agent_card_url' setter, or add custom routes to the returned app." + ) + def to_starlette_app(self, *, app_kwargs: dict[str, Any] | None = None) -> Starlette: """Create a Starlette application for serving this agent via HTTP. @@ -243,13 +263,22 @@ def to_starlette_app(self, *, app_kwargs: dict[str, Any] | None = None) -> Starl Args: app_kwargs: Additional keyword arguments to pass to the Starlette constructor. + Must not include ``routes`` (managed internally). Keys such as ``rpc_url`` + or ``agent_card_url`` that were accepted by the v0.3 ``build()`` method are + no longer valid here — use the constructor parameters instead. Returns: Starlette: A Starlette application configured to serve this agent. + + Raises: + ValueError: If ``app_kwargs`` contains keys that were valid in v0.3's ``build()`` + but are now managed internally (``routes``, ``rpc_url``, ``agent_card_url``, + ``extended_agent_card_url``). """ - a2a_app = A2AStarletteApplication(agent_card=self.public_agent_card, http_handler=self.request_handler).build( - **app_kwargs or {} - ) + self._validate_app_kwargs(app_kwargs) + routes = create_agent_card_routes(self.public_agent_card) + routes.extend(create_jsonrpc_routes(self.request_handler, rpc_url="/", enable_v0_3_compat=True)) + a2a_app = Starlette(routes=routes, **(app_kwargs or {})) if self.mount_path: # Create parent app and mount the A2A app at the specified path @@ -268,12 +297,24 @@ def to_fastapi_app(self, *, app_kwargs: dict[str, Any] | None = None) -> FastAPI Args: app_kwargs: Additional keyword arguments to pass to the FastAPI constructor. + Must not include ``routes`` (managed internally). Keys such as ``rpc_url`` + or ``agent_card_url`` that were accepted by the v0.3 ``build()`` method are + no longer valid here — use the constructor parameters instead. Returns: FastAPI: A FastAPI application configured to serve this agent. + + Raises: + ValueError: If ``app_kwargs`` contains keys that were valid in v0.3's ``build()`` + but are now managed internally (``routes``, ``rpc_url``, ``agent_card_url``, + ``extended_agent_card_url``). """ - a2a_app = A2AFastAPIApplication(agent_card=self.public_agent_card, http_handler=self.request_handler).build( - **app_kwargs or {} + self._validate_app_kwargs(app_kwargs) + a2a_app = FastAPI(**(app_kwargs or {})) + add_a2a_routes_to_fastapi( + a2a_app, + agent_card_routes=create_agent_card_routes(self.public_agent_card), + jsonrpc_routes=create_jsonrpc_routes(self.request_handler, rpc_url="/", enable_v0_3_compat=True), ) if self.mount_path: diff --git a/strands-py/src/strands/types/a2a.py b/strands-py/src/strands/types/a2a.py index 2ca444cb04..25d57c73c0 100644 --- a/strands-py/src/strands/types/a2a.py +++ b/strands-py/src/strands/types/a2a.py @@ -1,22 +1,22 @@ """Additional A2A types.""" -from typing import Any, TypeAlias +from typing import TypeAlias -from a2a.types import Message, Task, TaskArtifactUpdateEvent, TaskStatusUpdateEvent +from a2a.types import StreamResponse from ._events import TypedEvent -A2AResponse: TypeAlias = tuple[Task, TaskStatusUpdateEvent | TaskArtifactUpdateEvent | None] | Message | Any +A2AResponse: TypeAlias = StreamResponse class A2AStreamEvent(TypedEvent): """Event emitted for every update received from the remote A2A server. - This event wraps all A2A response types during streaming, including: - - Partial task updates (TaskArtifactUpdateEvent) - - Status updates (TaskStatusUpdateEvent) - - Complete messages (Message) - - Final task completions + This event wraps every ``StreamResponse`` received during streaming, including: + - The initial ``Task`` (``task`` field) + - Partial task updates (``artifact_update`` field) + - Status updates (``status_update`` field) + - Complete messages (``message`` field) The event is emitted for EVERY update from the server, regardless of whether it represents a complete or partial response. When streaming completes, an @@ -28,7 +28,7 @@ def __init__(self, a2a_event: A2AResponse) -> None: """Initialize with A2A event. Args: - a2a_event: The original A2A event (Task tuple or Message) + a2a_event: The original A2A StreamResponse event. """ super().__init__( { diff --git a/strands-py/tests/strands/agent/test_a2a_agent.py b/strands-py/tests/strands/agent/test_a2a_agent.py index 9c3be79177..97c5072410 100644 --- a/strands-py/tests/strands/agent/test_a2a_agent.py +++ b/strands-py/tests/strands/agent/test_a2a_agent.py @@ -7,7 +7,21 @@ import pytest from a2a.client import ClientConfig -from a2a.types import AgentCard, Message, Part, Role, TaskState, TextPart +from a2a.types import ( + AgentCapabilities, + AgentCard, + AgentInterface, + Artifact, + Part, + Role, + StreamResponse, + Task, + TaskArtifactUpdateEvent, + TaskState, + TaskStatus, + TaskStatusUpdateEvent, +) +from a2a.types import Message as A2AMessage from strands.agent.a2a_agent import A2AAgent from strands.agent.agent_result import AgentResult @@ -19,9 +33,9 @@ def mock_agent_card(): return AgentCard( name="test-agent", description="Test agent", - url="http://localhost:8000", + supported_interfaces=[AgentInterface(protocol_binding="JSONRPC", url="http://localhost:8000")], version="1.0.0", - capabilities={}, + capabilities=AgentCapabilities(), default_input_modes=["text/plain"], default_output_modes=["text/plain"], skills=[], @@ -184,24 +198,32 @@ async def test_get_agent_card_preserves_custom_name_and_description(mock_agent_c @pytest.mark.asyncio async def test_get_agent_card_handles_empty_string_name_and_description(mock_httpx_client): - """Test that empty string name/description from card are preserved (not treated as None).""" - mock_card = MagicMock(spec=AgentCard) - mock_card.name = "" - mock_card.description = "" + """Test that an AgentCard with unset (empty-string) name/description leaves the agent's unset. + + Protobuf string fields have no None: an unset ``name``/``description`` reads back as ``""``, + indistinguishable from an explicit empty value. Treating it as "unset" (rather than adopting the + empty string) is the safer default for a real AgentCard, where these fields are mandatory. + """ + card = AgentCard( + name="", + description="", + supported_interfaces=[AgentInterface(protocol_binding="JSONRPC", url="http://localhost:8000")], + version="1.0.0", + capabilities=AgentCapabilities(), + ) agent = A2AAgent(endpoint="http://localhost:8000") with patch("strands.agent.a2a_agent.httpx.AsyncClient", return_value=mock_httpx_client): with patch("strands.agent.a2a_agent.A2ACardResolver") as mock_resolver_class: mock_resolver = AsyncMock() - mock_resolver.get_agent_card = AsyncMock(return_value=mock_card) + mock_resolver.get_agent_card = AsyncMock(return_value=card) mock_resolver_class.return_value = mock_resolver await agent.get_agent_card() - # Empty strings should be set (not treated as falsy/None) - assert agent.name == "" - assert agent.description == "" + assert agent.name is None + assert agent.description is None @pytest.mark.asyncio @@ -313,7 +335,7 @@ async def test_get_a2a_client_with_client_config_preserves_user_settings(mock_ag httpx_client=mock_auth_client, streaming=False, # user set this to False polling=True, - supported_transports=["jsonrpc"], + supported_protocol_bindings=["JSONRPC"], ) agent = A2AAgent(endpoint="http://localhost:8000", client_config=config) @@ -333,7 +355,7 @@ async def test_get_a2a_client_with_client_config_preserves_user_settings(mock_ag assert created_config.httpx_client is mock_auth_client assert created_config.streaming is True # overridden to True assert created_config.polling is True # preserved - assert created_config.supported_transports == ["jsonrpc"] # preserved + assert created_config.supported_protocol_bindings == ["JSONRPC"] # preserved @pytest.mark.asyncio @@ -362,7 +384,7 @@ async def test_get_a2a_client_config_without_httpx_delegates_to_factory(mock_age ClientFactory handles creating a default httpx client internally. We just pass the config with streaming=True and let the factory do its job. """ - config = ClientConfig(polling=True, supported_transports=["jsonrpc"]) + config = ClientConfig(polling=True, supported_protocol_bindings=["JSONRPC"]) agent = A2AAgent(endpoint="http://localhost:8000", client_config=config, timeout=600) with patch.object(agent, "get_agent_card", return_value=mock_agent_card): @@ -378,7 +400,7 @@ async def test_get_a2a_client_config_without_httpx_delegates_to_factory(mock_age created_config = mock_factory_class.call_args[0][0] assert created_config.streaming is True assert created_config.polling is True - assert created_config.supported_transports == ["jsonrpc"] + assert created_config.supported_protocol_bindings == ["JSONRPC"] assert created_config.httpx_client is None # factory handles default @@ -389,7 +411,7 @@ async def test_send_message_uses_provided_factory(mock_agent_card): mock_a2a_client = MagicMock() async def mock_send_message(*args, **kwargs): - yield MagicMock() + yield StreamResponse(message=A2AMessage(message_id=uuid4().hex, role=Role.ROLE_AGENT)) mock_a2a_client.send_message = mock_send_message external_factory.create.return_value = mock_a2a_client @@ -417,7 +439,7 @@ async def test_send_message_uses_client_config_httpx_client(mock_agent_card): mock_a2a_client = MagicMock() async def mock_send(*args, **kwargs): - yield MagicMock() + yield StreamResponse(message=A2AMessage(message_id=uuid4().hex, role=Role.ROLE_AGENT)) mock_a2a_client.send_message = mock_send @@ -440,10 +462,8 @@ async def mock_send(*args, **kwargs): @pytest.mark.asyncio async def test_send_message_creates_per_call_client(a2a_agent, mock_agent_card): """Test _send_message creates a fresh httpx client for each call when no factory provided.""" - mock_response = Message( - message_id=uuid4().hex, - role=Role.agent, - parts=[Part(TextPart(kind="text", text="Response"))], + mock_response = StreamResponse( + message=A2AMessage(message_id=uuid4().hex, role=Role.ROLE_AGENT, parts=[Part(text="Response")]) ) async def mock_send_message(*args, **kwargs): @@ -493,10 +513,8 @@ async def test_get_a2a_client_no_config_creates_managed_httpx(): @pytest.mark.asyncio async def test_invoke_async_success(a2a_agent, mock_agent_card): """Test successful async invocation.""" - mock_response = Message( - message_id=uuid4().hex, - role=Role.agent, - parts=[Part(TextPart(kind="text", text="Response"))], + mock_response = StreamResponse( + message=A2AMessage(message_id=uuid4().hex, role=Role.ROLE_AGENT, parts=[Part(text="Response")]) ) async def mock_send_message(*args, **kwargs): @@ -552,10 +570,8 @@ def test_call_sync(a2a_agent): @pytest.mark.asyncio async def test_stream_async_success(a2a_agent, mock_agent_card): """Test successful async streaming.""" - mock_response = Message( - message_id=uuid4().hex, - role=Role.agent, - parts=[Part(TextPart(kind="text", text="Response"))], + mock_response = StreamResponse( + message=A2AMessage(message_id=uuid4().hex, role=Role.ROLE_AGENT, parts=[Part(text="Response")]) ) async def mock_send_message(*args, **kwargs): @@ -585,292 +601,50 @@ async def test_stream_async_no_prompt(a2a_agent): pass -# === Complete Event Tests === - - -def test_is_complete_event_message(a2a_agent): - """Test _is_complete_event returns True for Message.""" - mock_message = MagicMock(spec=Message) - - assert a2a_agent._is_complete_event(mock_message) is True - - -def test_is_complete_event_tuple_with_none_update(a2a_agent): - """Test _is_complete_event returns True for tuple with None update event.""" - mock_task = MagicMock() - - assert a2a_agent._is_complete_event((mock_task, None)) is True - - -def test_is_complete_event_artifact_last_chunk(a2a_agent): - """Test _is_complete_event handles TaskArtifactUpdateEvent last_chunk flag.""" - from a2a.types import TaskArtifactUpdateEvent - - mock_task = MagicMock() - - # last_chunk=True -> complete - event_complete = MagicMock(spec=TaskArtifactUpdateEvent) - event_complete.last_chunk = True - assert a2a_agent._is_complete_event((mock_task, event_complete)) is True - - # last_chunk=False -> not complete - event_incomplete = MagicMock(spec=TaskArtifactUpdateEvent) - event_incomplete.last_chunk = False - assert a2a_agent._is_complete_event((mock_task, event_incomplete)) is False - - # last_chunk=None -> not complete - event_none = MagicMock(spec=TaskArtifactUpdateEvent) - event_none.last_chunk = None - assert a2a_agent._is_complete_event((mock_task, event_none)) is False - - -def test_is_complete_event_status_update(a2a_agent): - """Test _is_complete_event handles TaskStatusUpdateEvent state.""" - from a2a.types import TaskState, TaskStatusUpdateEvent - - mock_task = MagicMock() - - # completed state -> complete - event_completed = MagicMock(spec=TaskStatusUpdateEvent) - event_completed.status = MagicMock() - event_completed.status.state = TaskState.completed - assert a2a_agent._is_complete_event((mock_task, event_completed)) is True - - # working state -> not complete - event_working = MagicMock(spec=TaskStatusUpdateEvent) - event_working.status = MagicMock() - event_working.status.state = TaskState.working - assert a2a_agent._is_complete_event((mock_task, event_working)) is False - - # no status -> not complete - event_no_status = MagicMock(spec=TaskStatusUpdateEvent) - event_no_status.status = None - assert a2a_agent._is_complete_event((mock_task, event_no_status)) is False - - -def test_is_complete_event_unknown_type(a2a_agent): - """Test _is_complete_event returns False for unknown event types.""" - assert a2a_agent._is_complete_event("unknown") is False - - @pytest.mark.asyncio -async def test_stream_async_tracks_complete_events(a2a_agent, mock_agent_card): - """Test stream_async uses last complete event for final result.""" - from a2a.types import TaskState, TaskStatusUpdateEvent - - mock_task = MagicMock() - mock_task.artifacts = None - - # First event: incomplete - incomplete_event = MagicMock(spec=TaskStatusUpdateEvent) - incomplete_event.status = MagicMock() - incomplete_event.status.state = TaskState.working - incomplete_event.status.message = None - - # Second event: complete - complete_event = MagicMock(spec=TaskStatusUpdateEvent) - complete_event.status = MagicMock() - complete_event.status.state = TaskState.completed - complete_event.status.message = MagicMock() - complete_event.status.message.parts = [] +async def test_stream_async_no_responses_emits_no_result(a2a_agent, mock_agent_card): + """Test that stream_async emits no AgentResultEvent when the server sends nothing at all.""" async def mock_send_message(*args, **kwargs): - yield (mock_task, incomplete_event) - yield (mock_task, complete_event) + return + yield # Make it an async generator with patch.object(a2a_agent, "get_agent_card", return_value=mock_agent_card): async with mock_a2a_client_context(mock_send_message): - events = [] - async for event in a2a_agent.stream_async("Hello"): - events.append(event) + events = [event async for event in a2a_agent.stream_async("Hello")] - # Should have 2 stream events + 1 result event - assert len(events) == 3 - assert "result" in events[2] + assert events == [] @pytest.mark.asyncio -async def test_stream_async_falls_back_to_last_event(a2a_agent, mock_agent_card): - """Test stream_async falls back to last event when no complete event.""" - from a2a.types import TaskState, TaskStatusUpdateEvent - - mock_task = MagicMock() - mock_task.artifacts = None - - incomplete_event = MagicMock(spec=TaskStatusUpdateEvent) - incomplete_event.status = MagicMock() - incomplete_event.status.state = TaskState.working - incomplete_event.status.message = None +async def test_stream_async_accumulates_task_lifecycle_events(a2a_agent, mock_agent_card): + """Test stream_async accumulates content across a full task, artifact_update, status_update stream.""" + task = Task(id="t1", context_id="c1", status=TaskStatus(state=TaskState.TASK_STATE_SUBMITTED)) + artifact_event = StreamResponse( + artifact_update=TaskArtifactUpdateEvent( + task_id="t1", context_id="c1", artifact=Artifact(artifact_id="a1", parts=[Part(text="final answer")]) + ) + ) + status_event = StreamResponse( + status_update=TaskStatusUpdateEvent( + task_id="t1", context_id="c1", status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED) + ) + ) async def mock_send_message(*args, **kwargs): - yield (mock_task, incomplete_event) + yield StreamResponse(task=task) + yield artifact_event + yield status_event with patch.object(a2a_agent, "get_agent_card", return_value=mock_agent_card): async with mock_a2a_client_context(mock_send_message): - events = [] - async for event in a2a_agent.stream_async("Hello"): - events.append(event) - - # Should have 1 stream event + 1 result event (falls back to last) - assert len(events) == 2 - assert "result" in events[1] - - -# ========================================================================= -# NEW TESTS: Client-side lifecycle state handling -# ========================================================================= - - -def test_is_complete_event_failed_state(a2a_agent): - """Test that failed state is recognized as complete.""" - from unittest.mock import MagicMock - - from a2a.types import TaskState, TaskStatusUpdateEvent - - task = MagicMock() - status = MagicMock() - status.state = TaskState.failed - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status - - assert a2a_agent._is_complete_event((task, update_event)) is True - - -def test_is_complete_event_canceled_state(a2a_agent): - """Test that canceled state is recognized as complete.""" - from unittest.mock import MagicMock - - from a2a.types import TaskState, TaskStatusUpdateEvent - - task = MagicMock() - status = MagicMock() - status.state = TaskState.canceled - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status - - assert a2a_agent._is_complete_event((task, update_event)) is True - - -def test_is_complete_event_rejected_state(a2a_agent): - """Test that rejected state is recognized as complete.""" - from unittest.mock import MagicMock - - from a2a.types import TaskState, TaskStatusUpdateEvent - - task = MagicMock() - status = MagicMock() - status.state = TaskState.rejected - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status - - assert a2a_agent._is_complete_event((task, update_event)) is True - - -def test_is_complete_event_input_required_state(a2a_agent): - """Test that input_required state is recognized as complete (pausing).""" - from unittest.mock import MagicMock - - from a2a.types import TaskState, TaskStatusUpdateEvent - - task = MagicMock() - status = MagicMock() - status.state = TaskState.input_required - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status - - assert a2a_agent._is_complete_event((task, update_event)) is True - - -def test_is_complete_event_auth_required_state(a2a_agent): - """Test that auth_required state is recognized as complete (pausing).""" - from unittest.mock import MagicMock - - from a2a.types import TaskState, TaskStatusUpdateEvent - - task = MagicMock() - status = MagicMock() - status.state = TaskState.auth_required - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status - - assert a2a_agent._is_complete_event((task, update_event)) is True - - -def test_is_complete_event_working_state_not_complete(a2a_agent): - """Test that working state is NOT recognized as complete.""" - from unittest.mock import MagicMock - - from a2a.types import TaskState, TaskStatusUpdateEvent - - task = MagicMock() - status = MagicMock() - status.state = TaskState.working - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status - - assert a2a_agent._is_complete_event((task, update_event)) is False - - -def test_is_complete_event_submitted_state_not_complete(a2a_agent): - """Test that submitted state is NOT recognized as complete.""" - from unittest.mock import MagicMock - - from a2a.types import TaskState, TaskStatusUpdateEvent - - task = MagicMock() - status = MagicMock() - status.state = TaskState.submitted - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status - - assert a2a_agent._is_complete_event((task, update_event)) is False - - -# ========================================================================= -# DEVIL'S ADVOCATE FINDINGS — Tests addressing review gaps -# ========================================================================= - - -@pytest.mark.parametrize( - "state,expected_complete", - [ - (TaskState.completed, True), - (TaskState.failed, True), - (TaskState.canceled, True), - (TaskState.rejected, True), - (TaskState.input_required, True), - (TaskState.auth_required, True), - (TaskState.working, False), - (TaskState.submitted, False), - (TaskState.unknown, False), - ], - ids=[ - "completed-is-complete", - "failed-is-complete", - "canceled-is-complete", - "rejected-is-complete", - "input_required-is-complete", - "auth_required-is-complete", - "working-not-complete", - "submitted-not-complete", - "unknown-not-complete", - ], -) -def test_is_complete_event_all_states_parametrized(a2a_agent, state, expected_complete): - """Minor Finding 7: Parametrized test covering ALL TaskState values. - - This replaces verbose individual tests with a single parameterized test that - covers all 9 TaskState values. When a2a-sdk adds new states, adding a row here - is trivial. - """ - from unittest.mock import MagicMock - - from a2a.types import TaskStatusUpdateEvent - - task = MagicMock() - status = MagicMock() - status.state = state - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status - - assert a2a_agent._is_complete_event((task, update_event)) is expected_complete + events = [event async for event in a2a_agent.stream_async("Hello")] + + # 3 stream events + 1 final result event + assert len(events) == 4 + assert all(event.get("type") == "a2a_stream" for event in events[:3]) + assert "result" in events[3] + result = events[3]["result"] + assert result.stop_reason == "end_turn" + assert result.message["content"] == [{"text": "final answer"}] + assert result.state.get("a2a_task_state") == "completed" diff --git a/strands-py/tests/strands/multiagent/a2a/test_converters.py b/strands-py/tests/strands/multiagent/a2a/test_converters.py index fff48653bf..7cde80a4e3 100644 --- a/strands-py/tests/strands/multiagent/a2a/test_converters.py +++ b/strands-py/tests/strands/multiagent/a2a/test_converters.py @@ -1,17 +1,28 @@ """Tests for A2A converter functions.""" -from unittest.mock import MagicMock from uuid import uuid4 import pytest +from a2a.types import ( + Artifact, + Part, + Role, + StreamResponse, + Task, + TaskArtifactUpdateEvent, + TaskState, + TaskStatus, + TaskStatusUpdateEvent, +) from a2a.types import Message as A2AMessage -from a2a.types import Part, Role, TaskArtifactUpdateEvent, TaskStatusUpdateEvent, TextPart from strands.agent.agent_result import AgentResult from strands.multiagent.a2a._converters import ( + _extract_task_state, + _parts_to_content, convert_content_blocks_to_parts, convert_input_to_message, - convert_response_to_agent_result, + convert_responses_to_agent_result, ) @@ -20,9 +31,9 @@ def test_convert_string_input(): message = convert_input_to_message("Hello") assert isinstance(message, A2AMessage) - assert message.role == Role.user + assert message.role == Role.ROLE_USER assert len(message.parts) == 1 - assert message.parts[0].root.text == "Hello" + assert message.parts[0].text == "Hello" def test_convert_message_list_input(): @@ -34,7 +45,7 @@ def test_convert_message_list_input(): message = convert_input_to_message(messages) assert isinstance(message, A2AMessage) - assert message.role == Role.user + assert message.role == Role.ROLE_USER assert len(message.parts) == 1 @@ -62,6 +73,19 @@ def test_convert_interrupt_response_raises_error(): convert_input_to_message(interrupt_responses) +def test_convert_message_list_finds_last_user_message(): + """Test that message list conversion finds the last user message.""" + messages = [ + {"role": "user", "content": [{"text": "First"}]}, + {"role": "assistant", "content": [{"text": "Response"}]}, + {"role": "user", "content": [{"text": "Second"}]}, + ] + + message = convert_input_to_message(messages) + + assert message.parts[0].text == "Second" + + def test_convert_content_blocks_to_parts(): """Test converting content blocks to A2A parts.""" content_blocks = [{"text": "Hello"}, {"text": "World"}] @@ -69,19 +93,28 @@ def test_convert_content_blocks_to_parts(): parts = convert_content_blocks_to_parts(content_blocks) assert len(parts) == 2 - assert parts[0].root.text == "Hello" - assert parts[1].root.text == "World" + assert parts[0].text == "Hello" + assert parts[1].text == "World" + + +def test_convert_content_blocks_skips_non_text(): + """Test that non-text content blocks are skipped.""" + content_blocks = [{"text": "Hello"}, {"image": "data"}, {"text": "World"}] + + parts = convert_content_blocks_to_parts(content_blocks) + + assert len(parts) == 2 def test_convert_a2a_message_response(): - """Test converting A2A message response to AgentResult.""" + """Test converting a bare A2A Message response to AgentResult.""" a2a_message = A2AMessage( message_id=uuid4().hex, - role=Role.agent, - parts=[Part(TextPart(kind="text", text="Response"))], + role=Role.ROLE_AGENT, + parts=[Part(text="Response")], ) - result = convert_response_to_agent_result(a2a_message) + result = convert_responses_to_agent_result([StreamResponse(message=a2a_message)]) assert isinstance(result, AgentResult) assert result.message["role"] == "assistant" @@ -89,417 +122,338 @@ def test_convert_a2a_message_response(): assert result.message["content"][0]["text"] == "Response" -def test_convert_task_response(): - """Test converting task response to AgentResult.""" - mock_task = MagicMock() - mock_artifact = MagicMock() - mock_part = MagicMock() - mock_part.root.text = "Task response" - mock_artifact.parts = [mock_part] - mock_task.artifacts = [mock_artifact] - - result = convert_response_to_agent_result((mock_task, None)) - - assert isinstance(result, AgentResult) - assert len(result.message["content"]) == 1 - assert result.message["content"][0]["text"] == "Task response" - - def test_convert_multiple_parts_response(): """Test converting response with multiple parts to separate content blocks.""" a2a_message = A2AMessage( message_id=uuid4().hex, - role=Role.agent, - parts=[ - Part(TextPart(kind="text", text="First")), - Part(TextPart(kind="text", text="Second")), - ], + role=Role.ROLE_AGENT, + parts=[Part(text="First"), Part(text="Second")], ) - result = convert_response_to_agent_result(a2a_message) + result = convert_responses_to_agent_result([StreamResponse(message=a2a_message)]) assert len(result.message["content"]) == 2 assert result.message["content"][0]["text"] == "First" assert result.message["content"][1]["text"] == "Second" -# --- New tests for coverage --- - - -def test_convert_message_list_finds_last_user_message(): - """Test that message list conversion finds the last user message.""" - messages = [ - {"role": "user", "content": [{"text": "First"}]}, - {"role": "assistant", "content": [{"text": "Response"}]}, - {"role": "user", "content": [{"text": "Second"}]}, - ] - - message = convert_input_to_message(messages) - - assert message.parts[0].root.text == "Second" - - -def test_convert_content_blocks_skips_non_text(): - """Test that non-text content blocks are skipped.""" - content_blocks = [{"text": "Hello"}, {"image": "data"}, {"text": "World"}] +def test_convert_bare_task_response_uses_task_artifacts(): + """Test that a bare `task` response (no separate update event) extracts task.artifacts.""" + task = Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_SUBMITTED), + artifacts=[Artifact(artifact_id="a1", parts=[Part(text="Task response")])], + ) - parts = convert_content_blocks_to_parts(content_blocks) + result = convert_responses_to_agent_result([StreamResponse(task=task)]) - assert len(parts) == 2 + assert isinstance(result, AgentResult) + assert len(result.message["content"]) == 1 + assert result.message["content"][0]["text"] == "Task response" def test_convert_task_artifact_update_event(): - """Test converting TaskArtifactUpdateEvent to AgentResult.""" - mock_task = MagicMock() - mock_part = MagicMock() - mock_part.root.text = "Streamed artifact" - mock_artifact = MagicMock() - mock_artifact.parts = [mock_part] - - mock_event = MagicMock(spec=TaskArtifactUpdateEvent) - mock_event.artifact = mock_artifact + """Test converting a TaskArtifactUpdateEvent response to AgentResult.""" + event = TaskArtifactUpdateEvent( + task_id="task-1", + context_id="ctx-1", + artifact=Artifact(artifact_id="a1", parts=[Part(text="Streamed artifact")]), + ) - result = convert_response_to_agent_result((mock_task, mock_event)) + result = convert_responses_to_agent_result([StreamResponse(artifact_update=event)]) assert result.message["content"][0]["text"] == "Streamed artifact" -def test_convert_task_status_update_event(): - """Test converting TaskStatusUpdateEvent to AgentResult.""" - mock_task = MagicMock() - mock_part = MagicMock() - mock_part.root.text = "Status message" - mock_message = MagicMock() - mock_message.parts = [mock_part] - mock_status = MagicMock() - mock_status.message = mock_message - - mock_event = MagicMock(spec=TaskStatusUpdateEvent) - mock_event.status = mock_status +def test_convert_task_status_update_event_with_message(): + """Test converting a TaskStatusUpdateEvent with a message to AgentResult.""" + event = TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus( + state=TaskState.TASK_STATE_FAILED, + message=A2AMessage(message_id=uuid4().hex, role=Role.ROLE_AGENT, parts=[Part(text="Status message")]), + ), + ) - result = convert_response_to_agent_result((mock_task, mock_event)) + result = convert_responses_to_agent_result([StreamResponse(status_update=event)]) assert result.message["content"][0]["text"] == "Status message" -def test_convert_task_status_update_event_no_message_falls_back_to_task_artifacts(): - """Test that TaskStatusUpdateEvent with no message falls back to task.artifacts.""" - mock_task = MagicMock() - mock_part = MagicMock() - mock_part.root.text = "Artifact content" - mock_artifact = MagicMock() - mock_artifact.parts = [mock_part] - mock_task.artifacts = [mock_artifact] - - mock_event = MagicMock(spec=TaskStatusUpdateEvent) - mock_status = MagicMock() - mock_status.message = None - mock_event.status = mock_status - - result = convert_response_to_agent_result((mock_task, mock_event)) - - assert len(result.message["content"]) == 1 - assert result.message["content"][0]["text"] == "Artifact content" - - -def test_convert_task_artifact_update_event_empty_parts_falls_back_to_task_artifacts(): - """Test that TaskArtifactUpdateEvent with empty parts falls back to task.artifacts.""" - mock_task = MagicMock() - mock_part = MagicMock() - mock_part.root.text = "Full artifact content" - mock_artifact = MagicMock() - mock_artifact.parts = [mock_part] - mock_task.artifacts = [mock_artifact] - - mock_event = MagicMock(spec=TaskArtifactUpdateEvent) - mock_event_artifact = MagicMock() - mock_event_artifact.parts = [] - mock_event.artifact = mock_event_artifact - - result = convert_response_to_agent_result((mock_task, mock_event)) - - assert len(result.message["content"]) == 1 - assert result.message["content"][0]["text"] == "Full artifact content" - - -def test_convert_response_handles_missing_data(): - """Test that response conversion handles missing/malformed data gracefully.""" - # TaskArtifactUpdateEvent with no artifact - mock_event = MagicMock(spec=TaskArtifactUpdateEvent) - mock_event.artifact = None - result = convert_response_to_agent_result((MagicMock(), mock_event)) - assert len(result.message["content"]) == 0 - - # TaskStatusUpdateEvent with no status - mock_event = MagicMock(spec=TaskStatusUpdateEvent) - mock_event.status = None - result = convert_response_to_agent_result((MagicMock(), mock_event)) - assert len(result.message["content"]) == 0 +def test_convert_status_update_without_message_has_no_content(): + """A terminal status_update with no message and no prior artifact yields empty content. - # Task artifact without parts attribute - mock_task = MagicMock() - mock_artifact = MagicMock(spec=[]) - del mock_artifact.parts - mock_task.artifacts = [mock_artifact] - result = convert_response_to_agent_result((mock_task, None)) - assert len(result.message["content"]) == 0 - - -# ========================================================================= -# NEW TESTS: Lifecycle State Mapping -# ========================================================================= - - -def test_convert_response_completed_state_maps_to_end_turn(): - """Test that completed state maps to end_turn stop_reason.""" - from unittest.mock import MagicMock - - from a2a.types import TaskState, TaskStatus, TaskStatusUpdateEvent - - task = MagicMock() - task.artifacts = None + This is the streaming pattern where the final content already arrived via a separate + artifact_update event earlier in the same call — covered by + test_artifact_update_then_status_update_does_not_duplicate below. + """ + event = TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ) - status = TaskStatus(state=TaskState.completed, message=None) - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status + result = convert_responses_to_agent_result([StreamResponse(status_update=event)]) - result = convert_response_to_agent_result((task, update_event)) + assert result.message["content"] == [] assert result.stop_reason == "end_turn" -def test_convert_response_failed_state_maps_to_end_turn(): - """Test that failed state maps to end_turn stop_reason with error content.""" - from unittest.mock import MagicMock - - from a2a.types import Message, TaskState, TaskStatus, TaskStatusUpdateEvent +def test_artifact_update_then_status_update_does_not_duplicate(): + """The common non-compliant-streaming pattern: artifact carries content, terminal status has none. - task = MagicMock() - task.artifacts = None + Regression test: summing artifact_update content across the stream must not also pull in the + terminal status_update's (nonexistent) message, and must not double-count. + """ + artifact_event = StreamResponse( + artifact_update=TaskArtifactUpdateEvent( + task_id="task-1", + context_id="ctx-1", + artifact=Artifact(artifact_id="a1", parts=[Part(text="final answer")]), + ) + ) + status_event = StreamResponse( + status_update=TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ) + ) - # Create a status message with error info - error_part = MagicMock() - error_part.root = MagicMock() - error_part.root.text = "Agent execution failed: timeout" + result = convert_responses_to_agent_result([artifact_event, status_event]) + + assert result.message["content"] == [{"text": "final answer"}] + assert result.state.get("a2a_task_state") == "completed" + + +def test_multiple_artifact_updates_accumulate_in_order(): + """Multiple artifact_update chunks (compliant-streaming deltas) accumulate in stream order.""" + responses = [ + StreamResponse( + artifact_update=TaskArtifactUpdateEvent( + task_id="task-1", + context_id="ctx-1", + artifact=Artifact(artifact_id="a1", parts=[Part(text="Hello, ")]), + ) + ), + StreamResponse( + artifact_update=TaskArtifactUpdateEvent( + task_id="task-1", + context_id="ctx-1", + artifact=Artifact(artifact_id="a1", parts=[Part(text="world!")]), + append=True, + last_chunk=True, + ) + ), + ] - error_message = MagicMock(spec=Message) - error_message.parts = [error_part] + result = convert_responses_to_agent_result(responses) - status = TaskStatus(state=TaskState.failed, message=error_message) - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status + assert result.message["content"] == [{"text": "Hello, "}, {"text": "world!"}] - result = convert_response_to_agent_result((task, update_event)) - assert result.stop_reason == "end_turn" - assert result.state.get("a2a_task_state") == "failed" - assert "Agent execution failed" in result.message["content"][0]["text"] +def test_artifact_replace_does_not_duplicate_on_resend(): + """A peer that re-sends its full cumulative artifact each turn (append=False) must not duplicate. -def test_convert_response_input_required_maps_to_interrupt(): - """Test that input_required state maps to interrupt stop_reason.""" - from unittest.mock import MagicMock + append=False (the default) means "replace this artifact's content", not "add to it". + """ + responses = [ + StreamResponse( + artifact_update=TaskArtifactUpdateEvent( + task_id="task-1", context_id="ctx-1", artifact=Artifact(artifact_id="a1", parts=[Part(text="Hel")]) + ) + ), + StreamResponse( + artifact_update=TaskArtifactUpdateEvent( + task_id="task-1", context_id="ctx-1", artifact=Artifact(artifact_id="a1", parts=[Part(text="Hello")]) + ) + ), + StreamResponse( + artifact_update=TaskArtifactUpdateEvent( + task_id="task-1", + context_id="ctx-1", + artifact=Artifact(artifact_id="a1", parts=[Part(text="Hello world")]), + ) + ), + ] - from a2a.types import Message, TaskState, TaskStatus, TaskStatusUpdateEvent + result = convert_responses_to_agent_result(responses) - task = MagicMock() - task.artifacts = None + assert result.message["content"] == [{"text": "Hello world"}] - input_part = MagicMock() - input_part.root = MagicMock() - input_part.root.text = "Agent requires input:\n- approval: Need confirmation" - input_message = MagicMock(spec=Message) - input_message.parts = [input_part] +def test_terminal_status_message_appended_after_artifact_content(): + """An actionable terminal status message (e.g. an approval prompt) is not dropped when + artifact content already streamed — both are surfaced, in order. + """ + artifact_event = StreamResponse( + artifact_update=TaskArtifactUpdateEvent( + task_id="task-1", + context_id="ctx-1", + artifact=Artifact(artifact_id="a1", parts=[Part(text="partial answer")]), + ) + ) + status_event = StreamResponse( + status_update=TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus( + state=TaskState.TASK_STATE_INPUT_REQUIRED, + message=A2AMessage(message_id=uuid4().hex, role=Role.ROLE_AGENT, parts=[Part(text="need approval")]), + ), + ) + ) - status = TaskStatus(state=TaskState.input_required, message=input_message) - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status + result = convert_responses_to_agent_result([artifact_event, status_event]) - result = convert_response_to_agent_result((task, update_event)) + assert result.message["content"] == [{"text": "partial answer"}, {"text": "need approval"}] assert result.stop_reason == "interrupt" - assert result.state.get("a2a_task_state") == "input-required" - assert "approval" in result.message["content"][0]["text"] -def test_convert_response_canceled_state_maps_to_end_turn(): - """Test that canceled state maps to end_turn stop_reason.""" - from unittest.mock import MagicMock +def test_parts_to_content_drops_empty_text_parts(): + """An empty-text part (the compliant-streaming last_chunk marker) yields no content block.""" + assert _parts_to_content([Part(text="")]) == [] + assert _parts_to_content([Part(text="real"), Part(text="")]) == [{"text": "real"}] - from a2a.types import TaskState, TaskStatus, TaskStatusUpdateEvent - task = MagicMock() - task.artifacts = None +def test_task_without_status_does_not_reset_observed_state(): + """A bare `task` snapshot with no status field must not overwrite an already-observed state. - status = TaskStatus(state=TaskState.canceled, message=None) - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status + `task.status.state` reads as TASK_STATE_UNSPECIFIED (0) when `status` was never set, which + is a real enum value rather than "no state" — extraction must gate on HasField("status"). + """ + status_event = StreamResponse( + status_update=TaskStatusUpdateEvent( + task_id="task-1", context_id="ctx-1", status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED) + ) + ) + task_without_status = StreamResponse(task=Task(id="task-1", context_id="ctx-1")) - result = convert_response_to_agent_result((task, update_event)) - assert result.stop_reason == "end_turn" - assert result.state.get("a2a_task_state") == "canceled" + result = convert_responses_to_agent_result([status_event, task_without_status]) + assert result.state.get("a2a_task_state") == "completed" -def test_convert_response_rejected_state_maps_to_end_turn(): - """Test that rejected state maps to end_turn stop_reason.""" - from unittest.mock import MagicMock - from a2a.types import TaskState, TaskStatus, TaskStatusUpdateEvent +# ========================================================================= +# Lifecycle state mapping +# ========================================================================= - task = MagicMock() - task.artifacts = None - status = TaskStatus(state=TaskState.rejected, message=None) - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status +@pytest.mark.parametrize( + ("task_state", "expected_stop_reason", "expected_state_str"), + [ + (TaskState.TASK_STATE_COMPLETED, "end_turn", "completed"), + (TaskState.TASK_STATE_FAILED, "end_turn", "failed"), + (TaskState.TASK_STATE_CANCELED, "end_turn", "canceled"), + (TaskState.TASK_STATE_REJECTED, "end_turn", "rejected"), + (TaskState.TASK_STATE_INPUT_REQUIRED, "interrupt", "input-required"), + (TaskState.TASK_STATE_AUTH_REQUIRED, "interrupt", "auth-required"), + (TaskState.TASK_STATE_WORKING, "end_turn", "working"), + (TaskState.TASK_STATE_SUBMITTED, "end_turn", "submitted"), + (TaskState.TASK_STATE_UNSPECIFIED, "end_turn", "unknown"), + ], +) +def test_convert_response_state_mapping(task_state, expected_stop_reason, expected_state_str): + """Test that each TaskState maps to the documented stop_reason and state string.""" + event = TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=task_state), + ) - result = convert_response_to_agent_result((task, update_event)) - assert result.stop_reason == "end_turn" - assert result.state.get("a2a_task_state") == "rejected" + result = convert_responses_to_agent_result([StreamResponse(status_update=event)]) + assert result.stop_reason == expected_stop_reason + assert result.state.get("a2a_task_state") == expected_state_str -def test_convert_response_auth_required_maps_to_interrupt(): - """Test that auth_required state maps to interrupt stop_reason.""" - from unittest.mock import MagicMock - from a2a.types import TaskState, TaskStatus, TaskStatusUpdateEvent +def test_convert_response_no_events_yields_end_turn_and_no_state(): + """An empty response list defaults to end_turn with no a2a_task_state entry.""" + result = convert_responses_to_agent_result([]) - task = MagicMock() - task.artifacts = None + assert result.stop_reason == "end_turn" + assert result.message["content"] == [] + assert "a2a_task_state" not in result.state - status = TaskStatus(state=TaskState.auth_required, message=None) - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status - result = convert_response_to_agent_result((task, update_event)) - assert result.stop_reason == "interrupt" - assert result.state.get("a2a_task_state") == "auth-required" +def test_extract_task_state_from_status_update(): + """Test _extract_task_state helper on a status_update response.""" + event = TaskStatusUpdateEvent(task_id="t", context_id="c", status=TaskStatus(state=TaskState.TASK_STATE_FAILED)) + state = _extract_task_state(StreamResponse(status_update=event)) -def test_extract_task_state_from_status_update(): - """Test _extract_task_state helper.""" - from unittest.mock import MagicMock + assert state == TaskState.TASK_STATE_FAILED - from a2a.types import TaskState, TaskStatus, TaskStatusUpdateEvent - from strands.multiagent.a2a._converters import _extract_task_state +def test_extract_task_state_from_task(): + """Test _extract_task_state helper on a bare task response.""" + task = Task(id="t", context_id="c", status=TaskStatus(state=TaskState.TASK_STATE_SUBMITTED)) - task = MagicMock() - status = TaskStatus(state=TaskState.failed, message=None) - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status + state = _extract_task_state(StreamResponse(task=task)) - state = _extract_task_state((task, update_event)) - assert state == TaskState.failed + assert state == TaskState.TASK_STATE_SUBMITTED def test_extract_task_state_from_message_returns_none(): """Test _extract_task_state returns None for Message responses.""" - from unittest.mock import MagicMock + message = A2AMessage(message_id=uuid4().hex, role=Role.ROLE_AGENT, parts=[Part(text="hi")]) - from a2a.types import Message + state = _extract_task_state(StreamResponse(message=message)) - from strands.multiagent.a2a._converters import _extract_task_state - - message = MagicMock(spec=Message) - state = _extract_task_state(message) assert state is None -# ========================================================================= -# DEVIL'S ADVOCATE FINDINGS — Tests addressing review gaps -# ========================================================================= - - -def test_convert_response_completed_state_includes_state_metadata(): - """Major Finding 3: The completed state test was missing state assertion. - - Every other state test asserts both stop_reason AND result.state, but the most - important one (completed — the happy path) was missing the state check. This ensures - downstream consumers relying on result.state["a2a_task_state"] won't break silently. - """ - from unittest.mock import MagicMock - - from a2a.types import TaskState, TaskStatus, TaskStatusUpdateEvent - - task = MagicMock() - task.artifacts = None - - status = TaskStatus(state=TaskState.completed, message=None) - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status - - result = convert_response_to_agent_result((task, update_event)) - assert result.stop_reason == "end_turn" - assert result.state.get("a2a_task_state") == "completed" # THIS WAS MISSING - - -def test_convert_response_unknown_state_defaults_to_end_turn(): - """Major Finding 4: TaskState.unknown should default to end_turn. - - The a2a-sdk has a TaskState.unknown value. Our code handles it via the .get() - default ("end_turn"). This test documents that this is an intentional design - decision: unknown states are treated as terminal completions rather than errors. - - Rationale: An unknown state from a remote server is ambiguous. Treating it as - end_turn (completed) is the safest default — the client won't hang waiting for - more events, and the result content (if any) is still accessible. - """ - from unittest.mock import MagicMock - - from a2a.types import TaskState, TaskStatus, TaskStatusUpdateEvent - - task = MagicMock() - task.artifacts = None +def test_extract_task_state_from_artifact_update_returns_none(): + """_extract_task_state returns None for artifact_update responses (they carry no state).""" + event = TaskArtifactUpdateEvent( + task_id="t", context_id="c", artifact=Artifact(artifact_id="a1", parts=[Part(text="x")]) + ) - status = TaskStatus(state=TaskState.unknown, message=None) - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status + state = _extract_task_state(StreamResponse(artifact_update=event)) - result = convert_response_to_agent_result((task, update_event)) - # unknown is NOT in _STATE_TO_STOP_REASON, so defaults to "end_turn" - assert result.stop_reason == "end_turn" - # state metadata should reflect the actual state value - assert result.state.get("a2a_task_state") == "unknown" + assert state is None -def test_convert_response_working_state_defaults_to_end_turn(): - """Test that working state (not in mapping) defaults to end_turn. +def test_task_with_no_artifacts_falls_back_to_status_message(): + """A Task whose answer lives in task.status.message (no artifacts) still produces content. - This covers the edge case where a TaskStatusUpdateEvent with state=working - somehow reaches the converter (shouldn't normally happen since _is_complete_event - filters these out, but defense-in-depth). + Third-party A2A servers may reply with a completed Task carrying text only in + task.status.message — the spec does not require artifacts. """ - from unittest.mock import MagicMock - - from a2a.types import TaskState, TaskStatus, TaskStatusUpdateEvent - - task = MagicMock() - task.artifacts = None + task = Task( + id="t1", + context_id="c1", + status=TaskStatus( + state=TaskState.TASK_STATE_COMPLETED, + message=A2AMessage(message_id=uuid4().hex, role=Role.ROLE_AGENT, parts=[Part(text="the answer")]), + ), + ) - status = TaskStatus(state=TaskState.working, message=None) - update_event = MagicMock(spec=TaskStatusUpdateEvent) - update_event.status = status + result = convert_responses_to_agent_result([StreamResponse(task=task)]) - result = convert_response_to_agent_result((task, update_event)) assert result.stop_reason == "end_turn" - assert result.state.get("a2a_task_state") == "working" - - -def test_extract_task_state_from_artifact_update_returns_none(): - """Minor Finding 5: _extract_task_state with TaskArtifactUpdateEvent returns None. - - This is the untested path where the update event is an artifact (not status). - """ - from unittest.mock import MagicMock - - from a2a.types import TaskArtifactUpdateEvent - - from strands.multiagent.a2a._converters import _extract_task_state + assert any(block.get("text") == "the answer" for block in result.message["content"]) + + +def test_task_with_artifacts_ignores_status_message(): + """When a Task has artifacts, task.status.message is not used (artifacts take precedence).""" + task = Task( + id="t1", + context_id="c1", + status=TaskStatus( + state=TaskState.TASK_STATE_COMPLETED, + message=A2AMessage(message_id=uuid4().hex, role=Role.ROLE_AGENT, parts=[Part(text="status text")]), + ), + artifacts=[Artifact(artifact_id="a1", parts=[Part(text="artifact text")])], + ) - task = MagicMock() - mock_event = MagicMock(spec=TaskArtifactUpdateEvent) + result = convert_responses_to_agent_result([StreamResponse(task=task)]) - state = _extract_task_state((task, mock_event)) - assert state is None + content_texts = [block["text"] for block in result.message["content"] if "text" in block] + assert "artifact text" in content_texts + assert "status text" not in content_texts def test_state_to_stop_reason_covers_all_lifecycle_states(): @@ -507,22 +461,20 @@ def test_state_to_stop_reason_covers_all_lifecycle_states(): Guards against future additions to the a2a-sdk that we miss. """ - from a2a.types import TaskState - from strands.multiagent.a2a._converters import _STATE_TO_STOP_REASON # These are the states we explicitly handle expected_mapped = { - TaskState.completed, - TaskState.failed, - TaskState.canceled, - TaskState.rejected, - TaskState.input_required, - TaskState.auth_required, + TaskState.TASK_STATE_COMPLETED, + TaskState.TASK_STATE_FAILED, + TaskState.TASK_STATE_CANCELED, + TaskState.TASK_STATE_REJECTED, + TaskState.TASK_STATE_INPUT_REQUIRED, + TaskState.TASK_STATE_AUTH_REQUIRED, } assert set(_STATE_TO_STOP_REASON.keys()) == expected_mapped # These should NOT be in the mapping (they're non-terminal progress states) - assert TaskState.working not in _STATE_TO_STOP_REASON - assert TaskState.submitted not in _STATE_TO_STOP_REASON - assert TaskState.unknown not in _STATE_TO_STOP_REASON + assert TaskState.TASK_STATE_WORKING not in _STATE_TO_STOP_REASON + assert TaskState.TASK_STATE_SUBMITTED not in _STATE_TO_STOP_REASON + assert TaskState.TASK_STATE_UNSPECIFIED not in _STATE_TO_STOP_REASON diff --git a/strands-py/tests/strands/multiagent/a2a/test_executor.py b/strands-py/tests/strands/multiagent/a2a/test_executor.py index 3d65605b75..cd0f320a8e 100644 --- a/strands-py/tests/strands/multiagent/a2a/test_executor.py +++ b/strands-py/tests/strands/multiagent/a2a/test_executor.py @@ -1,12 +1,13 @@ """Tests for the StrandsA2AExecutor class.""" -import base64 from typing import Any from unittest.mock import AsyncMock, MagicMock, patch import pytest -from a2a.types import DataPart, FilePart, InternalError, InvalidParamsError, TextPart, UnsupportedOperationError -from a2a.utils.errors import ServerError +from a2a.helpers import new_data_part +from a2a.types import Part +from a2a.utils.errors import InternalError, InvalidParamsError, UnsupportedOperationError +from google.protobuf.json_format import MessageToDict from strands.agent.agent_result import AgentResult as SAAgentResult from strands.multiagent.a2a.executor import StrandsA2AExecutor, _StreamState @@ -89,18 +90,10 @@ def test_strip_file_extension(): def test_convert_a2a_parts_to_content_blocks_text_part(): - """Test conversion of TextPart to ContentBlock.""" - from a2a.types import TextPart - + """Test conversion of a text Part to ContentBlock.""" executor = StrandsA2AExecutor(MagicMock()) - # Mock TextPart with proper spec - text_part = MagicMock(spec=TextPart) - text_part.text = "Hello, world!" - - # Mock Part with TextPart root - part = MagicMock() - part.root = text_part + part = Part(text="Hello, world!") result = executor._convert_a2a_parts_to_content_blocks([part]) @@ -109,25 +102,10 @@ def test_convert_a2a_parts_to_content_blocks_text_part(): def test_convert_a2a_parts_to_content_blocks_file_part_image_bytes(): - """Test conversion of FilePart with image bytes to ContentBlock.""" + """Test conversion of a raw Part with image bytes to ContentBlock.""" executor = StrandsA2AExecutor(MagicMock()) - base64_bytes = base64.b64encode(VALID_PNG_BYTES).decode("utf-8") - - # Mock file object - file_obj = MagicMock() - file_obj.name = "test_image.png" - file_obj.mime_type = "image/png" - file_obj.bytes = base64_bytes - file_obj.uri = None - - # Mock FilePart with proper spec - file_part = MagicMock(spec=FilePart) - file_part.file = file_obj - - # Mock Part with FilePart root - part = MagicMock() - part.root = file_part + part = Part(raw=VALID_PNG_BYTES, media_type="image/png", filename="test_image.png") result = executor._convert_a2a_parts_to_content_blocks([part]) @@ -139,25 +117,10 @@ def test_convert_a2a_parts_to_content_blocks_file_part_image_bytes(): def test_convert_a2a_parts_to_content_blocks_file_part_video_bytes(): - """Test conversion of FilePart with video bytes to ContentBlock.""" + """Test conversion of a raw Part with video bytes to ContentBlock.""" executor = StrandsA2AExecutor(MagicMock()) - base64_bytes = base64.b64encode(VALID_MP4_BYTES).decode("utf-8") - - # Mock file object - file_obj = MagicMock() - file_obj.name = "test_video.mp4" - file_obj.mime_type = "video/mp4" - file_obj.bytes = base64_bytes - file_obj.uri = None - - # Mock FilePart with proper spec - file_part = MagicMock(spec=FilePart) - file_part.file = file_obj - - # Mock Part with FilePart root - part = MagicMock() - part.root = file_part + part = Part(raw=VALID_MP4_BYTES, media_type="video/mp4", filename="test_video.mp4") result = executor._convert_a2a_parts_to_content_blocks([part]) @@ -169,25 +132,10 @@ def test_convert_a2a_parts_to_content_blocks_file_part_video_bytes(): def test_convert_a2a_parts_to_content_blocks_file_part_document_bytes(): - """Test conversion of FilePart with document bytes to ContentBlock.""" + """Test conversion of a raw Part with document bytes to ContentBlock.""" executor = StrandsA2AExecutor(MagicMock()) - base64_bytes = base64.b64encode(VALID_DOCUMENT_BYTES).decode("utf-8") - - # Mock file object - file_obj = MagicMock() - file_obj.name = "test_document.pdf" - file_obj.mime_type = "application/pdf" - file_obj.bytes = base64_bytes - file_obj.uri = None - - # Mock FilePart with proper spec - file_part = MagicMock(spec=FilePart) - file_part.file = file_obj - - # Mock Part with FilePart root - part = MagicMock() - part.root = file_part + part = Part(raw=VALID_DOCUMENT_BYTES, media_type="application/pdf", filename="test_document.pdf") result = executor._convert_a2a_parts_to_content_blocks([part]) @@ -200,25 +148,10 @@ def test_convert_a2a_parts_to_content_blocks_file_part_document_bytes(): def test_convert_a2a_parts_to_content_blocks_file_part_uri(): - """Test conversion of FilePart with URI to ContentBlock.""" - from a2a.types import FilePart - + """Test conversion of a URL Part to ContentBlock.""" executor = StrandsA2AExecutor(MagicMock()) - # Mock file object with URI - file_obj = MagicMock() - file_obj.name = "test_image.png" - file_obj.mime_type = "image/png" - file_obj.bytes = None - file_obj.uri = "https://example.com/image.png" - - # Mock FilePart with proper spec - file_part = MagicMock(spec=FilePart) - file_part.file = file_obj - - # Mock Part with FilePart root - part = MagicMock() - part.root = file_part + part = Part(url="https://example.com/image.png", media_type="image/png", filename="test_image.png") result = executor._convert_a2a_parts_to_content_blocks([part]) @@ -229,76 +162,12 @@ def test_convert_a2a_parts_to_content_blocks_file_part_uri(): assert "https://example.com/image.png" in content_block["text"] -def test_convert_a2a_parts_to_content_blocks_file_part_with_bytes(): - """Test conversion of FilePart with bytes data.""" - executor = StrandsA2AExecutor(MagicMock()) - - base64_bytes = base64.b64encode(VALID_PNG_BYTES).decode("utf-8") - - # Mock file object with bytes (no validation needed since no decoding) - file_obj = MagicMock() - file_obj.name = "test_image.png" - file_obj.mime_type = "image/png" - file_obj.bytes = base64_bytes - file_obj.uri = None - - # Mock FilePart with proper spec - file_part = MagicMock(spec=FilePart) - file_part.file = file_obj - - # Mock Part with FilePart root - part = MagicMock() - part.root = file_part - - result = executor._convert_a2a_parts_to_content_blocks([part]) - - assert len(result) == 1 - content_block = result[0] - assert "image" in content_block - assert content_block["image"]["source"]["bytes"] == VALID_PNG_BYTES - - -def test_convert_a2a_parts_to_content_blocks_file_part_invalid_base64(): - """Test conversion of FilePart with invalid base64 data raises ValueError.""" - executor = StrandsA2AExecutor(MagicMock()) - - # Invalid base64 string - contains invalid characters - invalid_base64 = "SGVsbG8gV29ybGQ@#$%" - - # Mock file object with invalid base64 bytes - file_obj = MagicMock() - file_obj.name = "test.txt" - file_obj.mime_type = "text/plain" - file_obj.bytes = invalid_base64 - file_obj.uri = None - - # Mock FilePart - file_part = MagicMock(spec=FilePart) - file_part.file = file_obj - part = MagicMock() - part.root = file_part - - # Should handle the base64 decode error gracefully and return empty list - result = executor._convert_a2a_parts_to_content_blocks([part]) - assert isinstance(result, list) - # The part should be skipped due to base64 decode error - assert len(result) == 0 - - def test_convert_a2a_parts_to_content_blocks_data_part(): - """Test conversion of DataPart to ContentBlock.""" - from a2a.types import DataPart - + """Test conversion of a data Part to ContentBlock.""" executor = StrandsA2AExecutor(MagicMock()) - # Mock DataPart with proper spec test_data = {"key": "value", "number": 42} - data_part = MagicMock(spec=DataPart) - data_part.data = test_data - - # Mock Part with DataPart root - part = MagicMock() - part.root = data_part + part = new_data_part(test_data) result = executor._convert_a2a_parts_to_content_blocks([part]) @@ -312,23 +181,9 @@ def test_convert_a2a_parts_to_content_blocks_data_part(): def test_convert_a2a_parts_to_content_blocks_mixed_parts(): """Test conversion of mixed A2A parts to ContentBlocks.""" - from a2a.types import DataPart, TextPart - executor = StrandsA2AExecutor(MagicMock()) - # Mock TextPart with proper spec - text_part = MagicMock(spec=TextPart) - text_part.text = "Text content" - text_part_mock = MagicMock() - text_part_mock.root = text_part - - # Mock DataPart with proper spec - data_part = MagicMock(spec=DataPart) - data_part.data = {"test": "data"} - data_part_mock = MagicMock() - data_part_mock.root = data_part - - parts = [text_part_mock, data_part_mock] + parts = [Part(text="Text content"), new_data_part({"test": "data"})] result = executor._convert_a2a_parts_to_content_blocks(parts) assert len(result) == 2 @@ -358,15 +213,8 @@ async def mock_stream(content_blocks): mock_task.context_id = "test-context-id" mock_request_context.current_task = mock_task - # Mock message with parts - from a2a.types import TextPart - mock_message = MagicMock() - text_part = MagicMock(spec=TextPart) - text_part.text = "Test input" - part = MagicMock() - part.root = text_part - mock_message.parts = [part] + mock_message.parts = [Part(text="Test input")] mock_request_context.message = mock_message await executor.execute(mock_request_context, mock_event_queue) @@ -402,15 +250,8 @@ async def mock_stream(content_blocks): mock_task.context_id = "test-context-id" mock_request_context.current_task = mock_task - # Mock message with parts - from a2a.types import TextPart - mock_message = MagicMock() - text_part = MagicMock(spec=TextPart) - text_part.text = "Test input" - part = MagicMock() - part.root = text_part - mock_message.parts = [part] + mock_message.parts = [Part(text="Test input")] mock_request_context.message = mock_message await executor.execute(mock_request_context, mock_event_queue) @@ -447,15 +288,8 @@ async def mock_stream(content_blocks): mock_task.context_id = "test-context-id" mock_request_context.current_task = mock_task - # Mock message with parts - from a2a.types import TextPart - mock_message = MagicMock() - text_part = MagicMock(spec=TextPart) - text_part.text = "Test input" - part = MagicMock() - part.root = text_part - mock_message.parts = [part] + mock_message.parts = [Part(text="Test input")] mock_request_context.message = mock_message await executor.execute(mock_request_context, mock_event_queue) @@ -492,15 +326,8 @@ async def mock_stream(content_blocks): mock_task.context_id = "test-context-id" mock_request_context.current_task = mock_task - # Mock message with parts - from a2a.types import TextPart - mock_message = MagicMock() - text_part = MagicMock(spec=TextPart) - text_part.text = "Test input" - part = MagicMock() - part.root = text_part - mock_message.parts = [part] + mock_message.parts = [Part(text="Test input")] mock_request_context.message = mock_message await executor.execute(mock_request_context, mock_event_queue) @@ -520,7 +347,7 @@ async def mock_stream(content_blocks): async def test_execute_streaming_mode_fallback_to_text_extraction( mock_strands_agent, mock_request_context, mock_event_queue ): - """Test that execute raises ServerError when no A2A parts are available.""" + """Test that execute raises the specific error when no A2A parts are available.""" # Create executor executor = StrandsA2AExecutor(mock_strands_agent) @@ -537,12 +364,9 @@ async def test_execute_streaming_mode_fallback_to_text_extraction( mock_request_context.message = mock_message mock_request_context.get_user_input.return_value = "Fallback input" - with pytest.raises(ServerError) as excinfo: + with pytest.raises(InternalError): await executor.execute(mock_request_context, mock_event_queue) - # Verify the error is a ServerError containing an InternalError - assert isinstance(excinfo.value.error, InternalError) - @pytest.mark.asyncio async def test_execute_creates_task_when_none_exists(mock_strands_agent, mock_request_context, mock_event_queue): @@ -562,18 +386,11 @@ async def mock_stream(content_blocks): # Mock no existing task mock_request_context.current_task = None - # Mock message with parts - from a2a.types import TextPart - mock_message = MagicMock() - text_part = MagicMock(spec=TextPart) - text_part.text = "Test input" - part = MagicMock() - part.root = text_part - mock_message.parts = [part] + mock_message.parts = [Part(text="Test input")] mock_request_context.message = mock_message - with patch("strands.multiagent.a2a.executor.new_task") as mock_new_task: + with patch("strands.multiagent.a2a.executor.new_task_from_user_message") as mock_new_task: mock_new_task.return_value = MagicMock(id="new-task-id", context_id="new-context-id") await executor.execute(mock_request_context, mock_event_queue) @@ -601,15 +418,8 @@ async def test_execute_streaming_mode_handles_agent_exception( mock_task.context_id = "test-context-id" mock_request_context.current_task = mock_task - # Mock message with parts - from a2a.types import TextPart - mock_message = MagicMock() - text_part = MagicMock(spec=TextPart) - text_part.text = "Test input" - part = MagicMock() - part.root = text_part - mock_message.parts = [part] + mock_message.parts = [Part(text="Test input")] mock_request_context.message = mock_message # Should NOT raise - instead transitions to failed state @@ -623,20 +433,19 @@ async def test_execute_streaming_mode_handles_agent_exception( from a2a.types import TaskState, TaskStatusUpdateEvent failed_events = [ - e for e in enqueued_events if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.failed + e + for e in enqueued_events + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_FAILED ] assert len(failed_events) == 1 - assert "Agent execution failed" in failed_events[0].status.message.parts[0].root.text + assert "Agent execution failed" in failed_events[0].status.message.parts[0].text executor = StrandsA2AExecutor(mock_strands_agent) # Cancel with no current_task raises UnsupportedOperationError mock_request_context.current_task = None - with pytest.raises(ServerError) as excinfo: + with pytest.raises(UnsupportedOperationError): await executor.cancel(mock_request_context, mock_event_queue) - # Verify the error is a ServerError containing an UnsupportedOperationError - assert isinstance(excinfo.value.error, UnsupportedOperationError) - @pytest.mark.asyncio async def test_handle_agent_result_with_none_result(mock_strands_agent, mock_request_context, mock_event_queue): @@ -714,16 +523,15 @@ async def test_handle_agent_result_with_content(mock_strands_agent): # Check that the artifact contains the expected content call_args = mock_updater.add_artifact.call_args[0][0] assert len(call_args) == 1 - assert call_args[0].root.text == "Test response content" + assert call_args[0].text == "Test response content" def test_handle_conversion_error(): """Test that conversion handles errors gracefully.""" executor = StrandsA2AExecutor(MagicMock()) - # Mock Part that will raise an exception during processing - problematic_part = MagicMock() - problematic_part.root = None # This should cause an AttributeError + # A bare object with no HasField method raises AttributeError during processing. + problematic_part = object() # Should not raise an exception, but return empty list or handle gracefully result = executor._convert_a2a_parts_to_content_blocks([problematic_part]) @@ -742,110 +550,56 @@ def test_convert_a2a_parts_to_content_blocks_empty_list(): def test_convert_a2a_parts_to_content_blocks_file_part_no_name(): - """Test conversion of FilePart with no file name.""" + """Test conversion of a raw Part with no filename.""" executor = StrandsA2AExecutor(MagicMock()) - base64_bytes = base64.b64encode(VALID_DOCUMENT_BYTES).decode("utf-8") - - # Mock file object without name - file_obj = MagicMock() - delattr(file_obj, "name") # Remove name attribute - file_obj.mime_type = "text/plain" - file_obj.bytes = base64_bytes - file_obj.uri = None - - # Mock FilePart with proper spec - file_part = MagicMock(spec=FilePart) - file_part.file = file_obj - - # Mock Part with FilePart root - part = MagicMock() - part.root = file_part + part = Part(raw=VALID_DOCUMENT_BYTES, media_type="text/plain") result = executor._convert_a2a_parts_to_content_blocks([part]) assert len(result) == 1 content_block = result[0] assert "document" in content_block - assert content_block["document"]["name"] == "FileNameNotProvided" # Should use default + # Should use default + assert content_block["document"]["name"] == "FileNameNotProvided" def test_convert_a2a_parts_to_content_blocks_file_part_no_mime_type(): - """Test conversion of FilePart with no MIME type.""" + """Test conversion of a raw Part with no media_type.""" executor = StrandsA2AExecutor(MagicMock()) - base64_bytes = base64.b64encode(VALID_DOCUMENT_BYTES).decode("utf-8") - - # Mock file object without MIME type - file_obj = MagicMock() - file_obj.name = "test_file" - delattr(file_obj, "mime_type") - file_obj.bytes = base64_bytes - file_obj.uri = None - - # Mock FilePart with proper spec - file_part = MagicMock(spec=FilePart) - file_part.file = file_obj - - # Mock Part with FilePart root - part = MagicMock() - part.root = file_part + part = Part(raw=VALID_DOCUMENT_BYTES, filename="test_file") result = executor._convert_a2a_parts_to_content_blocks([part]) assert len(result) == 1 content_block = result[0] assert "document" in content_block # Should default to document with unknown type - assert content_block["document"]["format"] == "txt" # Should use default format for unknown file type + # Should use default format for unknown file type + assert content_block["document"]["format"] == "txt" def test_convert_a2a_parts_to_content_blocks_file_part_no_bytes_no_uri(): - """Test conversion of FilePart with neither bytes nor URI.""" - from a2a.types import FilePart - + """Test conversion of a Part with none of text/raw/url/data set.""" executor = StrandsA2AExecutor(MagicMock()) - # Mock file object without bytes or URI - file_obj = MagicMock() - file_obj.name = "test_file.txt" - file_obj.mime_type = "text/plain" - file_obj.bytes = None - file_obj.uri = None - - # Mock FilePart with proper spec - file_part = MagicMock(spec=FilePart) - file_part.file = file_obj - - # Mock Part with FilePart root - part = MagicMock() - part.root = file_part + part = Part() result = executor._convert_a2a_parts_to_content_blocks([part]) - # Should return empty list since no fallback case exists + # Should return empty list since no field matches any conversion branch assert len(result) == 0 def test_convert_a2a_parts_to_content_blocks_data_part_serialization_error(): - """Test conversion of DataPart with non-serializable data.""" - from a2a.types import DataPart - + """Test conversion of a data Part when structured-data conversion fails.""" executor = StrandsA2AExecutor(MagicMock()) - # Create non-serializable data (e.g., a function) - def non_serializable(): - pass - - # Mock DataPart with proper spec - data_part = MagicMock(spec=DataPart) - data_part.data = {"function": non_serializable} # This will cause JSON serialization to fail + part = new_data_part({"key": "value"}) - # Mock Part with DataPart root - part = MagicMock() - part.root = data_part - - # Should not raise an exception, should handle gracefully - result = executor._convert_a2a_parts_to_content_blocks([part]) + with patch("strands.multiagent.a2a.executor.MessageToDict", side_effect=Exception("boom")): + # Should not raise an exception, should handle gracefully + result = executor._convert_a2a_parts_to_content_blocks([part]) # The error handling should result in an empty list or the part being skipped assert isinstance(result, list) @@ -855,23 +609,21 @@ def non_serializable(): async def test_execute_streaming_mode_raises_error_for_empty_content_blocks( mock_strands_agent, mock_event_queue, mock_request_context ): - """Test that execute raises ServerError when content blocks are empty after conversion.""" + """Test that execute raises the specific error when content blocks are empty after conversion.""" executor = StrandsA2AExecutor(mock_strands_agent) # Create a mock message with parts that will result in empty content blocks # This could happen if all parts fail to convert or are invalid mock_message = MagicMock() - mock_message.parts = [MagicMock()] # Has parts but they won't convert to valid content blocks + # Has parts but they won't convert to valid content blocks + mock_message.parts = [MagicMock()] mock_request_context.message = mock_message # Mock the conversion to return empty list with patch.object(executor, "_convert_a2a_parts_to_content_blocks", return_value=[]): - with pytest.raises(ServerError) as excinfo: + with pytest.raises(InternalError): await executor.execute(mock_request_context, mock_event_queue) - # Verify the error is a ServerError containing an InternalError - assert isinstance(excinfo.value.error, InternalError) - @pytest.mark.asyncio async def test_execute_with_mixed_part_types(mock_strands_agent, mock_request_context, mock_event_queue): @@ -895,27 +647,13 @@ async def mock_stream(content_blocks): mock_request_context.current_task = mock_task # Create mixed parts - text_part = MagicMock(spec=TextPart) - text_part.text = "Hello" - text_part_mock = MagicMock() - text_part_mock.root = text_part - - # File part with bytes - file_obj = MagicMock() - file_obj.name = "image.png" - file_obj.mime_type = "image/png" - file_obj.bytes = base64.b64encode(VALID_PNG_BYTES).decode("utf-8") - file_obj.uri = None - file_part = MagicMock(spec=FilePart) - file_part.file = file_obj - file_part_mock = MagicMock() - file_part_mock.root = file_part + text_part_mock = Part(text="Hello") + + # File part with raw bytes + file_part_mock = Part(raw=VALID_PNG_BYTES, media_type="image/png", filename="image.png") # Data part - data_part = MagicMock(spec=DataPart) - data_part.data = {"key": "value"} - data_part_mock = MagicMock() - data_part_mock.root = data_part + data_part_mock = new_data_part({"key": "value"}) # Mock message with mixed parts mock_message = MagicMock() @@ -948,42 +686,16 @@ def test_integration_example(): executor = StrandsA2AExecutor(MagicMock()) # Example 1: Text content - text_part = MagicMock(spec=TextPart) - text_part.text = "Hello, this is a text message" - text_part_mock = MagicMock() - text_part_mock.root = text_part + text_part_mock = Part(text="Hello, this is a text message") # Example 2: Image file - image_bytes = base64.b64encode(VALID_PNG_BYTES).decode("utf-8") - image_file = MagicMock() - image_file.name = "photo.jpg" - image_file.mime_type = "image/jpeg" - image_file.bytes = image_bytes - image_file.uri = None - - image_part = MagicMock(spec=FilePart) - image_part.file = image_file - image_part_mock = MagicMock() - image_part_mock.root = image_part + image_part_mock = Part(raw=VALID_PNG_BYTES, media_type="image/jpeg", filename="photo.jpg") # Example 3: Document file - doc_bytes = base64.b64encode(VALID_DOCUMENT_BYTES).decode("utf-8") - doc_file = MagicMock() - doc_file.name = "report.pdf" - doc_file.mime_type = "application/pdf" - doc_file.bytes = doc_bytes - doc_file.uri = None - - doc_part = MagicMock(spec=FilePart) - doc_part.file = doc_file - doc_part_mock = MagicMock() - doc_part_mock.root = doc_part + doc_part_mock = Part(raw=VALID_DOCUMENT_BYTES, media_type="application/pdf", filename="report.pdf") # Example 4: Structured data - data_part = MagicMock(spec=DataPart) - data_part.data = {"user": "john_doe", "action": "upload_file", "timestamp": "2023-12-01T10:00:00Z"} - data_part_mock = MagicMock() - data_part_mock.root = data_part + data_part_mock = new_data_part({"user": "john_doe", "action": "upload_file", "timestamp": "2023-12-01T10:00:00Z"}) # Convert all parts to ContentBlocks parts = [text_part_mock, image_part_mock, doc_part_mock, data_part_mock] @@ -1003,7 +715,8 @@ def test_integration_example(): # Document part becomes document ContentBlock assert "document" in content_blocks[2] assert content_blocks[2]["document"]["format"] == "pdf" - assert content_blocks[2]["document"]["name"] == "report" # Extension stripped + # Extension stripped + assert content_blocks[2]["document"]["name"] == "report" assert content_blocks[2]["document"]["source"]["bytes"] == VALID_DOCUMENT_BYTES # Data part becomes text ContentBlock with JSON representation @@ -1043,8 +756,6 @@ def test_default_formats_modularization(): @pytest.mark.asyncio async def test_legacy_mode_emits_deprecation_warning(mock_strands_agent, mock_request_context, mock_event_queue): """Test that legacy streaming (default) emits deprecation warning.""" - from a2a.types import TextPart - executor = StrandsA2AExecutor(mock_strands_agent) # Default is False # Mock stream_async @@ -1059,13 +770,8 @@ async def mock_stream(content_blocks): mock_task.context_id = "test-context-id" mock_request_context.current_task = mock_task - # Mock message - mock_text_part = MagicMock(spec=TextPart) - mock_text_part.text = "test" - mock_part = MagicMock() - mock_part.root = mock_text_part mock_message = MagicMock() - mock_message.parts = [mock_part] + mock_message.parts = [Part(text="test")] mock_request_context.message = mock_message with pytest.warns(UserWarning, match="does not conform to what is expected in the A2A spec"): @@ -1077,8 +783,6 @@ async def test_a2a_compliant_mode_no_warning(mock_strands_agent, mock_request_co """Test that A2A-compliant mode does not emit warning.""" import warnings - from a2a.types import TextPart - executor = StrandsA2AExecutor(mock_strands_agent, enable_a2a_compliant_streaming=True) # Mock stream_async @@ -1093,13 +797,8 @@ async def mock_stream(content_blocks): mock_task.context_id = "test-context-id" mock_request_context.current_task = mock_task - # Mock message - mock_text_part = MagicMock(spec=TextPart) - mock_text_part.text = "test" - mock_part = MagicMock() - mock_part.root = mock_text_part mock_message = MagicMock() - mock_message.parts = [mock_part] + mock_message.parts = [Part(text="test")] mock_request_context.message = mock_message with warnings.catch_warnings(): @@ -1147,7 +846,7 @@ async def test_a2a_compliant_handle_result_first_chunk_with_content(mock_strands mock_updater.add_artifact.assert_called_once() parts = mock_updater.add_artifact.call_args[0][0] assert len(parts) == 1 - assert parts[0].root.text == "Final response" + assert parts[0].text == "Final response" assert mock_updater.add_artifact.call_args[1]["artifact_id"] == "artifact-456" assert mock_updater.add_artifact.call_args[1]["last_chunk"] is True mock_updater.complete.assert_called_once() @@ -1172,7 +871,7 @@ async def test_a2a_compliant_handle_result_first_chunk_with_none_result(mock_str mock_updater.add_artifact.assert_called_once() parts = mock_updater.add_artifact.call_args[0][0] assert len(parts) == 1 - assert parts[0].root.text == "" + assert parts[0].text == "" assert mock_updater.add_artifact.call_args[1]["artifact_id"] == "artifact-789" assert mock_updater.add_artifact.call_args[1]["last_chunk"] is True mock_updater.complete.assert_called_once() @@ -1200,7 +899,7 @@ async def test_a2a_compliant_handle_result_not_first_chunk(mock_strands_agent): mock_updater.add_artifact.assert_called_once() parts = mock_updater.add_artifact.call_args[0][0] assert len(parts) == 1 - assert parts[0].root.text == "" + assert parts[0].text == "" assert mock_updater.add_artifact.call_args[1]["artifact_id"] == "artifact-abc" assert mock_updater.add_artifact.call_args[1]["append"] is True assert mock_updater.add_artifact.call_args[1]["last_chunk"] is True @@ -1225,13 +924,8 @@ async def mock_stream(content_blocks: list, **kwargs: Any) -> Any: mock_strands_agent.stream_async = MagicMock(side_effect=mock_stream) - # Set up message with a text part - mock_text_part = MagicMock(spec=TextPart) - mock_text_part.text = "test input" - mock_part = MagicMock() - mock_part.root = mock_text_part mock_message = MagicMock() - mock_message.parts = [mock_part] + mock_message.parts = [Part(text="test input")] mock_request_context.message = mock_message @@ -1308,7 +1002,7 @@ async def test_invocation_state_context_when_no_task(mock_strands_agent, mock_re executor = StrandsA2AExecutor(mock_strands_agent) - with patch("strands.multiagent.a2a.executor.new_task") as mock_new_task: + with patch("strands.multiagent.a2a.executor.new_task_from_user_message") as mock_new_task: mock_new_task.return_value = MagicMock(id="generated-id", context_id="generated-ctx") await executor.execute(mock_request_context, mock_event_queue) @@ -1350,7 +1044,7 @@ async def test_execute_transitions_to_failed_on_streaming_error( mock_strands_agent, mock_request_context, mock_event_queue ): """Test that errors during streaming transition task to failed state.""" - from a2a.types import TaskState, TaskStatusUpdateEvent, TextPart + from a2a.types import TaskState, TaskStatusUpdateEvent async def mock_stream(content_blocks, **kwargs): """Mock streaming that raises mid-stream.""" @@ -1366,12 +1060,8 @@ async def mock_stream(content_blocks, **kwargs): mock_task.context_id = "ctx-fail" mock_request_context.current_task = mock_task - mock_text_part = MagicMock(spec=TextPart) - mock_text_part.text = "test" - mock_part = MagicMock() - mock_part.root = mock_text_part mock_message = MagicMock() - mock_message.parts = [mock_part] + mock_message.parts = [Part(text="test")] mock_request_context.message = mock_message # Should not raise @@ -1380,10 +1070,52 @@ async def mock_stream(content_blocks, **kwargs): # Verify failed state was enqueued enqueued_events = [call[0][0] for call in mock_event_queue.enqueue_event.call_args_list] failed_events = [ - e for e in enqueued_events if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.failed + e + for e in enqueued_events + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_FAILED ] assert len(failed_events) == 1 - assert "Agent execution failed" in failed_events[0].status.message.parts[0].root.text + assert "Agent execution failed" in failed_events[0].status.message.parts[0].text + + +@pytest.mark.asyncio +async def test_execute_a2a_client_error_transitions_to_failed( + mock_strands_agent, mock_request_context, mock_event_queue +): + """An A2AError subclass raised by a nested A2A call transitions the task to failed.""" + from a2a.types import TaskState, TaskStatusUpdateEvent + from a2a.utils.errors import A2AError + + class FakeClientError(A2AError): + """Simulates an A2A client error from a nested remote call.""" + + async def mock_stream(content_blocks, **kwargs): + yield {"data": "partial"} + raise FakeClientError(message="Remote agent timed out") + + mock_strands_agent.stream_async = MagicMock(side_effect=mock_stream) + + executor = StrandsA2AExecutor(mock_strands_agent) + + mock_task = MagicMock() + mock_task.id = "task-client-err" + mock_task.context_id = "ctx-client-err" + mock_request_context.current_task = mock_task + + mock_message = MagicMock() + mock_message.parts = [Part(text="call remote agent")] + mock_request_context.message = mock_message + + await executor.execute(mock_request_context, mock_event_queue) + + enqueued_events = [call[0][0] for call in mock_event_queue.enqueue_event.call_args_list] + failed_events = [ + e + for e in enqueued_events + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_FAILED + ] + assert len(failed_events) == 1 + assert "Agent execution failed" in failed_events[0].status.message.parts[0].text @pytest.mark.asyncio @@ -1403,10 +1135,12 @@ async def test_cancel_with_valid_task(mock_strands_agent, mock_request_context, # Verify canceled state was enqueued enqueued_events = [call[0][0] for call in mock_event_queue.enqueue_event.call_args_list] canceled_events = [ - e for e in enqueued_events if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.canceled + e + for e in enqueued_events + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_CANCELED ] assert len(canceled_events) == 1 - assert "cancelled" in canceled_events[0].status.message.parts[0].root.text.lower() + assert "cancelled" in canceled_events[0].status.message.parts[0].text.lower() @pytest.mark.asyncio @@ -1415,18 +1149,16 @@ async def test_cancel_without_task_raises_unsupported(mock_strands_agent, mock_r executor = StrandsA2AExecutor(mock_strands_agent) mock_request_context.current_task = None - with pytest.raises(ServerError) as excinfo: + with pytest.raises(UnsupportedOperationError): await executor.cancel(mock_request_context, mock_event_queue) - assert isinstance(excinfo.value.error, UnsupportedOperationError) - @pytest.mark.asyncio async def test_execute_with_interrupt_transitions_to_input_required( mock_strands_agent, mock_request_context, mock_event_queue ): """Test that agent interrupts map to input_required state.""" - from a2a.types import TaskState, TaskStatusUpdateEvent, TextPart + from a2a.types import TaskState, TaskStatusUpdateEvent from strands.interrupt import Interrupt @@ -1449,12 +1181,8 @@ async def mock_stream(content_blocks, **kwargs): mock_task.context_id = "ctx-interrupt" mock_request_context.current_task = mock_task - mock_text_part = MagicMock(spec=TextPart) - mock_text_part.text = "delete file X" - mock_part = MagicMock() - mock_part.root = mock_text_part mock_message = MagicMock() - mock_message.parts = [mock_part] + mock_message.parts = [Part(text="delete file X")] mock_request_context.message = mock_message await executor.execute(mock_request_context, mock_event_queue) @@ -1464,10 +1192,10 @@ async def mock_stream(content_blocks, **kwargs): input_required_events = [ e for e in enqueued_events - if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.input_required + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_INPUT_REQUIRED ] assert len(input_required_events) == 1 - msg_text = input_required_events[0].status.message.parts[0].root.text + msg_text = input_required_events[0].status.message.parts[0].text assert "approval" in msg_text assert "Need user approval" in msg_text @@ -1475,7 +1203,7 @@ async def mock_stream(content_blocks, **kwargs): @pytest.mark.asyncio async def test_execute_with_multiple_interrupts(mock_strands_agent, mock_request_context, mock_event_queue): """Test handling of multiple interrupts in a single result.""" - from a2a.types import TaskState, TaskStatusUpdateEvent, TextPart + from a2a.types import TaskState, TaskStatusUpdateEvent from strands.interrupt import Interrupt @@ -1498,12 +1226,8 @@ async def mock_stream(content_blocks, **kwargs): mock_task.context_id = "ctx-multi-int" mock_request_context.current_task = mock_task - mock_text_part = MagicMock(spec=TextPart) - mock_text_part.text = "delete with backup" - mock_part = MagicMock() - mock_part.root = mock_text_part mock_message = MagicMock() - mock_message.parts = [mock_part] + mock_message.parts = [Part(text="delete with backup")] mock_request_context.message = mock_message await executor.execute(mock_request_context, mock_event_queue) @@ -1512,10 +1236,10 @@ async def mock_stream(content_blocks, **kwargs): input_required_events = [ e for e in enqueued_events - if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.input_required + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_INPUT_REQUIRED ] assert len(input_required_events) == 1 - msg_text = input_required_events[0].status.message.parts[0].root.text + msg_text = input_required_events[0].status.message.parts[0].text assert "confirm_delete" in msg_text assert "select_backup" in msg_text assert "Confirm deletion of file X" in msg_text @@ -1525,7 +1249,7 @@ async def mock_stream(content_blocks, **kwargs): @pytest.mark.asyncio async def test_execute_normal_completion_no_interrupts(mock_strands_agent, mock_request_context, mock_event_queue): """Test that normal completion (no interrupts) still works as before.""" - from a2a.types import TaskState, TaskStatusUpdateEvent, TextPart + from a2a.types import TaskState, TaskStatusUpdateEvent mock_result = MagicMock(spec=SAAgentResult) mock_result.stop_reason = "end_turn" @@ -1545,12 +1269,8 @@ async def mock_stream(content_blocks, **kwargs): mock_task.context_id = "ctx-normal" mock_request_context.current_task = mock_task - mock_text_part = MagicMock(spec=TextPart) - mock_text_part.text = "do something" - mock_part = MagicMock() - mock_part.root = mock_text_part mock_message = MagicMock() - mock_message.parts = [mock_part] + mock_message.parts = [Part(text="do something")] mock_request_context.message = mock_message await executor.execute(mock_request_context, mock_event_queue) @@ -1558,7 +1278,9 @@ async def mock_stream(content_blocks, **kwargs): # Verify completed state was enqueued (not input_required) enqueued_events = [call[0][0] for call in mock_event_queue.enqueue_event.call_args_list] completed_events = [ - e for e in enqueued_events if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.completed + e + for e in enqueued_events + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_COMPLETED ] assert len(completed_events) == 1 @@ -1566,7 +1288,7 @@ async def mock_stream(content_blocks, **kwargs): input_required_events = [ e for e in enqueued_events - if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.input_required + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_INPUT_REQUIRED ] assert len(input_required_events) == 0 @@ -1584,17 +1306,13 @@ async def test_execute_setup_failure_raises_server_error(mock_strands_agent, moc # No message at all mock_request_context.message = None - with pytest.raises(ServerError) as excinfo: + with pytest.raises(InternalError): await executor.execute(mock_request_context, mock_event_queue) - assert isinstance(excinfo.value.error, InternalError) - @pytest.mark.asyncio async def test_execute_error_when_task_already_terminal(mock_strands_agent, mock_request_context, mock_event_queue): """Test that error during execution is handled gracefully when task is already in terminal state.""" - from a2a.types import TextPart - # Make stream_async raise to trigger the error path mock_strands_agent.stream_async = MagicMock(side_effect=Exception("Agent error")) @@ -1605,12 +1323,8 @@ async def test_execute_error_when_task_already_terminal(mock_strands_agent, mock mock_task.context_id = "ctx-already-done" mock_request_context.current_task = mock_task - mock_text_part = MagicMock(spec=TextPart) - mock_text_part.text = "test" - mock_part = MagicMock() - mock_part.root = mock_text_part mock_message = MagicMock() - mock_message.parts = [mock_part] + mock_message.parts = [Part(text="test")] mock_request_context.message = mock_message # Patch TaskUpdater.failed to raise RuntimeError (simulating task already in terminal state) @@ -1650,7 +1364,9 @@ async def test_cancel_calls_agent_cancel_method(mock_strands_agent, mock_request # Verify task state is canceled enqueued_events = [call[0][0] for call in mock_event_queue.enqueue_event.call_args_list] canceled_events = [ - e for e in enqueued_events if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.canceled + e + for e in enqueued_events + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_CANCELED ] assert len(canceled_events) == 1 @@ -1676,14 +1392,16 @@ async def test_cancel_handles_agent_cancel_exception(mock_strands_agent, mock_re # Task should still be transitioned to canceled enqueued_events = [call[0][0] for call in mock_event_queue.enqueue_event.call_args_list] canceled_events = [ - e for e in enqueued_events if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.canceled + e + for e in enqueued_events + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_CANCELED ] assert len(canceled_events) == 1 @pytest.mark.asyncio async def test_cancel_raises_when_task_already_terminal(mock_strands_agent, mock_request_context, mock_event_queue): - """Test that cancel() raises ServerError when task is already in a terminal state.""" + """Test that cancel() raises the specific error when task is already in a terminal state.""" executor = StrandsA2AExecutor(mock_strands_agent) mock_task = MagicMock() @@ -1698,10 +1416,8 @@ async def test_cancel_raises_when_task_already_terminal(mock_strands_agent, mock mock_updater.new_agent_message = MagicMock(return_value=MagicMock()) MockTaskUpdater.return_value = mock_updater - with pytest.raises(ServerError) as excinfo: + with pytest.raises(UnsupportedOperationError): await executor.cancel(mock_request_context, mock_event_queue) - - assert isinstance(excinfo.value.error, UnsupportedOperationError) mock_updater.cancel.assert_called_once() @@ -1722,7 +1438,7 @@ async def test_execute_handles_asyncio_cancelled_error(mock_strands_agent, mock_ """ import asyncio - from a2a.types import TaskState, TaskStatusUpdateEvent, TextPart + from a2a.types import TaskState, TaskStatusUpdateEvent async def mock_stream(content_blocks, **kwargs): """Mock streaming that gets cancelled mid-stream.""" @@ -1738,12 +1454,8 @@ async def mock_stream(content_blocks, **kwargs): mock_task.context_id = "ctx-cancelled" mock_request_context.current_task = mock_task - mock_text_part = MagicMock(spec=TextPart) - mock_text_part.text = "test" - mock_part = MagicMock() - mock_part.root = mock_text_part mock_message = MagicMock() - mock_message.parts = [mock_part] + mock_message.parts = [Part(text="test")] mock_request_context.message = mock_message # CancelledError should be re-raised (framework needs to know task was cancelled) @@ -1753,12 +1465,14 @@ async def mock_stream(content_blocks, **kwargs): # But BEFORE re-raising, the task should have been transitioned to canceled enqueued_events = [call[0][0] for call in mock_event_queue.enqueue_event.call_args_list] canceled_events = [ - e for e in enqueued_events if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.canceled + e + for e in enqueued_events + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_CANCELED ] assert len(canceled_events) == 1 assert ( - "cancelled" in canceled_events[0].status.message.parts[0].root.text.lower() - or "connection termination" in canceled_events[0].status.message.parts[0].root.text.lower() + "cancelled" in canceled_events[0].status.message.parts[0].text.lower() + or "connection termination" in canceled_events[0].status.message.parts[0].text.lower() ) @@ -1773,8 +1487,6 @@ async def test_execute_asyncio_cancelled_when_task_already_terminal( """ import asyncio - from a2a.types import TextPart - async def mock_stream(content_blocks, **kwargs): """Async generator that immediately raises CancelledError.""" yield {"data": "partial"} # Must yield to be async generator @@ -1789,12 +1501,8 @@ async def mock_stream(content_blocks, **kwargs): mock_task.context_id = "ctx-cancelled-terminal" mock_request_context.current_task = mock_task - mock_text_part = MagicMock(spec=TextPart) - mock_text_part.text = "test" - mock_part = MagicMock() - mock_part.root = mock_text_part mock_message = MagicMock() - mock_message.parts = [mock_part] + mock_message.parts = [Part(text="test")] mock_request_context.message = mock_message # Patch TaskUpdater to simulate task already in terminal state @@ -1826,11 +1534,12 @@ async def test_execute_with_interrupt_empty_list_transitions_to_input_required( no interrupt details. This should STILL transition to input_required — the stop_reason is the authoritative signal. Previously this would silently complete the task. """ - from a2a.types import TaskState, TaskStatusUpdateEvent, TextPart + from a2a.types import TaskState, TaskStatusUpdateEvent mock_result = MagicMock(spec=SAAgentResult) mock_result.stop_reason = "interrupt" - mock_result.interrupts = [] # Empty list — previously this was falsy and caused completion! + # Empty list — previously this was falsy and caused completion! + mock_result.interrupts = [] async def mock_stream(content_blocks, **kwargs): yield {"result": mock_result} @@ -1844,12 +1553,8 @@ async def mock_stream(content_blocks, **kwargs): mock_task.context_id = "ctx-empty-interrupts" mock_request_context.current_task = mock_task - mock_text_part = MagicMock(spec=TextPart) - mock_text_part.text = "do something" - mock_part = MagicMock() - mock_part.root = mock_text_part mock_message = MagicMock() - mock_message.parts = [mock_part] + mock_message.parts = [Part(text="do something")] mock_request_context.message = mock_message await executor.execute(mock_request_context, mock_event_queue) @@ -1859,16 +1564,18 @@ async def mock_stream(content_blocks, **kwargs): input_required_events = [ e for e in enqueued_events - if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.input_required + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_INPUT_REQUIRED ] completed_events = [ - e for e in enqueued_events if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.completed + e + for e in enqueued_events + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_COMPLETED ] assert len(input_required_events) == 1, "Empty interrupts list should still trigger input_required" assert len(completed_events) == 0, "Should NOT complete when stop_reason='interrupt'" # Verify the fallback message is used - assert "additional input" in input_required_events[0].status.message.parts[0].root.text.lower() + assert "additional input" in input_required_events[0].status.message.parts[0].text.lower() @pytest.mark.asyncio @@ -1880,7 +1587,7 @@ async def test_execute_with_interrupt_none_list_transitions_to_input_required( Same logic — the stop_reason is authoritative. None interrupts should still result in input_required transition. """ - from a2a.types import TaskState, TaskStatusUpdateEvent, TextPart + from a2a.types import TaskState, TaskStatusUpdateEvent mock_result = MagicMock(spec=SAAgentResult) mock_result.stop_reason = "interrupt" @@ -1898,12 +1605,8 @@ async def mock_stream(content_blocks, **kwargs): mock_task.context_id = "ctx-none-interrupts" mock_request_context.current_task = mock_task - mock_text_part = MagicMock(spec=TextPart) - mock_text_part.text = "do something" - mock_part = MagicMock() - mock_part.root = mock_text_part mock_message = MagicMock() - mock_message.parts = [mock_part] + mock_message.parts = [Part(text="do something")] mock_request_context.message = mock_message await executor.execute(mock_request_context, mock_event_queue) @@ -1912,7 +1615,7 @@ async def mock_stream(content_blocks, **kwargs): input_required_events = [ e for e in enqueued_events - if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.input_required + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_INPUT_REQUIRED ] assert len(input_required_events) == 1 @@ -1938,7 +1641,9 @@ async def test_cancel_without_hasattr_cancel(mock_strands_agent, mock_request_co # Task should still be transitioned to canceled enqueued_events = [call[0][0] for call in mock_event_queue.enqueue_event.call_args_list] canceled_events = [ - e for e in enqueued_events if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.canceled + e + for e in enqueued_events + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_CANCELED ] assert len(canceled_events) == 1 @@ -2006,12 +1711,7 @@ async def stream( def _make_request_context(context_id: str, message_id: str, text: str): """Build a RequestContext-like object the executor can consume.""" - from a2a.types import TextPart - - text_part = MagicMock(spec=TextPart) - text_part.text = text - part = MagicMock() - part.root = text_part + part = Part(text=text) message = MagicMock() message.parts = [part] @@ -2037,8 +1737,8 @@ def _artifact_texts(mock_event_queue): event = call[0][0] if isinstance(event, TaskArtifactUpdateEvent): for part in event.artifact.parts: - if hasattr(part.root, "text"): - texts.append(part.root.text) + if part.HasField("text"): + texts.append(part.text) return texts @@ -2076,8 +1776,10 @@ async def test_single_agent_evicts_least_recently_used_context(mock_event_queue) await executor.execute(_make_request_context("ctx-A", "a-1", "hi"), mock_event_queue) await executor.execute(_make_request_context("ctx-B", "b-1", "hi"), mock_event_queue) - await executor.execute(_make_request_context("ctx-A", "a-2", "again"), mock_event_queue) # touch A - await executor.execute(_make_request_context("ctx-C", "c-1", "hi"), mock_event_queue) # evicts B + # touch A + await executor.execute(_make_request_context("ctx-A", "a-2", "again"), mock_event_queue) + # evicts B + await executor.execute(_make_request_context("ctx-C", "c-1", "hi"), mock_event_queue) assert set(executor._snapshots.keys()) == {"ctx-A", "ctx-C"} @@ -2095,9 +1797,8 @@ async def test_execute_raises_when_context_id_missing(mock_strands_agent, mock_e context = _make_request_context("ignored", "m-1", "hello") context.context_id = None # simulate the should-never-happen absent id - with pytest.raises(ServerError) as excinfo: + with pytest.raises(InternalError): await executor.execute(context, mock_event_queue) - assert isinstance(excinfo.value.error, InternalError) @pytest.mark.asyncio @@ -2185,8 +1886,10 @@ async def test_factory_mode_evicts_least_recently_used_context(mock_event_queue) await executor.execute(_make_request_context("ctx-A", "a-1", "hi"), mock_event_queue) await executor.execute(_make_request_context("ctx-B", "b-1", "hi"), mock_event_queue) - await executor.execute(_make_request_context("ctx-A", "a-2", "again"), mock_event_queue) # touch A - await executor.execute(_make_request_context("ctx-C", "c-1", "hi"), mock_event_queue) # evicts B + # touch A + await executor.execute(_make_request_context("ctx-A", "a-2", "again"), mock_event_queue) + # evicts B + await executor.execute(_make_request_context("ctx-C", "c-1", "hi"), mock_event_queue) assert set(executor._contexts.keys()) == {"ctx-A", "ctx-C"} @@ -2231,7 +1934,8 @@ async def stream( ) -> "AsyncGenerator": yield {"messageStart": {"role": "assistant"}} yield {"contentBlockStart": {"start": {}}} - await asyncio.sleep(0.01) # force interleaving with the other request + # force interleaving with the other request + await asyncio.sleep(0.01) yield {"contentBlockDelta": {"delta": {"text": f"chunk-{context_id}"}}} yield {"contentBlockStop": {}} yield {"messageStop": {"stopReason": "end_turn"}} @@ -2280,30 +1984,18 @@ def artifact_ids(events): def _interrupt_response_part(interrupt_id, response): - """Build an A2A DataPart carrying an interrupt response for the given id.""" - data_part = MagicMock(spec=DataPart) - data_part.data = {"interruptResponse": {"interruptId": interrupt_id, "response": response}} - part = MagicMock() - part.root = data_part - return part + """Build an A2A data Part carrying an interrupt response for the given id.""" + return new_data_part({"interruptResponse": {"interruptId": interrupt_id, "response": response}}) def _data_part(data): - """Build an A2A DataPart carrying arbitrary structured data.""" - data_part = MagicMock(spec=DataPart) - data_part.data = data - part = MagicMock() - part.root = data_part - return part + """Build an A2A data Part carrying arbitrary structured data.""" + return new_data_part(data) def _text_part(text): - """Build an A2A TextPart.""" - text_part = MagicMock(spec=TextPart) - text_part.text = text - part = MagicMock() - part.root = text_part - return part + """Build an A2A text Part.""" + return Part(text=text) def _park_interrupt(agent, *interrupt_ids): @@ -2388,6 +2080,41 @@ async def mock_stream(agent_input, **kwargs): assert tru_input == exp_input +@pytest.mark.asyncio +async def test_execute_interrupt_response_numeric_values_arrive_as_float( + mock_strands_agent, mock_request_context, mock_event_queue +): + """Protobuf Value has no integer type — all numbers round-trip as float. + + This is a v1 wire-format limitation. Integers sent by a peer (e.g. 3) arrive as 3.0. + """ + _park_interrupt(mock_strands_agent, "int-1") + + mock_result = MagicMock(spec=SAAgentResult) + mock_result.stop_reason = "end_turn" + mock_result.interrupts = None + mock_result.__str__ = MagicMock(return_value="done") + + async def mock_stream(agent_input, **kwargs): + yield {"result": mock_result} + + mock_strands_agent.stream_async = MagicMock(side_effect=mock_stream) + executor = StrandsA2AExecutor(mock_strands_agent) + + response = {"count": 3, "ratio": 1.5, "values": [1, 2, 3], "ok": True} + _request_with_parts(mock_request_context, [_interrupt_response_part("int-1", response)]) + await executor.execute(mock_request_context, mock_event_queue) + + tru_input = mock_strands_agent.stream_async.call_args[0][0] + tru_response = tru_input[0]["interruptResponse"]["response"] + assert tru_response["count"] == 3.0 + assert isinstance(tru_response["count"], float) + assert tru_response["ratio"] == 1.5 + assert tru_response["values"] == [1.0, 2.0, 3.0] + assert all(isinstance(v, float) for v in tru_response["values"]) + assert tru_response["ok"] is True + + @pytest.mark.asyncio async def test_execute_null_interrupt_response_rejected(mock_strands_agent, mock_request_context, mock_event_queue): """A null answer leaves the interrupt unsatisfied, so it is refused instead of silently re-firing.""" @@ -2397,10 +2124,8 @@ async def test_execute_null_interrupt_response_rejected(mock_strands_agent, mock _request_with_parts(mock_request_context, [_interrupt_response_part("int-1", None)]) - with pytest.raises(ServerError) as exc_info: + with pytest.raises(InvalidParamsError): await executor.execute(mock_request_context, mock_event_queue) - - assert isinstance(exc_info.value.error, InvalidParamsError) mock_strands_agent.stream_async.assert_not_called() @@ -2439,10 +2164,8 @@ async def test_execute_interrupt_response_for_unknown_id_fails_closed( _request_with_parts(mock_request_context, [_interrupt_response_part("int-other", {"approved": True})]) - with pytest.raises(ServerError) as exc_info: + with pytest.raises(InvalidParamsError): await executor.execute(mock_request_context, mock_event_queue) - - assert isinstance(exc_info.value.error, InvalidParamsError) mock_strands_agent.stream_async.assert_not_called() @@ -2456,10 +2179,8 @@ async def test_execute_interrupt_response_when_not_parked_fails_closed( _request_with_parts(mock_request_context, [_interrupt_response_part("int-1", {"approved": True})]) - with pytest.raises(ServerError) as exc_info: + with pytest.raises(InvalidParamsError): await executor.execute(mock_request_context, mock_event_queue) - - assert isinstance(exc_info.value.error, InvalidParamsError) mock_strands_agent.stream_async.assert_not_called() @@ -2474,10 +2195,8 @@ async def test_execute_rejected_resume_leaves_interrupt_parked( _request_with_parts(mock_request_context, [_interrupt_response_part("int-other", "yes")]) - with pytest.raises(ServerError) as exc_info: + with pytest.raises(InvalidParamsError): await executor.execute(mock_request_context, mock_event_queue) - - assert isinstance(exc_info.value.error, InvalidParamsError) assert state.activated assert "int-1" in state.interrupts @@ -2496,10 +2215,8 @@ async def test_execute_duplicate_interrupt_responses_rejected( [_interrupt_response_part("int-1", {"approved": True}), _interrupt_response_part("int-1", {"approved": False})], ) - with pytest.raises(ServerError) as exc_info: + with pytest.raises(InvalidParamsError): await executor.execute(mock_request_context, mock_event_queue) - - assert isinstance(exc_info.value.error, InvalidParamsError) mock_strands_agent.stream_async.assert_not_called() @@ -2524,10 +2241,8 @@ async def test_execute_malformed_interrupt_response_rejected( _request_with_parts(mock_request_context, [_data_part(malformed)]) - with pytest.raises(ServerError) as exc_info: + with pytest.raises(InvalidParamsError): await executor.execute(mock_request_context, mock_event_queue) - - assert isinstance(exc_info.value.error, InvalidParamsError) mock_strands_agent.stream_async.assert_not_called() @@ -2545,10 +2260,8 @@ async def test_execute_interrupt_response_mixed_with_other_parts_rejected( [_interrupt_response_part("int-1", {"approved": True}), _text_part("and also do something else")], ) - with pytest.raises(ServerError) as exc_info: + with pytest.raises(InvalidParamsError): await executor.execute(mock_request_context, mock_event_queue) - - assert isinstance(exc_info.value.error, InvalidParamsError) mock_strands_agent.stream_async.assert_not_called() @@ -2563,10 +2276,8 @@ async def test_execute_new_message_while_parked_fails_closed( _request_with_parts(mock_request_context, [_text_part("never mind, do something else")]) - with pytest.raises(ServerError) as exc_info: + with pytest.raises(InvalidParamsError): await executor.execute(mock_request_context, mock_event_queue) - - assert isinstance(exc_info.value.error, InvalidParamsError) mock_strands_agent.stream_async.assert_not_called() @@ -2608,7 +2319,9 @@ def _input_required_message(mock_event_queue): events = [call[0][0] for call in mock_event_queue.enqueue_event.call_args_list] input_required = [ - e for e in events if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.input_required + e + for e in events + if isinstance(e, TaskStatusUpdateEvent) and e.status.state == TaskState.TASK_STATE_INPUT_REQUIRED ] assert len(input_required) == 1 return input_required[0].status.message @@ -2642,7 +2355,7 @@ async def test_interrupt_advertises_ids_in_data_part(mock_strands_agent, mock_re [Interrupt(id="v1:tool_call:tu-1:abc", name="approve_campaign", reason={"name": "spring"})], ) - data_parts = [p.root.data for p in _input_required_message(mock_event_queue).parts if isinstance(p.root, DataPart)] + data_parts = [MessageToDict(p.data) for p in _input_required_message(mock_event_queue).parts if p.HasField("data")] tru_data = data_parts exp_data = [ { @@ -2668,8 +2381,8 @@ async def test_interrupt_keeps_human_readable_text_part(mock_strands_agent, mock [Interrupt(id="int-1", name="approval", reason="Need user approval")], ) - first_part = _input_required_message(mock_event_queue).parts[0].root - assert isinstance(first_part, TextPart) + first_part = _input_required_message(mock_event_queue).parts[0] + assert first_part.HasField("text") assert "approval" in first_part.text assert "Need user approval" in first_part.text @@ -2688,7 +2401,7 @@ async def test_interrupt_advertises_every_pending_id(mock_strands_agent, mock_re [Interrupt(id="int-1", name="first"), Interrupt(id="int-2", name="second")], ) - data_parts = [p.root.data for p in _input_required_message(mock_event_queue).parts if isinstance(p.root, DataPart)] + data_parts = [MessageToDict(p.data) for p in _input_required_message(mock_event_queue).parts if p.HasField("data")] tru_ids = [entry["interruptId"] for entry in data_parts[0]["interrupts"]] exp_ids = ["int-1", "int-2"] assert tru_ids == exp_ids @@ -2714,7 +2427,7 @@ def __str__(self): [Interrupt(id="int-1", name="approval", reason=Unserializable())], ) - data_parts = [p.root.data for p in _input_required_message(mock_event_queue).parts if isinstance(p.root, DataPart)] + data_parts = [MessageToDict(p.data) for p in _input_required_message(mock_event_queue).parts if p.HasField("data")] tru_reason = data_parts[0]["interrupts"][0]["reason"] exp_reason = "" assert tru_reason == exp_reason @@ -2727,5 +2440,5 @@ async def test_interrupt_without_details_sends_no_data_part(mock_strands_agent, await _run_until_interrupt(executor, mock_strands_agent, mock_request_context, mock_event_queue, []) parts = _input_required_message(mock_event_queue).parts - assert not [p for p in parts if isinstance(p.root, DataPart)] - assert isinstance(parts[0].root, TextPart) + assert not [p for p in parts if p.HasField("data")] + assert parts[0].HasField("text") diff --git a/strands-py/tests/strands/multiagent/a2a/test_server.py b/strands-py/tests/strands/multiagent/a2a/test_server.py index 3cfa9f34ba..b6435ef21c 100644 --- a/strands-py/tests/strands/multiagent/a2a/test_server.py +++ b/strands-py/tests/strands/multiagent/a2a/test_server.py @@ -88,7 +88,7 @@ def test_public_agent_card(mock_strands_agent): assert isinstance(card, AgentCard) assert card.name == "Test Agent" assert card.description == "A test agent for unit testing" - assert card.url == "http://127.0.0.1:9000/" + assert card.supported_interfaces[0].url == "http://127.0.0.1:9000/" assert card.version == "0.0.1" assert card.default_input_modes == ["text"] assert card.default_output_modes == ["text"] @@ -99,19 +99,17 @@ def test_public_agent_card(mock_strands_agent): def test_public_agent_card_with_missing_name(mock_strands_agent): """Test that public_agent_card raises ValueError when name is missing.""" mock_strands_agent.name = "" - a2a_agent = A2AServer(mock_strands_agent, skills=[]) with pytest.raises(ValueError, match="A2A agent name cannot be None or empty"): - _ = a2a_agent.public_agent_card + A2AServer(mock_strands_agent, skills=[]) def test_public_agent_card_with_missing_description(mock_strands_agent): """Test that public_agent_card raises ValueError when description is missing.""" mock_strands_agent.description = "" - a2a_agent = A2AServer(mock_strands_agent, skills=[]) with pytest.raises(ValueError, match="A2A agent description cannot be None or empty"): - _ = a2a_agent.public_agent_card + A2AServer(mock_strands_agent, skills=[]) def test_agent_skills_empty_registry(mock_strands_agent): @@ -241,11 +239,10 @@ def test_agent_skills_handles_missing_description(mock_strands_agent): } mock_strands_agent.tool_registry.get_all_tools_config.return_value = mock_tool_config - a2a_agent = A2AServer(mock_strands_agent) - - # This should raise a KeyError when accessing agent_skills due to missing description + # A2AServer builds its AgentCard (and thus agent_skills) eagerly during construction, + # because DefaultRequestHandler requires the AgentCard upfront. with pytest.raises(KeyError): - _ = a2a_agent.agent_skills + A2AServer(mock_strands_agent) def test_agent_skills_handles_missing_name(mock_strands_agent): @@ -259,11 +256,8 @@ def test_agent_skills_handles_missing_name(mock_strands_agent): } mock_strands_agent.tool_registry.get_all_tools_config.return_value = mock_tool_config - a2a_agent = A2AServer(mock_strands_agent) - - # This should raise a KeyError when accessing agent_skills due to missing name with pytest.raises(KeyError): - _ = a2a_agent.agent_skills + A2AServer(mock_strands_agent) def test_agent_skills_setter(mock_strands_agent): @@ -387,25 +381,25 @@ def test_explicit_skills_override_tools(mock_strands_agent): def test_skills_not_loaded_during_initialization(mock_strands_agent): - """Test that skills are not loaded from tools during initialization.""" - # Create a mock that would raise an exception if called - mock_strands_agent.tool_registry.get_all_tools_config.side_effect = Exception("Should not be called during init") + """Test that agent_skills is not cached and is recomputed from tools on each access. + + A2AServer.__init__ builds the AgentCard (and thus agent_skills) once eagerly, because + DefaultRequestHandler requires the AgentCard upfront. ``agent_skills`` itself is not + cached, though: ``_agent_skills`` stays None and every access re-derives skills from the + tool registry. + """ + mock_tool_config = {"test_tool": {"name": "test_tool", "description": "A test tool"}} + mock_strands_agent.tool_registry.get_all_tools_config.return_value = mock_tool_config - # This should not raise an exception because tools are not accessed during initialization a2a_agent = A2AServer(mock_strands_agent) - # Verify that _agent_skills is None assert a2a_agent._agent_skills is None - # Reset the mock to return proper data for when skills are actually accessed - mock_tool_config = {"test_tool": {"name": "test_tool", "description": "A test tool"}} - mock_strands_agent.tool_registry.get_all_tools_config.side_effect = None - mock_strands_agent.tool_registry.get_all_tools_config.return_value = mock_tool_config - - # Now accessing skills should work skills = a2a_agent.agent_skills + assert len(skills) == 1 assert skills[0].name == "test_tool" + assert a2a_agent._agent_skills is None def test_public_agent_card_with_custom_skills(mock_strands_agent): @@ -618,7 +612,7 @@ def test_public_agent_card_with_http_url(mock_strands_agent): card = a2a_agent.public_agent_card assert isinstance(card, AgentCard) - assert card.url == "https://my-alb.amazonaws.com/agent1/" + assert card.supported_interfaces[0].url == "https://my-alb.amazonaws.com/agent1/" assert card.name == "Test Agent" assert card.description == "A test agent for unit testing" @@ -640,7 +634,7 @@ def test_agent_card_url_override(mock_strands_agent): assert a2a_agent.agent_card_url == "https://my-alb.amazonaws.com/agent1" card = a2a_agent.public_agent_card - assert card.url == "https://my-alb.amazonaws.com/agent1" + assert card.supported_interfaces[0].url == "https://my-alb.amazonaws.com/agent1" def test_to_starlette_app_with_mounting(mock_strands_agent): @@ -702,7 +696,7 @@ def test_backwards_compatibility_without_http_url(mock_strands_agent): # Agent card should use the traditional URL card = a2a_agent.public_agent_card - assert card.url == "http://localhost:9000/" + assert card.supported_interfaces[0].url == "http://localhost:9000/" def test_mount_path_logging(mock_strands_agent, caplog): @@ -793,11 +787,11 @@ def test_serve_at_root_fastapi_mounting_behavior(mock_strands_agent): client_mounted = TestClient(app_mounted) # Should work at mounted path - response = client_mounted.get("/agent1/.well-known/agent.json") + response = client_mounted.get("/agent1/.well-known/agent-card.json") assert response.status_code == 200 # Should not work at root - response = client_mounted.get("/.well-known/agent.json") + response = client_mounted.get("/.well-known/agent-card.json") assert response.status_code == 404 @@ -813,11 +807,11 @@ def test_serve_at_root_fastapi_root_behavior(mock_strands_agent): client_root = TestClient(app_root) # Should work at root - response = client_root.get("/.well-known/agent.json") + response = client_root.get("/.well-known/agent-card.json") assert response.status_code == 200 # Should not work at mounted path (since we're serving at root) - response = client_root.get("/agent1/.well-known/agent.json") + response = client_root.get("/agent1/.well-known/agent-card.json") assert response.status_code == 404 @@ -833,7 +827,7 @@ def test_serve_at_root_starlette_behavior(mock_strands_agent): client_mounted = TestClient(app_mounted) # Should work at mounted path - response = client_mounted.get("/agent1/.well-known/agent.json") + response = client_mounted.get("/agent1/.well-known/agent-card.json") assert response.status_code == 200 # Serve at root @@ -842,7 +836,7 @@ def test_serve_at_root_starlette_behavior(mock_strands_agent): client_root = TestClient(app_root) # Should work at root - response = client_root.get("/.well-known/agent.json") + response = client_root.get("/.well-known/agent-card.json") assert response.status_code == 200 @@ -857,8 +851,8 @@ def test_serve_at_root_alb_scenarios(mock_strands_agent): app_preserved = server_preserved.to_fastapi_app() client_preserved = TestClient(app_preserved) - # Container receives /agent1/.well-known/agent.json - response = client_preserved.get("/agent1/.well-known/agent.json") + # Container receives /agent1/.well-known/agent-card.json + response = client_preserved.get("/agent1/.well-known/agent-card.json") assert response.status_code == 200 agent_data = response.json() assert agent_data["url"] == "http://my-alb.amazonaws.com/agent1/" @@ -870,8 +864,8 @@ def test_serve_at_root_alb_scenarios(mock_strands_agent): app_stripped = server_stripped.to_fastapi_app() client_stripped = TestClient(app_stripped) - # Container receives /.well-known/agent.json (path stripped by ALB) - response = client_stripped.get("/.well-known/agent.json") + # Container receives /.well-known/agent-card.json (path stripped by ALB) + response = client_stripped.get("/.well-known/agent-card.json") assert response.status_code == 200 agent_data = response.json() assert agent_data["url"] == "http://my-alb.amazonaws.com/agent1/" @@ -947,7 +941,7 @@ def test_serve_with_overridden_host_port_updates_agent_card_url(mock_run, mock_s # Verify the agent card reflects the updated URL card = a2a_agent.public_agent_card - assert card.url == "http://localhost:9210/" + assert card.supported_interfaces[0].url == "http://localhost:9210/" # Verify uvicorn was called with the overridden parameters mock_run.assert_called_once() @@ -1033,7 +1027,7 @@ def test_serve_with_explicit_http_url_does_not_override_url(mock_run, mock_stran # Verify the agent card still shows the public URL card = a2a_agent.public_agent_card - assert card.url == "https://my-alb.amazonaws.com/agent1/" + assert card.supported_interfaces[0].url == "https://my-alb.amazonaws.com/agent1/" @patch("uvicorn.run") diff --git a/strands-py/tests_integ/test_a2a_executor.py b/strands-py/tests_integ/test_a2a_executor.py index 7ae10efc29..599543b20e 100644 --- a/strands-py/tests_integ/test_a2a_executor.py +++ b/strands-py/tests_integ/test_a2a_executor.py @@ -1,17 +1,42 @@ """Integration tests for A2A executor with real file processing.""" -import base64 import os import threading import time +from uuid import uuid4 import pytest import requests import uvicorn +from a2a.helpers import new_data_part +from a2a.types import Message, Part, Role, SendMessageRequest +from google.protobuf.json_format import MessageToDict from strands import Agent from strands.multiagent.a2a import A2AServer +_A2A_VERSION_HEADER = {"A2A-Version": "1.0"} + + +def _send_message(message: Message, request_id: str, *, port: int) -> dict: + """POST a SendMessage JSON-RPC request and return its `result`, failing on a JSON-RPC error.""" + payload = { + "jsonrpc": "2.0", + "id": request_id, + "method": "SendMessage", + "params": MessageToDict(SendMessageRequest(message=message)), + } + response = requests.post( + f"http://127.0.0.1:{port}", + headers={"Content-Type": "application/json", **_A2A_VERSION_HEADER}, + json=payload, + timeout=30, + ) + assert response.status_code == 200 + response_data = response.json() + assert "error" not in response_data, response_data + return response_data["result"] + @pytest.mark.asyncio async def test_a2a_executor_with_real_image(): @@ -21,9 +46,6 @@ async def test_a2a_executor_with_real_image(): with open(test_image_path, "rb") as f: original_image_bytes = f.read() - # Encode as base64 (A2A format) - base64_image = base64.b64encode(original_image_bytes).decode("utf-8") - # Create real Strands agent strands_agent = Agent(name="Test Image Agent", description="Agent for testing image processing") @@ -36,69 +58,194 @@ async def test_a2a_executor_with_real_image(): server_thread.start() time.sleep(1) # Give server time to start - try: - # Create A2A message with real image - message_payload = { - "jsonrpc": "2.0", - "id": "test-image-request", - "method": "message/send", - "params": { - "message": { - "messageId": "msg-123", - "role": "user", - "parts": [ - { - "kind": "text", - "text": "What primary color is this image, respond with NONE if you are unsure", - "metadata": None, - }, - { - "kind": "file", - "file": {"name": "image.png", "mimeType": "image/png", "bytes": base64_image}, - "metadata": None, - }, - ], - } - }, - } - - # Send request to A2A server - response = requests.post( - "http://127.0.0.1:9001", headers={"Content-Type": "application/json"}, json=message_payload, timeout=30 - ) - - # Verify response - assert response.status_code == 200 - response_data = response.json() - assert "completed" == response_data["result"]["status"]["state"] - all_text = " ".join( - part["text"] - for artifact in response_data["result"]["artifacts"] - for part in artifact["parts"] - if part.get("kind") == "text" - ).lower() - assert "yellow" in all_text - - except Exception as e: - pytest.fail(f"Integration test failed: {e}") - - -def test_a2a_executor_image_roundtrip(): - """Test that image data survives the A2A base64 encoding/decoding roundtrip.""" - # Read the test image - test_image_path = os.path.join(os.path.dirname(__file__), "resources/yellow.png") - with open(test_image_path, "rb") as f: - original_bytes = f.read() + message = Message( + message_id=str(uuid4()), + role=Role.ROLE_USER, + parts=[ + Part(text="What primary color is this image, respond with NONE if you are unsure"), + Part(raw=original_image_bytes, media_type="image/png", filename="image.png"), + ], + ) + result = _send_message(message, "test-image-request", port=9001) + + task = result["task"] + assert task["status"]["state"] == "TASK_STATE_COMPLETED" + all_text = " ".join( + part["text"] for artifact in task["artifacts"] for part in artifact["parts"] if "text" in part + ).lower() + assert "yellow" in all_text + + +@pytest.mark.asyncio +async def test_a2a_executor_interrupt_resume_over_http(): + """Park a tool interrupt and resume it over real JSON-RPC, matching production wire shapes. + + Uses a `MockedModelProvider` (no live model call) so the interrupt/resume turn sequence is + deterministic, while still exercising the real HTTP + JSON-RPC + protobuf wire round trip. + """ + from strands import tool + from strands.types.tools import ToolContext + from tests.fixtures.mocked_model_provider import MockedModelProvider + + @tool(name="approval_tool", context=True) + def approval_tool(tool_context: ToolContext) -> str: + return tool_context.interrupt("approval_interrupt", reason="need approval") + + tool_use_message = { + "role": "assistant", + "content": [{"toolUse": {"toolUseId": "t1", "name": "approval_tool", "input": {}}}], + } + final_message = {"role": "assistant", "content": [{"text": "done"}]} + model = MockedModelProvider([tool_use_message, final_message]) + strands_agent = Agent( + name="Test Interrupt Agent", + description="Agent for testing interrupt park/resume over A2A", + model=model, + tools=[approval_tool], + callback_handler=None, + ) + + a2a_server = A2AServer(agent=strands_agent, port=9002) + fastapi_app = a2a_server.to_fastapi_app() + + server_thread = threading.Thread(target=lambda: uvicorn.run(fastapi_app, port=9002), daemon=True) + server_thread.start() + time.sleep(1) + + initial_message = Message( + message_id=str(uuid4()), + role=Role.ROLE_USER, + parts=[Part(text="Use the approval_tool now.")], + ) + task = _send_message(initial_message, "test-interrupt-request", port=9002)["task"] + assert task["status"]["state"] == "TASK_STATE_INPUT_REQUIRED" + + interrupt_data = next( + part["data"] for part in task["status"]["message"]["parts"] if "data" in part and "interrupts" in part["data"] + ) + interrupt_id = interrupt_data["interrupts"][0]["interruptId"] + + resume_message = Message( + message_id=str(uuid4()), + task_id=task["id"], + context_id=task["contextId"], + role=Role.ROLE_USER, + parts=[new_data_part({"interruptResponse": {"interruptId": interrupt_id, "response": "APPROVE"}})], + ) + resumed_task = _send_message(resume_message, "test-interrupt-resume", port=9002)["task"] + assert resumed_task["status"]["state"] == "TASK_STATE_COMPLETED" + - # Simulate A2A protocol: encode to base64 string - base64_string = base64.b64encode(original_bytes).decode("utf-8") +@pytest.mark.asyncio +async def test_a2a_agent_content_round_trip_streaming_disabled(): + """End-to-end: A2AAgent → A2AServer with enable_a2a_compliant_streaming=False. + + Asserts the reply *content* is preserved through convert_responses_to_agent_result, + not just stop_reason. + """ + from strands.agent.a2a_agent import A2AAgent + from tests.fixtures.mocked_model_provider import MockedModelProvider + + reply_text = "The capital of France is Paris." + model = MockedModelProvider([{"role": "assistant", "content": [{"text": reply_text}]}]) + strands_agent = Agent( + name="Content Agent", + description="Agent that returns a known answer", + model=model, + callback_handler=None, + ) + + port = 9003 + a2a_server = A2AServer(agent=strands_agent, port=port, enable_a2a_compliant_streaming=False) + fastapi_app = a2a_server.to_fastapi_app() + + server_thread = threading.Thread(target=lambda: uvicorn.run(fastapi_app, port=port), daemon=True) + server_thread.start() + time.sleep(1) - # Simulate executor decoding - decoded_bytes = base64.b64decode(base64_string) + a2a_agent = A2AAgent(endpoint=f"http://127.0.0.1:{port}") + result = await a2a_agent.invoke_async("What is the capital of France?") + + assert result.stop_reason == "end_turn" + content_text = " ".join(block["text"] for block in result.message["content"] if "text" in block) + assert reply_text in content_text + + +@pytest.mark.asyncio +async def test_a2a_agent_content_round_trip_streaming_enabled(): + """End-to-end: A2AAgent → A2AServer with enable_a2a_compliant_streaming=True. + + Asserts content is correctly accumulated from artifact_update deltas through + convert_responses_to_agent_result. + """ + from strands.agent.a2a_agent import A2AAgent + from tests.fixtures.mocked_model_provider import MockedModelProvider + + reply_text = "The capital of France is Paris." + model = MockedModelProvider([{"role": "assistant", "content": [{"text": reply_text}]}]) + strands_agent = Agent( + name="Streaming Content Agent", + description="Agent that returns a known answer via compliant streaming", + model=model, + callback_handler=None, + ) + + port = 9004 + a2a_server = A2AServer(agent=strands_agent, port=port, enable_a2a_compliant_streaming=True) + fastapi_app = a2a_server.to_fastapi_app() + + server_thread = threading.Thread(target=lambda: uvicorn.run(fastapi_app, port=port), daemon=True) + server_thread.start() + time.sleep(1) + + a2a_agent = A2AAgent(endpoint=f"http://127.0.0.1:{port}") + result = await a2a_agent.invoke_async("What is the capital of France?") + + assert result.stop_reason == "end_turn" + content_text = " ".join(block["text"] for block in result.message["content"] if "text" in block) + assert reply_text in content_text + + +@pytest.mark.asyncio +async def test_a2a_agent_interrupt_round_trip(): + """End-to-end: A2AAgent sees interrupt content through convert_responses_to_agent_result. + + Parks on input_required with a reason message, asserts the A2AAgent receives both + stop_reason="interrupt" AND the interrupt question text in the result content. + """ + from strands import tool + from strands.agent.a2a_agent import A2AAgent + from strands.types.tools import ToolContext + from tests.fixtures.mocked_model_provider import MockedModelProvider + + @tool(name="ask_user", context=True) + def ask_user(tool_context: ToolContext) -> str: + return tool_context.interrupt("ask_interrupt", reason="Do you approve this action?") + + tool_use_message = { + "role": "assistant", + "content": [{"toolUse": {"toolUseId": "t1", "name": "ask_user", "input": {}}}], + } + model = MockedModelProvider([tool_use_message]) + strands_agent = Agent( + name="Interrupt Agent", + description="Agent that asks for approval", + model=model, + tools=[ask_user], + callback_handler=None, + ) + + port = 9005 + a2a_server = A2AServer(agent=strands_agent, port=port, enable_a2a_compliant_streaming=True) + fastapi_app = a2a_server.to_fastapi_app() + + server_thread = threading.Thread(target=lambda: uvicorn.run(fastapi_app, port=port), daemon=True) + server_thread.start() + time.sleep(1) - # Verify perfect roundtrip - assert decoded_bytes == original_bytes - assert len(decoded_bytes) == len(original_bytes) + a2a_agent = A2AAgent(endpoint=f"http://127.0.0.1:{port}") + result = await a2a_agent.invoke_async("Please ask for approval.") - # Verify it's actually image data (PNG signature) - assert decoded_bytes.startswith(b"\x89PNG\r\n\x1a\n") + assert result.stop_reason == "interrupt" + content_text = " ".join(block["text"] for block in result.message["content"] if "text" in block) + assert "approve" in content_text.lower() or "interrupt" in content_text.lower()