diff --git a/doc/code/targets/0_prompt_targets.md b/doc/code/targets/0_prompt_targets.md index 85de4c9f50..4a37974761 100644 --- a/doc/code/targets/0_prompt_targets.md +++ b/doc/code/targets/0_prompt_targets.md @@ -73,8 +73,100 @@ Here are some examples: | **OpenAIChatTarget** (e.g., GPT-4) | **Yes** (multi-turn + editable history) | Designed for conversational prompts (system messages, conversation history, etc.). | | **OpenAIImageTarget** | **No** | Used for image generation; does not manage conversation history. | | **HTTPTarget** | **No** | Generic HTTP target. Some apps might allow conversation history, but this target doesn't handle it. | +| **A2ATarget** | **No** (multi-turn, but no editable history) | Text-only Agent-to-Agent endpoints through the official a2a-sdk (v0.3 / v1.0). | | **AzureBlobStorageTarget** | **No** | Used primarily for storage; not for conversation-based AI. | +## A2A agents + +Install the optional client with `pip install "pyrit[a2a]"` (also included in +`pyrit[all]`). `A2ATarget` uses the official `a2a-sdk` 1.x client for JSON-RPC. +Set `protocol_version="0.3"` (the default) or `"1.0"` for a known endpoint. +Set `"auto"` to discover the agent card. Discovery errors are not hidden by +a fallback to another protocol. Use `agent_card_path` for a non-standard +relative card path. + +The adapter sends one text piece per turn. The agent owns its context and +history. The target cannot edit or replay that history, send system-role +messages, enforce native JSON output, or send audio, images, or files. +Custom capabilities cannot enable these unsupported features or disable +multi-turn support. Disabling it would squash local history while the agent +retains the same history. To restore an +existing upstream context on a new target instance, use +`set_conversation_context`. A local history without an upstream context ID +is rejected. `reset_conversation_async` forgets the local mapping; it does +not delete remote state. Sends and resets for the same conversation are +serialized. Different conversations can run concurrently. +`set_conversation_context` raises an error if that conversation is in use. + +The target retries submission only after an explicit HTTP 429. Retries use +the same A2A message ID. Rate-limit errors during polling retry only the poll. +Polling honors `Retry-After` (seconds or HTTP date), with exponential backoff +when that header is absent or invalid. The polling deadline bounds these waits. +`max_requests_per_minute` applies to submission attempts and task polls, not +agent-card discovery. +Empty results, task failures, and agent errors that mention a downstream +rate limit do not cause automatic resubmission. A polling timeout does not +cancel the remote task. Do not blindly resubmit an action after a timeout. +`request_timeout_seconds` sets the HTTPX connect, read, write, and pool +timeouts and defaults to `task_timeout_seconds`. These are per-phase +timeouts, not a total turn deadline. Instead, pass HTTPX `timeout` for +per-phase settings or `timeout=None` to disable HTTP timeouts. Do not pass +both timeout options. Timeouts must be finite and positive. +`task_timeout_seconds` is a separate deadline that starts after submission +returns a pending task. It bounds polling, rate-limit waits, and in-flight +poll requests, even with `timeout=None`. `poll_interval_seconds` must be +finite and non-negative. + +Identifiers include the endpoint, requested protocol, card path, and optional +`routing_identifier`. Set this non-secret deployment label when HTTP headers +select a different agent at the same endpoint. Do not put tokens in the URL, +card path, or routing label. Credentials, HTTP headers, and timeout settings +are not copied into identifier parameters. + +`auth_token` accepts a static token or an async callable that returns a token. +The callable runs before each HTTP request, including discovery, submission +retries, and polls. Use a provider that caches tokens and refreshes them +before expiry. Do not combine `auth_token` with HTTPX `auth` or an +`Authorization` header. Custom HTTPX authentication is supported when +`auth_token` is not set. + +For Foundry, keep the credential open for the full target operation: + +```python +from azure.identity.aio import DefaultAzureCredential, get_bearer_token_provider +from pyrit.prompt_target import A2ATarget + +async with DefaultAzureCredential() as credential: + token_provider = get_bearer_token_provider(credential, "https://ai.azure.com/.default") + target = A2ATarget( + endpoint=agent_endpoint, + protocol_version="auto", + agent_card_path="agentCard/v1.0", + auth_token=token_provider, + ) + responses = await target.send_prompt_async(message=message) +``` + +### A2A integration tests + +The local tests run an official SDK server on an ephemeral loopback port. +They cover 0.3 and 1.0, discovery, persisted multi-turn history, task +continuation, failure, and prevention of duplicate submissions. From a +development environment with the `all` extra, set `RUN_ALL_TESTS=true` and run: + +```text +uv run pytest tests/integration/targets/test_a2a_target_integration.py -k "not foundry" +``` + +The separate Foundry test also requires `RUN_ALL_TESTS=true` and +`A2A_FOUNDRY_ENDPOINT`. It uses `A2A_FOUNDRY_AUTH_TOKEN` if supplied; otherwise, +it uses a refreshable Entra token provider through `DefaultAzureCredential` for +`https://ai.azure.com/.default`. The identity must have access to the agent +(for example, the Agent Consumer role). `A2A_FOUNDRY_CARD_PATH` +defaults to `agentCard/v1.0`; set it to +`agentCard/v0.3` for a preview endpoint. The test requires a successful text +response, not just a non-empty error. It does not create Azure resources. + ## Target Capabilities Every `PromptTarget` exposes a `TargetConfiguration` (via `target.configuration`) that declares what the target natively supports. This lets attacks, converters, and scorers reason about whether a given target is suitable for a given workflow — and, where possible, adapt automatically when a capability is missing. diff --git a/pyproject.toml b/pyproject.toml index df88176d1d..bbcb5f2f81 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -94,6 +94,7 @@ dev = [ "pytest-timeout>=2.4.0", "pytest-xdist>=3.6.1", "respx>=0.22.0", + "sse-starlette>=3.0.0", "ruff>=0.14.4", "types-aiofiles>=24.1.0", "types-PyYAML>=6.0.12.20250516", @@ -136,8 +137,13 @@ litellm = [ "litellm>=1.84.0", ] +a2a = [ + "a2a-sdk>=1.1.5,<2", +] + # all includes all functional dependencies excluding the ones from the "dev" dependency group all = [ + "a2a-sdk>=1.1.5,<2", "accelerate>=1.7.0", "azure-ai-ml>=1.32.0", "azure-cognitiveservices-speech>=1.44.0", diff --git a/pyrit/prompt_target/__init__.py b/pyrit/prompt_target/__init__.py index b33d87d21b..25cde5eaae 100644 --- a/pyrit/prompt_target/__init__.py +++ b/pyrit/prompt_target/__init__.py @@ -14,6 +14,7 @@ from pyrit.common.lazy_imports import get_lazy_dir, resolve_lazy_export if TYPE_CHECKING: + from pyrit.prompt_target.a2a_target import A2ATarget from pyrit.prompt_target.azure_blob_storage_target import AzureBlobStorageTarget from pyrit.prompt_target.azure_ml_chat_target import AzureMLChatTarget from pyrit.prompt_target.common.conversation_normalization_pipeline import ConversationNormalizationPipeline @@ -73,6 +74,7 @@ from pyrit.prompt_target.websocket_target import WebsocketTarget _LAZY_EXPORTS: dict[str, str | tuple[str, str | None]] = { + "A2ATarget": "pyrit.prompt_target.a2a_target", "TargetTraceConfig": "pyrit.prompt_target.common.target_trace_config", "AzureBlobStorageTarget": "pyrit.prompt_target.azure_blob_storage_target", "AzureMLChatTarget": "pyrit.prompt_target.azure_ml_chat_target", diff --git a/pyrit/prompt_target/a2a_target.py b/pyrit/prompt_target/a2a_target.py new file mode 100644 index 0000000000..a434e9e9bb --- /dev/null +++ b/pyrit/prompt_target/a2a_target.py @@ -0,0 +1,520 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +import asyncio +import logging +import math +import time +import uuid +from collections.abc import AsyncGenerator +from dataclasses import dataclass +from email.utils import parsedate_to_datetime +from typing import TYPE_CHECKING, Any, Literal, cast +from weakref import WeakValueDictionary + +import httpx + +from pyrit.common.net_utility import get_httpx_client +from pyrit.exceptions import EmptyResponseException, RateLimitException, pyrit_target_retry +from pyrit.exceptions.exception_classes import CONTENT_FILTER_MARKERS +from pyrit.models import ( + ComponentIdentifier, + Message, + MessagePiece, + construct_response_from_request, +) +from pyrit.prompt_target.common.prompt_target import PromptTarget +from pyrit.prompt_target.common.target_capabilities import CapabilityName, TargetCapabilities +from pyrit.prompt_target.common.target_configuration import TargetConfiguration +from pyrit.prompt_target.common.utils import limit_requests_per_minute + +if TYPE_CHECKING: + from collections.abc import Awaitable, Callable + + from a2a.client import Client + from a2a.types import a2a_pb2 + +logger = logging.getLogger(__name__) + + +class _BearerTokenAuth(httpx.Auth): + """Resolve bearer credentials for each HTTP request, including polls and retries.""" + + def __init__(self, token: str | Callable[[], Awaitable[str]]) -> None: + self._token = token + + async def async_auth_flow( # pyrit-async-suffix-exempt + self, request: httpx.Request + ) -> AsyncGenerator[httpx.Request, httpx.Response]: + token = self._token if isinstance(self._token, str) else await self._token() + if not isinstance(token, str) or not token.strip(): + raise ValueError("The A2A token provider must return a non-empty token string.") + request.headers["Authorization"] = f"Bearer {token}" + yield request + + +@dataclass +class _A2AConversationState: + """Server-issued identifiers that continue one PyRIT conversation on the agent.""" + + context_id: str | None = None + open_task_id: str | None = None + + +class A2ATarget(PromptTarget): + """ + A PromptTarget for interacting with agents speaking the Agent-to-Agent (A2A) protocol. + + The Agent-to-Agent protocol defines task-based message exchange between autonomous agents. + This target adapts PyRIT's prompt target interface to the official `a2a-sdk`, supporting: + - A2A v0.3 compatibility (JSON-RPC `message/send` and `tasks/get`) + - A2A v1.0 (JSON-RPC `SendMessage` and `GetTask`) + - Agent card discovery (`protocol_version="auto"`) + + Each PyRIT conversation maps to an upstream A2A context (`context_id`), so the agent + keeps its own server-side state across turns. Tasks left in `input-required` or + `auth-required` states are continued with their task ID. + + Rate limits (429) during polling are retried against the existing task without resubmitting + the prompt. Failed or canceled tasks produce error responses, and questions from + `input-required` tasks are extracted from the task status message. + """ + + _DEFAULT_CONFIGURATION: TargetConfiguration = TargetConfiguration( + capabilities=TargetCapabilities( + supports_multi_turn=True, + input_modalities=frozenset({frozenset(["text"])}), + ) + ) + + def __init__( + self, + *, + endpoint: str, + auth_token: str | Callable[[], Awaitable[str]] | None = None, + api_key: str | None = None, + api_key_header: str = "X-API-Key", + protocol_version: Literal["auto", "1.0", "0.3"] = "0.3", + agent_card_path: str | None = None, + routing_identifier: str | None = None, + task_timeout_seconds: float = 120.0, + request_timeout_seconds: float | None = None, + poll_interval_seconds: float = 1.0, + max_requests_per_minute: int | None = None, + custom_configuration: TargetConfiguration | None = None, + **httpx_client_kwargs: Any, + ) -> None: + """ + Initialize the A2ATarget. + + Args: + endpoint (str): The target URL of the A2A agent endpoint. + auth_token (str | Callable[[], Awaitable[str]] | None): Bearer token or async token provider. + The provider is called before each HTTP request and owns token refresh and credential cleanup. + api_key (str | None): Custom API key for the agent. + api_key_header (str): Header name for the API key (defaults to "X-API-Key"). + protocol_version (Literal["auto", "1.0", "0.3"]): A2A protocol version. Defaults to "0.3". + agent_card_path (str | None): Relative card path for automatic discovery. + routing_identifier (str | None): Non-secret deployment label for header-based routing. + This is stored in the identifier; do not put credentials here. + task_timeout_seconds (float): Polling deadline after submission returns. Must be finite and positive. + request_timeout_seconds (float | None): Timeout for individual HTTP requests. Defaults to + task_timeout_seconds. + poll_interval_seconds (float): Finite, non-negative delay between task polls. Defaults to 1.0. + max_requests_per_minute (int | None): Rate limit for submission attempts and polls, excluding discovery. + custom_configuration (TargetConfiguration | None): Custom target capabilities override. + **httpx_client_kwargs: Additional HTTPX options. ``timeout`` supports HTTPX's per-phase + settings or None to disable HTTP timeouts. Do not combine it with request_timeout_seconds. + + Raises: + ImportError: If the a2a-sdk package is not installed. + ValueError: If the protocol, capabilities, or timeouts are invalid, or authentication options conflict. + """ + try: + from a2a.client import ClientFactory # noqa: F401 + except ImportError as exc: + raise ImportError( + "The a2a-sdk package is required for A2ATarget. Install it with `pip install pyrit[a2a]`." + ) from exc + + if protocol_version not in ("auto", "1.0", "0.3"): + raise ValueError( + f"Unsupported A2A protocol_version '{protocol_version}'. Expected 'auto', '1.0', or '0.3'." + ) + + if custom_configuration: + self._validate_capabilities(custom_configuration.capabilities) + + for name, value in ( + ("task_timeout_seconds", task_timeout_seconds), + ("request_timeout_seconds", request_timeout_seconds), + ("poll_interval_seconds", poll_interval_seconds), + ): + allow_zero = name == "poll_interval_seconds" + if value is not None and (not math.isfinite(value) or value < 0 or (value == 0 and not allow_zero)): + requirement = "non-negative" if allow_zero else "positive" + raise ValueError(f"{name} must be finite and {requirement}.") + if request_timeout_seconds is not None and "timeout" in httpx_client_kwargs: + raise ValueError("Specify either request_timeout_seconds or HTTPX timeout, not both.") + if auth_token is not None: + if "auth" in httpx_client_kwargs: + raise ValueError("Specify either auth_token or HTTPX auth, not both.") + if "Authorization" in httpx.Headers(httpx_client_kwargs.get("headers")) or ( + api_key and api_key_header.lower() == "authorization" + ): + raise ValueError("Specify either auth_token or an Authorization header, not both.") + if "timeout" in httpx_client_kwargs: + timeout = httpx.Timeout(httpx_client_kwargs["timeout"]) + for value in timeout.as_dict().values(): + if value is not None and (not math.isfinite(value) or value <= 0): + raise ValueError("HTTPX timeouts must be finite and positive, or None.") + + super().__init__( + endpoint=endpoint, + max_requests_per_minute=max_requests_per_minute, + custom_configuration=custom_configuration, + ) + + self._endpoint = endpoint.rstrip("/") + self._auth_token = auth_token + self._api_key = api_key + self._api_key_header = api_key_header + self._protocol_version: Literal["auto", "1.0", "0.3"] = protocol_version + self._agent_card_path = agent_card_path + self._routing_identifier = routing_identifier + self._task_timeout_seconds = task_timeout_seconds + self._request_timeout_seconds = ( + request_timeout_seconds if request_timeout_seconds is not None else task_timeout_seconds + ) + self._poll_interval_seconds = poll_interval_seconds + self._conversations: dict[str, _A2AConversationState] = {} + self._conversation_locks: WeakValueDictionary[str, asyncio.Lock] = WeakValueDictionary() + self._httpx_client_kwargs = httpx_client_kwargs + + def _build_identifier(self) -> ComponentIdentifier: + return self._create_identifier( + params={ + "protocol_version": self._protocol_version, + "agent_card_path": self._agent_card_path, + "routing_identifier": self._routing_identifier, + } + ) + + @staticmethod + def _validate_capabilities(capabilities: TargetCapabilities) -> None: + if not capabilities.supports_multi_turn: + raise ValueError("A2ATarget requires supports_multi_turn=True because it maintains upstream context.") + for name, modalities in ( + ("input", capabilities.input_modalities), + ("output", capabilities.output_modalities), + ): + if modalities != frozenset({frozenset({"text"})}): + raise ValueError(f"A2ATarget only supports text {name} modality.") + for capability in CapabilityName: + if capability != CapabilityName.MULTI_TURN and capabilities.includes(capability=capability): + raise ValueError(f"A2ATarget does not support {capability.value}.") + + def apply_capabilities(self, *, capabilities: TargetCapabilities) -> None: + """Replace capabilities only if the text-only adapter can implement them.""" + self._validate_capabilities(capabilities) + super().apply_capabilities(capabilities=capabilities) + + def _validate_request(self, *, normalized_conversation: list[Message]) -> None: + self._validate_capabilities(self.capabilities) + super()._validate_request(normalized_conversation=normalized_conversation) + + def _build_headers(self) -> dict[str, str]: + headers: dict[str, str] = { + "Accept": "application/json", + } + if self._api_key: + headers[self._api_key_header] = self._api_key + return headers + + def set_conversation_context( + self, + *, + conversation_id: str, + context_id: str, + open_task_id: str | None = None, + ) -> None: + """ + Explicitly set or restore the upstream A2A context for a conversation. + + Args: + conversation_id (str): PyRIT conversation ID. + context_id (str): Upstream A2A context ID. + open_task_id (str | None): Open task ID waiting for input. + + Raises: + ValueError: If context_id is empty. + RuntimeError: If a send or reset is in progress for this conversation. + """ + if not context_id.strip(): + raise ValueError("A2ATarget requires a non-empty upstream context ID.") + lock = self._conversation_locks.get(conversation_id) + if lock is not None and lock.locked(): + raise RuntimeError("Cannot replace an A2A context while its conversation is in use.") + self._conversations[conversation_id] = _A2AConversationState(context_id=context_id, open_task_id=open_task_id) + + async def reset_conversation_async(self, *, conversation_id: str) -> None: + """ + Forget the A2A context held for a conversation. + + Args: + conversation_id (str): PyRIT conversation ID. + """ + lock = self._conversation_locks.setdefault(conversation_id, asyncio.Lock()) + async with lock: + self._conversations.pop(conversation_id, None) + + @staticmethod + def _rate_limit_response(exc: Exception) -> httpx.Response | None: + """Return only an explicit HTTP 429, not a rate-limit string in an agent error.""" + cause = exc if isinstance(exc, httpx.HTTPStatusError) else exc.__cause__ + if isinstance(cause, httpx.HTTPStatusError) and cause.response.status_code == 429: + return cause.response + return None + + @staticmethod + def _retry_after_seconds(response: httpx.Response) -> float | None: + value = response.headers.get("Retry-After") + if value is None: + return None + try: + delay = float(value) + except ValueError: + try: + delay = parsedate_to_datetime(value).timestamp() - time.time() + except (ValueError, TypeError, OverflowError): + logger.warning("Invalid A2A Retry-After header; using polling backoff.") + return None + if not math.isfinite(delay): + logger.warning("Non-finite A2A Retry-After header; using polling backoff.") + return None + return max(0.0, delay) + + @staticmethod + def _extract_message_text(msg: a2a_pb2.Message) -> str: + texts = [part.text for part in msg.parts if part.HasField("text")] + return "\n".join(texts).strip() + + @staticmethod + def _extract_task_artifacts_text(task: a2a_pb2.Task) -> str: + texts: list[str] = [] + for artifact in task.artifacts: + texts.extend(part.text for part in artifact.parts if part.HasField("text")) + return "\n".join(texts).strip() + + async def _create_a2a_client_async(self, http_client: httpx.AsyncClient) -> Client: + from a2a.client import ClientConfig, ClientFactory + from a2a.types import a2a_pb2 + from a2a.utils.constants import TransportProtocol + + config = ClientConfig(streaming=False, httpx_client=http_client, accepted_output_modes=["text/plain"]) + factory = ClientFactory(config) + + if self._protocol_version == "auto": + return await factory.create_from_url(self._endpoint, relative_card_path=self._agent_card_path) + + card = a2a_pb2.AgentCard( + name="A2AAgent", + supported_interfaces=[ + a2a_pb2.AgentInterface( + url=self._endpoint, + protocol_binding=TransportProtocol.JSONRPC, + protocol_version=self._protocol_version, + ) + ], + ) + return factory.create(card) + + async def _await_task_async(self, *, client: Client, task: a2a_pb2.Task) -> a2a_pb2.Task: + from a2a.types import a2a_pb2 + from a2a.utils.errors import A2AError + + delay = self._poll_interval_seconds + try: + async with asyncio.timeout(self._task_timeout_seconds): + while task.status.state in (a2a_pb2.TASK_STATE_SUBMITTED, a2a_pb2.TASK_STATE_WORKING): + await asyncio.sleep(delay) + try: + task = cast("a2a_pb2.Task", await self._get_task_async(client=client, task_id=task.id)) + delay = self._poll_interval_seconds + except (A2AError, httpx.HTTPError) as exc: + response = self._rate_limit_response(exc) + if response is None: + raise + retry_after = self._retry_after_seconds(response) + delay = ( + max(self._poll_interval_seconds, retry_after) + if retry_after is not None + else min(self._task_timeout_seconds, max(1.0, delay * 2)) + ) + logger.warning( + "Rate limit polling A2A task %s; retrying only the poll in %s seconds.", task.id, delay + ) + except TimeoutError as exc: + state_name = a2a_pb2.TaskState.Name(task.status.state) + raise TimeoutError( + f"A2A task {task.id} still in state {state_name} after {self._task_timeout_seconds} seconds." + ) from exc + return task + + @limit_requests_per_minute + async def _get_task_async(self, *, client: Client, task_id: str) -> a2a_pb2.Task: + from a2a.types import a2a_pb2 + + return await client.get_task(a2a_pb2.GetTaskRequest(id=task_id)) + + @pyrit_target_retry + @limit_requests_per_minute + async def _submit_message_async( + self, *, client: Client, request: a2a_pb2.SendMessageRequest + ) -> a2a_pb2.StreamResponse | None: + from a2a.utils.errors import A2AError + + stream = client.send_message(request) + try: + return await anext(stream, None) + except (A2AError, httpx.HTTPError) as exc: + if self._rate_limit_response(exc) is not None: + raise RateLimitException(message=f"A2A endpoint rate limited: {exc}") from exc + raise + finally: + if isinstance(stream, AsyncGenerator): + await stream.aclose() + + async def _send_prompt_to_target_async(self, *, normalized_conversation: list[Message]) -> list[Message]: + """ + Send the latest message to the agent, continuing the conversation's A2A context. + + Args: + normalized_conversation (list[Message]): Normalized conversation with the current request last. + + Returns: + list[Message]: A list containing the agent's response. + + Raises: + ValueError: If no conversation is provided or upstream context is missing for multi-turn history. + RateLimitException: If submission still returns HTTP 429 after retries. + TimeoutError: If task execution times out. + EmptyResponseException: If the agent returns an empty response without text. + """ + if not normalized_conversation: + raise ValueError("No conversation provided to A2ATarget.") + conversation_id = normalized_conversation[-1].get_piece().conversation_id + if not conversation_id: + raise ValueError("A2ATarget requires a conversation ID.") + lock = self._conversation_locks.setdefault(conversation_id, asyncio.Lock()) + async with lock: + return [ + await self._send_message_async( + normalized_conversation=normalized_conversation, conversation_id=conversation_id + ) + ] + + async def _send_message_async(self, *, normalized_conversation: list[Message], conversation_id: str) -> Message: + from a2a.types import a2a_pb2 + from a2a.utils.errors import A2AError + + message_piece = normalized_conversation[-1].get_piece() + state = self._conversations.get(conversation_id) + if (state is None or not state.context_id) and len(normalized_conversation) > 1: + raise ValueError( + f"A2ATarget has no upstream context for conversation '{conversation_id}'. " + "The target requires server-side context continuity for multi-turn conversations, " + "and earlier turns cannot be restored." + ) + state = self._conversations.setdefault(conversation_id, _A2AConversationState()) + prompt_text = message_piece.converted_value + + req = a2a_pb2.SendMessageRequest( + message=a2a_pb2.Message( + role=a2a_pb2.ROLE_USER, + parts=[a2a_pb2.Part(text=prompt_text)], + message_id=str(uuid.uuid4()), + context_id=state.context_id or "", + task_id=state.open_task_id or "", + ), + configuration=a2a_pb2.SendMessageConfiguration(return_immediately=False), + ) + + client_kwargs = dict(self._httpx_client_kwargs) + headers = httpx.Headers(self._build_headers()) + headers.update(client_kwargs.pop("headers", {})) + client_kwargs.setdefault("timeout", self._request_timeout_seconds) + if self._auth_token is not None: + client_kwargs["auth"] = _BearerTokenAuth(self._auth_token) + async with get_httpx_client(use_async=True, headers=headers, **client_kwargs) as http_client: + client = await self._create_a2a_client_async(http_client) + + try: + resp_event = await self._submit_message_async(client=client, request=req) + except (A2AError, httpx.HTTPError) as exc: + return self._error_response(request=message_piece, error_text=str(exc)) + + if resp_event is None: + raise EmptyResponseException(message="A2A agent returned an empty response stream.") + + if resp_event.HasField("message"): + msg = resp_event.message + if msg.context_id: + state.context_id = msg.context_id + state.open_task_id = None + reply_text = self._extract_message_text(msg) + if not reply_text: + raise EmptyResponseException(message=f"A2A message {msg.message_id} contained no text content.") + return construct_response_from_request(request=message_piece, response_text_pieces=[reply_text]) + + if resp_event.HasField("task"): + task = resp_event.task + if task.context_id: + state.context_id = task.context_id + task = await self._await_task_async(client=client, task=task) + return self._task_response(request=message_piece, task=task, state=state) + + raise EmptyResponseException(message="A2A response had neither message nor task.") + + def _error_response(self, *, request: MessagePiece, error_text: str) -> Message: + logger.warning("A2A agent at %s returned error: %s", self._endpoint, error_text) + return construct_response_from_request( + request=request, + response_text_pieces=[error_text], + response_type="error", + error="blocked" if any(marker in error_text for marker in CONTENT_FILTER_MARKERS) else "unknown", + ) + + def _task_response(self, *, request: MessagePiece, task: a2a_pb2.Task, state: _A2AConversationState) -> Message: + from a2a.types import a2a_pb2 + + if task.context_id: + state.context_id = task.context_id + status_text = self._extract_message_text(task.status.message) + state_name = a2a_pb2.TaskState.Name(task.status.state) + if task.status.state in ( + a2a_pb2.TASK_STATE_FAILED, + a2a_pb2.TASK_STATE_CANCELED, + a2a_pb2.TASK_STATE_REJECTED, + ): + state.open_task_id = None + return self._error_response( + request=request, error_text=status_text or f"A2A task {task.id} ended with state {state_name}" + ) + + artifacts_text = self._extract_task_artifacts_text(task) + if task.status.state in (a2a_pb2.TASK_STATE_INPUT_REQUIRED, a2a_pb2.TASK_STATE_AUTH_REQUIRED): + state.open_task_id = task.id + text = status_text or artifacts_text + empty_message = f"A2A task {task.id} is in {state_name} state but returned no text prompt." + elif task.status.state == a2a_pb2.TASK_STATE_COMPLETED: + state.open_task_id = None + text = artifacts_text or status_text + empty_message = f"A2A task {task.id} completed but returned no text response." + else: + raise ValueError(f"A2A task {task.id} returned unsupported state {state_name}.") + if not text: + raise EmptyResponseException(message=empty_message) + return construct_response_from_request(request=request, response_text_pieces=[text]) diff --git a/tests/integration/targets/test_a2a_target_integration.py b/tests/integration/targets/test_a2a_target_integration.py new file mode 100644 index 0000000000..caaabf9f54 --- /dev/null +++ b/tests/integration/targets/test_a2a_target_integration.py @@ -0,0 +1,327 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +import asyncio +import os +import socket +import uuid +from contextlib import AsyncExitStack +from typing import TYPE_CHECKING, Literal + +import httpx +import pytest +import uvicorn +from starlette.applications import Starlette +from starlette.middleware.base import BaseHTTPMiddleware +from starlette.responses import Response + +pytest.importorskip("a2a") + +from a2a.server.agent_execution import AgentExecutor +from a2a.server.request_handlers import DefaultRequestHandler +from a2a.server.routes.agent_card_routes import create_agent_card_routes +from a2a.server.routes.jsonrpc_routes import create_jsonrpc_routes +from a2a.server.tasks import InMemoryTaskStore +from a2a.types import a2a_pb2 as a2a +from a2a.utils.constants import TransportProtocol +from a2a.utils.errors import UnsupportedOperationError + +from pyrit.exceptions import EmptyResponseException +from pyrit.models import Conversation, Message, MessagePiece +from pyrit.prompt_target import A2ATarget + +if TYPE_CHECKING: + from collections.abc import AsyncIterator + + from a2a.server.agent_execution import RequestContext + from a2a.server.context import ServerCallContext + from a2a.server.events.event_queue_v2 import EventQueue + from starlette.middleware.base import RequestResponseEndpoint + from starlette.requests import Request + + from pyrit.memory import SQLiteMemory + +pytestmark = pytest.mark.run_only_if_all_tests + + +def _agent_message(*, text: str, context_id: str) -> a2a.Message: + return a2a.Message( + role=a2a.ROLE_AGENT, + message_id=str(uuid.uuid4()), + context_id=context_id, + parts=[a2a.Part(text=text)], + ) + + +class _TestExecutor(AgentExecutor): + def __init__(self) -> None: + self.contexts: dict[str, str] = {} + self.requests: list[a2a.Message] = [] + self.release_task = asyncio.Event() + self.task_finished = asyncio.Event() + + async def execute(self, context: RequestContext, event_queue: EventQueue) -> None: # pyrit-async-suffix-exempt + assert context.message is not None + assert context.context_id is not None + assert context.task_id is not None + message = a2a.Message() + message.CopyFrom(context.message) + self.requests.append(message) + text = context.get_user_input() + context_id = context.context_id + + if text in ("ask confirmation", "yes", "fail", "empty", "poll"): + task = a2a.Task( + id=context.task_id, + context_id=context_id, + status=a2a.TaskStatus(state=a2a.TASK_STATE_COMPLETED), + ) + if text == "ask confirmation": + task.status.state = a2a.TASK_STATE_INPUT_REQUIRED + task.status.message.CopyFrom(_agent_message(text="Are you sure?", context_id=context_id)) + task.artifacts.append(a2a.Artifact(artifact_id="partial", parts=[a2a.Part(text="partial preview")])) + elif text == "yes": + assert context.current_task is not None + assert context.current_task.status.state == a2a.TASK_STATE_INPUT_REQUIRED + task.artifacts.append(a2a.Artifact(artifact_id="done", parts=[a2a.Part(text="Confirmed.")])) + elif text == "fail": + task.status.state = a2a.TASK_STATE_FAILED + task.status.message.CopyFrom(_agent_message(text="Execution failed.", context_id=context_id)) + elif text == "poll": + task.status.state = a2a.TASK_STATE_WORKING + await event_queue.enqueue_event(task) + await self.release_task.wait() + await event_queue.enqueue_event( + a2a.TaskArtifactUpdateEvent( + task_id=context.task_id, + context_id=context_id, + artifact=a2a.Artifact(artifact_id="done", parts=[a2a.Part(text="Finished after polling.")]), + ) + ) + await event_queue.enqueue_event( + a2a.TaskStatusUpdateEvent( + task_id=context.task_id, + context_id=context_id, + status=a2a.TaskStatus(state=a2a.TASK_STATE_COMPLETED), + ) + ) + self.task_finished.set() + return + await event_queue.enqueue_event(task) + self.task_finished.set() + return + + if text.startswith("remember "): + self.contexts[context_id] = text.removeprefix("remember ") + reply = "Stored." + elif text == "recall": + reply = self.contexts[context_id] + else: + reply = f"Echo: {text}" + await event_queue.enqueue_event(_agent_message(text=reply, context_id=context_id)) + + async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None: # pyrit-async-suffix-exempt + raise UnsupportedOperationError("Cancellation is not used by this test agent.") + + +class _PollingHandler(DefaultRequestHandler): + async def on_message_send( # pyrit-async-suffix-exempt + self, params: a2a.SendMessageRequest, context: ServerCallContext + ) -> a2a.Task | a2a.Message: + if params.message.parts[0].text == "poll": + # Simulate an agent that returns a pending task even for a blocking request. + params.configuration.return_immediately = True + result = await super().on_message_send(params, context) + assert isinstance(result, (a2a.Task, a2a.Message)) + return result + + +_TestServer = tuple[str, _TestExecutor, dict[str, int]] + + +@pytest.fixture(params=["0.3", "1.0"]) +async def a2a_test_server( + request: pytest.FixtureRequest, +) -> AsyncIterator[_TestServer]: + executor = _TestExecutor() + counts = {"polls": 0, "cards": 0} + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + url = f"http://127.0.0.1:{sock.getsockname()[1]}" + card = a2a.AgentCard( + name="PyRIT integration agent", + description="Local deterministic SDK agent", + version="1.0", + supported_interfaces=[ + a2a.AgentInterface( + url=url + "/", + protocol_binding=TransportProtocol.JSONRPC, + protocol_version=request.param, + ) + ], + default_input_modes=["text/plain"], + default_output_modes=["text/plain"], + capabilities=a2a.AgentCapabilities(), + skills=[a2a.AgentSkill(id="test", name="Test", description="Test agent", tags=["test"])], + ) + handler = _PollingHandler(agent_executor=executor, task_store=InMemoryTaskStore(), agent_card=card) + + async def card_modifier_async(card: a2a.AgentCard) -> a2a.AgentCard: + counts["cards"] += 1 + return card + + app = Starlette( + routes=[ + *create_jsonrpc_routes(handler, "/", enable_v0_3_compat=True), + *create_agent_card_routes(card, card_modifier=card_modifier_async), + *create_agent_card_routes(card, card_url="/agentCard/custom", card_modifier=card_modifier_async), + ] + ) + + async def rate_limit_poll_async(request: Request, call_next: RequestResponseEndpoint) -> Response: + if request.method == "POST": + body = await request.json() + if body["method"] in ("tasks/get", "GetTask"): + counts["polls"] += 1 + if counts["polls"] == 1: + return Response("Poll rate limited", status_code=429) + executor.release_task.set() + await asyncio.wait_for(executor.task_finished.wait(), timeout=5) + return await call_next(request) + + app.add_middleware(BaseHTTPMiddleware, dispatch=rate_limit_poll_async) + server = uvicorn.Server(uvicorn.Config(app, log_level="error", lifespan="off", ws="none")) + serving = asyncio.create_task(server.serve(sockets=[sock])) + try: + async with asyncio.timeout(10): + while not server.started: + if serving.done(): + await serving + pytest.fail("SDK server stopped before startup.") + await asyncio.sleep(0.01) + async with httpx.AsyncClient() as client: + response = await client.get(url + "/.well-known/agent-card.json") + response.raise_for_status() + counts["cards"] = 0 + yield url, executor, counts + finally: + executor.release_task.set() + server.should_exit = True + await asyncio.wait_for(serving, timeout=10) + await handler.aclose() + + +def _make_msg(*, text: str, conversation_id: str) -> Message: + return MessagePiece(role="user", original_value=text, conversation_id=conversation_id).to_message() + + +def _persist_turn(*, memory: SQLiteMemory, target: A2ATarget, request: Message, response: Message) -> None: + conversation_id = request.get_piece().conversation_id + assert conversation_id is not None + memory.add_conversation_to_memory( + conversation=Conversation( + conversation_id=conversation_id, + target_identifier=target.get_identifier(), + ) + ) + memory.add_message_to_memory(request=request) + memory.add_message_to_memory(request=response) + + +@pytest.mark.parametrize("protocol_version", ["0.3", "1.0", "auto"]) +async def test_a2a_http_multi_turn_continuity( + sqlite_instance: SQLiteMemory, a2a_test_server: _TestServer, protocol_version: Literal["0.3", "1.0", "auto"] +) -> None: + url, executor, counts = a2a_test_server + target = A2ATarget(endpoint=url, protocol_version=protocol_version) + cid = str(uuid.uuid4()) + request = _make_msg(text="remember avocado-42", conversation_id=cid) + first = await target.send_prompt_async(message=request) + assert first[0].get_value() == "Stored." + _persist_turn(memory=sqlite_instance, target=target, request=request, response=first[0]) + assert len(sqlite_instance.get_conversation_messages(conversation_id=cid)) == 2 + second = await target.send_prompt_async(message=_make_msg(text="recall", conversation_id=cid)) + assert second[0].get_value() == "avocado-42" + assert executor.requests[0].context_id == executor.requests[1].context_id + assert counts["cards"] == (2 if protocol_version == "auto" else 0) + + +async def test_a2a_http_custom_card_discovery(sqlite_instance: SQLiteMemory, a2a_test_server: _TestServer) -> None: + url, executor, counts = a2a_test_server + target = A2ATarget(endpoint=url, protocol_version="auto", agent_card_path="/agentCard/custom") + result = await target.send_prompt_async(message=_make_msg(text="hello", conversation_id=str(uuid.uuid4()))) + assert result[0].get_value() == "Echo: hello" + assert counts["cards"] == 1 + + +async def test_a2a_http_input_required_continuation( + sqlite_instance: SQLiteMemory, a2a_test_server: _TestServer +) -> None: + url, executor, _ = a2a_test_server + target = A2ATarget(endpoint=url, protocol_version="auto") + cid = str(uuid.uuid4()) + request = _make_msg(text="ask confirmation", conversation_id=cid) + first = await target.send_prompt_async(message=request) + assert first[0].get_value() == "Are you sure?" + _persist_turn(memory=sqlite_instance, target=target, request=request, response=first[0]) + second = await target.send_prompt_async(message=_make_msg(text="yes", conversation_id=cid)) + assert second[0].get_value() == "Confirmed." + assert executor.requests[1].task_id == executor.requests[0].task_id + assert executor.requests[1].context_id == executor.requests[0].context_id + assert target._conversations[cid].open_task_id is None + + +async def test_a2a_http_failed_task(sqlite_instance: SQLiteMemory, a2a_test_server: _TestServer) -> None: + url, _, _ = a2a_test_server + target = A2ATarget(endpoint=url, protocol_version="auto") + responses = await target.send_prompt_async(message=_make_msg(text="fail", conversation_id=str(uuid.uuid4()))) + piece = responses[0].get_piece() + assert piece.converted_value_data_type == "error" + assert piece.response_error == "unknown" + assert piece.converted_value == "Execution failed." + + +async def test_a2a_http_polling_rate_limit_no_resubmission( + sqlite_instance: SQLiteMemory, a2a_test_server: _TestServer +) -> None: + url, executor, counts = a2a_test_server + target = A2ATarget(endpoint=url, protocol_version="auto", poll_interval_seconds=0.01, task_timeout_seconds=5) + responses = await target.send_prompt_async(message=_make_msg(text="poll", conversation_id=str(uuid.uuid4()))) + assert responses[0].get_value() == "Finished after polling." + assert counts["polls"] >= 2 + assert len(executor.requests) == 1 + + +async def test_a2a_http_empty_task_no_resubmission(sqlite_instance: SQLiteMemory, a2a_test_server: _TestServer) -> None: + url, executor, _ = a2a_test_server + target = A2ATarget(endpoint=url, protocol_version="auto") + with pytest.raises(EmptyResponseException): + await target.send_prompt_async(message=_make_msg(text="empty", conversation_id=str(uuid.uuid4()))) + assert len(executor.requests) == 1 + + +@pytest.mark.skipif(not os.getenv("A2A_FOUNDRY_ENDPOINT"), reason="A2A_FOUNDRY_ENDPOINT is not set") +async def test_a2a_foundry_live_integration(sqlite_instance: SQLiteMemory) -> None: + from azure.identity.aio import DefaultAzureCredential, get_bearer_token_provider + + async with AsyncExitStack() as stack: + token = os.environ.get("A2A_FOUNDRY_AUTH_TOKEN") + provider = None + if not token: + credential = await stack.enter_async_context(DefaultAzureCredential()) + provider = get_bearer_token_provider(credential, "https://ai.azure.com/.default") + target = A2ATarget( + endpoint=os.environ["A2A_FOUNDRY_ENDPOINT"], + auth_token=token or provider, + protocol_version="auto", + agent_card_path=os.getenv("A2A_FOUNDRY_CARD_PATH", "agentCard/v1.0"), + ) + responses = await target.send_prompt_async( + message=_make_msg(text="Reply with a short greeting.", conversation_id=str(uuid.uuid4())) + ) + piece = responses[0].get_piece() + assert piece.response_error == "none" + assert piece.converted_value_data_type == "text" + assert piece.converted_value.strip() diff --git a/tests/unit/prompt_target/target/test_a2a_target.py b/tests/unit/prompt_target/target/test_a2a_target.py new file mode 100644 index 0000000000..2369a0a7b2 --- /dev/null +++ b/tests/unit/prompt_target/target/test_a2a_target.py @@ -0,0 +1,855 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import asyncio +import json +from collections.abc import AsyncIterator, Iterator +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +from pyrit.exceptions import EmptyResponseException +from pyrit.models import Message, MessagePiece +from pyrit.prompt_target import A2ATarget +from pyrit.prompt_target.a2a_target import _A2AConversationState +from pyrit.prompt_target.common.target_capabilities import TargetCapabilities +from pyrit.prompt_target.common.target_configuration import TargetConfiguration + +ENDPOINT = "https://agent.example.com/a2a" + + +def _user_message(text: str = "Hello A2A Agent", conversation_id: str = "conv-1234") -> Message: + piece = MessagePiece( + role="user", + original_value=text, + converted_value=text, + conversation_id=conversation_id, + ) + return Message(message_pieces=[piece]) + + +def _rpc_response(payload: dict[str, Any], status_code: int = 200) -> httpx.Response: + req = httpx.Request("POST", ENDPOINT) + return httpx.Response(status_code=status_code, json=payload, request=req) + + +def _task_payload( + text: str | None, + *, + state: str = "completed", + task_id: str = "task-1", + context_id: str = "ctx-1", + status_message: str | None = None, +) -> dict[str, Any]: + status_obj: dict[str, Any] = {"state": state} + if status_message is not None: + status_obj["message"] = { + "role": "agent", + "messageId": "msg-status-1", + "parts": [{"kind": "text", "text": status_message}], + } + result: dict[str, Any] = { + "kind": "task", + "id": task_id, + "contextId": context_id, + "status": status_obj, + } + if text is not None: + result["artifacts"] = [ + { + "artifactId": "art-1", + "parts": [{"kind": "text", "text": text}], + } + ] + else: + result["artifacts"] = [] + return {"jsonrpc": "2.0", "id": "1", "result": result} + + +def _message_payload(text: str, *, message_id: str = "m-1", context_id: str = "ctx-1") -> dict[str, Any]: + return { + "jsonrpc": "2.0", + "id": "1", + "result": { + "kind": "message", + "role": "agent", + "messageId": message_id, + "contextId": context_id, + "parts": [{"kind": "text", "text": text}], + }, + } + + +@pytest.fixture(autouse=True) +def a2a_dependencies(request: pytest.FixtureRequest) -> Iterator[None]: + if request.node.name != "test_a2a_target_missing_sdk_raises_import_error": + pytest.importorskip("a2a") + with patch.dict("os.environ", {"RETRY_WAIT_MIN_SECONDS": "0", "RETRY_WAIT_MAX_SECONDS": "0"}): + yield + + +@pytest.mark.usefixtures("patch_central_database") +def test_a2a_target_initialization(): + target = A2ATarget(endpoint=ENDPOINT + "/", auth_token="test-token", api_key="test-api-key") + assert target._endpoint == ENDPOINT + assert target._protocol_version == "0.3" + assert "Authorization" not in target._build_headers() + assert target._build_headers()["X-API-Key"] == "test-api-key" + assert target._build_headers()["Accept"] == "application/json" + + +@pytest.mark.usefixtures("patch_central_database") +def test_a2a_target_identifier(): + target = A2ATarget(endpoint=ENDPOINT, auth_token="secret-token", api_key="secret-key") + identifier = target.get_identifier() + assert identifier.params["endpoint"] == ENDPOINT + assert identifier.params["protocol_version"] == "0.3" + assert "task_timeout_seconds" not in identifier.params + assert "secret-token" not in str(identifier.params) + assert "secret-key" not in str(identifier.params) + + +@pytest.mark.usefixtures("patch_central_database") +def test_a2a_target_rejects_unknown_protocol_version(): + with pytest.raises(ValueError, match="Unsupported A2A protocol_version"): + A2ATarget(endpoint=ENDPOINT, protocol_version="v9") # type: ignore[arg-type] + + +@pytest.mark.usefixtures("patch_central_database") +def test_a2a_target_rejects_unsupported_capabilities(): + invalid_config = TargetConfiguration( + capabilities=TargetCapabilities( + supports_multi_turn=True, + input_modalities=frozenset({frozenset(["image_path"])}), + ) + ) + with pytest.raises(ValueError, match="only supports text input modality"): + A2ATarget(endpoint=ENDPOINT, custom_configuration=invalid_config) + + system_prompt_config = TargetConfiguration( + capabilities=TargetCapabilities( + supports_multi_turn=True, + supports_system_prompt=True, + ) + ) + with pytest.raises(ValueError, match="does not support supports_system_prompt"): + A2ATarget(endpoint=ENDPOINT, custom_configuration=system_prompt_config) + + +@pytest.mark.usefixtures("patch_central_database") +def test_a2a_target_missing_sdk_raises_import_error() -> None: + real_import = __import__ + + def mock_import(name, *args, **kwargs): + if name == "a2a" or name.startswith("a2a."): + raise ImportError("No module named 'a2a'") + return real_import(name, *args, **kwargs) + + with ( + patch("builtins.__import__", side_effect=mock_import), + pytest.raises(ImportError, match="pip install pyrit\\[a2a\\]"), + ): + A2ATarget(endpoint=ENDPOINT) + + +@pytest.mark.usefixtures("patch_central_database") +async def test_a2a_target_missing_upstream_context_raises(): + target = A2ATarget(endpoint=ENDPOINT) + msg1 = Message(message_pieces=[MessagePiece(role="user", original_value="turn 1", conversation_id="conv-1")]) + msg2 = Message(message_pieces=[MessagePiece(role="user", original_value="turn 2", conversation_id="conv-1")]) + with pytest.raises(ValueError, match="has no upstream context for conversation"): + await target._send_prompt_to_target_async(normalized_conversation=[msg1, msg2]) + + +@pytest.mark.usefixtures("patch_central_database") +async def test_a2a_target_set_and_reset_conversation_context(): + target = A2ATarget(endpoint=ENDPOINT) + target.set_conversation_context(conversation_id="conv-1", context_id="ctx-restored") + assert target._conversations["conv-1"].context_id == "ctx-restored" + + await target.reset_conversation_async(conversation_id="conv-1") + assert "conv-1" not in target._conversations + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_send_prompt_message_result(mock_send): + mock_send.return_value = _rpc_response(_message_payload("Direct reply", context_id="ctx-abc")) + + target = A2ATarget(endpoint=ENDPOINT) + responses = await target.send_prompt_async(message=_user_message()) + + assert responses[0].message_pieces[0].converted_value == "Direct reply" + assert target._conversations["conv-1234"].context_id == "ctx-abc" + + sent_req = mock_send.call_args[0][0] + payload = json.loads(sent_req.content) + assert payload["method"] == "message/send" + assert payload["params"]["message"]["parts"][0]["text"] == "Hello A2A Agent" + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_send_prompt_task_completed(mock_send): + mock_send.return_value = _rpc_response(_task_payload("Task answer", state="completed", context_id="ctx-task")) + + target = A2ATarget(endpoint=ENDPOINT) + responses = await target.send_prompt_async(message=_user_message()) + + assert responses[0].message_pieces[0].converted_value == "Task answer" + assert target._conversations["conv-1234"].context_id == "ctx-task" + + +@pytest.mark.usefixtures("patch_central_database") +@patch("asyncio.sleep", new_callable=AsyncMock) +@patch("httpx.AsyncClient.send") +async def test_a2a_polls_pending_task(mock_send, mock_sleep): + mock_send.side_effect = [ + _rpc_response(_task_payload(None, state="working", task_id="task-42")), + _rpc_response(_task_payload("Polled reply", state="completed", task_id="task-42")), + ] + + target = A2ATarget(endpoint=ENDPOINT) + responses = await target.send_prompt_async(message=_user_message()) + + assert responses[0].message_pieces[0].converted_value == "Polled reply" + assert mock_send.call_count == 2 + second_req = mock_send.call_args_list[1][0][0] + second_payload = json.loads(second_req.content) + assert second_payload["method"] == "tasks/get" + assert second_payload["params"]["id"] == "task-42" + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_pending_task_times_out(mock_send): + mock_send.return_value = _rpc_response(_task_payload(None, state="working", task_id="task-stuck")) + + target = A2ATarget(endpoint=ENDPOINT, task_timeout_seconds=0.01) + with pytest.raises(TimeoutError, match="still in state TASK_STATE_WORKING"): + await target.send_prompt_async(message=_user_message()) + assert mock_send.call_count == 1 + assert target._conversations["conv-1234"].context_id == "ctx-1" + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_input_required_task_extracts_question_from_status_message(mock_send): + mock_send.return_value = _rpc_response( + _task_payload( + text="partial artifact", + state="input-required", + task_id="task-ask", + context_id="ctx-ask", + status_message="What is your confirmation code?", + ) + ) + + target = A2ATarget(endpoint=ENDPOINT) + responses = await target.send_prompt_async(message=_user_message()) + + # Prioritizes status message question over partial artifacts + assert responses[0].message_pieces[0].converted_value == "What is your confirmation code?" + assert target._conversations["conv-1234"].open_task_id == "task-ask" + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_input_required_continuation(mock_send): + mock_send.side_effect = [ + _rpc_response( + _task_payload( + None, + state="input-required", + task_id="task-open", + context_id="ctx-open", + status_message="Confirm?", + ) + ), + _rpc_response(_task_payload("Confirmed!", state="completed", task_id="task-open", context_id="ctx-open")), + ] + + target = A2ATarget(endpoint=ENDPOINT) + r1 = await target.send_prompt_async(message=_user_message("start", conversation_id="c1")) + assert r1[0].message_pieces[0].converted_value == "Confirm?" + + r2 = await target.send_prompt_async(message=_user_message("yes", conversation_id="c1")) + assert r2[0].message_pieces[0].converted_value == "Confirmed!" + + second_req = mock_send.call_args_list[1][0][0] + payload = json.loads(second_req.content) + assert payload["params"]["message"]["taskId"] == "task-open" + assert payload["params"]["message"]["contextId"] == "ctx-open" + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_task_failed_state_returns_error_response(mock_send): + mock_send.return_value = _rpc_response( + _task_payload(None, state="failed", task_id="task-fail", status_message="Quota exceeded") + ) + + target = A2ATarget(endpoint=ENDPOINT) + responses = await target.send_prompt_async(message=_user_message()) + + piece = responses[0].message_pieces[0] + assert piece.converted_value_data_type == "error" + assert piece.response_error == "unknown" + assert "Quota exceeded" in piece.converted_value + assert target._conversations["conv-1234"].open_task_id is None + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_task_canceled_state_returns_error_response(mock_send): + mock_send.return_value = _rpc_response(_task_payload(None, state="canceled", task_id="task-cancel")) + + target = A2ATarget(endpoint=ENDPOINT) + responses = await target.send_prompt_async(message=_user_message()) + + piece = responses[0].message_pieces[0] + assert piece.converted_value_data_type == "error" + assert "CANCELED" in piece.converted_value + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_content_filter_error_is_blocked(mock_send): + mock_send.return_value = _rpc_response( + { + "jsonrpc": "2.0", + "id": "1", + "error": {"code": -32000, "message": "content_filter: prompt was blocked"}, + } + ) + + target = A2ATarget(endpoint=ENDPOINT) + responses = await target.send_prompt_async(message=_user_message()) + + piece = responses[0].message_pieces[0] + assert piece.converted_value_data_type == "error" + assert piece.response_error == "blocked" + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_initial_send_rate_limit_retried(mock_send): + mock_send.side_effect = [ + httpx.Response(429, request=httpx.Request("POST", ENDPOINT)), + _rpc_response(_task_payload("Success after retry")), + ] + + target = A2ATarget(endpoint=ENDPOINT) + responses = await target.send_prompt_async(message=_user_message()) + + assert responses[0].message_pieces[0].converted_value == "Success after retry" + assert mock_send.call_count == 2 + first_message = json.loads(mock_send.call_args_list[0][0][0].content)["params"]["message"] + second_message = json.loads(mock_send.call_args_list[1][0][0].content)["params"]["message"] + assert first_message == second_message + + +@pytest.mark.usefixtures("patch_central_database") +@patch("asyncio.sleep", new_callable=AsyncMock) +@patch("httpx.AsyncClient.send") +async def test_a2a_polling_rate_limit_does_not_resubmit_prompt(mock_send, mock_sleep): + mock_send.side_effect = [ + _rpc_response(_task_payload(None, state="working", task_id="t-safe")), + httpx.Response(429, request=httpx.Request("POST", ENDPOINT)), # poll fails with 429 + _rpc_response(_task_payload("Poll success", state="completed", task_id="t-safe")), + ] + + target = A2ATarget(endpoint=ENDPOINT) + responses = await target.send_prompt_async(message=_user_message()) + + assert responses[0].message_pieces[0].converted_value == "Poll success" + assert mock_send.call_count == 3 + # First call: message/send + assert json.loads(mock_send.call_args_list[0][0][0].content)["method"] == "message/send" + # Second call: tasks/get (got 429) + assert json.loads(mock_send.call_args_list[1][0][0].content)["method"] == "tasks/get" + # Third call: tasks/get retry (did NOT call message/send again!) + assert json.loads(mock_send.call_args_list[2][0][0].content)["method"] == "tasks/get" + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_empty_response_raises(mock_send): + mock_send.return_value = _rpc_response(_task_payload(None, state="completed", task_id="t-empty")) + + target = A2ATarget(endpoint=ENDPOINT) + with pytest.raises(EmptyResponseException, match="completed but returned no text response"): + await target.send_prompt_async(message=_user_message()) + assert mock_send.call_count == 1 + + +@pytest.mark.usefixtures("patch_central_database") +async def test_a2a_auth_headers_sent() -> None: + requests: list[httpx.Request] = [] + + async def handle_async(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(200, json=_message_payload("Authed")) + + target = A2ATarget( + endpoint=ENDPOINT, + auth_token="secret-bearer-token", + api_key="custom-key", + transport=httpx.MockTransport(handle_async), + ) + await target.send_prompt_async(message=_user_message()) + assert len(requests) == 1 + assert requests[0].headers["Authorization"] == "Bearer secret-bearer-token" + assert requests[0].headers["X-API-Key"] == "custom-key" + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + "flag", + [ + "supports_multi_message_pieces", + "supports_editable_history", + "supports_system_prompt", + "supports_json_output", + "supports_json_schema", + "supports_streaming_audio", + ], +) +def test_a2a_capability_updates_reject_unsupported_flags(flag: str) -> None: + capabilities = TargetCapabilities.model_validate({"supports_multi_turn": True, flag: True}) + with pytest.raises(ValueError, match=flag): + A2ATarget(endpoint=ENDPOINT, custom_configuration=TargetConfiguration(capabilities=capabilities)) + + target = A2ATarget(endpoint=ENDPOINT) + original = target.capabilities + with pytest.raises(ValueError, match=flag): + target.apply_capabilities(capabilities=capabilities) + assert target.capabilities == original + + +@pytest.mark.usefixtures("patch_central_database") +def test_a2a_identity_excludes_operational_settings() -> None: + target = A2ATarget(endpoint=ENDPOINT, routing_identifier="agent-one") + same = A2ATarget( + endpoint=ENDPOINT, + routing_identifier="agent-one", + task_timeout_seconds=40, + request_timeout_seconds=10, + poll_interval_seconds=0.1, + auth_token="different-secret", + headers={"X-Agent": "one"}, + ) + assert target.get_identifier().hash == same.get_identifier().hash + for kwargs in ( + {"routing_identifier": "agent-two"}, + {"routing_identifier": "agent-one", "protocol_version": "1.0"}, + {"routing_identifier": "agent-one", "agent_card_path": "agentCard/v1.0"}, + ): + assert target.get_identifier().hash != A2ATarget(endpoint=ENDPOINT, **kwargs).get_identifier().hash + + +@pytest.mark.usefixtures("patch_central_database") +async def test_a2a_context_entry_without_context_id_does_not_restore_history() -> None: + target = A2ATarget(endpoint=ENDPOINT) + target._conversations["conv-1234"] = _A2AConversationState() + with pytest.raises(ValueError, match="has no upstream context"): + await target._send_prompt_to_target_async(normalized_conversation=[_user_message(), _user_message()]) + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_multiple_pieces_never_silently_discarded(mock_send: AsyncMock) -> None: + target = A2ATarget(endpoint=ENDPOINT) + message = Message(message_pieces=[_user_message("one").get_piece(), _user_message("two").get_piece()]) + with pytest.raises(ValueError, match="single message piece"): + await target.send_prompt_async(message=message) + mock_send.assert_not_called() + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + "payload", + [ + _message_payload(""), + _task_payload(None, state="completed"), + _task_payload(None, state="input-required"), + _task_payload(None, state="auth-required"), + ], +) +@patch("httpx.AsyncClient.send") +async def test_a2a_accepted_empty_result_never_resubmitted(mock_send: AsyncMock, payload: dict[str, Any]) -> None: + mock_send.return_value = _rpc_response(payload) + target = A2ATarget(endpoint=ENDPOINT) + with pytest.raises(EmptyResponseException): + await target.send_prompt_async(message=_user_message()) + mock_send.assert_called_once() + assert target._conversations["conv-1234"].context_id == "ctx-1" + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_relayed_rate_limit_is_not_safe_to_resubmit(mock_send: AsyncMock) -> None: + mock_send.return_value = _rpc_response( + {"jsonrpc": "2.0", "id": "1", "error": {"code": -32603, "message": "Model 429: rate limit"}} + ) + response = await A2ATarget(endpoint=ENDPOINT).send_prompt_async(message=_user_message()) + assert response[0].get_piece().response_error == "unknown" + mock_send.assert_called_once() + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_request_timeout_is_explicit(mock_send: AsyncMock) -> None: + mock_send.return_value = _rpc_response(_message_payload("ok")) + await A2ATarget(endpoint=ENDPOINT, request_timeout_seconds=35).send_prompt_async(message=_user_message()) + assert mock_send.call_args[0][0].extensions["timeout"]["read"] == 35 + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_discovery_error_does_not_fall_back_to_submission(mock_send: AsyncMock) -> None: + from a2a.client.errors import AgentCardResolutionError + + mock_send.return_value = httpx.Response(404, request=httpx.Request("GET", ENDPOINT)) + with pytest.raises(AgentCardResolutionError): + await A2ATarget(endpoint=ENDPOINT, protocol_version="auto").send_prompt_async(message=_user_message()) + assert all(call.args[0].method == "GET" for call in mock_send.call_args_list) + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_non_text_output_never_resubmitted(mock_send: AsyncMock) -> None: + payload = _message_payload("ignored") + payload["result"]["parts"] = [{"kind": "data", "data": {"answer": 42}}] + mock_send.return_value = _rpc_response(payload) + with pytest.raises(EmptyResponseException, match="no text"): + await A2ATarget(endpoint=ENDPOINT).send_prompt_async(message=_user_message()) + mock_send.assert_called_once() + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_programming_errors_propagate_without_retry(mock_send: AsyncMock) -> None: + mock_send.side_effect = TypeError("unexpected client bug") + with pytest.raises(TypeError, match="unexpected client bug"): + await A2ATarget(endpoint=ENDPOINT).send_prompt_async(message=_user_message()) + mock_send.assert_called_once() + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_poll_error_never_resubmits(mock_send: AsyncMock) -> None: + from a2a.client.errors import A2AClientError + + mock_send.side_effect = [ + _rpc_response(_task_payload(None, state="working")), + httpx.Response(503, request=httpx.Request("POST", ENDPOINT)), + ] + with pytest.raises(A2AClientError, match="503"): + await A2ATarget(endpoint=ENDPOINT, poll_interval_seconds=0).send_prompt_async(message=_user_message()) + assert mock_send.call_count == 2 + assert json.loads(mock_send.call_args_list[1].args[0].content)["method"] == "tasks/get" + + +@pytest.mark.usefixtures("patch_central_database") +async def test_a2a_empty_sdk_iterator_never_resubmitted() -> None: + from a2a.client import Client + + client = MagicMock(spec=Client) + iterator = MagicMock(spec=AsyncIterator) + iterator.__anext__ = AsyncMock(side_effect=StopAsyncIteration) + client.send_message.return_value = iterator + target = A2ATarget(endpoint=ENDPOINT) + with ( + patch.object(target, "_create_a2a_client_async", return_value=client), + pytest.raises(EmptyResponseException, match="empty response stream"), + ): + await target.send_prompt_async(message=_user_message()) + client.send_message.assert_called_once() + + +@pytest.mark.usefixtures("patch_central_database") +def test_a2a_rejects_single_turn_override() -> None: + capabilities = TargetCapabilities(supports_multi_turn=False) + with pytest.raises(ValueError, match="requires supports_multi_turn=True"): + A2ATarget(endpoint=ENDPOINT, custom_configuration=TargetConfiguration(capabilities=capabilities)) + target = A2ATarget(endpoint=ENDPOINT) + with pytest.raises(ValueError, match="requires supports_multi_turn=True"): + target.apply_capabilities(capabilities=capabilities) + assert target.capabilities.supports_multi_turn + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("name", ["task_timeout_seconds", "request_timeout_seconds", "poll_interval_seconds"]) +@pytest.mark.parametrize("value", [-1, float("inf"), float("-inf"), float("nan")]) +def test_a2a_rejects_invalid_waits(*, name: str, value: float) -> None: + with pytest.raises(ValueError, match=name): + A2ATarget(endpoint=ENDPOINT, **{name: value}) + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("name", ["task_timeout_seconds", "request_timeout_seconds"]) +def test_a2a_rejects_zero_timeout(name: str) -> None: + with pytest.raises(ValueError, match=name): + A2ATarget(endpoint=ENDPOINT, **{name: 0}) + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("timeout", [None, 12.0, httpx.Timeout(10, read=30)]) +@patch("httpx.AsyncClient.send") +async def test_a2a_httpx_timeout_is_preserved(mock_send: AsyncMock, timeout: httpx.Timeout | float | None) -> None: + mock_send.return_value = _rpc_response(_message_payload("ok")) + target = A2ATarget(endpoint=ENDPOINT, timeout=timeout) + await target.send_prompt_async(message=_user_message()) + assert mock_send.call_args.args[0].extensions["timeout"] == httpx.Timeout(timeout).as_dict() + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("timeout", [-1, 0, float("inf"), float("nan"), httpx.Timeout(10, read=-1)]) +def test_a2a_rejects_invalid_httpx_timeout(timeout: httpx.Timeout | float) -> None: + with pytest.raises(ValueError, match="HTTPX timeouts"): + A2ATarget(endpoint=ENDPOINT, timeout=timeout) + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize( + "kwargs", + [ + {"request_timeout_seconds": 10, "timeout": None}, + {"auth_token": "token", "auth": httpx.BasicAuth("user", "password")}, + {"auth_token": "token", "headers": {"authorization": "Bearer other"}}, + {"auth_token": "token", "api_key": "key", "api_key_header": "Authorization"}, + ], +) +def test_a2a_rejects_conflicting_options(kwargs: dict[str, Any]) -> None: + with pytest.raises(ValueError, match="Specify either"): + A2ATarget(endpoint=ENDPOINT, **kwargs) + + +@pytest.mark.usefixtures("patch_central_database") +async def test_a2a_provider_refreshes_discovery_submission_retry_and_polls() -> None: + requests: list[httpx.Request] = [] + provider = AsyncMock(side_effect=["card-token", "send-token", "retry-token", "poll-token"]) + + async def handle_async(request: httpx.Request) -> httpx.Response: + requests.append(request) + if request.method == "GET": + return httpx.Response( + 200, + json={ + "name": "Test", + "description": "Test", + "version": "1", + "protocolVersion": "0.3.0", + "url": ENDPOINT, + "capabilities": {}, + "defaultInputModes": ["text/plain"], + "defaultOutputModes": ["text/plain"], + "skills": [], + }, + ) + if len(requests) == 2: + return httpx.Response(429) + return httpx.Response( + 200, json=(_task_payload(None, state="working") if len(requests) == 3 else _task_payload("done")) + ) + + target = A2ATarget( + endpoint=ENDPOINT, + protocol_version="auto", + auth_token=provider, + transport=httpx.MockTransport(handle_async), + poll_interval_seconds=0, + ) + result = await target.send_prompt_async(message=_user_message()) + assert result[0].get_value() == "done" + assert provider.await_count == 4 + assert [request.headers["Authorization"] for request in requests] == [ + "Bearer card-token", + "Bearer send-token", + "Bearer retry-token", + "Bearer poll-token", + ] + first = json.loads(requests[1].content)["params"]["message"] + retry = json.loads(requests[2].content)["params"]["message"] + assert first["messageId"] == retry["messageId"] + assert "token" not in str(target.get_identifier().params) + + +@pytest.mark.usefixtures("patch_central_database") +async def test_a2a_provider_failure_is_not_retried() -> None: + provider = AsyncMock(side_effect=ValueError("credential failed")) + handler = AsyncMock() + target = A2ATarget(endpoint=ENDPOINT, auth_token=provider, transport=httpx.MockTransport(handler)) + with pytest.raises(ValueError, match="credential failed"): + await target.send_prompt_async(message=_user_message()) + provider.assert_awaited_once() + handler.assert_not_called() + + +@pytest.mark.usefixtures("patch_central_database") +async def test_a2a_empty_provider_token_fails_before_http() -> None: + provider = AsyncMock(return_value=" ") + handler = AsyncMock() + target = A2ATarget(endpoint=ENDPOINT, auth_token=provider, transport=httpx.MockTransport(handler)) + with pytest.raises(ValueError, match="non-empty token"): + await target.send_prompt_async(message=_user_message()) + provider.assert_awaited_once() + handler.assert_not_called() + + +@pytest.mark.usefixtures("patch_central_database") +async def test_a2a_custom_httpx_auth_is_supported() -> None: + auth = httpx.BasicAuth("user", "password") + requests: list[httpx.Request] = [] + + async def handle_async(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(200, json=_message_payload("ok")) + + target = A2ATarget(endpoint=ENDPOINT, auth=auth, transport=httpx.MockTransport(handle_async)) + await target.send_prompt_async(message=_user_message()) + assert requests[0].headers["Authorization"].startswith("Basic ") + + +@pytest.mark.usefixtures("patch_central_database") +@pytest.mark.parametrize("action", ["send", "reset", "cancel"]) +async def test_a2a_serializes_conversation_state(action: str) -> None: + started = asyncio.Event() + release = asyncio.Event() + requests: list[dict[str, Any]] = [] + + async def handle_async(request: httpx.Request) -> httpx.Response: + requests.append(json.loads(request.content)["params"]["message"]) + if len(requests) == 1: + started.set() + await release.wait() + return httpx.Response(200, json=_message_payload("ok", context_id="ctx-locked")) + + target = A2ATarget(endpoint=ENDPOINT, transport=httpx.MockTransport(handle_async)) + async with asyncio.timeout(5), asyncio.TaskGroup() as tasks: + first = tasks.create_task(target.send_prompt_async(message=_user_message("first"))) + await started.wait() + with pytest.raises(RuntimeError, match="conversation is in use"): + target.set_conversation_context(conversation_id="conv-1234", context_id="replacement") + if action == "reset": + second = tasks.create_task(target.reset_conversation_async(conversation_id="conv-1234")) + else: + second = tasks.create_task(target.send_prompt_async(message=_user_message("second"))) + await asyncio.sleep(0) + assert not second.done() + assert len(requests) == 1 + if action == "cancel": + first.cancel() + release.set() + if action == "reset": + assert "conv-1234" not in target._conversations + assert len(requests) == 1 + else: + assert len(requests) == 2 + if action == "send": + assert requests[1]["contextId"] == "ctx-locked" + assert all(not lock.locked() for lock in target._conversation_locks.values()) + + +@pytest.mark.usefixtures("patch_central_database") +async def test_a2a_different_conversations_remain_concurrent() -> None: + both_started = asyncio.Event() + requests: list[httpx.Request] = [] + + async def handle_async(request: httpx.Request) -> httpx.Response: + requests.append(request) + if len(requests) == 2: + both_started.set() + await both_started.wait() + return httpx.Response(200, json=_message_payload("ok")) + + target = A2ATarget(endpoint=ENDPOINT, transport=httpx.MockTransport(handle_async)) + async with asyncio.timeout(5): + await asyncio.gather( + target.send_prompt_async(message=_user_message(conversation_id="one")), + target.send_prompt_async(message=_user_message(conversation_id="two")), + ) + assert len(requests) == 2 + + +@pytest.mark.usefixtures("patch_central_database") +@patch("asyncio.sleep", new_callable=AsyncMock) +@patch("httpx.AsyncClient.send") +async def test_a2a_poll_respects_retry_after(mock_send: AsyncMock, mock_sleep: AsyncMock) -> None: + mock_send.side_effect = [ + _rpc_response(_task_payload(None, state="working")), + httpx.Response(429, headers={"Retry-After": "7"}, request=httpx.Request("POST", ENDPOINT)), + _rpc_response(_task_payload(None, state="working")), + _rpc_response(_task_payload("done")), + ] + await A2ATarget(endpoint=ENDPOINT).send_prompt_async(message=_user_message()) + assert [call.args[0] for call in mock_sleep.await_args_list] == [1, 7, 1] + assert [json.loads(call.args[0].content)["method"] for call in mock_send.call_args_list] == [ + "message/send", + "tasks/get", + "tasks/get", + "tasks/get", + ] + + +@pytest.mark.parametrize( + ("header", "expected"), + [("5", 5), ("Wed, 21 Oct 2015 07:28:00 GMT", 5), ("-1", 0), ("invalid", None), ("inf", None)], +) +def test_a2a_retry_after_parsing(*, header: str, expected: float | None) -> None: + with patch("pyrit.prompt_target.a2a_target.time.time", return_value=1445412475): + assert A2ATarget._retry_after_seconds(httpx.Response(429, headers={"Retry-After": header})) == expected + + +@pytest.mark.usefixtures("patch_central_database") +@patch("asyncio.sleep", new_callable=AsyncMock) +@patch("httpx.AsyncClient.send") +async def test_a2a_submission_and_poll_are_rate_limited(mock_send: AsyncMock, mock_sleep: AsyncMock) -> None: + mock_send.side_effect = [ + _rpc_response(_task_payload(None, state="working")), + _rpc_response(_task_payload("done")), + ] + target = A2ATarget(endpoint=ENDPOINT, max_requests_per_minute=60, poll_interval_seconds=0) + await target.send_prompt_async(message=_user_message()) + assert [call.args[0] for call in mock_sleep.await_args_list] == [1, 0, 1] + + +@pytest.mark.usefixtures("patch_central_database") +async def test_a2a_poll_deadline_cancels_inflight_http_request() -> None: + requests: list[httpx.Request] = [] + + async def handle_async(request: httpx.Request) -> httpx.Response: + requests.append(request) + if len(requests) == 1: + return httpx.Response(200, json=_task_payload(None, state="working")) + await asyncio.Event().wait() + raise AssertionError("The poll must be cancelled.") + + target = A2ATarget( + endpoint=ENDPOINT, + task_timeout_seconds=0.02, + poll_interval_seconds=0, + timeout=None, + transport=httpx.MockTransport(handle_async), + ) + with pytest.raises(TimeoutError, match="TASK_STATE_WORKING"): + await target.send_prompt_async(message=_user_message()) + assert len(requests) == 2 + assert all(not lock.locked() for lock in target._conversation_locks.values()) + + +@pytest.mark.usefixtures("patch_central_database") +@patch("httpx.AsyncClient.send") +async def test_a2a_poll_deadline_bounds_retry_after(mock_send: AsyncMock) -> None: + mock_send.side_effect = [ + _rpc_response(_task_payload(None, state="working")), + httpx.Response(429, headers={"Retry-After": "3600"}, request=httpx.Request("POST", ENDPOINT)), + ] + target = A2ATarget(endpoint=ENDPOINT, task_timeout_seconds=0.02, poll_interval_seconds=0) + async with asyncio.timeout(5): + with pytest.raises(TimeoutError, match="TASK_STATE_WORKING"): + await target.send_prompt_async(message=_user_message()) + assert mock_send.call_count == 2 diff --git a/uv.lock b/uv.lock index 6f8a57063f..f34f41160e 100644 --- a/uv.lock +++ b/uv.lock @@ -38,6 +38,25 @@ constraints = [ { name = "werkzeug", specifier = ">=3.1.6" }, ] +[[package]] +name = "a2a-sdk" +version = "1.1.5" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "culsans", marker = "python_full_version < '3.13'" }, + { name = "google-api-core" }, + { name = "googleapis-common-protos" }, + { name = "httpx" }, + { name = "json-rpc" }, + { name = "packaging" }, + { name = "protobuf" }, + { name = "pydantic" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/75/b3/db19deba777c2ac8d963d2d0c123023730502d21c62b725eb4bdcb351c33/a2a_sdk-1.1.5.tar.gz", hash = "sha256:f81837f86a9bbb3ed4909539be642d61e8539ce8300a6578fa71dc901ee65b97", size = 402715, upload-time = "2026-09-21T09:52:26.352Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/0b/91/087e4281bae91e19c4bb0736dc0807725e2e1f641dcc1ac86e625137c32c/a2a_sdk-1.1.5-py3-none-any.whl", hash = "sha256:c1d63d36b79a097c62dd9be001fce21d5939f5a811a9c5afd0ba8082f13f68bf", size = 255110, upload-time = "2026-09-21T09:52:24.741Z" }, +] + [[package]] name = "accelerate" version = "1.15.0" @@ -192,6 +211,20 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/a4/37/cfd1ed540a4d318da025590d96b728e63713c09e9377950fc655dadeb856/aiohttp-3.14.3-cp314-cp314t-win_arm64.whl", hash = "sha256:2e1161602f45a54de2ce0905243a95f58cb42dcd378402f3697f5e0b21e9d2e7", size = 469280, upload-time = "2026-07-23T01:57:24.241Z" }, ] +[[package]] +name = "aiologic" +version = "0.17.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "sniffio", marker = "python_full_version < '3.13'" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'" }, + { name = "wrapt", marker = "python_full_version < '3.13'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/f1/7a/d51f2fde1e8ae8a83431f8e97b7a71e9358cdb1d4d2ce6be387fa44d68de/aiologic-0.17.1.tar.gz", hash = "sha256:2e1b93b9e88ced318c2a63ad7b382688f40cbfe40e3d42258d49dc9c5aea179d", size = 252354, upload-time = "2026-06-27T20:41:33.25Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/5b/d3/2d310b1b839014034dba0cba685e492df8a5c7ad32c19cab7e979eed6554/aiologic-0.17.1-py3-none-any.whl", hash = "sha256:c66b319830fedb7ca3d2b2125fa6f5b653f89418c2a27ea76f259ec5f00943c0", size = 161331, upload-time = "2026-06-27T20:41:31.877Z" }, +] + [[package]] name = "aiosignal" version = "1.4.0" @@ -1238,7 +1271,7 @@ name = "cuda-bindings" version = "13.3.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "cuda-pathfinder" }, + { name = "cuda-pathfinder", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/51/6b/457ca12dad3ee9bfcc9a545cfd6b64b359ba49de40f776f6e028e678f262/cuda_bindings-13.3.1-cp311-cp311-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:c5879712accf6e14bb01aa5e67440eb84998b8d104b509cc7a6dc0b8f656a474", size = 6053539, upload-time = "2026-05-29T23:11:43.19Z" }, @@ -1271,43 +1304,56 @@ wheels = [ [package.optional-dependencies] cublas = [ - { name = "nvidia-cublas", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, - { name = "nvidia-cuda-nvrtc", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, + { name = "nvidia-cublas", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, + { name = "nvidia-cuda-nvrtc", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, ] cudart = [ - { name = "nvidia-cuda-runtime", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, + { name = "nvidia-cuda-runtime", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, ] cufft = [ - { name = "nvidia-cufft", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, - { name = "nvidia-nvjitlink", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, + { name = "nvidia-cufft", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, + { name = "nvidia-nvjitlink", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, ] cufile = [ - { name = "nvidia-cufile", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, + { name = "nvidia-cufile", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, ] cupti = [ - { name = "nvidia-cuda-cupti", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, + { name = "nvidia-cuda-cupti", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, ] curand = [ - { name = "nvidia-curand", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, + { name = "nvidia-curand", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, ] cusolver = [ - { name = "nvidia-cublas", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, - { name = "nvidia-cusolver", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, - { name = "nvidia-cusparse", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, - { name = "nvidia-nvjitlink", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, + { name = "nvidia-cublas", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, + { name = "nvidia-cusolver", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, + { name = "nvidia-cusparse", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, + { name = "nvidia-nvjitlink", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, ] cusparse = [ - { name = "nvidia-cusparse", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, - { name = "nvidia-nvjitlink", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, + { name = "nvidia-cusparse", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, + { name = "nvidia-nvjitlink", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, ] nvjitlink = [ - { name = "nvidia-nvjitlink", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, + { name = "nvidia-nvjitlink", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, ] nvrtc = [ - { name = "nvidia-cuda-nvrtc", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, + { name = "nvidia-cuda-nvrtc", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, ] nvtx = [ - { name = "nvidia-nvtx", marker = "platform_machine == 'aarch64' or platform_machine == 'x86_64'" }, + { name = "nvidia-nvtx", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, +] + +[[package]] +name = "culsans" +version = "0.11.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "aiologic", marker = "python_full_version < '3.13'" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/d9/e3/49afa1bc180e0d28008ec6bcdf82a4072d1c7a41032b5b759b60814ca4b0/culsans-0.11.0.tar.gz", hash = "sha256:0b43d0d05dce6106293d114c86e3fb4bfc63088cfe8ff08ed3fe36891447fe33", size = 107546, upload-time = "2025-12-31T23:15:38.196Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/e0/5d/9fb19fb38f6d6120422064279ea5532e22b84aa2be8831d49607194feda3/culsans-0.11.0-py3-none-any.whl", hash = "sha256:278d118f63fc75b9db11b664b436a1b83cc30d9577127848ba41420e66eb5a47", size = 21811, upload-time = "2025-12-31T23:15:37.189Z" }, ] [[package]] @@ -1808,6 +1854,48 @@ http = [ { name = "aiohttp" }, ] +[[package]] +name = "google-api-core" +version = "2.38.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "google-auth" }, + { name = "googleapis-common-protos" }, + { name = "opentelemetry-api" }, + { name = "proto-plus" }, + { name = "protobuf" }, + { name = "requests" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/59/e5/18aeff14213db86267a0f79d869352954394ab4ab3a3866d9b9b0e82dc6a/google_api_core-2.38.0.tar.gz", hash = "sha256:31e326eafa31b34f1db7a715f50f61dbca7e6e37277762f252ed90e6ce94c246", size = 206784, upload-time = "2026-09-17T20:30:08.766Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/18/55/e545e4b0040cb812f65d2abc545c6420f57b687d2af130a447c86d14d00b/google_api_core-2.38.0-py3-none-any.whl", hash = "sha256:8db2730375eba434bd75041fbe06b5374287407689ce132ac5650c8c8f042fcc", size = 187996, upload-time = "2026-09-17T20:29:32.813Z" }, +] + +[[package]] +name = "google-auth" +version = "2.58.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "cryptography" }, + { name = "pyasn1-modules" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/ac/ca/f398a483ce5aad18ca2f735646e45ccee2439bd94a41a4ad0cfa646bd495/google_auth-2.58.0.tar.gz", hash = "sha256:55e30cf15e737de92c5323d78cda8a83fcd57e7ffbaf900c4600039fd60a80fd", size = 380018, upload-time = "2026-09-09T20:49:38.043Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/59/13/477d90d09591b3938b45c4e11f4d8a51291682112cb5efcac961e815d562/google_auth-2.58.0-py3-none-any.whl", hash = "sha256:8a9c4645bb4c8e91668fb1934b95ae6a8687084232753639220ba9bf04a1610d", size = 262404, upload-time = "2026-09-09T20:49:33.951Z" }, +] + +[[package]] +name = "googleapis-common-protos" +version = "1.75.3" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "protobuf" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/8a/c5/4353a188e2c335aee33269e8b654af228278cca8e5f0b4b5f11e5d0e9adb/googleapis_common_protos-1.75.3.tar.gz", hash = "sha256:57c435ac2c68b108999b6db075d9053e4d7a936ba57b4a3d45667b1346f1738a", size = 153905, upload-time = "2026-09-03T22:31:21.869Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/1a/7a/7d79170c6ce6f12e109df2b3879d6b934010cf4f99aea8de8b7e5408c174/googleapis_common_protos-1.75.3-py3-none-any.whl", hash = "sha256:a018d2bf098ca9fb6faa08d5bb780e2a2c2f73c566f069761331386c9596d3f2", size = 306984, upload-time = "2026-09-03T22:30:45.133Z" }, +] + [[package]] name = "greenlet" version = "3.3.0" @@ -2321,6 +2409,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/14/2f/967ba146e6d58cf6a652da73885f52fc68001525b4197effc174321d70b4/jmespath-1.1.0-py3-none-any.whl", hash = "sha256:a5663118de4908c91729bea0acadca56526eb2698e83de10cd116ae0f4e97c64", size = 20419, upload-time = "2026-01-22T16:35:24.919Z" }, ] +[[package]] +name = "json-rpc" +version = "1.15.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/6d/9e/59f4a5b7855ced7346ebf40a2e9a8942863f644378d956f68bcef2c88b90/json-rpc-1.15.0.tar.gz", hash = "sha256:e6441d56c1dcd54241c937d0a2dcd193bdf0bdc539b5316524713f554b7f85b9", size = 28854, upload-time = "2023-06-11T09:45:49.078Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/94/9e/820c4b086ad01ba7d77369fb8b11470a01fac9b4977f02e18659cf378b6b/json_rpc-1.15.0-py2.py3-none-any.whl", hash = "sha256:4a4668bbbe7116feb4abbd0f54e64a4adcf4b8f648f19ffa0848ad0f6606a9bf", size = 39450, upload-time = "2023-06-11T09:45:47.136Z" }, +] + [[package]] name = "json5" version = "0.13.0" @@ -3519,7 +3616,7 @@ name = "nvidia-cublas" version = "13.1.1.3" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-cuda-nvrtc" }, + { name = "nvidia-cuda-nvrtc", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/a7/a1/0bd24ee8c8d03adac032fd2909426a00c88f8c57961b1277ded97f91119f/nvidia_cublas-13.1.1.3-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:b7a210458267ac818974c53038fbec2e969d5c99f305ab15c72522fa9f001dd5", size = 542848918, upload-time = "2026-04-08T18:46:22.985Z" }, @@ -3558,7 +3655,7 @@ name = "nvidia-cudnn-cu13" version = "9.24.0.43" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-cublas" }, + { name = "nvidia-cublas", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/ca/30/7c257e3d5cb4fecb147b93895c66e29c93f8e76d74b45bb418ff0587c4ec/nvidia_cudnn_cu13-9.24.0.43-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:a6812a554a1ff0413e9c52b84c26c050380649ab9615f9c16bded368ce9f421f", size = 650976863, upload-time = "2026-07-02T16:23:39.248Z" }, @@ -3570,7 +3667,7 @@ name = "nvidia-cufft" version = "12.0.0.61" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-nvjitlink" }, + { name = "nvidia-nvjitlink", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/8b/ae/f417a75c0259e85c1d2f83ca4e960289a5f814ed0cea74d18c353d3e989d/nvidia_cufft-12.0.0.61-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:2708c852ef8cd89d1d2068bdbece0aa188813a0c934db3779b9b1faa8442e5f5", size = 214053554, upload-time = "2025-09-04T08:31:38.196Z" }, @@ -3600,9 +3697,9 @@ name = "nvidia-cusolver" version = "12.0.4.66" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-cublas" }, - { name = "nvidia-cusparse" }, - { name = "nvidia-nvjitlink" }, + { name = "nvidia-cublas", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "nvidia-cusparse", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "nvidia-nvjitlink", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/c8/c3/b30c9e935fc01e3da443ec0116ed1b2a009bb867f5324d3f2d7e533e776b/nvidia_cusolver-12.0.4.66-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:02c2457eaa9e39de20f880f4bd8820e6a1cfb9f9a34f820eb12a155aa5bc92d2", size = 223467760, upload-time = "2025-09-04T08:33:04.222Z" }, @@ -3614,7 +3711,7 @@ name = "nvidia-cusparse" version = "12.6.3.3" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-nvjitlink" }, + { name = "nvidia-nvjitlink", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/f8/94/5c26f33738ae35276672f12615a64bd008ed5be6d1ebcb23579285d960a9/nvidia_cusparse-12.6.3.3-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:80bcc4662f23f1054ee334a15c72b8940402975e0eab63178fc7e670aa59472c", size = 162155568, upload-time = "2025-09-04T08:33:42.864Z" }, @@ -4063,7 +4160,7 @@ name = "pexpect" version = "4.9.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "ptyprocess" }, + { name = "ptyprocess", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/42/92/cc564bf6381ff43ce1f4d06852fc19a2f11d180f23dc32d9588bee2f149d/pexpect-4.9.0.tar.gz", hash = "sha256:ee7d41123f3c9911050ea2c2dac107568dc43b2d3b0c7557a33212c398ead30f", size = 166450, upload-time = "2023-11-25T09:07:26.339Z" } wheels = [ @@ -4359,6 +4456,33 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/5b/5a/bc7b4a4ef808fa59a816c17b20c4bef6884daebbdf627ff2a161da67da19/propcache-0.4.1-py3-none-any.whl", hash = "sha256:af2a6052aeb6cf17d3e46ee169099044fd8224cbaf75c76a2ef596e8163e2237", size = 13305, upload-time = "2025-10-08T19:49:00.792Z" }, ] +[[package]] +name = "proto-plus" +version = "1.28.4" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "protobuf" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/40/a6/4fbadcc2044034449b3f8f0ce82dcf3005d53f37c136642103fd4836a31c/proto_plus-1.28.4.tar.gz", hash = "sha256:5ff7ecad828e032a491fcb86947801768e32237f99dd049b649965b892ae9a63", size = 58679, upload-time = "2026-08-25T19:19:15.102Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/41/5d/0f04b85dafdc3250ced7f2592efc17dce7f40712e941e9632202481e600d/proto_plus-1.28.4-py3-none-any.whl", hash = "sha256:4b01341272f8a348db3f003b6143109f83ab43091019d5181b3fcdf500ab32aa", size = 50797, upload-time = "2026-08-25T19:18:12.338Z" }, +] + +[[package]] +name = "protobuf" +version = "7.36.2" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/d9/89/5b8517baa72f84a67b8a307ba953c91057af618bf40bf676f3c03551f8f0/protobuf-7.36.2.tar.gz", hash = "sha256:497d0463ff3316681da6c0b9e8d06cb465d61abce00b613ab42226175644d1bb", size = 512737, upload-time = "2026-09-17T20:07:59.326Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/32/72/98342feb672507c8f3a69e34b4fa8961f608edba5c1a48a6f47156d92cb5/protobuf-7.36.2-cp310-abi3-macosx_10_9_universal2.whl", hash = "sha256:cbc70b17ee27e28894c7fee8bb04be1abead49e936bc70eb60052531eee2079e", size = 456039, upload-time = "2026-09-17T20:07:51.542Z" }, + { url = "https://files.pythonhosted.org/packages/b6/ea/91fdf7c2b8bbd49cde056f00a9df6773532987e1c00fe2830b895af95c7e/protobuf-7.36.2-cp310-abi3-manylinux2014_aarch64.whl", hash = "sha256:e11e1f0180583a2af89db6a2ecd9e8dc40aa6d2988ca175bfd0e6d12ea72d74e", size = 344219, upload-time = "2026-09-17T20:07:52.914Z" }, + { url = "https://files.pythonhosted.org/packages/17/ab/5fd5f8ece73fad885c5a09aa849b32d70472f954ba3a92d3bb5974ea953b/protobuf-7.36.2-cp310-abi3-manylinux2014_s390x.whl", hash = "sha256:f4fee11ec330d238b34a05c9b675f693c20415d1c5bd7d5320cc2f8a798eb9cf", size = 357223, upload-time = "2026-09-17T20:07:53.985Z" }, + { url = "https://files.pythonhosted.org/packages/db/f3/3996583dd2906297a637af12114deddf7658af6e683fedb83be061983fb5/protobuf-7.36.2-cp310-abi3-manylinux2014_x86_64.whl", hash = "sha256:89f23aa53c24553a2416fd4fd1ec06f74fa42b14b546d8883128813f775bbfd2", size = 343223, upload-time = "2026-09-17T20:07:54.931Z" }, + { url = "https://files.pythonhosted.org/packages/fc/1b/dcc64f358fcb51811b58ae40b3d28f820725f116d86487cc20bd4b130701/protobuf-7.36.2-cp310-abi3-win32.whl", hash = "sha256:912c1221170e16c08d1f086762f563dd61ff83c18b5fa6652952dfaded66f728", size = 442998, upload-time = "2026-09-17T20:07:55.826Z" }, + { url = "https://files.pythonhosted.org/packages/8a/55/b77bda4e5e5f5971fb51b07663694690e9afdb9402136c16a522bd621cad/protobuf-7.36.2-cp310-abi3-win_amd64.whl", hash = "sha256:a300819d441e078a5608c0d3c709796bb548136058fda017ae51d425b44fd353", size = 456514, upload-time = "2026-09-17T20:07:57.188Z" }, + { url = "https://files.pythonhosted.org/packages/e4/04/d52c7016b04b6c5108f26691f9d33ec82a9b65d041f1a9c771137693d618/protobuf-7.36.2-py3-none-any.whl", hash = "sha256:bdb3a345d48db958e6ce1f18e508beb0cc981d64f24088427549c866cd039f1e", size = 179806, upload-time = "2026-09-17T20:07:58.211Z" }, +] + [[package]] name = "psutil" version = "7.2.1" @@ -4448,6 +4572,27 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/93/c0/37d4a7e8e2f7a6076283673d5298018ca26478b934c6ee369e10505ab32c/pyarrow-25.0.1-cp314-cp314t-win_amd64.whl", hash = "sha256:4288f27577352d608ca08553b0865e4a9b3aa14820c5d95b53337218d609835b", size = 28753071, upload-time = "2026-08-10T12:40:46.623Z" }, ] +[[package]] +name = "pyasn1" +version = "0.6.4" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/a4/9a/23310166d960def5897e91fe20e5b724601b02a22e84ba1f94232c0b7f67/pyasn1-0.6.4.tar.gz", hash = "sha256:9c447d8431c947fe4c8febc4ed9e760bc29011a5b01e5c74b67025bd9fb8ce81", size = 151262, upload-time = "2026-07-09T01:12:33.988Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/9a/3b/6163796d69c3977d1e4287bea4a6979161cbbdd170ebb430511e8e1999ce/pyasn1-0.6.4-py3-none-any.whl", hash = "sha256:deda9277cfd454080ec40b207fb6df82206a3a2688735233cdcd8d3d565f088b", size = 84410, upload-time = "2026-07-09T01:12:32.92Z" }, +] + +[[package]] +name = "pyasn1-modules" +version = "0.4.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pyasn1" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/e9/e6/78ebbb10a8c8e4b61a59249394a4a594c1a7af95593dc933a349c8d00964/pyasn1_modules-0.4.2.tar.gz", hash = "sha256:677091de870a80aae844b1ca6134f54652fa2c8c5a52aa396440ac3106e941e6", size = 307892, upload-time = "2025-03-28T02:41:22.17Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/47/8d/d529b5d697919ba8c11ad626e835d4039be708a35b0d22de83a269a6682c/pyasn1_modules-0.4.2-py3-none-any.whl", hash = "sha256:29253a9207ce32b64c3ac6600edc75368f98473906e8fd1043bd6b5b1de2c14a", size = 181259, upload-time = "2025-03-28T02:41:19.028Z" }, +] + [[package]] name = "pycparser" version = "2.23" @@ -4768,7 +4913,11 @@ dependencies = [ ] [package.optional-dependencies] +a2a = [ + { name = "a2a-sdk" }, +] all = [ + { name = "a2a-sdk" }, { name = "accelerate" }, { name = "azure-ai-ml" }, { name = "azure-cognitiveservices-speech" }, @@ -4833,6 +4982,7 @@ dev = [ { name = "pytest-xdist" }, { name = "respx" }, { name = "ruff" }, + { name = "sse-starlette" }, { name = "ty" }, { name = "types-aiofiles" }, { name = "types-pyyaml" }, @@ -4841,6 +4991,8 @@ dev = [ [package.metadata] requires-dist = [ + { name = "a2a-sdk", marker = "extra == 'a2a'", specifier = ">=1.1.5,<2" }, + { name = "a2a-sdk", marker = "extra == 'all'", specifier = ">=1.1.5,<2" }, { name = "accelerate", marker = "extra == 'all'", specifier = ">=1.7.0" }, { name = "accelerate", marker = "extra == 'gcg'", specifier = ">=1.7.0" }, { name = "aiofiles", specifier = ">=24,<26" }, @@ -4916,7 +5068,7 @@ requires-dist = [ { name = "uvicorn", extras = ["standard"], specifier = ">=0.32.0" }, { name = "websockets", specifier = ">=14.0" }, ] -provides-extras = ["huggingface", "gcg", "playwright", "fairness-bias", "opencv", "speech", "litellm", "all"] +provides-extras = ["huggingface", "gcg", "playwright", "fairness-bias", "opencv", "speech", "litellm", "a2a", "all"] [package.metadata.requires-dev] dev = [ @@ -4938,6 +5090,7 @@ dev = [ { name = "pytest-xdist", specifier = ">=3.6.1" }, { name = "respx", specifier = ">=0.22.0" }, { name = "ruff", specifier = ">=0.14.4" }, + { name = "sse-starlette", specifier = ">=3.0.0" }, { name = "ty", specifier = ">=0.0.32" }, { name = "types-aiofiles", specifier = ">=24.1.0" }, { name = "types-pyyaml", specifier = ">=6.0.12.20250516" }, @@ -5233,9 +5386,9 @@ resolution-markers = [ "python_full_version < '3.12' and sys_platform != 'emscripten' and sys_platform != 'win32'", ] dependencies = [ - { name = "attrs" }, - { name = "rpds-py" }, - { name = "typing-extensions" }, + { name = "attrs", marker = "python_full_version < '3.12'" }, + { name = "rpds-py", marker = "python_full_version < '3.12'" }, + { name = "typing-extensions", marker = "python_full_version < '3.12'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/2f/db/98b5c277be99dd18bfd91dd04e1b759cad18d1a338188c936e92f921c7e2/referencing-0.36.2.tar.gz", hash = "sha256:df2e89862cd09deabbdba16944cc3f10feb6b3e6f18e902f7cc25609a34775aa", size = 74744, upload-time = "2025-01-25T08:48:16.138Z" } wheels = [ @@ -5258,9 +5411,9 @@ resolution-markers = [ "python_full_version == '3.12.*' and sys_platform != 'emscripten' and sys_platform != 'win32'", ] dependencies = [ - { name = "attrs" }, - { name = "rpds-py" }, - { name = "typing-extensions", marker = "python_full_version < '3.13'" }, + { name = "attrs", marker = "python_full_version >= '3.12'" }, + { name = "rpds-py", marker = "python_full_version >= '3.12'" }, + { name = "typing-extensions", marker = "python_full_version == '3.12.*'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/22/f5/df4e9027acead3ecc63e50fe1e36aca1523e1719559c499951bb4b53188f/referencing-0.37.0.tar.gz", hash = "sha256:44aefc3142c5b842538163acb373e24cce6632bd54bdb01b21ad5863489f50d8", size = 78036, upload-time = "2025-10-13T15:30:48.871Z" } wheels = [