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/README.md b/README.md index dca3d8690..627bf24e4 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 406dbd992..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 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 d0c9ace7c..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 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..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 @@ -15,11 +16,8 @@ InternalError, Part, TaskState, - TextPart, UnsupportedOperationError, ) -from a2a.utils import new_agent_text_message, new_task -from a2a.utils.errors import ServerError try: from .agent import OrchestratorAgent @@ -58,8 +56,10 @@ async def execute( query = context.get_user_input() task = context.current_task - if not task: - task = new_task(context.message) # type: ignore + if task is None: + if context.message is None: + raise ValueError("A2A request is missing its 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,30 +72,26 @@ 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( - content, - task.context_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.input_required, - new_agent_text_message( + TaskState.TASK_STATE_INPUT_REQUIRED, + new_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 +99,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 +115,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..36a06aa38 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,90 @@ from uuid import uuid4 import httpx -from a2a.client import A2ACardResolver, A2AClient -from a2a.types import ( - MessageSendParams, - Part, - SendMessageRequest, - SendMessageSuccessResponse, - TextPart, +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, ) logger = logging.getLogger(__name__) -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) - - 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 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: + 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) - - return "No response from peer" + raise RuntimeError(error_msg) from exc 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 268741448..98f3118a2 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" } @@ -24,8 +24,8 @@ 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.apps import A2AStarletteApplication from a2a.server.events import EventQueue from a2a.server.request_handlers import DefaultRequestHandler from a2a.server.tasks import ( @@ -33,17 +33,18 @@ 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 dotenv import load_dotenv from setup_logging import setup_logging @@ -77,21 +78,22 @@ async def execute( request_text = context.get_user_input() task = context.current_task - if not task: - task = new_task(context.message) # type: ignore[arg-type] + if task is None: + if context.message is None: + raise ValueError("A2A request is missing its 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.working, - new_agent_text_message( + await updater.start_work( + new_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 +103,7 @@ async def cancel( context: RequestContext, event_queue: EventQueue, ) -> None: - raise ServerError(error=UnsupportedOperationError()) + raise UnsupportedOperationError() def main() -> None: @@ -116,8 +118,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 +145,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 0410ddad6..29b8a922b 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" } @@ -24,8 +24,8 @@ 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.apps import A2AStarletteApplication from a2a.server.events import EventQueue from a2a.server.request_handlers import DefaultRequestHandler from a2a.server.tasks import ( @@ -33,17 +33,18 @@ 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 dotenv import load_dotenv from setup_logging import setup_logging @@ -78,21 +79,22 @@ async def execute( request_text = context.get_user_input() task = context.current_task - if not task: - task = new_task(context.message) # type: ignore[arg-type] + if task is None: + if context.message is None: + raise ValueError("A2A request is missing its 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.working, - new_agent_text_message( + await updater.start_work( + new_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 +104,7 @@ async def cancel( context: RequestContext, event_queue: EventQueue, ) -> None: - raise ServerError(error=UnsupportedOperationError()) + raise UnsupportedOperationError() def main() -> None: @@ -117,8 +119,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 +144,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/pyproject.toml b/pyproject.toml index 903f54aeb..f2e6e8f27 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/adapter.py b/src/band/integrations/a2a/adapter.py index a852e70b4..b9f295ca1 100644 --- a/src/band/integrations/a2a/adapter.py +++ b/src/band/integrations/a2a/adapter.py @@ -3,71 +3,36 @@ 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.helpers import get_message_text, new_text_message +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, + apply_task_stream_event, + state_name, + task_id_from_stream_event, + task_response_text, +) 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 +90,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 +103,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 +147,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,91 +164,100 @@ 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 text: - await tools.send_message( - content=text, - mentions=[{"id": sender_id, "name": sender_name or ""}], - ) + if getattr(event, "HasField", lambda _name: False)("message"): + await self._deliver_message(event.message, tools, sender_id, sender_name) return - # Unpack task event - task, update = event + task = self._reduce_task_event(room_id, event) + if task is None: + return 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.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: - # 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.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.failed, - TaskState.canceled, - TaskState.rejected, - TaskState.auth_required, - ): - # Error states - send as error event (no mentions needed) - error_text = self._get_status_text(task) or f"Task {state.value}" - await tools.send_event( - content=error_text, - message_type="error", - metadata={"a2a_state": state.value}, - ) - finally: - # Clean up on terminal states - if state in TERMINAL_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) - # 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.""" @@ -296,12 +269,11 @@ 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 new_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: @@ -318,28 +290,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 isinstance(part.root, TextPart): - return part.root.text - - # Fallback: check status message - if task.status.message: - text = get_message_text(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 text: - return text - - return "" + return task_response_text(task) async def on_cleanup(self, room_id: str) -> None: """Clean up A2A context for room.""" @@ -349,6 +300,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 +316,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 +345,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 +364,27 @@ 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) + ): + 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._remember_task(room_id, task) + 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/__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..fa9237b09 100644 --- a/src/band/integrations/a2a/gateway/adapter.py +++ b/src/band/integrations/a2a/gateway/adapter.py @@ -3,23 +3,18 @@ from __future__ import annotations import asyncio +from dataclasses import dataclass +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.types import ( - Message as A2AMessage, - Part, - Role, - Task, - TaskState, - TaskStatus, - TaskStatusUpdateEvent, - TextPart, -) -from a2a.utils import get_message_text +from a2a.server.agent_execution import AgentExecutor, RequestContext +from a2a.server.events import EventQueue +from a2a.types import Task, TaskState, TaskStatus from band.client.rest import ( AsyncRestClient, @@ -35,16 +30,38 @@ 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.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.""" + + 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 +118,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 +128,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 +136,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) @@ -132,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. @@ -150,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: @@ -166,22 +179,23 @@ 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() 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) @@ -189,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, @@ -252,16 +235,20 @@ 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 - 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] + logger.debug( + "A2A response received: room=%s task=%s type=%s", + room_id, + pending.task.id, + msg.message_type, + ) + await self._publish_band_response(pending, msg) + 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 +256,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: + 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: @@ -295,75 +283,145 @@ 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. + 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.TASK_STATE_WORKING), + ) - Args: - peer_id: Target peer slug or UUID. - message: A2A message from remote agent. + async def _execute_a2a( + self, peer_id: str, context: RequestContext, event_queue: EventQueue + ) -> None: + """Bridge one official A2A execution to a Band room.""" + 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, + request.room_id, + request.context_id, + request.pending.task.id, + ) + try: + 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", + request.room_id, + request.pending.task.id, + ) + raise + except Exception: + logger.exception( + "A2A request failed: room=%s context=%s task=%s", + request.room_id, + request.context_id, + request.pending.task.id, + ) + raise + else: + logger.info( + "A2A request completed: room=%s task=%s", + request.room_id, + request.pending.task.id, + ) - Yields: - TaskStatusUpdateEvent for SSE streaming. - """ - # Resolve peer from slug or UUID + 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.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.id ) - - # 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=task, - sse_queue=sse_queue, - peer_id=peer_uuid, + 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), ) - # 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 + 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=room_id, + chat_id=request.room_id, message=ChatMessageRequest( - content=f"@{peer_name} {content}", - mentions=[ChatMessageRequestMentionsItem(id=peer_uuid, name=peer_name)], + 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( + 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) + await PendingA2ATask(task=task, event_queue=event_queue).cancel() async def _get_or_create_room( self, context_id: str | None, target_peer_id: str @@ -452,63 +510,16 @@ 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: - """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.failed - final = True - elif message_type in ("thought", "tool_call", "tool_result"): - state = TaskState.working - final = False + 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.completed - final = True - - # Update task status - task.status = TaskStatus( - state=state, - message=A2AMessage( - role=Role.agent, - message_id=str(uuid4()), - parts=[Part(root=TextPart(text=msg.content))], - ), - ) - - return TaskStatusUpdateEvent( - task_id=task.id, - context_id=task.context_id, - status=task.status, - final=final, - ) + 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/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..a4b6c79bf 100644 --- a/src/band/integrations/a2a/gateway/server.py +++ b/src/band/integrations/a2a/gateway/server.py @@ -1,50 +1,34 @@ -"""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 Awaitable, 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.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 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.requests import Request -from starlette.responses import JSONResponse, StreamingResponse -from starlette.routing import Route +from starlette.responses import JSONResponse +from starlette.routing import BaseRoute, 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 +36,33 @@ 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"], - ), - # 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 - - 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 + supported_interfaces=[ + AgentInterface( + 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), skills=[ @@ -180,283 +76,104 @@ async def _handle_agent_card(self, request: Request) -> JSONResponse: default_input_modes=["text/plain"], default_output_modes=["text/plain"], ) - 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", {}) + def _build_app(self) -> Starlette: + 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(): + 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, + ) + protocol_routes.extend( + create_agent_card_routes( + card, + card_url=f"/agents/{alias}/.well-known/agent-card.json", + ) + ) + protocol_routes.extend( + [ + Route( + f"/agents/{alias}/.well-known/agent.json", + self._legacy_agent_card(card), + methods=["GET"], + ) + ] + ) + 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}", + ) + ) - logger.debug( - "Received JSON-RPC %s for peer %s (%s), request_id=%s", - method, - peer.name, - slug, - request_id, + 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, ) - # 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, - ) - - # 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, - ) - - # 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, - ) - return JSONResponse( - { - "jsonrpc": "2.0", - "result": task.model_dump(mode="json", by_alias=True), - "id": request_id, - } - ) - 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, - ) - - async def _handle_jsonrpc_stream( - self, slug: str, params: dict[str, Any], request_id: str | None - ) -> StreamingResponse: - """Handle streaming message/stream JSON-RPC request. + async def response(_request: Any) -> JSONResponse: + return JSONResponse(payload) - Args: - slug: Peer slug. - params: JSON-RPC params containing message. - request_id: JSON-RPC request ID. + return response - 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, - ) - - 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..082ab444c 100644 --- a/src/band/integrations/a2a/gateway/types.py +++ b/src/band/integrations/a2a/gateway/types.py @@ -5,7 +5,10 @@ import asyncio from dataclasses import dataclass, field -from a2a.types import Task, TaskStatusUpdateEvent +from a2a.server.events import EventQueue +from a2a.server.tasks import TaskUpdater +from a2a.helpers import new_text_message +from a2a.types import Task @dataclass @@ -34,10 +37,55 @@ class PendingA2ATask: Attributes: task: The A2A Task object tracking this request. - sse_queue: Queue for streaming TaskStatusUpdateEvent to the client. - peer_id: The target peer this request is for. + event_queue: Official A2A event queue owned by DefaultRequestHandler. + done: Set when the final Band reply has been emitted or the room is + cleaned up. """ task: Task - sse_queue: asyncio.Queue[TaskStatusUpdateEvent] - peer_id: str + event_queue: EventQueue + 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 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 new file mode 100644 index 000000000..b980b266d --- /dev/null +++ b/src/band/integrations/a2a/protocol.py @@ -0,0 +1,101 @@ +"""Small protocol helpers shared by the A2A integrations.""" + +from __future__ import annotations + +from copy import deepcopy + +from a2a.helpers import get_artifact_text, get_message_text +from a2a.types import Role, StreamResponse, Task, TaskState + +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 snapshot_task(task: Task) -> Task: + """Return an independent task snapshot for event queue publication.""" + 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) + 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 + + +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 = get_artifact_text(artifact) + if text: + return text + + if 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 = get_message_text(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 + + +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/helpers.py b/tests/integrations/a2a/gateway/helpers.py new file mode 100644 index 000000000..1ff13dc3f --- /dev/null +++ b/tests/integrations/a2a/gateway/helpers.py @@ -0,0 +1,18 @@ +"""Shared test-data builders for the A2A gateway.""" + +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..6d4880b57 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,28 +8,32 @@ from uuid import uuid4 import pytest +from a2a.server.agent_execution import RequestContext +from a2a.server.events import EventQueueLegacy from a2a.types import ( - Message as A2AMessage, + Message, Part, Role, + SendMessageRequest, + Task, TaskState, - TextPart, + TaskStatus, ) from band.core.types import PlatformMessage -from band.integrations.a2a.gateway import ( - A2AGatewayAdapter, - GatewaySessionState, -) +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 import Peer -from band_rest.core.api_error import ApiError +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, @@ -43,646 +47,299 @@ 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: - """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_init_default_values(self) -> None: - """Should initialize with default values.""" - adapter = A2AGatewayAdapter() - - 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, - ) - - 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", - ) +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, + ) - assert adapter._rest is not None - def test_init_sets_history_converter(self) -> None: - """Should set GatewayHistoryConverter.""" - adapter = A2AGatewayAdapter() +class TestGatewayConfiguration: + def test_timeout_is_adapter_configuration(self) -> None: + config = A2AGatewayAdapterConfig(response_timeout_s=12) + adapter = A2AGatewayAdapter(config=config) - assert adapter.history_converter is not None + 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"), - ] + 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 - - await adapter.on_started("Gateway", "A2A Gateway Agent") + ) as server_type: + server = MagicMock() + server.start = AsyncMock() + server_type.return_value = server - # 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 + await adapter.on_started("Gateway", "A2A Gateway") - @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")] - adapter._rest.agent_api_peers.list_agent_peers = AsyncMock( - return_value=mock_response + assert adapter._peers["weather-agent"].id == "weather" + server.start.assert_awaited_once() + assert ( + adapter._rest.agent_api_peers.list_agent_peers.call_args.kwargs[ + "request_options" + ] + == DEFAULT_REQUEST_OPTIONS ) - # 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") - - mock_server_class.assert_called_once() - mock_server.start.assert_called_once() - assert adapter._server is mock_server +class TestGatewayExecution: @pytest.mark.asyncio - async def test_on_started_stores_agent_info(self) -> None: - """Should store agent name and description.""" + 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) + tools = FakeAgentTools() - # Mock REST client - mock_response = MagicMock() - mock_response.data = [] - adapter._rest.agent_api_peers.list_agent_peers = AsyncMock( - return_value=mock_response + 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", + ) + + adapter._rest.agent_api_messages.create_agent_chat_message = AsyncMock( + side_effect=send_message ) + queue = EventQueueLegacy() - # 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") + await BandAgentExecutor(adapter, "weather").execute(make_request(), queue) - assert adapter.agent_name == "Test Gateway" - assert adapter.agent_description == "A test gateway" + 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_on_started_retries_peer_discovery_on_rate_limit(self) -> None: - """Should retry peer discovery when startup hits HTTP 429.""" - adapter = A2AGatewayAdapter() - - mock_response = MagicMock() - mock_response.data = [make_peer("weather", "Weather Agent")] - 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, - ] - ) - - 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( - "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 - - -class TestA2AGatewayAdapterOnMessage: - """Tests for A2AGatewayAdapter.on_message().""" - - @pytest.fixture - def adapter_with_mocks(self) -> A2AGatewayAdapter: - """Create adapter with mocked dependencies.""" + async def test_posts_to_band_and_returns_terminal_response(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 + configure_room_creation(adapter) + sent = asyncio.Event() - @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") + async def send_message(**_kwargs: object) -> None: + sent.set() - history = GatewaySessionState( - context_to_room={"ctx-1": "room-1", "ctx-2": "room-2"}, - room_participants={"room-1": {"peer-a"}, "room-2": {"peer-b"}}, + adapter._rest.agent_api_messages.create_agent_chat_message = AsyncMock( + side_effect=send_message + ) + queue = EventQueueLegacy() + execution = asyncio.create_task( + BandAgentExecutor(adapter, "weather").execute(make_request(), queue) ) + await asyncio.wait_for(sent.wait(), timeout=1) - await adapter_with_mocks.on_message( - msg, - tools, - history, + initial = await queue.dequeue_event() + assert initial.status.state == TaskState.TASK_STATE_WORKING + + await adapter.on_message( + make_platform_message("Sunny"), + FakeAgentTools(), + GatewaySessionState(), None, None, - is_session_bootstrap=True, + is_session_bootstrap=False, room_id="room-123", ) + await asyncio.wait_for(execution, timeout=1) + final = await queue.dequeue_event() - 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"}, - } + 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_on_message_correlates_pending_task( - self, adapter_with_mocks: A2AGatewayAdapter - ) -> None: - """Should push event to pending task's SSE queue.""" - tools = FakeAgentTools() - msg = make_platform_message("Weather is sunny", room_id="room-123") + 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() - # Set up pending task - from band.integrations.a2a.gateway.types import PendingA2ATask - from a2a.types import Task, TaskStatus + async def send_message(**_kwargs: object) -> None: + sent.set() - sse_queue: asyncio.Queue = asyncio.Queue() - task = Task( - id="task-123", - context_id="ctx-123", - status=TaskStatus(state=TaskState.working), + adapter._rest.agent_api_messages.create_agent_chat_message = AsyncMock( + side_effect=send_message ) - adapter_with_mocks._pending_tasks["room-123"] = PendingA2ATask( - task=task, - sse_queue=sse_queue, - peer_id="weather", + execution = asyncio.create_task( + BandAgentExecutor(adapter, "weather").execute(make_request(), queue) ) + await asyncio.wait_for(sent.wait(), timeout=1) + await queue.dequeue_event() - await adapter_with_mocks.on_message( - msg, - tools, + await adapter.on_message( + make_platform_message("Checking", message_type="thought"), + FakeAgentTools(), GatewaySessionState(), None, None, is_session_bootstrap=False, room_id="room-123", ) + update = await queue.dequeue_event() + assert update.status.state == TaskState.TASK_STATE_WORKING + assert not execution.done() - # Event should be in queue - assert not sse_queue.empty() - event = sse_queue.get_nowait() - assert event.task_id == "task-123" - assert event.final is True # text message = completed - - @pytest.mark.asyncio - async def test_on_message_cleans_up_on_final_event( - self, adapter_with_mocks: A2AGatewayAdapter - ) -> None: - """Should clean up pending task on final event.""" - 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 - - sse_queue: asyncio.Queue = asyncio.Queue() - task = Task( - id="task-123", - context_id="ctx-123", - status=TaskStatus(state=TaskState.working), - ) - adapter_with_mocks._pending_tasks["room-123"] = PendingA2ATask( - task=task, - sse_queue=sse_queue, - peer_id="weather", - ) - - await adapter_with_mocks.on_message( - msg, - tools, + await adapter.on_message( + make_platform_message("Sunny"), + FakeAgentTools(), GatewaySessionState(), None, None, is_session_bootstrap=False, room_id="room-123", ) - - # Pending task should be cleaned up - assert "room-123" not in adapter_with_mocks._pending_tasks - - -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() + await asyncio.wait_for(execution, timeout=1) + 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 - - sse_queue: asyncio.Queue = asyncio.Queue() - task = Task( - id="task-123", - context_id="ctx-123", - status=TaskStatus(state=TaskState.working), - ) - adapter._pending_tasks["room-123"] = PendingA2ATask( - task=task, - sse_queue=sse_queue, - peer_id="weather", - ) + 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 - - @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 - + terminal = await queue.dequeue_event() + assert terminal.status.state == TaskState.TASK_STATE_FAILED + assert pending.done.is_set() + assert adapter._pending_tasks == {} -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 - - # 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 - ) + room, context = await adapter._get_or_create_room("ctx", "weather") + same_room, same_context = await adapter._get_or_create_room(context, "data") - # 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 + 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 ) - # 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 + def test_rehydrate_merges_without_overwriting_live_context(self) -> None: + adapter = A2AGatewayAdapter() + adapter._context_to_room["ctx"] = "live-room" - @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 + adapter._rehydrate( + GatewaySessionState( + context_to_room={"ctx": "old-room", "new": "new-room"}, + 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") + assert adapter._context_to_room == { + "ctx": "live-room", + "new": "new-room", + } + assert adapter._room_participants["new-room"] == {"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 - @pytest.mark.asyncio - async def test_same_context_different_peers_same_room( - self, adapter_with_tracking: A2AGatewayAdapter +class TestGatewayResponses: + @pytest.mark.parametrize( + ("message_type", "state"), + [ + ("thought", TaskState.TASK_STATE_WORKING), + ("text", TaskState.TASK_STATE_COMPLETED), + ("error", TaskState.TASK_STATE_FAILED), + ], + ) + async def test_publishes_band_message_with_matching_task_state( + self, message_type: str, state: int ) -> 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 + adapter = A2AGatewayAdapter() + queue = EventQueueLegacy() + pending = make_pending(queue) + + 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 ac1247d65..79a67e639 100644 --- a/tests/integrations/a2a/gateway/test_server.py +++ b/tests/integrations/a2a/gateway/test_server.py @@ -1,829 +1,195 @@ -"""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.types import ( - Message as A2AMessage, - TaskState, - TaskStatus, - TaskStatusUpdateEvent, -) -from starlette.testclient import TestClient +import httpx +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_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", +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: + if context.message is None: + raise ValueError("A2A request is missing its 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( + TaskStatusUpdateEvent( + task_id=task.id, + context_id=task.context_id, + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ) + ) + + async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None: + raise NotImplementedError + + +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(), ) -class TestGatewayServerInit: - """Tests for GatewayServer initialization.""" +@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_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.""" +async def test_agent_cards_use_the_schema_expected_by_each_protocol_version( + gateway_client: httpx.AsyncClient, +) -> None: + 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" + assert standard_card["supportedInterfaces"][0]["protocolBinding"] == "JSONRPC" + assert standard_card["supportedInterfaces"][0]["url"].endswith( + "/agents/weather-agent" + ) - @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, + 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" + assert legacy_card["protocolVersion"] == PROTOCOL_VERSION_0_3 + assert legacy_card["url"].endswith("/agents/weather-agent") + assert "supportedInterfaces" not in legacy_card + + +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"] == [ + { + "slug": "weather-agent", + "id": "uuid-weather", + "name": "Weather Agent", + "description": "Gets weather info", } - 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", - status=TaskStatus(state=TaskState.completed), - final=True, - ) +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 - 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 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"}], - }, - ) - - assert "text/event-stream" in response.headers["content-type"] - - 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 - - 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"}], - }, - ) - - # Response should be SSE formatted - content = response.text - assert content.startswith("data: ") - assert '"taskId":"task-123"' in content - assert '"final":true' in content - - 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"}, - ) - - 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 - - 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, - ) - - 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"}], - }, - ) - - # 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 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 - 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, - ) +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": {}}, + ) - 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"}], - } - }, + assert response.status_code == 200 + assert response.json()["error"]["code"] == -32601 + + +async def test_jsonrpc_send_runs_through_official_handler_and_executor( + gateway_client: httpx.AsyncClient, +) -> None: + response = await gateway_client.post( + "/agents/weather-agent", + headers={"A2A-Version": "1.0"}, + json={ + "jsonrpc": "2.0", + "id": "request-1", + "method": "SendMessage", + "params": { + "message": { + "role": "ROLE_USER", + "messageId": "message-1", + "parts": [{"text": "Hello"}], + } }, - ) + }, + ) - # 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 response.status_code == 200 + body = response.json() + assert body["id"] == "request-1" + assert body["result"]["task"]["status"]["state"] == "TASK_STATE_COMPLETED" + + +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={ + "message": { + "messageId": "message-1", + "role": "ROLE_USER", + "parts": [{"text": "Hello"}], + } + }, + ) - 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"}], - } + 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 + + +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", + "id": "request-1", + "method": "message/stream", + "params": { + "message": { + "messageId": "message-1", + "role": "user", + "parts": [{"type": "text", "text": "Hello"}], }, }, - ) - - # 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 - 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"] 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/tests/integrations/a2a/test_adapter.py b/tests/integrations/a2a/test_adapter.py index f0d4e48b9..00d1ffbbc 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.helpers import new_text_message 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.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,308 @@ 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))], + 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)]) ) + return task - artifacts = None - 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, - ) +def task_event(task: Task) -> StreamResponse: + return StreamResponse(task=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 status_event(task: Task) -> StreamResponse: + return StreamResponse( + status_update={ + "task_id": task.id, + "context_id": task.context_id, + "status": task.status, + } + ) - 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 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, + } + ) - 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().""" - - @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 - +class TestA2AAdapterStartup: @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", - } - } + client = MagicMock() - @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") + with patch("band.integrations.a2a.adapter.ClientFactory") as factory_type: + factory = factory_type.return_value + factory.create_from_url = AsyncMock(return_value=client) - with patch("band.integrations.a2a.adapter.ClientFactory") as mock_factory: - mock_factory.connect = AsyncMock(return_value=MagicMock()) + await adapter.on_started("Agent", "Description") - await adapter.on_started("Test Agent", "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" - assert adapter.agent_name == "Test Agent" - assert adapter.agent_description == "A test agent" - - -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.""" - tools = FakeAgentTools() - msg = make_platform_message("Hello") - - task = make_task( - state=TaskState.failed, - status_message="Currency API unavailable", - ) - - 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", + 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), + ) ) - - # 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 - - @pytest.mark.asyncio - async def test_exception_sends_error_event(self, adapter_with_client): - """Should send error event on exception.""" 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.on_message( + make_platform_message(), tools, A2ASessionState(), None, None, - is_session_bootstrap=True, + is_session_bootstrap=False, room_id="room-123", ) - # 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"] == "Part one. \nPart two." @pytest.mark.asyncio - async def test_direct_message_reply(self, adapter_with_client): - """Should handle direct A2A Message reply.""" + 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) - a2a_reply = A2AMessage( - role=Role.agent, - message_id=str(uuid4()), - parts=[Part(root=TextPart(text="Hello! How can I help?"))], + await adapter._handle_event( + task_event(task), tools, "room-123", "user-456", "Test User" ) - - 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, - 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"] == "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 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" - - @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", + task.status.CopyFrom( + TaskStatus( + state=TaskState.TASK_STATE_COMPLETED, + message=new_text_message("Sunny"), + ) ) - - 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", + await adapter._handle_event( + status_event(task), tools, "room-123", "user-456", "Test User" ) - 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.""" - adapter = A2AAdapter(remote_url="http://localhost:10000") - adapter._client = MagicMock() - return adapter + assert tools.messages_sent[-1]["content"] == "Sunny" + assert adapter._tasks == {} @pytest.mark.asyncio - async def test_emits_task_event_on_completed(self, adapter_with_client): - """Should emit task event when task completes.""" + async def test_terminal_task_is_retained_when_band_delivery_fails( + self, adapter: A2AAdapter + ) -> None: tools = FakeAgentTools() - msg = make_platform_message("Hello") - - task = make_task( - state=TaskState.completed, - artifact_text="Response", - ) - - async def mock_send_message(*args, **kwargs): - yield (task, None) + tools.send_message = AsyncMock(side_effect=RuntimeError("Band unavailable")) + task = make_task(artifact_text="Final response") - 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", - ) + with pytest.raises(RuntimeError, match="Band unavailable"): + await adapter._handle_event( + task_event(task), tools, "room-123", "user-456", "Test User" + ) - # 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" + 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_emits_task_event_on_input_required(self, adapter_with_client): - """Should emit task event when input is required.""" + async def test_input_required_is_forwarded_and_persisted( + self, adapter: A2AAdapter + ) -> None: 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, + 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", ) - # 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.""" - 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", + assert tools.messages_sent[-1]["content"] == "Which city?" + assert tools.events_sent[-1]["metadata"]["a2a_task_state"] == ( + "TASK_STATE_INPUT_REQUIRED" ) - # 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.""" + async def test_direct_message_response_is_forwarded( + self, adapter: A2AAdapter + ) -> None: 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, + await adapter._handle_event( + StreamResponse(message=new_text_message("Hello")), tools, - history, - None, - None, - is_session_bootstrap=False, - room_id="room-123", + "room-123", + "user-456", + "Test User", ) - # Resubscribe should not be called on non-bootstrap - adapter_with_client._client.resubscribe.assert_not_called() + assert tools.messages_sent[-1]["content"] == "Hello" - @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", +class TestA2AAdapterSession: + @pytest.mark.asyncio + async def test_rehydrates_context_and_resubscribes_active_task(self) -> None: + adapter = A2AAdapter(remote_url="http://localhost:10000") + adapter._client = MagicMock() + adapter._client.subscribe = MagicMock( + return_value=stream(task_event(make_task(TaskState.TASK_STATE_WORKING))) ) - 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", + await adapter._rehydrate_from_history( + "room-123", + A2ASessionState( + context_id="ctx-123", + task_id="task-123", + task_state="TASK_STATE_WORKING", + ), ) - # 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" + assert adapter._contexts["room-123"] == "ctx-123" + assert adapter._tasks["room-123"] == "task-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", + async def test_does_not_resubscribe_terminal_task(self) -> None: + adapter = A2AAdapter(remote_url="http://localhost:10000") + adapter._client = MagicMock() + 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() diff --git a/tests/integrations/a2a/test_protocol.py b/tests/integrations/a2a/test_protocol.py new file mode 100644 index 000000000..8ca0e3a38 --- /dev/null +++ b/tests/integrations/a2a/test_protocol.py @@ -0,0 +1,95 @@ +"""Tests for shared protobuf A2A protocol helpers.""" + +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, +) + + +def test_task_stream_updates_build_a_task_from_deltas() -> None: + status = TaskStatus( + state=TaskState.TASK_STATE_COMPLETED, + message=new_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_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( + 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 diff --git a/uv.lock b/uv.lock index f7db2d10f..bbe2d413a 100644 --- a/uv.lock +++ b/uv.lock @@ -100,18 +100,21 @@ 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" }, ] -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]] @@ -318,6 +321,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" @@ -686,10 +703,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" }, @@ -1715,6 +1732,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" @@ -3427,6 +3457,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"