From a846ec58a2f33d2d857b6de7aeb16993eb75ef6a Mon Sep 17 00:00:00 2001 From: Alexander Zaikman Date: Sun, 26 Jul 2026 13:01:57 +0300 Subject: [PATCH 1/7] update A2A SDK gateway to 1.1.2 --- .github/dependabot.yml | 5 + examples/a2a_gateway/02_with_demo_agent.py | 2 +- .../a2a_gateway/demo_orchestrator/__main__.py | 2 +- examples/mixed/03_fact_checker_a2a.py | 2 +- examples/mixed/04_risk_reviewer_a2a.py | 2 +- pyproject.toml | 8 +- src/band/adapters/__init__.py | 5 + src/band/adapters/a2a_gateway.py | 2 + src/band/integrations/a2a/gateway/__init__.py | 2 + src/band/integrations/a2a/gateway/adapter.py | 240 +++-- src/band/integrations/a2a/gateway/config.py | 16 + src/band/integrations/a2a/gateway/server.py | 458 ++------- src/band/integrations/a2a/gateway/types.py | 14 +- tests/integrations/a2a/gateway/fixtures.py | 18 + .../integrations/a2a/gateway/test_adapter.py | 271 ++++-- tests/integrations/a2a/gateway/test_server.py | 891 ++---------------- tests/integrations/a2a/gateway/test_types.py | 125 --- uv.lock | 55 +- 18 files changed, 645 insertions(+), 1473 deletions(-) create mode 100644 src/band/integrations/a2a/gateway/config.py create mode 100644 tests/integrations/a2a/gateway/fixtures.py delete mode 100644 tests/integrations/a2a/gateway/test_types.py diff --git a/.github/dependabot.yml b/.github/dependabot.yml index 742a403e9..bcf7d8245 100644 --- a/.github/dependabot.yml +++ b/.github/dependabot.yml @@ -9,6 +9,11 @@ updates: day: "monday" time: "09:00" open-pull-requests-limit: 10 + groups: + uv-minor-and-patch: + update-types: + - "minor" + - "patch" labels: - "dependencies" - "python" diff --git a/examples/a2a_gateway/02_with_demo_agent.py b/examples/a2a_gateway/02_with_demo_agent.py index 406dbd992..948c70a29 100644 --- a/examples/a2a_gateway/02_with_demo_agent.py +++ b/examples/a2a_gateway/02_with_demo_agent.py @@ -68,7 +68,7 @@ sys.path.insert(0, str(Path(__file__).parent)) import uvicorn -from a2a.server.apps import A2AStarletteApplication +from a2a.server.apps.jsonrpc.starlette_app import A2AStarletteApplication from a2a.server.request_handlers import DefaultRequestHandler from a2a.server.tasks import ( InMemoryPushNotificationConfigStore, diff --git a/examples/a2a_gateway/demo_orchestrator/__main__.py b/examples/a2a_gateway/demo_orchestrator/__main__.py index d0c9ace7c..1d5029571 100644 --- a/examples/a2a_gateway/demo_orchestrator/__main__.py +++ b/examples/a2a_gateway/demo_orchestrator/__main__.py @@ -24,7 +24,7 @@ import click import uvicorn -from a2a.server.apps import A2AStarletteApplication +from a2a.server.apps.jsonrpc.starlette_app import A2AStarletteApplication from a2a.server.request_handlers import DefaultRequestHandler from a2a.server.tasks import ( InMemoryPushNotificationConfigStore, diff --git a/examples/mixed/03_fact_checker_a2a.py b/examples/mixed/03_fact_checker_a2a.py index 268741448..20fd7f129 100644 --- a/examples/mixed/03_fact_checker_a2a.py +++ b/examples/mixed/03_fact_checker_a2a.py @@ -25,7 +25,7 @@ import uvicorn from a2a.server.agent_execution import AgentExecutor, RequestContext -from a2a.server.apps import A2AStarletteApplication +from a2a.server.apps.jsonrpc.starlette_app import A2AStarletteApplication from a2a.server.events import EventQueue from a2a.server.request_handlers import DefaultRequestHandler from a2a.server.tasks import ( diff --git a/examples/mixed/04_risk_reviewer_a2a.py b/examples/mixed/04_risk_reviewer_a2a.py index 0410ddad6..ea248bf77 100644 --- a/examples/mixed/04_risk_reviewer_a2a.py +++ b/examples/mixed/04_risk_reviewer_a2a.py @@ -25,7 +25,7 @@ import uvicorn from a2a.server.agent_execution import AgentExecutor, RequestContext -from a2a.server.apps import A2AStarletteApplication +from a2a.server.apps.jsonrpc.starlette_app import A2AStarletteApplication from a2a.server.events import EventQueue from a2a.server.request_handlers import DefaultRequestHandler from a2a.server.tasks import ( diff --git a/pyproject.toml b/pyproject.toml index 49977ce90..a1f382644 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -87,16 +87,16 @@ gemini = [ "google-genai>=1.43.0", # httpx is a transitive dep used for retry exception types ] a2a = [ - "a2a-sdk>=0.3.22", # Official A2A Python SDK; brings protobuf>=5.29.5 transitively + "a2a-sdk>=1.1.2", # Official A2A Python SDK; brings protobuf transitively ] a2a_gateway = [ - "a2a-sdk>=0.3.22", # Official A2A Python SDK; brings protobuf>=5.29.5 transitively + "a2a-sdk>=1.1.2", # Official A2A Python SDK; brings protobuf transitively "starlette>=0.40.0", # HTTP server framework "uvicorn>=0.32.0", # ASGI server "python-multipart>=0.0.22", ] a2a_gateway_demo = [ - "a2a-sdk>=0.3.22", # A2A client for calling gateway peers + "a2a-sdk>=1.1.2", # A2A client for calling gateway peers "starlette>=0.40.0", # HTTP server "uvicorn>=0.32.0", # ASGI server "langgraph>=1.0.0", # LangGraph for orchestrator agent @@ -177,7 +177,7 @@ dev = [ "openai>=2.0.0", "beautifulsoup4>=4.12.0", # Include a2a-sdk for testing - "a2a-sdk>=0.3.22", + "a2a-sdk>=1.1.2", # Include a2a_gateway deps for testing "starlette>=0.40.0", "uvicorn>=0.32.0", diff --git a/src/band/adapters/__init__.py b/src/band/adapters/__init__.py index 62e257178..3f5620cd4 100644 --- a/src/band/adapters/__init__.py +++ b/src/band/adapters/__init__.py @@ -45,6 +45,9 @@ ) from band.adapters.a2a import A2AAdapter as A2AAdapter from band.adapters.a2a_gateway import A2AGatewayAdapter as A2AGatewayAdapter + from band.adapters.a2a_gateway import ( + A2AGatewayAdapterConfig as A2AGatewayAdapterConfig, + ) from band.adapters.codex import CodexAdapter as CodexAdapter from band.adapters.codex import CodexAdapterConfig as CodexAdapterConfig from band.adapters.acp import ( @@ -77,6 +80,7 @@ "CrewAIFlowAdapter", "A2AAdapter", "A2AGatewayAdapter", + "A2AGatewayAdapterConfig", "CodexAdapter", "CodexAdapterConfig", "ACPClientAdapter", @@ -110,6 +114,7 @@ "CrewAIFlowAdapter": "crewai_flow", "A2AAdapter": "a2a", "A2AGatewayAdapter": "a2a_gateway", + "A2AGatewayAdapterConfig": "a2a_gateway", "CodexAdapter": "codex", "CodexAdapterConfig": "codex", "ACPClientAdapter": "acp", diff --git a/src/band/adapters/a2a_gateway.py b/src/band/adapters/a2a_gateway.py index cce07b1fb..995a3b1ce 100644 --- a/src/band/adapters/a2a_gateway.py +++ b/src/band/adapters/a2a_gateway.py @@ -1,11 +1,13 @@ """A2A Gateway adapter - re-exports from integrations module.""" from band.integrations.a2a.gateway.adapter import A2AGatewayAdapter +from band.integrations.a2a.gateway.config import A2AGatewayAdapterConfig from band.integrations.a2a.gateway.server import GatewayServer from band.integrations.a2a.gateway.types import GatewaySessionState, PendingA2ATask __all__ = [ "A2AGatewayAdapter", + "A2AGatewayAdapterConfig", "GatewayServer", "GatewaySessionState", "PendingA2ATask", diff --git a/src/band/integrations/a2a/gateway/__init__.py b/src/band/integrations/a2a/gateway/__init__.py index 385624844..0f0ddf89a 100644 --- a/src/band/integrations/a2a/gateway/__init__.py +++ b/src/band/integrations/a2a/gateway/__init__.py @@ -1,11 +1,13 @@ """A2A Gateway adapter for exposing Band peers as A2A endpoints.""" from band.integrations.a2a.gateway.adapter import A2AGatewayAdapter +from band.integrations.a2a.gateway.config import A2AGatewayAdapterConfig from band.integrations.a2a.gateway.server import GatewayServer from band.integrations.a2a.gateway.types import GatewaySessionState, PendingA2ATask __all__ = [ "A2AGatewayAdapter", + "A2AGatewayAdapterConfig", "GatewayServer", "GatewaySessionState", "PendingA2ATask", diff --git a/src/band/integrations/a2a/gateway/adapter.py b/src/band/integrations/a2a/gateway/adapter.py index 46a40fec6..420bfba97 100644 --- a/src/band/integrations/a2a/gateway/adapter.py +++ b/src/band/integrations/a2a/gateway/adapter.py @@ -3,12 +3,16 @@ from __future__ import annotations import asyncio +from functools import partial import logging import re +from contextlib import asynccontextmanager from collections.abc import AsyncIterator from typing import ClassVar from uuid import uuid4 +from a2a.server.agent_execution import AgentExecutor, RequestContext +from a2a.server.events import EventQueue from a2a.types import ( Message as A2AMessage, Part, @@ -35,6 +39,7 @@ from band.core.simple_adapter import SimpleAdapter from band.core.types import AdapterFeatures, Capability, Emit, PlatformMessage from band.integrations.a2a.gateway.server import GatewayServer +from band.integrations.a2a.gateway.config import A2AGatewayAdapterConfig from band.integrations.a2a.gateway.types import GatewaySessionState, PendingA2ATask from band_rest import Peer from band_rest.agent_api_peers.types.list_agent_peers_response import ( @@ -45,6 +50,20 @@ logger = logging.getLogger(__name__) +class BandAgentExecutor(AgentExecutor): + """Adapt one official A2A handler execution to a Band peer.""" + + def __init__(self, adapter: A2AGatewayAdapter, peer_slug: str) -> None: + self.adapter = adapter + self.peer_slug = peer_slug + + async def execute(self, context: RequestContext, event_queue: EventQueue) -> None: + await self.adapter._execute_a2a(self.peer_slug, context, event_queue) + + async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None: + await self.adapter._cancel_a2a(context, event_queue) + + def slugify(name: str) -> str: """Convert name to URL-safe slug. @@ -101,6 +120,7 @@ def __init__( api_key: str = "", gateway_url: str = "http://localhost:10000", port: int = 10000, + config: A2AGatewayAdapterConfig | None = None, features: AdapterFeatures | None = None, ) -> None: """Initialize gateway adapter. @@ -110,6 +130,7 @@ def __init__( api_key: API key for authentication (same as Agent.create()). gateway_url: Base URL for A2A endpoints exposed by this gateway. port: Port for HTTP server to listen on. + config: A2A Gateway runtime configuration. """ super().__init__( history_converter=GatewayHistoryConverter(), @@ -117,6 +138,7 @@ def __init__( ) self.gateway_url = gateway_url self.port = port + self.config = config or A2AGatewayAdapterConfig() # Direct REST client for room/message operations self._rest = AsyncRestClient(base_url=rest_url, api_key=api_key) @@ -166,7 +188,7 @@ async def on_started(self, agent_name: str, agent_description: str) -> None: peers_by_uuid=self._peers_by_uuid, gateway_url=self.gateway_url, port=self.port, - on_request=self._handle_a2a_request, + executor_factory=partial(BandAgentExecutor, self), ) await self._server.start() @@ -252,16 +274,21 @@ async def on_message( if is_session_bootstrap and history: self._rehydrate(history) - # Find pending task for this room + # Find pending task for this room. pending = self._pending_tasks.get(room_id) if pending: - # Convert to A2A event and push to SSE queue + logger.debug( + "A2A response received: room=%s task=%s type=%s", + room_id, + pending.task.id, + msg.message_type, + ) event = self._translate_to_a2a(msg, pending.task) - await pending.sse_queue.put(event) - - # Clean up on terminal state - if event.final: - del self._pending_tasks[room_id] + await pending.publish_response(event) + else: + logger.debug( + "Ignoring Band message without pending A2A task: room=%s", room_id + ) async def on_cleanup(self, room_id: str) -> None: """Clean up resources for a room. @@ -269,8 +296,9 @@ async def on_cleanup(self, room_id: str) -> None: Args: room_id: The room identifier. """ - # Clean up pending task if exists - self._pending_tasks.pop(room_id, None) + pending = self._pending_tasks.pop(room_id, None) + if pending: + pending.done.set() logger.debug("Cleaned up gateway resources for room %s", room_id) async def stop(self) -> None: @@ -295,75 +323,144 @@ def _resolve_peer(self, peer_id: str) -> Peer | None: # Try UUID fallback return self._peers_by_uuid.get(peer_id) - async def _handle_a2a_request( - self, peer_id: str, message: A2AMessage - ) -> AsyncIterator[TaskStatusUpdateEvent]: - """Handle incoming A2A request from remote agent. - - Args: - peer_id: Target peer slug or UUID. - message: A2A message from remote agent. + def _make_task(self, context: RequestContext) -> Task: + """Create the task event emitted before Band execution starts.""" + return Task( + id=context.task_id, + context_id=context.context_id, + status=TaskStatus(state=TaskState.working), + ) - Yields: - TaskStatusUpdateEvent for SSE streaming. - """ - # Resolve peer from slug or UUID + async def _execute_a2a( + self, peer_id: str, context: RequestContext, event_queue: EventQueue + ) -> None: + """Bridge one official A2A execution to a Band room.""" peer = self._resolve_peer(peer_id) if not peer: - logger.error("Peer not found: %s", peer_id) - return + logger.warning("A2A request target not found: peer=%s", peer_id) + raise ValueError(f"Peer not found: {peer_id}") - # Use the peer's actual UUID for Band API calls peer_uuid = peer.id - - # Get or create room for context room_id, context_id = await self._get_or_create_room( - message.context_id, peer_uuid + context.context_id, peer_uuid ) - - # Create A2A task - task = self._create_task(context_id) - - # Register pending task with SSE queue - sse_queue: asyncio.Queue[TaskStatusUpdateEvent] = asyncio.Queue() - self._pending_tasks[room_id] = PendingA2ATask( + task = self._make_task(context) + pending = PendingA2ATask( task=task, - sse_queue=sse_queue, + event_queue=event_queue, peer_id=peer_uuid, + done=asyncio.Event(), ) - - # Emit task event to track context mapping in history - await self._emit_context_event(room_id, context_id) - - # Send message to Band via REST client - content = get_message_text(message) or "" - - # Use peer name for mention - peer_name = peer.name - - await self._rest.agent_api_messages.create_agent_chat_message( - chat_id=room_id, - message=ChatMessageRequest( - content=f"@{peer_name} {content}", - mentions=[ChatMessageRequestMentionsItem(id=peer_uuid, name=peer_name)], - ), - request_options=DEFAULT_REQUEST_OPTIONS, + logger.info( + "A2A request started: peer=%s room=%s context=%s task=%s", + peer_id, + room_id, + context_id, + task.id, ) + try: + async with self.pending_task(room_id, pending): + await event_queue.enqueue_event(task) + await self._emit_context_event(room_id, context_id) + content = get_message_text(context.message) or "" + await self._rest.agent_api_messages.create_agent_chat_message( + chat_id=room_id, + message=ChatMessageRequest( + content=f"@{peer.name} {content}", + mentions=[ + ChatMessageRequestMentionsItem(id=peer_uuid, name=peer.name) + ], + ), + request_options=DEFAULT_REQUEST_OPTIONS, + ) + logger.debug( + "A2A request sent to Band: room=%s task=%s", + room_id, + task.id, + ) + try: + if self.config.response_timeout_s is None: + await pending.done.wait() + else: + async with asyncio.timeout(self.config.response_timeout_s): + await pending.done.wait() + except TimeoutError: + logger.warning( + "A2A response timed out: room=%s task=%s timeout=%ss", + room_id, + task.id, + self.config.response_timeout_s, + ) + task.status = TaskStatus( + state=TaskState.failed, + message=A2AMessage( + role=Role.agent, + message_id=str(uuid4()), + parts=[ + Part( + root=TextPart( + text="Timed out waiting for a Band response" + ) + ) + ], + ), + ) + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id=task.id, + context_id=task.context_id, + status=task.status, + final=True, + ) + ) + except asyncio.CancelledError: + logger.debug("A2A request cancelled: room=%s task=%s", room_id, task.id) + raise + except Exception: + logger.exception( + "A2A request failed: room=%s context=%s task=%s", + room_id, + context_id, + task.id, + ) + raise + else: + logger.info("A2A request completed: room=%s task=%s", room_id, task.id) + + @asynccontextmanager + async def pending_task( + self, room_id: str, pending: PendingA2ATask + ) -> AsyncIterator[PendingA2ATask]: + """Register one pending request and always release its room slot.""" + if room_id in self._pending_tasks: + raise RuntimeError(f"Room already has a pending A2A task: {room_id}") + self._pending_tasks[room_id] = pending logger.debug( - "Sent message to peer %s (%s) in room %s (context=%s)", - peer_name, - peer_uuid, - room_id, - context_id, + "Registered pending A2A task: room=%s task=%s", room_id, pending.task.id ) + try: + yield pending + finally: + self._pending_tasks.pop(room_id, None) + logger.debug( + "Released pending A2A task: room=%s task=%s", room_id, pending.task.id + ) - # Stream events from queue (populated by on_message()) - while True: - event = await sse_queue.get() - yield event - if event.final: - break + async def _cancel_a2a( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + """Publish the official terminal cancellation event.""" + task = context.current_task or self._make_task(context) + task.status = TaskStatus(state=TaskState.canceled) + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id=task.id, + context_id=task.context_id, + status=task.status, + final=True, + ) + ) async def _get_or_create_room( self, context_id: str | None, target_peer_id: str @@ -452,21 +549,6 @@ def _rehydrate(self, history: GatewaySessionState) -> None: len(self._room_participants), ) - def _create_task(self, context_id: str) -> Task: - """Create a new A2A Task for tracking. - - Args: - context_id: A2A context ID. - - Returns: - New Task instance. - """ - return Task( - id=str(uuid4()), - context_id=context_id, - status=TaskStatus(state=TaskState.working), - ) - def _translate_to_a2a( self, msg: PlatformMessage, task: Task ) -> TaskStatusUpdateEvent: diff --git a/src/band/integrations/a2a/gateway/config.py b/src/band/integrations/a2a/gateway/config.py new file mode 100644 index 000000000..d944e7dc7 --- /dev/null +++ b/src/band/integrations/a2a/gateway/config.py @@ -0,0 +1,16 @@ +"""Configuration for the A2A Gateway adapter.""" + +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass +class A2AGatewayAdapterConfig: + """Runtime policy for the A2A Gateway adapter.""" + + response_timeout_s: float | None = 300.0 + + def __post_init__(self) -> None: + if self.response_timeout_s is not None and self.response_timeout_s <= 0: + raise ValueError("response_timeout_s must be positive or None") diff --git a/src/band/integrations/a2a/gateway/server.py b/src/band/integrations/a2a/gateway/server.py index e9a304a7d..882d9fc0e 100644 --- a/src/band/integrations/a2a/gateway/server.py +++ b/src/band/integrations/a2a/gateway/server.py @@ -1,50 +1,30 @@ -"""HTTP server for A2A Gateway adapter.""" +"""HTTP server for the A2A Gateway adapter.""" from __future__ import annotations import asyncio -import json import logging -from collections.abc import AsyncIterator, Callable +from collections.abc import Callable from typing import Any -from a2a.types import ( - AgentCapabilities, - AgentCard, - AgentSkill, - Message as A2AMessage, - Task, - TaskStatusUpdateEvent, -) -from pydantic import ValidationError +from a2a.server.agent_execution import AgentExecutor +from a2a.server.apps.jsonrpc.starlette_app import A2AStarletteApplication +from a2a.server.apps.rest.rest_adapter import RESTAdapter +from a2a.server.request_handlers import DefaultRequestHandler +from a2a.server.tasks import InMemoryTaskStore +from a2a.types import AgentCapabilities, AgentCard, AgentInterface, AgentSkill from starlette.applications import Starlette -from starlette.requests import Request -from starlette.responses import JSONResponse, StreamingResponse from starlette.routing import Route from band_rest import Peer logger = logging.getLogger(__name__) -# Type alias for the on_request callback -OnRequestCallback = Callable[[str, A2AMessage], AsyncIterator[TaskStatusUpdateEvent]] +ExecutorFactory = Callable[[str], AgentExecutor] class GatewayServer: - """Starlette HTTP server exposing A2A endpoints for each peer. - - Creates per-peer routes for AgentCard discovery and message streaming. - Routes are created at startup based on the discovered peers. - - Peers are addressed by slug (e.g., "weather-agent") with UUID fallback. - - Attributes: - peers: Dict mapping slug to Peer objects (primary lookup). - peers_by_uuid: Dict mapping UUID to Peer objects (fallback lookup). - gateway_url: Base URL for the gateway (used in AgentCard URLs). - port: Port to listen on. - on_request: Callback invoked when an A2A message is received. - """ + """Expose each discovered Band peer through the official A2A server routes.""" def __init__( self, @@ -52,121 +32,37 @@ def __init__( peers_by_uuid: dict[str, Peer], gateway_url: str, port: int, - on_request: OnRequestCallback, + executor_factory: ExecutorFactory, ) -> None: - """Initialize gateway server. - - Args: - peers: Dict of slug → Peer objects (primary lookup). - peers_by_uuid: Dict of UUID → Peer objects (fallback lookup). - gateway_url: Base URL for AgentCard URLs (e.g., "http://localhost:10000"). - port: Port to listen on. - on_request: Async callback invoked with (peer_id, message) → events. - """ self.peers = peers self.peers_by_uuid = peers_by_uuid - self.gateway_url = gateway_url + self.gateway_url = gateway_url.rstrip("/") self.port = port - self.on_request = on_request + self.executor_factory = executor_factory self._app: Starlette | None = None self._server_task: asyncio.Task[Any] | None = None def _resolve_peer(self, peer_id: str) -> tuple[str, Peer] | None: - """Resolve peer by slug or UUID. - - Args: - peer_id: Peer slug or UUID from URL path. - - Returns: - Tuple of (slug, Peer) if found, None otherwise. - """ - # Try slug first (primary) if peer_id in self.peers: return peer_id, self.peers[peer_id] - # Try UUID fallback - if peer_id in self.peers_by_uuid: - peer = self.peers_by_uuid[peer_id] - # Find the slug for this peer - for slug, p in self.peers.items(): - if p.id == peer.id: - return slug, peer - return None - - def _build_app(self) -> Starlette: - """Build Starlette application with routes.""" - routes = [ - # List all available peers - Route( - "/peers", - self._handle_list_peers, - methods=["GET"], - ), - # Per-peer agent card discovery (support both naming conventions) - Route( - "/agents/{peer_id}/.well-known/agent.json", - self._handle_agent_card, - methods=["GET"], - ), - Route( - "/agents/{peer_id}/.well-known/agent-card.json", - self._handle_agent_card, - methods=["GET"], + peer = self.peers_by_uuid.get(peer_id) + if peer is None: + return None + return next( + ( + (slug, candidate) + for slug, candidate in self.peers.items() + if candidate.id == peer.id ), - # Per-peer message streaming (legacy REST endpoint) - Route( - "/agents/{peer_id}/v1/message:stream", - self._handle_message_stream, - methods=["POST"], - ), - # JSON-RPC endpoint (A2A SDK posts here with method field) - Route( - "/agents/{peer_id}", - self._handle_jsonrpc, - methods=["POST"], - ), - ] - return Starlette(routes=routes) - - async def _handle_list_peers(self, request: Request) -> JSONResponse: - """Return list of all available peers. - - Args: - request: Starlette request. - - Returns: - JSONResponse with list of peers (slug, id, name, description). - """ - peers_list = [ - { - "slug": slug, # Primary identifier for URLs - "id": peer.id, # UUID fallback - "name": peer.name, # Display name - "description": peer.description or "", - } - for slug, peer in self.peers.items() - ] - return JSONResponse({"peers": peers_list, "count": len(peers_list)}) - - async def _handle_agent_card(self, request: Request) -> JSONResponse: - """Return AgentCard for the specified peer. - - Args: - request: Starlette request with peer_id path parameter (slug or UUID). - - Returns: - JSONResponse with AgentCard or 404 if peer not found. - """ - peer_id = request.path_params["peer_id"] - resolved = self._resolve_peer(peer_id) - if not resolved: - return JSONResponse({"error": "Not found"}, status_code=404) - - slug, peer = resolved + None, + ) - card = AgentCard( + def _agent_card(self, slug: str, peer: Peer) -> AgentCard: + rpc_url = f"{self.gateway_url}/agents/{slug}" + return AgentCard( name=peer.name, description=peer.description or "", - url=f"{self.gateway_url}/agents/{slug}", # Use slug in URL + url=rpc_url, version="1.0.0", capabilities=AgentCapabilities(streaming=True), skills=[ @@ -179,284 +75,74 @@ async def _handle_agent_card(self, request: Request) -> JSONResponse: ], default_input_modes=["text/plain"], default_output_modes=["text/plain"], + preferred_transport="JSONRPC", + additional_interfaces=[AgentInterface(transport="JSONRPC", url=rpc_url)], ) - return JSONResponse(card.model_dump(mode="json", by_alias=True)) - - async def _handle_message_stream( - self, request: Request - ) -> JSONResponse | StreamingResponse: - """Handle incoming A2A message and stream response events. - - Args: - request: Starlette request with peer_id path parameter (slug or UUID) and JSON body. - - Returns: - StreamingResponse with SSE events. - """ - peer_id = request.path_params["peer_id"] - - # Resolve peer by slug or UUID - resolved = self._resolve_peer(peer_id) - if not resolved: - return JSONResponse({"error": "Not found"}, status_code=404) - - slug, peer = resolved - - body = await request.json() - try: - message = A2AMessage(**body) - except ValidationError as exc: - return JSONResponse( - { - "error": "Invalid A2A message payload", - "details": exc.errors(), - }, - status_code=400, - ) - - logger.debug( - "Received A2A message for peer %s (%s): %s", - peer.name, - slug, - message.message_id, - ) - - async def event_stream() -> AsyncIterator[str]: - """Generate SSE events from on_request callback.""" - try: - # Pass slug to callback (adapter resolves to peer) - async for event in self.on_request(slug, message): - yield f"data: {event.model_dump_json()}\n\n" - except Exception as e: - logger.exception("Error in A2A request handler: %s", e) - # Send error event - yield f'data: {{"error": "{e!s}"}}\n\n' - - return StreamingResponse(event_stream(), media_type="text/event-stream") - - async def _handle_jsonrpc( - self, request: Request - ) -> JSONResponse | StreamingResponse: - """Handle JSON-RPC requests from A2A SDK. - - The A2A SDK posts JSON-RPC requests to the base agent URL with - a method field to indicate the operation (message/send, message/stream, etc.). - - Args: - request: Starlette request with peer_id path parameter and JSON-RPC body. - - Returns: - JSONResponse for sync methods, StreamingResponse for streaming methods. - """ - peer_id = request.path_params["peer_id"] - - # Resolve peer by slug or UUID - resolved = self._resolve_peer(peer_id) - if not resolved: - return JSONResponse( - { - "jsonrpc": "2.0", - "error": {"code": -32001, "message": "Peer not found"}, - "id": None, - }, - status_code=404, - ) - - slug, peer = resolved - body = await request.json() - - # Extract JSON-RPC fields - method = body.get("method", "") - request_id = body.get("id") - params = body.get("params", {}) - - logger.debug( - "Received JSON-RPC %s for peer %s (%s), request_id=%s", - method, - peer.name, - slug, - request_id, - ) - - # Route based on method - if method == "message/send": - return await self._handle_jsonrpc_send(slug, params, request_id) - elif method == "message/stream": - return await self._handle_jsonrpc_stream(slug, params, request_id) - else: - return JSONResponse( - { - "jsonrpc": "2.0", - "error": {"code": -32601, "message": f"Method not found: {method}"}, - "id": request_id, - }, - status_code=400, - ) - - async def _handle_jsonrpc_send( - self, slug: str, params: dict[str, Any], request_id: str | None - ) -> JSONResponse: - """Handle synchronous message/send JSON-RPC request. - - Args: - slug: Peer slug. - params: JSON-RPC params containing message. - request_id: JSON-RPC request ID. - Returns: - JSONResponse with JSON-RPC result (Task). - """ - # Extract message from params - message_data = params.get("message", {}) - try: - message = A2AMessage(**message_data) - except ValidationError as exc: - return JSONResponse( - { - "jsonrpc": "2.0", - "error": { - "code": -32602, - "message": "Invalid params", - "data": exc.errors(), - }, - "id": request_id, - }, - status_code=400, - ) + def _build_app(self) -> Starlette: + routes: list[Route] = [ + Route("/peers", self._handle_list_peers, methods=["GET"]), + ] - # Collect all events until final - final_event: TaskStatusUpdateEvent | None = None - try: - async for event in self.on_request(slug, message): - final_event = event - if event.final: - break - except Exception as e: - logger.exception("Error in JSON-RPC send handler: %s", e) - return JSONResponse( - { - "jsonrpc": "2.0", - "error": {"code": -32000, "message": str(e)}, - "id": request_id, - }, - status_code=500, + for slug, peer in self.peers.items(): + card = self._agent_card(slug, peer) + handler = DefaultRequestHandler( + agent_executor=self.executor_factory(slug), + task_store=InMemoryTaskStore(), ) - - # Build Task result from final event - if final_event: - task = Task( - id=final_event.task_id, - context_id=final_event.context_id, - status=final_event.status, + jsonrpc_app = A2AStarletteApplication( + agent_card=card, + http_handler=handler, ) - return JSONResponse( - { - "jsonrpc": "2.0", - "result": task.model_dump(mode="json", by_alias=True), - "id": request_id, - } + rest_adapter = RESTAdapter(agent_card=card, http_handler=handler) + routes.extend( + jsonrpc_app.routes( + agent_card_url=f"/agents/{slug}/.well-known/agent.json", + rpc_url=f"/agents/{slug}", + ) ) - else: - # No events received - create empty task - return JSONResponse( - { - "jsonrpc": "2.0", - "error": {"code": -32000, "message": "No response from peer"}, - "id": request_id, - }, - status_code=500, + routes.extend( + Route( + f"/agents/{slug}{path}", + endpoint, + methods=[method], + ) + for (path, method), endpoint in rest_adapter.routes().items() ) - async def _handle_jsonrpc_stream( - self, slug: str, params: dict[str, Any], request_id: str | None - ) -> StreamingResponse: - """Handle streaming message/stream JSON-RPC request. - - Args: - slug: Peer slug. - params: JSON-RPC params containing message. - request_id: JSON-RPC request ID. - - Returns: - StreamingResponse with SSE events in JSON-RPC format. - """ - # Extract message from params - message_data = params.get("message", {}) - try: - message = A2AMessage(**message_data) - except ValidationError as exc: - return StreamingResponse( - iter( - [ - "data: " - + json.dumps( - { - "jsonrpc": "2.0", - "error": { - "code": -32602, - "message": "Invalid params", - "data": exc.errors(), - }, - "id": request_id, - } - ) - + "\n\n" - ] - ), - media_type="text/event-stream", - status_code=400, - ) + return Starlette(routes=routes) - async def event_stream() -> AsyncIterator[str]: - """Generate SSE events in JSON-RPC format.""" - try: - async for event in self.on_request(slug, message): - # Wrap event in JSON-RPC response - jsonrpc_response = { - "jsonrpc": "2.0", - "result": event.model_dump(mode="json", by_alias=True), - "id": request_id, - } - yield f"data: {json.dumps(jsonrpc_response)}\n\n" - except Exception as e: - logger.exception("Error in JSON-RPC stream handler: %s", e) - error_response = { - "jsonrpc": "2.0", - "error": {"code": -32000, "message": str(e)}, - "id": request_id, - } - yield f"data: {json.dumps(error_response)}\n\n" + async def _handle_list_peers(self, _request: Any) -> Any: + from starlette.responses import JSONResponse - return StreamingResponse(event_stream(), media_type="text/event-stream") + peers = [ + { + "slug": slug, + "id": peer.id, + "name": peer.name, + "description": peer.description or "", + } + for slug, peer in self.peers.items() + ] + return JSONResponse({"peers": peers, "count": len(peers)}) async def start(self) -> None: - """Start the HTTP server. - - Creates the Starlette app and starts serving on the configured port. - """ import uvicorn self._app = self._build_app() - - config = uvicorn.Config( - self._app, - host="0.0.0.0", - port=self.port, - log_level="warning", + server = uvicorn.Server( + uvicorn.Config( + self._app, host="0.0.0.0", port=self.port, log_level="warning" + ) ) - server = uvicorn.Server(config) - + self._server_task = asyncio.create_task(server.serve()) logger.info( "Starting A2A Gateway server on port %d with %d peers", self.port, len(self.peers), ) - # Run server in background task - self._server_task = asyncio.create_task(server.serve()) - async def stop(self) -> None: - """Stop the HTTP server.""" if self._server_task: self._server_task.cancel() try: diff --git a/src/band/integrations/a2a/gateway/types.py b/src/band/integrations/a2a/gateway/types.py index c26e9d3fc..9196a0334 100644 --- a/src/band/integrations/a2a/gateway/types.py +++ b/src/band/integrations/a2a/gateway/types.py @@ -5,6 +5,7 @@ import asyncio from dataclasses import dataclass, field +from a2a.server.events import EventQueue from a2a.types import Task, TaskStatusUpdateEvent @@ -34,10 +35,19 @@ class PendingA2ATask: Attributes: task: The A2A Task object tracking this request. - sse_queue: Queue for streaming TaskStatusUpdateEvent to the client. + event_queue: Official A2A event queue owned by DefaultRequestHandler. peer_id: The target peer this request is for. + done: Set when the final Band reply has been emitted or the room is + cleaned up. """ task: Task - sse_queue: asyncio.Queue[TaskStatusUpdateEvent] + event_queue: EventQueue peer_id: str + done: asyncio.Event + + async def publish_response(self, event: TaskStatusUpdateEvent) -> None: + """Publish a response and release the executor on terminal events.""" + await self.event_queue.enqueue_event(event) + if event.final: + self.done.set() diff --git a/tests/integrations/a2a/gateway/fixtures.py b/tests/integrations/a2a/gateway/fixtures.py new file mode 100644 index 000000000..66a0f4c01 --- /dev/null +++ b/tests/integrations/a2a/gateway/fixtures.py @@ -0,0 +1,18 @@ +"""Shared fixtures and builders for gateway tests.""" + +from __future__ import annotations + +from band_rest import Peer + + +def make_peer(peer_id: str, name: str, description: str = "") -> Peer: + """Build a representative registry peer for gateway tests.""" + return Peer( + id=peer_id, + name=name, + type="Agent", + description=description, + handle=f"test/{name.lower().replace(' ', '-')}", + is_contact=False, + source="registry", + ) diff --git a/tests/integrations/a2a/gateway/test_adapter.py b/tests/integrations/a2a/gateway/test_adapter.py index bf2f298a9..6938563c6 100644 --- a/tests/integrations/a2a/gateway/test_adapter.py +++ b/tests/integrations/a2a/gateway/test_adapter.py @@ -8,8 +8,11 @@ from uuid import uuid4 import pytest +from a2a.server.events import EventQueue +from a2a.server.agent_execution import RequestContext from a2a.types import ( Message as A2AMessage, + MessageSendParams, Part, Role, TaskState, @@ -19,11 +22,13 @@ from band.core.types import PlatformMessage from band.integrations.a2a.gateway import ( A2AGatewayAdapter, + A2AGatewayAdapterConfig, GatewaySessionState, ) +from band.integrations.a2a.gateway.adapter import BandAgentExecutor from band.testing import FakeAgentTools -from band_rest import Peer from band_rest.core.api_error import ApiError +from tests.integrations.a2a.gateway.fixtures import make_peer def make_platform_message( @@ -43,19 +48,6 @@ def make_platform_message( ) -def make_peer(peer_id: str, name: str, description: str = "") -> Peer: - """Create a mock Peer object.""" - return Peer( - id=peer_id, - name=name, - type="Agent", - description=description, - handle=f"test/{name.lower().replace(' ', '-')}", - is_contact=False, - source="registry", - ) - - def make_a2a_message( content: str, context_id: str | None = None, task_id: str | None = None ) -> A2AMessage: @@ -72,6 +64,11 @@ def make_a2a_message( class TestA2AGatewayAdapterInit: """Tests for A2AGatewayAdapter initialization.""" + def test_config_rejects_non_positive_response_timeout(self) -> None: + """Should reject a timeout that cannot provide a deadline.""" + with pytest.raises(ValueError, match="response_timeout_s"): + A2AGatewayAdapterConfig(response_timeout_s=0) + def test_init_default_values(self) -> None: """Should initialize with default values.""" adapter = A2AGatewayAdapter() @@ -96,21 +93,6 @@ def test_init_with_custom_values(self) -> None: assert adapter.gateway_url == "http://localhost:9000" assert adapter.port == 9000 - def test_init_creates_rest_client(self) -> None: - """Should create AsyncRestClient.""" - adapter = A2AGatewayAdapter( - rest_url="https://api.example.com", - api_key="my-key", - ) - - assert adapter._rest is not None - - def test_init_sets_history_converter(self) -> None: - """Should set GatewayHistoryConverter.""" - adapter = A2AGatewayAdapter() - - assert adapter.history_converter is not None - class TestA2AGatewayAdapterOnStarted: """Tests for A2AGatewayAdapter.on_started().""" @@ -174,31 +156,6 @@ async def test_on_started_starts_http_server(self) -> None: mock_server.start.assert_called_once() assert adapter._server is mock_server - @pytest.mark.asyncio - async def test_on_started_stores_agent_info(self) -> None: - """Should store agent name and description.""" - adapter = A2AGatewayAdapter() - - # Mock REST client - mock_response = MagicMock() - mock_response.data = [] - adapter._rest.agent_api_peers.list_agent_peers = AsyncMock( - return_value=mock_response - ) - - # Mock server - with patch( - "band.integrations.a2a.gateway.adapter.GatewayServer" - ) as mock_server_class: - mock_server = MagicMock() - mock_server.start = AsyncMock() - mock_server_class.return_value = mock_server - - await adapter.on_started("Test Gateway", "A test gateway") - - assert adapter.agent_name == "Test Gateway" - assert adapter.agent_description == "A test gateway" - @pytest.mark.asyncio async def test_on_started_retries_peer_discovery_on_rate_limit(self) -> None: """Should retry peer discovery when startup hits HTTP 429.""" @@ -280,15 +237,16 @@ async def test_on_message_rehydrates_on_bootstrap( async def test_on_message_correlates_pending_task( self, adapter_with_mocks: A2AGatewayAdapter ) -> None: - """Should push event to pending task's SSE queue.""" + """Should publish non-final updates without completing the task.""" tools = FakeAgentTools() - msg = make_platform_message("Weather is sunny", room_id="room-123") + msg = make_platform_message( + "Checking the forecast", room_id="room-123", message_type="thought" + ) - # Set up pending task from band.integrations.a2a.gateway.types import PendingA2ATask from a2a.types import Task, TaskStatus - sse_queue: asyncio.Queue = asyncio.Queue() + event_queue = EventQueue() task = Task( id="task-123", context_id="ctx-123", @@ -296,8 +254,9 @@ async def test_on_message_correlates_pending_task( ) adapter_with_mocks._pending_tasks["room-123"] = PendingA2ATask( task=task, - sse_queue=sse_queue, + event_queue=event_queue, peer_id="weather", + done=asyncio.Event(), ) await adapter_with_mocks.on_message( @@ -310,17 +269,17 @@ async def test_on_message_correlates_pending_task( room_id="room-123", ) - # Event should be in queue - assert not sse_queue.empty() - event = sse_queue.get_nowait() + pending = adapter_with_mocks._pending_tasks["room-123"] + event = await event_queue.dequeue_event() assert event.task_id == "task-123" - assert event.final is True # text message = completed + assert event.final is False + assert not pending.done.is_set() @pytest.mark.asyncio - async def test_on_message_cleans_up_on_final_event( + async def test_on_message_completes_pending_task( self, adapter_with_mocks: A2AGatewayAdapter ) -> None: - """Should clean up pending task on final event.""" + """Should complete the pending task on a final response.""" tools = FakeAgentTools() msg = make_platform_message("Done", room_id="room-123") @@ -328,7 +287,7 @@ async def test_on_message_cleans_up_on_final_event( from band.integrations.a2a.gateway.types import PendingA2ATask from a2a.types import Task, TaskStatus - sse_queue: asyncio.Queue = asyncio.Queue() + event_queue = EventQueue() task = Task( id="task-123", context_id="ctx-123", @@ -336,8 +295,9 @@ async def test_on_message_cleans_up_on_final_event( ) adapter_with_mocks._pending_tasks["room-123"] = PendingA2ATask( task=task, - sse_queue=sse_queue, + event_queue=event_queue, peer_id="weather", + done=asyncio.Event(), ) await adapter_with_mocks.on_message( @@ -350,8 +310,172 @@ async def test_on_message_cleans_up_on_final_event( room_id="room-123", ) - # Pending task should be cleaned up - assert "room-123" not in adapter_with_mocks._pending_tasks + pending = adapter_with_mocks._pending_tasks["room-123"] + assert pending.done.is_set() + + +class TestA2AGatewayExecutor: + """Tests the Band-backed execution boundary used by official A2A handlers.""" + + @pytest.mark.asyncio + async def test_execute_posts_and_emits_ordered_task_response(self) -> None: + adapter = A2AGatewayAdapter() + adapter._peers = {"weather": make_peer("weather", "Weather Agent")} + + chat_response = MagicMock() + chat_response.data.id = "room-123" + adapter._rest.agent_api_chats.create_agent_chat = AsyncMock( + return_value=chat_response + ) + adapter._rest.agent_api_participants.add_agent_chat_participant = AsyncMock() + message_sent = asyncio.Event() + + async def send_message(**_kwargs) -> None: + message_sent.set() + + adapter._rest.agent_api_messages.create_agent_chat_message = AsyncMock( + side_effect=send_message + ) + adapter._rest.agent_api_events.create_agent_chat_event = AsyncMock() + + context = RequestContext( + request=MessageSendParams(message=make_a2a_message("What is the weather?")) + ) + event_queue = EventQueue() + executor = BandAgentExecutor(adapter, "weather") + execution = asyncio.create_task(executor.execute(context, event_queue)) + + await asyncio.wait_for(message_sent.wait(), timeout=1) + initial = await event_queue.dequeue_event() + assert initial.kind == "task" + assert initial.status.state == TaskState.working + + await adapter.on_message( + make_platform_message("Sunny", room_id="room-123"), + FakeAgentTools(), + GatewaySessionState(), + None, + None, + is_session_bootstrap=False, + room_id="room-123", + ) + await asyncio.wait_for(execution, timeout=1) + + final = await event_queue.dequeue_event() + assert final.kind == "status-update" + assert final.final is True + assert final.status.state == TaskState.completed + assert final.status.message is not None + assert final.status.message.parts[0].root.text == "Sunny" + + request = adapter._rest.agent_api_messages.create_agent_chat_message.await_args + assert request.kwargs["chat_id"] == "room-123" + assert ( + request.kwargs["message"].content == "@Weather Agent What is the weather?" + ) + assert adapter._pending_tasks == {} + + @pytest.mark.asyncio + async def test_execute_fails_when_band_response_times_out(self) -> None: + adapter = A2AGatewayAdapter( + config=A2AGatewayAdapterConfig(response_timeout_s=0.01) + ) + adapter._peers = {"weather": make_peer("weather", "Weather Agent")} + + chat_response = MagicMock() + chat_response.data.id = "room-123" + adapter._rest.agent_api_chats.create_agent_chat = AsyncMock( + return_value=chat_response + ) + adapter._rest.agent_api_participants.add_agent_chat_participant = AsyncMock() + adapter._rest.agent_api_messages.create_agent_chat_message = AsyncMock() + adapter._rest.agent_api_events.create_agent_chat_event = AsyncMock() + + context = RequestContext( + request=MessageSendParams(message=make_a2a_message("What is the weather?")) + ) + event_queue = EventQueue() + executor = BandAgentExecutor(adapter, "weather") + + await executor.execute(context, event_queue) + + initial = await event_queue.dequeue_event() + terminal = await event_queue.dequeue_event() + assert initial.kind == "task" + assert terminal.kind == "status-update" + assert terminal.final is True + assert terminal.status.state == TaskState.failed + assert terminal.status.message is not None + assert "Timed out" in terminal.status.message.parts[0].root.text + assert adapter._pending_tasks == {} + + @pytest.mark.asyncio + async def test_execute_keeps_stream_open_for_updates_until_final_response( + self, + ) -> None: + adapter = A2AGatewayAdapter( + config=A2AGatewayAdapterConfig(response_timeout_s=1) + ) + adapter._peers = {"weather": make_peer("weather", "Weather Agent")} + + chat_response = MagicMock() + chat_response.data.id = "room-123" + adapter._rest.agent_api_chats.create_agent_chat = AsyncMock( + return_value=chat_response + ) + adapter._rest.agent_api_participants.add_agent_chat_participant = AsyncMock() + message_sent = asyncio.Event() + + async def send_message(**_kwargs) -> None: + message_sent.set() + + adapter._rest.agent_api_messages.create_agent_chat_message = AsyncMock( + side_effect=send_message + ) + adapter._rest.agent_api_events.create_agent_chat_event = AsyncMock() + + context = RequestContext( + request=MessageSendParams(message=make_a2a_message("What is the weather?")) + ) + event_queue = EventQueue() + execution = asyncio.create_task( + BandAgentExecutor(adapter, "weather").execute(context, event_queue) + ) + + await asyncio.wait_for(message_sent.wait(), timeout=1) + await event_queue.dequeue_event() + + await adapter.on_message( + make_platform_message( + "Checking the forecast", + room_id="room-123", + message_type="thought", + ), + FakeAgentTools(), + GatewaySessionState(), + None, + None, + is_session_bootstrap=False, + room_id="room-123", + ) + update = await event_queue.dequeue_event() + assert update.final is False + assert not execution.done() + + await adapter.on_message( + make_platform_message("Sunny", room_id="room-123"), + FakeAgentTools(), + GatewaySessionState(), + None, + None, + is_session_bootstrap=False, + room_id="room-123", + ) + await asyncio.wait_for(execution, timeout=1) + + final = await event_queue.dequeue_event() + assert final.final is True + assert final.status.state == TaskState.completed class TestA2AGatewayAdapterRoomManagement: @@ -564,21 +688,24 @@ async def test_on_cleanup_removes_pending_task(self) -> None: from band.integrations.a2a.gateway.types import PendingA2ATask from a2a.types import Task, TaskStatus - sse_queue: asyncio.Queue = asyncio.Queue() + event_queue = EventQueue() task = Task( id="task-123", context_id="ctx-123", status=TaskStatus(state=TaskState.working), ) - adapter._pending_tasks["room-123"] = PendingA2ATask( + pending = PendingA2ATask( task=task, - sse_queue=sse_queue, + event_queue=event_queue, peer_id="weather", + done=asyncio.Event(), ) + adapter._pending_tasks["room-123"] = pending await adapter.on_cleanup("room-123") assert "room-123" not in adapter._pending_tasks + assert pending.done.is_set() @pytest.mark.asyncio async def test_stop_stops_server(self) -> None: diff --git a/tests/integrations/a2a/gateway/test_server.py b/tests/integrations/a2a/gateway/test_server.py index ac1247d65..78cd7e1d4 100644 --- a/tests/integrations/a2a/gateway/test_server.py +++ b/tests/integrations/a2a/gateway/test_server.py @@ -1,829 +1,134 @@ -"""Tests for GatewayServer.""" +"""ASGI-level tests for the official A2A gateway routes.""" from __future__ import annotations -from collections.abc import AsyncIterator -from unittest.mock import AsyncMock +from uuid import uuid4 -import pytest +from a2a.server.agent_execution import AgentExecutor, RequestContext +from a2a.server.events import EventQueue from a2a.types import ( - Message as A2AMessage, + AgentCard, TaskState, TaskStatus, TaskStatusUpdateEvent, ) from starlette.testclient import TestClient +from a2a.utils import new_task from band.integrations.a2a.gateway.server import GatewayServer -from band_rest import Peer - - -def make_peer(peer_id: str, name: str, description: str = "") -> Peer: - """Create a mock Peer object.""" - return Peer( - id=peer_id, - name=name, - type="Agent", - description=description, - handle=f"test/{name.lower().replace(' ', '-')}", - is_contact=False, - source="registry", - ) - - -class TestGatewayServerInit: - """Tests for GatewayServer initialization.""" - - def test_init_stores_config(self) -> None: - """Should store configuration.""" - peer = make_peer("uuid-weather", "Weather Agent") - peers = {"weather-agent": peer} - peers_by_uuid = {"uuid-weather": peer} - on_request = AsyncMock() - - server = GatewayServer( - peers=peers, - peers_by_uuid=peers_by_uuid, - gateway_url="http://localhost:9000", - port=9000, - on_request=on_request, - ) - - assert server.peers == peers - assert server.peers_by_uuid == peers_by_uuid - assert server.gateway_url == "http://localhost:9000" - assert server.port == 9000 - assert server.on_request == on_request - - def test_init_defaults(self) -> None: - """Should start with no app or server task.""" - server = GatewayServer( - peers={}, - peers_by_uuid={}, - gateway_url="http://localhost:10000", - port=10000, - on_request=AsyncMock(), - ) - - assert server._app is None - assert server._server_task is None - - -class TestGatewayServerBuildApp: - """Tests for GatewayServer._build_app().""" - - def test_build_app_creates_starlette_app(self) -> None: - """Should create a Starlette application.""" - peer = make_peer("uuid-weather", "Weather Agent") - server = GatewayServer( - peers={"weather-agent": peer}, - peers_by_uuid={"uuid-weather": peer}, - gateway_url="http://localhost:10000", - port=10000, - on_request=AsyncMock(), - ) - - app = server._build_app() - - from starlette.applications import Starlette - - assert isinstance(app, Starlette) - - def test_build_app_has_agent_card_route(self) -> None: - """Should have agent card discovery route.""" - peer = make_peer("uuid-weather", "Weather Agent") - server = GatewayServer( - peers={"weather-agent": peer}, - peers_by_uuid={"uuid-weather": peer}, - gateway_url="http://localhost:10000", - port=10000, - on_request=AsyncMock(), - ) - - app = server._build_app() - - # Check routes - route_paths = [r.path for r in app.routes] - assert "/agents/{peer_id}/.well-known/agent.json" in route_paths - - def test_build_app_has_message_stream_route(self) -> None: - """Should have message streaming route.""" - peer = make_peer("uuid-weather", "Weather Agent") - server = GatewayServer( - peers={"weather-agent": peer}, - peers_by_uuid={"uuid-weather": peer}, - gateway_url="http://localhost:10000", - port=10000, - on_request=AsyncMock(), - ) - - app = server._build_app() - - # Check routes - route_paths = [r.path for r in app.routes] - assert "/agents/{peer_id}/v1/message:stream" in route_paths - - -class TestGatewayServerAgentCard: - """Tests for agent card endpoint.""" - - @pytest.fixture - def server_with_peers(self) -> GatewayServer: - """Create server with test peers.""" - weather = make_peer("uuid-weather", "Weather Agent", "Gets weather info") - servicenow = make_peer("uuid-servicenow", "ServiceNow Agent", "Creates tickets") - peers = { - "weather-agent": weather, - "servicenow-agent": servicenow, - } - peers_by_uuid = { - "uuid-weather": weather, - "uuid-servicenow": servicenow, - } - return GatewayServer( - peers=peers, - peers_by_uuid=peers_by_uuid, - gateway_url="http://localhost:10000", - port=10000, - on_request=AsyncMock(), - ) - - def test_agent_card_returns_valid_json( - self, server_with_peers: GatewayServer - ) -> None: - """Should return valid AgentCard JSON for known peer (by slug).""" - app = server_with_peers._build_app() - client = TestClient(app) - - response = client.get("/agents/weather-agent/.well-known/agent.json") - - assert response.status_code == 200 - data = response.json() - assert data["name"] == "Weather Agent" - assert data["description"] == "Gets weather info" - assert data["url"] == "http://localhost:10000/agents/weather-agent" - - def test_agent_card_includes_capabilities( - self, server_with_peers: GatewayServer - ) -> None: - """Should include streaming capability.""" - app = server_with_peers._build_app() - client = TestClient(app) - - response = client.get("/agents/weather-agent/.well-known/agent.json") - - data = response.json() - assert data["capabilities"]["streaming"] is True - - def test_agent_card_includes_skills(self, server_with_peers: GatewayServer) -> None: - """Should include default skill.""" - app = server_with_peers._build_app() - client = TestClient(app) - - response = client.get("/agents/weather-agent/.well-known/agent.json") - - data = response.json() - assert len(data["skills"]) == 1 - assert data["skills"][0]["name"] == "Weather Agent" - assert "band" in data["skills"][0]["tags"] - - def test_agent_card_returns_404_for_unknown_peer( - self, server_with_peers: GatewayServer - ) -> None: - """Should return 404 for unknown peer.""" - app = server_with_peers._build_app() - client = TestClient(app) - - response = client.get("/agents/unknown/.well-known/agent.json") - - assert response.status_code == 404 - assert response.json()["error"] == "Not found" - - def test_agent_card_resolves_by_uuid( - self, server_with_peers: GatewayServer - ) -> None: - """Should resolve agent by UUID as fallback.""" - app = server_with_peers._build_app() - client = TestClient(app) - - response = client.get("/agents/uuid-weather/.well-known/agent.json") - - assert response.status_code == 200 - data = response.json() - assert data["name"] == "Weather Agent" - # URL should use the slug, not the UUID - assert data["url"] == "http://localhost:10000/agents/weather-agent" - - -class TestGatewayServerMessageStream: - """Tests for message streaming endpoint.""" - - @pytest.fixture - def server_with_callback(self) -> tuple[GatewayServer, AsyncMock]: - """Create server with mock callback.""" - callback = AsyncMock() - peer = make_peer("uuid-weather", "Weather Agent") - peers = {"weather-agent": peer} - peers_by_uuid = {"uuid-weather": peer} - server = GatewayServer( - peers=peers, - peers_by_uuid=peers_by_uuid, - gateway_url="http://localhost:10000", - port=10000, - on_request=callback, - ) - return server, callback - - def test_message_stream_returns_404_for_unknown_peer( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """Should return 404 for unknown peer.""" - server, _ = server_with_callback - app = server._build_app() - client = TestClient(app) - - response = client.post( - "/agents/unknown/v1/message:stream", - json={ - "role": "user", - "messageId": "msg-123", - "parts": [{"text": "Hello"}], - }, - ) - - assert response.status_code == 404 - - def test_message_stream_calls_on_request( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """Should call on_request callback with slug.""" - server, _ = server_with_callback - - # Track calls manually - calls: list[tuple[str, A2AMessage]] = [] - - async def mock_callback( - peer_id: str, message: A2AMessage - ) -> AsyncIterator[TaskStatusUpdateEvent]: - calls.append((peer_id, message)) - yield TaskStatusUpdateEvent( - taskId="task-123", - contextId="ctx-123", +from tests.integrations.a2a.gateway.fixtures import make_peer + + +class FakeExecutor(AgentExecutor): + async def execute(self, context: RequestContext, event_queue: EventQueue) -> None: + task = context.current_task + if task is None: + task = new_task(context.message) + await event_queue.enqueue_event(task) + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id=task.id, + context_id=task.context_id, status=TaskStatus(state=TaskState.completed), final=True, ) + ) - server.on_request = mock_callback - app = server._build_app() - client = TestClient(app) - - with client: - client.post( - "/agents/weather-agent/v1/message:stream", - json={ - "role": "user", - "messageId": "msg-123", - "parts": [{"kind": "text", "text": "What is the weather?"}], - }, - ) - - # Callback should have been called with slug - assert len(calls) == 1 - assert calls[0][0] == "weather-agent" # slug - assert isinstance(calls[0][1], A2AMessage) # message - - def test_message_stream_returns_sse_content_type( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """Should return SSE content type.""" - server, _ = server_with_callback + async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None: + raise NotImplementedError - async def mock_callback( - peer_id: str, message: A2AMessage - ) -> AsyncIterator[TaskStatusUpdateEvent]: - yield TaskStatusUpdateEvent( - taskId="task-123", - contextId="ctx-123", - status=TaskStatus(state=TaskState.completed), - final=True, - ) - server.on_request = mock_callback - app = server._build_app() - client = TestClient(app) +def build_server() -> GatewayServer: + peer = make_peer("uuid-weather", "Weather Agent", "Gets weather info") + return GatewayServer( + peers={"weather-agent": peer}, + peers_by_uuid={peer.id: peer}, + gateway_url="http://localhost:10000", + port=10000, + executor_factory=lambda _slug: FakeExecutor(), + ) - with client: - response = client.post( - "/agents/weather-agent/v1/message:stream", - json={ - "role": "user", - "messageId": "msg-123", - "parts": [{"kind": "text", "text": "Hello"}], - }, - ) - assert "text/event-stream" in response.headers["content-type"] +def test_agent_card_is_served_by_upstream_handler() -> None: + response = TestClient(build_server()._build_app()).get( + "/agents/weather-agent/.well-known/agent.json" + ) - def test_message_stream_yields_events_as_json( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """Should yield events as SSE-formatted JSON.""" - server, _ = server_with_callback + assert response.status_code == 200 + card = AgentCard.model_validate(response.json()) + assert card.name == "Weather Agent" + assert card.additional_interfaces is not None + assert card.additional_interfaces[0].transport == "JSONRPC" + assert card.additional_interfaces[0].url.endswith("/agents/weather-agent") - async def mock_callback( - peer_id: str, message: A2AMessage - ) -> AsyncIterator[TaskStatusUpdateEvent]: - yield TaskStatusUpdateEvent( - taskId="task-123", - contextId="ctx-123", - status=TaskStatus(state=TaskState.completed), - final=True, - ) - - server.on_request = mock_callback - app = server._build_app() - client = TestClient(app) - with client: - response = client.post( - "/agents/weather-agent/v1/message:stream", - json={ - "role": "user", - "messageId": "msg-123", - "parts": [{"kind": "text", "text": "Hello"}], - }, - ) +def test_peers_listing_remains_gateway_owned() -> None: + response = TestClient(build_server()._build_app()).get("/peers") - # Response should be SSE formatted - content = response.text - assert content.startswith("data: ") - assert '"taskId":"task-123"' in content - assert '"final":true' in content + assert response.status_code == 200 + assert response.json()["peers"] == [ + { + "slug": "weather-agent", + "id": "uuid-weather", + "name": "Weather Agent", + "description": "Gets weather info", + } + ] - def test_message_stream_returns_400_for_invalid_payload( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """Legacy stream endpoint should reject malformed message payloads.""" - server, _ = server_with_callback - app = server._build_app() - client = TestClient(app) - response = client.post( - "/agents/weather-agent/v1/message:stream", - json={"jsonrpc": "2.0", "method": "message/send"}, - ) +def test_unknown_peer_is_not_resolved_by_a2a_routes() -> None: + response = TestClient(build_server()._build_app()).get( + "/agents/missing/.well-known/agent.json" + ) + assert response.status_code == 404 - assert response.status_code == 400 - data = response.json() - assert data["error"] == "Invalid A2A message payload" - assert data["details"] - def test_message_stream_handles_multiple_events( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """Should yield multiple events.""" - server, _ = server_with_callback +def test_jsonrpc_method_errors_are_upstream_owned() -> None: + response = TestClient(build_server()._build_app()).post( + "/agents/weather-agent", + json={"jsonrpc": "2.0", "id": str(uuid4()), "method": "missing", "params": {}}, + ) - async def mock_callback( - peer_id: str, message: A2AMessage - ) -> AsyncIterator[TaskStatusUpdateEvent]: - yield TaskStatusUpdateEvent( - taskId="task-123", - contextId="ctx-123", - status=TaskStatus(state=TaskState.working), - final=False, - ) - yield TaskStatusUpdateEvent( - taskId="task-123", - contextId="ctx-123", - status=TaskStatus(state=TaskState.completed), - final=True, - ) + assert response.status_code == 200 + assert response.json()["error"]["code"] == -32601 - server.on_request = mock_callback - app = server._build_app() - client = TestClient(app) - with client: - response = client.post( - "/agents/weather-agent/v1/message:stream", - json={ +def test_jsonrpc_send_runs_through_official_handler_and_executor() -> None: + response = TestClient(build_server()._build_app()).post( + "/agents/weather-agent", + json={ + "jsonrpc": "2.0", + "id": "request-1", + "method": "message/send", + "params": { + "message": { "role": "user", - "messageId": "msg-123", + "messageId": "message-1", "parts": [{"kind": "text", "text": "Hello"}], - }, - ) - - # Should have two events - content = response.text - events = [line for line in content.split("\n") if line.startswith("data: ")] - assert len(events) == 2 - assert '"working"' in events[0] - assert '"completed"' in events[1] - - -class TestGatewayServerJsonRpc: - """Tests for JSON-RPC endpoint.""" - - @pytest.fixture - def server_with_callback(self) -> tuple[GatewayServer, AsyncMock]: - """Create server with mock callback.""" - callback = AsyncMock() - peer = make_peer("uuid-weather", "Weather Agent") - peers = {"weather-agent": peer} - peers_by_uuid = {"uuid-weather": peer} - server = GatewayServer( - peers=peers, - peers_by_uuid=peers_by_uuid, - gateway_url="http://localhost:10000", - port=10000, - on_request=callback, - ) - return server, callback - - def test_jsonrpc_returns_404_for_unknown_peer( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """Should return 404 JSON-RPC error for unknown peer.""" - server, _ = server_with_callback - app = server._build_app() - client = TestClient(app) - - response = client.post( - "/agents/unknown", - json={ - "jsonrpc": "2.0", - "id": "req-123", - "method": "message/send", - "params": {"message": {"role": "user", "parts": []}}, - }, - ) - - assert response.status_code == 404 - data = response.json() - assert data["jsonrpc"] == "2.0" - assert data["error"]["code"] == -32001 - assert "not found" in data["error"]["message"].lower() - - def test_jsonrpc_returns_error_for_unknown_method( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """Should return error for unknown JSON-RPC method.""" - server, _ = server_with_callback - app = server._build_app() - client = TestClient(app) - - response = client.post( - "/agents/weather-agent", - json={ - "jsonrpc": "2.0", - "id": "req-123", - "method": "unknown/method", - "params": {}, - }, - ) - - assert response.status_code == 400 - data = response.json() - assert data["jsonrpc"] == "2.0" - assert data["error"]["code"] == -32601 - assert "unknown/method" in data["error"]["message"] - assert data["id"] == "req-123" - - def test_jsonrpc_send_returns_task_result( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """Should return Task result for message/send.""" - server, _ = server_with_callback - - async def mock_callback( - peer_id: str, message: A2AMessage - ) -> AsyncIterator[TaskStatusUpdateEvent]: - yield TaskStatusUpdateEvent( - taskId="task-123", - contextId="ctx-123", - status=TaskStatus(state=TaskState.completed), - final=True, - ) - - server.on_request = mock_callback - app = server._build_app() - client = TestClient(app) - - with client: - response = client.post( - "/agents/weather-agent", - json={ - "jsonrpc": "2.0", - "id": "req-123", - "method": "message/send", - "params": { - "message": { - "role": "user", - "messageId": "msg-123", - "parts": [{"kind": "text", "text": "Hello"}], - } - }, - }, - ) - - assert response.status_code == 200 - data = response.json() - assert data["jsonrpc"] == "2.0" - assert data["id"] == "req-123" - assert data["result"]["id"] == "task-123" - assert data["result"]["contextId"] == "ctx-123" - - def test_jsonrpc_stream_returns_sse( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """Should return SSE response for message/stream.""" - server, _ = server_with_callback - - async def mock_callback( - peer_id: str, message: A2AMessage - ) -> AsyncIterator[TaskStatusUpdateEvent]: - yield TaskStatusUpdateEvent( - taskId="task-123", - contextId="ctx-123", - status=TaskStatus(state=TaskState.completed), - final=True, - ) - - server.on_request = mock_callback - app = server._build_app() - client = TestClient(app) - - with client: - response = client.post( - "/agents/weather-agent", - json={ - "jsonrpc": "2.0", - "id": "req-123", - "method": "message/stream", - "params": { - "message": { - "role": "user", - "messageId": "msg-123", - "parts": [{"kind": "text", "text": "Hello"}], - } - }, - }, - ) - - # Should be SSE format with JSON-RPC wrapped events - assert "text/event-stream" in response.headers["content-type"] - content = response.text - assert content.startswith("data: ") - assert '"jsonrpc":' in content - assert "2.0" in content - assert '"id":' in content - assert "req-123" in content - assert '"taskId":' in content - assert "task-123" in content - - def test_jsonrpc_send_returns_invalid_params_for_bad_message( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """JSON-RPC send should return an invalid params error for bad message data.""" - server, _ = server_with_callback - app = server._build_app() - client = TestClient(app) - - response = client.post( - "/agents/weather-agent", - json={ - "jsonrpc": "2.0", - "id": "req-invalid", - "method": "message/send", - "params": {"message": {"parts": []}}, - }, - ) - - assert response.status_code == 400 - data = response.json() - assert data["jsonrpc"] == "2.0" - assert data["id"] == "req-invalid" - assert data["error"]["code"] == -32602 - assert data["error"]["message"] == "Invalid params" - assert data["error"]["data"] - - def test_jsonrpc_stream_returns_invalid_params_for_bad_message( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """JSON-RPC stream should surface invalid params as SSE instead of crashing.""" - server, _ = server_with_callback - app = server._build_app() - client = TestClient(app) - - response = client.post( - "/agents/weather-agent", - json={ - "jsonrpc": "2.0", - "id": "req-invalid-stream", - "method": "message/stream", - "params": {"message": {"parts": []}}, + } }, - ) - - assert response.status_code == 400 - assert "text/event-stream" in response.headers["content-type"] - content = response.text - assert '"code": -32602' in content - assert "req-invalid-stream" in content - - def test_jsonrpc_resolves_by_uuid( - self, server_with_callback: tuple[GatewayServer, AsyncMock] - ) -> None: - """Should resolve peer by UUID fallback.""" - server, _ = server_with_callback - - async def mock_callback( - peer_id: str, message: A2AMessage - ) -> AsyncIterator[TaskStatusUpdateEvent]: - yield TaskStatusUpdateEvent( - taskId="task-456", - contextId="ctx-456", - status=TaskStatus(state=TaskState.completed), - final=True, - ) - - server.on_request = mock_callback - app = server._build_app() - client = TestClient(app) - - with client: - response = client.post( - "/agents/uuid-weather", # Using UUID instead of slug - json={ - "jsonrpc": "2.0", - "id": "req-456", - "method": "message/send", - "params": { - "message": { - "role": "user", - "messageId": "msg-456", - "parts": [{"kind": "text", "text": "Hello"}], - } - }, - }, - ) - - assert response.status_code == 200 - data = response.json() - assert data["result"]["id"] == "task-456" - - -class TestGatewayHttpContextRouting: - """Tests that gateway HTTP endpoint routes same contextId to same room.""" - - @pytest.fixture - def server_with_context_tracking(self) -> tuple[GatewayServer, dict]: - """Create server that tracks context_id from requests.""" - weather = make_peer("uuid-weather", "Weather Agent") - peers = {"weather-agent": weather} - peers_by_uuid = {"uuid-weather": weather} - - # Track which context_ids are received - context_ids_received: list[str | None] = [] - - async def mock_on_request( - peer_id: str, message: A2AMessage - ) -> AsyncIterator[TaskStatusUpdateEvent]: - ctx = message.context_id - context_ids_received.append(ctx) - - yield TaskStatusUpdateEvent( - taskId="task-1", - contextId=ctx or "generated-ctx", - status=TaskStatus(state=TaskState.completed), - final=True, - ) - - server = GatewayServer( - peers=peers, - peers_by_uuid=peers_by_uuid, - gateway_url="http://localhost:10000", - port=10000, - on_request=mock_on_request, - ) - return server, context_ids_received - - def test_json_rpc_same_context_id_twice_tracked( - self, server_with_context_tracking: tuple[GatewayServer, list] - ) -> None: - """JSON-RPC: Same contextId twice should be received consistently.""" - server, context_ids = server_with_context_tracking - app = server._build_app() - client = TestClient(app) - - # First request with contextId - response1 = client.post( - "/agents/weather-agent", - json={ - "jsonrpc": "2.0", - "id": "1", - "method": "message/send", - "params": { - "message": { - "role": "user", - "messageId": "msg-1", - "contextId": "ctx-shared-123", - "parts": [{"kind": "text", "text": "First message"}], - } - }, - }, - ) - - # Second request with SAME contextId - response2 = client.post( - "/agents/weather-agent", - json={ - "jsonrpc": "2.0", - "id": "2", - "method": "message/send", - "params": { - "message": { - "role": "user", - "messageId": "msg-2", - "contextId": "ctx-shared-123", - "parts": [{"kind": "text", "text": "Second message"}], - } - }, - }, - ) - - assert response1.status_code == 200 - assert response2.status_code == 200 - # Both requests should have same context_id - assert len(context_ids) == 2 - assert context_ids[0] == "ctx-shared-123" - assert context_ids[1] == "ctx-shared-123" - - def test_json_rpc_different_context_ids_tracked( - self, server_with_context_tracking: tuple[GatewayServer, list] - ) -> None: - """JSON-RPC: Different contextIds should be tracked separately.""" - server, context_ids = server_with_context_tracking - app = server._build_app() - client = TestClient(app) - - client.post( - "/agents/weather-agent", - json={ - "jsonrpc": "2.0", - "id": "1", - "method": "message/send", - "params": { - "message": { - "role": "user", - "messageId": "m1", - "contextId": "ctx-a", - "parts": [{"kind": "text", "text": "A"}], - } - }, - }, - ) - client.post( - "/agents/weather-agent", - json={ - "jsonrpc": "2.0", - "id": "2", - "method": "message/send", - "params": { - "message": { - "role": "user", - "messageId": "m2", - "contextId": "ctx-b", - "parts": [{"kind": "text", "text": "B"}], - } - }, - }, - ) - - # Two different contexts received - assert len(context_ids) == 2 - assert "ctx-a" in context_ids - assert "ctx-b" in context_ids - - def test_legacy_stream_endpoint_receives_context_id( - self, server_with_context_tracking: tuple[GatewayServer, list] - ) -> None: - """Legacy REST stream endpoint should receive contextId.""" - server, context_ids = server_with_context_tracking - app = server._build_app() - client = TestClient(app) + }, + ) - response = client.post( - "/agents/weather-agent/v1/message:stream", - json={ - "role": "user", - "messageId": "msg-123", - "contextId": "ctx-legacy-456", - "parts": [{"kind": "text", "text": "Hello from legacy"}], - }, - ) + assert response.status_code == 200 + body = response.json() + assert body["id"] == "request-1" + assert body["result"]["status"]["state"] == "completed" + + +def test_rest_stream_runs_through_upstream_adapter() -> None: + response = TestClient(build_server()._build_app()).post( + "/agents/weather-agent/v1/message:stream", + json={ + "message": { + "role": "ROLE_USER", + "messageId": "message-1", + "content": [{"text": "Hello"}], + } + }, + ) - assert response.status_code == 200 - assert len(context_ids) == 1 - assert context_ids[0] == "ctx-legacy-456" + assert response.status_code == 200 + assert "text/event-stream" in response.headers["content-type"] + assert '"task":' in response.text + assert '"state": "TASK_STATE_COMPLETED"' in response.text diff --git a/tests/integrations/a2a/gateway/test_types.py b/tests/integrations/a2a/gateway/test_types.py deleted file mode 100644 index c849a724a..000000000 --- a/tests/integrations/a2a/gateway/test_types.py +++ /dev/null @@ -1,125 +0,0 @@ -"""Tests for A2A Gateway types.""" - -from __future__ import annotations - -import asyncio -from uuid import uuid4 - -import pytest -from a2a.types import Task, TaskState, TaskStatus, TaskStatusUpdateEvent - -from band.integrations.a2a.gateway.types import GatewaySessionState, PendingA2ATask - - -class TestGatewaySessionState: - """Tests for GatewaySessionState dataclass.""" - - def test_init_default_empty_dicts(self) -> None: - """Empty state should have empty dicts.""" - state = GatewaySessionState() - assert state.context_to_room == {} - assert state.room_participants == {} - - def test_init_with_context_mapping(self) -> None: - """Can initialize with context_to_room mapping.""" - state = GatewaySessionState( - context_to_room={"ctx-1": "room-1", "ctx-2": "room-2"} - ) - assert state.context_to_room["ctx-1"] == "room-1" - assert state.context_to_room["ctx-2"] == "room-2" - assert state.room_participants == {} - - def test_init_with_participants(self) -> None: - """Can initialize with room_participants mapping.""" - state = GatewaySessionState( - room_participants={"room-1": {"peer-1", "peer-2"}, "room-2": {"peer-3"}} - ) - assert state.context_to_room == {} - assert state.room_participants["room-1"] == {"peer-1", "peer-2"} - assert state.room_participants["room-2"] == {"peer-3"} - - def test_init_with_both(self) -> None: - """Can initialize with both context_to_room and room_participants.""" - state = GatewaySessionState( - context_to_room={"ctx-1": "room-1"}, - room_participants={"room-1": {"peer-1"}}, - ) - assert state.context_to_room == {"ctx-1": "room-1"} - assert state.room_participants == {"room-1": {"peer-1"}} - - def test_context_to_room_dict_behavior(self) -> None: - """context_to_room behaves like a normal dict.""" - state = GatewaySessionState() - state.context_to_room["new-ctx"] = "new-room" - assert state.context_to_room["new-ctx"] == "new-room" - - def test_room_participants_set_behavior(self) -> None: - """room_participants values behave like sets.""" - state = GatewaySessionState(room_participants={"room-1": {"peer-1"}}) - state.room_participants["room-1"].add("peer-2") - assert state.room_participants["room-1"] == {"peer-1", "peer-2"} - - -class TestPendingA2ATask: - """Tests for PendingA2ATask dataclass.""" - - @pytest.fixture - def sample_task(self) -> Task: - """Create a sample A2A Task for testing.""" - return Task( - id=str(uuid4()), - context_id=str(uuid4()), - status=TaskStatus(state=TaskState.working), - ) - - @pytest.fixture - def sample_queue(self) -> asyncio.Queue[TaskStatusUpdateEvent]: - """Create a sample asyncio Queue for testing.""" - return asyncio.Queue() - - def test_init_with_task_and_queue( - self, sample_task: Task, sample_queue: asyncio.Queue[TaskStatusUpdateEvent] - ) -> None: - """Can initialize with task and queue.""" - pending = PendingA2ATask( - task=sample_task, - sse_queue=sample_queue, - peer_id="weather", - ) - assert pending.task == sample_task - assert pending.sse_queue == sample_queue - - def test_init_stores_peer_id( - self, sample_task: Task, sample_queue: asyncio.Queue[TaskStatusUpdateEvent] - ) -> None: - """Stores peer_id for correlation.""" - pending = PendingA2ATask( - task=sample_task, - sse_queue=sample_queue, - peer_id="servicenow", - ) - assert pending.peer_id == "servicenow" - - def test_task_property_access( - self, sample_task: Task, sample_queue: asyncio.Queue[TaskStatusUpdateEvent] - ) -> None: - """Can access task properties.""" - pending = PendingA2ATask( - task=sample_task, - sse_queue=sample_queue, - peer_id="weather", - ) - assert pending.task.id == sample_task.id - assert pending.task.context_id == sample_task.context_id - assert pending.task.status.state == TaskState.working - - def test_sse_queue_is_async_queue( - self, sample_task: Task, sample_queue: asyncio.Queue[TaskStatusUpdateEvent] - ) -> None: - """Queue instance is asyncio.Queue.""" - pending = PendingA2ATask( - task=sample_task, - sse_queue=sample_queue, - peer_id="weather", - ) - assert isinstance(pending.sse_queue, asyncio.Queue) diff --git a/uv.lock b/uv.lock index 88aba8ca7..654dffc86 100644 --- a/uv.lock +++ b/uv.lock @@ -100,19 +100,22 @@ conflicts = [[ [[package]] name = "a2a-sdk" -version = "0.3.26" +version = "1.1.2" source = { registry = "https://pypi.org/simple" } dependencies = [ + { name = "culsans", marker = "python_full_version < '3.13' or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-dev') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-pydantic-ai') or (extra == 'extra-8-band-sdk-dev' and extra == 'extra-8-band-sdk-dev-crewai') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-pydantic-ai')" }, { name = "google-api-core" }, + { name = "googleapis-common-protos" }, { name = "httpx" }, - { name = "httpx-sse" }, + { name = "json-rpc" }, + { name = "packaging" }, { name = "protobuf" }, { name = "pydantic", version = "2.11.9", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.14' or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-dev') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-pydantic-ai') or (extra == 'extra-8-band-sdk-dev' and extra == 'extra-8-band-sdk-dev-crewai') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-pydantic-ai')" }, { name = "pydantic", version = "2.12.5", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.14' or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-dev') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-pydantic-ai') or (extra == 'extra-8-band-sdk-dev' and extra == 'extra-8-band-sdk-dev-crewai') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-pydantic-ai')" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/be/97/a6840e01795b182ce751ca165430d46459927cde9bfab838087cbb24aef7/a2a_sdk-0.3.26.tar.gz", hash = "sha256:44068e2d037afbb07ab899267439e9bc7eaa7ac2af94f1e8b239933c993ad52d", size = 274598, upload-time = "2026-04-09T15:21:13.902Z" } +sdist = { url = "https://files.pythonhosted.org/packages/38/cc/59b35c518d8289bd59d20d9d216ca29ccb41c4697eb85971efe41d1adaf3/a2a_sdk-1.1.2.tar.gz", hash = "sha256:f928d8bf9a0dc0a473ee8258bd5226c1b8520a85c70e7f233c3b9b0e30a1bd10", size = 379627, upload-time = "2026-07-22T13:40:51.409Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/dd/d5/51f4ee1bf3b736add42a542d3c8a3fd3fa85f3d36c17972127defc46c26f/a2a_sdk-0.3.26-py3-none-any.whl", hash = "sha256:754e0573f6d33b225c1d8d51f640efa69cbbed7bdfb06ce9c3540ea9f58d4a91", size = 151016, upload-time = "2026-04-09T15:21:12.35Z" }, + { url = "https://files.pythonhosted.org/packages/c1/33/d414a03d5ef5a7d6d09a20dc9f8094e7fb9edb488662ff99c535a2a62cf1/a2a_sdk-1.1.2-py3-none-any.whl", hash = "sha256:eb9a698f526dc45a9ea967d7804e7aa363ff273d2c44fbed699bc68c9399f9c9", size = 245981, upload-time = "2026-07-22T13:40:49.484Z" }, ] [[package]] @@ -321,6 +324,20 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/62/29/2f8418269e46454a26171bfdd6a055d74febf32234e474930f2f60a17145/aiohttp-3.13.5-cp314-cp314t-win_amd64.whl", hash = "sha256:18a2f6c1182c51baa1d28d68fea51513cb2a76612f038853c0ad3c145423d3d9", size = 505441, upload-time = "2026-03-31T22:00:12.791Z" }, ] +[[package]] +name = "aiologic" +version = "0.17.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "sniffio", marker = "python_full_version < '3.13' or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-dev') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-pydantic-ai') or (extra == 'extra-8-band-sdk-dev' and extra == 'extra-8-band-sdk-dev-crewai') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-pydantic-ai')" }, + { name = "typing-extensions", marker = "python_full_version < '3.13' or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-dev') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-pydantic-ai') or (extra == 'extra-8-band-sdk-dev' and extra == 'extra-8-band-sdk-dev-crewai') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-pydantic-ai')" }, + { name = "wrapt", marker = "python_full_version < '3.13' or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-dev') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-pydantic-ai') or (extra == 'extra-8-band-sdk-dev' and extra == 'extra-8-band-sdk-dev-crewai') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-pydantic-ai')" }, +] +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 = "aiopenapi3" version = "0.8.1" @@ -693,10 +710,10 @@ slack = [ [package.metadata] requires-dist = [ - { name = "a2a-sdk", marker = "extra == 'a2a'", specifier = ">=0.3.22" }, - { name = "a2a-sdk", marker = "extra == 'a2a-gateway'", specifier = ">=0.3.22" }, - { name = "a2a-sdk", marker = "extra == 'a2a-gateway-demo'", specifier = ">=0.3.22" }, - { name = "a2a-sdk", marker = "extra == 'dev'", specifier = ">=0.3.22" }, + { name = "a2a-sdk", marker = "extra == 'a2a'", specifier = ">=1.1.2" }, + { name = "a2a-sdk", marker = "extra == 'a2a-gateway'", specifier = ">=1.1.2" }, + { name = "a2a-sdk", marker = "extra == 'a2a-gateway-demo'", specifier = ">=1.1.2" }, + { name = "a2a-sdk", marker = "extra == 'dev'", specifier = ">=1.1.2" }, { name = "agent-client-protocol", marker = "extra == 'acp'", specifier = ">=0.9.0" }, { name = "agent-client-protocol", marker = "extra == 'dev'", specifier = ">=0.9.0" }, { name = "agno", marker = "extra == 'agno'", specifier = ">=2.6.0" }, @@ -1727,6 +1744,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/aa/50/a9caea39ad19c431c1a3f8a31114df65b260cdfe67786b6c7e7c040c4c44/cryptography-49.0.0-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:be9fcb48a55f023493482827d4f459bd263cc20efde64f204b97c123201850c6", size = 3783731, upload-time = "2026-06-12T20:02:43.319Z" }, ] +[[package]] +name = "culsans" +version = "0.11.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "aiologic", marker = "python_full_version < '3.13' or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-dev') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-pydantic-ai') or (extra == 'extra-8-band-sdk-dev' and extra == 'extra-8-band-sdk-dev-crewai') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-pydantic-ai')" }, + { name = "typing-extensions", marker = "python_full_version < '3.13' or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-dev') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-crewai' and extra == 'extra-8-band-sdk-pydantic-ai') or (extra == 'extra-8-band-sdk-dev' and extra == 'extra-8-band-sdk-dev-crewai') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-parlant') or (extra == 'extra-8-band-sdk-dev-crewai' and extra == 'extra-8-band-sdk-pydantic-ai')" }, +] +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]] name = "cycler" version = "0.12.1" @@ -3420,6 +3450,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/f0/9e/2ab68cc0ff030e1ef78329d7b933473d3ad2c7d0e66aede6a7c87f74753c/json_repair-0.25.3-py3-none-any.whl", hash = "sha256:f00b510dd21b31ebe72581bdb07e66381df2883d6f640c89605e482882c12b17", size = 12812, upload-time = "2024-07-10T13:42:16.918Z" }, ] +[[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.10.0" From 9bcebec2d1a53da24abaeaa9db0e973b14060cdb Mon Sep 17 00:00:00 2001 From: Alexander Zaikman Date: Sun, 26 Jul 2026 14:18:28 +0300 Subject: [PATCH 2/7] Migrate A2A integration to SDK 1.1.2 --- README.md | 7 +- examples/a2a_bridge/01_basic_agent.py | 2 +- examples/a2a_bridge/README.md | 6 - examples/a2a_gateway/01_basic_gateway.py | 4 +- examples/a2a_gateway/02_with_demo_agent.py | 39 +- examples/a2a_gateway/README.md | 27 +- .../a2a_gateway/demo_orchestrator/__main__.py | 35 +- .../demo_orchestrator/agent_executor.py | 31 +- .../demo_orchestrator/remote_agent.py | 202 ++-- examples/mixed/03_fact_checker_a2a.py | 50 +- examples/mixed/04_risk_reviewer_a2a.py | 50 +- examples/mixed/README.md | 8 +- examples/run_agent.py | 2 +- src/band/integrations/a2a/adapter.py | 214 ++--- src/band/integrations/a2a/gateway/adapter.py | 90 +- src/band/integrations/a2a/gateway/server.py | 78 +- src/band/integrations/a2a/gateway/types.py | 4 +- src/band/integrations/a2a/protocol.py | 78 ++ .../a2a/gateway/{fixtures.py => helpers.py} | 2 +- .../integrations/a2a/gateway/test_adapter.py | 839 ++++------------ tests/integrations/a2a/gateway/test_server.py | 90 +- tests/integrations/a2a/test_adapter.py | 909 +++--------------- 22 files changed, 904 insertions(+), 1863 deletions(-) create mode 100644 src/band/integrations/a2a/protocol.py rename tests/integrations/a2a/gateway/{fixtures.py => helpers.py} (88%) diff --git a/README.md b/README.md index db44b0da3..4bbff79e8 100644 --- a/README.md +++ b/README.md @@ -569,7 +569,7 @@ import asyncio import os from band import Agent, configure_logging -from band.adapters.a2a_gateway import A2AGatewayAdapter +from band.adapters.a2a_gateway import A2AGatewayAdapter, A2AGatewayAdapterConfig configure_logging() @@ -582,6 +582,8 @@ async def main() -> None: api_key=os.environ["GATEWAY_API_KEY"], gateway_url=gateway_url, port=gateway_port, + # The default is 300 seconds. Use None for no response deadline. + config=A2AGatewayAdapterConfig(response_timeout_s=300), ) agent = Agent.create( @@ -601,6 +603,9 @@ Discovery endpoints include: ```bash curl http://localhost:10000/peers +curl http://localhost:10000/agents/weather-agent/.well-known/agent-card.json + +# Legacy card discovery remains available for older clients: curl http://localhost:10000/agents/weather-agent/.well-known/agent.json ``` diff --git a/examples/a2a_bridge/01_basic_agent.py b/examples/a2a_bridge/01_basic_agent.py index b70784ac8..d920add35 100644 --- a/examples/a2a_bridge/01_basic_agent.py +++ b/examples/a2a_bridge/01_basic_agent.py @@ -25,7 +25,7 @@ python -m app --host localhost --port 10000 2. Verify the agent is running: - curl http://localhost:10000/.well-known/agent.json + curl http://localhost:10000/.well-known/agent-card.json Run with: uv run examples/a2a_bridge/01_basic_agent.py diff --git a/examples/a2a_bridge/README.md b/examples/a2a_bridge/README.md index 758b5feeb..9f298ad2c 100644 --- a/examples/a2a_bridge/README.md +++ b/examples/a2a_bridge/README.md @@ -63,12 +63,6 @@ export A2A_AGENT_URL=http://localhost:10000 Before starting the bridge, confirm the remote agent card is reachable: -```bash -curl http://localhost:10000/.well-known/agent.json -``` - -or: - ```bash curl http://localhost:10000/.well-known/agent-card.json ``` diff --git a/examples/a2a_gateway/01_basic_gateway.py b/examples/a2a_gateway/01_basic_gateway.py index d2b53dc6d..f32637576 100644 --- a/examples/a2a_gateway/01_basic_gateway.py +++ b/examples/a2a_gateway/01_basic_gateway.py @@ -42,7 +42,7 @@ uv run examples/a2a_gateway/01_basic_gateway.py Then remote agents can connect: - - Discovery: GET http://localhost:10000/agents/weather/.well-known/agent.json + - Discovery: GET http://localhost:10000/agents/weather/.well-known/agent-card.json - JSON-RPC: POST http://localhost:10000/agents/weather - Stream: POST http://localhost:10000/agents/weather/v1/message:stream """ @@ -108,7 +108,7 @@ async def main() -> None: logger.info("Starting A2A Gateway on %s...", gateway_url) logger.info("Peers will be exposed at:") logger.info( - " - %s/agents/{peer_id}/.well-known/agent.json (discovery)", gateway_url + " - %s/agents/{peer_id}/.well-known/agent-card.json (discovery)", gateway_url ) logger.info(" - %s/agents/{peer_id}/v1/message:stream (messaging)", gateway_url) logger.info("Waiting for peers to be discovered...") diff --git a/examples/a2a_gateway/02_with_demo_agent.py b/examples/a2a_gateway/02_with_demo_agent.py index 948c70a29..a1bf2c7ec 100644 --- a/examples/a2a_gateway/02_with_demo_agent.py +++ b/examples/a2a_gateway/02_with_demo_agent.py @@ -36,7 +36,7 @@ Test the demo: # Check orchestrator agent card - curl http://localhost:10001/.well-known/agent.json + curl http://localhost:10001/.well-known/agent-card.json # Send a JSON-RPC message to the orchestrator (it will route to gateway peers) curl -X POST http://localhost:10001/ \\ @@ -44,7 +44,7 @@ -d '{ "jsonrpc": "2.0", "id": "1", - "method": "message/send", + "method": "SendMessage", "params": { "message": { "role": "user", @@ -68,14 +68,17 @@ sys.path.insert(0, str(Path(__file__).parent)) import uvicorn -from a2a.server.apps.jsonrpc.starlette_app import A2AStarletteApplication +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.routes.rest_routes import create_rest_routes from a2a.server.request_handlers import DefaultRequestHandler from a2a.server.tasks import ( InMemoryPushNotificationConfigStore, InMemoryTaskStore, ) -from a2a.types import AgentCapabilities, AgentCard, AgentSkill +from a2a.types import AgentCapabilities, AgentCard, AgentInterface, AgentSkill from dotenv import load_dotenv +from starlette.applications import Starlette from demo_orchestrator.agent import OrchestratorAgent from demo_orchestrator.agent_executor import OrchestratorAgentExecutor @@ -175,7 +178,13 @@ def run_orchestrator() -> None: agent_card = AgentCard( name="Demo Orchestrator", description="Routes user requests to Band platform peers via A2A Gateway", - url=f"http://{ORCHESTRATOR_HOST}:{ORCHESTRATOR_PORT}/", + supported_interfaces=[ + AgentInterface( + protocol_binding="JSONRPC", + protocol_version="1.0", + url=f"http://{ORCHESTRATOR_HOST}:{ORCHESTRATOR_PORT}/", + ) + ], version="1.0.0", default_input_modes=OrchestratorAgent.SUPPORTED_CONTENT_TYPES, default_output_modes=OrchestratorAgent.SUPPORTED_CONTENT_TYPES, @@ -187,12 +196,20 @@ def run_orchestrator() -> None: request_handler = DefaultRequestHandler( agent_executor=OrchestratorAgentExecutor(agent), task_store=InMemoryTaskStore(), + agent_card=agent_card, push_config_store=InMemoryPushNotificationConfigStore(), ) - server = A2AStarletteApplication( - agent_card=agent_card, - http_handler=request_handler, + server = Starlette( + routes=( + create_agent_card_routes(agent_card) + + create_jsonrpc_routes( + request_handler, + rpc_url="/", + enable_v0_3_compat=True, + ) + + create_rest_routes(request_handler, enable_v0_3_compat=True) + ) ) logger.info( @@ -202,7 +219,7 @@ def run_orchestrator() -> None: ) # Run uvicorn (blocking) - uvicorn.run(server.build(), host=ORCHESTRATOR_HOST, port=ORCHESTRATOR_PORT) + uvicorn.run(server, host=ORCHESTRATOR_HOST, port=ORCHESTRATOR_PORT) async def main() -> None: @@ -221,7 +238,9 @@ async def main() -> None: ) logger.info("") logger.info("Test with:") - logger.info(" curl http://localhost:%s/.well-known/agent.json", ORCHESTRATOR_PORT) + logger.info( + " curl http://localhost:%s/.well-known/agent-card.json", ORCHESTRATOR_PORT + ) logger.info("") # Run gateway in background, orchestrator in foreground diff --git a/examples/a2a_gateway/README.md b/examples/a2a_gateway/README.md index 94d17efa0..8e267ada0 100644 --- a/examples/a2a_gateway/README.md +++ b/examples/a2a_gateway/README.md @@ -95,7 +95,7 @@ uv run examples/a2a_gateway/01_basic_gateway.py Once it is up, the most useful endpoints are: - `GET http://localhost:10000/peers` -- `GET http://localhost:10000/agents//.well-known/agent.json` +- `GET http://localhost:10000/agents//.well-known/agent-card.json` - `POST http://localhost:10000/agents/` - `POST http://localhost:10000/agents//v1/message:stream` @@ -118,22 +118,24 @@ Pick one peer id from the response. ### 3. Fetch the peer's agent card ```bash -curl http://localhost:10000/agents//.well-known/agent.json +curl http://localhost:10000/agents//.well-known/agent-card.json ``` ### 4. Send a message Use a stable `contextId` so you can test room reuse. -For JSON-RPC clients, send requests to the base agent route: +For A2A 1.0 JSON-RPC clients, send requests to the base agent route with the +`A2A-Version: 1.0` header. The 1.0 method is `SendMessage`: ```bash curl -X POST http://localhost:10000/agents/ \ -H "Content-Type: application/json" \ + -H "A2A-Version: 1.0" \ -d '{ "jsonrpc": "2.0", "id": "1", - "method": "message/send", + "method": "SendMessage", "params": { "message": { "role": "user", @@ -145,16 +147,19 @@ curl -X POST http://localhost:10000/agents/ \ }' ``` -If you want the legacy streaming route, send a raw A2A `Message` body instead: +For the REST streaming binding, use the versioned route and wrap the message +inside a `SendMessageRequest`: ```bash curl -N -X POST http://localhost:10000/agents//v1/message:stream \ -H "Content-Type: application/json" \ -d '{ - "role": "user", - "parts": [{"kind": "text", "text": "Hello from an A2A client"}], - "messageId": "msg-1", - "contextId": "ctx-1" + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "Hello from an A2A client"}], + "messageId": "msg-1", + "contextId": "ctx-1" + } }' ``` @@ -191,7 +196,7 @@ uv run examples/a2a_gateway/02_with_demo_agent.py Then check the orchestrator card: ```bash -curl http://localhost:10001/.well-known/agent.json +curl http://localhost:10001/.well-known/agent-card.json ``` The orchestrator's job is to accept an incoming A2A request and route it to one of the peers exposed by the gateway. @@ -201,7 +206,7 @@ The orchestrator's job is to accept an incoming A2A request and route it to one ### Basic gateway - `/peers` returns real Band peers -- `/.well-known/agent.json` returns a valid card for a peer +- `/.well-known/agent-card.json` returns a valid card for a peer - a message call returns a peer response - the same `contextId` reuses the same room diff --git a/examples/a2a_gateway/demo_orchestrator/__main__.py b/examples/a2a_gateway/demo_orchestrator/__main__.py index 1d5029571..3207410b1 100644 --- a/examples/a2a_gateway/demo_orchestrator/__main__.py +++ b/examples/a2a_gateway/demo_orchestrator/__main__.py @@ -6,8 +6,8 @@ uv run python examples/a2a_gateway/demo_orchestrator/__main__.py --gateway-url http://localhost:10000 This starts an A2A-compliant server that: -1. Exposes itself at /.well-known/agent.json -2. Accepts messages at /v1/message:stream +1. Exposes itself at /.well-known/agent-card.json +2. Accepts messages at /message:stream and / (JSON-RPC) 3. Routes requests to Band peers via the A2A Gateway """ @@ -24,14 +24,17 @@ import click import uvicorn -from a2a.server.apps.jsonrpc.starlette_app import A2AStarletteApplication +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.routes.rest_routes import create_rest_routes from a2a.server.request_handlers import DefaultRequestHandler from a2a.server.tasks import ( InMemoryPushNotificationConfigStore, InMemoryTaskStore, ) -from a2a.types import AgentCapabilities, AgentCard, AgentSkill +from a2a.types import AgentCapabilities, AgentCard, AgentInterface, AgentSkill from dotenv import load_dotenv +from starlette.applications import Starlette from agent import OrchestratorAgent from agent_executor import OrchestratorAgentExecutor @@ -135,7 +138,13 @@ def main(host: str, port: int, gateway_url: str, peers: str, model: str) -> None "Band platform peers via the A2A Gateway. It intelligently " "determines which peer can best handle each request." ), - url=f"http://{host}:{port}/", + supported_interfaces=[ + AgentInterface( + protocol_binding="JSONRPC", + protocol_version="1.0", + url=f"http://{host}:{port}/", + ) + ], version="1.0.0", default_input_modes=OrchestratorAgent.SUPPORTED_CONTENT_TYPES, default_output_modes=OrchestratorAgent.SUPPORTED_CONTENT_TYPES, @@ -149,16 +158,24 @@ def main(host: str, port: int, gateway_url: str, peers: str, model: str) -> None request_handler = DefaultRequestHandler( agent_executor=OrchestratorAgentExecutor(agent), task_store=InMemoryTaskStore(), + agent_card=agent_card, push_config_store=push_config_store, ) - server = A2AStarletteApplication( - agent_card=agent_card, - http_handler=request_handler, + server = Starlette( + routes=( + create_agent_card_routes(agent_card) + + create_jsonrpc_routes( + request_handler, + rpc_url="/", + enable_v0_3_compat=True, + ) + + create_rest_routes(request_handler, enable_v0_3_compat=True) + ) ) # Run server - uvicorn.run(server.build(), host=host, port=port) + uvicorn.run(server, host=host, port=port) except Exception as e: logger.error("Error starting server: %s", e) diff --git a/examples/a2a_gateway/demo_orchestrator/agent_executor.py b/examples/a2a_gateway/demo_orchestrator/agent_executor.py index d68505347..07a051a23 100644 --- a/examples/a2a_gateway/demo_orchestrator/agent_executor.py +++ b/examples/a2a_gateway/demo_orchestrator/agent_executor.py @@ -15,11 +15,9 @@ InternalError, Part, TaskState, - TextPart, UnsupportedOperationError, ) -from a2a.utils import new_agent_text_message, new_task -from a2a.utils.errors import ServerError +from band.integrations.a2a.protocol import new_task, text_message try: from .agent import OrchestratorAgent @@ -59,7 +57,9 @@ async def execute( task = context.current_task if not task: - task = new_task(context.message) # type: ignore + if context.message is None: + raise ValueError("A2A request is missing its message") + task = new_task(context.message) await event_queue.enqueue_event(task) updater = TaskUpdater(event_queue, task.id, task.context_id) @@ -73,29 +73,28 @@ async def execute( if not is_task_complete and not require_user_input: # Working status update await updater.update_status( - TaskState.working, - new_agent_text_message( + TaskState.TASK_STATE_WORKING, + text_message( content, - task.context_id, - task.id, + context_id=task.context_id, + task_id=task.id, ), ) elif require_user_input: # Need more input from user await updater.update_status( - TaskState.input_required, - new_agent_text_message( + TaskState.TASK_STATE_INPUT_REQUIRED, + text_message( content, - task.context_id, - task.id, + context_id=task.context_id, + task_id=task.id, ), - final=True, ) break else: # Task complete - add artifact and finish await updater.add_artifact( - [Part(root=TextPart(text=content))], + [Part(text=content)], name="orchestrator_result", ) await updater.complete() @@ -103,7 +102,7 @@ async def execute( except Exception as e: logger.error("Error executing orchestrator agent: %s", e) - raise ServerError(error=InternalError()) from e + raise InternalError() from e async def cancel( self, @@ -119,4 +118,4 @@ async def cancel( Raises: ServerError: Cancellation not supported """ - raise ServerError(error=UnsupportedOperationError()) + raise UnsupportedOperationError() diff --git a/examples/a2a_gateway/demo_orchestrator/remote_agent.py b/examples/a2a_gateway/demo_orchestrator/remote_agent.py index e1e88508e..2b520155a 100644 --- a/examples/a2a_gateway/demo_orchestrator/remote_agent.py +++ b/examples/a2a_gateway/demo_orchestrator/remote_agent.py @@ -1,8 +1,4 @@ -"""Gateway client for calling A2A Gateway peers. - -This module provides a client wrapper for calling Band platform peers -via the A2A Gateway using the standard A2A protocol. -""" +"""Gateway client for calling A2A Gateway peers.""" from __future__ import annotations @@ -10,189 +6,113 @@ from uuid import uuid4 import httpx -from a2a.client import A2ACardResolver, A2AClient -from a2a.types import ( - MessageSendParams, - Part, - SendMessageRequest, - SendMessageSuccessResponse, - TextPart, -) - -logger = logging.getLogger(__name__) +from a2a.client import ClientConfig, ClientFactory +from a2a.types import Message, Part, Role, SendMessageRequest, Task +from band.integrations.a2a.protocol import text_from_message -def extract_text_from_parts(parts: list[Part]) -> str: - """Extract text content from A2A message parts.""" - texts = [] - for part in parts: - if isinstance(part.root, TextPart): - texts.append(part.root.text) - return " ".join(texts) +logger = logging.getLogger(__name__) class GatewayClient: - """Client for calling A2A Gateway peers. - - This client wraps the A2A SDK's client to provide a simple interface - for calling peers exposed by the A2A Gateway. - - Example: - client = GatewayClient("http://localhost:10000") - response = await client.call_peer("weather", "What's the weather in NYC?") - """ + """Small client wrapper around the official A2A 1.1 client.""" def __init__(self, gateway_url: str, timeout: float = 60.0): - """Initialize the gateway client. - - Args: - gateway_url: Base URL of the A2A Gateway (e.g., "http://localhost:10000") - timeout: Request timeout in seconds - """ self.gateway_url = gateway_url.rstrip("/") self.timeout = timeout self._http_client: httpx.AsyncClient | None = None async def _get_http_client(self) -> httpx.AsyncClient: - """Get or create the HTTP client.""" if self._http_client is None or self._http_client.is_closed: self._http_client = httpx.AsyncClient(timeout=self.timeout) return self._http_client async def close(self) -> None: - """Close the HTTP client.""" if self._http_client and not self._http_client.is_closed: await self._http_client.aclose() self._http_client = None async def list_peers(self) -> list[dict]: - """Fetch list of available peers from gateway. - - Returns: - List of peer dicts with id, name, description - """ http_client = await self._get_http_client() try: response = await http_client.get(f"{self.gateway_url}/peers") response.raise_for_status() - data = response.json() - return data.get("peers", []) - except Exception as e: - logger.warning("Could not fetch peers from gateway: %s", e) + return response.json().get("peers", []) + except Exception as exc: + logger.warning("Could not fetch peers from gateway: %s", exc) return [] async def discover_peer(self, peer_id: str) -> bool: - """Check if a peer is available via the gateway. - - Args: - peer_id: The ID of the peer to check - - Returns: - True if the peer is available, False otherwise - """ - peer_url = f"{self.gateway_url}/agents/{peer_id}" - http_client = await self._get_http_client() - try: - resolver = A2ACardResolver(http_client, peer_url) - await resolver.get_agent_card() + client = await self._create_client(peer_id) + await client.close() return True - except Exception as e: - logger.debug("Peer %s not available: %s", peer_id, e) + except Exception as exc: + logger.debug("Peer %s not available: %s", peer_id, exc) return False + async def _create_client(self, peer_id: str): + http_client = await self._get_http_client() + factory = ClientFactory(ClientConfig(streaming=True, httpx_client=http_client)) + return await factory.create_from_url(f"{self.gateway_url}/agents/{peer_id}") + async def call_peer( self, peer_id: str, message: str, context_id: str | None = None, ) -> str: - """Call a peer via the A2A Gateway. - - Args: - peer_id: The ID of the peer to call (e.g., "weather", "servicenow") - message: The message to send to the peer - context_id: Optional context ID for conversation continuity - - Returns: - The peer's response text - - Raises: - RuntimeError: If the peer is not available or the call fails - """ - peer_url = f"{self.gateway_url}/agents/{peer_id}" - http_client = await self._get_http_client() - - logger.info("Calling peer '%s' via gateway: %s", peer_id, peer_url) - + logger.info("Calling peer '%s' via gateway", peer_id) try: - # Resolve agent card - resolver = A2ACardResolver(http_client, peer_url) - card = await resolver.get_agent_card() - logger.debug("Resolved agent card for peer '%s': %s", peer_id, card.name) - - # Create A2A client - client = A2AClient(http_client, card, url=peer_url) - - # Build message - message_id = str(uuid4()) - ctx_id = context_id or str(uuid4()) - + client = await self._create_client(peer_id) request = SendMessageRequest( - id=str(uuid4()), - params=MessageSendParams( - message={ - "role": "user", - "parts": [{"kind": "text", "text": message}], - "messageId": message_id, - "contextId": ctx_id, - } - ), + message=Message( + role=Role.ROLE_USER, + message_id=str(uuid4()), + context_id=context_id or str(uuid4()), + parts=[Part(text=message)], + ) ) - - # Send message and get response - response = await client.send_message(request) - - # Extract response text - return self._extract_response(response) - - except Exception as e: - error_msg = f"Failed to call peer '{peer_id}': {e}" + task: Task | None = None + async for event in client.send_message(request): + if event.HasField("message"): + return text_from_message(event.message) + if event.HasField("task"): + task = Task() + task.CopyFrom(event.task) + elif event.HasField("status_update"): + if task is None: + task = Task( + id=event.status_update.task_id, + context_id=event.status_update.context_id, + ) + task.status.CopyFrom(event.status_update.status) + elif event.HasField("artifact_update"): + if task is None: + task = Task( + id=event.artifact_update.task_id, + context_id=event.artifact_update.context_id, + ) + task.artifacts.add().CopyFrom(event.artifact_update.artifact) + return self._extract_response(task) if task else "No response from peer" + except Exception as exc: + error_msg = f"Failed to call peer '{peer_id}': {exc}" logger.error(error_msg) - raise RuntimeError(error_msg) from e - - def _extract_response(self, response: SendMessageSuccessResponse) -> str: - """Extract text response from A2A response. - - Args: - response: The A2A send message response - - Returns: - The extracted text response - """ - # The response contains a Task with status and/or artifacts - task = response.root.result - - # Try artifacts first (completed tasks) - if task.artifacts: - for artifact in task.artifacts: - if artifact.parts: - text = extract_text_from_parts(artifact.parts) - if text: - return text - - # Fall back to status message - if task.status and task.status.message: - parts = task.status.message.parts - if parts: - return extract_text_from_parts(parts) - + raise RuntimeError(error_msg) from exc + + def _extract_response(self, task: Task) -> str: + for artifact in task.artifacts: + text = "\n".join(part.text for part in artifact.parts if part.text) + if text: + return text + if task.status.message: + text = text_from_message(task.status.message) + if text: + return text return "No response from peer" async def __aenter__(self) -> GatewayClient: - """Async context manager entry.""" return self async def __aexit__(self, exc_type, exc_val, exc_tb) -> None: - """Async context manager exit.""" await self.close() diff --git a/examples/mixed/03_fact_checker_a2a.py b/examples/mixed/03_fact_checker_a2a.py index 20fd7f129..ab626c397 100644 --- a/examples/mixed/03_fact_checker_a2a.py +++ b/examples/mixed/03_fact_checker_a2a.py @@ -1,6 +1,6 @@ # /// script # requires-python = ">=3.11" -# dependencies = ["band-sdk[a2a]"] +# dependencies = ["band-sdk[a2a_gateway]"] # # [tool.uv.sources] # band-sdk = { git = "https://github.com/band-ai/band-sdk-python.git" } @@ -25,7 +25,6 @@ import uvicorn from a2a.server.agent_execution import AgentExecutor, RequestContext -from a2a.server.apps.jsonrpc.starlette_app import A2AStarletteApplication from a2a.server.events import EventQueue from a2a.server.request_handlers import DefaultRequestHandler from a2a.server.tasks import ( @@ -33,17 +32,20 @@ InMemoryTaskStore, TaskUpdater, ) +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.routes.rest_routes import create_rest_routes +from starlette.applications import Starlette from a2a.types import ( AgentCapabilities, AgentCard, + AgentInterface, AgentSkill, Part, TaskState, - TextPart, UnsupportedOperationError, ) -from a2a.utils import new_agent_text_message, new_task -from a2a.utils.errors import ServerError +from band.integrations.a2a.protocol import new_task, text_message from dotenv import load_dotenv from setup_logging import setup_logging @@ -78,20 +80,22 @@ async def execute( task = context.current_task if not task: - task = new_task(context.message) # type: ignore[arg-type] + if context.message is None: + raise ValueError("A2A request is missing its message") + task = new_task(context.message) await event_queue.enqueue_event(task) updater = TaskUpdater(event_queue, task.id, task.context_id) await updater.update_status( - TaskState.working, - new_agent_text_message( + TaskState.TASK_STATE_WORKING, + text_message( "Reviewing the request for API, config, and test-surface details...", - task.context_id, - task.id, + context_id=task.context_id, + task_id=task.id, ), ) await updater.add_artifact( - [Part(root=TextPart(text=_fact_check_response(request_text)))], + [Part(text=_fact_check_response(request_text))], name="fact_check_report", ) await updater.complete() @@ -101,7 +105,7 @@ async def cancel( context: RequestContext, event_queue: EventQueue, ) -> None: - raise ServerError(error=UnsupportedOperationError()) + raise UnsupportedOperationError() def main() -> None: @@ -116,8 +120,14 @@ def main() -> None: agent_card = AgentCard( name="Mixed Contract Checker", description="Deterministic contract-checking A2A service for the mixed example", - url=f"{base_url}/", version="1.0.0", + supported_interfaces=[ + AgentInterface( + url=base_url, + protocol_binding="JSONRPC", + protocol_version="1.0", + ) + ], default_input_modes=["text/plain"], default_output_modes=["text/plain"], capabilities=AgentCapabilities(streaming=True, push_notifications=False), @@ -137,12 +147,18 @@ def main() -> None: request_handler = DefaultRequestHandler( agent_executor=FactCheckerExecutor(), task_store=InMemoryTaskStore(), + agent_card=agent_card, push_config_store=InMemoryPushNotificationConfigStore(), ) - app = A2AStarletteApplication( - agent_card=agent_card, - http_handler=request_handler, - ).build() + app = Starlette( + routes=( + create_agent_card_routes(agent_card) + + create_jsonrpc_routes( + request_handler, rpc_url="/", enable_v0_3_compat=True + ) + + create_rest_routes(request_handler, enable_v0_3_compat=True) + ) + ) logger.info("Starting mixed contract checker A2A server on %s", base_url) uvicorn.run(app, host=host, port=port, log_level="warning") diff --git a/examples/mixed/04_risk_reviewer_a2a.py b/examples/mixed/04_risk_reviewer_a2a.py index ea248bf77..eb69e12fb 100644 --- a/examples/mixed/04_risk_reviewer_a2a.py +++ b/examples/mixed/04_risk_reviewer_a2a.py @@ -1,6 +1,6 @@ # /// script # requires-python = ">=3.11" -# dependencies = ["band-sdk[a2a]"] +# dependencies = ["band-sdk[a2a_gateway]"] # # [tool.uv.sources] # band-sdk = { git = "https://github.com/band-ai/band-sdk-python.git" } @@ -25,7 +25,6 @@ import uvicorn from a2a.server.agent_execution import AgentExecutor, RequestContext -from a2a.server.apps.jsonrpc.starlette_app import A2AStarletteApplication from a2a.server.events import EventQueue from a2a.server.request_handlers import DefaultRequestHandler from a2a.server.tasks import ( @@ -33,17 +32,20 @@ InMemoryTaskStore, TaskUpdater, ) +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.routes.rest_routes import create_rest_routes +from starlette.applications import Starlette from a2a.types import ( AgentCapabilities, AgentCard, + AgentInterface, AgentSkill, Part, TaskState, - TextPart, UnsupportedOperationError, ) -from a2a.utils import new_agent_text_message, new_task -from a2a.utils.errors import ServerError +from band.integrations.a2a.protocol import new_task, text_message from dotenv import load_dotenv from setup_logging import setup_logging @@ -79,20 +81,22 @@ async def execute( task = context.current_task if not task: - task = new_task(context.message) # type: ignore[arg-type] + if context.message is None: + raise ValueError("A2A request is missing its message") + task = new_task(context.message) await event_queue.enqueue_event(task) updater = TaskUpdater(event_queue, task.id, task.context_id) await updater.update_status( - TaskState.working, - new_agent_text_message( + TaskState.TASK_STATE_WORKING, + text_message( "Reviewing the request for rollout, compatibility, and rollback risks...", - task.context_id, - task.id, + context_id=task.context_id, + task_id=task.id, ), ) await updater.add_artifact( - [Part(root=TextPart(text=_risk_review_response(request_text)))], + [Part(text=_risk_review_response(request_text))], name="risk_review_report", ) await updater.complete() @@ -102,7 +106,7 @@ async def cancel( context: RequestContext, event_queue: EventQueue, ) -> None: - raise ServerError(error=UnsupportedOperationError()) + raise UnsupportedOperationError() def main() -> None: @@ -117,8 +121,14 @@ def main() -> None: agent_card = AgentCard( name="Mixed Risk Reviewer", description="Deterministic rollout-risk A2A service for the mixed example", - url=f"{base_url}/", version="1.0.0", + supported_interfaces=[ + AgentInterface( + url=base_url, + protocol_binding="JSONRPC", + protocol_version="1.0", + ) + ], default_input_modes=["text/plain"], default_output_modes=["text/plain"], capabilities=AgentCapabilities(streaming=True, push_notifications=False), @@ -136,12 +146,18 @@ def main() -> None: request_handler = DefaultRequestHandler( agent_executor=RiskReviewerExecutor(), task_store=InMemoryTaskStore(), + agent_card=agent_card, push_config_store=InMemoryPushNotificationConfigStore(), ) - app = A2AStarletteApplication( - agent_card=agent_card, - http_handler=request_handler, - ).build() + app = Starlette( + routes=( + create_agent_card_routes(agent_card) + + create_jsonrpc_routes( + request_handler, rpc_url="/", enable_v0_3_compat=True + ) + + create_rest_routes(request_handler, enable_v0_3_compat=True) + ) + ) logger.info("Starting mixed risk reviewer A2A server on %s", base_url) uvicorn.run(app, host=host, port=port, log_level="warning") diff --git a/examples/mixed/README.md b/examples/mixed/README.md index 928bb7d3d..f53a100e7 100644 --- a/examples/mixed/README.md +++ b/examples/mixed/README.md @@ -130,8 +130,8 @@ uv run examples/mixed/04_risk_reviewer_a2a.py Optional smoke checks before bridging: ```bash -curl http://127.0.0.1:10121/.well-known/agent.json -curl http://127.0.0.1:10122/.well-known/agent.json +curl http://127.0.0.1:10121/.well-known/agent-card.json +curl http://127.0.0.1:10122/.well-known/agent-card.json ``` At this point they are still remote A2A services only. They are not Band participants yet. @@ -203,8 +203,8 @@ The bridge is the missing piece. Starting `03_fact_checker_a2a.py` and `04_risk_ Check that: -- `http://127.0.0.1:10121/.well-known/agent.json` works -- `http://127.0.0.1:10122/.well-known/agent.json` works +- `http://127.0.0.1:10121/.well-known/agent-card.json` works +- `http://127.0.0.1:10122/.well-known/agent-card.json` works - the bridge process is using the right `MIXED_FACT_URL` and `MIXED_RISK_URL` ### The CrewAI agents start but do not respond diff --git a/examples/run_agent.py b/examples/run_agent.py index 89e47f078..d6a5689fe 100644 --- a/examples/run_agent.py +++ b/examples/run_agent.py @@ -805,7 +805,7 @@ async def run_a2a_gateway_agent( logger.info("Starting A2A Gateway on %s...", gateway_url) logger.info("Peers will be exposed at:") logger.info( - " - %s/agents/{peer_id}/.well-known/agent.json (discovery)", gateway_url + " - %s/agents/{peer_id}/.well-known/agent-card.json (discovery)", gateway_url ) logger.info(" - %s/agents/{peer_id}/v1/message:stream (messaging)", gateway_url) await agent.run() diff --git a/src/band/integrations/a2a/adapter.py b/src/band/integrations/a2a/adapter.py index a852e70b4..24007b911 100644 --- a/src/band/integrations/a2a/adapter.py +++ b/src/band/integrations/a2a/adapter.py @@ -3,71 +3,34 @@ from __future__ import annotations import logging -from typing import ClassVar, Any -from uuid import uuid4 +from typing import ClassVar +import httpx from a2a.client import Client, ClientConfig, ClientFactory -from a2a.client.middleware import ClientCallContext, ClientCallInterceptor +from a2a.types import SendMessageRequest from a2a.types import ( - AgentCard, Message as A2AMessage, - Part, Role, + SubscribeToTaskRequest, + StreamResponse, Task, - TaskArtifactUpdateEvent, - TaskIdParams, TaskState, - TaskStatusUpdateEvent, - TextPart, ) -from a2a.utils import get_message_text from band.converters.a2a import A2AHistoryConverter from band.core.protocols import AgentToolsProtocol from band.core.simple_adapter import SimpleAdapter from band.core.types import AdapterFeatures, Capability, Emit, PlatformMessage +from band.integrations.a2a.protocol import ( + TERMINAL_TASK_STATES, + state_name, + text_from_message, + text_message, +) from band.integrations.a2a.types import A2AAuth, A2ASessionState logger = logging.getLogger(__name__) -# Terminal states where task cannot be resumed -TERMINAL_STATES = ( - TaskState.completed, - TaskState.failed, - TaskState.canceled, - TaskState.rejected, - TaskState.auth_required, -) - -# String values of terminal states for comparison with A2ASessionState.task_state -TERMINAL_STATE_VALUES = frozenset(s.value for s in TERMINAL_STATES) - - -class _StaticHeadersInterceptor(ClientCallInterceptor): - """Attach static headers to all outbound A2A requests.""" - - def __init__(self, headers: dict[str, str]) -> None: - self._headers = dict(headers) - - async def intercept( - self, - method_name: str, - request_payload: dict[str, Any], - http_kwargs: dict[str, Any], - agent_card: AgentCard | None, - context: ClientCallContext | None, - ) -> tuple[dict[str, Any], dict[str, Any]]: - del method_name, agent_card, context - - if not self._headers: - return request_payload, http_kwargs - - updated_http_kwargs = dict(http_kwargs) - headers = dict(updated_http_kwargs.get("headers", {})) - headers.update(self._headers) - updated_http_kwargs["headers"] = headers - return request_payload, updated_http_kwargs - class A2AAdapter(SimpleAdapter[A2ASessionState]): """Adapter that forwards messages to a remote A2A agent. @@ -125,8 +88,10 @@ def __init__( self.auth = auth self.streaming = streaming self._client: Client | None = None + self._http_client: httpx.AsyncClient | None = None self._contexts: dict[str, str] = {} # room_id → A2A context_id self._tasks: dict[str, str] = {} # room_id → last task_id + self._task_cache: dict[tuple[str, str], Task] = {} # Track sender per task for mentions: (room_id, task_id) → sender info self._task_senders: dict[tuple[str, str], dict[str, str]] = {} @@ -136,16 +101,11 @@ async def on_started(self, agent_name: str, agent_description: str) -> None: headers = self.auth.to_headers() if self.auth else {} - # Build client configuration - config = ClientConfig(streaming=self.streaming) - - # Connect to remote A2A agent - self._client = await ClientFactory.connect( - agent=self.remote_url, - client_config=config, - interceptors=[_StaticHeadersInterceptor(headers)] if headers else None, - resolver_http_kwargs={"headers": headers} if headers else None, + self._http_client = httpx.AsyncClient(headers=headers) + factory = ClientFactory( + ClientConfig(streaming=self.streaming, httpx_client=self._http_client) ) + self._client = await factory.create_from_url(self.remote_url) logger.info( "Connected to A2A agent at %s (streaming=%s, auth=%s)", @@ -185,7 +145,9 @@ async def on_message( try: # Send to remote A2A agent and process events - async for event in self._client.send_message(a2a_message): + async for event in self._client.send_message( + SendMessageRequest(message=a2a_message) + ): await self._handle_event( event, tools, room_id, msg.sender_id, msg.sender_name ) @@ -200,17 +162,15 @@ async def on_message( async def _handle_event( self, - event: tuple[Task, TaskStatusUpdateEvent | TaskArtifactUpdateEvent | None] - | A2AMessage, + event: StreamResponse, tools: AgentToolsProtocol, room_id: str, sender_id: str, sender_name: str | None, ) -> None: """Handle A2A event and forward to Band platform.""" - # Handle direct message reply (rare - most responses come via Task) - if isinstance(event, A2AMessage): - text = get_message_text(event) + if getattr(event, "HasField", lambda _name: False)("message"): + text = text_from_message(event.message) if text: await tools.send_message( content=text, @@ -218,8 +178,27 @@ async def _handle_event( ) return - # Unpack task event - task, update = event + if event.HasField("task"): + task = event.task + self._task_cache[(room_id, task.id)] = Task() + self._task_cache[(room_id, task.id)].CopyFrom(task) + elif event.HasField("status_update"): + update = event.status_update + task = self._task_cache.setdefault( + (room_id, update.task_id), + Task(id=update.task_id, context_id=update.context_id), + ) + task.status.CopyFrom(update.status) + elif event.HasField("artifact_update"): + update = event.artifact_update + task = self._task_cache.setdefault( + (room_id, update.task_id), + Task(id=update.task_id, context_id=update.context_id), + ) + task.artifacts.add().CopyFrom(update.artifact) + else: + return + key = (room_id, task.id) # Store sender info on first event for this task @@ -235,13 +214,13 @@ async def _handle_event( try: # Handle based on task state - if state == TaskState.working: + if state == TaskState.TASK_STATE_WORKING: # Stream progress as thought event (no mentions needed for events) status_text = self._get_status_text(task) if status_text: await tools.send_event(content=status_text, message_type="thought") - elif state == TaskState.input_required: + elif state == TaskState.TASK_STATE_INPUT_REQUIRED: # Agent needs more info - send as message with mention text = self._get_status_text(task) or "Please provide more information." sender = self._task_senders.get(key) @@ -252,7 +231,7 @@ async def _handle_event( # Emit task event for rehydration (input_required is resumable) await self._emit_task_event(tools, task, state) - elif state == TaskState.completed: + elif state == TaskState.TASK_STATE_COMPLETED: # Extract and send final response with mention response = self._extract_response(task) if response: @@ -263,25 +242,26 @@ async def _handle_event( ) elif state in ( - TaskState.failed, - TaskState.canceled, - TaskState.rejected, - TaskState.auth_required, + TaskState.TASK_STATE_FAILED, + TaskState.TASK_STATE_CANCELED, + TaskState.TASK_STATE_REJECTED, + TaskState.TASK_STATE_AUTH_REQUIRED, ): # Error states - send as error event (no mentions needed) - error_text = self._get_status_text(task) or f"Task {state.value}" + error_text = self._get_status_text(task) or f"Task {state_name(state)}" await tools.send_event( content=error_text, message_type="error", - metadata={"a2a_state": state.value}, + metadata={"a2a_state": state_name(state)}, ) finally: # Clean up on terminal states - if state in TERMINAL_STATES: + if state in TERMINAL_TASK_STATES: # Emit task event for rehydration (records final state) await self._emit_task_event(tools, task, state) # Clean up sender tracking self._task_senders.pop(key, None) + self._task_cache.pop(key, None) # Clear task_id so next message starts a new task # (context_id is preserved for multi-turn conversation) self._tasks.pop(room_id, None) @@ -296,18 +276,17 @@ def _to_a2a_message(self, msg: PlatformMessage, room_id: str) -> A2AMessage: context_id, task_id, ) - return A2AMessage( - role=Role.user, - message_id=str(uuid4()), - parts=[Part(root=TextPart(text=msg.content))], - context_id=context_id, # Existing context or None - task_id=task_id, # Continue existing task + return text_message( + msg.content, + role=Role.ROLE_USER, + context_id=context_id, + task_id=task_id, ) def _get_status_text(self, task: Task) -> str | None: """Extract text from task status message.""" if task.status.message: - return get_message_text(task.status.message) + return text_from_message(task.status.message) return None def _extract_response(self, task: Task) -> str: @@ -322,20 +301,20 @@ def _extract_response(self, task: Task) -> str: if task.artifacts: for artifact in task.artifacts: for part in artifact.parts: - if isinstance(part.root, TextPart): - return part.root.text + if part.text: + return part.text # Fallback: check status message if task.status.message: - text = get_message_text(task.status.message) + text = text_from_message(task.status.message) if text: return text # Last resort: check history for last agent message if task.history: for msg in reversed(task.history): - if msg.role == Role.agent: - text = get_message_text(msg) + if msg.role == Role.ROLE_AGENT: + text = text_from_message(msg) if text: return text @@ -349,6 +328,7 @@ async def on_cleanup(self, room_id: str) -> None: keys_to_remove = [key for key in self._task_senders if key[0] == room_id] for key in keys_to_remove: self._task_senders.pop(key, None) + self._task_cache.pop(key, None) logger.debug("Cleaned up A2A context for room %s", room_id) async def _emit_task_event( @@ -364,12 +344,12 @@ async def _emit_task_event( state: Current task state. """ await tools.send_event( - content=f"A2A task {state.value}", + content=f"A2A task {state_name(state)}", message_type="task", metadata={ "a2a_context_id": task.context_id, "a2a_task_id": task.id, - "a2a_task_state": state.value, + "a2a_task_state": state_name(state), }, ) @@ -393,7 +373,9 @@ async def _rehydrate_from_history( ) # Try to resume task if it was in a resumable state - if state.task_id and state.task_state not in TERMINAL_STATE_VALUES: + if state.task_id and state.task_state not in { + state_name(value) for value in TERMINAL_TASK_STATES + }: await self._try_resubscribe(room_id, state.task_id) async def _try_resubscribe(self, room_id: str, task_id: str) -> None: @@ -410,25 +392,37 @@ async def _try_resubscribe(self, room_id: str, task_id: str) -> None: return try: - async for event in self._client.resubscribe(TaskIdParams(id=task_id)): - if isinstance(event, tuple): - task, _ = event - current_state = task.status.state - if current_state not in TERMINAL_STATES: - self._tasks[room_id] = task_id - if task.context_id: - self._contexts[room_id] = task.context_id - logger.info( - "Resumed A2A task %s (state=%s)", - task_id, - current_state.value, - ) - else: - logger.info( - "A2A task %s already terminal (state=%s)", - task_id, - current_state.value, - ) - break # Only need first event to get current state + async for event in self._client.subscribe( + SubscribeToTaskRequest(id=task_id) + ): + if event.HasField("task"): + task = event.task + elif event.HasField("status_update"): + update = event.status_update + task = self._task_cache.setdefault( + (room_id, update.task_id), + Task(id=update.task_id, context_id=update.context_id), + ) + task.status.CopyFrom(update.status) + else: + continue + + current_state = task.status.state + if current_state not in TERMINAL_TASK_STATES: + self._tasks[room_id] = task_id + if task.context_id: + self._contexts[room_id] = task.context_id + logger.info( + "Resumed A2A task %s (state=%s)", + task_id, + state_name(current_state), + ) + else: + logger.info( + "A2A task %s already terminal (state=%s)", + task_id, + state_name(current_state), + ) + break # Only need first event to get current state except Exception as e: logger.warning("Could not resubscribe to A2A task %s: %s", task_id, e) diff --git a/src/band/integrations/a2a/gateway/adapter.py b/src/band/integrations/a2a/gateway/adapter.py index 420bfba97..c0223c404 100644 --- a/src/band/integrations/a2a/gateway/adapter.py +++ b/src/band/integrations/a2a/gateway/adapter.py @@ -14,16 +14,11 @@ from a2a.server.agent_execution import AgentExecutor, RequestContext from a2a.server.events import EventQueue from a2a.types import ( - Message as A2AMessage, - Part, - Role, Task, TaskState, TaskStatus, TaskStatusUpdateEvent, - TextPart, ) -from a2a.utils import get_message_text from band.client.rest import ( AsyncRestClient, @@ -41,6 +36,12 @@ from band.integrations.a2a.gateway.server import GatewayServer from band.integrations.a2a.gateway.config import A2AGatewayAdapterConfig from band.integrations.a2a.gateway.types import GatewaySessionState, PendingA2ATask +from band.integrations.a2a.protocol import ( + is_terminal_state, + snapshot_task, + text_from_message, + text_message, +) from band_rest import Peer from band_rest.agent_api_peers.types.list_agent_peers_response import ( ListAgentPeersResponse, @@ -298,7 +299,26 @@ async def on_cleanup(self, room_id: str) -> None: """ pending = self._pending_tasks.pop(room_id, None) if pending: - pending.done.set() + if not is_terminal_state(pending.task.status.state): + pending.task.status.CopyFrom( + TaskStatus( + state=TaskState.TASK_STATE_FAILED, + message=text_message( + "Band room closed before the A2A response completed", + context_id=pending.task.context_id, + task_id=pending.task.id, + ), + ) + ) + await pending.publish_response( + TaskStatusUpdateEvent( + task_id=pending.task.id, + context_id=pending.task.context_id, + status=pending.task.status, + ) + ) + else: + pending.done.set() logger.debug("Cleaned up gateway resources for room %s", room_id) async def stop(self) -> None: @@ -328,7 +348,7 @@ def _make_task(self, context: RequestContext) -> Task: return Task( id=context.task_id, context_id=context.context_id, - status=TaskStatus(state=TaskState.working), + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), ) async def _execute_a2a( @@ -360,9 +380,9 @@ async def _execute_a2a( ) try: async with self.pending_task(room_id, pending): - await event_queue.enqueue_event(task) + await event_queue.enqueue_event(snapshot_task(task)) await self._emit_context_event(room_id, context_id) - content = get_message_text(context.message) or "" + content = text_from_message(context.message) await self._rest.agent_api_messages.create_agent_chat_message( chat_id=room_id, @@ -392,26 +412,21 @@ async def _execute_a2a( task.id, self.config.response_timeout_s, ) - task.status = TaskStatus( - state=TaskState.failed, - message=A2AMessage( - role=Role.agent, - message_id=str(uuid4()), - parts=[ - Part( - root=TextPart( - text="Timed out waiting for a Band response" - ) - ) - ], - ), + task.status.CopyFrom( + TaskStatus( + state=TaskState.TASK_STATE_FAILED, + message=text_message( + "Timed out waiting for a Band response", + context_id=task.context_id, + task_id=task.id, + ), + ) ) await event_queue.enqueue_event( TaskStatusUpdateEvent( task_id=task.id, context_id=task.context_id, status=task.status, - final=True, ) ) except asyncio.CancelledError: @@ -452,13 +467,12 @@ async def _cancel_a2a( ) -> None: """Publish the official terminal cancellation event.""" task = context.current_task or self._make_task(context) - task.status = TaskStatus(state=TaskState.canceled) + task.status.CopyFrom(TaskStatus(state=TaskState.TASK_STATE_CANCELED)) await event_queue.enqueue_event( TaskStatusUpdateEvent( task_id=task.id, context_id=task.context_id, status=task.status, - final=True, ) ) @@ -565,31 +579,29 @@ def _translate_to_a2a( message_type = getattr(msg, "message_type", "text") if message_type == "error": - state = TaskState.failed - final = True + state = TaskState.TASK_STATE_FAILED elif message_type in ("thought", "tool_call", "tool_result"): - state = TaskState.working - final = False + state = TaskState.TASK_STATE_WORKING else: # Regular text message = completed response - state = TaskState.completed - final = True + state = TaskState.TASK_STATE_COMPLETED # Update task status - task.status = TaskStatus( - state=state, - message=A2AMessage( - role=Role.agent, - message_id=str(uuid4()), - parts=[Part(root=TextPart(text=msg.content))], - ), + task.status.CopyFrom( + TaskStatus( + state=state, + message=text_message( + msg.content, + context_id=task.context_id, + task_id=task.id, + ), + ) ) return TaskStatusUpdateEvent( task_id=task.id, context_id=task.context_id, status=task.status, - final=final, ) async def _emit_context_event(self, room_id: str, context_id: str) -> None: diff --git a/src/band/integrations/a2a/gateway/server.py b/src/band/integrations/a2a/gateway/server.py index 882d9fc0e..82a2a8208 100644 --- a/src/band/integrations/a2a/gateway/server.py +++ b/src/band/integrations/a2a/gateway/server.py @@ -8,13 +8,14 @@ from typing import Any from a2a.server.agent_execution import AgentExecutor -from a2a.server.apps.jsonrpc.starlette_app import A2AStarletteApplication -from a2a.server.apps.rest.rest_adapter import RESTAdapter 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.routes.rest_routes import create_rest_routes from a2a.server.tasks import InMemoryTaskStore from a2a.types import AgentCapabilities, AgentCard, AgentInterface, AgentSkill from starlette.applications import Starlette -from starlette.routing import Route +from starlette.routing import BaseRoute, Route from band_rest import Peer @@ -62,7 +63,13 @@ def _agent_card(self, slug: str, peer: Peer) -> AgentCard: return AgentCard( name=peer.name, description=peer.description or "", - url=rpc_url, + supported_interfaces=[ + AgentInterface( + protocol_binding="JSONRPC", + protocol_version="1.0", + url=rpc_url, + ) + ], version="1.0.0", capabilities=AgentCapabilities(streaming=True), skills=[ @@ -75,42 +82,53 @@ def _agent_card(self, slug: str, peer: Peer) -> AgentCard: ], default_input_modes=["text/plain"], default_output_modes=["text/plain"], - preferred_transport="JSONRPC", - additional_interfaces=[AgentInterface(transport="JSONRPC", url=rpc_url)], ) def _build_app(self) -> Starlette: - routes: list[Route] = [ + routes: list[BaseRoute] = [ Route("/peers", self._handle_list_peers, methods=["GET"]), ] + protocol_routes: list[BaseRoute] = [] + rest_routes: list[BaseRoute] = [] for slug, peer in self.peers.items(): - card = self._agent_card(slug, peer) - handler = DefaultRequestHandler( - agent_executor=self.executor_factory(slug), - task_store=InMemoryTaskStore(), - ) - jsonrpc_app = A2AStarletteApplication( - agent_card=card, - http_handler=handler, - ) - rest_adapter = RESTAdapter(agent_card=card, http_handler=handler) - routes.extend( - jsonrpc_app.routes( - agent_card_url=f"/agents/{slug}/.well-known/agent.json", - rpc_url=f"/agents/{slug}", + aliases = dict.fromkeys((slug, peer.id)) + for alias in aliases: + card = self._agent_card(alias, peer) + executor = self.executor_factory(slug) + handler = DefaultRequestHandler( + agent_executor=executor, + task_store=InMemoryTaskStore(), + agent_card=card, ) - ) - routes.extend( - Route( - f"/agents/{slug}{path}", - endpoint, - methods=[method], + protocol_routes.extend( + create_agent_card_routes( + card, + card_url=f"/agents/{alias}/.well-known/agent-card.json", + ) + ) + protocol_routes.extend( + create_agent_card_routes( + card, + card_url=f"/agents/{alias}/.well-known/agent.json", + ) + ) + protocol_routes.extend( + create_jsonrpc_routes( + handler, + rpc_url=f"/agents/{alias}", + enable_v0_3_compat=True, + ) + ) + rest_routes.extend( + create_rest_routes( + handler, + enable_v0_3_compat=True, + path_prefix=f"/agents/{alias}", + ) ) - for (path, method), endpoint in rest_adapter.routes().items() - ) - return Starlette(routes=routes) + return Starlette(routes=routes + protocol_routes + rest_routes) async def _handle_list_peers(self, _request: Any) -> Any: from starlette.responses import JSONResponse diff --git a/src/band/integrations/a2a/gateway/types.py b/src/band/integrations/a2a/gateway/types.py index 9196a0334..b3eb3fcb2 100644 --- a/src/band/integrations/a2a/gateway/types.py +++ b/src/band/integrations/a2a/gateway/types.py @@ -8,6 +8,8 @@ from a2a.server.events import EventQueue from a2a.types import Task, TaskStatusUpdateEvent +from band.integrations.a2a.protocol import is_terminal_state + @dataclass class GatewaySessionState: @@ -49,5 +51,5 @@ class PendingA2ATask: async def publish_response(self, event: TaskStatusUpdateEvent) -> None: """Publish a response and release the executor on terminal events.""" await self.event_queue.enqueue_event(event) - if event.final: + if is_terminal_state(event.status.state): self.done.set() diff --git a/src/band/integrations/a2a/protocol.py b/src/band/integrations/a2a/protocol.py new file mode 100644 index 000000000..8d9bfef25 --- /dev/null +++ b/src/band/integrations/a2a/protocol.py @@ -0,0 +1,78 @@ +"""Small protocol helpers shared by the A2A integrations.""" + +from __future__ import annotations + +from copy import deepcopy +from uuid import uuid4 + +from a2a.types import ( + Message, + Part, + Role, + SendMessageRequest, + Task, + TaskState, + TaskStatus, +) + +TERMINAL_TASK_STATES = frozenset( + { + TaskState.TASK_STATE_COMPLETED, + TaskState.TASK_STATE_FAILED, + TaskState.TASK_STATE_CANCELED, + TaskState.TASK_STATE_REJECTED, + TaskState.TASK_STATE_AUTH_REQUIRED, + } +) + + +def text_from_message(message: Message | None) -> str: + """Return the text parts from a protobuf A2A message.""" + if message is None: + return "" + return "\n".join(part.text for part in message.parts if part.text) + + +def text_message( + content: str, + *, + role: Role = Role.ROLE_AGENT, + context_id: str | None = None, + task_id: str | None = None, +) -> Message: + """Build a text-only protobuf A2A message.""" + message = Message( + message_id=str(uuid4()), + role=role, + parts=[Part(text=content)], + ) + if context_id: + message.context_id = context_id + if task_id: + message.task_id = task_id + return message + + +def new_task(request: SendMessageRequest | Message) -> Task: + """Create a working task from an incoming A2A message.""" + message = request.message if isinstance(request, SendMessageRequest) else request + return Task( + id=message.task_id or str(uuid4()), + context_id=message.context_id or str(uuid4()), + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + + +def snapshot_task(task: Task) -> Task: + """Return an independent task snapshot for event queue publication.""" + return deepcopy(task) + + +def is_terminal_state(state: int) -> bool: + """Return whether an A2A task state ends execution.""" + return state in TERMINAL_TASK_STATES + + +def state_name(state: int) -> str: + """Return the stable enum name for a protobuf task state.""" + return TaskState.Name(state) diff --git a/tests/integrations/a2a/gateway/fixtures.py b/tests/integrations/a2a/gateway/helpers.py similarity index 88% rename from tests/integrations/a2a/gateway/fixtures.py rename to tests/integrations/a2a/gateway/helpers.py index 66a0f4c01..1ff13dc3f 100644 --- a/tests/integrations/a2a/gateway/fixtures.py +++ b/tests/integrations/a2a/gateway/helpers.py @@ -1,4 +1,4 @@ -"""Shared fixtures and builders for gateway tests.""" +"""Shared test-data builders for the A2A gateway.""" from __future__ import annotations diff --git a/tests/integrations/a2a/gateway/test_adapter.py b/tests/integrations/a2a/gateway/test_adapter.py index 6938563c6..8b554b59a 100644 --- a/tests/integrations/a2a/gateway/test_adapter.py +++ b/tests/integrations/a2a/gateway/test_adapter.py @@ -1,4 +1,4 @@ -"""Tests for A2AGatewayAdapter.""" +"""Behavior tests for the Band-backed A2A gateway executor.""" from __future__ import annotations @@ -8,33 +8,32 @@ from uuid import uuid4 import pytest -from a2a.server.events import EventQueue from a2a.server.agent_execution import RequestContext +from a2a.server.events import EventQueueLegacy from a2a.types import ( - Message as A2AMessage, - MessageSendParams, + Message, Part, Role, + SendMessageRequest, + Task, TaskState, - TextPart, + TaskStatus, ) from band.core.types import PlatformMessage -from band.integrations.a2a.gateway import ( - A2AGatewayAdapter, - A2AGatewayAdapterConfig, - GatewaySessionState, -) +from band.integrations.a2a.gateway import A2AGatewayAdapter, A2AGatewayAdapterConfig from band.integrations.a2a.gateway.adapter import BandAgentExecutor +from band.integrations.a2a.gateway.types import GatewaySessionState, PendingA2ATask from band.testing import FakeAgentTools from band_rest.core.api_error import ApiError -from tests.integrations.a2a.gateway.fixtures import make_peer +from tests.integrations.a2a.gateway.helpers import make_peer def make_platform_message( - content: str, room_id: str = "room-123", message_type: str = "text" + content: str, + room_id: str = "room-123", + message_type: str = "text", ) -> PlatformMessage: - """Create a test PlatformMessage.""" return PlatformMessage( id=str(uuid4()), room_id=room_id, @@ -48,310 +47,150 @@ def make_platform_message( ) -def make_a2a_message( - content: str, context_id: str | None = None, task_id: str | None = None -) -> A2AMessage: - """Create an A2A message for testing.""" - return A2AMessage( - role=Role.user, +def make_request(content: str = "What is the weather?") -> RequestContext: + message = Message( message_id=str(uuid4()), - parts=[Part(root=TextPart(text=content))], - context_id=context_id, - task_id=task_id, + role=Role.ROLE_USER, + parts=[Part(text=content)], ) + return RequestContext(None, request=SendMessageRequest(message=message)) -class TestA2AGatewayAdapterInit: - """Tests for A2AGatewayAdapter initialization.""" +def configure_room_creation(adapter: A2AGatewayAdapter) -> None: + response = MagicMock() + response.data.id = "room-123" + adapter._rest.agent_api_chats.create_agent_chat = AsyncMock(return_value=response) + adapter._rest.agent_api_participants.add_agent_chat_participant = AsyncMock() + adapter._rest.agent_api_messages.create_agent_chat_message = AsyncMock() + adapter._rest.agent_api_events.create_agent_chat_event = AsyncMock() - def test_config_rejects_non_positive_response_timeout(self) -> None: - """Should reject a timeout that cannot provide a deadline.""" - with pytest.raises(ValueError, match="response_timeout_s"): - A2AGatewayAdapterConfig(response_timeout_s=0) - def test_init_default_values(self) -> None: - """Should initialize with default values.""" - adapter = A2AGatewayAdapter() +def make_pending(event_queue: EventQueueLegacy) -> PendingA2ATask: + return PendingA2ATask( + task=Task( + id="task-123", + context_id="ctx-123", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ), + event_queue=event_queue, + peer_id="weather", + done=asyncio.Event(), + ) - assert adapter.gateway_url == "http://localhost:10000" - assert adapter.port == 10000 - assert adapter._peers == {} - assert adapter._server is None - assert adapter._context_to_room == {} - assert adapter._room_participants == {} - assert adapter._pending_tasks == {} - def test_init_with_custom_values(self) -> None: - """Should accept custom configuration.""" - adapter = A2AGatewayAdapter( - rest_url="https://custom.api.com", - api_key="test-key", - gateway_url="http://localhost:9000", - port=9000, - ) +class TestGatewayConfiguration: + def test_timeout_is_adapter_configuration(self) -> None: + config = A2AGatewayAdapterConfig(response_timeout_s=12) + adapter = A2AGatewayAdapter(config=config) - assert adapter.gateway_url == "http://localhost:9000" - assert adapter.port == 9000 + assert adapter.config is config + assert adapter.config.response_timeout_s == 12 + def test_timeout_must_be_positive(self) -> None: + with pytest.raises(ValueError, match="response_timeout_s"): + A2AGatewayAdapterConfig(response_timeout_s=0) -class TestA2AGatewayAdapterOnStarted: - """Tests for A2AGatewayAdapter.on_started().""" +class TestGatewayStartup: @pytest.mark.asyncio - async def test_on_started_fetches_peers_via_rest(self) -> None: - """Should fetch peers from REST API.""" + async def test_discovers_peers_and_starts_server(self) -> None: adapter = A2AGatewayAdapter() - - # Mock REST client - mock_response = MagicMock() - mock_response.data = [ - make_peer("weather", "Weather Agent"), - make_peer("servicenow", "ServiceNow Agent"), - ] - adapter._rest.agent_api_peers.list_agent_peers = AsyncMock( - return_value=mock_response - ) - - # Mock server - with patch( - "band.integrations.a2a.gateway.adapter.GatewayServer" - ) as mock_server_class: - mock_server = MagicMock() - mock_server.start = AsyncMock() - mock_server_class.return_value = mock_server - - await adapter.on_started("Gateway", "A2A Gateway Agent") - - # Peers are now keyed by slug, not UUID - assert len(adapter._peers) == 2 - assert "weather-agent" in adapter._peers - assert "servicenow-agent" in adapter._peers - # UUID fallback should also be populated - assert "weather" in adapter._peers_by_uuid - assert "servicenow" in adapter._peers_by_uuid - - @pytest.mark.asyncio - async def test_on_started_starts_http_server(self) -> None: - """Should start HTTP server with peer routes.""" - adapter = A2AGatewayAdapter(port=10001) - - # Mock REST client - mock_response = MagicMock() - mock_response.data = [make_peer("weather", "Weather Agent")] + response = MagicMock() + response.data = [make_peer("weather", "Weather Agent")] adapter._rest.agent_api_peers.list_agent_peers = AsyncMock( - return_value=mock_response + return_value=response ) - # Mock server with patch( "band.integrations.a2a.gateway.adapter.GatewayServer" - ) as mock_server_class: - mock_server = MagicMock() - mock_server.start = AsyncMock() - mock_server_class.return_value = mock_server + ) as server_type: + server = MagicMock() + server.start = AsyncMock() + server_type.return_value = server - await adapter.on_started("Gateway", "A2A Gateway Agent") + await adapter.on_started("Gateway", "A2A Gateway") - mock_server_class.assert_called_once() - mock_server.start.assert_called_once() - assert adapter._server is mock_server + assert adapter._peers["weather-agent"].id == "weather" + server.start.assert_awaited_once() @pytest.mark.asyncio - async def test_on_started_retries_peer_discovery_on_rate_limit(self) -> None: - """Should retry peer discovery when startup hits HTTP 429.""" + async def test_retries_peer_discovery_only_for_rate_limits(self) -> None: adapter = A2AGatewayAdapter() - - mock_response = MagicMock() - mock_response.data = [make_peer("weather", "Weather Agent")] + response = MagicMock() + response.data = [] adapter._rest.agent_api_peers.list_agent_peers = AsyncMock( - side_effect=[ - ApiError(status_code=429, headers={}, body=""), - ApiError(status_code=429, headers={}, body=""), - mock_response, - ] + side_effect=[ApiError(status_code=429, headers={}, body=""), response] ) - with patch( - "band.integrations.a2a.gateway.adapter.GatewayServer" - ) as mock_server_class: - mock_server = MagicMock() - mock_server.start = AsyncMock() - mock_server_class.return_value = mock_server - with patch( + with ( + patch("band.integrations.a2a.gateway.adapter.GatewayServer") as server_type, + patch( "band.integrations.a2a.gateway.adapter.asyncio.sleep", new=AsyncMock(), - ) as mock_sleep: - await adapter.on_started("Gateway", "A2A Gateway Agent") - - assert adapter._rest.agent_api_peers.list_agent_peers.await_count == 3 - assert mock_sleep.await_count == 2 + ) as sleep, + ): + server_type.return_value.start = AsyncMock() + await adapter.on_started("Gateway", "A2A Gateway") + assert adapter._rest.agent_api_peers.list_agent_peers.await_count == 2 + sleep.assert_awaited_once() -class TestA2AGatewayAdapterOnMessage: - """Tests for A2AGatewayAdapter.on_message().""" - @pytest.fixture - def adapter_with_mocks(self) -> A2AGatewayAdapter: - """Create adapter with mocked dependencies.""" +class TestGatewayExecution: + @pytest.mark.asyncio + async def test_initial_task_snapshot_stays_working_if_reply_is_immediate(self) -> None: adapter = A2AGatewayAdapter() adapter._peers = {"weather": make_peer("weather", "Weather Agent")} - adapter._rest.agent_api_chats.create_agent_chat = AsyncMock() - adapter._rest.agent_api_participants.add_agent_chat_participant = AsyncMock() - adapter._rest.agent_api_messages.create_agent_chat_message = AsyncMock() - adapter._rest.agent_api_events.create_agent_chat_event = AsyncMock() - return adapter - - @pytest.mark.asyncio - async def test_on_message_rehydrates_on_bootstrap( - self, adapter_with_mocks: A2AGatewayAdapter - ) -> None: - """Should rehydrate session state on bootstrap.""" - tools = FakeAgentTools() - msg = make_platform_message("Hello", room_id="room-123") - - history = GatewaySessionState( - context_to_room={"ctx-1": "room-1", "ctx-2": "room-2"}, - room_participants={"room-1": {"peer-a"}, "room-2": {"peer-b"}}, - ) - - await adapter_with_mocks.on_message( - msg, - tools, - history, - None, - None, - is_session_bootstrap=True, - room_id="room-123", - ) - - assert adapter_with_mocks._context_to_room == { - "ctx-1": "room-1", - "ctx-2": "room-2", - } - assert adapter_with_mocks._room_participants == { - "room-1": {"peer-a"}, - "room-2": {"peer-b"}, - } - - @pytest.mark.asyncio - async def test_on_message_correlates_pending_task( - self, adapter_with_mocks: A2AGatewayAdapter - ) -> None: - """Should publish non-final updates without completing the task.""" - tools = FakeAgentTools() - msg = make_platform_message( - "Checking the forecast", room_id="room-123", message_type="thought" - ) - - from band.integrations.a2a.gateway.types import PendingA2ATask - from a2a.types import Task, TaskStatus - - event_queue = EventQueue() - task = Task( - id="task-123", - context_id="ctx-123", - status=TaskStatus(state=TaskState.working), - ) - adapter_with_mocks._pending_tasks["room-123"] = PendingA2ATask( - task=task, - event_queue=event_queue, - peer_id="weather", - done=asyncio.Event(), - ) - - await adapter_with_mocks.on_message( - msg, - tools, - GatewaySessionState(), - None, - None, - is_session_bootstrap=False, - room_id="room-123", - ) - - pending = adapter_with_mocks._pending_tasks["room-123"] - event = await event_queue.dequeue_event() - assert event.task_id == "task-123" - assert event.final is False - assert not pending.done.is_set() - - @pytest.mark.asyncio - async def test_on_message_completes_pending_task( - self, adapter_with_mocks: A2AGatewayAdapter - ) -> None: - """Should complete the pending task on a final response.""" + configure_room_creation(adapter) tools = FakeAgentTools() - msg = make_platform_message("Done", room_id="room-123") - # Set up pending task - from band.integrations.a2a.gateway.types import PendingA2ATask - from a2a.types import Task, TaskStatus - - event_queue = EventQueue() - task = Task( - id="task-123", - context_id="ctx-123", - status=TaskStatus(state=TaskState.working), - ) - adapter_with_mocks._pending_tasks["room-123"] = PendingA2ATask( - task=task, - event_queue=event_queue, - peer_id="weather", - done=asyncio.Event(), - ) + async def send_message(**_kwargs: object) -> None: + await adapter.on_message( + make_platform_message("Sunny"), + tools, + GatewaySessionState(), + None, + None, + is_session_bootstrap=False, + room_id="room-123", + ) - await adapter_with_mocks.on_message( - msg, - tools, - GatewaySessionState(), - None, - None, - is_session_bootstrap=False, - room_id="room-123", + adapter._rest.agent_api_messages.create_agent_chat_message = AsyncMock( + side_effect=send_message ) + queue = EventQueueLegacy() - pending = adapter_with_mocks._pending_tasks["room-123"] - assert pending.done.is_set() - + await BandAgentExecutor(adapter, "weather").execute(make_request(), queue) -class TestA2AGatewayExecutor: - """Tests the Band-backed execution boundary used by official A2A handlers.""" + initial = await queue.dequeue_event() + terminal = await queue.dequeue_event() + assert initial.status.state == TaskState.TASK_STATE_WORKING + assert terminal.status.state == TaskState.TASK_STATE_COMPLETED @pytest.mark.asyncio - async def test_execute_posts_and_emits_ordered_task_response(self) -> None: + async def test_posts_to_band_and_returns_terminal_response(self) -> None: adapter = A2AGatewayAdapter() adapter._peers = {"weather": make_peer("weather", "Weather Agent")} + configure_room_creation(adapter) + sent = asyncio.Event() - chat_response = MagicMock() - chat_response.data.id = "room-123" - adapter._rest.agent_api_chats.create_agent_chat = AsyncMock( - return_value=chat_response - ) - adapter._rest.agent_api_participants.add_agent_chat_participant = AsyncMock() - message_sent = asyncio.Event() - - async def send_message(**_kwargs) -> None: - message_sent.set() + async def send_message(**_kwargs: object) -> None: + sent.set() adapter._rest.agent_api_messages.create_agent_chat_message = AsyncMock( side_effect=send_message ) - adapter._rest.agent_api_events.create_agent_chat_event = AsyncMock() - - context = RequestContext( - request=MessageSendParams(message=make_a2a_message("What is the weather?")) + queue = EventQueueLegacy() + execution = asyncio.create_task( + BandAgentExecutor(adapter, "weather").execute(make_request(), queue) ) - event_queue = EventQueue() - executor = BandAgentExecutor(adapter, "weather") - execution = asyncio.create_task(executor.execute(context, event_queue)) + await asyncio.wait_for(sent.wait(), timeout=1) - await asyncio.wait_for(message_sent.wait(), timeout=1) - initial = await event_queue.dequeue_event() - assert initial.kind == "task" - assert initial.status.state == TaskState.working + initial = await queue.dequeue_event() + assert initial.status.state == TaskState.TASK_STATE_WORKING await adapter.on_message( - make_platform_message("Sunny", room_id="room-123"), + make_platform_message("Sunny"), FakeAgentTools(), GatewaySessionState(), None, @@ -360,97 +199,36 @@ async def send_message(**_kwargs) -> None: room_id="room-123", ) await asyncio.wait_for(execution, timeout=1) + final = await queue.dequeue_event() - final = await event_queue.dequeue_event() - assert final.kind == "status-update" - assert final.final is True - assert final.status.state == TaskState.completed - assert final.status.message is not None - assert final.status.message.parts[0].root.text == "Sunny" - - request = adapter._rest.agent_api_messages.create_agent_chat_message.await_args - assert request.kwargs["chat_id"] == "room-123" - assert ( - request.kwargs["message"].content == "@Weather Agent What is the weather?" - ) + assert final.status.state == TaskState.TASK_STATE_COMPLETED + assert final.status.message.parts[0].text == "Sunny" assert adapter._pending_tasks == {} @pytest.mark.asyncio - async def test_execute_fails_when_band_response_times_out(self) -> None: - adapter = A2AGatewayAdapter( - config=A2AGatewayAdapterConfig(response_timeout_s=0.01) - ) - adapter._peers = {"weather": make_peer("weather", "Weather Agent")} - - chat_response = MagicMock() - chat_response.data.id = "room-123" - adapter._rest.agent_api_chats.create_agent_chat = AsyncMock( - return_value=chat_response - ) - adapter._rest.agent_api_participants.add_agent_chat_participant = AsyncMock() - adapter._rest.agent_api_messages.create_agent_chat_message = AsyncMock() - adapter._rest.agent_api_events.create_agent_chat_event = AsyncMock() - - context = RequestContext( - request=MessageSendParams(message=make_a2a_message("What is the weather?")) - ) - event_queue = EventQueue() - executor = BandAgentExecutor(adapter, "weather") - - await executor.execute(context, event_queue) - - initial = await event_queue.dequeue_event() - terminal = await event_queue.dequeue_event() - assert initial.kind == "task" - assert terminal.kind == "status-update" - assert terminal.final is True - assert terminal.status.state == TaskState.failed - assert terminal.status.message is not None - assert "Timed out" in terminal.status.message.parts[0].root.text - assert adapter._pending_tasks == {} - - @pytest.mark.asyncio - async def test_execute_keeps_stream_open_for_updates_until_final_response( - self, - ) -> None: + async def test_keeps_stream_open_for_non_final_updates(self) -> None: adapter = A2AGatewayAdapter( config=A2AGatewayAdapterConfig(response_timeout_s=1) ) adapter._peers = {"weather": make_peer("weather", "Weather Agent")} + configure_room_creation(adapter) + queue = EventQueueLegacy() + sent = asyncio.Event() - chat_response = MagicMock() - chat_response.data.id = "room-123" - adapter._rest.agent_api_chats.create_agent_chat = AsyncMock( - return_value=chat_response - ) - adapter._rest.agent_api_participants.add_agent_chat_participant = AsyncMock() - message_sent = asyncio.Event() - - async def send_message(**_kwargs) -> None: - message_sent.set() + async def send_message(**_kwargs: object) -> None: + sent.set() adapter._rest.agent_api_messages.create_agent_chat_message = AsyncMock( side_effect=send_message ) - adapter._rest.agent_api_events.create_agent_chat_event = AsyncMock() - - context = RequestContext( - request=MessageSendParams(message=make_a2a_message("What is the weather?")) - ) - event_queue = EventQueue() execution = asyncio.create_task( - BandAgentExecutor(adapter, "weather").execute(context, event_queue) + BandAgentExecutor(adapter, "weather").execute(make_request(), queue) ) - - await asyncio.wait_for(message_sent.wait(), timeout=1) - await event_queue.dequeue_event() + await asyncio.wait_for(sent.wait(), timeout=1) + await queue.dequeue_event() await adapter.on_message( - make_platform_message( - "Checking the forecast", - room_id="room-123", - message_type="thought", - ), + make_platform_message("Checking", message_type="thought"), FakeAgentTools(), GatewaySessionState(), None, @@ -458,12 +236,12 @@ async def send_message(**_kwargs) -> None: is_session_bootstrap=False, room_id="room-123", ) - update = await event_queue.dequeue_event() - assert update.final is False + update = await queue.dequeue_event() + assert update.status.state == TaskState.TASK_STATE_WORKING assert not execution.done() await adapter.on_message( - make_platform_message("Sunny", room_id="room-123"), + make_platform_message("Sunny"), FakeAgentTools(), GatewaySessionState(), None, @@ -472,344 +250,107 @@ async def send_message(**_kwargs) -> None: room_id="room-123", ) await asyncio.wait_for(execution, timeout=1) - - final = await event_queue.dequeue_event() - assert final.final is True - assert final.status.state == TaskState.completed - - -class TestA2AGatewayAdapterRoomManagement: - """Tests for room creation and management.""" - - @pytest.fixture - def adapter_with_mocks(self) -> A2AGatewayAdapter: - """Create adapter with mocked REST client.""" - adapter = A2AGatewayAdapter() - adapter._peers = {"weather": make_peer("weather", "Weather Agent")} - - # Mock create_agent_chat to return room with ID - mock_chat_response = MagicMock() - mock_chat_response.data = MagicMock() - mock_chat_response.data.id = "new-room-123" - adapter._rest.agent_api_chats.create_agent_chat = AsyncMock( - return_value=mock_chat_response - ) - adapter._rest.agent_api_participants.add_agent_chat_participant = AsyncMock() - return adapter - - @pytest.mark.asyncio - async def test_get_or_create_room_new_context( - self, adapter_with_mocks: A2AGatewayAdapter - ) -> None: - """Should create new room for new context.""" - room_id, context_id = await adapter_with_mocks._get_or_create_room( - None, "weather" - ) - - assert room_id == "new-room-123" - assert context_id is not None # UUID generated - assert adapter_with_mocks._context_to_room[context_id] == room_id - assert "weather" in adapter_with_mocks._room_participants[room_id] - - # REST calls should be made - adapter_with_mocks._rest.agent_api_chats.create_agent_chat.assert_called_once() - adapter_with_mocks._rest.agent_api_participants.add_agent_chat_participant.assert_called_once() - - @pytest.mark.asyncio - async def test_get_or_create_room_existing_context( - self, adapter_with_mocks: A2AGatewayAdapter - ) -> None: - """Should reuse existing room for known context.""" - # Pre-populate context mapping - adapter_with_mocks._context_to_room["existing-ctx"] = "existing-room" - adapter_with_mocks._room_participants["existing-room"] = {"weather"} - - room_id, context_id = await adapter_with_mocks._get_or_create_room( - "existing-ctx", "weather" - ) - - assert room_id == "existing-room" - assert context_id == "existing-ctx" - - # No new room should be created - adapter_with_mocks._rest.agent_api_chats.create_agent_chat.assert_not_called() + assert ( + await queue.dequeue_event() + ).status.state == TaskState.TASK_STATE_COMPLETED @pytest.mark.asyncio - async def test_get_or_create_room_adds_new_participant( - self, adapter_with_mocks: A2AGatewayAdapter - ) -> None: - """Should add new participant to existing room.""" - # Pre-populate context mapping with different peer - adapter_with_mocks._context_to_room["ctx-1"] = "room-1" - adapter_with_mocks._room_participants["room-1"] = {"other-peer"} - - room_id, context_id = await adapter_with_mocks._get_or_create_room( - "ctx-1", "weather" - ) - - assert room_id == "room-1" - assert "weather" in adapter_with_mocks._room_participants["room-1"] - - # Should add participant but not create room - adapter_with_mocks._rest.agent_api_chats.create_agent_chat.assert_not_called() - adapter_with_mocks._rest.agent_api_participants.add_agent_chat_participant.assert_called_once() - - -class TestA2AGatewayAdapterRehydration: - """Tests for session rehydration.""" - - def test_rehydrate_restores_context_mapping(self) -> None: - """Should restore context → room mappings.""" - adapter = A2AGatewayAdapter() - - history = GatewaySessionState( - context_to_room={"ctx-1": "room-1", "ctx-2": "room-2"}, - room_participants={}, - ) - - adapter._rehydrate(history) - - assert adapter._context_to_room == {"ctx-1": "room-1", "ctx-2": "room-2"} - - def test_rehydrate_restores_participants(self) -> None: - """Should restore room participants.""" - adapter = A2AGatewayAdapter() - - history = GatewaySessionState( - context_to_room={}, - room_participants={"room-1": {"peer-a", "peer-b"}}, - ) - - adapter._rehydrate(history) - - assert adapter._room_participants["room-1"] == {"peer-a", "peer-b"} - - def test_rehydrate_merges_with_existing(self) -> None: - """Should merge with existing state, not replace.""" - adapter = A2AGatewayAdapter() - adapter._context_to_room["existing-ctx"] = "existing-room" - adapter._room_participants["existing-room"] = {"existing-peer"} - - history = GatewaySessionState( - context_to_room={"new-ctx": "new-room"}, - room_participants={"new-room": {"new-peer"}}, - ) - - adapter._rehydrate(history) - - # Both old and new should be present - assert adapter._context_to_room["existing-ctx"] == "existing-room" - assert adapter._context_to_room["new-ctx"] == "new-room" - - def test_rehydrate_does_not_overwrite_existing_context(self) -> None: - """Should not overwrite existing context mappings.""" - adapter = A2AGatewayAdapter() - adapter._context_to_room["ctx-1"] = "current-room" - - history = GatewaySessionState( - context_to_room={"ctx-1": "old-room"}, # Same context, different room - room_participants={}, - ) - - adapter._rehydrate(history) - - # Should keep current mapping - assert adapter._context_to_room["ctx-1"] == "current-room" - - -class TestA2AGatewayAdapterTranslation: - """Tests for message translation.""" - - def test_translate_to_a2a_text_message(self) -> None: - """Should translate text message to completed event.""" - adapter = A2AGatewayAdapter() - msg = make_platform_message("Hello world") - - from a2a.types import Task, TaskStatus - - task = Task( - id="task-123", - context_id="ctx-123", - status=TaskStatus(state=TaskState.working), - ) - - event = adapter._translate_to_a2a(msg, task) - - assert event.task_id == "task-123" - assert event.context_id == "ctx-123" - assert event.status.state == TaskState.completed - assert event.final is True - - def test_translate_to_a2a_thought_message(self) -> None: - """Should translate thought message to working event.""" - adapter = A2AGatewayAdapter() - msg = make_platform_message("Thinking...", message_type="thought") - - from a2a.types import Task, TaskStatus - - task = Task( - id="task-123", - context_id="ctx-123", - status=TaskStatus(state=TaskState.working), - ) - - event = adapter._translate_to_a2a(msg, task) - - assert event.status.state == TaskState.working - assert event.final is False - - def test_translate_to_a2a_error_message(self) -> None: - """Should translate error message to failed event.""" - adapter = A2AGatewayAdapter() - msg = make_platform_message("Something went wrong", message_type="error") - - from a2a.types import Task, TaskStatus - - task = Task( - id="task-123", - context_id="ctx-123", - status=TaskStatus(state=TaskState.working), + async def test_timeout_returns_terminal_failure(self) -> None: + adapter = A2AGatewayAdapter( + config=A2AGatewayAdapterConfig(response_timeout_s=0.01) ) + adapter._peers = {"weather": make_peer("weather", "Weather Agent")} + configure_room_creation(adapter) + queue = EventQueueLegacy() - event = adapter._translate_to_a2a(msg, task) - - assert event.status.state == TaskState.failed - assert event.final is True - + await BandAgentExecutor(adapter, "weather").execute(make_request(), queue) -class TestA2AGatewayAdapterCleanup: - """Tests for cleanup methods.""" + await queue.dequeue_event() + terminal = await queue.dequeue_event() + assert terminal.status.state == TaskState.TASK_STATE_FAILED + assert adapter._pending_tasks == {} @pytest.mark.asyncio - async def test_on_cleanup_removes_pending_task(self) -> None: - """Should remove pending task for room.""" + async def test_room_cleanup_returns_terminal_failure(self) -> None: adapter = A2AGatewayAdapter() - - from band.integrations.a2a.gateway.types import PendingA2ATask - from a2a.types import Task, TaskStatus - - event_queue = EventQueue() - task = Task( - id="task-123", - context_id="ctx-123", - status=TaskStatus(state=TaskState.working), - ) - pending = PendingA2ATask( - task=task, - event_queue=event_queue, - peer_id="weather", - done=asyncio.Event(), - ) + queue = EventQueueLegacy() + pending = make_pending(queue) adapter._pending_tasks["room-123"] = pending await adapter.on_cleanup("room-123") - assert "room-123" not in adapter._pending_tasks + terminal = await queue.dequeue_event() + assert terminal.status.state == TaskState.TASK_STATE_FAILED assert pending.done.is_set() + assert adapter._pending_tasks == {} - @pytest.mark.asyncio - async def test_stop_stops_server(self) -> None: - """Should stop HTTP server.""" - adapter = A2AGatewayAdapter() - mock_server = MagicMock() - mock_server.stop = AsyncMock() - adapter._server = mock_server - - await adapter.stop() - - mock_server.stop.assert_called_once() - assert adapter._server is None - - -class TestGatewayContextIdRoomMapping: - """Tests that gateway maps context_id to rooms correctly.""" +class TestGatewayRoomState: @pytest.fixture - def adapter_with_tracking(self) -> A2AGatewayAdapter: - """Create adapter with mocked REST client that tracks room creation.""" + def adapter(self) -> A2AGatewayAdapter: adapter = A2AGatewayAdapter() - adapter._peers = {"weather": make_peer("weather", "Weather Agent")} - - # Track room creation with unique IDs - rooms_created: list[str] = [] - - def create_room_side_effect(*args, **kwargs): - room_id = f"room-{len(rooms_created) + 1}" - rooms_created.append(room_id) - mock_response = MagicMock() - mock_response.data = MagicMock() - mock_response.data.id = room_id - return mock_response - + adapter._peers = { + "weather": make_peer("weather", "Weather Agent"), + "data": make_peer("data", "Data Agent"), + } + response = MagicMock() + response.data.id = "new-room" adapter._rest.agent_api_chats.create_agent_chat = AsyncMock( - side_effect=create_room_side_effect + return_value=response ) adapter._rest.agent_api_participants.add_agent_chat_participant = AsyncMock() - adapter._rooms_created = rooms_created # Expose for assertions return adapter @pytest.mark.asyncio - async def test_same_context_id_twice_reuses_room( - self, adapter_with_tracking: A2AGatewayAdapter + async def test_context_reuses_room_and_adds_new_peer( + self, adapter: A2AGatewayAdapter ) -> None: - """Same context_id should map to same room - no new room created.""" - adapter = adapter_with_tracking + room, context = await adapter._get_or_create_room("ctx", "weather") + same_room, same_context = await adapter._get_or_create_room(context, "data") - # First request with context_id - room_id_1, ctx_1 = await adapter._get_or_create_room("ctx-user-123", "weather") - create_calls_after_first = ( - adapter._rest.agent_api_chats.create_agent_chat.call_count + assert (room, context) == ("new-room", "ctx") + assert (same_room, same_context) == (room, context) + assert adapter._room_participants[room] == {"weather", "data"} + assert adapter._rest.agent_api_chats.create_agent_chat.await_count == 1 + assert ( + adapter._rest.agent_api_participants.add_agent_chat_participant.await_count + == 2 ) - # Second request with SAME context_id - room_id_2, ctx_2 = await adapter._get_or_create_room("ctx-user-123", "weather") - create_calls_after_second = ( - adapter._rest.agent_api_chats.create_agent_chat.call_count - ) + def test_rehydrate_merges_without_overwriting_live_context(self) -> None: + adapter = A2AGatewayAdapter() + adapter._context_to_room["ctx"] = "live-room" - # Same room, same context - assert room_id_1 == room_id_2 - assert ctx_1 == ctx_2 == "ctx-user-123" - # Only created once, not twice - assert create_calls_after_first == 1 - assert create_calls_after_second == 1 + adapter._rehydrate( + GatewaySessionState( + context_to_room={"ctx": "old-room", "new": "new-room"}, + room_participants={"new-room": {"weather"}}, + ) + ) - @pytest.mark.asyncio - async def test_different_context_ids_create_different_rooms( - self, adapter_with_tracking: A2AGatewayAdapter - ) -> None: - """Different context_ids should create different rooms.""" - adapter = adapter_with_tracking + assert adapter._context_to_room == { + "ctx": "live-room", + "new": "new-room", + } + assert adapter._room_participants["new-room"] == {"weather"} - room_id_1, ctx_1 = await adapter._get_or_create_room("ctx-first", "weather") - room_id_2, ctx_2 = await adapter._get_or_create_room("ctx-second", "weather") - # Different rooms for different contexts - assert room_id_1 != room_id_2 - assert ctx_1 == "ctx-first" - assert ctx_2 == "ctx-second" - # Created twice - assert adapter._rest.agent_api_chats.create_agent_chat.call_count == 2 +class TestGatewayTranslation: + @pytest.mark.parametrize( + ("message_type", "state"), + [ + ("thought", TaskState.TASK_STATE_WORKING), + ("text", TaskState.TASK_STATE_COMPLETED), + ("error", TaskState.TASK_STATE_FAILED), + ], + ) + def test_translates_band_message_state(self, message_type: str, state: int) -> None: + adapter = A2AGatewayAdapter() + task = make_pending(EventQueueLegacy()).task - @pytest.mark.asyncio - async def test_same_context_different_peers_same_room( - self, adapter_with_tracking: A2AGatewayAdapter - ) -> None: - """Same context_id with different peers should use same room, add peers.""" - adapter = adapter_with_tracking - adapter._peers["data"] = make_peer("data", "Data Agent") - - room_id_1, _ = await adapter._get_or_create_room("ctx-multi-agent", "weather") - room_id_2, _ = await adapter._get_or_create_room("ctx-multi-agent", "data") - - # Same room - assert room_id_1 == room_id_2 - # Both peers added to room - assert "weather" in adapter._room_participants[room_id_1] - assert "data" in adapter._room_participants[room_id_1] - # Room created once, but participant added twice - assert adapter._rest.agent_api_chats.create_agent_chat.call_count == 1 - assert ( - adapter._rest.agent_api_participants.add_agent_chat_participant.call_count - == 2 + event = adapter._translate_to_a2a( + make_platform_message("response", message_type=message_type), task ) + + assert event.status.state == state + assert event.status.message.parts[0].text == "response" diff --git a/tests/integrations/a2a/gateway/test_server.py b/tests/integrations/a2a/gateway/test_server.py index 78cd7e1d4..2b06f894c 100644 --- a/tests/integrations/a2a/gateway/test_server.py +++ b/tests/integrations/a2a/gateway/test_server.py @@ -6,31 +6,24 @@ from a2a.server.agent_execution import AgentExecutor, RequestContext from a2a.server.events import EventQueue -from a2a.types import ( - AgentCard, - TaskState, - TaskStatus, - TaskStatusUpdateEvent, -) +from a2a.types import TaskState, TaskStatus, TaskStatusUpdateEvent from starlette.testclient import TestClient -from a2a.utils import new_task from band.integrations.a2a.gateway.server import GatewayServer -from tests.integrations.a2a.gateway.fixtures import make_peer +from band.integrations.a2a.protocol import new_task +from tests.integrations.a2a.gateway.helpers import make_peer class FakeExecutor(AgentExecutor): async def execute(self, context: RequestContext, event_queue: EventQueue) -> None: - task = context.current_task - if task is None: - task = new_task(context.message) + task = context.current_task or new_task(context.message) + if context.current_task is None: await event_queue.enqueue_event(task) await event_queue.enqueue_event( TaskStatusUpdateEvent( task_id=task.id, context_id=task.context_id, - status=TaskStatus(state=TaskState.completed), - final=True, + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), ) ) @@ -49,17 +42,19 @@ def build_server() -> GatewayServer: ) -def test_agent_card_is_served_by_upstream_handler() -> None: - response = TestClient(build_server()._build_app()).get( - "/agents/weather-agent/.well-known/agent.json" - ) +def test_agent_card_is_served_at_standard_and_legacy_paths() -> None: + client = TestClient(build_server()._build_app()) - assert response.status_code == 200 - card = AgentCard.model_validate(response.json()) - assert card.name == "Weather Agent" - assert card.additional_interfaces is not None - assert card.additional_interfaces[0].transport == "JSONRPC" - assert card.additional_interfaces[0].url.endswith("/agents/weather-agent") + for path in ( + "/agents/weather-agent/.well-known/agent-card.json", + "/agents/weather-agent/.well-known/agent.json", + ): + response = client.get(path) + assert response.status_code == 200 + card = response.json() + assert card["name"] == "Weather Agent" + assert card["supportedInterfaces"][0]["protocolBinding"] == "JSONRPC" + assert card["supportedInterfaces"][0]["url"].endswith("/agents/weather-agent") def test_peers_listing_remains_gateway_owned() -> None: @@ -78,14 +73,22 @@ def test_peers_listing_remains_gateway_owned() -> None: def test_unknown_peer_is_not_resolved_by_a2a_routes() -> None: response = TestClient(build_server()._build_app()).get( - "/agents/missing/.well-known/agent.json" + "/agents/missing/.well-known/agent-card.json" ) assert response.status_code == 404 +def test_uuid_peer_alias_serves_the_same_agent_card() -> None: + response = TestClient(build_server()._build_app()).get( + "/agents/uuid-weather/.well-known/agent-card.json" + ) + assert response.status_code == 200 + + def test_jsonrpc_method_errors_are_upstream_owned() -> None: response = TestClient(build_server()._build_app()).post( "/agents/weather-agent", + headers={"A2A-Version": "1.0"}, json={"jsonrpc": "2.0", "id": str(uuid4()), "method": "missing", "params": {}}, ) @@ -96,15 +99,16 @@ def test_jsonrpc_method_errors_are_upstream_owned() -> None: def test_jsonrpc_send_runs_through_official_handler_and_executor() -> None: response = TestClient(build_server()._build_app()).post( "/agents/weather-agent", + headers={"A2A-Version": "1.0"}, json={ "jsonrpc": "2.0", "id": "request-1", - "method": "message/send", + "method": "SendMessage", "params": { "message": { - "role": "user", + "role": "ROLE_USER", "messageId": "message-1", - "parts": [{"kind": "text", "text": "Hello"}], + "parts": [{"text": "Hello"}], } }, }, @@ -113,17 +117,18 @@ def test_jsonrpc_send_runs_through_official_handler_and_executor() -> None: assert response.status_code == 200 body = response.json() assert body["id"] == "request-1" - assert body["result"]["status"]["state"] == "completed" + assert body["result"]["task"]["status"]["state"] == "TASK_STATE_COMPLETED" -def test_rest_stream_runs_through_upstream_adapter() -> None: +def test_rest_stream_runs_through_upstream_handler() -> None: response = TestClient(build_server()._build_app()).post( - "/agents/weather-agent/v1/message:stream", + "/agents/weather-agent/message:stream", + headers={"A2A-Version": "1.0"}, json={ "message": { - "role": "ROLE_USER", "messageId": "message-1", - "content": [{"text": "Hello"}], + "role": "ROLE_USER", + "parts": [{"text": "Hello"}], } }, ) @@ -132,3 +137,24 @@ def test_rest_stream_runs_through_upstream_adapter() -> None: assert "text/event-stream" in response.headers["content-type"] assert '"task":' in response.text assert '"state": "TASK_STATE_COMPLETED"' in response.text + + +def test_v03_jsonrpc_stream_accepts_legacy_payload() -> None: + response = TestClient(build_server()._build_app()).post( + "/agents/weather-agent", + json={ + "jsonrpc": "2.0", + "id": "request-1", + "method": "message/stream", + "params": { + "message": { + "messageId": "message-1", + "role": "user", + "parts": [{"type": "text", "text": "Hello"}], + }, + }, + }, + ) + + assert response.status_code == 200 + assert "text/event-stream" in response.headers["content-type"] diff --git a/tests/integrations/a2a/test_adapter.py b/tests/integrations/a2a/test_adapter.py index f0d4e48b9..abb928b95 100644 --- a/tests/integrations/a2a/test_adapter.py +++ b/tests/integrations/a2a/test_adapter.py @@ -1,35 +1,33 @@ -"""Tests for A2AAdapter.""" +"""Behavior tests for the outbound A2A adapter.""" from __future__ import annotations from datetime import datetime -from typing import AsyncIterator from unittest.mock import AsyncMock, MagicMock, patch from uuid import uuid4 import pytest from a2a.types import ( Artifact, - Message as A2AMessage, Part, Role, + SendMessageRequest, + StreamResponse, Task, TaskState, TaskStatus, - TextPart, ) -from band.converters.a2a import A2AHistoryConverter from band.core.types import PlatformMessage from band.integrations.a2a import A2AAdapter, A2AAuth, A2ASessionState +from band.integrations.a2a.protocol import text_message from band.testing import FakeAgentTools -def make_platform_message(content: str, room_id: str = "room-123") -> PlatformMessage: - """Create a test PlatformMessage.""" +def make_platform_message(content: str = "Hello") -> PlatformMessage: return PlatformMessage( id=str(uuid4()), - room_id=room_id, + room_id="room-123", content=content, sender_id="user-456", sender_type="User", @@ -41,850 +39,231 @@ def make_platform_message(content: str, room_id: str = "room-123") -> PlatformMe def make_task( - state: TaskState, - task_id: str = "task-123", - context_id: str = "ctx-123", + state: int = TaskState.TASK_STATE_COMPLETED, + *, status_message: str | None = None, artifact_text: str | None = None, ) -> Task: - """Create a mock A2A Task.""" - status_msg = None + task = Task( + id="task-123", + context_id="ctx-123", + status=TaskStatus(state=state), + ) if status_message: - status_msg = A2AMessage( - role=Role.agent, - message_id=str(uuid4()), - parts=[Part(root=TextPart(text=status_message))], - ) - - artifacts = None + task.status.message.CopyFrom(text_message(status_message)) if artifact_text: - artifacts = [ - Artifact( - artifact_id=str(uuid4()), - parts=[Part(root=TextPart(text=artifact_text))], - ) - ] - - return Task( - id=task_id, - context_id=context_id, - status=TaskStatus(state=state, message=status_msg), - artifacts=artifacts, - history=None, - ) - + task.artifacts.append( + Artifact(artifact_id="artifact-1", parts=[Part(text=artifact_text)]) + ) + return task -class TestA2AAuth: - """Tests for A2AAuth.""" - def test_to_headers_with_api_key(self): - """Should add X-API-Key header.""" - auth = A2AAuth(api_key="my-secret-key") - headers = auth.to_headers() +def task_event(task: Task) -> StreamResponse: + return StreamResponse(task=task) - assert headers == {"X-API-Key": "my-secret-key"} - def test_to_headers_with_bearer_token(self): - """Should add Authorization Bearer header.""" - auth = A2AAuth(bearer_token="eyJ...") - headers = auth.to_headers() +def status_event(task: Task) -> StreamResponse: + return StreamResponse( + status_update={ + "task_id": task.id, + "context_id": task.context_id, + "status": task.status, + } + ) - assert headers == {"Authorization": "Bearer eyJ..."} - def test_to_headers_with_custom_headers(self): - """Should include custom headers.""" - auth = A2AAuth(headers={"X-Custom": "value"}) - headers = auth.to_headers() +async def stream(*events: StreamResponse): + for event in events: + yield event - assert headers == {"X-Custom": "value"} - def test_to_headers_combined(self): - """Should combine all auth methods.""" +class TestA2AAuth: + def test_to_headers_combines_authentication_methods(self) -> None: auth = A2AAuth( api_key="key", bearer_token="token", headers={"X-Custom": "value"}, ) - headers = auth.to_headers() - assert headers == { + assert auth.to_headers() == { "X-API-Key": "key", "Authorization": "Bearer token", "X-Custom": "value", } -class TestA2AAdapterInit: - """Tests for A2AAdapter initialization.""" - - def test_init_default_values(self): - """Should initialize with default values.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - - assert adapter.remote_url == "http://localhost:10000" - assert adapter.auth is None - assert adapter.streaming is True - assert adapter._client is None - assert adapter._contexts == {} - assert adapter._tasks == {} - assert adapter._task_senders == {} - - def test_init_with_auth(self): - """Should accept auth configuration.""" - auth = A2AAuth(api_key="test-key") - adapter = A2AAdapter(remote_url="http://localhost:10000", auth=auth) - - assert adapter.auth is auth - - def test_init_with_streaming_disabled(self): - """Should accept streaming=False.""" - adapter = A2AAdapter(remote_url="http://localhost:10000", streaming=False) - - assert adapter.streaming is False - - -class TestA2AAdapterOnStarted: - """Tests for A2AAdapter.on_started().""" - +class TestA2AAdapterStartup: @pytest.mark.asyncio - async def test_on_started_connects_to_agent(self): - """Should connect to remote A2A agent.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - - with patch("band.integrations.a2a.adapter.ClientFactory") as mock_factory: - mock_client = MagicMock() - mock_factory.connect = AsyncMock(return_value=mock_client) - - await adapter.on_started("Test Agent", "A test agent") - - mock_factory.connect.assert_called_once() - assert adapter._client is mock_client - - @pytest.mark.asyncio - async def test_on_started_passes_auth_to_discovery_and_requests(self): - """Should apply auth headers to discovery and RPC requests.""" - auth = A2AAuth( - api_key="test-key", - bearer_token="test-token", - headers={"X-Custom-Auth": "custom"}, + async def test_creates_client_with_auth_headers(self) -> None: + adapter = A2AAdapter( + remote_url="http://localhost:10000", + auth=A2AAuth(api_key="key"), ) - adapter = A2AAdapter(remote_url="http://localhost:10000", auth=auth) - - with patch("band.integrations.a2a.adapter.ClientFactory") as mock_factory: - mock_factory.connect = AsyncMock(return_value=MagicMock()) - - await adapter.on_started("Test Agent", "A test agent") - - _, kwargs = mock_factory.connect.call_args - assert kwargs["resolver_http_kwargs"] == { - "headers": { - "X-API-Key": "test-key", - "Authorization": "Bearer test-token", - "X-Custom-Auth": "custom", - } - } - - interceptors = kwargs["interceptors"] - assert interceptors is not None - assert len(interceptors) == 1 - - request_payload, http_kwargs = await interceptors[0].intercept( - "message/send", - {"jsonrpc": "2.0"}, - {}, - None, - None, - ) - assert request_payload == {"jsonrpc": "2.0"} - assert http_kwargs == { - "headers": { - "X-API-Key": "test-key", - "Authorization": "Bearer test-token", - "X-Custom-Auth": "custom", - } - } - - @pytest.mark.asyncio - async def test_on_started_sets_agent_name(self): - """Should store agent name and description.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") + client = MagicMock() - with patch("band.integrations.a2a.adapter.ClientFactory") as mock_factory: - mock_factory.connect = AsyncMock(return_value=MagicMock()) + with patch("band.integrations.a2a.adapter.ClientFactory") as factory_type: + factory = factory_type.return_value + factory.create_from_url = AsyncMock(return_value=client) - await adapter.on_started("Test Agent", "A test agent") + await adapter.on_started("Agent", "Description") - assert adapter.agent_name == "Test Agent" - assert adapter.agent_description == "A test agent" + assert adapter._client is client + config = factory_type.call_args.args[0] + assert config.streaming is True + assert adapter._http_client is not None + assert adapter._http_client.headers["X-API-Key"] == "key" -class TestA2AAdapterOnMessage: - """Tests for A2AAdapter.on_message().""" - +class TestA2AAdapterMessageFlow: @pytest.fixture - def adapter_with_client(self): - """Create adapter with mocked client.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - adapter._client = MagicMock() - return adapter + def adapter(self) -> A2AAdapter: + return A2AAdapter(remote_url="http://localhost:10000") @pytest.mark.asyncio - async def test_raises_if_client_not_initialized(self): - """Should raise if on_started not called.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - tools = FakeAgentTools() - msg = make_platform_message("Hello") - - with pytest.raises(RuntimeError, match="client not initialized"): - await adapter.on_message( - msg, - tools, - A2ASessionState(), - None, - None, - is_session_bootstrap=True, - room_id="room-123", + async def test_forwards_band_message_as_a2a_request( + self, adapter: A2AAdapter + ) -> None: + adapter._client = MagicMock() + adapter._client.send_message = MagicMock( + return_value=stream( + task_event(make_task(TaskState.TASK_STATE_WORKING)), + status_event(make_task()), ) - - @pytest.mark.asyncio - async def test_completed_task_sends_message(self, adapter_with_client): - """Should send message when task completes with artifact.""" - tools = FakeAgentTools() - msg = make_platform_message("What is 10 USD in EUR?") - - task = make_task( - state=TaskState.completed, - artifact_text="10 USD is approximately 9.20 EUR.", ) - - async def mock_send_message(*args, **kwargs) -> AsyncIterator: - yield (task, None) - - adapter_with_client._client.send_message = mock_send_message - - await adapter_with_client.on_message( - msg, - tools, - A2ASessionState(), - None, - None, - is_session_bootstrap=True, - room_id="room-123", - ) - - # Should have sent a message with mention - assert len(tools.messages_sent) == 1 - assert tools.messages_sent[0]["content"] == "10 USD is approximately 9.20 EUR." - assert tools.messages_sent[0]["mentions"] == [ - {"id": "user-456", "name": "Test User"} - ] - - # Context should be tracked for multi-turn - assert adapter_with_client._contexts["room-123"] == "ctx-123" - - # Task ID should be cleared after completion (new message = new task) - assert "room-123" not in adapter_with_client._tasks - - # Task sender should be cleaned up after completion - assert ("room-123", "task-123") not in adapter_with_client._task_senders - - @pytest.mark.asyncio - async def test_working_task_sends_thought_event(self, adapter_with_client): - """Should send thought event when task is working.""" tools = FakeAgentTools() - msg = make_platform_message("Processing request") - task = make_task( - state=TaskState.working, - status_message="Fetching exchange rates...", - ) - - async def mock_send_message(*args, **kwargs) -> AsyncIterator: - yield (task, None) - - adapter_with_client._client.send_message = mock_send_message - - await adapter_with_client.on_message( - msg, + await adapter.on_message( + make_platform_message("What is the weather?"), tools, A2ASessionState(), None, None, - is_session_bootstrap=True, + is_session_bootstrap=False, room_id="room-123", ) - # Should have sent a thought event - assert len(tools.events_sent) == 1 - assert tools.events_sent[0]["content"] == "Fetching exchange rates..." - assert tools.events_sent[0]["message_type"] == "thought" + request = adapter._client.send_message.call_args.args[0] + assert isinstance(request, SendMessageRequest) + assert request.message.role == Role.ROLE_USER + assert request.message.parts[0].text == "What is the weather?" @pytest.mark.asyncio - async def test_input_required_sends_message(self, adapter_with_client): - """Should send message when agent needs more input.""" + async def test_completed_task_posts_artifact_response( + self, adapter: A2AAdapter + ) -> None: tools = FakeAgentTools() - msg = make_platform_message("Convert currency") - - task = make_task( - state=TaskState.input_required, - status_message="What currency do you want to convert to?", - ) - - async def mock_send_message(*args, **kwargs) -> AsyncIterator: - yield (task, None) - adapter_with_client._client.send_message = mock_send_message - - await adapter_with_client.on_message( - msg, + await adapter._handle_event( + task_event(make_task(artifact_text="Final response")), tools, - A2ASessionState(), - None, - None, - is_session_bootstrap=True, - room_id="room-123", + "room-123", + "user-456", + "Test User", ) - # Should have sent a message asking for more info with mention - assert len(tools.messages_sent) == 1 - assert ( - tools.messages_sent[0]["content"] - == "What currency do you want to convert to?" + assert tools.messages_sent[-1]["content"] == "Final response" + assert tools.events_sent[-1]["metadata"]["a2a_task_state"] == ( + "TASK_STATE_COMPLETED" ) - assert tools.messages_sent[0]["mentions"] == [ - {"id": "user-456", "name": "Test User"} - ] @pytest.mark.asyncio - async def test_failed_task_sends_error_event(self, adapter_with_client): - """Should send error event when task fails.""" + async def test_status_update_is_applied_to_task_and_completes_flow( + self, adapter: A2AAdapter + ) -> None: tools = FakeAgentTools() - msg = make_platform_message("Hello") + task = make_task(TaskState.TASK_STATE_WORKING) - task = make_task( - state=TaskState.failed, - status_message="Currency API unavailable", + await adapter._handle_event( + task_event(task), tools, "room-123", "user-456", "Test User" ) - - async def mock_send_message(*args, **kwargs) -> AsyncIterator: - yield (task, None) - - adapter_with_client._client.send_message = mock_send_message - - await adapter_with_client.on_message( - msg, - tools, - A2ASessionState(), - None, - None, - is_session_bootstrap=True, - room_id="room-123", + task.status.CopyFrom( + TaskStatus( + state=TaskState.TASK_STATE_COMPLETED, + message=text_message("Sunny"), + ) + ) + await adapter._handle_event( + status_event(task), tools, "room-123", "user-456", "Test User" ) - # Should have sent an error event + task event - error_events = [e for e in tools.events_sent if e["message_type"] == "error"] - assert len(error_events) == 1 - assert error_events[0]["content"] == "Currency API unavailable" - assert error_events[0]["metadata"]["a2a_state"] == "failed" - - # Task event should also be emitted for rehydration - task_events = [e for e in tools.events_sent if e["message_type"] == "task"] - assert len(task_events) == 1 + assert tools.messages_sent[-1]["content"] == "Sunny" + assert adapter._tasks == {} @pytest.mark.asyncio - async def test_exception_sends_error_event(self, adapter_with_client): - """Should send error event on exception.""" + async def test_input_required_is_forwarded_and_persisted( + self, adapter: A2AAdapter + ) -> None: tools = FakeAgentTools() - msg = make_platform_message("Hello") - - async def mock_send_message(*args, **kwargs) -> AsyncIterator: - raise ConnectionError("Connection failed") - yield # Make it an async generator - adapter_with_client._client.send_message = mock_send_message - - await adapter_with_client.on_message( - msg, + await adapter._handle_event( + task_event( + make_task( + TaskState.TASK_STATE_INPUT_REQUIRED, + status_message="Which city?", + ) + ), tools, - A2ASessionState(), - None, - None, - is_session_bootstrap=True, - room_id="room-123", + "room-123", + "user-456", + "Test User", ) - # Should have sent an error event - assert len(tools.events_sent) == 1 - assert "Connection failed" in tools.events_sent[0]["content"] - assert tools.events_sent[0]["message_type"] == "error" + assert tools.messages_sent[-1]["content"] == "Which city?" + assert tools.events_sent[-1]["metadata"]["a2a_task_state"] == ( + "TASK_STATE_INPUT_REQUIRED" + ) @pytest.mark.asyncio - async def test_direct_message_reply(self, adapter_with_client): - """Should handle direct A2A Message reply.""" + async def test_direct_message_response_is_forwarded( + self, adapter: A2AAdapter + ) -> None: tools = FakeAgentTools() - msg = make_platform_message("Hello") - - a2a_reply = A2AMessage( - role=Role.agent, - message_id=str(uuid4()), - parts=[Part(root=TextPart(text="Hello! How can I help?"))], - ) - - async def mock_send_message(*args, **kwargs) -> AsyncIterator: - yield a2a_reply - - adapter_with_client._client.send_message = mock_send_message - await adapter_with_client.on_message( - msg, + await adapter._handle_event( + StreamResponse(message=text_message("Hello")), tools, - A2ASessionState(), - None, - None, - is_session_bootstrap=True, - room_id="room-123", + "room-123", + "user-456", + "Test User", ) - # Should have sent a message with mention - assert len(tools.messages_sent) == 1 - assert tools.messages_sent[0]["content"] == "Hello! How can I help?" - assert tools.messages_sent[0]["mentions"] == [ - {"id": "user-456", "name": "Test User"} - ] - - -class TestA2AAdapterContextManagement: - """Tests for context and task tracking.""" - - def test_tracks_context_per_room(self): - """Should track different contexts for different rooms.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - - adapter._contexts["room-1"] = "ctx-1" - adapter._contexts["room-2"] = "ctx-2" + assert tools.messages_sent[-1]["content"] == "Hello" - assert adapter._contexts["room-1"] == "ctx-1" - assert adapter._contexts["room-2"] == "ctx-2" - - def test_tracks_task_per_room(self): - """Should track different tasks for different rooms.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - - adapter._tasks["room-1"] = "task-1" - adapter._tasks["room-2"] = "task-2" - - assert adapter._tasks["room-1"] == "task-1" - assert adapter._tasks["room-2"] == "task-2" +class TestA2AAdapterSession: @pytest.mark.asyncio - async def test_on_cleanup_removes_context(self): - """Should clean up context, task, and sender tracking for room.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - adapter._contexts["room-1"] = "ctx-1" - adapter._tasks["room-1"] = "task-1" - adapter._task_senders[("room-1", "task-1")] = {"id": "u1", "name": "User1"} - adapter._task_senders[("room-1", "task-2")] = {"id": "u2", "name": "User2"} - adapter._contexts["room-2"] = "ctx-2" - adapter._tasks["room-2"] = "task-2" - adapter._task_senders[("room-2", "task-3")] = {"id": "u3", "name": "User3"} - - await adapter.on_cleanup("room-1") - - assert "room-1" not in adapter._contexts - assert "room-1" not in adapter._tasks - assert ("room-1", "task-1") not in adapter._task_senders - assert ("room-1", "task-2") not in adapter._task_senders - # Other room unaffected - assert adapter._contexts["room-2"] == "ctx-2" - assert adapter._tasks["room-2"] == "task-2" - assert adapter._task_senders[("room-2", "task-3")] == { - "id": "u3", - "name": "User3", - } - - -class TestA2AAdapterMessageConversion: - """Tests for message conversion.""" - - def test_to_a2a_message_basic(self): - """Should convert Band message to A2A format.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - msg = make_platform_message("Hello world") - - a2a_msg = adapter._to_a2a_message(msg, "room-123") - - assert a2a_msg.role == Role.user - assert len(a2a_msg.parts) == 1 - assert isinstance(a2a_msg.parts[0].root, TextPart) - assert a2a_msg.parts[0].root.text == "Hello world" - assert a2a_msg.context_id is None - assert a2a_msg.task_id is None - - def test_to_a2a_message_with_existing_context(self): - """Should include existing context_id and task_id.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - adapter._contexts["room-123"] = "existing-ctx" - adapter._tasks["room-123"] = "existing-task" - - msg = make_platform_message("Follow up") - - a2a_msg = adapter._to_a2a_message(msg, "room-123") - - assert a2a_msg.context_id == "existing-ctx" - assert a2a_msg.task_id == "existing-task" - - -class TestA2AAdapterResponseExtraction: - """Tests for response extraction from Task.""" - - def test_extract_response_from_artifact(self): - """Should extract response from artifact.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - task = make_task( - state=TaskState.completed, - artifact_text="Response from artifact", - ) - - response = adapter._extract_response(task) - - assert response == "Response from artifact" - - def test_extract_response_from_status_message(self): - """Should fallback to status message if no artifact.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - task = make_task( - state=TaskState.completed, - status_message="Response from status", - ) - - response = adapter._extract_response(task) - - assert response == "Response from status" - - def test_extract_response_empty_if_no_content(self): - """Should return empty string if no content found.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - task = make_task(state=TaskState.completed) - - response = adapter._extract_response(task) - - assert response == "" - - -class TestA2AHistoryConverter: - """Tests for A2AHistoryConverter.""" - - def test_convert_empty_history(self): - """Should return empty session state for empty history.""" - converter = A2AHistoryConverter() - result = converter.convert([]) - - assert result.context_id is None - assert result.task_id is None - assert result.task_state is None - - def test_convert_finds_task_event(self): - """Should extract A2A metadata from task event.""" - converter = A2AHistoryConverter() - raw_history = [ - {"message_type": "text", "content": "Hello"}, - { - "message_type": "task", - "content": "A2A task completed", - "metadata": { - "a2a_context_id": "ctx-abc", - "a2a_task_id": "task-xyz", - "a2a_task_state": "completed", - }, - }, - ] - - result = converter.convert(raw_history) - - assert result.context_id == "ctx-abc" - assert result.task_id == "task-xyz" - assert result.task_state == "completed" - - def test_convert_finds_latest_task_event(self): - """Should find most recent A2A task event.""" - converter = A2AHistoryConverter() - raw_history = [ - { - "message_type": "task", - "metadata": { - "a2a_context_id": "ctx-old", - "a2a_task_id": "task-old", - "a2a_task_state": "completed", - }, - }, - {"message_type": "text", "content": "New message"}, - { - "message_type": "task", - "metadata": { - "a2a_context_id": "ctx-new", - "a2a_task_id": "task-new", - "a2a_task_state": "input_required", - }, - }, - ] - - result = converter.convert(raw_history) - - assert result.context_id == "ctx-new" - assert result.task_id == "task-new" - assert result.task_state == "input_required" - - def test_convert_ignores_non_a2a_task_events(self): - """Should ignore task events without A2A metadata.""" - converter = A2AHistoryConverter() - raw_history = [ - { - "message_type": "task", - "metadata": {"other_key": "value"}, # No A2A metadata - }, - ] - - result = converter.convert(raw_history) - - assert result.context_id is None - assert result.task_id is None - assert result.task_state is None - - -class TestA2AAdapterTaskEventEmission: - """Tests for task event emission.""" - - @pytest.fixture - def adapter_with_client(self): - """Create adapter with mocked client.""" + async def test_rehydrates_context_and_resubscribes_active_task(self) -> None: adapter = A2AAdapter(remote_url="http://localhost:10000") adapter._client = MagicMock() - return adapter - - @pytest.mark.asyncio - async def test_emits_task_event_on_completed(self, adapter_with_client): - """Should emit task event when task completes.""" - tools = FakeAgentTools() - msg = make_platform_message("Hello") - - task = make_task( - state=TaskState.completed, - artifact_text="Response", + adapter._client.subscribe = MagicMock( + return_value=stream(task_event(make_task(TaskState.TASK_STATE_WORKING))) ) - async def mock_send_message(*args, **kwargs): - yield (task, None) - - adapter_with_client._client.send_message = mock_send_message - - await adapter_with_client.on_message( - msg, - tools, - A2ASessionState(), - None, - None, - is_session_bootstrap=True, - room_id="room-123", + await adapter._rehydrate_from_history( + "room-123", + A2ASessionState( + context_id="ctx-123", + task_id="task-123", + task_state="TASK_STATE_WORKING", + ), ) - # Find the task event - task_events = [e for e in tools.events_sent if e["message_type"] == "task"] - assert len(task_events) == 1 - assert task_events[0]["metadata"]["a2a_context_id"] == "ctx-123" - assert task_events[0]["metadata"]["a2a_task_id"] == "task-123" - assert task_events[0]["metadata"]["a2a_task_state"] == "completed" + assert adapter._contexts["room-123"] == "ctx-123" + assert adapter._tasks["room-123"] == "task-123" @pytest.mark.asyncio - async def test_emits_task_event_on_input_required(self, adapter_with_client): - """Should emit task event when input is required.""" - tools = FakeAgentTools() - msg = make_platform_message("Hello") - - task = make_task( - state=TaskState.input_required, - status_message="What currency?", - ) - - async def mock_send_message(*args, **kwargs): - yield (task, None) - - adapter_with_client._client.send_message = mock_send_message - - await adapter_with_client.on_message( - msg, - tools, - A2ASessionState(), - None, - None, - is_session_bootstrap=True, - room_id="room-123", - ) - - # Find the task event - task_events = [e for e in tools.events_sent if e["message_type"] == "task"] - assert len(task_events) == 1 - assert task_events[0]["metadata"]["a2a_task_state"] == "input-required" - - -class TestA2AAdapterSessionRehydration: - """Tests for session rehydration.""" - - @pytest.fixture - def adapter_with_client(self): - """Create adapter with mocked client.""" + async def test_does_not_resubscribe_terminal_task(self) -> None: adapter = A2AAdapter(remote_url="http://localhost:10000") adapter._client = MagicMock() - return adapter - - @pytest.mark.asyncio - async def test_rehydrates_context_on_bootstrap(self, adapter_with_client): - """Should restore context_id on session bootstrap.""" - tools = FakeAgentTools() - msg = make_platform_message("Hello") - - # Create history with A2A session state - history = A2ASessionState( - context_id="restored-ctx", - task_id="old-task", - task_state="completed", # Terminal state, won't try to resubscribe - ) - - task = make_task(state=TaskState.completed, artifact_text="Response") - - async def mock_send_message(*args, **kwargs): - yield (task, None) - - adapter_with_client._client.send_message = mock_send_message - - await adapter_with_client.on_message( - msg, - tools, - history, - None, - None, - is_session_bootstrap=True, - room_id="room-123", - ) - - # Context should be restored before the message was processed - # Note: The new task will update context_id, so we check the initial restoration happened - assert "room-123" in adapter_with_client._contexts - - @pytest.mark.asyncio - async def test_no_rehydration_on_non_bootstrap(self, adapter_with_client): - """Should not rehydrate on non-bootstrap messages.""" - tools = FakeAgentTools() - msg = make_platform_message("Hello") - - # History would normally trigger rehydration - history = A2ASessionState( - context_id="restored-ctx", - task_id="old-task", - task_state="input_required", - ) - - task = make_task(state=TaskState.completed, artifact_text="Response") - - async def mock_send_message(*args, **kwargs): - yield (task, None) - - adapter_with_client._client.send_message = mock_send_message - adapter_with_client._client.resubscribe = MagicMock() - - await adapter_with_client.on_message( - msg, - tools, - history, - None, - None, - is_session_bootstrap=False, - room_id="room-123", - ) - - # Resubscribe should not be called on non-bootstrap - adapter_with_client._client.resubscribe.assert_not_called() - - @pytest.mark.asyncio - async def test_resubscribes_to_resumable_task(self, adapter_with_client): - """Should try to resubscribe to non-terminal task.""" - tools = FakeAgentTools() - msg = make_platform_message("Hello") - - # History with resumable task - history = A2ASessionState( - context_id="ctx-123", - task_id="resumable-task", - task_state="input_required", # Non-terminal - ) - - # Mock resubscribe to return current task state - resumed_task = make_task( - state=TaskState.input_required, - task_id="resumable-task", - ) - - async def mock_resubscribe(*args, **kwargs): - yield (resumed_task, None) - - adapter_with_client._client.resubscribe = mock_resubscribe - - # Mock send_message for the actual message processing - task = make_task(state=TaskState.completed, artifact_text="Response") - - async def mock_send_message(*args, **kwargs): - yield (task, None) - - adapter_with_client._client.send_message = mock_send_message - - await adapter_with_client.on_message( - msg, - tools, - history, - None, - None, - is_session_bootstrap=True, - room_id="room-123", - ) - - # Task should have been restored - # Note: The new completed task clears it, so we verify resubscribe was called - assert adapter_with_client._contexts.get("room-123") == "ctx-123" - - @pytest.mark.asyncio - async def test_handles_resubscribe_failure(self, adapter_with_client): - """Should handle resubscribe failure gracefully.""" - tools = FakeAgentTools() - msg = make_platform_message("Hello") - - history = A2ASessionState( - context_id="ctx-123", - task_id="old-task", - task_state="input_required", - ) - - async def mock_resubscribe(*args, **kwargs): - raise Exception("Task not found") - yield # Make it an async generator - - adapter_with_client._client.resubscribe = mock_resubscribe - - task = make_task(state=TaskState.completed, artifact_text="Response") - - async def mock_send_message(*args, **kwargs): - yield (task, None) - - adapter_with_client._client.send_message = mock_send_message - - # Should not raise - await adapter_with_client.on_message( - msg, - tools, - history, - None, - None, - is_session_bootstrap=True, - room_id="room-123", + adapter._client.subscribe = MagicMock() + + await adapter._rehydrate_from_history( + "room-123", + A2ASessionState( + context_id="ctx-123", + task_id="task-123", + task_state="TASK_STATE_COMPLETED", + ), ) - # Context should still be restored - assert adapter_with_client._contexts.get("room-123") == "ctx-123" + adapter._client.subscribe.assert_not_called() From 91b6fd954a6300e5fcfa8b6fc90955772ec95927 Mon Sep 17 00:00:00 2001 From: Alexander Zaikman Date: Sun, 26 Jul 2026 14:21:03 +0300 Subject: [PATCH 3/7] style(tests): format A2A gateway test --- tests/integrations/a2a/gateway/test_adapter.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/integrations/a2a/gateway/test_adapter.py b/tests/integrations/a2a/gateway/test_adapter.py index 8b554b59a..9a677293e 100644 --- a/tests/integrations/a2a/gateway/test_adapter.py +++ b/tests/integrations/a2a/gateway/test_adapter.py @@ -138,7 +138,9 @@ async def test_retries_peer_discovery_only_for_rate_limits(self) -> None: class TestGatewayExecution: @pytest.mark.asyncio - async def test_initial_task_snapshot_stays_working_if_reply_is_immediate(self) -> None: + async def test_initial_task_snapshot_stays_working_if_reply_is_immediate( + self, + ) -> None: adapter = A2AGatewayAdapter() adapter._peers = {"weather": make_peer("weather", "Weather Agent")} configure_room_creation(adapter) From 4a7c12d0e8161c1ccfaf88610cebdd85b6d22386 Mon Sep 17 00:00:00 2001 From: Alexander Zaikman Date: Sun, 26 Jul 2026 14:24:21 +0300 Subject: [PATCH 4/7] refactor(integrations): share A2A stream reducer --- .../demo_orchestrator/remote_agent.py | 37 +++-------- src/band/integrations/a2a/adapter.py | 52 ++++------------ src/band/integrations/a2a/protocol.py | 56 +++++++++++++++++ tests/integrations/a2a/test_protocol.py | 61 +++++++++++++++++++ 4 files changed, 135 insertions(+), 71 deletions(-) create mode 100644 tests/integrations/a2a/test_protocol.py diff --git a/examples/a2a_gateway/demo_orchestrator/remote_agent.py b/examples/a2a_gateway/demo_orchestrator/remote_agent.py index 2b520155a..2aeb54196 100644 --- a/examples/a2a_gateway/demo_orchestrator/remote_agent.py +++ b/examples/a2a_gateway/demo_orchestrator/remote_agent.py @@ -9,7 +9,11 @@ from a2a.client import ClientConfig, ClientFactory from a2a.types import Message, Part, Role, SendMessageRequest, Task -from band.integrations.a2a.protocol import text_from_message +from band.integrations.a2a.protocol import ( + apply_task_stream_event, + task_response_text, + text_from_message, +) logger = logging.getLogger(__name__) @@ -77,40 +81,13 @@ async def call_peer( async for event in client.send_message(request): if event.HasField("message"): return text_from_message(event.message) - if event.HasField("task"): - task = Task() - task.CopyFrom(event.task) - elif event.HasField("status_update"): - if task is None: - task = Task( - id=event.status_update.task_id, - context_id=event.status_update.context_id, - ) - task.status.CopyFrom(event.status_update.status) - elif event.HasField("artifact_update"): - if task is None: - task = Task( - id=event.artifact_update.task_id, - context_id=event.artifact_update.context_id, - ) - task.artifacts.add().CopyFrom(event.artifact_update.artifact) - return self._extract_response(task) if task else "No response from peer" + task = apply_task_stream_event(task, event) or task + return task_response_text(task) or "No response from peer" except Exception as exc: error_msg = f"Failed to call peer '{peer_id}': {exc}" logger.error(error_msg) raise RuntimeError(error_msg) from exc - def _extract_response(self, task: Task) -> str: - for artifact in task.artifacts: - text = "\n".join(part.text for part in artifact.parts if part.text) - if text: - return text - if task.status.message: - text = text_from_message(task.status.message) - if text: - return text - return "No response from peer" - async def __aenter__(self) -> GatewayClient: return self diff --git a/src/band/integrations/a2a/adapter.py b/src/band/integrations/a2a/adapter.py index 24007b911..212e8adfc 100644 --- a/src/band/integrations/a2a/adapter.py +++ b/src/band/integrations/a2a/adapter.py @@ -23,7 +23,10 @@ from band.core.types import AdapterFeatures, Capability, Emit, PlatformMessage from band.integrations.a2a.protocol import ( TERMINAL_TASK_STATES, + apply_task_stream_event, state_name, + task_id_from_stream_event, + task_response_text, text_from_message, text_message, ) @@ -178,26 +181,14 @@ async def _handle_event( ) return - if event.HasField("task"): - task = event.task - self._task_cache[(room_id, task.id)] = Task() - self._task_cache[(room_id, task.id)].CopyFrom(task) - elif event.HasField("status_update"): - update = event.status_update - task = self._task_cache.setdefault( - (room_id, update.task_id), - Task(id=update.task_id, context_id=update.context_id), - ) - task.status.CopyFrom(update.status) - elif event.HasField("artifact_update"): - update = event.artifact_update - task = self._task_cache.setdefault( - (room_id, update.task_id), - Task(id=update.task_id, context_id=update.context_id), - ) - task.artifacts.add().CopyFrom(update.artifact) - else: + task_id = task_id_from_stream_event(event) + if task_id is None: return + key = (room_id, task_id) + task = apply_task_stream_event(self._task_cache.get(key), event) + if task is None: + return + self._task_cache[(room_id, task.id)] = task key = (room_id, task.id) @@ -297,28 +288,7 @@ def _extract_response(self, task: Task) -> str: 2. Status message 3. Last agent message in history """ - # First: check artifacts - if task.artifacts: - for artifact in task.artifacts: - for part in artifact.parts: - if part.text: - return part.text - - # Fallback: check status message - if task.status.message: - text = text_from_message(task.status.message) - if text: - return text - - # Last resort: check history for last agent message - if task.history: - for msg in reversed(task.history): - if msg.role == Role.ROLE_AGENT: - text = text_from_message(msg) - if text: - return text - - return "" + return task_response_text(task) async def on_cleanup(self, room_id: str) -> None: """Clean up A2A context for room.""" diff --git a/src/band/integrations/a2a/protocol.py b/src/band/integrations/a2a/protocol.py index 8d9bfef25..4b83d9fef 100644 --- a/src/band/integrations/a2a/protocol.py +++ b/src/band/integrations/a2a/protocol.py @@ -10,6 +10,7 @@ Part, Role, SendMessageRequest, + StreamResponse, Task, TaskState, TaskStatus, @@ -68,6 +69,61 @@ def snapshot_task(task: Task) -> Task: return deepcopy(task) +def apply_task_stream_event(task: Task | None, event: StreamResponse) -> Task | None: + """Apply one task-bearing stream event and return the current task state.""" + if event.HasField("task"): + return snapshot_task(event.task) + + if event.HasField("status_update"): + update = event.status_update + task = task or Task(id=update.task_id, context_id=update.context_id) + task.status.CopyFrom(update.status) + return task + + if event.HasField("artifact_update"): + update = event.artifact_update + task = task or Task(id=update.task_id, context_id=update.context_id) + task.artifacts.add().CopyFrom(update.artifact) + return task + + return None + + +def task_id_from_stream_event(event: StreamResponse) -> str | None: + """Return the task ID carried by a task-related stream event.""" + if event.HasField("task"): + return event.task.id + if event.HasField("status_update"): + return event.status_update.task_id + if event.HasField("artifact_update"): + return event.artifact_update.task_id + return None + + +def task_response_text(task: Task | None) -> str: + """Extract a task's best available text response.""" + if task is None: + return "" + + for artifact in task.artifacts: + text = "\n".join(part.text for part in artifact.parts if part.text) + if text: + return text + + if task.status.message: + text = text_from_message(task.status.message) + if text: + return text + + for message in reversed(task.history): + if message.role == Role.ROLE_AGENT: + text = text_from_message(message) + if text: + return text + + return "" + + def is_terminal_state(state: int) -> bool: """Return whether an A2A task state ends execution.""" return state in TERMINAL_TASK_STATES diff --git a/tests/integrations/a2a/test_protocol.py b/tests/integrations/a2a/test_protocol.py new file mode 100644 index 000000000..5695eac74 --- /dev/null +++ b/tests/integrations/a2a/test_protocol.py @@ -0,0 +1,61 @@ +"""Tests for shared protobuf A2A protocol helpers.""" + +from __future__ import annotations + +from a2a.types import Artifact, Part, StreamResponse, Task, TaskState, TaskStatus + +from band.integrations.a2a.protocol import ( + apply_task_stream_event, + task_id_from_stream_event, + task_response_text, + text_message, +) + + +def test_task_stream_updates_build_a_task_from_deltas() -> None: + status = TaskStatus( + state=TaskState.TASK_STATE_COMPLETED, + message=text_message("Sunny"), + ) + status_event = StreamResponse( + status_update={ + "task_id": "task-123", + "context_id": "context-123", + "status": status, + } + ) + artifact_event = StreamResponse( + artifact_update={ + "task_id": "task-123", + "context_id": "context-123", + "artifact": Artifact( + artifact_id="artifact-123", + parts=[Part(text="Detailed forecast")], + ), + } + ) + + task = apply_task_stream_event(None, status_event) + task = apply_task_stream_event(task, artifact_event) + + assert task is not None + assert task_id_from_stream_event(status_event) == "task-123" + assert task.status.state == TaskState.TASK_STATE_COMPLETED + assert task_response_text(task) == "Detailed forecast" + + +def test_task_stream_snapshot_does_not_alias_the_event() -> None: + event = StreamResponse( + task=Task( + id="task-123", + context_id="context-123", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + ) + + task = apply_task_stream_event(None, event) + event.task.status.state = TaskState.TASK_STATE_COMPLETED + + assert task is not None + assert task_id_from_stream_event(event) == "task-123" + assert task.status.state == TaskState.TASK_STATE_WORKING From c45f27f1ca9d5589adc46b909a23795afc49ac05 Mon Sep 17 00:00:00 2001 From: Alexander Zaikman Date: Sun, 26 Jul 2026 14:28:59 +0300 Subject: [PATCH 5/7] fix(integrations): preserve A2A streaming compatibility --- src/band/integrations/a2a/gateway/server.py | 40 ++++++++-- src/band/integrations/a2a/protocol.py | 15 +++- tests/integrations/a2a/gateway/test_server.py | 73 +++++++++++-------- tests/integrations/a2a/test_adapter.py | 60 +++++++++++++++ tests/integrations/a2a/test_protocol.py | 34 +++++++++ 5 files changed, 185 insertions(+), 37 deletions(-) diff --git a/src/band/integrations/a2a/gateway/server.py b/src/band/integrations/a2a/gateway/server.py index 82a2a8208..1050c144c 100644 --- a/src/band/integrations/a2a/gateway/server.py +++ b/src/band/integrations/a2a/gateway/server.py @@ -4,7 +4,7 @@ import asyncio import logging -from collections.abc import Callable +from collections.abc import Awaitable, Callable from typing import Any from a2a.server.agent_execution import AgentExecutor @@ -14,7 +14,10 @@ from a2a.server.routes.rest_routes import create_rest_routes from a2a.server.tasks import InMemoryTaskStore from a2a.types import AgentCapabilities, AgentCard, AgentInterface, AgentSkill +from a2a.compat.v0_3.conversions import to_compat_agent_card +from a2a.utils.constants import PROTOCOL_VERSION_0_3 from starlette.applications import Starlette +from starlette.responses import JSONResponse from starlette.routing import BaseRoute, Route from band_rest import Peer @@ -68,7 +71,12 @@ def _agent_card(self, slug: str, peer: Peer) -> AgentCard: protocol_binding="JSONRPC", protocol_version="1.0", url=rpc_url, - ) + ), + AgentInterface( + protocol_binding="JSONRPC", + protocol_version=PROTOCOL_VERSION_0_3, + url=rpc_url, + ), ], version="1.0.0", capabilities=AgentCapabilities(streaming=True), @@ -108,10 +116,13 @@ def _build_app(self) -> Starlette: ) ) protocol_routes.extend( - create_agent_card_routes( - card, - card_url=f"/agents/{alias}/.well-known/agent.json", - ) + [ + Route( + f"/agents/{alias}/.well-known/agent.json", + self._legacy_agent_card(card), + methods=["GET"], + ) + ] ) protocol_routes.extend( create_jsonrpc_routes( @@ -130,6 +141,23 @@ def _build_app(self) -> Starlette: return Starlette(routes=routes + protocol_routes + rest_routes) + @staticmethod + def _legacy_agent_card( + card: AgentCard, + ) -> Callable[[Any], Awaitable[JSONResponse]]: + """Serve the SDK's v0.3 card representation for legacy discovery.""" + legacy_card = to_compat_agent_card(card) + payload = legacy_card.model_dump( + by_alias=True, + mode="json", + exclude_none=True, + ) + + async def response(_request: Any) -> JSONResponse: + return JSONResponse(payload) + + return response + async def _handle_list_peers(self, _request: Any) -> Any: from starlette.responses import JSONResponse diff --git a/src/band/integrations/a2a/protocol.py b/src/band/integrations/a2a/protocol.py index 4b83d9fef..d091609f0 100644 --- a/src/band/integrations/a2a/protocol.py +++ b/src/band/integrations/a2a/protocol.py @@ -83,7 +83,20 @@ def apply_task_stream_event(task: Task | None, event: StreamResponse) -> Task | if event.HasField("artifact_update"): update = event.artifact_update task = task or Task(id=update.task_id, context_id=update.context_id) - task.artifacts.add().CopyFrom(update.artifact) + existing_artifact = next( + ( + artifact + for artifact in task.artifacts + if artifact.artifact_id == update.artifact.artifact_id + ), + None, + ) + if existing_artifact is None: + task.artifacts.add().CopyFrom(update.artifact) + elif update.append: + existing_artifact.parts.extend(update.artifact.parts) + else: + existing_artifact.CopyFrom(update.artifact) return task return None diff --git a/tests/integrations/a2a/gateway/test_server.py b/tests/integrations/a2a/gateway/test_server.py index 2b06f894c..66f92d614 100644 --- a/tests/integrations/a2a/gateway/test_server.py +++ b/tests/integrations/a2a/gateway/test_server.py @@ -2,11 +2,14 @@ from __future__ import annotations +from collections.abc import Iterator from uuid import uuid4 +import pytest from a2a.server.agent_execution import AgentExecutor, RequestContext from a2a.server.events import EventQueue from a2a.types import TaskState, TaskStatus, TaskStatusUpdateEvent +from a2a.utils.constants import PROTOCOL_VERSION_0_3 from starlette.testclient import TestClient from band.integrations.a2a.gateway.server import GatewayServer @@ -42,23 +45,35 @@ def build_server() -> GatewayServer: ) -def test_agent_card_is_served_at_standard_and_legacy_paths() -> None: - client = TestClient(build_server()._build_app()) +@pytest.fixture +def gateway_client() -> Iterator[TestClient]: + with TestClient(build_server()._build_app()) as client: + yield client - for path in ( - "/agents/weather-agent/.well-known/agent-card.json", - "/agents/weather-agent/.well-known/agent.json", - ): - response = client.get(path) - assert response.status_code == 200 - card = response.json() - assert card["name"] == "Weather Agent" - assert card["supportedInterfaces"][0]["protocolBinding"] == "JSONRPC" - assert card["supportedInterfaces"][0]["url"].endswith("/agents/weather-agent") + +def test_agent_cards_use_the_schema_expected_by_each_protocol_version( + gateway_client: TestClient, +) -> None: + standard = gateway_client.get("/agents/weather-agent/.well-known/agent-card.json") + assert standard.status_code == 200 + standard_card = standard.json() + assert standard_card["name"] == "Weather Agent" + assert standard_card["supportedInterfaces"][0]["protocolBinding"] == "JSONRPC" + assert standard_card["supportedInterfaces"][0]["url"].endswith( + "/agents/weather-agent" + ) + + legacy = gateway_client.get("/agents/weather-agent/.well-known/agent.json") + assert legacy.status_code == 200 + legacy_card = legacy.json() + assert legacy_card["name"] == "Weather Agent" + assert legacy_card["protocolVersion"] == PROTOCOL_VERSION_0_3 + assert legacy_card["url"].endswith("/agents/weather-agent") + assert "supportedInterfaces" not in legacy_card -def test_peers_listing_remains_gateway_owned() -> None: - response = TestClient(build_server()._build_app()).get("/peers") +def test_peers_listing_remains_gateway_owned(gateway_client: TestClient) -> None: + response = gateway_client.get("/peers") assert response.status_code == 200 assert response.json()["peers"] == [ @@ -71,22 +86,18 @@ def test_peers_listing_remains_gateway_owned() -> None: ] -def test_unknown_peer_is_not_resolved_by_a2a_routes() -> None: - response = TestClient(build_server()._build_app()).get( - "/agents/missing/.well-known/agent-card.json" - ) +def test_unknown_peer_is_not_resolved_by_a2a_routes(gateway_client: TestClient) -> None: + response = gateway_client.get("/agents/missing/.well-known/agent-card.json") assert response.status_code == 404 -def test_uuid_peer_alias_serves_the_same_agent_card() -> None: - response = TestClient(build_server()._build_app()).get( - "/agents/uuid-weather/.well-known/agent-card.json" - ) +def test_uuid_peer_alias_serves_the_same_agent_card(gateway_client: TestClient) -> None: + response = gateway_client.get("/agents/uuid-weather/.well-known/agent-card.json") assert response.status_code == 200 -def test_jsonrpc_method_errors_are_upstream_owned() -> None: - response = TestClient(build_server()._build_app()).post( +def test_jsonrpc_method_errors_are_upstream_owned(gateway_client: TestClient) -> None: + response = gateway_client.post( "/agents/weather-agent", headers={"A2A-Version": "1.0"}, json={"jsonrpc": "2.0", "id": str(uuid4()), "method": "missing", "params": {}}, @@ -96,8 +107,10 @@ def test_jsonrpc_method_errors_are_upstream_owned() -> None: assert response.json()["error"]["code"] == -32601 -def test_jsonrpc_send_runs_through_official_handler_and_executor() -> None: - response = TestClient(build_server()._build_app()).post( +def test_jsonrpc_send_runs_through_official_handler_and_executor( + gateway_client: TestClient, +) -> None: + response = gateway_client.post( "/agents/weather-agent", headers={"A2A-Version": "1.0"}, json={ @@ -120,8 +133,8 @@ def test_jsonrpc_send_runs_through_official_handler_and_executor() -> None: assert body["result"]["task"]["status"]["state"] == "TASK_STATE_COMPLETED" -def test_rest_stream_runs_through_upstream_handler() -> None: - response = TestClient(build_server()._build_app()).post( +def test_rest_stream_runs_through_upstream_handler(gateway_client: TestClient) -> None: + response = gateway_client.post( "/agents/weather-agent/message:stream", headers={"A2A-Version": "1.0"}, json={ @@ -139,8 +152,8 @@ def test_rest_stream_runs_through_upstream_handler() -> None: assert '"state": "TASK_STATE_COMPLETED"' in response.text -def test_v03_jsonrpc_stream_accepts_legacy_payload() -> None: - response = TestClient(build_server()._build_app()).post( +def test_v03_jsonrpc_stream_accepts_legacy_payload(gateway_client: TestClient) -> None: + response = gateway_client.post( "/agents/weather-agent", json={ "jsonrpc": "2.0", diff --git a/tests/integrations/a2a/test_adapter.py b/tests/integrations/a2a/test_adapter.py index abb928b95..76cdc823e 100644 --- a/tests/integrations/a2a/test_adapter.py +++ b/tests/integrations/a2a/test_adapter.py @@ -72,6 +72,27 @@ def status_event(task: Task) -> StreamResponse: ) +def artifact_event( + task: Task, + text: str, + *, + append: bool, + last_chunk: bool, +) -> StreamResponse: + return StreamResponse( + artifact_update={ + "task_id": task.id, + "context_id": task.context_id, + "artifact": Artifact( + artifact_id="artifact-123", + parts=[Part(text=text)], + ), + "append": append, + "last_chunk": last_chunk, + } + ) + + async def stream(*events: StreamResponse): for event in events: yield event @@ -166,6 +187,45 @@ async def test_completed_task_posts_artifact_response( "TASK_STATE_COMPLETED" ) + @pytest.mark.asyncio + async def test_streamed_artifact_chunks_are_posted_as_one_response( + self, adapter: A2AAdapter + ) -> None: + working = make_task(TaskState.TASK_STATE_WORKING) + completed = make_task(TaskState.TASK_STATE_COMPLETED) + adapter._client = MagicMock() + adapter._client.send_message = MagicMock( + return_value=stream( + task_event(working), + artifact_event( + working, + "Part one. ", + append=False, + last_chunk=False, + ), + artifact_event( + working, + "Part two.", + append=True, + last_chunk=True, + ), + status_event(completed), + ) + ) + tools = FakeAgentTools() + + await adapter.on_message( + make_platform_message(), + tools, + A2ASessionState(), + None, + None, + is_session_bootstrap=False, + room_id="room-123", + ) + + assert tools.messages_sent[-1]["content"] == "Part one. \nPart two." + @pytest.mark.asyncio async def test_status_update_is_applied_to_task_and_completes_flow( self, adapter: A2AAdapter diff --git a/tests/integrations/a2a/test_protocol.py b/tests/integrations/a2a/test_protocol.py index 5695eac74..31b3e9145 100644 --- a/tests/integrations/a2a/test_protocol.py +++ b/tests/integrations/a2a/test_protocol.py @@ -44,6 +44,40 @@ def test_task_stream_updates_build_a_task_from_deltas() -> None: assert task_response_text(task) == "Detailed forecast" +def test_appended_artifact_chunks_are_combined_before_response_extraction() -> None: + first_chunk = StreamResponse( + artifact_update={ + "task_id": "task-123", + "context_id": "context-123", + "artifact": Artifact( + artifact_id="artifact-123", + parts=[Part(text="Part one. ")], + ), + "append": False, + "last_chunk": False, + } + ) + final_chunk = StreamResponse( + artifact_update={ + "task_id": "task-123", + "context_id": "context-123", + "artifact": Artifact( + artifact_id="artifact-123", + parts=[Part(text="Part two.")], + ), + "append": True, + "last_chunk": True, + } + ) + + task = apply_task_stream_event(None, first_chunk) + task = apply_task_stream_event(task, final_chunk) + + assert task is not None + assert len(task.artifacts) == 1 + assert task_response_text(task) == "Part one. \nPart two." + + def test_task_stream_snapshot_does_not_alias_the_event() -> None: event = StreamResponse( task=Task( From f1646c168a03110f232c52c893bfc4033dcb40ef Mon Sep 17 00:00:00 2001 From: Alexander Zaikman Date: Sun, 26 Jul 2026 14:31:10 +0300 Subject: [PATCH 6/7] test(integrations): use async gateway transport --- tests/integrations/a2a/gateway/test_server.py | 74 ++++++++++++------- 1 file changed, 48 insertions(+), 26 deletions(-) diff --git a/tests/integrations/a2a/gateway/test_server.py b/tests/integrations/a2a/gateway/test_server.py index 66f92d614..c1bcd2308 100644 --- a/tests/integrations/a2a/gateway/test_server.py +++ b/tests/integrations/a2a/gateway/test_server.py @@ -2,15 +2,16 @@ from __future__ import annotations -from collections.abc import Iterator +from collections.abc import AsyncIterator from uuid import uuid4 -import pytest +import httpx +import pytest_asyncio from a2a.server.agent_execution import AgentExecutor, RequestContext from a2a.server.events import EventQueue from a2a.types import TaskState, TaskStatus, TaskStatusUpdateEvent from a2a.utils.constants import PROTOCOL_VERSION_0_3 -from starlette.testclient import TestClient +from httpx import ASGITransport from band.integrations.a2a.gateway.server import GatewayServer from band.integrations.a2a.protocol import new_task @@ -19,7 +20,11 @@ class FakeExecutor(AgentExecutor): async def execute(self, context: RequestContext, event_queue: EventQueue) -> None: - task = context.current_task or new_task(context.message) + task = context.current_task + if task is None: + if context.message is None: + raise ValueError("A2A request is missing its message") + task = new_task(context.message) if context.current_task is None: await event_queue.enqueue_event(task) await event_queue.enqueue_event( @@ -45,16 +50,19 @@ def build_server() -> GatewayServer: ) -@pytest.fixture -def gateway_client() -> Iterator[TestClient]: - with TestClient(build_server()._build_app()) as client: +@pytest_asyncio.fixture +async def gateway_client() -> AsyncIterator[httpx.AsyncClient]: + transport = ASGITransport(app=build_server()._build_app()) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: yield client -def test_agent_cards_use_the_schema_expected_by_each_protocol_version( - gateway_client: TestClient, +async def test_agent_cards_use_the_schema_expected_by_each_protocol_version( + gateway_client: httpx.AsyncClient, ) -> None: - standard = gateway_client.get("/agents/weather-agent/.well-known/agent-card.json") + standard = await gateway_client.get( + "/agents/weather-agent/.well-known/agent-card.json" + ) assert standard.status_code == 200 standard_card = standard.json() assert standard_card["name"] == "Weather Agent" @@ -63,7 +71,7 @@ def test_agent_cards_use_the_schema_expected_by_each_protocol_version( "/agents/weather-agent" ) - legacy = gateway_client.get("/agents/weather-agent/.well-known/agent.json") + legacy = await gateway_client.get("/agents/weather-agent/.well-known/agent.json") assert legacy.status_code == 200 legacy_card = legacy.json() assert legacy_card["name"] == "Weather Agent" @@ -72,8 +80,10 @@ def test_agent_cards_use_the_schema_expected_by_each_protocol_version( assert "supportedInterfaces" not in legacy_card -def test_peers_listing_remains_gateway_owned(gateway_client: TestClient) -> None: - response = gateway_client.get("/peers") +async def test_peers_listing_remains_gateway_owned( + gateway_client: httpx.AsyncClient, +) -> None: + response = await gateway_client.get("/peers") assert response.status_code == 200 assert response.json()["peers"] == [ @@ -86,18 +96,26 @@ def test_peers_listing_remains_gateway_owned(gateway_client: TestClient) -> None ] -def test_unknown_peer_is_not_resolved_by_a2a_routes(gateway_client: TestClient) -> None: - response = gateway_client.get("/agents/missing/.well-known/agent-card.json") +async def test_unknown_peer_is_not_resolved_by_a2a_routes( + gateway_client: httpx.AsyncClient, +) -> None: + response = await gateway_client.get("/agents/missing/.well-known/agent-card.json") assert response.status_code == 404 -def test_uuid_peer_alias_serves_the_same_agent_card(gateway_client: TestClient) -> None: - response = gateway_client.get("/agents/uuid-weather/.well-known/agent-card.json") +async def test_uuid_peer_alias_serves_the_same_agent_card( + gateway_client: httpx.AsyncClient, +) -> None: + response = await gateway_client.get( + "/agents/uuid-weather/.well-known/agent-card.json" + ) assert response.status_code == 200 -def test_jsonrpc_method_errors_are_upstream_owned(gateway_client: TestClient) -> None: - response = gateway_client.post( +async def test_jsonrpc_method_errors_are_upstream_owned( + gateway_client: httpx.AsyncClient, +) -> None: + response = await gateway_client.post( "/agents/weather-agent", headers={"A2A-Version": "1.0"}, json={"jsonrpc": "2.0", "id": str(uuid4()), "method": "missing", "params": {}}, @@ -107,10 +125,10 @@ def test_jsonrpc_method_errors_are_upstream_owned(gateway_client: TestClient) -> assert response.json()["error"]["code"] == -32601 -def test_jsonrpc_send_runs_through_official_handler_and_executor( - gateway_client: TestClient, +async def test_jsonrpc_send_runs_through_official_handler_and_executor( + gateway_client: httpx.AsyncClient, ) -> None: - response = gateway_client.post( + response = await gateway_client.post( "/agents/weather-agent", headers={"A2A-Version": "1.0"}, json={ @@ -133,8 +151,10 @@ def test_jsonrpc_send_runs_through_official_handler_and_executor( assert body["result"]["task"]["status"]["state"] == "TASK_STATE_COMPLETED" -def test_rest_stream_runs_through_upstream_handler(gateway_client: TestClient) -> None: - response = gateway_client.post( +async def test_rest_stream_runs_through_upstream_handler( + gateway_client: httpx.AsyncClient, +) -> None: + response = await gateway_client.post( "/agents/weather-agent/message:stream", headers={"A2A-Version": "1.0"}, json={ @@ -152,8 +172,10 @@ def test_rest_stream_runs_through_upstream_handler(gateway_client: TestClient) - assert '"state": "TASK_STATE_COMPLETED"' in response.text -def test_v03_jsonrpc_stream_accepts_legacy_payload(gateway_client: TestClient) -> None: - response = gateway_client.post( +async def test_v03_jsonrpc_stream_accepts_legacy_payload( + gateway_client: httpx.AsyncClient, +) -> None: + response = await gateway_client.post( "/agents/weather-agent", json={ "jsonrpc": "2.0", From 74602ec7e7257183a748acf8531ebf6432af8092 Mon Sep 17 00:00:00 2001 From: Alexander Zaikman Date: Sun, 26 Jul 2026 14:52:12 +0300 Subject: [PATCH 7/7] refactor(integrations): lean on A2A SDK helpers --- .../demo_orchestrator/agent_executor.py | 19 +- .../demo_orchestrator/remote_agent.py | 4 +- examples/mixed/03_fact_checker_a2a.py | 12 +- examples/mixed/04_risk_reviewer_a2a.py | 12 +- src/band/integrations/a2a/adapter.py | 164 +++++---- src/band/integrations/a2a/gateway/adapter.py | 319 +++++++----------- src/band/integrations/a2a/gateway/server.py | 15 - src/band/integrations/a2a/gateway/types.py | 56 ++- src/band/integrations/a2a/protocol.py | 58 +--- .../integrations/a2a/gateway/test_adapter.py | 45 +-- tests/integrations/a2a/gateway/test_server.py | 4 +- tests/integrations/a2a/test_adapter.py | 25 +- tests/integrations/a2a/test_protocol.py | 4 +- 13 files changed, 309 insertions(+), 428 deletions(-) diff --git a/examples/a2a_gateway/demo_orchestrator/agent_executor.py b/examples/a2a_gateway/demo_orchestrator/agent_executor.py index 07a051a23..42ce98ee0 100644 --- a/examples/a2a_gateway/demo_orchestrator/agent_executor.py +++ b/examples/a2a_gateway/demo_orchestrator/agent_executor.py @@ -8,6 +8,7 @@ import logging +from a2a.helpers import new_task_from_user_message, new_text_message from a2a.server.agent_execution import AgentExecutor, RequestContext from a2a.server.events import EventQueue from a2a.server.tasks import TaskUpdater @@ -17,7 +18,6 @@ TaskState, UnsupportedOperationError, ) -from band.integrations.a2a.protocol import new_task, text_message try: from .agent import OrchestratorAgent @@ -56,10 +56,10 @@ async def execute( query = context.get_user_input() task = context.current_task - if not task: + if task is None: if context.message is None: raise ValueError("A2A request is missing its message") - task = new_task(context.message) + task = new_task_from_user_message(context.message) await event_queue.enqueue_event(task) updater = TaskUpdater(event_queue, task.id, task.context_id) @@ -72,19 +72,16 @@ async def execute( if not is_task_complete and not require_user_input: # Working status update - await updater.update_status( - TaskState.TASK_STATE_WORKING, - text_message( - content, - context_id=task.context_id, - task_id=task.id, - ), + await updater.start_work( + new_text_message( + content, context_id=task.context_id, task_id=task.id + ) ) elif require_user_input: # Need more input from user await updater.update_status( TaskState.TASK_STATE_INPUT_REQUIRED, - text_message( + new_text_message( content, context_id=task.context_id, task_id=task.id, diff --git a/examples/a2a_gateway/demo_orchestrator/remote_agent.py b/examples/a2a_gateway/demo_orchestrator/remote_agent.py index 2aeb54196..36a06aa38 100644 --- a/examples/a2a_gateway/demo_orchestrator/remote_agent.py +++ b/examples/a2a_gateway/demo_orchestrator/remote_agent.py @@ -7,12 +7,12 @@ import httpx from a2a.client import ClientConfig, ClientFactory +from a2a.helpers import get_message_text from a2a.types import Message, Part, Role, SendMessageRequest, Task from band.integrations.a2a.protocol import ( apply_task_stream_event, task_response_text, - text_from_message, ) logger = logging.getLogger(__name__) @@ -80,7 +80,7 @@ async def call_peer( task: Task | None = None async for event in client.send_message(request): if event.HasField("message"): - return text_from_message(event.message) + return get_message_text(event.message) task = apply_task_stream_event(task, event) or task return task_response_text(task) or "No response from peer" except Exception as exc: diff --git a/examples/mixed/03_fact_checker_a2a.py b/examples/mixed/03_fact_checker_a2a.py index ab626c397..98f3118a2 100644 --- a/examples/mixed/03_fact_checker_a2a.py +++ b/examples/mixed/03_fact_checker_a2a.py @@ -24,6 +24,7 @@ import os import uvicorn +from a2a.helpers import new_task_from_user_message, new_text_message from a2a.server.agent_execution import AgentExecutor, RequestContext from a2a.server.events import EventQueue from a2a.server.request_handlers import DefaultRequestHandler @@ -42,10 +43,8 @@ AgentInterface, AgentSkill, Part, - TaskState, UnsupportedOperationError, ) -from band.integrations.a2a.protocol import new_task, text_message from dotenv import load_dotenv from setup_logging import setup_logging @@ -79,16 +78,15 @@ async def execute( request_text = context.get_user_input() task = context.current_task - if not task: + if task is None: if context.message is None: raise ValueError("A2A request is missing its message") - task = new_task(context.message) + task = new_task_from_user_message(context.message) await event_queue.enqueue_event(task) updater = TaskUpdater(event_queue, task.id, task.context_id) - await updater.update_status( - TaskState.TASK_STATE_WORKING, - text_message( + await updater.start_work( + new_text_message( "Reviewing the request for API, config, and test-surface details...", context_id=task.context_id, task_id=task.id, diff --git a/examples/mixed/04_risk_reviewer_a2a.py b/examples/mixed/04_risk_reviewer_a2a.py index eb69e12fb..29b8a922b 100644 --- a/examples/mixed/04_risk_reviewer_a2a.py +++ b/examples/mixed/04_risk_reviewer_a2a.py @@ -24,6 +24,7 @@ import os import uvicorn +from a2a.helpers import new_task_from_user_message, new_text_message from a2a.server.agent_execution import AgentExecutor, RequestContext from a2a.server.events import EventQueue from a2a.server.request_handlers import DefaultRequestHandler @@ -42,10 +43,8 @@ AgentInterface, AgentSkill, Part, - TaskState, UnsupportedOperationError, ) -from band.integrations.a2a.protocol import new_task, text_message from dotenv import load_dotenv from setup_logging import setup_logging @@ -80,16 +79,15 @@ async def execute( request_text = context.get_user_input() task = context.current_task - if not task: + if task is None: if context.message is None: raise ValueError("A2A request is missing its message") - task = new_task(context.message) + task = new_task_from_user_message(context.message) await event_queue.enqueue_event(task) updater = TaskUpdater(event_queue, task.id, task.context_id) - await updater.update_status( - TaskState.TASK_STATE_WORKING, - text_message( + await updater.start_work( + new_text_message( "Reviewing the request for rollout, compatibility, and rollback risks...", context_id=task.context_id, task_id=task.id, diff --git a/src/band/integrations/a2a/adapter.py b/src/band/integrations/a2a/adapter.py index 212e8adfc..b9f295ca1 100644 --- a/src/band/integrations/a2a/adapter.py +++ b/src/band/integrations/a2a/adapter.py @@ -7,6 +7,7 @@ import httpx from a2a.client import Client, ClientConfig, ClientFactory +from a2a.helpers import get_message_text, new_text_message from a2a.types import SendMessageRequest from a2a.types import ( Message as A2AMessage, @@ -27,8 +28,6 @@ state_name, task_id_from_stream_event, task_response_text, - text_from_message, - text_message, ) from band.integrations.a2a.types import A2AAuth, A2ASessionState @@ -173,89 +172,92 @@ async def _handle_event( ) -> None: """Handle A2A event and forward to Band platform.""" if getattr(event, "HasField", lambda _name: False)("message"): - text = text_from_message(event.message) - if text: - await tools.send_message( - content=text, - mentions=[{"id": sender_id, "name": sender_name or ""}], - ) + await self._deliver_message(event.message, tools, sender_id, sender_name) return - task_id = task_id_from_stream_event(event) - if task_id is None: - return - key = (room_id, task_id) - task = apply_task_stream_event(self._task_cache.get(key), event) + task = self._reduce_task_event(room_id, event) if task is None: return - self._task_cache[(room_id, task.id)] = task - key = (room_id, task.id) + self._remember_task(room_id, task) + self._task_senders.setdefault(key, {"id": sender_id, "name": sender_name or ""}) - # Store sender info on first event for this task - if key not in self._task_senders: - self._task_senders[key] = {"id": sender_id, "name": sender_name or ""} + await self._deliver_task_update(task, tools, self._task_senders[key]) + if task.status.state == TaskState.TASK_STATE_INPUT_REQUIRED: + await self._emit_task_event(tools, task, task.status.state) + elif task.status.state in TERMINAL_TASK_STATES: + await self._emit_task_event(tools, task, task.status.state) + self._finalize_task(room_id, task.id) - # Track task/context for multi-turn and resumption + async def _deliver_message( + self, + message: A2AMessage, + tools: AgentToolsProtocol, + sender_id: str, + sender_name: str | None, + ) -> None: + """Forward a direct A2A message to its Band sender.""" + text = get_message_text(message) + if text: + await tools.send_message( + content=text, + mentions=[{"id": sender_id, "name": sender_name or ""}], + ) + + def _reduce_task_event(self, room_id: str, event: StreamResponse) -> Task | None: + """Reduce a raw task delta and retain its current snapshot.""" + task_id = task_id_from_stream_event(event) + if task_id is None: + return None + task = apply_task_stream_event(self._task_cache.get((room_id, task_id)), event) + if task is not None: + self._task_cache[(room_id, task.id)] = task + return task + + def _remember_task(self, room_id: str, task: Task) -> None: + """Record task identity needed for subsequent turns and resumption.""" self._tasks[room_id] = task.id if task.context_id: self._contexts[room_id] = task.context_id + async def _deliver_task_update( + self, + task: Task, + tools: AgentToolsProtocol, + sender: dict[str, str], + ) -> None: + """Translate one reduced A2A task state into Band output.""" state = task.status.state + if state == TaskState.TASK_STATE_WORKING: + status_text = self._get_status_text(task) + if status_text: + await tools.send_event(content=status_text, message_type="thought") + return - try: - # Handle based on task state - if state == TaskState.TASK_STATE_WORKING: - # Stream progress as thought event (no mentions needed for events) - status_text = self._get_status_text(task) - if status_text: - await tools.send_event(content=status_text, message_type="thought") - - elif state == TaskState.TASK_STATE_INPUT_REQUIRED: - # Agent needs more info - send as message with mention - text = self._get_status_text(task) or "Please provide more information." - sender = self._task_senders.get(key) - await tools.send_message( - content=text, - mentions=[sender] if sender else None, - ) - # Emit task event for rehydration (input_required is resumable) - await self._emit_task_event(tools, task, state) - - elif state == TaskState.TASK_STATE_COMPLETED: - # Extract and send final response with mention - response = self._extract_response(task) - if response: - sender = self._task_senders.get(key) - await tools.send_message( - content=response, - mentions=[sender] if sender else None, - ) + if state == TaskState.TASK_STATE_INPUT_REQUIRED: + text = self._get_status_text(task) or "Please provide more information." + await tools.send_message(content=text, mentions=[sender]) + return - elif state in ( - TaskState.TASK_STATE_FAILED, - TaskState.TASK_STATE_CANCELED, - TaskState.TASK_STATE_REJECTED, - TaskState.TASK_STATE_AUTH_REQUIRED, - ): - # Error states - send as error event (no mentions needed) - error_text = self._get_status_text(task) or f"Task {state_name(state)}" - await tools.send_event( - content=error_text, - message_type="error", - metadata={"a2a_state": state_name(state)}, - ) - finally: - # Clean up on terminal states - if state in TERMINAL_TASK_STATES: - # Emit task event for rehydration (records final state) - await self._emit_task_event(tools, task, state) - # Clean up sender tracking - self._task_senders.pop(key, None) - self._task_cache.pop(key, None) - # Clear task_id so next message starts a new task - # (context_id is preserved for multi-turn conversation) - self._tasks.pop(room_id, None) + if state == TaskState.TASK_STATE_COMPLETED: + response = self._extract_response(task) + if response: + await tools.send_message(content=response, mentions=[sender]) + return + + if state in TERMINAL_TASK_STATES: + error_text = self._get_status_text(task) or f"Task {state_name(state)}" + await tools.send_event( + content=error_text, + message_type="error", + metadata={"a2a_state": state_name(state)}, + ) + + def _finalize_task(self, room_id: str, task_id: str) -> None: + """Release a terminal task after its Band output and state are persisted.""" + self._task_senders.pop((room_id, task_id), None) + self._task_cache.pop((room_id, task_id), None) + self._tasks.pop(room_id, None) def _to_a2a_message(self, msg: PlatformMessage, room_id: str) -> A2AMessage: """Convert Band message to A2A format.""" @@ -267,7 +269,7 @@ def _to_a2a_message(self, msg: PlatformMessage, room_id: str) -> A2AMessage: context_id, task_id, ) - return text_message( + return new_text_message( msg.content, role=Role.ROLE_USER, context_id=context_id, @@ -277,7 +279,7 @@ def _to_a2a_message(self, msg: PlatformMessage, room_id: str) -> A2AMessage: def _get_status_text(self, task: Task) -> str | None: """Extract text from task status message.""" if task.status.message: - return text_from_message(task.status.message) + return get_message_text(task.status.message) return None def _extract_response(self, task: Task) -> str: @@ -365,23 +367,13 @@ async def _try_resubscribe(self, room_id: str, task_id: str) -> None: async for event in self._client.subscribe( SubscribeToTaskRequest(id=task_id) ): - if event.HasField("task"): - task = event.task - elif event.HasField("status_update"): - update = event.status_update - task = self._task_cache.setdefault( - (room_id, update.task_id), - Task(id=update.task_id, context_id=update.context_id), - ) - task.status.CopyFrom(update.status) - else: + task = self._reduce_task_event(room_id, event) + if task is None: continue current_state = task.status.state if current_state not in TERMINAL_TASK_STATES: - self._tasks[room_id] = task_id - if task.context_id: - self._contexts[room_id] = task.context_id + self._remember_task(room_id, task) logger.info( "Resumed A2A task %s (state=%s)", task_id, diff --git a/src/band/integrations/a2a/gateway/adapter.py b/src/band/integrations/a2a/gateway/adapter.py index c0223c404..fa9237b09 100644 --- a/src/band/integrations/a2a/gateway/adapter.py +++ b/src/band/integrations/a2a/gateway/adapter.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +from dataclasses import dataclass from functools import partial import logging import re @@ -13,12 +14,7 @@ from a2a.server.agent_execution import AgentExecutor, RequestContext from a2a.server.events import EventQueue -from a2a.types import ( - Task, - TaskState, - TaskStatus, - TaskStatusUpdateEvent, -) +from a2a.types import Task, TaskState, TaskStatus from band.client.rest import ( AsyncRestClient, @@ -36,21 +32,22 @@ from band.integrations.a2a.gateway.server import GatewayServer from band.integrations.a2a.gateway.config import A2AGatewayAdapterConfig from band.integrations.a2a.gateway.types import GatewaySessionState, PendingA2ATask -from band.integrations.a2a.protocol import ( - is_terminal_state, - snapshot_task, - text_from_message, - text_message, -) +from band.integrations.a2a.protocol import snapshot_task from band_rest import Peer -from band_rest.agent_api_peers.types.list_agent_peers_response import ( - ListAgentPeersResponse, -) -from band_rest.core.api_error import ApiError logger = logging.getLogger(__name__) +@dataclass +class GatewayRequest: + """Band routing and A2A state for one gateway request.""" + + peer: Peer + room_id: str + context_id: str + pending: PendingA2ATask + + class BandAgentExecutor(AgentExecutor): """Adapt one official A2A handler execution to a Band peer.""" @@ -155,13 +152,6 @@ def __init__( # Request/response correlation self._pending_tasks: dict[str, PendingA2ATask] = {} # room_id → task - self._peer_discovery_retry_delays_seconds: tuple[float, ...] = ( - 1.0, - 2.0, - 4.0, - 8.0, - 16.0, - ) async def on_started(self, agent_name: str, agent_description: str) -> None: """Fetch peers via REST and start HTTP server. @@ -173,7 +163,7 @@ async def on_started(self, agent_name: str, agent_description: str) -> None: await super().on_started(agent_name, agent_description) # Fetch ALL peers at startup using REST client (with pagination) - all_peers = await self._fetch_all_peers_with_retry() + all_peers = await self._fetch_all_peers() # Build slug and UUID mappings for peer in all_peers: @@ -195,16 +185,17 @@ async def on_started(self, agent_name: str, agent_description: str) -> None: logger.info("Gateway HTTP server started on port %d", self.port) - async def _fetch_all_peers_with_retry(self) -> list[Peer]: - """Fetch all peer pages, retrying if the platform rate-limits startup.""" + async def _fetch_all_peers(self) -> list[Peer]: + """Fetch every peer page using the REST client's retry policy.""" all_peers: list[Peer] = [] page = 1 page_size = 100 while True: - response = await self._list_peers_page_with_retry( + response = await self._rest.agent_api_peers.list_agent_peers( page=page, page_size=page_size, + request_options=DEFAULT_REQUEST_OPTIONS, ) all_peers.extend(response.data) @@ -212,37 +203,6 @@ async def _fetch_all_peers_with_retry(self) -> list[Peer]: return all_peers page += 1 - async def _list_peers_page_with_retry( - self, *, page: int, page_size: int - ) -> ListAgentPeersResponse: - """Fetch one peer page with explicit backoff for live 429s.""" - attempts = len(self._peer_discovery_retry_delays_seconds) + 1 - for attempt, delay in enumerate( - (0.0, *self._peer_discovery_retry_delays_seconds), start=1 - ): - if delay > 0: - logger.warning( - "Rate limited discovering peers for gateway; retrying page %s in %.1fs " - "(attempt %s/%s)", - page, - delay, - attempt, - attempts, - ) - await asyncio.sleep(delay) - - try: - return await self._rest.agent_api_peers.list_agent_peers( - page=page, - page_size=page_size, - request_options=DEFAULT_REQUEST_OPTIONS, - ) - except ApiError as exc: - if exc.status_code != 429 or attempt == attempts: - raise - - raise RuntimeError("Peer discovery retry loop exited unexpectedly") - async def on_message( self, msg: PlatformMessage, @@ -284,8 +244,7 @@ async def on_message( pending.task.id, msg.message_type, ) - event = self._translate_to_a2a(msg, pending.task) - await pending.publish_response(event) + await self._publish_band_response(pending, msg) else: logger.debug( "Ignoring Band message without pending A2A task: room=%s", room_id @@ -299,26 +258,7 @@ async def on_cleanup(self, room_id: str) -> None: """ pending = self._pending_tasks.pop(room_id, None) if pending: - if not is_terminal_state(pending.task.status.state): - pending.task.status.CopyFrom( - TaskStatus( - state=TaskState.TASK_STATE_FAILED, - message=text_message( - "Band room closed before the A2A response completed", - context_id=pending.task.context_id, - task_id=pending.task.id, - ), - ) - ) - await pending.publish_response( - TaskStatusUpdateEvent( - task_id=pending.task.id, - context_id=pending.task.context_id, - status=pending.task.status, - ) - ) - else: - pending.done.set() + await pending.fail("Band room closed before the A2A response completed") logger.debug("Cleaned up gateway resources for room %s", room_id) async def stop(self) -> None: @@ -355,93 +295,107 @@ async def _execute_a2a( self, peer_id: str, context: RequestContext, event_queue: EventQueue ) -> None: """Bridge one official A2A execution to a Band room.""" - peer = self._resolve_peer(peer_id) - if not peer: - logger.warning("A2A request target not found: peer=%s", peer_id) - raise ValueError(f"Peer not found: {peer_id}") - - peer_uuid = peer.id - room_id, context_id = await self._get_or_create_room( - context.context_id, peer_uuid - ) - task = self._make_task(context) - pending = PendingA2ATask( - task=task, - event_queue=event_queue, - peer_id=peer_uuid, - done=asyncio.Event(), - ) + request = await self._establish_request(peer_id, context, event_queue) logger.info( "A2A request started: peer=%s room=%s context=%s task=%s", peer_id, - room_id, - context_id, - task.id, + request.room_id, + request.context_id, + request.pending.task.id, ) try: - async with self.pending_task(room_id, pending): - await event_queue.enqueue_event(snapshot_task(task)) - await self._emit_context_event(room_id, context_id) - content = text_from_message(context.message) - - await self._rest.agent_api_messages.create_agent_chat_message( - chat_id=room_id, - message=ChatMessageRequest( - content=f"@{peer.name} {content}", - mentions=[ - ChatMessageRequestMentionsItem(id=peer_uuid, name=peer.name) - ], - ), - request_options=DEFAULT_REQUEST_OPTIONS, - ) - logger.debug( - "A2A request sent to Band: room=%s task=%s", - room_id, - task.id, - ) - try: - if self.config.response_timeout_s is None: - await pending.done.wait() - else: - async with asyncio.timeout(self.config.response_timeout_s): - await pending.done.wait() - except TimeoutError: - logger.warning( - "A2A response timed out: room=%s task=%s timeout=%ss", - room_id, - task.id, - self.config.response_timeout_s, - ) - task.status.CopyFrom( - TaskStatus( - state=TaskState.TASK_STATE_FAILED, - message=text_message( - "Timed out waiting for a Band response", - context_id=task.context_id, - task_id=task.id, - ), - ) - ) - await event_queue.enqueue_event( - TaskStatusUpdateEvent( - task_id=task.id, - context_id=task.context_id, - status=task.status, - ) - ) + async with self.pending_task(request.room_id, request.pending): + await self._announce_request(request) + await self._send_to_band(request, context) + await self._await_response(request) except asyncio.CancelledError: - logger.debug("A2A request cancelled: room=%s task=%s", room_id, task.id) + logger.debug( + "A2A request cancelled: room=%s task=%s", + request.room_id, + request.pending.task.id, + ) raise except Exception: logger.exception( "A2A request failed: room=%s context=%s task=%s", - room_id, - context_id, - task.id, + request.room_id, + request.context_id, + request.pending.task.id, ) raise else: - logger.info("A2A request completed: room=%s task=%s", room_id, task.id) + logger.info( + "A2A request completed: room=%s task=%s", + request.room_id, + request.pending.task.id, + ) + + async def _establish_request( + self, peer_id: str, context: RequestContext, event_queue: EventQueue + ) -> GatewayRequest: + """Resolve the peer and create its Band room and A2A task state.""" + peer = self._resolve_peer(peer_id) + if not peer: + logger.warning("A2A request target not found: peer=%s", peer_id) + raise ValueError(f"Peer not found: {peer_id}") + + room_id, context_id = await self._get_or_create_room( + context.context_id, peer.id + ) + task = self._make_task(context) + return GatewayRequest( + peer=peer, + room_id=room_id, + context_id=context_id, + pending=PendingA2ATask(task=task, event_queue=event_queue), + ) + + async def _announce_request(self, request: GatewayRequest) -> None: + """Publish the initial working task and retain its Band context.""" + await request.pending.event_queue.enqueue_event( + snapshot_task(request.pending.task) + ) + await self._emit_context_event(request.room_id, request.context_id) + + async def _send_to_band( + self, request: GatewayRequest, context: RequestContext + ) -> None: + """Send the A2A request text to the selected Band peer.""" + content = context.get_user_input() + await self._rest.agent_api_messages.create_agent_chat_message( + chat_id=request.room_id, + message=ChatMessageRequest( + content=f"@{request.peer.name} {content}", + mentions=[ + ChatMessageRequestMentionsItem( + id=request.peer.id, name=request.peer.name + ) + ], + ), + request_options=DEFAULT_REQUEST_OPTIONS, + ) + logger.debug( + "A2A request sent to Band: room=%s task=%s", + request.room_id, + request.pending.task.id, + ) + + async def _await_response(self, request: GatewayRequest) -> None: + """Wait for a terminal Band reply, publishing timeout failure if needed.""" + try: + if self.config.response_timeout_s is None: + await request.pending.done.wait() + else: + async with asyncio.timeout(self.config.response_timeout_s): + await request.pending.done.wait() + except TimeoutError: + logger.warning( + "A2A response timed out: room=%s task=%s timeout=%ss", + request.room_id, + request.pending.task.id, + self.config.response_timeout_s, + ) + await request.pending.fail("Timed out waiting for a Band response") @asynccontextmanager async def pending_task( @@ -467,14 +421,7 @@ async def _cancel_a2a( ) -> None: """Publish the official terminal cancellation event.""" task = context.current_task or self._make_task(context) - task.status.CopyFrom(TaskStatus(state=TaskState.TASK_STATE_CANCELED)) - await event_queue.enqueue_event( - TaskStatusUpdateEvent( - task_id=task.id, - context_id=task.context_id, - status=task.status, - ) - ) + await PendingA2ATask(task=task, event_queue=event_queue).cancel() async def _get_or_create_room( self, context_id: str | None, target_peer_id: str @@ -563,46 +510,16 @@ def _rehydrate(self, history: GatewaySessionState) -> None: len(self._room_participants), ) - def _translate_to_a2a( - self, msg: PlatformMessage, task: Task - ) -> TaskStatusUpdateEvent: - """Convert platform message to A2A TaskStatusUpdateEvent. - - Args: - msg: Platform message from peer. - task: Associated A2A task. - - Returns: - TaskStatusUpdateEvent for SSE streaming. - """ - # Determine task state based on message type - message_type = getattr(msg, "message_type", "text") - - if message_type == "error": - state = TaskState.TASK_STATE_FAILED - elif message_type in ("thought", "tool_call", "tool_result"): - state = TaskState.TASK_STATE_WORKING + async def _publish_band_response( + self, pending: PendingA2ATask, msg: PlatformMessage + ) -> None: + """Translate Band's message category into an A2A task intent.""" + if msg.message_type == "error": + await pending.fail(msg.content) + elif msg.message_type in ("thought", "tool_call", "tool_result"): + await pending.report_progress(msg.content) else: - # Regular text message = completed response - state = TaskState.TASK_STATE_COMPLETED - - # Update task status - task.status.CopyFrom( - TaskStatus( - state=state, - message=text_message( - msg.content, - context_id=task.context_id, - task_id=task.id, - ), - ) - ) - - return TaskStatusUpdateEvent( - task_id=task.id, - context_id=task.context_id, - status=task.status, - ) + await pending.complete_with_message(msg.content) async def _emit_context_event(self, room_id: str, context_id: str) -> None: """Emit a task event to persist context mapping in history. diff --git a/src/band/integrations/a2a/gateway/server.py b/src/band/integrations/a2a/gateway/server.py index 1050c144c..a4b6c79bf 100644 --- a/src/band/integrations/a2a/gateway/server.py +++ b/src/band/integrations/a2a/gateway/server.py @@ -46,21 +46,6 @@ def __init__( self._app: Starlette | None = None self._server_task: asyncio.Task[Any] | None = None - def _resolve_peer(self, peer_id: str) -> tuple[str, Peer] | None: - if peer_id in self.peers: - return peer_id, self.peers[peer_id] - peer = self.peers_by_uuid.get(peer_id) - if peer is None: - return None - return next( - ( - (slug, candidate) - for slug, candidate in self.peers.items() - if candidate.id == peer.id - ), - None, - ) - def _agent_card(self, slug: str, peer: Peer) -> AgentCard: rpc_url = f"{self.gateway_url}/agents/{slug}" return AgentCard( diff --git a/src/band/integrations/a2a/gateway/types.py b/src/band/integrations/a2a/gateway/types.py index b3eb3fcb2..082ab444c 100644 --- a/src/band/integrations/a2a/gateway/types.py +++ b/src/band/integrations/a2a/gateway/types.py @@ -6,9 +6,9 @@ from dataclasses import dataclass, field from a2a.server.events import EventQueue -from a2a.types import Task, TaskStatusUpdateEvent - -from band.integrations.a2a.protocol import is_terminal_state +from a2a.server.tasks import TaskUpdater +from a2a.helpers import new_text_message +from a2a.types import Task @dataclass @@ -38,18 +38,54 @@ class PendingA2ATask: Attributes: task: The A2A Task object tracking this request. event_queue: Official A2A event queue owned by DefaultRequestHandler. - peer_id: The target peer this request is for. done: Set when the final Band reply has been emitted or the room is cleaned up. """ task: Task event_queue: EventQueue - peer_id: str - done: asyncio.Event + done: asyncio.Event = field(default_factory=asyncio.Event) + _updater: TaskUpdater = field(init=False, repr=False) + _lock: asyncio.Lock = field(default_factory=asyncio.Lock, init=False, repr=False) + + def __post_init__(self) -> None: + self._updater = TaskUpdater( + self.event_queue, self.task.id, self.task.context_id + ) + + async def report_progress(self, content: str) -> None: + """Publish a non-terminal progress update from Band.""" + async with self._lock: + if not self.done.is_set(): + await self._updater.start_work(self._message(content)) + + async def complete_with_message(self, content: str) -> None: + """Publish Band's final response and release the request.""" + async with self._lock: + if self.done.is_set(): + return + await self._updater.complete(self._message(content)) + self.done.set() + + async def fail(self, reason: str) -> None: + """Publish a terminal failure and release the request.""" + async with self._lock: + if self.done.is_set(): + return + await self._updater.failed(self._message(reason)) + self.done.set() - async def publish_response(self, event: TaskStatusUpdateEvent) -> None: - """Publish a response and release the executor on terminal events.""" - await self.event_queue.enqueue_event(event) - if is_terminal_state(event.status.state): + async def cancel(self) -> None: + """Publish a terminal cancellation and release the request.""" + async with self._lock: + if self.done.is_set(): + return + await self._updater.cancel() self.done.set() + + def _message(self, content: str): + return new_text_message( + content, + context_id=self.task.context_id, + task_id=self.task.id, + ) diff --git a/src/band/integrations/a2a/protocol.py b/src/band/integrations/a2a/protocol.py index d091609f0..b980b266d 100644 --- a/src/band/integrations/a2a/protocol.py +++ b/src/band/integrations/a2a/protocol.py @@ -3,18 +3,9 @@ from __future__ import annotations from copy import deepcopy -from uuid import uuid4 - -from a2a.types import ( - Message, - Part, - Role, - SendMessageRequest, - StreamResponse, - Task, - TaskState, - TaskStatus, -) + +from a2a.helpers import get_artifact_text, get_message_text +from a2a.types import Role, StreamResponse, Task, TaskState TERMINAL_TASK_STATES = frozenset( { @@ -27,43 +18,6 @@ ) -def text_from_message(message: Message | None) -> str: - """Return the text parts from a protobuf A2A message.""" - if message is None: - return "" - return "\n".join(part.text for part in message.parts if part.text) - - -def text_message( - content: str, - *, - role: Role = Role.ROLE_AGENT, - context_id: str | None = None, - task_id: str | None = None, -) -> Message: - """Build a text-only protobuf A2A message.""" - message = Message( - message_id=str(uuid4()), - role=role, - parts=[Part(text=content)], - ) - if context_id: - message.context_id = context_id - if task_id: - message.task_id = task_id - return message - - -def new_task(request: SendMessageRequest | Message) -> Task: - """Create a working task from an incoming A2A message.""" - message = request.message if isinstance(request, SendMessageRequest) else request - return Task( - id=message.task_id or str(uuid4()), - context_id=message.context_id or str(uuid4()), - status=TaskStatus(state=TaskState.TASK_STATE_WORKING), - ) - - def snapshot_task(task: Task) -> Task: """Return an independent task snapshot for event queue publication.""" return deepcopy(task) @@ -119,18 +73,18 @@ def task_response_text(task: Task | None) -> str: return "" for artifact in task.artifacts: - text = "\n".join(part.text for part in artifact.parts if part.text) + text = get_artifact_text(artifact) if text: return text if task.status.message: - text = text_from_message(task.status.message) + text = get_message_text(task.status.message) if text: return text for message in reversed(task.history): if message.role == Role.ROLE_AGENT: - text = text_from_message(message) + text = get_message_text(message) if text: return text diff --git a/tests/integrations/a2a/gateway/test_adapter.py b/tests/integrations/a2a/gateway/test_adapter.py index 9a677293e..6d4880b57 100644 --- a/tests/integrations/a2a/gateway/test_adapter.py +++ b/tests/integrations/a2a/gateway/test_adapter.py @@ -21,11 +21,11 @@ ) from band.core.types import PlatformMessage +from band.client.rest import DEFAULT_REQUEST_OPTIONS from band.integrations.a2a.gateway import A2AGatewayAdapter, A2AGatewayAdapterConfig from band.integrations.a2a.gateway.adapter import BandAgentExecutor from band.integrations.a2a.gateway.types import GatewaySessionState, PendingA2ATask from band.testing import FakeAgentTools -from band_rest.core.api_error import ApiError from tests.integrations.a2a.gateway.helpers import make_peer @@ -73,8 +73,6 @@ def make_pending(event_queue: EventQueueLegacy) -> PendingA2ATask: status=TaskStatus(state=TaskState.TASK_STATE_WORKING), ), event_queue=event_queue, - peer_id="weather", - done=asyncio.Event(), ) @@ -112,29 +110,13 @@ async def test_discovers_peers_and_starts_server(self) -> None: assert adapter._peers["weather-agent"].id == "weather" server.start.assert_awaited_once() - - @pytest.mark.asyncio - async def test_retries_peer_discovery_only_for_rate_limits(self) -> None: - adapter = A2AGatewayAdapter() - response = MagicMock() - response.data = [] - adapter._rest.agent_api_peers.list_agent_peers = AsyncMock( - side_effect=[ApiError(status_code=429, headers={}, body=""), response] + assert ( + adapter._rest.agent_api_peers.list_agent_peers.call_args.kwargs[ + "request_options" + ] + == DEFAULT_REQUEST_OPTIONS ) - with ( - patch("band.integrations.a2a.gateway.adapter.GatewayServer") as server_type, - patch( - "band.integrations.a2a.gateway.adapter.asyncio.sleep", - new=AsyncMock(), - ) as sleep, - ): - server_type.return_value.start = AsyncMock() - await adapter.on_started("Gateway", "A2A Gateway") - - assert adapter._rest.agent_api_peers.list_agent_peers.await_count == 2 - sleep.assert_awaited_once() - class TestGatewayExecution: @pytest.mark.asyncio @@ -337,7 +319,7 @@ def test_rehydrate_merges_without_overwriting_live_context(self) -> None: assert adapter._room_participants["new-room"] == {"weather"} -class TestGatewayTranslation: +class TestGatewayResponses: @pytest.mark.parametrize( ("message_type", "state"), [ @@ -346,13 +328,18 @@ class TestGatewayTranslation: ("error", TaskState.TASK_STATE_FAILED), ], ) - def test_translates_band_message_state(self, message_type: str, state: int) -> None: + async def test_publishes_band_message_with_matching_task_state( + self, message_type: str, state: int + ) -> None: adapter = A2AGatewayAdapter() - task = make_pending(EventQueueLegacy()).task + queue = EventQueueLegacy() + pending = make_pending(queue) - event = adapter._translate_to_a2a( - make_platform_message("response", message_type=message_type), task + await adapter._publish_band_response( + pending, + make_platform_message("response", message_type=message_type), ) + event = await queue.dequeue_event() assert event.status.state == state assert event.status.message.parts[0].text == "response" diff --git a/tests/integrations/a2a/gateway/test_server.py b/tests/integrations/a2a/gateway/test_server.py index c1bcd2308..79a67e639 100644 --- a/tests/integrations/a2a/gateway/test_server.py +++ b/tests/integrations/a2a/gateway/test_server.py @@ -9,12 +9,12 @@ import pytest_asyncio from a2a.server.agent_execution import AgentExecutor, RequestContext from a2a.server.events import EventQueue +from a2a.helpers import new_task_from_user_message from a2a.types import TaskState, TaskStatus, TaskStatusUpdateEvent from a2a.utils.constants import PROTOCOL_VERSION_0_3 from httpx import ASGITransport from band.integrations.a2a.gateway.server import GatewayServer -from band.integrations.a2a.protocol import new_task from tests.integrations.a2a.gateway.helpers import make_peer @@ -24,7 +24,7 @@ async def execute(self, context: RequestContext, event_queue: EventQueue) -> Non if task is None: if context.message is None: raise ValueError("A2A request is missing its message") - task = new_task(context.message) + task = new_task_from_user_message(context.message) if context.current_task is None: await event_queue.enqueue_event(task) await event_queue.enqueue_event( diff --git a/tests/integrations/a2a/test_adapter.py b/tests/integrations/a2a/test_adapter.py index 76cdc823e..00d1ffbbc 100644 --- a/tests/integrations/a2a/test_adapter.py +++ b/tests/integrations/a2a/test_adapter.py @@ -7,6 +7,7 @@ from uuid import uuid4 import pytest +from a2a.helpers import new_text_message from a2a.types import ( Artifact, Part, @@ -20,7 +21,6 @@ from band.core.types import PlatformMessage from band.integrations.a2a import A2AAdapter, A2AAuth, A2ASessionState -from band.integrations.a2a.protocol import text_message from band.testing import FakeAgentTools @@ -50,7 +50,7 @@ def make_task( status=TaskStatus(state=state), ) if status_message: - task.status.message.CopyFrom(text_message(status_message)) + task.status.message.CopyFrom(new_text_message(status_message)) if artifact_text: task.artifacts.append( Artifact(artifact_id="artifact-1", parts=[Part(text=artifact_text)]) @@ -239,7 +239,7 @@ async def test_status_update_is_applied_to_task_and_completes_flow( task.status.CopyFrom( TaskStatus( state=TaskState.TASK_STATE_COMPLETED, - message=text_message("Sunny"), + message=new_text_message("Sunny"), ) ) await adapter._handle_event( @@ -249,6 +249,23 @@ async def test_status_update_is_applied_to_task_and_completes_flow( assert tools.messages_sent[-1]["content"] == "Sunny" assert adapter._tasks == {} + @pytest.mark.asyncio + async def test_terminal_task_is_retained_when_band_delivery_fails( + self, adapter: A2AAdapter + ) -> None: + tools = FakeAgentTools() + tools.send_message = AsyncMock(side_effect=RuntimeError("Band unavailable")) + task = make_task(artifact_text="Final response") + + with pytest.raises(RuntimeError, match="Band unavailable"): + await adapter._handle_event( + task_event(task), tools, "room-123", "user-456", "Test User" + ) + + retained = adapter._task_cache[("room-123", task.id)] + assert retained.status.state == TaskState.TASK_STATE_COMPLETED + assert adapter._tasks["room-123"] == task.id + @pytest.mark.asyncio async def test_input_required_is_forwarded_and_persisted( self, adapter: A2AAdapter @@ -280,7 +297,7 @@ async def test_direct_message_response_is_forwarded( tools = FakeAgentTools() await adapter._handle_event( - StreamResponse(message=text_message("Hello")), + StreamResponse(message=new_text_message("Hello")), tools, "room-123", "user-456", diff --git a/tests/integrations/a2a/test_protocol.py b/tests/integrations/a2a/test_protocol.py index 31b3e9145..8ca0e3a38 100644 --- a/tests/integrations/a2a/test_protocol.py +++ b/tests/integrations/a2a/test_protocol.py @@ -2,20 +2,20 @@ from __future__ import annotations +from a2a.helpers import new_text_message from a2a.types import Artifact, Part, StreamResponse, Task, TaskState, TaskStatus from band.integrations.a2a.protocol import ( apply_task_stream_event, task_id_from_stream_event, task_response_text, - text_message, ) def test_task_stream_updates_build_a_task_from_deltas() -> None: status = TaskStatus( state=TaskState.TASK_STATE_COMPLETED, - message=text_message("Sunny"), + message=new_text_message("Sunny"), ) status_event = StreamResponse( status_update={