diff --git a/dev/a2a/.gitignore b/dev/a2a/.gitignore new file mode 100644 index 000000000..3a82ddea6 --- /dev/null +++ b/dev/a2a/.gitignore @@ -0,0 +1,3 @@ +.tck/ +*.log + diff --git a/dev/a2a/README.md b/dev/a2a/README.md new file mode 100644 index 000000000..2280bbaf8 --- /dev/null +++ b/dev/a2a/README.md @@ -0,0 +1,52 @@ +# A2A TCK development harness + +This directory contains a local A2A agent and a PowerShell runner for the +official [A2A Technology Compatibility Kit][tck]. + +## Prerequisites + +From the repository root, prepare the development environment: + +```powershell +.\scripts\dev_setup.ps1 +``` + +Install [uv](https://docs.astral.sh/uv/) and ensure `git` is available on +`PATH`. + +## Run the TCK + +From the repository root: + +```powershell +.\dev\a2a\run_tck.ps1 +``` + +The script: + +1. Clones or updates the TCK under `dev\a2a\.tck`. +2. Creates the TCK virtual environment and installs its dependencies. +3. Starts the local agent on `http://127.0.0.1:41241`. +4. Runs the TCK against the JSON-RPC and HTTP+JSON interfaces. +5. Stops the local agent. + +By default, the complete suite runs for both supported transports. Examples: + +```powershell +# Run only MUST requirements over JSON-RPC. +.\dev\a2a\run_tck.ps1 -Transport jsonrpc -Level must + +# Run only HTTP+JSON against a specific TCK revision. +.\dev\a2a\run_tck.ps1 -Transport http_json -TckRevision main + +# Forward additional arguments to pytest through the TCK. +.\dev\a2a\run_tck.ps1 -PytestArgs "-x", "--pdb" +``` + +TCK reports are written under `dev\a2a\.tck\reports`. + +The test agent intentionally disables JWT authentication. It is only intended +for local protocol compatibility testing. + +[tck]: https://github.com/a2aproject/a2a-tck + diff --git a/dev/a2a/agent.py b/dev/a2a/agent.py new file mode 100644 index 000000000..ddd2ce572 --- /dev/null +++ b/dev/a2a/agent.py @@ -0,0 +1,236 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from __future__ import annotations + +import argparse +import asyncio +from uuid import uuid4 + +import uvicorn +from a2a.types import AgentInterface, AgentSkill +from a2a.utils.constants import TransportProtocol +from fastapi import FastAPI +from starlette.routing import Route + +from microsoft_agents.activity import ( + Activity, + ActivityTypes, + Attachment, + EndOfConversationCodes, + InputHints, + StreamInfo, +) +from microsoft_agents.hosting.a2a import A2AAdapter, add_a2a +from microsoft_agents.hosting.core import ( + AgentApplication, + AgentAuthConfiguration, + AnonymousTokenProvider, + ApplicationOptions, + ConnectionManager, + MemoryStorage, + TurnContext, + TurnState, +) + + +def create_agent() -> AgentApplication[TurnState]: + connection_manager = ConnectionManager( + provider_factory=lambda _: AnonymousTokenProvider(), + connections_configurations={ + "SERVICE_CONNECTION": AgentAuthConfiguration( + anonymous_allowed=True, + ) + }, + ) + application = AgentApplication[TurnState]( + options=ApplicationOptions(storage=MemoryStorage()), + connection_manager=connection_manager, + ) + + @application.activity("message") + async def on_message(context: TurnContext, _state: TurnState) -> None: + message_id = getattr(context.activity.channel_data, "message_id", "") + + async def send_artifact( + *, + text: str | None = None, + value: dict | None = None, + attachment: Attachment | None = None, + ) -> None: + await context.send_activity( + Activity( + type=ActivityTypes.message, + text=text, + value=value, + attachments=[attachment] if attachment else [], + entities=[ + StreamInfo( + stream_id=str(uuid4()), + stream_type="content", + stream_sequence=1, + ) + ], + ) + ) + + if message_id.startswith("test-resubscribe-message-id"): + await context.send_activity("Working") + await asyncio.sleep(4) + elif message_id.startswith("tck-input-required"): + await context.send_activity( + Activity( + type=ActivityTypes.message, + input_hint=InputHints.expecting_input, + ) + ) + return + elif message_id.startswith("tck-message-response"): + await context.send_activity("Direct message response") + return + elif message_id.startswith("tck-artifact-file-url"): + await send_artifact( + attachment=Attachment( + content_type="text/plain", + content_url="https://example.com/output.txt", + name="output.txt", + ) + ) + elif message_id.startswith("tck-artifact-file"): + await send_artifact( + attachment=Attachment( + content_type="text/plain", + content=b"tck", + name="output.txt", + ) + ) + elif message_id.startswith("tck-artifact-data"): + await send_artifact(value={"key": "value", "count": 42}) + elif message_id.startswith("tck-artifact-text"): + await send_artifact(text="Generated text content") + elif message_id.startswith("tck-stream-artifact-chunked"): + response = context.streaming_response + response.queue_text_chunk("chunk-1 ") + response.queue_text_chunk("chunk-2") + await response.end_stream() + return + elif message_id.startswith("tck-stream-artifact-file"): + await send_artifact( + attachment=Attachment( + content_type="text/plain", + content=b"tck", + name="output.txt", + ) + ) + elif message_id.startswith("tck-stream-artifact-text"): + response = context.streaming_response + response.queue_text_chunk("Streamed text content") + await response.end_stream() + return + elif message_id.startswith("tck-stream-ordering-001"): + response = context.streaming_response + response.queue_informative_update("Working") + response.queue_text_chunk("Ordered output") + await response.end_stream() + return + elif message_id.startswith("tck-stream-001"): + response = context.streaming_response + response.queue_text_chunk("Stream hello from TCK") + await response.end_stream() + return + elif message_id.startswith("tck-stream-002"): + pass + elif message_id.startswith("tck-stream-003"): + response = context.streaming_response + response.queue_informative_update("Working") + response.queue_text_chunk("Stream task lifecycle") + await response.end_stream() + return + else: + await context.send_activity("Hello from TCK") + + await context.send_activity( + Activity( + type=ActivityTypes.end_of_conversation, + code=EndOfConversationCodes.completed_successfully, + ) + ) + + return application + + +def create_app() -> FastAPI: + agent = create_agent() + adapter = A2AAdapter( + agent, + agent_card_name="Microsoft Agents SDK TCK Agent", + agent_card_description="A deterministic agent for A2A TCK validation.", + agent_card_version="1.0.0", + agent_interfaces=[ + AgentInterface( + url="/rpc", + protocol_binding=TransportProtocol.JSONRPC, + ), + AgentInterface( + url="/rest", + protocol_binding=TransportProtocol.HTTP_JSON, + ), + ], + skills=[ + AgentSkill( + id="echo", + name="Echo", + description="Echoes incoming text.", + tags=["test", "echo"], + examples=["Hello"], + input_modes=["text/plain"], + output_modes=["text/plain"], + ) + ], + ) + + app = FastAPI(title="Microsoft Agents SDK A2A TCK Agent", version="1.0.0") + add_a2a( + app, + agent, + adapter=adapter, + use_jwt_middleware=False, + ) + + rpc_route = next( + route + for route in app.router.routes + if isinstance(route, Route) and route.path == "/rpc" + ) + mount_index = next( + index + for index, route in enumerate(app.router.routes) + if route.path == "/{tenant}" + ) + app.router.routes.insert( + mount_index, + Route( + path="/rpc/", + endpoint=rpc_route.endpoint, + methods=["POST"], + ), + ) + + @app.get("/health") + async def health() -> dict[str, str]: + return {"status": "ok"} + + return app + + +def main() -> None: + parser = argparse.ArgumentParser(description="Run the local A2A TCK agent.") + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", default=41241, type=int) + args = parser.parse_args() + + uvicorn.run(create_app(), host=args.host, port=args.port) + + +if __name__ == "__main__": + main() diff --git a/dev/a2a/run_tck.ps1 b/dev/a2a/run_tck.ps1 new file mode 100644 index 000000000..c366e99f9 --- /dev/null +++ b/dev/a2a/run_tck.ps1 @@ -0,0 +1,155 @@ +param( + [ValidateSet("jsonrpc", "http_json", "all")] + [string]$Transport = "all", + + [ValidateSet("must", "should", "may", "all")] + [string]$Level = "all", + + [string]$TckRevision = "main", + + [string]$HostAddress = "127.0.0.1", + + [int]$Port = 41241, + + [string[]]$PytestArgs = @() +) + +$ErrorActionPreference = "Stop" + +$repositoryRoot = (Resolve-Path (Join-Path $PSScriptRoot "..\..")).Path +$repositoryPython = Join-Path $repositoryRoot "venv\Scripts\python.exe" +$agentPath = Join-Path $PSScriptRoot "agent.py" +$tckPath = Join-Path $PSScriptRoot ".tck" +$tckPython = Join-Path $tckPath ".venv\Scripts\python.exe" +$stdoutLog = Join-Path $PSScriptRoot "agent.stdout.log" +$stderrLog = Join-Path $PSScriptRoot "agent.stderr.log" +$sutHost = "http://${HostAddress}:$Port/rpc" + +if (-not (Test-Path $repositoryPython)) { + throw "Repository virtual environment not found. Run .\scripts\dev_setup.ps1 first." +} + +if (-not (Get-Command git -ErrorAction SilentlyContinue)) { + throw "git was not found on PATH." +} + +if (-not (Get-Command uv -ErrorAction SilentlyContinue)) { + throw "uv was not found on PATH. Install it from https://docs.astral.sh/uv/." +} + +if (-not (Test-Path (Join-Path $tckPath ".git"))) { + git clone https://github.com/a2aproject/a2a-tck.git $tckPath + if ($LASTEXITCODE -ne 0) { + throw "Failed to clone the A2A TCK repository." + } +} + +git -C $tckPath fetch origin $TckRevision +if ($LASTEXITCODE -ne 0) { + throw "Failed to fetch TCK revision '$TckRevision'." +} + +git -C $tckPath checkout --detach FETCH_HEAD +if ($LASTEXITCODE -ne 0) { + throw "Failed to check out TCK revision '$TckRevision'." +} + +if (-not (Test-Path $tckPython)) { + uv venv (Join-Path $tckPath ".venv") --python 3.11 + if ($LASTEXITCODE -ne 0) { + throw "Failed to create the TCK virtual environment." + } +} + +& $tckPython -m ensurepip --upgrade +if ($LASTEXITCODE -ne 0) { + throw "Failed to bootstrap pip in the TCK virtual environment." +} + +& $tckPython -m pip install -e $tckPath +if ($LASTEXITCODE -ne 0) { + throw "Failed to install the TCK." +} + +Remove-Item $stdoutLog, $stderrLog -ErrorAction SilentlyContinue + +$agentProcess = $null +$exitCode = 1 + +try { + $agentProcess = Start-Process ` + -FilePath $repositoryPython ` + -ArgumentList @($agentPath, "--host", $HostAddress, "--port", $Port) ` + -WorkingDirectory $repositoryRoot ` + -RedirectStandardOutput $stdoutLog ` + -RedirectStandardError $stderrLog ` + -PassThru + + $healthUrl = "http://${HostAddress}:$Port/health" + $deadline = (Get-Date).AddSeconds(30) + $healthy = $false + + while ((Get-Date) -lt $deadline) { + if ($agentProcess.HasExited) { + break + } + + try { + $response = Invoke-WebRequest -Uri $healthUrl -UseBasicParsing + if ($response.StatusCode -eq 200) { + $healthy = $true + break + } + } + catch { + Start-Sleep -Milliseconds 250 + } + } + + if (-not $healthy) { + if (Test-Path $stderrLog) { + Get-Content $stderrLog + } + throw "The A2A test agent did not become healthy at $healthUrl." + } + + $tckArguments = @( + (Join-Path $tckPath "run_tck.py"), + "--sut-host", + $sutHost, + "-v" + ) + + if ($Transport -eq "all") { + $tckArguments += @("--transport", "jsonrpc,http_json") + } + else { + $tckArguments += @("--transport", $Transport) + } + + if ($Level -ne "all") { + $tckArguments += @("--level", $Level) + } + + if ($PytestArgs.Count -gt 0) { + $tckArguments += "--" + $tckArguments += $PytestArgs + } + + Push-Location $tckPath + try { + & $tckPython @tckArguments + $exitCode = $LASTEXITCODE + } + finally { + Pop-Location + } +} +finally { + if ($null -ne $agentProcess -and -not $agentProcess.HasExited) { + Stop-Process -Id $agentProcess.Id + $agentProcess.WaitForExit() + } +} + +exit $exitCode diff --git a/dev/integration/pyproject.toml b/dev/integration/pyproject.toml index c29a72124..345d28c01 100644 --- a/dev/integration/pyproject.toml +++ b/dev/integration/pyproject.toml @@ -7,6 +7,7 @@ dependencies = [ "pytest-asyncio", "microsoft-agents-authentication-msal", "microsoft-agents-hosting-aiohttp", + "microsoft-agents-hosting-a2a @ file:///${PROJECT_ROOT}/../../libraries/microsoft-agents-hosting-a2a", "microsoft-agents-hosting-dialogs", "microsoft-agents-hosting-fastapi", "microsoft-agents-hosting-testing @ file:///${PROJECT_ROOT}/../microsoft-agents-hosting-testing", diff --git a/dev/integration/tests/jwt_validation/test_a2a_jwt_validation.py b/dev/integration/tests/jwt_validation/test_a2a_jwt_validation.py new file mode 100644 index 000000000..ddb50b4b4 --- /dev/null +++ b/dev/integration/tests/jwt_validation/test_a2a_jwt_validation.py @@ -0,0 +1,129 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import uuid +from unittest.mock import MagicMock + +import pytest +from a2a.types import AgentInterface, Message, Part, Role, SendMessageRequest +from a2a.utils.constants import TransportProtocol +from fastapi import FastAPI +from fastapi.testclient import TestClient +from google.protobuf.json_format import MessageToDict + +from microsoft_agents.activity import Activity, ActivityTypes, EndOfConversationCodes +from microsoft_agents.hosting.a2a import A2AAdapter, add_a2a +from microsoft_agents.hosting.core import ( + AgentApplication, + ApplicationOptions, + MemoryStorage, + TurnContext, + TurnState, +) + +from tests.utils.config import REAL_SERVICE_CONNECTION_ENV_VARS +from tests.utils.pytest import skip_if_no_var + +from ._helpers import ( + acquire_real_service_connection_token, + auth_config_with_invalid_audience, +) + +_requires_real_service_connection = skip_if_no_var( + *REAL_SERVICE_CONNECTION_ENV_VARS, load_root_env_file=True +) + + +def _rpc_request() -> dict: + request = SendMessageRequest( + message=Message( + role=Role.ROLE_USER, + message_id=str(uuid.uuid4()), + parts=[Part(text="hello")], + ) + ) + return { + "jsonrpc": "2.0", + "id": "request-1", + "method": "SendMessage", + "params": MessageToDict(request), + } + + +def _create_app(auth_config): + observed = {} + agent = AgentApplication[TurnState]( + options=ApplicationOptions(storage=MemoryStorage()), + connection_manager=MagicMock(), + ) + + @agent.activity(ActivityTypes.message) + async def on_message(context: TurnContext, _state: TurnState) -> None: + observed["identity"] = context.identity + await context.send_activity("Authenticated") + await context.send_activity( + Activity( + type=ActivityTypes.end_of_conversation, + code=EndOfConversationCodes.completed_successfully, + ) + ) + + app = FastAPI() + app.state.agent_configuration = auth_config + adapter = A2AAdapter( + agent, + agent_interfaces=[ + AgentInterface( + url="/a2a", + protocol_binding=TransportProtocol.JSONRPC, + ) + ], + ) + add_a2a(app, agent, adapter=adapter) + return app, observed + + +@_requires_real_service_connection +@pytest.mark.asyncio +async def test_a2a_accepts_real_service_connection_token_and_forwards_identity(): + token, auth_config = await acquire_real_service_connection_token() + app, observed = _create_app(auth_config) + + with TestClient(app) as client: + response = client.post( + "/a2a", + json=_rpc_request(), + headers={ + "A2A-Version": "1.0", + "Authorization": f"Bearer {token}", + }, + ) + + assert response.status_code == 200 + assert response.json()["result"]["task"]["status"]["state"] == ( + "TASK_STATE_COMPLETED" + ) + identity = observed["identity"] + assert identity.allow_anonymous is False + assert identity.claims + + +@_requires_real_service_connection +@pytest.mark.asyncio +async def test_a2a_rejects_real_token_with_invalid_audience(): + token, auth_config = await acquire_real_service_connection_token() + app, observed = _create_app(auth_config_with_invalid_audience(auth_config)) + + with TestClient(app) as client: + response = client.post( + "/a2a", + json=_rpc_request(), + headers={ + "A2A-Version": "1.0", + "Authorization": f"Bearer {token}", + }, + ) + + assert response.status_code == 401 + assert response.json() == {"error": "Invalid token or authentication failed."} + assert "identity" not in observed diff --git a/libraries/microsoft-agents-activity/microsoft_agents/activity/activity.py b/libraries/microsoft-agents-activity/microsoft_agents/activity/activity.py index 74b1cf167..a75c27e9c 100644 --- a/libraries/microsoft-agents-activity/microsoft_agents/activity/activity.py +++ b/libraries/microsoft-agents-activity/microsoft_agents/activity/activity.py @@ -38,6 +38,7 @@ ClientCitation, ProductInfo, SensitivityUsageInfo, + StreamInfo, ) from .entity._validate_known_entities import _validate_known_entities from .conversation_reference import ConversationReference @@ -199,7 +200,7 @@ class Activity(AgentsModel): text_highlights: list[TextHighlight] = None semantic_action: SemanticAction = None caller_id: NonEmptyString = None - request_id: str | None = Field(None, exclude=True) + request_id: Annotated[str | None, Field(None, exclude=True)] = None @field_validator("entities", mode="before") @classmethod @@ -1079,7 +1080,11 @@ def _convert_entity_list( entities.append(Activity._convert_entity(e, entity_cls)) return entities - def get_product_info_entity(self) -> Optional[ProductInfo]: + def get_product_info_entity(self) -> ProductInfo | None: + """Gets the product info entity from the activity's entities. + + :return: The product info entity if found, otherwise None. + """ if not self.entities: return None target = EntityTypes.PRODUCT_INFO.lower() @@ -1092,6 +1097,21 @@ def get_product_info_entity(self) -> Optional[ProductInfo]: return None return Activity._convert_entity(raw_product_info, ProductInfo) + def get_streaming_entity(self) -> StreamInfo | None: + """Gets the streaming entity from the activity's entities. + + :return: The streaming entity if found, otherwise None. + """ + if not self.entities: + return None + target = EntityTypes.STREAM_INFO.lower() + raw_stream_info = next( + filter(lambda e: e.type.lower() == target, self.entities), None + ) + if raw_stream_info is None: + return None + return Activity._convert_entity(raw_stream_info, StreamInfo) + def get_mentions(self) -> list[Mention]: """ Resolves the mentions from the entities of this activity. diff --git a/libraries/microsoft-agents-activity/microsoft_agents/activity/attachment.py b/libraries/microsoft-agents-activity/microsoft_agents/activity/attachment.py index af702f5af..6d31ca3c0 100644 --- a/libraries/microsoft-agents-activity/microsoft_agents/activity/attachment.py +++ b/libraries/microsoft-agents-activity/microsoft_agents/activity/attachment.py @@ -2,7 +2,6 @@ # Licensed under the MIT License. from .agents_model import AgentsModel -from ._type_aliases import NonEmptyString class Attachment(AgentsModel): @@ -20,8 +19,8 @@ class Attachment(AgentsModel): :type thumbnail_url: str """ - content_type: NonEmptyString - content_url: NonEmptyString = None + content_type: str + content_url: str | None = None content: object = None - name: NonEmptyString = None - thumbnail_url: NonEmptyString = None + name: str | None = None + thumbnail_url: str | None = None diff --git a/libraries/microsoft-agents-activity/microsoft_agents/activity/channels.py b/libraries/microsoft-agents-activity/microsoft_agents/activity/channels.py index 3896e1547..b2c8d9842 100644 --- a/libraries/microsoft-agents-activity/microsoft_agents/activity/channels.py +++ b/libraries/microsoft-agents-activity/microsoft_agents/activity/channels.py @@ -13,6 +13,9 @@ class Channels(str, Enum): Ids of channels supported by ABS. """ + a2a = "a2a" + """A2A protocol""" + agents = "agents" """Agents channel.""" diff --git a/libraries/microsoft-agents-activity/microsoft_agents/activity/conversation_reference.py b/libraries/microsoft-agents-activity/microsoft_agents/activity/conversation_reference.py index dc69d2b6a..131736926 100644 --- a/libraries/microsoft-agents-activity/microsoft_agents/activity/conversation_reference.py +++ b/libraries/microsoft-agents-activity/microsoft_agents/activity/conversation_reference.py @@ -52,7 +52,7 @@ class ConversationReference(AgentsModel): conversation: ConversationAccount channel_id: Optional[ChannelId] = None locale: Optional[NonEmptyString] = None - service_url: NonEmptyString = None + service_url: str | None = None request_id: str | None = None def get_continuation_activity(self) -> Activity: diff --git a/libraries/microsoft-agents-activity/microsoft_agents/activity/end_of_conversation_codes.py b/libraries/microsoft-agents-activity/microsoft_agents/activity/end_of_conversation_codes.py index 425cc603f..1dcebf856 100644 --- a/libraries/microsoft-agents-activity/microsoft_agents/activity/end_of_conversation_codes.py +++ b/libraries/microsoft-agents-activity/microsoft_agents/activity/end_of_conversation_codes.py @@ -11,3 +11,4 @@ class EndOfConversationCodes(str, Enum): timed_out = "botTimedOut" issued_invalid_message = "botIssuedInvalidMessage" channel_failed = "channelFailed" + error = "error" diff --git a/libraries/microsoft-agents-hosting-a2a/LICENSE b/libraries/microsoft-agents-hosting-a2a/LICENSE new file mode 100644 index 000000000..9e841e7a2 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/LICENSE @@ -0,0 +1,21 @@ + MIT License + + Copyright (c) Microsoft Corporation. + + Permission is hereby granted, free of charge, to any person obtaining a copy + of this software and associated documentation files (the "Software"), to deal + in the Software without restriction, including without limitation the rights + to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + copies of the Software, and to permit persons to whom the Software is + furnished to do so, subject to the following conditions: + + The above copyright notice and this permission notice shall be included in all + copies or substantial portions of the Software. + + THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE + SOFTWARE diff --git a/libraries/microsoft-agents-hosting-a2a/MANIFEST.in b/libraries/microsoft-agents-hosting-a2a/MANIFEST.in new file mode 100644 index 000000000..43a71d9ed --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/MANIFEST.in @@ -0,0 +1 @@ +include VERSION.txt \ No newline at end of file diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/__init__.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/__init__.py new file mode 100644 index 000000000..bc05aff99 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/__init__.py @@ -0,0 +1,30 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from .activity import A2AActivity +from .request_handling import A2AHttpAdapter +from .server import ( + create_jsonrpc_routes, + create_rest_routes, + SDKServerCallContextBuilder, +) +from .extension import ( + A2AAgentExtension, + A2AClient, + A2ATurnContext, +) +from .a2a_adapter import A2AAdapter +from .add_a2a import add_a2a + +__all__ = [ + "A2AActivity", + "A2AHttpAdapter", + "create_jsonrpc_routes", + "create_rest_routes", + "SDKServerCallContextBuilder", + "A2AAdapter", + "A2AAgentExtension", + "A2AClient", + "A2ATurnContext", + "add_a2a", +] diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/a2a_adapter.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/a2a_adapter.py new file mode 100644 index 000000000..655b602f7 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/a2a_adapter.py @@ -0,0 +1,503 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import logging + +from datetime import datetime, timezone +from urllib.parse import urlsplit +from uuid import uuid4 + + +from a2a.types import ( + AgentCard, + AgentInterface, + TaskStatusUpdateEvent, + TaskArtifactUpdateEvent, + TaskStatus, + TaskState, + AgentCapabilities, + AgentSkill, + HTTPAuthSecurityScheme, + SecurityScheme, +) +from a2a.server.request_handlers import RequestHandler +from a2a.server.agent_execution import RequestContext +from a2a.server.events import EventQueue +from a2a.server.tasks import TaskStore, InMemoryTaskStore +from a2a.utils.constants import TransportProtocol + + +from microsoft_agents.activity import ( + Activity, + ActivityTypes, + CallerIdConstants, + Channels, + ChannelAccount, + ConversationAccount, + ChannelId, + EndOfConversationCodes, + InvokeResponse, + ResourceResponse, + RoleTypes, + StreamInfo, +) +from microsoft_agents.hosting.core import ( + AgentApplication, + AuthenticationConstants, + ChannelAdapter, + ChannelServiceAdapter, + ClaimsIdentity, + TurnContext, + MiddlewareSet, +) +from microsoft_agents.hosting.core.channel_adapter_protocol import ( + ChannelAdapterProtocol, +) +from microsoft_agents.hosting.core.http._http_request_protocol import ( + HttpRequestProtocol, +) +from microsoft_agents.hosting.core.http._http_response import HttpResponse + +from .request_handling import ( + A2AHttpAdapter, + A2ARequestHandler, +) + +from .activity import utils, A2AActivity +from .activity.a2a_activity import _DEFAULT_USER_ID +from .extension.a2a_turn_context import A2ATurnContext + +from .server._constants import _CLAIMS_IDENTITY_KEY + +logger = logging.getLogger(__name__) + + +class A2AAdapter(A2AHttpAdapter, ChannelAdapter, ChannelAdapterProtocol): + """Adapter for handling Agent-to-Agent (A2A) communication within the Microsoft Agents framework.""" + + def __init__( + self, + agent: AgentApplication, + *, + agent_card_name: str = "A2AAdapter", + agent_card_description: str = "Agents SDK A2A", + agent_card_version: str = "0.0.0", + agent_interfaces: list[AgentInterface] | None = None, + skills: list[AgentSkill] | None = None, + task_store: TaskStore | None = None, + ): + """Initializes the A2AAdapter with the given agent and optional task store. + + :param agent: The agent instance to be used by the adapter. + :param agent_card_name: The name of the agent card. + :param agent_card_description: The description of the agent card. + :param agent_card_version: The version of the agent card. + :param agent_interfaces: The list of agent interfaces associated with the adapter. + :param skills: The list of skills associated with the adapter. + :param task_store: Optional task store for managing tasks. If not provided, an in-memory task store will be used. + """ + self.middleware_set = MiddlewareSet() + + self._agent = agent + self._agent_card_name = agent_card_name + self._agent_card_description = agent_card_description + self._agent_card_version = agent_card_version + self._task_store = task_store or InMemoryTaskStore() + + self._skills: list[AgentSkill] = skills or [] + self._agent_interfaces: list[AgentInterface] = agent_interfaces or [] + + self._a2a_request_handler = A2ARequestHandler( + self, + task_store=self._task_store, + agent_card=self._get_basic_agent_card(), + ) + + @property + def skills(self) -> list[AgentSkill]: + """Get a copy of the list of skills associated with the adapter.""" + return list(self._skills) + + @property + def agent_interfaces(self) -> list[AgentInterface]: + """Get a copy of the list of agent interfaces associated with the adapter.""" + return list(self._agent_interfaces) + + @property + def a2a_request_handler(self) -> RequestHandler: + """Get the A2A request handler.""" + return self._a2a_request_handler + + async def execute_agent_turn( + self, + context: RequestContext, + event_queue: EventQueue, + ) -> None: + """Execute an agent turn given the request context and event queue. + + :param context: The request context for the agent turn. + :param event_queue: The event queue for the agent turn. + """ + + if not context.message: + raise ValueError("Context message is required.") + + identity = context.call_context.state.get( + _CLAIMS_IDENTITY_KEY, ClaimsIdentity() + ) + if not isinstance(identity, ClaimsIdentity): + raise RuntimeError("Invalid identity in context call state.") + + request_id = str(uuid4()) + + activity = A2AActivity.from_message( + request_id, + context.task_id, + context.message, + ) + activity.request_id = request_id + + await self._process_activity_with_a2a( + identity, + activity, + context, + event_queue, + ) + + async def cancel_agent_turn( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + """Cancel an ongoing agent turn given the request context and event queue. + + :param context: The request context for the agent turn. + :param event_queue: The event queue for the agent turn. + """ + + end_of_conv_activity = Activity( + type=ActivityTypes.end_of_conversation, + code=EndOfConversationCodes.user_cancelled, + channel_id=ChannelId(Channels.a2a), + conversation=ConversationAccount(id=context.task_id or ""), + recipient=ChannelAccount(id="assistant", role=RoleTypes.agent), + from_property=ChannelAccount(id=_DEFAULT_USER_ID, role=RoleTypes.user), + ) + + identity = context.call_context.state.get( + _CLAIMS_IDENTITY_KEY, ClaimsIdentity() + ) + if not isinstance(identity, ClaimsIdentity): + raise RuntimeError("Invalid identity in context call state.") + + await self._process_activity_with_a2a( + identity, + end_of_conv_activity, + context, + event_queue, + ) + + def _create_turn_context( + self, + claims_identity: ClaimsIdentity, + oauth_scope: str | None, + activity: Activity, + request_context: RequestContext, + event_queue: EventQueue, + task_store: TaskStore, + ) -> TurnContext: + """Create a turn context for the given claims identity, OAuth scope, and activity. + + :param claims_identity: The claims identity for the turn context. + :param oauth_scope: The OAuth scope for the turn context. + :param activity: The activity for the turn context. + :param request_context: The request context for the turn context. + :param event_queue: The event queue for the turn context. + :param task_store: The task store for the turn context. + :return: The created turn context. + """ + context = A2ATurnContext( + self, + app=self._agent, + activity=activity, + identity=claims_identity, + request_context=request_context, + event_queue=event_queue, + task_store=task_store, + ) + context.turn_state[ChannelServiceAdapter.OAUTH_SCOPE_KEY] = oauth_scope + context.turn_state[ChannelServiceAdapter.AGENT_IDENTITY_KEY] = ( + claims_identity # for back-compat + ) + return context + + async def _process_activity_with_a2a( + self, + identity: ClaimsIdentity, + activity: Activity, + request_context: RequestContext, + event_queue: EventQueue, + ) -> InvokeResponse | None: + """Process an activity with the A2A adapter. + + :param identity: The claims identity of the agent. + :param activity: The activity to process. + :param request_context: The request context for the activity. + :param event_queue: The event queue for the activity. + :return: An InvokeResponse if applicable, otherwise None. + """ + + if activity.channel_id != Channels.a2a: + raise ValueError("Activity channel_id must be 'a2a'") + + outgoing_audience: str | None = None + + if identity.is_agent_claim(): + outgoing_audience = identity.get_token_audience() + activity.caller_id = f"{CallerIdConstants.agent_to_agent_prefix}{identity.get_outgoing_app_id()}" + else: + outgoing_audience = AuthenticationConstants.AGENTS_SDK_SCOPE + + # Create a turn context and run the pipeline. + context = self._create_turn_context( + identity, + outgoing_audience, + activity=activity, + request_context=request_context, + event_queue=event_queue, + task_store=self._task_store, + ) + + await self.run_pipeline(context, self._agent.on_turn) + + async def send_activities( + self, context: TurnContext, activities: list[Activity] + ) -> list[ResourceResponse]: + """Send activities through the A2A adapter. + + :param context: The turn context for the activities. + :param activities: The list of activities to send. + :return: A list of resource responses. + """ + + for activity in activities: + + if activity.channel_id != Channels.a2a: + continue + + entity = activity.get_streaming_entity() + + if entity is not None: + await self._on_streaming_response(context, activity, entity) + elif activity.type == ActivityTypes.message: + await self._on_message_response(context, activity) + elif activity.type == ActivityTypes.end_of_conversation: + await self._on_end_of_conversation_response(context, activity) + else: + logger.debug("A2AAdapter: Unhandled Activity Type: %s", activity.type) + + return [] + + async def _on_streaming_response( + self, context: TurnContext, activity: Activity, entity: StreamInfo + ): + """Handle a streaming response activity. + + :param context: The turn context for the activity. + :param activity: The activity containing the streaming response. + :param entity: The streaming entity associated with the activity. + """ + message = utils.get_incoming_message(context) + is_informative = entity.stream_type == "informative" + + event_queue = context.services.get(EventQueue, raise_if_missing=True) + + if is_informative: + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id=message.task_id, + context_id=message.context_id, + status=TaskStatus( + state=TaskState.TASK_STATE_WORKING, + timestamp=datetime.now(timezone.utc), + message=utils.create_message( + message.context_id, message.task_id, activity + ), + ), + ) + ) + else: + artifact = utils.activity_to_artifact(activity, entity.stream_id) + await event_queue.enqueue_event( + TaskArtifactUpdateEvent( + task_id=message.task_id, + context_id=message.context_id, + artifact=artifact, + append=False, + last_chunk=True, + ) + ) + + async def _on_message_response(self, context: TurnContext, activity: Activity): + """Handle a message response activity. + + :param context: The turn context for the activity. + :param activity: The activity containing the message response. + """ + message = utils.get_incoming_message(context) + state = utils.get_task_state(activity) + response = utils.create_message(message.context_id, message.task_id, activity) + + event_queue = context.services.get(EventQueue, raise_if_missing=True) + + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id=message.task_id, + context_id=message.context_id, + status=TaskStatus( + state=state, timestamp=datetime.now(timezone.utc), message=response + ), + ) + ) + + async def _on_end_of_conversation_response( + self, context: TurnContext, activity: Activity + ): + """Handle an end-of-conversation response activity. + + :param context: The turn context for the activity. + :param activity: The activity containing the end-of-conversation response. + """ + message = utils.get_incoming_message(context) + event_queue = context.services.get(EventQueue, raise_if_missing=True) + + if isinstance(activity.value, dict): + artifact = utils.create_artifact_from_data( + activity.value, name="Result", description="Task completion result" + ) + await event_queue.enqueue_event( + TaskArtifactUpdateEvent( + task_id=message.task_id, + context_id=message.context_id, + artifact=artifact, + append=False, + last_chunk=True, + ) + ) + + task_state: TaskState + if activity.code == EndOfConversationCodes.error: + task_state = TaskState.TASK_STATE_FAILED + elif activity.code == EndOfConversationCodes.user_cancelled: + task_state = TaskState.TASK_STATE_CANCELED + else: + task_state = TaskState.TASK_STATE_COMPLETED + + status_message: Activity | None = None + if utils.has_message_content(activity): + status_message = activity.model_copy() + status_message.value = None + + response = utils.create_message( + message.context_id, + message.task_id, + status_message, + ) + + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id=message.task_id, + context_id=message.context_id, + status=TaskStatus( + state=task_state, + timestamp=datetime.now(timezone.utc), + message=response, + ), + ) + ) + + def _get_basic_agent_card(self) -> AgentCard: + """Get the basic agent card with default settings.""" + return AgentCard( + name=self._agent_card_name, + description=self._agent_card_description, + version=self._agent_card_version, + security_schemes={ + "jwt": SecurityScheme( + http_auth_security_scheme=HTTPAuthSecurityScheme(scheme="bearer") + ) + }, + default_input_modes=["application/json"], + default_output_modes=["application/json"], + skills=[], + capabilities=AgentCapabilities( + extended_agent_card=False, + streaming=True, + ), + supported_interfaces=[], + ) + + async def get_agent_card( + self, request: HttpRequestProtocol, path_prefix: str + ) -> AgentCard: + """Get the agent card for the current agent, potentially customized based on the request. + + Set as asynchronous because in some implementations, fetching or customizing the agent card might involve I/O operations, such as querying a database or an external service. + + :param request: The HTTP request object conforming to HttpRequestProtocol. + :param path_prefix: The prefix to be used for constructing the agent interface URL. + :return: An AgentCard instance representing the agent's capabilities. + """ + agent_card = self._get_basic_agent_card() + + url_parts = urlsplit(request.url) + request_origin = f"{url_parts.scheme}://{url_parts.netloc}" + + if not self._agent_interfaces: + agent_card.supported_interfaces.append( + AgentInterface( + protocol_binding=TransportProtocol.JSONRPC, + url=f"{request_origin}{path_prefix}/", + protocol_version="1.0", + ) + ) + else: + for agent_interface in self._agent_interfaces: + if agent_interface.protocol_binding in ( + TransportProtocol.JSONRPC, + TransportProtocol.HTTP_JSON, + ): + interface_url = agent_interface.url + if interface_url.startswith("/"): + interface_url = f"{request_origin}{interface_url}" + + agent_card.supported_interfaces.append( + AgentInterface( + protocol_binding=agent_interface.protocol_binding, + url=interface_url, + protocol_version="1.0", + ) + ) + else: + logger.info( + "Unsupported protocol: %s", agent_interface.protocol_binding + ) + + if self._skills: + for skill_info in self._skills: + agent_card.skills.append( + AgentSkill( + id=skill_info.id, + name=skill_info.name, + description=skill_info.description, + tags=skill_info.tags, + examples=skill_info.examples, + input_modes=skill_info.input_modes, + output_modes=skill_info.output_modes, + ) + ) + return agent_card + + async def update_activity(self, context: TurnContext, activity: Activity) -> None: + raise NotImplementedError("A2AAdapter.update_activity is not implemented.") + + async def delete_activity(self, context: TurnContext, activity_id: str) -> None: + raise NotImplementedError("A2AAdapter.delete_activity is not implemented.") diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/activity/__init__.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/activity/__init__.py new file mode 100644 index 000000000..fa2508d30 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/activity/__init__.py @@ -0,0 +1,6 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from .a2a_activity import A2AActivity + +__all__ = ["A2AActivity"] diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/activity/a2a_activity.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/activity/a2a_activity.py new file mode 100644 index 000000000..77bcad5e1 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/activity/a2a_activity.py @@ -0,0 +1,165 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +"""A2A-aware :class:`Activity` subclass exposing A2A protocol helpers.""" + +from __future__ import annotations + +from typing import Iterable +from uuid import uuid4 + +from a2a.types import ( + Artifact, + Message, + Part, + TaskState, +) +from google.protobuf.json_format import MessageToDict + +from pydantic import Field + +from microsoft_agents.activity import ( + Activity, + ActivityTypes, + Attachment, + Channels, + ChannelId, + ChannelAccount, + ConversationAccount, + DeliveryModes, + RoleTypes, +) + +from . import utils + +_DEFAULT_USER_ID = "unknown" + + +class A2AActivity(Activity): + """A2A-aware :class:`Activity` subclass exposing A2A protocol data helpers.""" + + channel_id: ChannelId = Field( + default_factory=lambda: ChannelId(Channels.a2a), frozen=True + ) + + @staticmethod + def from_message( + request_id: str, task_id: str | None, message: Message + ) -> A2AActivity: + """Create an A2AActivity from a Message object. + + :param request_id: The request ID for the activity. + :param task_id: The task ID for the activity. If None, the task ID from the message will be used. + :param message: The Message object to convert. + :return: An A2AActivity object representing the message. + """ + + if not task_id and not message.task_id: + raise ValueError("Task ID is required.") + + task_id = task_id or message.task_id + + activity = A2AActivity._create_activity(task_id, message.parts, True, True) + + activity.request_id = request_id or str(uuid4()) # formatting + activity.channel_data = message + message.context_id = message.context_id or str(uuid4()) # Formatting + message.task_id = task_id + + return activity + + def to_message( + self, context_id: str, task_id: str, include_entities: bool = True + ) -> Message: + """Convert the activity to a Message object. + + :param context_id: The context ID for the message. + :param task_id: The task ID for the message. + :param include_entities: Whether to include entities in the message. Defaults to True. + :return: A Message object representing the activity. + """ + return utils.create_message(context_id, task_id, self, include_entities) + + def to_artifact( + self, artifact_id: str | None = None, include_entities: bool = True + ) -> Artifact: + """Convert the activity to an Artifact object. + + :param artifact_id: The ID of the artifact. If None, a new ID will be generated. + :param include_entities: Whether to include entities in the artifact. Defaults to True. + :return: An Artifact object representing the activity. + """ + return utils.activity_to_artifact(self, artifact_id, include_entities) + + def get_task_state(self) -> TaskState: + """Get the current task state of the activity. + + :return: The TaskState of the activity. + """ + return utils.get_task_state(self) + + @staticmethod + def _create_activity( + conversation_id: str, + parts: Iterable[Part], + is_ingress: bool, + is_streaming: bool, + ) -> A2AActivity: + """Create an Activity representing an A2A concept. + + :param conversation_id: The ID of the conversation. + :param parts: The parts to include in the activity. + :param is_ingress: Whether the activity is incoming. + :param is_streaming: Whether the activity is streaming. + :return: An A2AActivity object representing the activity. + """ + + agent = ChannelAccount(id="assistant", role=RoleTypes.agent) + user = ChannelAccount(id=_DEFAULT_USER_ID, role=RoleTypes.user) + + activity = A2AActivity( + type=ActivityTypes.message, + id=str(uuid4()), + delivery_mode=( + DeliveryModes.stream if is_streaming else DeliveryModes.expect_replies + ), + conversation=ConversationAccount(id=conversation_id), + recipient=agent if is_ingress else user, + from_property=user if is_ingress else agent, + attachments=[], + ) + + for part in parts: + content_kind = part.WhichOneof("content") + + if content_kind == "text": + if activity.text is None: + activity.text = part.text + else: + activity.text += part.text + elif content_kind == "url": + activity.attachments.append( + Attachment( + content_type=part.media_type, + content_url=part.url, + name=part.filename, + ) + ) + elif content_kind == "raw": + activity.attachments.append( + Attachment( + content_type=part.media_type, + content=bytes(part.raw), + name=part.filename, + ) + ) + elif content_kind == "data": + activity.attachments.append( + Attachment( + content_type=part.media_type or "application/json", + content=MessageToDict(part.data), + name=part.filename, + ) + ) + + return activity diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/activity/utils.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/activity/utils.py new file mode 100644 index 000000000..dd59bde15 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/activity/utils.py @@ -0,0 +1,216 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from functools import lru_cache +from typing import Mapping, cast, Sequence, Any +from uuid import uuid4 + +from pydantic import BaseModel, TypeAdapter +from pydantic.errors import PydanticSchemaGenerationError + +from a2a.types import ( + Artifact, + Message, + Part, + Role, + TaskState, +) +from google.protobuf.json_format import ParseDict +from google.protobuf.struct_pb2 import Value + +from microsoft_agents.activity import ( + Activity, + Entity, + InputHints, + StreamInfo, +) +from microsoft_agents.hosting.core import TurnContext + +_ENTITY_TYPE_TEMPLATE = "application/vnd.microsoft.entity.{0}" +_SCHEMAS: dict[str, Mapping] = {} + + +def _to_protobuf_value(data: dict) -> Value: + return ParseDict(data, Value()) + + +def activity_to_artifact( + activity: Activity, artifact_id: str | None = None, include_entities: bool = True +) -> Artifact: + """Convert an Activity object to an Artifact object. + + :param activity: The activity to convert. + :param artifact_id: The ID of the artifact. If None, a new ID will be generated. + :param include_entities: Whether to include entities in the artifact. Defaults to True. + :return: An Artifact object representing the activity. + """ + + artifact = Artifact( + artifact_id=artifact_id if artifact_id else str(uuid4()), + ) + + if activity.text is not None: + artifact.parts.append(Part(text=activity.text)) + + if activity.value is not None and isinstance(activity.value, dict): + artifact.parts.append(Part(data=_to_protobuf_value(activity.value))) + + for attachment in activity.attachments or []: + + part: Part + if attachment.content_url: + part = Part( + url=attachment.content_url, + media_type=attachment.content_type, + filename=attachment.name, + ) + elif isinstance(attachment.content, dict): + part = Part( + data=_to_protobuf_value(attachment.content), + media_type=attachment.content_type, + filename=attachment.name, + ) + elif isinstance(attachment.content, (bytes, bytearray, memoryview)): + part = Part( + raw=bytes(attachment.content), + media_type=attachment.content_type, + filename=attachment.name, + ) + else: + continue + + artifact.parts.append(part) + + if include_entities: + for entity in activity.entities or []: + if not isinstance(entity, StreamInfo): + artifact.parts.append( + Part( + metadata=_get_a2a_metadata(entity), + data=_to_protobuf_value(entity.model_dump(exclude_none=True)), + ) + ) + + return artifact + + +@lru_cache +def _try_get_json_schema(data_type: type) -> dict[str, Any] | None: + try: + if issubclass(data_type, BaseModel): + return data_type.model_json_schema() + return TypeAdapter(data_type).json_schema() + except PydanticSchemaGenerationError: + return None + + +def _get_a2a_metadata(data: object) -> dict: + """Convert the given data to A2A metadata. + + :param data: The data to convert to A2A metadata. + :param content_type: The content type of the data. + :return: A dictionary representing the A2A metadata. + """ + if data is None: + raise ValueError("Data cannot be None.") + + metadata: dict[str, Any] = { + "mimeType": "application/json", + "type": "object", + } + schema = _try_get_json_schema(type(data)) + if schema is not None: + metadata["schema"] = schema + + return metadata + + +def create_artifact_from_data( + data: dict, + name: str | None = None, + description: str | None = None, + artifact_id: str | None = None, +) -> Artifact: + """Create an Artifact object from the given data. + + :param data: The data to include in the artifact. + :param name: The name of the artifact. Defaults to None. + :param description: The description of the artifact. Defaults to None. + :param artifact_id: The ID of the artifact. If None, a new ID will be generated. + :return: An Artifact object representing the data. + """ + return Artifact( + artifact_id=artifact_id if artifact_id else str(uuid4()), # check .NET + name=name, + description=description, + parts=[ + Part( + data=_to_protobuf_value(data), + metadata=_get_a2a_metadata(data), + ) + ], + ) + + +def create_message( + context_id: str, + task_id: str, + activity: Activity | None, + include_entities: bool = True, +) -> Message: + """Create a Message object. Uses the given activity to populate the message parts. + + :param context_id: The context ID for the message. + :param task_id: The task ID for the message. + :param activity: The activity to convert. + :param include_entities: Whether to include entities in the message. Defaults to True. + :return: A Message object representing the activity. + """ + + parts: Sequence[Part] | None + + if not activity: + parts = None + else: + artifact = activity_to_artifact(activity, include_entities=include_entities) + parts = artifact.parts + + return Message( + task_id=task_id, + context_id=context_id, + message_id=str(uuid4()), + parts=parts, + role=Role.ROLE_AGENT, + ) + + +def has_message_content(activity: Activity) -> bool: + """Check if the activity has message content. + + :param activity: The activity to check. + :return: True if the activity has message content, False otherwise. + """ + return bool(activity.text) or bool(activity.attachments) + + +def get_task_state(activity: Activity) -> TaskState: + """Get the current task state of the activity. + + :return: The TaskState of the activity. + """ + if activity.input_hint == InputHints.expecting_input: + return TaskState.TASK_STATE_INPUT_REQUIRED + return TaskState.TASK_STATE_WORKING + + +def get_incoming_message(context: TurnContext) -> Message: + """Get the incoming message from the turn context. + + :param context: The turn context containing the activity. + :return: The Message object from the channel data. + :raises TypeError: If the channel data is not of type Message. + """ + data = context.activity.channel_data + if isinstance(data, Message): + return cast(Message, data) + raise TypeError("Expected channel_data to be of type Message") diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/add_a2a.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/add_a2a.py new file mode 100644 index 000000000..ba02ad0ee --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/add_a2a.py @@ -0,0 +1,157 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import logging + +from fastapi import FastAPI +from starlette.routing import BaseRoute, Route + +from a2a.server.routes import add_a2a_routes_to_fastapi +from a2a.types import AgentInterface +from a2a.utils.constants import TransportProtocol, AGENT_CARD_WELL_KNOWN_PATH + +from microsoft_agents.hosting.core import AgentApplication + +from .a2a_adapter import A2AAdapter +from .server import ( + create_jsonrpc_routes, + create_rest_routes, + create_agent_card_routes, + use_jwt_middleware as _use_jwt_middleware, +) + +logger = logging.getLogger(__name__) + + +def _create_jsonrpc_interface_routes( + adapter: A2AAdapter, + interface: AgentInterface, + _agent_card_cache_enabled: bool = False, + agent_card_cache_max_age: int = 3600, +) -> tuple[list[Route], list[BaseRoute]]: + """Create JSON-RPC interface routes for the given agent interface. + + :param adapter: The A2AAdapter instance used to handle requests. + :param interface: The agent interface to add routes for. + :param _agent_card_cache_enabled: Whether to enable caching for the agent card. Default is False. + :param agent_card_cache_max_age: The maximum age (in seconds) for caching the agent card. Default is 3600 seconds (1 hour). + """ + + jsonrpc_routes = create_jsonrpc_routes( + adapter.a2a_request_handler, rpc_url=interface.url + ) + agent_card_routes = create_agent_card_routes( + adapter.get_agent_card, + f"{interface.url}{AGENT_CARD_WELL_KNOWN_PATH}", + _cache_enabled=_agent_card_cache_enabled, + cache_max_age=agent_card_cache_max_age, + ) + return jsonrpc_routes, agent_card_routes + + +def _create_http_interface_routes( + adapter: A2AAdapter, + interface: AgentInterface, + _agent_card_cache_enabled: bool = False, + agent_card_cache_max_age: int = 3600, +) -> tuple[list[BaseRoute], list[BaseRoute]]: + """Create HTTP interface routes for the given agent interface. + + :param adapter: The A2AAdapter instance used to handle requests. + :param interface: The agent interface to add routes for. + :param _agent_card_cache_enabled: Whether to enable caching for the agent card. Default is False. + :param agent_card_cache_max_age: The maximum age (in seconds) for caching the agent card. Default is 3600 seconds (1 hour). + """ + http_routes = create_rest_routes( + adapter.a2a_request_handler, path_prefix=interface.url + ) + agent_card_routes = create_agent_card_routes( + adapter.get_agent_card, + f"{interface.url}{AGENT_CARD_WELL_KNOWN_PATH}", + _cache_enabled=_agent_card_cache_enabled, + cache_max_age=agent_card_cache_max_age, + ) + return http_routes, agent_card_routes + + +def add_a2a( + app: FastAPI, + agent: AgentApplication, + adapter: A2AAdapter | None = None, + *, + use_jwt_middleware: bool = True, + _agent_card_cache_enabled: bool = False, + agent_card_cache_max_age: int = 3600, +): + """Add A2A support to the given FastAPI app using the specified agent and adapter. + + :param app: The FastAPI app to which A2A support will be added. + :param agent: The agent to use for A2A communication. + :param adapter: An optional A2AAdapter instance. If not provided, a new one will be created using the agent. + :param use_jwt_middleware: Whether to use JWT middleware for the routes. + :param _agent_card_cache_enabled: Whether to enable caching for the agent card. Default is False. + :param agent_card_cache_max_age: The maximum age (in seconds) for caching the agent card. Default is 3600 seconds (1 hour). + """ + + adapter = adapter or A2AAdapter( + agent, + agent_interfaces=[ + AgentInterface( + url="/a2a", + protocol_binding=TransportProtocol.JSONRPC, + ) + ], + ) + + agent_card_routes: list[BaseRoute] = [] + jsonrpc_routes: list[BaseRoute] = [] + rest_routes: list[BaseRoute] = [] + + interfaces = adapter.agent_interfaces + if not interfaces: + raise ValueError( + "No agent interfaces found. Cannot add A2A routes to application." + ) + + unsupported_counter = 0 + for interface in interfaces: + if interface.protocol_binding == TransportProtocol.JSONRPC: + _jsonrpc_routes, _agent_card_routes = _create_jsonrpc_interface_routes( + adapter, + interface, + _agent_card_cache_enabled=_agent_card_cache_enabled, + agent_card_cache_max_age=agent_card_cache_max_age, + ) + jsonrpc_routes.extend(_jsonrpc_routes) + agent_card_routes.extend(_agent_card_routes) + elif interface.protocol_binding == TransportProtocol.HTTP_JSON: + _http_routes, _agent_card_routes = _create_http_interface_routes( + adapter, + interface, + _agent_card_cache_enabled=_agent_card_cache_enabled, + agent_card_cache_max_age=agent_card_cache_max_age, + ) + rest_routes.extend(_http_routes) + agent_card_routes.extend(_agent_card_routes) + else: + unsupported_counter += 1 + logger.warning( + "Unsupported protocol binding: %s", interface.protocol_binding + ) + + if unsupported_counter == len(interfaces): + raise ValueError( + "All agent interfaces have unsupported protocol bindings. Cannot add A2A routes to application." + ) + + if use_jwt_middleware: + _use_jwt_middleware(agent_card_routes) + _use_jwt_middleware(jsonrpc_routes) + _use_jwt_middleware(rest_routes) + + add_a2a_routes_to_fastapi( + app, + agent_card_routes=agent_card_routes, + jsonrpc_routes=jsonrpc_routes, + rest_routes=rest_routes, + ) diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/blob_task_store.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/blob_task_store.py new file mode 100644 index 000000000..1481c6ce0 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/blob_task_store.py @@ -0,0 +1,196 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import urllib.parse + +from asyncio import Lock +from typing import overload + +from a2a.server.context import ServerCallContext +from a2a.server.tasks import TaskStore +from a2a.types.a2a_pb2 import ( + ListTasksRequest, + ListTasksResponse, + Task, + TaskState, +) + + +from azure.storage.blob.aio import ( + ContainerClient, + BlobServiceClient, +) +from microsoft_agents.hosting.core.storage.error_handling import ( + ignore_error, + is_status_code_error, +) + +_TASK_PREFIX = "a2atask/" + + +class BlobTaskStore(TaskStore): + """A task store implementation backed by a blob storage.""" + + @overload + def __init__( + self, + container_client: ContainerClient, + ) -> None: ... + @overload + def __init__( + self, + *, + data_connection_string: str, + container_name: str, + ) -> None: ... + def __init__( + self, + container_client: ContainerClient | None = None, + *, + data_connection_string: str | None = None, + container_name: str | None = None, + ) -> None: + """Initializes the blob task store with either a container client or a connection string and container name. + + :param container_client: An optional ContainerClient instance for accessing the blob storage. + :param data_connection_string: An optional connection string for the blob storage. + :param container_name: An optional container name within the blob storage. + """ + + self._initialized = False + self._init_lock = Lock() + + if container_client is not None: + self._container_client = container_client + return + elif data_connection_string is not None and container_name is not None: + blob_service_client = BlobServiceClient.from_connection_string( + data_connection_string + ) + self._container_client = blob_service_client.get_container_client( + container_name + ) + return + else: + raise ValueError("Invalid combination of parameters provided.") + + async def _ensure_container_exists(self) -> None: + """Ensures that the blob container exists, creating it if necessary.""" + + if not self._initialized: + async with self._init_lock: + if not self._initialized: + await ignore_error( + self._container_client.create_container(), + is_status_code_error(409), + ) + self._initialized = True + + @staticmethod + def _should_include_task(task: Task, request: ListTasksRequest) -> bool: + """Determines whether a task should be included in the list based on the request parameters.""" + + if ( + request.status != TaskState.TASK_STATE_UNSPECIFIED + and request.status != task.status.state + ): + return False + if request.context_id and request.context_id != task.context_id: + return False + return True + + @staticmethod + def _get_blob_name(task_id: str) -> str: + """Generates a blob name for the given task ID.""" + + if not task_id: + raise ValueError("Task ID cannot be empty.") + return f"{_TASK_PREFIX}{urllib.parse.quote_plus(task_id)}" + + async def save(self, task: Task, context: ServerCallContext) -> None: + """Saves or updates a task in the blob store.""" + await self._ensure_container_exists() + + blob_name = self._get_blob_name(task.id) + blob_client = self._container_client.get_blob_client(blob_name) + serialized_task = task.SerializeToString() + + await blob_client.upload_blob( + data=serialized_task, + overwrite=True, + length=len(serialized_task), + ) + + async def get(self, task_id: str, context: ServerCallContext) -> Task | None: + """Retrieves a task from the blob store by its ID. + + :param task_id: The ID of the task to retrieve. + :param context: The server call context. + :return: The task if found, otherwise None. + """ + await self._ensure_container_exists() + + return await self._download_task(task_id) + + async def list( + self, params: ListTasksRequest, context: ServerCallContext + ) -> ListTasksResponse: + """Retrives a list of tasks from the store. + + :param params: The parameters for listing tasks. + :param context: The server call context. + :return: A response containing the list of tasks. + """ + await self._ensure_container_exists() + + tasks: list[Task] = [] + + kwargs = {} + if params.page_size > 0: + kwargs["results_per_page"] = params.page_size + + items = self._container_client.list_blobs( + name_starts_with=_TASK_PREFIX, **kwargs + ) + + iterator = items.by_page(params.page_token) + first_page = await anext(iterator) + + if not first_page: + return ListTasksResponse(tasks=[], next_page_token=None) + + async for blob in first_page: + task = await self._download_task_blob(blob.name) + if task and self._should_include_task(task, params): + tasks.append(task) + + next_page_token = getattr(iterator, "continuation_token", None) + return ListTasksResponse(tasks=tasks, next_page_token=next_page_token) + + async def delete(self, task_id: str, context: ServerCallContext) -> None: + """Deletes a task from the blob store by its ID. + + :param task_id: The ID of the task to delete. + :param context: The server call context.s + """ + await self._ensure_container_exists() + + blob_name = self._get_blob_name(task_id) + await self._container_client.get_blob_client(blob_name).delete_blob() + + async def _download_task(self, task_id: str) -> Task | None: + return await self._download_task_blob(self._get_blob_name(task_id)) + + async def _download_task_blob(self, blob_name: str) -> Task | None: + item = await ignore_error( + self._container_client.download_blob(blob=blob_name, timeout=5), + is_status_code_error(404), + ) + if not item: + return None + + serialized_task: bytes = await item.readall() + task = Task() + task.ParseFromString(serialized_task) + + return task diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/__init__.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/__init__.py new file mode 100644 index 000000000..136d19cb0 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/__init__.py @@ -0,0 +1,12 @@ +from .a2a_agent_extension import A2AAgentExtension +from .a2a_client import A2AClient +from .a2a_turn_context import A2ATurnContext +from .route_handlers import wrap_a2a_route_handler, A2ARouteHandler + +__all__ = [ + "A2AAgentExtension", + "A2AClient", + "A2ATurnContext", + "wrap_a2a_route_handler", + "A2ARouteHandler", +] diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/a2a_agent_extension.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/a2a_agent_extension.py new file mode 100644 index 000000000..1c5f30073 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/a2a_agent_extension.py @@ -0,0 +1,72 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from typing import Callable, Generic, Protocol +from re import Pattern + +from microsoft_agents.hosting.core.app import ( + AgentApplication, + RouteHandler, +) +from microsoft_agents.hosting.core.app._type_defs import StateT + +from .route_handlers import A2ARouteHandler, wrap_a2a_route_handler + + +class _AppRouteDecorator(Protocol[StateT]): + """Protocol for a decorator returned by :class:`TeamsAgentExtension` route methods.""" + + def __call__(self, func: A2ARouteHandler[StateT], /) -> RouteHandler[StateT]: + """Register *func* as an A2A route handler. + + :param func: A2A-aware handler to register. + :return: The wrapped core route handler. + """ + ... + + +class A2AAgentExtension(Generic[StateT]): + """Extension for adding A2A-specific routes to an AgentApplication.""" + + def __init__(self, app: AgentApplication[StateT]) -> None: + """Initialize the A2AAgentExtension with the given AgentApplication. + + :param app: The AgentApplication instance to extend. + """ + self._app = app + + def _wrap_decorator( + self, decorator: Callable[[RouteHandler[StateT]], RouteHandler[StateT]] + ) -> Callable[[A2ARouteHandler[StateT]], RouteHandler[StateT]]: + """Wrap a core route decorator so it accepts a :class:`A2ARouteHandler`. + + The returned decorator converts the A2A handler via + :func:`wrap_a2a_route_handler` before passing it to *decorator*, keeping the + A2A context upgrade transparent to callers. + + :param decorator: A core route decorator from :class:`AgentApplication`. + :return: A decorator that accepts and registers a :class:`A2ARouteHandler`. + """ + + def __call(func: A2ARouteHandler[StateT]) -> RouteHandler[StateT]: + return decorator(wrap_a2a_route_handler(func, self._app)) + + return __call + + def message( + self, + select: str | Pattern[str] | list[str | Pattern[str]], + *, + auth_handlers: list[str] | None = None, + **kwargs, + ) -> _AppRouteDecorator[StateT]: + """Register a handler for message activities matching *select*. + + :param select: A literal string, regex pattern, or list of either to match against + the message text. + :param auth_handlers: Optional list of auth handler names to run before the route. + :return: A decorator that registers the handler and returns it. + """ + return self._wrap_decorator( + self._app.message(select, auth_handlers=auth_handlers, **kwargs) + ) diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/a2a_client.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/a2a_client.py new file mode 100644 index 000000000..a22c02bfd --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/a2a_client.py @@ -0,0 +1,44 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from a2a.server.events import EventQueue +from a2a.server.agent_execution import RequestContext +from a2a.server.tasks import TaskStore + +from microsoft_agents.hosting.core import TurnContext + + +class A2AClient: + """A client for interacting with the A2A services within a TurnContext.""" + + def __init__(self, context: TurnContext): + """Initialize the A2AClient with the given TurnContext. + + :param context: The TurnContext containing the required services. + """ + + event_queue = context.services.get(EventQueue) + request_context = context.services.get(RequestContext) + task_store = context.services.get(TaskStore) + + if not event_queue or not request_context or not task_store: + raise ValueError("Missing required services in TurnContext.") + + self._event_queue = event_queue + self._request_context = request_context + self._task_store = task_store + + @property + def event_queue(self) -> EventQueue: + """Get the event queue associated with this client.""" + return self._event_queue + + @property + def request_context(self) -> RequestContext: + """Get the request context associated with this client.""" + return self._request_context + + @property + def task_store(self) -> TaskStore: + """Get the task store associated with this client.""" + return self._task_store diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/a2a_turn_context.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/a2a_turn_context.py new file mode 100644 index 000000000..d4026678a --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/a2a_turn_context.py @@ -0,0 +1,112 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from __future__ import annotations + +from typing import cast + +from a2a.server.agent_execution import RequestContext +from a2a.server.events import EventQueue +from a2a.server.tasks import TaskStore + +from microsoft_agents.activity import Activity +from microsoft_agents.hosting.core import ( + AgentApplication, + TurnContext, + ChannelAdapter, + ClaimsIdentity, +) + +from ..activity import A2AActivity +from .a2a_client import A2AClient + + +class A2ATurnContext(TurnContext): + """A context object for handling A2A-specific turn functionality. + + Wraps a plain :class:`TurnContext` so that Teams-aware route handlers + receive a typed context without changing the core routing engine. + """ + + def __init__( + self, + adapter: ChannelAdapter, + app: AgentApplication, + activity: Activity, + identity: ClaimsIdentity, + *, + request_context: RequestContext, + event_queue: EventQueue, + task_store: TaskStore, + ) -> None: + """Initialize the A2A turn context with the given adapter, app, activity, and identity. + + :param adapter: The channel service adapter. + :param app: The agent application instance. + :param activity: The activity for the turn context. + :param identity: The claims identity for the turn context. + """ + super().__init__(adapter, activity, identity) + + self._original = self + self._app = app + self._activity.__class__ = A2AActivity + self._a2a_activity = cast(A2AActivity, self._activity) + + self._services.set(RequestContext, request_context) + self._services.set(EventQueue, event_queue) + self._services.set(TaskStore, task_store) + + self._client = A2AClient(self) + + @staticmethod + def from_existing( + context: TurnContext, + app: AgentApplication, + ) -> A2ATurnContext: + """Initialize the A2A turn context. + + :param context: The existing turn context. + :param app: The agent application instance. + :return: An instance of A2ATurnContext. + """ + + obj = A2ATurnContext.__new__(A2ATurnContext) + + super(A2ATurnContext, obj).__init__(context) + + obj._original = context + obj._turn_state = context.turn_state + obj._app = app + obj._activity.__class__ = A2AActivity + obj._a2a_activity = cast(A2AActivity, obj._activity) + obj._client = A2AClient(obj) + + return obj + + @property + def client(self) -> A2AClient: + """Get the A2A client associated with this turn context.""" + return self._client + + @property + def responded(self) -> bool: + """Check if the turn context has already sent a response.""" + return self._original._responded + + @responded.setter + def responded(self, value: bool): + """Set the responded status for the turn context.""" + self._original._responded = value + + @property + def streaming_response(self): + """Get the streaming response associated with the turn context.""" + if self._original is self: + return super().streaming_response + return self._original.streaming_response + + @property + def activity(self) -> A2AActivity: + """Get the A2A activity associated with the turn context.""" + return self._a2a_activity diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/route_handlers.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/route_handlers.py new file mode 100644 index 000000000..00783e3c3 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/extension/route_handlers.py @@ -0,0 +1,52 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from __future__ import annotations + +from typing import ( + Awaitable, + Protocol, + cast, +) + +from microsoft_agents.hosting.core import AgentApplication, TurnContext +from microsoft_agents.hosting.core.app._type_defs import RouteHandler, _StateContra + +from .a2a_turn_context import A2ATurnContext + + +class A2ARouteHandler(Protocol[_StateContra]): + """Protocol for a A2A route handler that receives a :class:`A2ATurnContext`.""" + + def __call__( + self, context: A2ATurnContext, state: _StateContra, / + ) -> Awaitable[None]: + """Handle a turn with A2A context. + + :param context: A2A-aware turn context. + :param state: The current turn state. + """ + ... + + +def wrap_a2a_route_handler( + handler: A2ARouteHandler[_StateContra], app: AgentApplication +) -> RouteHandler[_StateContra]: + """Adapt a :class:`A2ARouteHandler` into a plain :class:`RouteHandler`. + + Wraps *handler* so that the core routing engine (which passes a plain + :class:`TurnContext`) receives a compatible callable. + + :param handler: The A2A-specific handler to wrap. + :param app: The agent application handling the turn. + :return: A :class:`RouteHandler` that upgrades the context before delegating. + """ + + async def __func(context: TurnContext, state: _StateContra) -> None: + if not isinstance(context, A2ATurnContext): + a2a_context = A2ATurnContext.from_existing(context, app) + else: + a2a_context = cast(A2ATurnContext, context) + await handler(a2a_context, state) + + return __func diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/py.typed b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/py.typed new file mode 100644 index 000000000..e69de29bb diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/request_handling/__init__.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/request_handling/__init__.py new file mode 100644 index 000000000..94fd72583 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/request_handling/__init__.py @@ -0,0 +1,12 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from .a2a_agent_executor import A2AAgentExecutor +from .a2a_http_adapter import A2AHttpAdapter +from .a2a_request_handler import A2ARequestHandler + +__all__ = [ + "A2AAgentExecutor", + "A2AHttpAdapter", + "A2ARequestHandler", +] diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/request_handling/a2a_agent_executor.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/request_handling/a2a_agent_executor.py new file mode 100644 index 000000000..a0c24f5dc --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/request_handling/a2a_agent_executor.py @@ -0,0 +1,81 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import logging + +from a2a.server.agent_execution import AgentExecutor, RequestContext +from a2a.server.tasks import TaskUpdater +from a2a.server.events import EventQueue +from a2a.types import Task, TaskState, TaskStatus +from google.protobuf.timestamp_pb2 import Timestamp + +from .a2a_http_adapter import A2AHttpAdapter + +logger = logging.getLogger(__name__) + + +class A2AAgentExecutor(AgentExecutor): + """Executor for handling A2A agent requests.""" + + def __init__(self, adapter: A2AHttpAdapter): + """Initialize the A2A agent executor. + + :param adapter: The A2A adapter instance to use for executing agent turns. + """ + self._adapter = adapter + + async def execute(self, context: RequestContext, event_queue: EventQueue) -> None: + """Execute an A2A agent turn. + + :param context: The request context containing the message and other metadata. + :param event_queue: The event queue for handling events during the turn. + """ + + if not context.message: + logger.warning("No message found in the request context. Dropping request.") + return + + # If there is no current task, this is not a continuation + if context.current_task is None: + timestamp = Timestamp() + timestamp.GetCurrentTime() + await event_queue.enqueue_event( + Task( + id=context.task_id or "", + context_id=context.context_id or "", + status=TaskStatus( + state=TaskState.TASK_STATE_SUBMITTED, + timestamp=timestamp, + ), + history=[context.message], + ) + ) + + await self._adapter.execute_agent_turn( + context, + event_queue, + ) + + async def cancel( + self, + context: RequestContext, + event_queue: EventQueue, + ) -> None: + """Cancel an ongoing A2A agent turn. + + :param context: The request context containing the message and other metadata. + :param event_queue: The event queue for handling events during the turn. + """ + task_id = context.task_id + + updater = TaskUpdater( + event_queue=event_queue, + task_id=task_id or "", + context_id=context.context_id or "", + ) + await updater.cancel() + + await self._adapter.cancel_agent_turn( + context, + event_queue, + ) diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/request_handling/a2a_http_adapter.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/request_handling/a2a_http_adapter.py new file mode 100644 index 000000000..5e57054fd --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/request_handling/a2a_http_adapter.py @@ -0,0 +1,65 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from typing import Protocol + +from a2a.server.agent_execution import RequestContext +from a2a.server.events import EventQueue +from a2a.server.request_handlers import RequestHandler +from a2a.types import ( + AgentCard, + AgentInterface, + AgentSkill, +) + +from microsoft_agents.hosting.core import HttpRequestProtocol + + +class A2AHttpAdapter(Protocol): + """Protocol for an A2A HTTP adapter.""" + + @property + def agent_interfaces(self) -> list[AgentInterface]: + """Get the list of agent interfaces.""" + ... + + @property + def skills(self) -> list[AgentSkill]: + """Get the list of skills.""" + ... + + @property + def a2a_request_handler(self) -> RequestHandler: + """Get the A2A request handler.""" + ... + + async def execute_agent_turn( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + """Execute an agent turn given the request context and event queue. + + :param context: The request context for the agent turn. + :param event_queue: The event queue for the agent turn. + """ + ... + + async def cancel_agent_turn( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + """Cancel an agent turn given the request context and event queue. + + :param context: The request context for the agent turn. + :param event_queue: The event queue for the agent turn. + """ + ... + + async def get_agent_card( + self, request: HttpRequestProtocol, path_prefix: str + ) -> AgentCard: + """Process a request for the agent card. + + :param request: The HTTP request for the agent card. + :param path_prefix: The path prefix to be used in the agent card URL. + :return: The agent card for the agent. + """ + ... diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/request_handling/a2a_request_handler.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/request_handling/a2a_request_handler.py new file mode 100644 index 000000000..77a59bb71 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/request_handling/a2a_request_handler.py @@ -0,0 +1,58 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from a2a.types import AgentCard, SubscribeToTaskRequest + +from a2a.server.tasks import TaskStore +from a2a.server.request_handlers import DefaultRequestHandlerV2 +from a2a.utils.errors import ( + TaskNotFoundError, + UnsupportedOperationError, +) +from a2a.server.context import ServerCallContext +from a2a.server.agent_execution.active_task import TERMINAL_TASK_STATES + +from .a2a_http_adapter import A2AHttpAdapter + +from .a2a_agent_executor import A2AAgentExecutor + + +class A2ARequestHandler(DefaultRequestHandlerV2): + """Request handler for A2A interactions.""" + + def __init__( + self, + adapter: A2AHttpAdapter, + task_store: TaskStore, + agent_card: AgentCard, + ): + self._adapter = adapter + super().__init__(A2AAgentExecutor(self._adapter), task_store, agent_card) + + def update_agent_card(self, agent_card: AgentCard) -> None: + """Update the agent card with the provided AgentCard instance. + + :param agent_card: The AgentCard instance to update the handler with. + """ + self._agent_card = agent_card + + async def on_subscribe_to_task( + self, params: SubscribeToTaskRequest, context: ServerCallContext + ): + """Handle subscription to a task. + + :param params: The parameters for the subscription request. + :param context: The context of the request. + :raises TaskNotFoundError: If the task does not exist. + :raises UnsupportedOperationError: If the task is in a terminal state. + :yield: Events related to the task subscription. + """ + task = await self.task_store.get(params.id, context) + if task is None: + raise TaskNotFoundError() + + if task.status.state in TERMINAL_TASK_STATES: + raise UnsupportedOperationError("Cannot subscribe to a terminal task.") + + async for event in super().on_subscribe_to_task(params, context): + yield event diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/__init__.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/__init__.py new file mode 100644 index 000000000..79aad513e --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/__init__.py @@ -0,0 +1,18 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from .routes import ( + create_jsonrpc_routes, + create_rest_routes, + create_agent_card_routes, + use_jwt_middleware, +) +from .sdk_server_call_context_builder import SDKServerCallContextBuilder + +__all__ = [ + "create_jsonrpc_routes", + "create_rest_routes", + "create_agent_card_routes", + "use_jwt_middleware", + "SDKServerCallContextBuilder", +] diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/_constants.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/_constants.py new file mode 100644 index 000000000..14eba25f8 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/_constants.py @@ -0,0 +1,4 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +_CLAIMS_IDENTITY_KEY = "__CLAIMS_IDENTITY_KEY" diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/_utils.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/_utils.py new file mode 100644 index 000000000..807fe2699 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/_utils.py @@ -0,0 +1,24 @@ +from urllib.parse import urlsplit + +from a2a.types import AgentInterface + + +def _get_interface_route_path(url: str) -> str: + parsed = urlsplit(url) + + if parsed.query or parsed.fragment: + raise ValueError( + "Agent interface URLs cannot contain a query string or fragment." + ) + + if parsed.scheme or parsed.netloc: + if parsed.scheme not in ("http", "https") or not parsed.netloc: + raise ValueError(f"Invalid HTTP agent interface URL: {url}") + path = parsed.path + else: + path = parsed.path + + if not path.startswith("/"): + raise ValueError("Relative agent interface URLs must start with '/'.") + + return path.rstrip("/") or "/" diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/routes.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/routes.py new file mode 100644 index 000000000..8136184b5 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/routes.py @@ -0,0 +1,147 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import hashlib + +from datetime import datetime, timezone +from email.utils import format_datetime +from typing import Awaitable, Callable, Sequence + +from a2a.server.routes import ( + create_jsonrpc_routes as _create_jsonrpc_routes, + create_rest_routes as _create_rest_routes, +) +from a2a.server.request_handlers import RequestHandler +from a2a.server.request_handlers.response_helpers import agent_card_to_dict +from a2a.types import AgentCard +from a2a.utils.constants import AGENT_CARD_WELL_KNOWN_PATH + +from starlette.requests import Request +from starlette.responses import JSONResponse, Response +from starlette.routing import BaseRoute, Mount, Route + +from microsoft_agents.hosting.core import HttpRequestProtocol +from microsoft_agents.hosting.fastapi import ( + JwtAuthorizationMiddleware, + jwt_authorization_decorator, +) +from microsoft_agents.hosting.fastapi._fastapi_request_adapter import ( + FastApiRequestAdapter, +) + +from .sdk_server_call_context_builder import SDKServerCallContextBuilder +from ._utils import _get_interface_route_path + +_AGENT_CARD_CACHE_CONTROL_TEMPLATE = "public, max-age={}" + + +def create_jsonrpc_routes( + request_handler: RequestHandler, + rpc_url: str, + enable_v0_3_compat: bool = False, +) -> list[Route]: + """Create JSON-RPC routes with the given request handler and RPC URL. + + :param request_handler: The request handler to use for the JSON-RPC routes. + :param rpc_url: The URL for the JSON-RPC endpoint. + :param enable_v0_3_compat: Whether to enable compatibility with version 0.3. + :return: A list of Route objects representing the JSON-RPC routes. + """ + return _create_jsonrpc_routes( + request_handler, + _get_interface_route_path(rpc_url), + context_builder=SDKServerCallContextBuilder(), + enable_v0_3_compat=enable_v0_3_compat, + ) + + +def create_rest_routes( + request_handler: RequestHandler, + path_prefix: str = "", + enable_v0_3_compat: bool = False, +) -> list[BaseRoute]: + """Create REST routes with the given request handler and path prefix. + + :param request_handler: The request handler to use for the REST routes. + :param path_prefix: The prefix to prepend to all REST route paths. + :param enable_v0_3_compat: Whether to enable compatibility with version 0.3. + :return: A list of BaseRoute objects representing the REST routes. + """ + routes = _create_rest_routes( + request_handler, + context_builder=SDKServerCallContextBuilder(), + enable_v0_3_compat=enable_v0_3_compat, + path_prefix=_get_interface_route_path(path_prefix), + ) + return routes + + +def create_agent_card_routes( + get_agent_card: Callable[[HttpRequestProtocol, str], Awaitable[AgentCard]], + card_url: str = AGENT_CARD_WELL_KNOWN_PATH, + *, + _cache_enabled: bool = False, + cache_max_age: int = 3600, # seconds, so 1 hour +) -> list[BaseRoute]: + """Create routes for serving the agent card. + + :param get_agent_card: A callable that takes an HttpRequestProtocol and returns an AgentCard. + :param card_url: The URL path for the agent card endpoint. + :param _cache_enabled: Whether to enable public caching for the agent card. Default is False. + :param cache_max_age: The maximum age, in seconds, for caching the agent card. + :return: A list of Route objects representing the agent card routes. + """ + + if cache_max_age < 0: + raise ValueError("cache_max_age must be non-negative") + + prefix = card_url + try: + i = card_url.index(AGENT_CARD_WELL_KNOWN_PATH) + prefix = card_url[:i] + except ValueError: + # not found + pass + + last_modified = format_datetime(datetime.now(timezone.utc), usegmt=True) + + async def _get_agent_card(request: Request) -> Response: + """Retruns the public AgentCard describing this agent's capabilities, supported transports, and skills.""" + card = await get_agent_card(FastApiRequestAdapter(request), prefix) + response = JSONResponse(agent_card_to_dict(card)) + if _cache_enabled: + etag = f'"{hashlib.sha256(response.body).hexdigest()}"' + headers = { + "Cache-Control": _AGENT_CARD_CACHE_CONTROL_TEMPLATE.format( + cache_max_age + ), + "ETag": etag, + "Last-Modified": last_modified, + } + response.headers.update(headers) + else: + response.headers["Cache-Control"] = "no-store" + return response + + return [Route(path=card_url, endpoint=_get_agent_card, methods=["GET"])] + + +def use_jwt_middleware(routes: Sequence[BaseRoute]) -> None: + """Wrap all routes with JWT authorization middleware. + + :param routes: A list of BaseRoute objects to wrap with JWT authorization middleware. + """ + wrapped: set[int] = set() + + def wrap(route: BaseRoute) -> None: + """Wrap the given route with JWT authorization middleware if it hasn't been wrapped already.""" + if isinstance(route, Mount): + for child in route.routes: + wrap(child) + elif isinstance(route, Route) and id(route) not in wrapped: + route.endpoint = jwt_authorization_decorator(route.endpoint) + route.app = JwtAuthorizationMiddleware(route.app) + wrapped.add(id(route)) + + for route in routes: + wrap(route) diff --git a/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/sdk_server_call_context_builder.py b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/sdk_server_call_context_builder.py new file mode 100644 index 000000000..c99e11fe2 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/microsoft_agents/hosting/a2a/server/sdk_server_call_context_builder.py @@ -0,0 +1,43 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from a2a.server.context import ServerCallContext +from a2a.server.routes import DefaultServerCallContextBuilder +from a2a.extensions.common import ( + HTTP_EXTENSION_HEADER, + get_requested_extensions, +) + +from starlette.requests import Request + +from ._constants import _CLAIMS_IDENTITY_KEY + + +class SDKServerCallContextBuilder(DefaultServerCallContextBuilder): + """A default implementation of ServerCallContextBuilder.""" + + def build(self, request: Request) -> ServerCallContext: + """Builds a ServerCallContext from a Starlette Request. + + Args: + request: The incoming Starlette Request object. + + Returns: + A ServerCallContext instance populated with user and state + information from the request. + """ + state = {} + if "auth" in request.scope: + state["auth"] = request.auth + state["headers"] = dict(request.headers) + + if getattr(request.state, "claims_identity", None) is not None: + state[_CLAIMS_IDENTITY_KEY] = request.state.claims_identity + + return ServerCallContext( + user=self.build_user(request), + state=state, + requested_extensions=get_requested_extensions( + request.headers.getlist(HTTP_EXTENSION_HEADER) + ), + ) diff --git a/libraries/microsoft-agents-hosting-a2a/pyproject.toml b/libraries/microsoft-agents-hosting-a2a/pyproject.toml new file mode 100644 index 000000000..a271caa56 --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/pyproject.toml @@ -0,0 +1,28 @@ +[build-system] +requires = ["setuptools>=77"] +build-backend = "setuptools.build_meta" + +[project] +name = "microsoft-agents-hosting-a2a" +dynamic = ["version", "dependencies", "optional-dependencies"] +description = "Library adding A2A handling to an SDK Agent" +readme = {file = "readme.md", content-type = "text/markdown"} +authors = [{name = "Microsoft Corporation"}] +license = "MIT" +license-files = ["LICENSE"] +requires-python = ">=3.10" +classifiers = [ + "Programming Language :: Python :: 3", + "Programming Language :: Python :: 3.10", + "Programming Language :: Python :: 3.11", + "Programming Language :: Python :: 3.12", + "Programming Language :: Python :: 3.13", + "Programming Language :: Python :: 3.14", + "Operating System :: OS Independent", +] + +[tool.setuptools.package-data] +"microsoft_agents.hosting.a2a" = ["py.typed"] + +[project.urls] +"Homepage" = "https://github.com/microsoft/Agents" diff --git a/libraries/microsoft-agents-hosting-a2a/readme.md b/libraries/microsoft-agents-hosting-a2a/readme.md new file mode 100644 index 000000000..e69de29bb diff --git a/libraries/microsoft-agents-hosting-a2a/setup.py b/libraries/microsoft-agents-hosting-a2a/setup.py new file mode 100644 index 000000000..19d63719c --- /dev/null +++ b/libraries/microsoft-agents-hosting-a2a/setup.py @@ -0,0 +1,22 @@ +from os import environ, path +from setuptools import setup + +# Try to read from VERSION.txt file first, fall back to environment variable +version_file = path.join(path.dirname(__file__), "VERSION.txt") +if path.exists(version_file): + with open(version_file, "r", encoding="utf-8") as f: + package_version = f.read().strip() +else: + package_version = environ.get("PackageVersion", "0.0.0") + +setup( + version=package_version, + install_requires=[ + f"microsoft-agents-hosting-core=={package_version}", + f"microsoft-agents-hosting-fastapi=={package_version}", + "a2a-sdk>=1.0.0", + ], + extras_require={ + "blob": ["azure-storage-blob"], + }, +) diff --git a/libraries/microsoft-agents-hosting-aiohttp/microsoft_agents/hosting/aiohttp/_aiohttp_request_adapter.py b/libraries/microsoft-agents-hosting-aiohttp/microsoft_agents/hosting/aiohttp/_aiohttp_request_adapter.py index 97362cee4..4f6b45761 100644 --- a/libraries/microsoft-agents-hosting-aiohttp/microsoft_agents/hosting/aiohttp/_aiohttp_request_adapter.py +++ b/libraries/microsoft-agents-hosting-aiohttp/microsoft_agents/hosting/aiohttp/_aiohttp_request_adapter.py @@ -3,8 +3,10 @@ from aiohttp.web import Request +from microsoft_agents.hosting.core import HttpRequestProtocol -class AiohttpRequestAdapter: + +class AiohttpRequestAdapter(HttpRequestProtocol): """Adapter to make aiohttp Request compatible with HttpRequestProtocol.""" def __init__(self, request: Request): @@ -18,6 +20,10 @@ def method(self) -> str: def headers(self): return self._request.headers + @property + def url(self) -> str: + return str(self._request.url) + async def json(self): return await self._request.json() diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/_utils/_service_set.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/_utils/_service_set.py index 283231dfa..5720d6427 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/_utils/_service_set.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/_utils/_service_set.py @@ -3,7 +3,7 @@ from __future__ import annotations -from typing import TypeVar, cast, Any +from typing import TypeVar, cast, Any, overload, Literal T = TypeVar("T") @@ -23,11 +23,19 @@ def __init__(self, service_set: _ServiceSet | None = None) -> None: if service_set is not None: self._state.update(service_set._state) - def get(self, key: type[T]) -> T | None: + @overload + def get(self, key: type[T], raise_if_missing: Literal[True]) -> T: ... + + @overload + def get( + self, key: type[T], raise_if_missing: Literal[False] = False + ) -> T | None: ... + def get(self, key: type[T], raise_if_missing: bool = False) -> T | None: """ Gets a value from the state collection. - :param key: - :return: + :param key: The type of the value to retrieve. + :param raise_if_missing: Whether to raise an exception if the value is missing. + :return: The value associated with the specified type, or None if not found and raise_if_missing is False. """ val = self._state.get(key) if val is not None: @@ -39,6 +47,10 @@ def get(self, key: type[T]) -> T | None: f"Value for key '{key.__name__}' is not of type {key.__name__} (got {type(val).__name__})" ) return cast(T, val) + + if raise_if_missing: + raise KeyError(f"Value for key '{key.__name__}' is missing") + return None def has(self, key: type) -> bool: diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/channel_adapter.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/channel_adapter.py index 6499eefbe..77e0a5f33 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/channel_adapter.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/channel_adapter.py @@ -9,7 +9,6 @@ from collections.abc import Callable from typing import Awaitable from microsoft_agents.hosting.core.authorization import ClaimsIdentity -from microsoft_agents.activity import ChannelAdapterProtocol from microsoft_agents.activity import ( Activity, ChannelId, @@ -21,6 +20,7 @@ from .turn_context import TurnContext from .middleware_set import MiddlewareSet, Middleware +from .channel_adapter_protocol import ChannelAdapterProtocol class ChannelAdapter(ABC, ChannelAdapterProtocol): diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/channel_adapter_protocol.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/channel_adapter_protocol.py new file mode 100644 index 000000000..c17cf83b0 --- /dev/null +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/channel_adapter_protocol.py @@ -0,0 +1,79 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from abc import abstractmethod +from typing import Protocol, Callable, Awaitable + +from typing_extensions import Self + +from .turn_context import TurnContext +from microsoft_agents.activity import ( + Activity, + ResourceResponse, + ConversationReference, + ConversationParameters, +) + + +class ChannelAdapterProtocol(Protocol): + on_turn_error: Callable[[TurnContext, Exception], Awaitable] | None + + @abstractmethod + async def send_activities( + self, context: TurnContext, activities: list[Activity] + ) -> list[ResourceResponse]: + pass + + @abstractmethod + async def update_activity(self, context: TurnContext, activity: Activity) -> None: + pass + + @abstractmethod + async def delete_activity( + self, context: TurnContext, reference: ConversationReference + ) -> None: + pass + + @abstractmethod + def use(self, middleware: object) -> Self: + pass + + @abstractmethod + async def continue_conversation( + self, + agent_id: str, + reference: ConversationReference, + callback: Callable[[TurnContext], Awaitable], + ) -> None: + pass + + # TODO: potentially move ClaimsIdentity to activity + @abstractmethod + async def continue_conversation_with_claims( + self, + claims_identity: dict, + continuation_activity: Activity, + callback: Callable[[TurnContext], Awaitable], + audience: str | None = None, + ): + pass + + @abstractmethod + async def create_conversation( + self, + agent_app_id: str, + channel_id: str, + service_url: str, + audience: str, + conversation_parameters: ConversationParameters, + callback: Callable[[TurnContext], Awaitable], + ) -> None: + pass + + @abstractmethod + async def run_pipeline( + self, + context: TurnContext, + callback: Callable[[TurnContext], Awaitable], + ) -> None: + pass diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/http/_http_request_protocol.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/http/_http_request_protocol.py index 0170c59f7..3e2513619 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/http/_http_request_protocol.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/http/_http_request_protocol.py @@ -23,6 +23,11 @@ def headers(self) -> Mapping[str, str]: """Request headers.""" ... + @property + def url(self) -> str: + """Request URL as a string.""" + ... + async def json(self) -> dict[str, Any]: """Parse request body as JSON.""" ... diff --git a/libraries/microsoft-agents-hosting-fastapi/microsoft_agents/hosting/fastapi/_fastapi_request_adapter.py b/libraries/microsoft-agents-hosting-fastapi/microsoft_agents/hosting/fastapi/_fastapi_request_adapter.py index d733a4502..ef0c1805d 100644 --- a/libraries/microsoft-agents-hosting-fastapi/microsoft_agents/hosting/fastapi/_fastapi_request_adapter.py +++ b/libraries/microsoft-agents-hosting-fastapi/microsoft_agents/hosting/fastapi/_fastapi_request_adapter.py @@ -3,8 +3,10 @@ from fastapi import Request +from microsoft_agents.hosting.core import HttpRequestProtocol -class FastApiRequestAdapter: + +class FastApiRequestAdapter(HttpRequestProtocol): """Adapter to make FastAPI Request compatible with HttpRequestProtocol.""" def __init__(self, request: Request): @@ -18,6 +20,10 @@ def method(self) -> str: def headers(self): return self._request.headers + @property + def url(self) -> str: + return str(self._request.url) + async def json(self): return await self._request.json() diff --git a/scripts/dev_setup.ps1 b/scripts/dev_setup.ps1 index 2b3f8d35d..109db212d 100644 --- a/scripts/dev_setup.ps1 +++ b/scripts/dev_setup.ps1 @@ -6,6 +6,7 @@ pip install -e ./libraries/microsoft-agents-activity/ --config-settings editable pip install -e ./libraries/microsoft-agents-authentication-msal/ --config-settings editable_mode=compat pip install -e ./libraries/microsoft-agents-authentication-entra-auth-sidecar/ --config-settings editable_mode=compat pip install -e ./libraries/microsoft-agents-copilotstudio-client/ --config-settings editable_mode=compat +pip install -e ./libraries/microsoft-agents-hosting-a2a/ --config-settings editable_mode=compat pip install -e ./libraries/microsoft-agents-hosting-aiohttp/ --config-settings editable_mode=compat pip install -e ./libraries/microsoft-agents-hosting-core/ --config-settings editable_mode=compat pip install -e ./libraries/microsoft-agents-hosting-teams/ --config-settings editable_mode=compat diff --git a/scripts/dev_setup.sh b/scripts/dev_setup.sh index b013fced4..ac6c42182 100644 --- a/scripts/dev_setup.sh +++ b/scripts/dev_setup.sh @@ -6,6 +6,7 @@ pip install -e ./libraries/microsoft-agents-activity/ --config-settings editable pip install -e ./libraries/microsoft-agents-authentication-msal/ --config-settings editable_mode=compat pip install -e ./libraries/microsoft-agents-authentication-entra-auth-sidecar/ --config-settings editable_mode=compat pip install -e ./libraries/microsoft-agents-copilotstudio-client/ --config-settings editable_mode=compat +pip install -e ./libraries/microsoft-agents-hosting-a2a/ --config-settings editable_mode=compat pip install -e ./libraries/microsoft-agents-hosting-aiohttp/ --config-settings editable_mode=compat pip install -e ./libraries/microsoft-agents-hosting-core/ --config-settings editable_mode=compat pip install -e ./libraries/microsoft-agents-hosting-teams/ --config-settings editable_mode=compat diff --git a/test_samples/a2a/README.md b/test_samples/a2a/README.md new file mode 100644 index 000000000..6b0c58cbe --- /dev/null +++ b/test_samples/a2a/README.md @@ -0,0 +1,45 @@ +# A2A echo sample + +This sample defines a custom Microsoft 365 Agents SDK agent, adds A2A +capabilities, and hosts it with FastAPI. + +## Install + +From this directory: + +```powershell +pip install -r requirements.txt +``` + +When developing from this repository, install the SDK libraries in editable +mode using `scripts\dev_setup.ps1`. + +## Run the agent + +From the repository root: + +```powershell +python -m test_samples.a2a.agent.main +``` + +The server listens on `http://127.0.0.1:41241` by default. Its endpoints are: + +- Agent card: `http://127.0.0.1:41241/a2a/.well-known/agent-card.json` +- JSON-RPC: `http://127.0.0.1:41241/a2a` +- Health check: `http://127.0.0.1:41241/health` + +Set `HOST` or `PORT` to override the defaults. + +## Run the client + +In another terminal: + +```powershell +python -m test_samples.a2a.client.cli --url http://127.0.0.1:41241/a2a +``` + +Enter a message to receive an echo response. Type `/quit` to exit. + +JWT authorization is intentionally disabled for this local sample. Production +applications should configure an `AgentAuthConfiguration` and enable the A2A +JWT middleware. \ No newline at end of file diff --git a/test_samples/a2a/agent/__init__.py b/test_samples/a2a/agent/__init__.py new file mode 100644 index 000000000..9989a770a --- /dev/null +++ b/test_samples/a2a/agent/__init__.py @@ -0,0 +1,13 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from .a2a import A2A_ADAPTER +from .agent import AGENT, EchoAgent +from .start_server import create_app + +__all__ = [ + "A2A_ADAPTER", + "AGENT", + "EchoAgent", + "create_app", +] \ No newline at end of file diff --git a/test_samples/a2a/agent/a2a.py b/test_samples/a2a/agent/a2a.py new file mode 100644 index 000000000..5b8c84b09 --- /dev/null +++ b/test_samples/a2a/agent/a2a.py @@ -0,0 +1,34 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from a2a.types import AgentInterface, AgentSkill +from a2a.utils.constants import TransportProtocol + +from microsoft_agents.hosting.a2a import A2AAdapter + +from .agent import AGENT + + +A2A_ADAPTER = A2AAdapter( + AGENT, + agent_card_name="SDK Echo Agent", + agent_card_description="A simple Microsoft 365 Agents SDK A2A sample.", + agent_card_version="1.0.0", + agent_interfaces=[ + AgentInterface( + url="/a2a", + protocol_binding=TransportProtocol.JSONRPC, + ) + ], + skills=[ + AgentSkill( + id="echo", + name="Echo", + description="Echoes a text message back to the caller.", + tags=["sample", "echo"], + examples=["Hello", "Repeat this message"], + input_modes=["text/plain"], + output_modes=["text/plain"], + ) + ], +) \ No newline at end of file diff --git a/test_samples/a2a/agent/agent.py b/test_samples/a2a/agent/agent.py new file mode 100644 index 000000000..8a18bfba5 --- /dev/null +++ b/test_samples/a2a/agent/agent.py @@ -0,0 +1,38 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from microsoft_agents.activity import ( + Activity, + ActivityTypes, + EndOfConversationCodes, + TurnContextProtocol, +) +from microsoft_agents.hosting.core import ( + AgentApplication, +) + +app = AgentApplication() + +class EchoAgent(ActivityHandler): + """A small SDK agent that echoes each A2A message.""" + + async def on_message_activity(self, turn_context: TurnContextProtocol): + text = (turn_context.activity.text or "").strip() + + if text.lower() == "/help": + response = "Send me a message and I will echo it back." + elif text: + response = f"Echo: {text}" + else: + response = "I received an empty message." + + await turn_context.send_activity(response) + await turn_context.send_activity( + Activity( + type=ActivityTypes.end_of_conversation, + code=EndOfConversationCodes.completed_successfully, + ) + ) + + +AGENT = EchoAgent() \ No newline at end of file diff --git a/test_samples/a2a/agent/main.py b/test_samples/a2a/agent/main.py new file mode 100644 index 000000000..fd6c91a7b --- /dev/null +++ b/test_samples/a2a/agent/main.py @@ -0,0 +1,10 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from .start_server import create_app, start_server + +app = create_app() + + +if __name__ == "__main__": + start_server(app) diff --git a/test_samples/a2a/agent/setup.py b/test_samples/a2a/agent/setup.py new file mode 100644 index 000000000..27d5b093a --- /dev/null +++ b/test_samples/a2a/agent/setup.py @@ -0,0 +1,22 @@ +import os + +from dotenv import load_dotenv + +from microsoft_agents.activity import ( + load_configuration_from_env, +) +from microsoft_agents.hosting.core import ( + Authorization, + MemoryStorage, +) +from microsoft_agents.authentication.msal import MsalConnectionManager + +load_dotenv() + +CONFIG = load_configuration_from_env(os.environ) + +STORAGE = MemoryStorage() +CONNECTIONS = MsalConnectionManager(**CONFIG) +ADAPTER = A2AAdapter() +AUTHORIZATION = Authorization(STORAGE, CONNECTIONS, **CONFIG) + diff --git a/test_samples/a2a/agent/start_server.py b/test_samples/a2a/agent/start_server.py new file mode 100644 index 000000000..78bba680c --- /dev/null +++ b/test_samples/a2a/agent/start_server.py @@ -0,0 +1,42 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from os import environ + +import uvicorn +from fastapi import FastAPI + +from microsoft_agents.hosting.a2a import add_a2a + +from .a2a import A2A_ADAPTER +from .agent import AGENT + + +def create_app() -> FastAPI: + """Create the FastAPI application and register the A2A routes.""" + + app = FastAPI(title="SDK A2A Echo Agent", version="1.0.0") + + # Authentication is disabled only to keep this local sample simple. + add_a2a( + app, + AGENT, + adapter=A2A_ADAPTER, + use_jwt_middleware=False, + ) + + @app.get("/health") + async def health(): + return {"status": "ok"} + + return app + + +def start_server(app: FastAPI | None = None) -> None: + """Start the local A2A server.""" + + uvicorn.run( + app or create_app(), + host=environ.get("HOST", "127.0.0.1"), + port=int(environ.get("PORT", "41241")), + ) \ No newline at end of file diff --git a/test_samples/a2a/client/__init__.py b/test_samples/a2a/client/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/test_samples/a2a/client/cli.py b/test_samples/a2a/client/cli.py new file mode 100644 index 000000000..30c3468bf --- /dev/null +++ b/test_samples/a2a/client/cli.py @@ -0,0 +1,143 @@ +# from https://github.com/a2aproject/a2a-python/blob/main/samples/cli.py + +import argparse +import asyncio +import os +import signal +import uuid + +from typing import Any + +import grpc +import httpx + +from a2a.client import A2ACardResolver, ClientConfig, create_client +from a2a.helpers import get_artifact_text, get_message_text +from a2a.helpers.agent_card import display_agent_card +from a2a.types import Message, Part, Role, SendMessageRequest, TaskState + + +async def _handle_stream( + stream: Any, current_task_id: str | None +) -> str | None: + async for event in stream: + if event.HasField('message'): + print('Message:', get_message_text(event.message, delimiter=' ')) + return None + + if not current_task_id: + if event.HasField('task'): + current_task_id = event.task.id + print('--- Task Started ---') + print(f'Task [state={TaskState.Name(event.task.status.state)}]') + elif event.HasField('status_update'): + current_task_id = event.status_update.task_id + print('--- Task Started ---') + elif event.HasField('artifact_update'): + current_task_id = event.artifact_update.task_id + print('--- Task Started ---') + else: + raise ValueError(f'Unexpected first event: {event}') + + if event.HasField('status_update'): + state_name = TaskState.Name(event.status_update.status.state) + message_text = ( + ': ' + + get_message_text( + event.status_update.status.message, delimiter=' ' + ) + if event.status_update.status.HasField('message') + else '' + ) + print(f'TaskStatusUpdate [state={state_name}]{message_text}') + if state_name in ( + 'TASK_STATE_COMPLETED', + 'TASK_STATE_FAILED', + 'TASK_STATE_CANCELED', + 'TASK_STATE_REJECTED', + ): + current_task_id = None + print('--- Task Finished ---') + elif event.HasField('artifact_update'): + print( + f'TaskArtifactUpdate [name={event.artifact_update.artifact.name}]:', + get_artifact_text( + event.artifact_update.artifact, delimiter=' ' + ), + ) + return current_task_id + + +async def main() -> None: + """Run the A2A terminal client.""" + parser = argparse.ArgumentParser(description='A2A Terminal Client') + parser.add_argument( + '--url', default='http://127.0.0.1:41241/a2a', help='Agent base URL' + ) + parser.add_argument( + '--transport', + default=None, + help='Preferred transport (JSONRPC, HTTP+JSON, GRPC)', + ) + args = parser.parse_args() + + config = ClientConfig( + grpc_channel_factory=grpc.aio.insecure_channel, + ) + if args.transport: + config.supported_protocol_bindings = [args.transport] + + print( + f'Connecting to {args.url} (preferred transport: {args.transport or "Any"})' + ) + + async with httpx.AsyncClient() as httpx_client: + resolver = A2ACardResolver(httpx_client, args.url) + card = await resolver.get_agent_card() + print('\n✓ Agent Card Found:') + display_agent_card(card) + + client = await create_client(card, client_config=config) + + actual_transport = getattr(client, '_transport', client) + print(f' Picked Transport: {actual_transport.__class__.__name__}') + + print('\nConnected! Send a message or type /quit to exit.') + + current_task_id = None + current_context_id = str(uuid.uuid4()) + + while True: + try: + loop = asyncio.get_running_loop() + user_input = await loop.run_in_executor(None, input, 'You: ') + except KeyboardInterrupt: + break + + if user_input.lower() in ('/quit', '/exit'): + break + if not user_input.strip(): + continue + + message = Message( + role=Role.ROLE_USER, + message_id=str(uuid.uuid4()), + parts=[Part(text=user_input)], + task_id=current_task_id, + context_id=current_context_id, + ) + + request = SendMessageRequest(message=message) + + try: + stream = client.send_message(request) + current_task_id = await _handle_stream(stream, current_task_id) + except (httpx.RequestError, grpc.RpcError) as e: + print(f'Error communicating with agent: {e}') + + await client.close() + + +if __name__ == '__main__': + signal.signal(signal.SIGINT, lambda sig, frame: os._exit(0)) + asyncio.run(main()) \ No newline at end of file diff --git a/test_samples/a2a/requirements.txt b/test_samples/a2a/requirements.txt new file mode 100644 index 000000000..b0cbd2946 --- /dev/null +++ b/test_samples/a2a/requirements.txt @@ -0,0 +1,3 @@ +microsoft-agents-hosting-a2a +a2a-sdk[fastapi] +uvicorn \ No newline at end of file diff --git a/tests/activity/test_activity.py b/tests/activity/test_activity.py index c1d37e0f3..34b6263e6 100644 --- a/tests/activity/test_activity.py +++ b/tests/activity/test_activity.py @@ -19,6 +19,7 @@ Thing, ProductInfo, RoleTypes, + StreamInfo, ) from tests.activity._common.my_channel_data import MyChannelData @@ -496,6 +497,46 @@ def test_get_product_info_entity_single(self, entities, expected): retrieved_product_info = activity.get_product_info_entity() assert retrieved_product_info == expected + @pytest.mark.parametrize( + "entities, expected", + [ + [ + [ + { + "type": "streaminfo", + "streamId": "stream_123", + "streamSequence": 1, + }, + {"type": "other"}, + ], + StreamInfo(stream_id="stream_123", stream_sequence=1), + ], + [ + [ + {"type": "other"}, + {"type": "mention", "text": "Another mention"}, + ], + None, + ], + [ + [ + {"type": "StreamInfo", "streamId": "stream_123"}, + {"type": "StreamInfo", "streamId": "stream_456"}, + ], + StreamInfo(stream_id="stream_123"), + ], + [[], None], + [None, None], + ], + ) + def test_get_streaming_entity(self, entities, expected): + activity = Activity(type="message", entities=entities) + for entity in activity.entities or []: + if entity.type.casefold() == EntityTypes.STREAM_INFO.value.casefold(): + assert isinstance(entity, StreamInfo) + + assert activity.get_streaming_entity() == expected + class TestActivityAsTypeHelpers: @pytest.mark.parametrize( diff --git a/tests/hosting_a2a/__init__.py b/tests/hosting_a2a/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/hosting_a2a/activity/__init__.py b/tests/hosting_a2a/activity/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/hosting_a2a/activity/test_a2a_activity.py b/tests/hosting_a2a/activity/test_a2a_activity.py new file mode 100644 index 000000000..3b6015696 --- /dev/null +++ b/tests/hosting_a2a/activity/test_a2a_activity.py @@ -0,0 +1,192 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import pytest +from a2a.types import Message, Part, Role, TaskState +from google.protobuf.json_format import ParseDict +from google.protobuf.json_format import MessageToDict +from google.protobuf.struct_pb2 import Value + +from microsoft_agents.activity import ( + ActivityTypes, + Channels, + DeliveryModes, + RoleTypes, +) +from microsoft_agents.hosting.a2a.activity import A2AActivity + + +def _message(*parts: Part, task_id: str = "task-1") -> Message: + return Message( + message_id="message-1", + context_id="context-1", + task_id=task_id, + role=Role.ROLE_USER, + parts=list(parts), + ) + + +def test_from_message_requires_task_id(): + message = _message(Part(text="hello"), task_id="") + + with pytest.raises(ValueError, match="Task ID is required"): + A2AActivity.from_message("request-1", None, message) + + +def test_from_message_maps_message_properties(): + message = _message(Part(text="hello")) + + activity = A2AActivity.from_message("request-1", None, message) + + assert activity.type == ActivityTypes.message + assert activity.channel_id == Channels.a2a + assert activity.request_id == "request-1" + assert activity.text == "hello" + assert activity.conversation.id == "task-1" + assert activity.delivery_mode == DeliveryModes.stream + assert activity.from_property.role == RoleTypes.user + assert activity.recipient.role == RoleTypes.agent + assert activity.channel_data is message + + +def test_from_message_uses_explicit_task_id_and_populates_missing_context_id(): + message = _message(Part(text="hello"), task_id="") + message.context_id = "" + + activity = A2AActivity.from_message("request-1", "explicit-task", message) + + assert activity.conversation.id == "explicit-task" + assert message.task_id == "explicit-task" + assert message.context_id + + +def test_from_message_maps_all_supported_part_types(): + message = _message( + Part(text="hello "), + Part(text="world"), + Part( + url="https://example.com/image.png", + media_type="image/png", + filename="image.png", + ), + Part( + raw=b"file contents", + media_type="application/octet-stream", + filename="data.bin", + ), + Part( + data=ParseDict({"answer": 42}, Value()), + media_type="application/json", + filename="result.json", + ), + ) + + activity = A2AActivity.from_message("request-1", None, message) + + assert activity.text == "hello world" + assert activity.delivery_mode == DeliveryModes.stream + assert len(activity.attachments) == 3 + + url_attachment, raw_attachment, data_attachment = activity.attachments + assert url_attachment.content_url == "https://example.com/image.png" + assert url_attachment.content_type == "image/png" + assert url_attachment.name == "image.png" + assert raw_attachment.content == b"file contents" + assert raw_attachment.content_type == "application/octet-stream" + assert raw_attachment.name == "data.bin" + assert data_attachment.content == {"answer": 42.0} + assert data_attachment.content_type == "application/json" + assert data_attachment.name == "result.json" + + +def test_from_message_accepts_parts_without_optional_filenames(): + message = _message( + Part( + url="https://example.com/image.png", + media_type="image/png", + ), + Part( + raw=b"file contents", + media_type="application/octet-stream", + ), + Part( + data=ParseDict({"answer": 42}, Value()), + media_type="application/json", + ), + ) + + activity = A2AActivity.from_message("request-1", None, message) + + assert [attachment.name for attachment in activity.attachments] == ["", "", ""] + + +def test_supported_parts_round_trip_through_activity_artifact(): + message = _message( + Part( + url="https://example.com/image.png", + media_type="image/png", + filename="image.png", + ), + Part( + raw=b"file contents", + media_type="application/octet-stream", + filename="data.bin", + ), + Part( + data=ParseDict({"answer": 42}, Value()), + media_type="application/json", + filename="result.json", + ), + ) + + activity = A2AActivity.from_message("request-1", None, message) + artifact = activity.to_artifact("artifact-1") + + assert artifact.parts[0].url == "https://example.com/image.png" + assert artifact.parts[0].media_type == "image/png" + assert artifact.parts[1].raw == b"file contents" + assert artifact.parts[1].media_type == "application/octet-stream" + assert MessageToDict(artifact.parts[2].data) == {"answer": 42.0} + assert artifact.parts[2].media_type == "application/json" + assert artifact.parts[2].filename == "result.json" + + +def test_empty_text_and_raw_parts_preserve_their_content_kind(): + message = _message( + Part(text=""), + Part( + raw=b"", + media_type="application/octet-stream", + filename="empty.bin", + ), + ) + + activity = A2AActivity.from_message("request-1", None, message) + artifact = activity.to_artifact("artifact-1") + + assert activity.text == "" + assert activity.attachments[0].content == b"" + assert artifact.parts[0].WhichOneof("content") == "text" + assert artifact.parts[0].text == "" + assert artifact.parts[1].WhichOneof("content") == "raw" + assert artifact.parts[1].raw == b"" + + +def test_to_message_and_artifact_delegate_activity_content(): + activity = A2AActivity(type=ActivityTypes.message, text="result") + + message = activity.to_message("context-1", "task-1") + artifact = activity.to_artifact("artifact-1") + + assert message.context_id == "context-1" + assert message.task_id == "task-1" + assert message.role == Role.ROLE_AGENT + assert [part.text for part in message.parts] == ["result"] + assert artifact.artifact_id == "artifact-1" + assert [part.text for part in artifact.parts] == ["result"] + + +def test_get_task_state_delegates_to_activity_state(): + activity = A2AActivity(type=ActivityTypes.message) + + assert activity.get_task_state() == TaskState.TASK_STATE_WORKING diff --git a/tests/hosting_a2a/activity/test_utils.py b/tests/hosting_a2a/activity/test_utils.py new file mode 100644 index 000000000..330c1ab6e --- /dev/null +++ b/tests/hosting_a2a/activity/test_utils.py @@ -0,0 +1,179 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from types import SimpleNamespace + +import pytest +from a2a.types import Message, Role, TaskState +from google.protobuf.json_format import MessageToDict + +from microsoft_agents.activity import ( + Activity, + Attachment, + Entity, + InputHints, + StreamInfo, +) +from microsoft_agents.hosting.a2a.activity import utils + + +def test_activity_to_artifact_maps_text_value_and_attachments(): + activity = Activity( + type="message", + text="hello", + value={"result": "ok"}, + attachments=[ + Attachment( + content_type="image/png", + content_url="https://example.com/image.png", + name="image.png", + ), + Attachment( + content_type="application/json", + content={"count": 2}, + name="result.json", + ), + ], + ) + + artifact = utils.activity_to_artifact(activity, artifact_id="artifact-1") + + assert artifact.artifact_id == "artifact-1" + assert artifact.parts[0].text == "hello" + assert MessageToDict(artifact.parts[1].data)["result"] == "ok" + assert artifact.parts[2].url == "https://example.com/image.png" + assert artifact.parts[2].media_type == "image/png" + assert artifact.parts[2].filename == "image.png" + assert MessageToDict(artifact.parts[3].data)["count"] == 2 + assert artifact.parts[3].media_type == "application/json" + assert artifact.parts[3].filename == "result.json" + + +def test_activity_to_artifact_adds_entities_and_excludes_stream_info(): + activity = Activity( + type="message", + entities=[ + Entity(type="citation"), + StreamInfo(stream_id="stream-1", stream_sequence=1), + ], + ) + + artifact = utils.activity_to_artifact(activity) + + assert len(artifact.parts) == 1 + assert MessageToDict(artifact.parts[0].data)["type"] == "citation" + assert artifact.parts[0].metadata["mimeType"] == "application/json" + assert artifact.parts[0].metadata["type"] == "object" + assert artifact.parts[0].metadata["schema"]["title"] == "Entity" + assert artifact.parts[0].metadata["schema"]["required"] == ["type"] + + +def test_activity_to_artifact_can_exclude_entities(): + activity = Activity(type="message", entities=[Entity(type="citation")]) + + artifact = utils.activity_to_artifact(activity, include_entities=False) + + assert list(artifact.parts) == [] + + +def test_activity_to_artifact_skips_unsupported_attachment_content(): + activity = Activity( + type="message", + attachments=[ + Attachment( + content_type="text/plain", + content="plain text", + ) + ], + ) + + artifact = utils.activity_to_artifact(activity) + + assert list(artifact.parts) == [] + + +def test_create_artifact_from_data_populates_artifact(): + artifact = utils.create_artifact_from_data( + {"answer": 42}, + name="result", + description="The result", + artifact_id="artifact-1", + ) + + assert artifact.artifact_id == "artifact-1" + assert artifact.name == "result" + assert artifact.description == "The result" + assert MessageToDict(artifact.parts[0].data)["answer"] == 42 + assert artifact.parts[0].metadata["mimeType"] == "application/json" + assert artifact.parts[0].metadata["type"] == "object" + assert artifact.parts[0].metadata["schema"] == { + "additionalProperties": True, + "type": "object", + } + + +def test_create_artifact_from_data_generates_artifact_id(): + artifact = utils.create_artifact_from_data({"answer": 42}) + + assert artifact.artifact_id + + +def test_create_artifact_from_data_rejects_none(): + with pytest.raises(ValueError, match="Data cannot be None"): + utils.create_artifact_from_data(None) + + +def test_create_message_without_activity_has_agent_role_and_no_parts(): + message = utils.create_message("context-1", "task-1", None) + + assert message.context_id == "context-1" + assert message.task_id == "task-1" + assert message.message_id + assert message.role == Role.ROLE_AGENT + assert list(message.parts) == [] + + +@pytest.mark.parametrize( + ("activity", "expected"), + [ + (Activity(type="message"), False), + (Activity(type="message", text="hello"), True), + ( + Activity( + type="message", + attachments=[Attachment(content_type="text/plain", content="hello")], + ), + True, + ), + ], +) +def test_has_message_content(activity, expected): + assert utils.has_message_content(activity) is expected + + +def test_get_task_state_maps_input_hint(): + assert ( + utils.get_task_state( + Activity(type="message", input_hint=InputHints.expecting_input) + ) + == TaskState.TASK_STATE_INPUT_REQUIRED + ) + assert ( + utils.get_task_state(Activity(type="message")) == TaskState.TASK_STATE_WORKING + ) + + +def test_get_incoming_message_returns_a2a_message(): + message = Message(message_id="message-1", role=Role.ROLE_USER) + context = SimpleNamespace(activity=Activity(type="message", channel_data=message)) + + assert utils.get_incoming_message(context) is message + + +def test_get_incoming_message_rejects_other_channel_data(): + context = SimpleNamespace( + activity=Activity(type="message", channel_data={"message": "hello"}) + ) + + with pytest.raises(TypeError, match="Expected channel_data"): + utils.get_incoming_message(context) diff --git a/tests/hosting_a2a/integration/__init__.py b/tests/hosting_a2a/integration/__init__.py new file mode 100644 index 000000000..5b7f7a925 --- /dev/null +++ b/tests/hosting_a2a/integration/__init__.py @@ -0,0 +1,2 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. diff --git a/tests/hosting_a2a/integration/conftest.py b/tests/hosting_a2a/integration/conftest.py new file mode 100644 index 000000000..f05723eb2 --- /dev/null +++ b/tests/hosting_a2a/integration/conftest.py @@ -0,0 +1,179 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from collections.abc import AsyncIterator +import asyncio +from dataclasses import dataclass + +import httpx +import pytest_asyncio +from a2a.types import AgentInterface, AgentSkill +from a2a.utils.constants import TransportProtocol +from fastapi import FastAPI + +from microsoft_agents.hosting.a2a import A2AAdapter, add_a2a +from microsoft_agents.activity import Activity, ActivityTypes, EndOfConversationCodes +from microsoft_agents.hosting.core import ( + AgentApplication, + ApplicationOptions, + MemoryStorage, + TurnContext, + TurnState, +) +from tests._common.testing_objects import TestingConnectionManager + +STREAMING_TRIGGER_TEXT = "stream-me" +"""Message text that causes the test agent to emit chunked streaming updates.""" + +INPUT_REQUIRED_TRIGGER_TEXT = "need-more-input" +"""Message text that leaves the task waiting for a continuation.""" + +STRUCTURED_RESULT_TRIGGER_TEXT = "structured-result" +"""Message text that completes with both text and structured result data.""" + +LONG_RUNNING_TRIGGER_TEXT = "long-running" +"""Message text that leaves a task running until it is canceled.""" + +FAILURE_TRIGGER_TEXT = "fail" +"""Message text that raises from the SDK agent handler.""" + +PARTS_TRIGGER_TEXT = "echo-parts" +"""Message text that echoes incoming attachments through an outgoing message.""" + + +@dataclass +class A2ATestHarness: + client: httpx.AsyncClient + long_running_started: asyncio.Event + release_long_running: asyncio.Event + + +def _create_agent_application( + long_running_started: asyncio.Event, + release_long_running: asyncio.Event, +) -> AgentApplication[TurnState]: + application = AgentApplication[TurnState]( + options=ApplicationOptions(storage=MemoryStorage()), + connection_manager=TestingConnectionManager(), + ) + + @application.activity("message") + async def on_message(context: TurnContext, _state: TurnState) -> None: + text = context.activity.text or "" + if text == STREAMING_TRIGGER_TEXT: + # Exercise the chunked streaming API so that informative and + # partial-content updates surface as A2A task status/artifact + # update events. + streaming_response = context.streaming_response + streaming_response.queue_informative_update("Thinking...") + streaming_response.queue_text_chunk("Echo: ") + streaming_response.queue_text_chunk(text) + await streaming_response.end_stream() + elif text == INPUT_REQUIRED_TRIGGER_TEXT: + await context.send_activity( + Activity( + type=ActivityTypes.message, + text="More information required", + input_hint="expectingInput", + ) + ) + return + elif text == STRUCTURED_RESULT_TRIGGER_TEXT: + await context.send_activity( + Activity( + type=ActivityTypes.end_of_conversation, + code=EndOfConversationCodes.completed_successfully, + text="Completed with structured data", + value={"answer": 42}, + ) + ) + return + elif text == LONG_RUNNING_TRIGGER_TEXT: + await context.send_activity( + Activity( + type=ActivityTypes.message, + text="Working", + ) + ) + long_running_started.set() + await release_long_running.wait() + await context.send_activity( + Activity( + type=ActivityTypes.end_of_conversation, + code=EndOfConversationCodes.completed_successfully, + ) + ) + return + elif text == FAILURE_TRIGGER_TEXT: + raise RuntimeError("Integration agent failed") + elif text == PARTS_TRIGGER_TEXT: + await context.send_activity( + Activity( + type=ActivityTypes.message, + text=text, + attachments=list(context.activity.attachments or []), + ) + ) + else: + await context.send_activity(f"Echo: {text}") + await context.send_activity( + Activity( + type=ActivityTypes.end_of_conversation, + code=EndOfConversationCodes.completed_successfully, + ) + ) + + return application + + +@pytest_asyncio.fixture +async def a2a_harness() -> AsyncIterator[A2ATestHarness]: + """Create an unauthenticated application exposing JSON-RPC and REST A2A routes.""" + + app = FastAPI() + long_running_started = asyncio.Event() + release_long_running = asyncio.Event() + agent = _create_agent_application(long_running_started, release_long_running) + adapter = A2AAdapter( + agent, + agent_card_name="Compatibility Agent", + agent_card_description="A deterministic A2A compatibility test agent.", + agent_card_version="1.0.0", + agent_interfaces=[ + AgentInterface( + url="/rpc", + protocol_binding=TransportProtocol.JSONRPC, + ), + AgentInterface( + url="/rest", + protocol_binding=TransportProtocol.HTTP_JSON, + ), + ], + skills=[ + AgentSkill( + id="echo", + name="Echo", + description="Echoes the incoming text.", + tags=["test"], + examples=["Hello"], + input_modes=["text/plain"], + output_modes=["text/plain"], + ) + ], + ) + add_a2a(app, agent, adapter=adapter, use_jwt_middleware=False) + + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url="http://testserver", + ) as client: + yield A2ATestHarness( + client=client, + long_running_started=long_running_started, + release_long_running=release_long_running, + ) + + +@pytest_asyncio.fixture +async def a2a_client(a2a_harness: A2ATestHarness) -> httpx.AsyncClient: + return a2a_harness.client diff --git a/tests/hosting_a2a/integration/test_blob_protocol.py b/tests/hosting_a2a/integration/test_blob_protocol.py new file mode 100644 index 000000000..8154beab7 --- /dev/null +++ b/tests/hosting_a2a/integration/test_blob_protocol.py @@ -0,0 +1,207 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import asyncio +import os +import uuid +from contextlib import asynccontextmanager + +import httpx +import pytest +from a2a.types import ( + AgentInterface, + ListTasksRequest, + Message, + Part, + Role, + SendMessageRequest, +) +from a2a.utils.constants import TransportProtocol +from azure.core.exceptions import ResourceNotFoundError +from azure.identity.aio import DefaultAzureCredential +from azure.storage.blob.aio import BlobServiceClient, ContainerClient +from dotenv import load_dotenv +from fastapi import FastAPI +from google.protobuf.json_format import MessageToDict + +from microsoft_agents.hosting.a2a import A2AAdapter, add_a2a +from microsoft_agents.hosting.a2a.blob_task_store import BlobTaskStore + +from .conftest import _create_agent_application + +pytestmark = [ + pytest.mark.blob, + pytest.mark.integration, + pytest.mark.filterwarnings( + "ignore:label\\(\\) is deprecated\\. Use is_required\\(\\) or is_repeated\\(\\) instead\\.:DeprecationWarning" + ), +] + + +def _send_request(text: str) -> SendMessageRequest: + return SendMessageRequest( + message=Message( + role=Role.ROLE_USER, + message_id=str(uuid.uuid4()), + parts=[Part(text=text)], + ) + ) + + +def _rpc_request(method: str, params: object, request_id: str) -> dict: + return { + "jsonrpc": "2.0", + "id": request_id, + "method": method, + "params": MessageToDict(params), + } + + +async def _post_rpc( + client: httpx.AsyncClient, + method: str, + params: object, + request_id: str, +) -> dict: + response = await client.post( + "/rpc", + json=_rpc_request(method, params, request_id), + headers={"A2A-Version": "1.0"}, + ) + assert response.status_code == 200 + return response.json() + + +@asynccontextmanager +async def _blob_container(): + load_dotenv() + connection_string = os.environ.get("TEST_BLOB_STORAGE_CONNECTION_STRING") + credential = None + + if connection_string: + service_client = BlobServiceClient.from_connection_string(connection_string) + else: + account_url = os.environ.get("TEST_BLOB_STORAGE_ACCOUNT_URL") + if not account_url: + pytest.skip( + "Set TEST_BLOB_STORAGE_CONNECTION_STRING or " + "TEST_BLOB_STORAGE_ACCOUNT_URL" + ) + credential = DefaultAzureCredential() + service_client = BlobServiceClient(account_url, credential=credential) + + container_name = f"asdka2aprotocol{uuid.uuid4().hex}" + container_client = service_client.get_container_client(container_name) + + try: + yield container_client + finally: + try: + await container_client.delete_container() + except ResourceNotFoundError: + pass + await container_client.close() + await service_client.close() + if credential: + await credential.close() + + +@asynccontextmanager +async def _blob_backed_a2a_client(container_client: ContainerClient): + app = FastAPI() + agent = _create_agent_application(asyncio.Event(), asyncio.Event()) + adapter = A2AAdapter( + agent, + agent_interfaces=[ + AgentInterface( + url="/rpc", + protocol_binding=TransportProtocol.JSONRPC, + ), + AgentInterface( + url="/rest", + protocol_binding=TransportProtocol.HTTP_JSON, + ), + ], + task_store=BlobTaskStore(container_client), + ) + add_a2a(app, agent, adapter=adapter, use_jwt_middleware=False) + + try: + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url="http://testserver", + ) as client: + yield client + finally: + await adapter.a2a_request_handler.aclose() + + +@pytest.mark.asyncio +async def test_task_persists_across_app_restart_and_protocols(): + async with _blob_container() as container_client: + async with _blob_backed_a2a_client(container_client) as client: + sent = await _post_rpc( + client, + "SendMessage", + _send_request("persisted"), + "send-task", + ) + task = sent["result"]["task"] + + async with _blob_backed_a2a_client(container_client) as restarted_client: + get_response = await restarted_client.get( + f"/rest/tasks/{task['id']}", + headers={"A2A-Version": "1.0"}, + ) + list_response = await restarted_client.get( + "/rest/tasks", + headers={"A2A-Version": "1.0"}, + ) + + assert get_response.status_code == 200 + assert get_response.json() == task + assert list_response.status_code == 200 + assert [item["id"] for item in list_response.json()["tasks"]] == [task["id"]] + + +@pytest.mark.asyncio +async def test_jsonrpc_list_tasks_uses_blob_continuation_tokens(): + async with _blob_container() as container_client: + async with _blob_backed_a2a_client(container_client) as client: + task_ids = set() + for index in range(3): + sent = await _post_rpc( + client, + "SendMessage", + _send_request(f"task-{index}"), + f"send-{index}", + ) + task_ids.add(sent["result"]["task"]["id"]) + + first_page = await _post_rpc( + client, + "ListTasks", + ListTasksRequest(page_size=2), + "list-first-page", + ) + first_result = first_page["result"] + second_page = await _post_rpc( + client, + "ListTasks", + ListTasksRequest( + page_size=2, + page_token=first_result["nextPageToken"], + ), + "list-second-page", + ) + second_result = second_page["result"] + + first_ids = {task["id"] for task in first_result["tasks"]} + second_ids = {task["id"] for task in second_result["tasks"]} + + assert len(first_ids) == 2 + assert first_result["nextPageToken"] + assert len(second_ids) == 1 + assert second_result.get("nextPageToken", "") == "" + assert first_ids.isdisjoint(second_ids) + assert first_ids | second_ids == task_ids diff --git a/tests/hosting_a2a/integration/test_protocol.py b/tests/hosting_a2a/integration/test_protocol.py new file mode 100644 index 000000000..0c35b7d87 --- /dev/null +++ b/tests/hosting_a2a/integration/test_protocol.py @@ -0,0 +1,821 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from __future__ import annotations + +import asyncio +import base64 +import json +import uuid + +import httpx +import pytest +from google.protobuf.json_format import MessageToDict, ParseDict +from google.protobuf.struct_pb2 import Value + +from a2a.types.a2a_pb2 import ( + CancelTaskRequest, + GetExtendedAgentCardRequest, + GetTaskRequest, + ListTasksRequest, + Message, + Part, + Role, + SendMessageRequest, + SubscribeToTaskRequest, +) + +from .conftest import ( + A2ATestHarness, + FAILURE_TRIGGER_TEXT, + INPUT_REQUIRED_TRIGGER_TEXT, + LONG_RUNNING_TRIGGER_TEXT, + PARTS_TRIGGER_TEXT, + STREAMING_TRIGGER_TEXT, + STRUCTURED_RESULT_TRIGGER_TEXT, +) + +pytestmark = pytest.mark.filterwarnings( + "ignore:label\\(\\) is deprecated\\. Use is_required\\(\\) or is_repeated\\(\\) instead\\.:DeprecationWarning" +) + + +def _message( + text: str = "hello", + *, + task_id: str = "", + context_id: str = "", + parts: list[Part] | None = None, +) -> Message: + return Message( + role=Role.ROLE_USER, + message_id=str(uuid.uuid4()), + task_id=task_id, + context_id=context_id, + parts=parts if parts is not None else [Part(text=text)], + ) + + +def _send_request( + text: str = "hello", + *, + task_id: str = "", + context_id: str = "", + parts: list[Part] | None = None, +) -> SendMessageRequest: + return SendMessageRequest( + message=_message( + text, + task_id=task_id, + context_id=context_id, + parts=parts, + ) + ) + + +def _rpc_request(method: str, params: object, request_id: str = "request-1") -> dict: + serialized_params = ( + MessageToDict(params) if hasattr(params, "DESCRIPTOR") else params + ) + return { + "jsonrpc": "2.0", + "id": request_id, + "method": method, + "params": serialized_params, + } + + +async def _post_rpc( + client: httpx.AsyncClient, + method: str, + params: object, + request_id: str = "request-1", +) -> dict: + response = await client.post( + "/rpc", + json=_rpc_request(method, params, request_id), + headers={"A2A-Version": "1.0"}, + ) + assert response.status_code == 200 + return response.json() + + +async def _stream_rpc( + client: httpx.AsyncClient, + method: str, + params: object, + request_id: str = "request-1", +) -> list[dict]: + events: list[dict] = [] + async with client.stream( + "POST", + "/rpc", + json=_rpc_request(method, params, request_id), + headers={"A2A-Version": "1.0", "Accept": "text/event-stream"}, + ) as response: + assert response.status_code == 200 + async for line in response.aiter_lines(): + if line.startswith("data:"): + events.append(json.loads(line[len("data:") :].strip())) + return events + + +async def _stream_rest( + client: httpx.AsyncClient, + params: SendMessageRequest, +) -> list[dict]: + events: list[dict] = [] + async with client.stream( + "POST", + "/rest/message:stream", + json=MessageToDict(params), + headers={"A2A-Version": "1.0", "Accept": "text/event-stream"}, + ) as response: + assert response.status_code == 200 + async for line in response.aiter_lines(): + if line.startswith("data:"): + events.append(json.loads(line[len("data:") :].strip())) + return events + + +async def _subscribe_rest( + client: httpx.AsyncClient, + task_id: str, +) -> list[dict]: + events: list[dict] = [] + async with client.stream( + "GET", + f"/rest/tasks/{task_id}:subscribe", + headers={"A2A-Version": "1.0", "Accept": "text/event-stream"}, + ) as response: + assert response.status_code == 200 + async for line in response.aiter_lines(): + if line.startswith("data:"): + events.append(json.loads(line[len("data:") :].strip())) + return events + + +async def _list_rest_tasks(client: httpx.AsyncClient) -> list[dict]: + response = await client.get( + "/rest/tasks", + headers={"A2A-Version": "1.0"}, + ) + assert response.status_code == 200 + return response.json()["tasks"] + + +async def _wait_for_task_state( + client: httpx.AsyncClient, + expected_state: str, +) -> dict: + async def wait() -> dict: + while True: + tasks = await _list_rest_tasks(client) + for task in tasks: + if task["status"]["state"] == expected_state: + return task + await asyncio.sleep(0) + + return await asyncio.wait_for(wait(), timeout=1) + + +def _history_text(task: dict) -> list[str]: + return [ + message["parts"][0]["text"] + for message in task["history"] + if message.get("parts") and "text" in message["parts"][0] + ] + + +def _task_id(response: dict) -> str: + return response["result"]["task"]["id"] + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_agent_card_describes_configured_interfaces_and_skill( + a2a_client: httpx.AsyncClient, +) -> None: + response = await a2a_client.get("/rpc/.well-known/agent-card.json") + + assert response.status_code == 200 + card = response.json() + assert card["name"] == "Compatibility Agent" + assert card["version"] == "1.0.0" + assert {item["url"] for item in card["supportedInterfaces"]} == { + "http://testserver/rpc", + "http://testserver/rest", + } + assert card["skills"][0]["id"] == "echo" + assert card["capabilities"]["streaming"] is True + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_jsonrpc_message_send_get_and_list( + a2a_client: httpx.AsyncClient, +) -> None: + response = await _post_rpc( + a2a_client, + "SendMessage", + _send_request("json-rpc"), + ) + + assert response["result"]["task"]["status"]["state"] == "TASK_STATE_COMPLETED" + task_id = _task_id(response) + context_id = response["result"]["task"]["contextId"] + assert response["result"]["task"]["history"][1]["parts"][0]["text"] == ( + "Echo: json-rpc" + ) + + get_response = await _post_rpc( + a2a_client, + "GetTask", + GetTaskRequest(id=task_id), + request_id="get-task", + ) + assert get_response["result"]["id"] == task_id + assert get_response["result"]["contextId"] == context_id + + list_response = await _post_rpc( + a2a_client, + "ListTasks", + ListTasksRequest(context_id=context_id), + request_id="list-tasks", + ) + assert [task["id"] for task in list_response["result"]["tasks"]] == [task_id] + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_rest_message_send_get_and_list_match_jsonrpc_semantics( + a2a_client: httpx.AsyncClient, +) -> None: + request = _send_request("rest") + response = await a2a_client.post( + "/rest/message:send", + json=MessageToDict(request), + headers={"A2A-Version": "1.0"}, + ) + + assert response.status_code == 200 + task = response.json()["task"] + assert task["status"]["state"] == "TASK_STATE_COMPLETED" + task_id = task["id"] + + get_response = await a2a_client.get( + f"/rest/tasks/{task_id}", + headers={"A2A-Version": "1.0"}, + ) + assert get_response.status_code == 200 + assert get_response.json()["id"] == task_id + + list_response = await a2a_client.get( + "/rest/tasks", + headers={"A2A-Version": "1.0"}, + ) + assert list_response.status_code == 200 + assert any(item["id"] == task_id for item in list_response.json()["tasks"]) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_jsonrpc_continues_input_required_task( + a2a_client: httpx.AsyncClient, +) -> None: + first_response = await _post_rpc( + a2a_client, + "SendMessage", + _send_request(INPUT_REQUIRED_TRIGGER_TEXT), + ) + first_task = first_response["result"]["task"] + assert first_task["status"]["state"] == "TASK_STATE_INPUT_REQUIRED" + + second_response = await _post_rpc( + a2a_client, + "SendMessage", + _send_request( + "additional input", + task_id=first_task["id"], + context_id=first_task["contextId"], + ), + request_id="continue-task", + ) + second_task = second_response["result"]["task"] + + assert second_task["id"] == first_task["id"] + assert second_task["contextId"] == first_task["contextId"] + assert second_task["status"]["state"] == "TASK_STATE_COMPLETED" + assert [message["parts"][0]["text"] for message in second_task["history"]] == [ + INPUT_REQUIRED_TRIGGER_TEXT, + "More information required", + "additional input", + "Echo: additional input", + ] + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_rest_continues_input_required_task( + a2a_client: httpx.AsyncClient, +) -> None: + first_response = await a2a_client.post( + "/rest/message:send", + json=MessageToDict(_send_request(INPUT_REQUIRED_TRIGGER_TEXT)), + headers={"A2A-Version": "1.0"}, + ) + assert first_response.status_code == 200 + first_task = first_response.json()["task"] + assert first_task["status"]["state"] == "TASK_STATE_INPUT_REQUIRED" + + second_response = await a2a_client.post( + "/rest/message:send", + json=MessageToDict( + _send_request( + "additional input", + task_id=first_task["id"], + context_id=first_task["contextId"], + ) + ), + headers={"A2A-Version": "1.0"}, + ) + assert second_response.status_code == 200 + second_task = second_response.json()["task"] + + assert second_task["id"] == first_task["id"] + assert second_task["contextId"] == first_task["contextId"] + assert second_task["status"]["state"] == "TASK_STATE_COMPLETED" + assert _history_text(second_task) == [ + INPUT_REQUIRED_TRIGGER_TEXT, + "More information required", + "additional input", + "Echo: additional input", + ] + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_jsonrpc_returns_structured_completion_artifact_and_status_message( + a2a_client: httpx.AsyncClient, +) -> None: + response = await _post_rpc( + a2a_client, + "SendMessage", + _send_request(STRUCTURED_RESULT_TRIGGER_TEXT), + ) + task = response["result"]["task"] + + assert task["status"]["state"] == "TASK_STATE_COMPLETED" + assert task["status"]["message"]["parts"] == [ + {"text": "Completed with structured data"} + ] + assert task["artifacts"][0]["name"] == "Result" + result_part = task["artifacts"][0]["parts"][0] + assert result_part["data"] == {"answer": 42.0} + assert result_part["metadata"]["mimeType"] == "application/json" + assert result_part["metadata"]["type"] == "object" + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_rest_returns_structured_completion_artifact_and_status_message( + a2a_client: httpx.AsyncClient, +) -> None: + response = await a2a_client.post( + "/rest/message:send", + json=MessageToDict(_send_request(STRUCTURED_RESULT_TRIGGER_TEXT)), + headers={"A2A-Version": "1.0"}, + ) + assert response.status_code == 200 + task = response.json()["task"] + + assert task["status"]["state"] == "TASK_STATE_COMPLETED" + assert task["status"]["message"]["parts"] == [ + {"text": "Completed with structured data"} + ] + result_part = task["artifacts"][0]["parts"][0] + assert result_part["data"] == {"answer": 42.0} + assert result_part["metadata"]["mimeType"] == "application/json" + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize("transport", ["jsonrpc", "rest"]) +async def test_message_parts_round_trip_through_protocol( + a2a_client: httpx.AsyncClient, + transport: str, +) -> None: + request = _send_request( + parts=[ + Part(text=PARTS_TRIGGER_TEXT), + Part( + url="https://example.com/image.png", + media_type="image/png", + filename="image.png", + ), + Part( + raw=b"file contents", + media_type="application/octet-stream", + filename="data.bin", + ), + Part( + data=ParseDict({"answer": 42}, Value()), + media_type="application/json", + filename="result.json", + ), + ] + ) + + if transport == "jsonrpc": + response = await _post_rpc(a2a_client, "SendMessage", request) + task = response["result"]["task"] + else: + response = await a2a_client.post( + "/rest/message:send", + json=MessageToDict(request), + headers={"A2A-Version": "1.0"}, + ) + assert response.status_code == 200 + task = response.json()["task"] + + parts = task["history"][1]["parts"] + assert parts[0] == {"text": PARTS_TRIGGER_TEXT} + assert parts[1] == { + "url": "https://example.com/image.png", + "mediaType": "image/png", + "filename": "image.png", + } + assert parts[2] == { + "raw": base64.b64encode(b"file contents").decode(), + "mediaType": "application/octet-stream", + "filename": "data.bin", + } + assert parts[3]["data"] == {"answer": 42.0} + assert parts[3]["mediaType"] == "application/json" + assert parts[3]["filename"] == "result.json" + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_jsonrpc_cancel_missing_task_returns_protocol_error( + a2a_client: httpx.AsyncClient, +) -> None: + response = await _post_rpc( + a2a_client, + "CancelTask", + CancelTaskRequest(id="missing-task"), + request_id="cancel-task", + ) + + assert response["id"] == "cancel-task" + assert response["error"]["code"] != 0 + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize("transport", ["jsonrpc", "rest"]) +async def test_cancel_active_task_persists_canceled_state( + a2a_harness: A2ATestHarness, + transport: str, +) -> None: + client = a2a_harness.client + request = _send_request(LONG_RUNNING_TRIGGER_TEXT) + + if transport == "jsonrpc": + send_call = asyncio.create_task(_post_rpc(client, "SendMessage", request)) + else: + send_call = asyncio.create_task( + client.post( + "/rest/message:send", + json=MessageToDict(request), + headers={"A2A-Version": "1.0"}, + ) + ) + + await asyncio.wait_for(a2a_harness.long_running_started.wait(), timeout=1) + active_task = await _wait_for_task_state(client, "TASK_STATE_WORKING") + + if transport == "jsonrpc": + cancel_response = await _post_rpc( + client, + "CancelTask", + CancelTaskRequest(id=active_task["id"]), + request_id="cancel-active", + ) + else: + cancel_response = await client.post( + f"/rest/tasks/{active_task['id']}:cancel", + headers={"A2A-Version": "1.0"}, + ) + + a2a_harness.release_long_running.set() + send_result = await asyncio.wait_for(send_call, timeout=1) + + if transport == "jsonrpc": + canceled_task = cancel_response["result"] + send_task = send_result["result"]["task"] + else: + assert cancel_response.status_code == 200 + canceled_task = cancel_response.json() + assert send_result.status_code == 200 + send_task = send_result.json()["task"] + + assert canceled_task["status"]["state"] == "TASK_STATE_CANCELED" + assert send_task["status"]["state"] == "TASK_STATE_CANCELED" + + persisted = await client.get( + f"/rest/tasks/{active_task['id']}", + headers={"A2A-Version": "1.0"}, + ) + assert persisted.status_code == 200 + assert persisted.json()["status"]["state"] == "TASK_STATE_CANCELED" + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize("transport", ["jsonrpc", "rest"]) +async def test_agent_failure_returns_protocol_error_and_persists_failed_task( + a2a_client: httpx.AsyncClient, + transport: str, +) -> None: + request = _send_request(FAILURE_TRIGGER_TEXT) + + if transport == "jsonrpc": + response = await _post_rpc(a2a_client, "SendMessage", request) + assert "error" in response + else: + response = await a2a_client.post( + "/rest/message:send", + json=MessageToDict(request), + headers={"A2A-Version": "1.0"}, + ) + assert response.status_code == 500 + + tasks = await _list_rest_tasks(a2a_client) + assert len(tasks) == 1 + assert tasks[0]["status"]["state"] == "TASK_STATE_FAILED" + assert _history_text(tasks[0]) == [FAILURE_TRIGGER_TEXT] + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_jsonrpc_rejects_unknown_method( + a2a_client: httpx.AsyncClient, +) -> None: + response = await a2a_client.post( + "/rpc", + json=_rpc_request("UnsupportedMethod", {}, request_id="unknown"), + headers={"A2A-Version": "1.0"}, + ) + + assert response.status_code == 200 + payload = response.json() + assert payload["id"] == "unknown" + assert payload["error"]["code"] == -32601 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_jsonrpc_streaming_message_emits_ordered_task_events( + a2a_client: httpx.AsyncClient, +) -> None: + events = await _stream_rpc( + a2a_client, + "SendStreamingMessage", + _send_request(STREAMING_TRIGGER_TEXT), + request_id="stream-1", + ) + + assert [event["id"] for event in events] == ["stream-1"] * len(events) + + # First event is the newly created task in the submitted state. + task = events[0]["result"]["task"] + task_id = task["id"] + assert task["status"]["state"] == "TASK_STATE_SUBMITTED" + + # The queued informative update surfaces as a "working" status update + # carrying the informative text as an agent message, before any content. + working_update = events[1]["result"]["statusUpdate"] + assert working_update["taskId"] == task_id + assert working_update["status"]["state"] == "TASK_STATE_WORKING" + assert working_update["status"]["message"]["parts"][0]["text"] == "Thinking..." + + # The queued text chunks are combined into a single artifact update. + artifact_update = events[2]["result"]["artifactUpdate"] + assert artifact_update["taskId"] == task_id + assert artifact_update["artifact"]["parts"][0]["text"] == ( + f"Echo: {STREAMING_TRIGGER_TEXT}" + ) + + # The stream ends with a "completed" status update once end_stream() + # finishes and the end-of-conversation activity is processed. + final_update = events[-1]["result"]["statusUpdate"] + assert final_update["taskId"] == task_id + assert final_update["status"]["state"] == "TASK_STATE_COMPLETED" + + # The persisted task reflects the same terminal state via a plain GetTask. + get_response = await _post_rpc( + a2a_client, + "GetTask", + GetTaskRequest(id=task_id), + request_id="get-task", + ) + assert get_response["result"]["status"]["state"] == "TASK_STATE_COMPLETED" + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_rest_streaming_message_emits_ordered_task_events( + a2a_client: httpx.AsyncClient, +) -> None: + events = await _stream_rest( + a2a_client, + _send_request(STREAMING_TRIGGER_TEXT), + ) + + task = events[0]["task"] + task_id = task["id"] + assert task["status"]["state"] == "TASK_STATE_SUBMITTED" + + working_update = events[1]["statusUpdate"] + assert working_update["taskId"] == task_id + assert working_update["status"]["state"] == "TASK_STATE_WORKING" + assert working_update["status"]["message"]["parts"][0]["text"] == "Thinking..." + + artifact_update = events[2]["artifactUpdate"] + assert artifact_update["taskId"] == task_id + assert artifact_update["artifact"]["parts"][0]["text"] == ( + f"Echo: {STREAMING_TRIGGER_TEXT}" + ) + + final_update = events[-1]["statusUpdate"] + assert final_update["taskId"] == task_id + assert final_update["status"]["state"] == "TASK_STATE_COMPLETED" + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize("transport", ["jsonrpc", "rest"]) +async def test_subscribe_to_active_task_streams_initial_and_completed_state( + a2a_harness: A2ATestHarness, + transport: str, +) -> None: + client = a2a_harness.client + send_call = asyncio.create_task( + client.post( + "/rest/message:send", + json=MessageToDict(_send_request(LONG_RUNNING_TRIGGER_TEXT)), + headers={"A2A-Version": "1.0"}, + ) + ) + + await asyncio.wait_for(a2a_harness.long_running_started.wait(), timeout=1) + active_task = await _wait_for_task_state(client, "TASK_STATE_WORKING") + + if transport == "jsonrpc": + subscribe_call = asyncio.create_task( + _stream_rpc( + client, + "SubscribeToTask", + SubscribeToTaskRequest(id=active_task["id"]), + request_id="subscribe-task", + ) + ) + else: + subscribe_call = asyncio.create_task(_subscribe_rest(client, active_task["id"])) + + await asyncio.sleep(0.05) + a2a_harness.release_long_running.set() + + events = await asyncio.wait_for(subscribe_call, timeout=1) + send_response = await asyncio.wait_for(send_call, timeout=1) + assert send_response.status_code == 200 + + if transport == "jsonrpc": + initial_task = events[0]["result"]["task"] + terminal_update = events[-1]["result"]["statusUpdate"] + else: + initial_task = events[0]["task"] + terminal_update = events[-1]["statusUpdate"] + + assert initial_task["id"] == active_task["id"] + assert initial_task["status"]["state"] == "TASK_STATE_WORKING" + assert terminal_update["status"]["state"] == "TASK_STATE_COMPLETED" + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize("transport", ["jsonrpc", "rest"]) +async def test_subscribe_to_completed_task_returns_unsupported_operation( + a2a_client: httpx.AsyncClient, + transport: str, +) -> None: + send_response = await _post_rpc( + a2a_client, + "SendMessage", + _send_request("complete-task"), + ) + task_id = send_response["result"]["task"]["id"] + + if transport == "jsonrpc": + response = await _post_rpc( + a2a_client, + "SubscribeToTask", + SubscribeToTaskRequest(id=task_id), + request_id="subscribe-completed", + ) + + assert response["error"]["code"] == -32004 + assert response["error"]["data"][0]["reason"] == "UNSUPPORTED_OPERATION" + else: + response = await a2a_client.get( + f"/rest/tasks/{task_id}:subscribe", + headers={ + "A2A-Version": "1.0", + "Accept": "text/event-stream", + }, + ) + + assert response.status_code == 400 + assert response.json()["error"]["details"][0]["reason"] == ( + "UNSUPPORTED_OPERATION" + ) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_jsonrpc_push_notification_config_is_explicitly_unsupported( + a2a_client: httpx.AsyncClient, +) -> None: + """The adapter's agent card does not advertise push-notification + capability, so every push-notification method must fail with the + protocol-defined PUSH_NOTIFICATION_NOT_SUPPORTED error rather than a + generic failure or a silent no-op.""" + + response = await _post_rpc( + a2a_client, + "CreateTaskPushNotificationConfig", + {"taskId": "missing-task", "url": "https://example.com/webhook"}, + request_id="push-create", + ) + + assert response["error"]["code"] == -32003 + assert response["error"]["message"] == ( + "Push notifications are not supported by the agent" + ) + error_details = response["error"]["data"][0] + assert error_details["reason"] == "PUSH_NOTIFICATION_NOT_SUPPORTED" + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_rest_push_notification_config_is_explicitly_unsupported( + a2a_client: httpx.AsyncClient, +) -> None: + response = await a2a_client.post( + "/rest/tasks/missing-task/pushNotificationConfigs", + json={"url": "https://example.com/webhook"}, + headers={"A2A-Version": "1.0"}, + ) + + assert response.status_code == 400 + payload = response.json() + assert payload["error"]["status"] == "FAILED_PRECONDITION" + assert payload["error"]["details"][0]["reason"] == ( + "PUSH_NOTIFICATION_NOT_SUPPORTED" + ) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_jsonrpc_extended_agent_card_is_explicitly_unsupported( + a2a_client: httpx.AsyncClient, +) -> None: + """The basic agent card does not advertise extended Agent Card support.""" + + response = await _post_rpc( + a2a_client, + "GetExtendedAgentCard", + GetExtendedAgentCardRequest(), + request_id="extended-card", + ) + + assert response["error"]["code"] == -32004 + error_details = response["error"]["data"][0] + assert error_details["reason"] == "UNSUPPORTED_OPERATION" + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_rest_extended_agent_card_is_explicitly_unsupported( + a2a_client: httpx.AsyncClient, +) -> None: + response = await a2a_client.get( + "/rest/extendedAgentCard", + headers={"A2A-Version": "1.0"}, + ) + + assert response.status_code == 400 + payload = response.json() + assert payload["error"]["status"] == "FAILED_PRECONDITION" + assert payload["error"]["details"][0]["reason"] == "UNSUPPORTED_OPERATION" diff --git a/tests/hosting_a2a/request_handling/__init__.py b/tests/hosting_a2a/request_handling/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/hosting_a2a/request_handling/test_request_handling.py b/tests/hosting_a2a/request_handling/test_request_handling.py new file mode 100644 index 000000000..1f3927062 --- /dev/null +++ b/tests/hosting_a2a/request_handling/test_request_handling.py @@ -0,0 +1,520 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import asyncio + +import pytest +from a2a.server.agent_execution import RequestContext +from a2a.server.context import ServerCallContext +from a2a.server.events import EventQueueLegacy +from a2a.server.tasks import InMemoryTaskStore, TaskUpdater +from a2a.types import ( + AgentCapabilities, + AgentCard, + AgentInterface, + CancelTaskRequest, + GetExtendedAgentCardRequest, + GetTaskRequest, + ListTasksRequest, + Message, + Part, + Role, + SendMessageRequest, + SubscribeToTaskRequest, + Task, + TaskState, + TaskStatus, + TaskStatusUpdateEvent, +) +from a2a.utils.errors import ( + ExtendedAgentCardNotConfiguredError, + TaskNotFoundError, + UnsupportedOperationError, +) + +from microsoft_agents.hosting.a2a.request_handling import ( + A2AAgentExecutor, + A2ARequestHandler, +) + +pytestmark = pytest.mark.filterwarnings( + "ignore:label\\(\\) is deprecated\\. Use is_required\\(\\) or is_repeated\\(\\) instead\\.:DeprecationWarning" +) + + +def _agent_card(name: str = "Test agent") -> AgentCard: + return AgentCard( + name=name, + description="Test agent description", + version="1.0.0", + supported_interfaces=[ + AgentInterface( + url="https://example.com/a2a", + protocol_binding="JSONRPC", + ) + ], + capabilities=AgentCapabilities(), + ) + + +def _message( + text: str, + *, + message_id: str, + role: Role = Role.ROLE_USER, + task_id: str = "", + context_id: str = "", +) -> Message: + return Message( + message_id=message_id, + role=role, + task_id=task_id, + context_id=context_id, + parts=[Part(text=text)], + ) + + +def _send_request( + text: str, + *, + message_id: str, + task_id: str = "", + context_id: str = "", +) -> SendMessageRequest: + return SendMessageRequest( + message=_message( + text, + message_id=message_id, + task_id=task_id, + context_id=context_id, + ) + ) + + +def _history_text(task: Task) -> list[str]: + return [message.parts[0].text for message in task.history] + + +def _status_text(task: Task) -> str | None: + if not task.status.HasField("message"): + return None + return task.status.message.parts[0].text + + +async def _wait_for_task_state(task_store, task_id, call_context, expected_state): + async def wait(): + while True: + task = await task_store.get(task_id, call_context) + if task is not None and task.status.state == expected_state: + return task + await asyncio.sleep(0) + + return await asyncio.wait_for(wait(), timeout=1) + + +class _CompletingAdapter: + def __init__(self) -> None: + self.contexts: list[RequestContext] = [] + + async def execute_agent_turn(self, context, event_queue) -> None: + self.contexts.append(context) + updater = TaskUpdater( + event_queue=event_queue, + task_id=context.task_id, + context_id=context.context_id, + ) + await updater.complete( + _message( + "completed", + message_id="response-1", + role=Role.ROLE_AGENT, + task_id=context.task_id, + context_id=context.context_id, + ) + ) + + async def cancel_agent_turn(self, context, event_queue) -> None: + raise AssertionError("Completed tasks should not be canceled") + + +class _InputThenCompleteAdapter: + def __init__(self) -> None: + self.contexts: list[RequestContext] = [] + self.current_task_states: list[int | None] = [] + + async def execute_agent_turn(self, context, event_queue) -> None: + self.contexts.append(context) + self.current_task_states.append( + context.current_task.status.state if context.current_task else None + ) + updater = TaskUpdater( + event_queue=event_queue, + task_id=context.task_id, + context_id=context.context_id, + ) + + if len(self.contexts) == 1: + await updater.update_status( + TaskState.TASK_STATE_INPUT_REQUIRED, + _message( + "More information required", + message_id="prompt-1", + role=Role.ROLE_AGENT, + task_id=context.task_id, + context_id=context.context_id, + ), + ) + return + + await updater.complete( + _message( + "Request completed", + message_id="response-2", + role=Role.ROLE_AGENT, + task_id=context.task_id, + context_id=context.context_id, + ) + ) + + async def cancel_agent_turn(self, context, event_queue) -> None: + raise AssertionError("This task should not be canceled") + + +class _CancelableAdapter: + def __init__(self) -> None: + self.started = asyncio.Event() + self.cancel_context: RequestContext | None = None + + async def execute_agent_turn(self, context, event_queue) -> None: + updater = TaskUpdater( + event_queue=event_queue, + task_id=context.task_id, + context_id=context.context_id, + ) + await updater.start_work() + self.started.set() + await asyncio.Event().wait() + + async def cancel_agent_turn(self, context, event_queue) -> None: + self.cancel_context = context + + +class _FailingAdapter: + async def execute_agent_turn(self, context, event_queue) -> None: + raise RuntimeError("agent failed") + + async def cancel_agent_turn(self, context, event_queue) -> None: + raise AssertionError("Failed tasks should not be canceled") + + +@pytest.mark.asyncio +async def test_executor_drops_request_without_message(): + adapter = _CompletingAdapter() + executor = A2AAgentExecutor(adapter) + context = RequestContext( + ServerCallContext(), + task_id="task-1", + context_id="context-1", + ) + event_queue = EventQueueLegacy() + + await executor.execute(context, event_queue) + + assert adapter.contexts == [] + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(event_queue.dequeue_event(), timeout=0.05) + + +@pytest.mark.asyncio +async def test_executor_establishes_task_mode_before_adapter_events(): + adapter = _CompletingAdapter() + executor = A2AAgentExecutor(adapter) + request = _send_request("hello", message_id="message-1") + context = RequestContext( + ServerCallContext(), + request=request, + task_id="task-1", + context_id="context-1", + ) + event_queue = EventQueueLegacy() + + await executor.execute(context, event_queue) + + initial_task = await event_queue.dequeue_event() + completion = await event_queue.dequeue_event() + + assert isinstance(initial_task, Task) + assert initial_task.id == "task-1" + assert initial_task.context_id == "context-1" + assert initial_task.status.state == TaskState.TASK_STATE_SUBMITTED + assert initial_task.status.HasField("timestamp") + assert list(initial_task.history) == [request.message] + + assert isinstance(completion, TaskStatusUpdateEvent) + assert completion.task_id == initial_task.id + assert completion.context_id == initial_task.context_id + assert completion.status.state == TaskState.TASK_STATE_COMPLETED + + +@pytest.mark.asyncio +async def test_executor_continuation_updates_existing_task_without_replacing_it(): + adapter = _CompletingAdapter() + executor = A2AAgentExecutor(adapter) + existing_task = Task( + id="task-1", + context_id="context-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ) + context = RequestContext( + ServerCallContext(), + request=_send_request( + "additional input", + message_id="message-2", + task_id=existing_task.id, + context_id=existing_task.context_id, + ), + task_id=existing_task.id, + context_id=existing_task.context_id, + task=existing_task, + ) + event_queue = EventQueueLegacy() + + await executor.execute(context, event_queue) + + event = await event_queue.dequeue_event() + assert isinstance(event, TaskStatusUpdateEvent) + assert event.task_id == existing_task.id + assert event.status.state == TaskState.TASK_STATE_COMPLETED + assert adapter.contexts[0].current_task is existing_task + + +@pytest.mark.asyncio +async def test_request_handler_completes_persists_and_queries_task(): + adapter = _CompletingAdapter() + task_store = InMemoryTaskStore() + handler = A2ARequestHandler(adapter, task_store, _agent_card()) + call_context = ServerCallContext() + request = _send_request("hello", message_id="message-1") + + try: + result = await handler.on_message_send(request, call_context) + + assert isinstance(result, Task) + assert result.status.state == TaskState.TASK_STATE_COMPLETED + assert request.message.task_id == result.id + assert request.message.context_id == result.context_id + assert _history_text(result) == ["hello"] + assert _status_text(result) == "completed" + assert adapter.contexts[0].current_task is None + + stored = await task_store.get(result.id, call_context) + assert stored == result + + without_history = await handler.on_get_task( + GetTaskRequest(id=result.id, history_length=0), + call_context, + ) + assert _history_text(without_history) == [] + assert _status_text(without_history) == "completed" + assert _history_text(stored) == ["hello"] + + listed = await handler.on_list_tasks( + ListTasksRequest(context_id=result.context_id), + call_context, + ) + assert [task.id for task in listed.tasks] == [result.id] + assert _status_text(listed.tasks[0]) == "completed" + finally: + await handler.aclose() + + +@pytest.mark.asyncio +async def test_request_handler_continues_interrupted_task_in_same_a2a_context(): + adapter = _InputThenCompleteAdapter() + task_store = InMemoryTaskStore() + handler = A2ARequestHandler(adapter, task_store, _agent_card()) + call_context = ServerCallContext() + + try: + first_result = await handler.on_message_send( + _send_request("initial request", message_id="message-1"), + call_context, + ) + + assert isinstance(first_result, Task) + assert first_result.status.state == TaskState.TASK_STATE_INPUT_REQUIRED + assert _history_text(first_result) == ["initial request"] + assert _status_text(first_result) == "More information required" + + second_result = await handler.on_message_send( + _send_request( + "additional input", + message_id="message-2", + task_id=first_result.id, + context_id=first_result.context_id, + ), + call_context, + ) + + assert isinstance(second_result, Task) + assert second_result.id == first_result.id + assert second_result.context_id == first_result.context_id + assert second_result.status.state == TaskState.TASK_STATE_COMPLETED + assert _history_text(second_result) == [ + "initial request", + "More information required", + "additional input", + ] + assert _status_text(second_result) == "Request completed" + assert adapter.current_task_states[0] is None + assert adapter.current_task_states[1] == TaskState.TASK_STATE_INPUT_REQUIRED + assert adapter.contexts[1].current_task.id == first_result.id + finally: + await handler.aclose() + + +@pytest.mark.asyncio +async def test_request_handler_cancels_and_persists_active_task(): + adapter = _CancelableAdapter() + task_store = InMemoryTaskStore() + handler = A2ARequestHandler(adapter, task_store, _agent_card()) + call_context = ServerCallContext() + request = _send_request("long-running request", message_id="message-1") + send_call = asyncio.create_task(handler.on_message_send(request, call_context)) + + try: + await asyncio.wait_for(adapter.started.wait(), timeout=1) + active_task = await _wait_for_task_state( + task_store, + request.message.task_id, + call_context, + TaskState.TASK_STATE_WORKING, + ) + canceled = await asyncio.wait_for( + handler.on_cancel_task( + CancelTaskRequest(id=request.message.task_id), + call_context, + ), + timeout=1, + ) + send_result = await asyncio.wait_for(send_call, timeout=1) + + assert canceled.status.state == TaskState.TASK_STATE_CANCELED + assert send_result == canceled + assert adapter.cancel_context is not None + assert adapter.cancel_context.task_id == canceled.id + assert adapter.cancel_context.context_id == canceled.context_id + assert active_task.status.state == TaskState.TASK_STATE_WORKING + assert adapter.cancel_context.current_task.id == active_task.id + + stored = await task_store.get(canceled.id, call_context) + assert stored == canceled + finally: + if not send_call.done(): + send_call.cancel() + await handler.aclose() + + +@pytest.mark.asyncio +async def test_request_handler_persists_failed_task_when_adapter_raises(): + task_store = InMemoryTaskStore() + handler = A2ARequestHandler(_FailingAdapter(), task_store, _agent_card()) + call_context = ServerCallContext() + request = _send_request("fail", message_id="message-1") + + try: + with pytest.raises(RuntimeError, match="agent failed"): + await handler.on_message_send(request, call_context) + + stored = await task_store.get(request.message.task_id, call_context) + assert stored is not None + assert stored.status.state == TaskState.TASK_STATE_FAILED + finally: + await handler.aclose() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "terminal_state", + [ + TaskState.TASK_STATE_COMPLETED, + TaskState.TASK_STATE_FAILED, + TaskState.TASK_STATE_CANCELED, + TaskState.TASK_STATE_REJECTED, + ], +) +async def test_request_handler_rejects_subscription_to_terminal_task( + terminal_state, +): + task_store = InMemoryTaskStore() + handler = A2ARequestHandler(_CompletingAdapter(), task_store, _agent_card()) + call_context = ServerCallContext() + task = Task( + id="task-1", + context_id="context-1", + status=TaskStatus(state=terminal_state), + ) + await task_store.save(task, call_context) + + try: + events = handler.on_subscribe_to_task( + SubscribeToTaskRequest(id=task.id), + call_context, + ) + + with pytest.raises( + UnsupportedOperationError, + match="Cannot subscribe to a terminal task", + ): + await anext(events) + finally: + await handler.aclose() + + +@pytest.mark.asyncio +async def test_request_handler_subscription_preserves_task_not_found_error(): + handler = A2ARequestHandler( + _CompletingAdapter(), + InMemoryTaskStore(), + _agent_card(), + ) + + try: + events = handler.on_subscribe_to_task( + SubscribeToTaskRequest(id="missing-task"), + ServerCallContext(), + ) + + with pytest.raises(TaskNotFoundError): + await anext(events) + finally: + await handler.aclose() + + +@pytest.mark.asyncio +async def test_request_handler_update_agent_card_changes_extended_card_capability(): + original_card = _agent_card() + updated_card = _agent_card("Updated agent") + updated_card.capabilities.extended_agent_card = True + handler = A2ARequestHandler( + _CompletingAdapter(), + InMemoryTaskStore(), + original_card, + ) + + try: + with pytest.raises(UnsupportedOperationError): + await handler.on_get_extended_agent_card( + GetExtendedAgentCardRequest(), + ServerCallContext(), + ) + + handler.update_agent_card(updated_card) + + with pytest.raises(ExtendedAgentCardNotConfiguredError): + await handler.on_get_extended_agent_card( + GetExtendedAgentCardRequest(), + ServerCallContext(), + ) + finally: + await handler.aclose() diff --git a/tests/hosting_a2a/server/__init__.py b/tests/hosting_a2a/server/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/hosting_a2a/server/test_routes.py b/tests/hosting_a2a/server/test_routes.py new file mode 100644 index 000000000..eeb9f5a45 --- /dev/null +++ b/tests/hosting_a2a/server/test_routes.py @@ -0,0 +1,204 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from a2a.server.request_handlers import RequestHandler +from a2a.types import AgentCapabilities, AgentCard, ListTasksResponse +from starlette.applications import Starlette +from starlette.requests import Request +from starlette.responses import JSONResponse +from starlette.routing import Mount, Route + +from microsoft_agents.hosting.core import AgentAuthConfiguration, ClaimsIdentity +from microsoft_agents.hosting.a2a.server import routes +from microsoft_agents.hosting.a2a.server._constants import _CLAIMS_IDENTITY_KEY + + +async def _request_with_identity(app, identity, method, path, **kwargs): + async def authenticated_app(scope, receive, send): + scope.setdefault("state", {})["claims_identity"] = identity + await app(scope, receive, send) + + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=authenticated_app), + base_url="http://testserver", + ) as client: + return await client.request(method, path, **kwargs) + + +@pytest.mark.asyncio +async def test_jsonrpc_routes_forward_authenticated_identity_to_request_handler(): + request_handler = MagicMock(spec=RequestHandler) + request_handler.on_list_tasks.return_value = ListTasksResponse(tasks=[]) + identity = ClaimsIdentity({"sub": "agent-1"}) + app = Starlette(routes=routes.create_jsonrpc_routes(request_handler, "/a2a")) + + response = await _request_with_identity( + app, + identity, + "POST", + "/a2a", + headers={"A2A-Version": "1.0"}, + json={ + "jsonrpc": "2.0", + "id": "request-1", + "method": "ListTasks", + "params": {}, + }, + ) + + assert response.status_code == 200 + assert response.json()["result"]["tasks"] == [] + context = request_handler.on_list_tasks.call_args.args[1] + assert context.state[_CLAIMS_IDENTITY_KEY] is identity + assert context.state["headers"]["a2a-version"] == "1.0" + + +@pytest.mark.asyncio +async def test_rest_routes_forward_authenticated_identity_to_request_handler(): + request_handler = MagicMock(spec=RequestHandler) + request_handler.on_list_tasks.return_value = ListTasksResponse(tasks=[]) + identity = ClaimsIdentity({"sub": "agent-1"}) + app = Starlette( + routes=routes.create_rest_routes( + request_handler, + path_prefix="/a2a", + ) + ) + + response = await _request_with_identity( + app, + identity, + "GET", + "/a2a/tasks", + headers={"A2A-Version": "1.0"}, + ) + + assert response.status_code == 200 + assert response.json()["tasks"] == [] + context = request_handler.on_list_tasks.call_args.args[1] + assert context.state[_CLAIMS_IDENTITY_KEY] is identity + assert context.state["headers"]["a2a-version"] == "1.0" + + +@pytest.mark.asyncio +async def test_create_agent_card_routes_adapts_request_and_uses_interface_prefix(): + observed = {} + + async def get_agent_card(request, prefix): + observed["request"] = request + observed["prefix"] = prefix + return AgentCard( + name="Test agent", + description="Test description", + version="1.0.0", + supported_interfaces=[], + capabilities=AgentCapabilities(), + ) + + route = routes.create_agent_card_routes( + get_agent_card, + "/a2a/.well-known/agent-card.json", + _cache_enabled=True, + )[0] + request = Request( + { + "type": "http", + "method": "GET", + "path": "/a2a/.well-known/agent-card.json", + "raw_path": b"/a2a/.well-known/agent-card.json", + "query_string": b"", + "headers": [], + "scheme": "https", + "server": ("example.com", 443), + "client": ("127.0.0.1", 1234), + } + ) + + response = await route.endpoint(request) + + assert route.path == "/a2a/.well-known/agent-card.json" + assert response.status_code == 200 + assert response.body == ( + b'{"name":"Test agent","description":"Test description",' + b'"version":"1.0.0","capabilities":{}}' + ) + assert ( + observed["request"].url == "https://example.com/a2a/.well-known/agent-card.json" + ) + assert observed["prefix"] == "/a2a" + assert response.headers["cache-control"] == "public, max-age=3600" + assert response.headers["etag"].startswith('"') + assert response.headers["etag"].endswith('"') + assert response.headers["last-modified"] + + +@pytest.mark.asyncio +async def test_create_agent_card_routes_uses_custom_url_as_prefix(): + observed = {} + + async def get_agent_card(request, prefix): + observed["prefix"] = prefix + return AgentCard( + name="Test agent", + description="Test description", + version="1.0.0", + supported_interfaces=[], + capabilities=AgentCapabilities(), + ) + + route = routes.create_agent_card_routes( + get_agent_card, + "/custom-agent-card", + )[0] + request = Request( + { + "type": "http", + "method": "GET", + "path": "/custom-agent-card", + "raw_path": b"/custom-agent-card", + "query_string": b"", + "headers": [], + "scheme": "https", + "server": ("example.com", 443), + "client": ("127.0.0.1", 1234), + } + ) + + response = await route.endpoint(request) + + assert response.status_code == 200 + assert observed["prefix"] == "/custom-agent-card" + assert response.headers["cache-control"] == "no-store" + assert "etag" not in response.headers + assert "last-modified" not in response.headers + + +@pytest.mark.asyncio +async def test_use_jwt_middleware_authorizes_shared_mounted_route_once(): + async def endpoint(request): + return JSONResponse({"ok": True}) + + route = Route("/messages", endpoint=endpoint) + mounted = Mount("/tenant", routes=[route]) + routes.use_jwt_middleware([route, mounted]) + app = Starlette(routes=[route]) + app.state.agent_configuration = AgentAuthConfiguration() + + with patch( + "microsoft_agents.hosting.fastapi.jwt_authorization_middleware." + "_authorize_request", + new=AsyncMock(return_value=ClaimsIdentity({"sub": "agent-1"})), + ) as authorize: + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url="http://testserver", + ) as client: + response = await client.get("/messages") + + assert response.status_code == 200 + assert response.json() == {"ok": True} + authorize.assert_awaited_once() diff --git a/tests/hosting_a2a/server/test_sdk_server_call_context_builder.py b/tests/hosting_a2a/server/test_sdk_server_call_context_builder.py new file mode 100644 index 000000000..c70c221dd --- /dev/null +++ b/tests/hosting_a2a/server/test_sdk_server_call_context_builder.py @@ -0,0 +1,67 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from a2a.auth.user import UnauthenticatedUser +from starlette.requests import Request + +from microsoft_agents.hosting.core import ClaimsIdentity +from microsoft_agents.hosting.a2a.server._constants import _CLAIMS_IDENTITY_KEY +from microsoft_agents.hosting.a2a.server.sdk_server_call_context_builder import ( + SDKServerCallContextBuilder, +) + + +def _request(*, claims_identity=None, auth=None): + scope = { + "type": "http", + "method": "GET", + "path": "/a2a", + "raw_path": b"/a2a", + "query_string": b"", + "headers": [ + (b"a2a-extensions", b"extension-1"), + (b"x-test", b"value"), + ], + "scheme": "https", + "server": ("example.com", 443), + "client": ("127.0.0.1", 1234), + "state": {}, + } + if claims_identity is not None: + scope["state"]["claims_identity"] = claims_identity + if auth is not None: + scope["auth"] = auth + return Request(scope) + + +def test_build_copies_headers_extensions_and_claims_identity(): + identity = ClaimsIdentity({"sub": "agent-1"}, True) + request = _request(claims_identity=identity) + builder = SDKServerCallContextBuilder() + + context = builder.build(request) + + assert isinstance(context.user, UnauthenticatedUser) + assert context.state["headers"]["x-test"] == "value" + assert context.state[_CLAIMS_IDENTITY_KEY] is identity + assert context.requested_extensions == {"extension-1"} + + +def test_build_omits_claims_identity_when_middleware_did_not_set_it(): + request = _request() + builder = SDKServerCallContextBuilder() + + context = builder.build(request) + + assert isinstance(context.user, UnauthenticatedUser) + assert _CLAIMS_IDENTITY_KEY not in context.state + + +def test_build_copies_asgi_auth_scope(): + auth = object() + request = _request(auth=auth) + builder = SDKServerCallContextBuilder() + + context = builder.build(request) + + assert context.state["auth"] is auth diff --git a/tests/hosting_a2a/server/test_utils.py b/tests/hosting_a2a/server/test_utils.py new file mode 100644 index 000000000..1ffd992ff --- /dev/null +++ b/tests/hosting_a2a/server/test_utils.py @@ -0,0 +1,25 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import pytest +from a2a.types import AgentInterface +from a2a.utils.constants import TransportProtocol + +from microsoft_agents.hosting.a2a.server._utils import _get_interface_route_path + + +@pytest.mark.parametrize( + ("url", "expected"), + [ + ("/a2a", "/a2a"), + ("https://example.com/a2a", "/a2a"), + ("https://example.com", "/"), + ], +) +def test_get_interface_route_path(url, expected): + interface = AgentInterface( + url=url, + protocol_binding=TransportProtocol.JSONRPC, + ) + + assert _get_interface_route_path(interface) == expected diff --git a/tests/hosting_a2a/test_a2a_adapter.py b/tests/hosting_a2a/test_a2a_adapter.py new file mode 100644 index 000000000..78496d160 --- /dev/null +++ b/tests/hosting_a2a/test_a2a_adapter.py @@ -0,0 +1,435 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import asyncio +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest +from a2a.server.agent_execution import RequestContext +from a2a.server.context import ServerCallContext +from a2a.server.events import EventQueue, EventQueueLegacy +from a2a.server.tasks import InMemoryTaskStore, TaskStore +from a2a.types import ( + AgentInterface, + AgentSkill, + Message, + Part, + Role, + SendMessageRequest, + TaskArtifactUpdateEvent, + TaskState, + TaskStatusUpdateEvent, +) +from a2a.utils.constants import TransportProtocol +from google.protobuf.json_format import MessageToDict + +from microsoft_agents.activity import ( + Activity, + ActivityTypes, + CallerIdConstants, + Channels, + EndOfConversationCodes, + InputHints, + StreamInfo, +) +from microsoft_agents.hosting.core import ( + AuthenticationConstants, + ChannelServiceAdapter, + ClaimsIdentity, + TurnContext, +) +from microsoft_agents.hosting.a2a import A2AAdapter +from microsoft_agents.hosting.a2a.activity import A2AActivity +from microsoft_agents.hosting.a2a.server._constants import _CLAIMS_IDENTITY_KEY + + +def _agent(): + return SimpleNamespace(on_turn=AsyncMock()) + + +def _message() -> Message: + return Message( + message_id="message-1", + task_id="task-1", + context_id="context-1", + role=Role.ROLE_USER, + parts=[Part(text="hello")], + ) + + +def _request_context(identity=None, *, include_message=True): + call_context = ServerCallContext( + state={} if identity is None else {_CLAIMS_IDENTITY_KEY: identity} + ) + request = SendMessageRequest(message=_message()) if include_message else None + return RequestContext( + call_context, + request=request, + task_id="task-1", + context_id="context-1", + ) + + +def _turn_context(adapter, event_queue): + context = TurnContext( + adapter, + A2AActivity( + type=ActivityTypes.message, + channel_data=_message(), + ), + ClaimsIdentity(), + ) + context.services.set(EventQueue, event_queue) + return context + + +def test_constructor_uses_defaults_and_creates_request_handler(): + adapter = A2AAdapter(_agent()) + + assert adapter.agent_interfaces == [] + assert adapter.skills == [] + assert adapter.a2a_request_handler is not None + + +def test_constructor_preserves_custom_interfaces_and_skills(): + interface = AgentInterface( + url="https://example.com/a2a", + protocol_binding=TransportProtocol.HTTP_JSON, + ) + skill = AgentSkill( + id="weather", + name="Weather", + description="Gets weather", + tags=["weather"], + ) + task_store = InMemoryTaskStore() + + adapter = A2AAdapter( + _agent(), + agent_interfaces=[interface], + skills=[skill], + task_store=task_store, + ) + + assert adapter.agent_interfaces == [interface] + assert adapter.skills == [skill] + + +@pytest.mark.asyncio +async def test_execute_agent_turn_requires_message(): + adapter = A2AAdapter(_agent()) + + with pytest.raises(ValueError, match="Context message is required"): + await adapter.execute_agent_turn( + _request_context(include_message=False), + MagicMock(spec=EventQueue), + ) + + +@pytest.mark.asyncio +async def test_execute_agent_turn_rejects_invalid_identity(): + adapter = A2AAdapter(_agent()) + context = _request_context(identity="invalid") + + with pytest.raises(RuntimeError, match="Invalid identity"): + await adapter.execute_agent_turn(context, MagicMock(spec=EventQueue)) + + +@pytest.mark.asyncio +async def test_execute_agent_turn_converts_message_and_exposes_a2a_services(): + identity = ClaimsIdentity({"sub": "agent-1"}) + context = _request_context(identity) + event_queue = EventQueueLegacy() + task_store = InMemoryTaskStore() + observed = {} + + async def on_turn(turn_context): + observed["context"] = turn_context + + adapter = A2AAdapter( + SimpleNamespace(on_turn=on_turn), + task_store=task_store, + ) + + await adapter.execute_agent_turn(context, event_queue) + + turn_context = observed["context"] + assert turn_context.identity is identity + assert isinstance(turn_context.activity, A2AActivity) + assert turn_context.activity.text == "hello" + assert turn_context.activity.request_id + assert turn_context.services.get(RequestContext) is context + assert turn_context.services.get(EventQueue) is event_queue + assert turn_context.services.get(TaskStore) is task_store + assert ( + turn_context.turn_state[ChannelServiceAdapter.OAUTH_SCOPE_KEY] + == AuthenticationConstants.AGENTS_SDK_SCOPE + ) + + +@pytest.mark.asyncio +async def test_cancel_agent_turn_delivers_user_cancelled_activity(): + identity = ClaimsIdentity({"sub": "agent-1"}) + context = _request_context(identity, include_message=False) + event_queue = EventQueueLegacy() + observed = {} + + async def on_turn(turn_context): + observed["context"] = turn_context + + adapter = A2AAdapter(SimpleNamespace(on_turn=on_turn)) + + await adapter.cancel_agent_turn(context, event_queue) + + turn_context = observed["context"] + activity = turn_context.activity + assert turn_context.identity is identity + assert activity.type == ActivityTypes.end_of_conversation + assert activity.code == EndOfConversationCodes.user_cancelled + assert activity.channel_id == Channels.a2a + assert activity.recipient.id == "assistant" + assert activity.from_property.id == "unknown" + assert turn_context.services.get(RequestContext) is context + assert turn_context.services.get(EventQueue) is event_queue + + +@pytest.mark.asyncio +async def test_cancel_agent_turn_uses_anonymous_identity_by_default(): + context = _request_context(include_message=False) + observed = {} + + async def on_turn(turn_context): + observed["identity"] = turn_context.identity + + adapter = A2AAdapter(SimpleNamespace(on_turn=on_turn)) + + await adapter.cancel_agent_turn(context, EventQueueLegacy()) + + identity = observed["identity"] + assert isinstance(identity, ClaimsIdentity) + assert identity.allow_anonymous is True + + +@pytest.mark.asyncio +async def test_cancel_agent_turn_rejects_invalid_identity(): + context = _request_context(identity="invalid", include_message=False) + adapter = A2AAdapter(_agent()) + + with pytest.raises(RuntimeError, match="Invalid identity"): + await adapter.cancel_agent_turn(context, EventQueueLegacy()) + + +@pytest.mark.asyncio +async def test_execute_agent_turn_sets_agent_caller_and_audience(): + identity = ClaimsIdentity( + { + AuthenticationConstants.VERSION_CLAIM: "2.0", + AuthenticationConstants.AUDIENCE_CLAIM: "target-agent", + AuthenticationConstants.AUTHORIZED_PARTY: "calling-agent", + } + ) + observed = {} + + async def on_turn(turn_context): + observed["context"] = turn_context + + adapter = A2AAdapter(SimpleNamespace(on_turn=on_turn)) + + await adapter.execute_agent_turn( + _request_context(identity), + EventQueueLegacy(), + ) + + context = observed["context"] + assert ( + context.turn_state[ChannelServiceAdapter.OAUTH_SCOPE_KEY] + == "app://calling-agent" + ) + assert context.activity.caller_id == ( + f"{CallerIdConstants.agent_to_agent_prefix}calling-agent" + ) + + +@pytest.mark.asyncio +async def test_send_activities_emits_protocol_events_for_supported_activities(): + adapter = A2AAdapter(_agent()) + event_queue = EventQueueLegacy() + context = _turn_context(adapter, event_queue) + stream_info = StreamInfo( + stream_id="stream-1", + stream_type="content", + stream_sequence=1, + ) + streaming = A2AActivity( + type=ActivityTypes.message, + text="chunk", + entities=[stream_info], + ) + message = A2AActivity( + type=ActivityTypes.message, + text="working", + input_hint=InputHints.expecting_input, + ) + completed = A2AActivity( + type=ActivityTypes.end_of_conversation, + code=EndOfConversationCodes.completed_successfully, + ) + ignored = Activity( + type=ActivityTypes.message, + channel_id=Channels.webchat, + ) + + result = await adapter.send_activities( + context, + [streaming, message, completed, ignored], + ) + + assert result == [] + artifact_event = await event_queue.dequeue_event() + message_event = await event_queue.dequeue_event() + completed_event = await event_queue.dequeue_event() + + assert isinstance(artifact_event, TaskArtifactUpdateEvent) + assert artifact_event.task_id == "task-1" + assert artifact_event.context_id == "context-1" + assert artifact_event.artifact.artifact_id == "stream-1" + assert artifact_event.artifact.parts[0].text == "chunk" + assert artifact_event.last_chunk is True + + assert isinstance(message_event, TaskStatusUpdateEvent) + assert message_event.status.state == TaskState.TASK_STATE_INPUT_REQUIRED + assert message_event.status.message.parts[0].text == "working" + + assert isinstance(completed_event, TaskStatusUpdateEvent) + assert completed_event.status.state == TaskState.TASK_STATE_COMPLETED + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(event_queue.dequeue_event(), timeout=0.05) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("code", "expected_state"), + [ + (EndOfConversationCodes.completed_successfully, TaskState.TASK_STATE_COMPLETED), + (EndOfConversationCodes.error, TaskState.TASK_STATE_FAILED), + (EndOfConversationCodes.user_cancelled, TaskState.TASK_STATE_CANCELED), + ], +) +async def test_end_of_conversation_maps_terminal_state(code, expected_state): + event_queue = EventQueueLegacy() + adapter = A2AAdapter(_agent()) + context = _turn_context(adapter, event_queue) + activity = A2AActivity( + type=ActivityTypes.end_of_conversation, + code=code, + ) + + await adapter.send_activities(context, [activity]) + + event = await event_queue.dequeue_event() + assert isinstance(event, TaskStatusUpdateEvent) + assert event.status.state == expected_state + + +@pytest.mark.asyncio +async def test_end_of_conversation_emits_result_artifact_before_status_message(): + event_queue = EventQueueLegacy() + adapter = A2AAdapter(_agent()) + context = _turn_context(adapter, event_queue) + activity = A2AActivity( + type=ActivityTypes.end_of_conversation, + code=EndOfConversationCodes.completed_successfully, + text="Completed with structured data", + value={"answer": 42}, + ) + + await adapter.send_activities(context, [activity]) + + artifact_event = await event_queue.dequeue_event() + status_event = await event_queue.dequeue_event() + + assert isinstance(artifact_event, TaskArtifactUpdateEvent) + assert artifact_event.last_chunk is True + assert artifact_event.artifact.name == "Result" + assert MessageToDict(artifact_event.artifact.parts[0].data) == {"answer": 42.0} + + assert isinstance(status_event, TaskStatusUpdateEvent) + assert status_event.status.state == TaskState.TASK_STATE_COMPLETED + assert status_event.status.message.parts[0].text == ( + "Completed with structured data" + ) + assert all(not part.HasField("data") for part in status_event.status.message.parts) + + +@pytest.mark.asyncio +async def test_get_agent_card_projects_interfaces_and_skills(): + interface = AgentInterface( + url="https://example.com/a2a", + protocol_binding=TransportProtocol.HTTP_JSON, + ) + skill = AgentSkill( + id="weather", + name="Weather", + description="Gets weather", + tags=["weather"], + examples=["Weather in Seattle"], + input_modes=["text/plain"], + output_modes=["application/json"], + ) + adapter = A2AAdapter( + _agent(), + agent_card_name="Test agent", + agent_card_description="Description", + agent_card_version="1.2.3", + agent_interfaces=[interface], + skills=[skill], + ) + + card = await adapter.get_agent_card( + SimpleNamespace(url="https://host.example/a2a/card"), + "/a2a", + ) + + assert card.name == "Test agent" + assert card.description == "Description" + assert card.version == "1.2.3" + assert card.supported_interfaces[0].url == "https://example.com/a2a" + assert card.supported_interfaces[0].protocol_binding == TransportProtocol.HTTP_JSON + assert card.supported_interfaces[0].protocol_version == "1.0" + assert card.skills[0].id == "weather" + assert card.skills[0].examples == ["Weather in Seattle"] + + +@pytest.mark.asyncio +async def test_get_agent_card_resolves_relative_interface_against_request_origin(): + adapter = A2AAdapter( + _agent(), + agent_interfaces=[ + AgentInterface( + url="/a2a", + protocol_binding=TransportProtocol.JSONRPC, + ) + ], + ) + + card = await adapter.get_agent_card( + SimpleNamespace(url="http://127.0.0.1:41241/a2a/.well-known/agent-card.json"), + "/a2a", + ) + + assert card.supported_interfaces[0].url == "http://127.0.0.1:41241/a2a" + + +@pytest.mark.asyncio +async def test_update_and_delete_activity_are_not_supported(): + adapter = A2AAdapter(_agent()) + + with pytest.raises(NotImplementedError): + await adapter.update_activity( + MagicMock(spec=TurnContext), + Activity(type=ActivityTypes.message), + ) + + with pytest.raises(NotImplementedError): + await adapter.delete_activity(MagicMock(spec=TurnContext), "activity-1") diff --git a/tests/hosting_a2a/test_a2a_agent_extension.py b/tests/hosting_a2a/test_a2a_agent_extension.py new file mode 100644 index 000000000..60a4b8f86 --- /dev/null +++ b/tests/hosting_a2a/test_a2a_agent_extension.py @@ -0,0 +1,92 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from unittest.mock import AsyncMock + +import pytest +from a2a.server.agent_execution import RequestContext +from a2a.server.context import ServerCallContext +from a2a.server.events import EventQueue, EventQueueLegacy +from a2a.server.tasks import InMemoryTaskStore, TaskStore + +from microsoft_agents.activity import Activity +from microsoft_agents.hosting.core import ChannelAdapter, ClaimsIdentity, TurnContext +from microsoft_agents.hosting.a2a.a2a_agent_extension import A2AAgentExtension +from microsoft_agents.hosting.a2a.extension.a2a_turn_context import A2ATurnContext + + +class _Adapter(ChannelAdapter): + async def send_activities(self, context, activities): + return [] + + async def update_activity(self, context, activity): + return None + + async def delete_activity(self, context, reference): + return None + + +class _RecordingApplication: + def message(self, select, *, auth_handlers=None, **kwargs): + self.registration = (select, auth_handlers, kwargs) + + def register(handler): + self.handler = handler + return handler + + return register + + +def _turn_context(): + context = TurnContext( + _Adapter(), + Activity(type="message", text="hello"), + ClaimsIdentity(), + ) + context.services.set(EventQueue, EventQueueLegacy()) + context.services.set( + RequestContext, + RequestContext( + ServerCallContext(), + task_id="task-1", + context_id="context-1", + ), + ) + context.services.set(TaskStore, InMemoryTaskStore()) + return context + + +@pytest.mark.asyncio +async def test_message_handler_receives_a2a_context_and_registration_options(): + app = _RecordingApplication() + extension = A2AAgentExtension(app) + handler = AsyncMock() + + registered_handler = extension.message( + ["hello", "help"], + auth_handlers=["auth"], + rank=10, + )(handler) + state = object() + + await registered_handler(_turn_context(), state) + + assert app.registration == (["hello", "help"], ["auth"], {"rank": 10}) + context = handler.call_args.args[0] + assert isinstance(context, A2ATurnContext) + assert context.activity.text == "hello" + assert handler.call_args.args[1] is state + + +@pytest.mark.asyncio +async def test_message_handler_reuses_existing_a2a_context(): + app = _RecordingApplication() + extension = A2AAgentExtension(app) + handler = AsyncMock() + registered_handler = extension.message([])(handler) + context = A2ATurnContext.from_existing(_turn_context(), app) + state = object() + + await registered_handler(context, state) + + assert handler.call_args.args == (context, state) diff --git a/tests/hosting_a2a/test_a2a_client.py b/tests/hosting_a2a/test_a2a_client.py new file mode 100644 index 000000000..1c0ddb103 --- /dev/null +++ b/tests/hosting_a2a/test_a2a_client.py @@ -0,0 +1,53 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from types import SimpleNamespace + +import pytest +from a2a.server.agent_execution import RequestContext +from a2a.server.context import ServerCallContext +from a2a.server.events import EventQueue, EventQueueLegacy +from a2a.server.tasks import InMemoryTaskStore, TaskStore + +from microsoft_agents.hosting.a2a.extension.a2a_client import A2AClient + + +class _Services: + def __init__(self, values): + self._values = values + + def get(self, key): + return self._values.get(key) + + +def _services(): + event_queue = EventQueueLegacy() + request_context = RequestContext( + ServerCallContext(), + task_id="task-1", + context_id="context-1", + ) + task_store = InMemoryTaskStore() + return { + EventQueue: event_queue, + RequestContext: request_context, + TaskStore: task_store, + } + + +def test_client_exposes_services_from_turn_context(): + services = _services() + client = A2AClient(SimpleNamespace(services=_Services(services))) + + assert client.event_queue is services[EventQueue] + assert client.request_context is services[RequestContext] + assert client.task_store is services[TaskStore] + + +@pytest.mark.parametrize("missing_service", [EventQueue, RequestContext, TaskStore]) +def test_client_requires_all_services(missing_service): + services = _services() + del services[missing_service] + + with pytest.raises(ValueError, match="Missing required services"): + A2AClient(SimpleNamespace(services=_Services(services))) diff --git a/tests/hosting_a2a/test_a2a_turn_context.py b/tests/hosting_a2a/test_a2a_turn_context.py new file mode 100644 index 000000000..3c72ef3c0 --- /dev/null +++ b/tests/hosting_a2a/test_a2a_turn_context.py @@ -0,0 +1,131 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from unittest.mock import MagicMock + +from a2a.server.agent_execution import RequestContext +from a2a.server.context import ServerCallContext +from a2a.server.events import EventQueue, EventQueueLegacy +from a2a.server.tasks import InMemoryTaskStore, TaskStore + +from microsoft_agents.activity import Activity, ResourceResponse +from microsoft_agents.hosting.core import ( + AgentApplication, + ChannelAdapter, + ClaimsIdentity, + TurnContext, +) +from microsoft_agents.hosting.a2a.activity import A2AActivity +from microsoft_agents.hosting.a2a.extension.a2a_turn_context import A2ATurnContext + + +class _Adapter(ChannelAdapter): + async def send_activities(self, context, activities): + return [ResourceResponse()] * len(activities) + + async def update_activity(self, context, activity): + return ResourceResponse(id=activity.id) + + async def delete_activity(self, context, reference): + return None + + +def _original_context(): + original = TurnContext( + _Adapter(), + Activity(type="message", text="hello"), + ClaimsIdentity(), + ) + original.services.set(EventQueue, EventQueueLegacy()) + original.services.set( + RequestContext, + RequestContext( + ServerCallContext(), + task_id="task-1", + context_id="context-1", + ), + ) + original.services.set(TaskStore, InMemoryTaskStore()) + return original + + +def test_wraps_existing_context_and_exposes_a2a_client(): + original = _original_context() + context = A2ATurnContext.from_existing( + original, + MagicMock(spec=AgentApplication), + ) + + assert isinstance(context.activity, A2AActivity) + assert context.activity.text == "hello" + assert context.client.event_queue is original.services.get(EventQueue) + assert context.client.request_context is original.services.get(RequestContext) + assert context.client.task_store is original.services.get(TaskStore) + + +def test_wrapped_context_preserves_original_turn_state(): + original = _original_context() + marker = object() + original.turn_state["marker"] = marker + + context = A2ATurnContext.from_existing( + original, + MagicMock(spec=AgentApplication), + ) + + assert context.turn_state is original.turn_state + assert context.turn_state["marker"] is marker + + +def test_can_be_constructed_directly_from_adapter_activity_and_identity(): + adapter = _Adapter() + activity = A2AActivity(type="message", text="hello") + identity = ClaimsIdentity({"sub": "agent-1"}) + request_context = RequestContext( + ServerCallContext(), + task_id="task-1", + context_id="context-1", + ) + event_queue = EventQueueLegacy() + task_store = InMemoryTaskStore() + + context = A2ATurnContext( + adapter, + MagicMock(spec=AgentApplication), + activity, + identity, + request_context=request_context, + event_queue=event_queue, + task_store=task_store, + ) + + assert context.adapter is adapter + assert context.activity is activity + assert isinstance(context.activity, A2AActivity) + assert context.identity is identity + assert context.client.request_context is request_context + assert context.client.event_queue is event_queue + assert context.client.task_store is task_store + + +def test_responded_property_is_forwarded_to_original_context(): + original = _original_context() + context = A2ATurnContext.from_existing( + original, + MagicMock(spec=AgentApplication), + ) + + context.responded = True + + assert original.responded is True + assert context.responded is True + + +def test_streaming_response_is_forwarded_to_original_context(): + original = _original_context() + context = A2ATurnContext.from_existing( + original, + MagicMock(spec=AgentApplication), + ) + + assert context.streaming_response is original.streaming_response diff --git a/tests/hosting_a2a/test_add_a2a.py b/tests/hosting_a2a/test_add_a2a.py new file mode 100644 index 000000000..f97fdbccc --- /dev/null +++ b/tests/hosting_a2a/test_add_a2a.py @@ -0,0 +1,211 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from a2a.types import AgentCapabilities, AgentCard, AgentInterface +from a2a.utils.constants import TransportProtocol +from fastapi import FastAPI + +from microsoft_agents.hosting.core import AgentAuthConfiguration +from microsoft_agents.hosting.core.http import HttpResponse +from microsoft_agents.hosting.a2a import add_a2a as exported_add_a2a +from microsoft_agents.hosting.a2a.add_a2a import add_a2a + + +def _adapter(*interfaces): + return SimpleNamespace( + agent_interfaces=list(interfaces), + a2a_request_handler=MagicMock(), + get_agent_card=AsyncMock( + return_value=AgentCard( + name="Test agent", + description="Test agent description", + version="1.0.0", + supported_interfaces=[], + capabilities=AgentCapabilities(), + ) + ), + ) + + +def test_add_a2a_uses_default_jsonrpc_interface(): + app = FastAPI() + + add_a2a( + app, + MagicMock(), + use_jwt_middleware=False, + ) + + paths = {route.path for route in app.routes} + assert "/a2a" in paths + assert "/a2a/.well-known/agent-card.json" in paths + + +@pytest.mark.asyncio +async def test_add_a2a_exposes_configured_jsonrpc_and_http_interfaces(): + app = FastAPI() + adapter = _adapter( + AgentInterface( + url="/rpc", + protocol_binding=TransportProtocol.JSONRPC, + ), + AgentInterface( + url="/rest", + protocol_binding=TransportProtocol.HTTP_JSON, + ), + ) + add_a2a(app, MagicMock(), adapter, use_jwt_middleware=False) + + paths = {route.path for route in app.routes} + assert { + "/rpc", + "/rpc/.well-known/agent-card.json", + "/rest/message:send", + "/rest/tasks", + "/rest/.well-known/agent-card.json", + }.issubset(paths) + + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url="http://testserver", + ) as client: + rpc_card = await client.get("/rpc/.well-known/agent-card.json") + rest_card = await client.get("/rest/.well-known/agent-card.json") + + assert rpc_card.status_code == 200 + assert rest_card.status_code == 200 + assert rpc_card.json()["name"] == "Test agent" + assert rest_card.json()["name"] == "Test agent" + assert rpc_card.headers["cache-control"] == "no-store" + assert rest_card.headers["cache-control"] == "no-store" + assert { + "/rpc", + "/rpc/.well-known/agent-card.json", + "/rest/message:send", + "/rest/tasks", + "/rest/.well-known/agent-card.json", + }.issubset(app.openapi()["paths"]) + + +@pytest.mark.asyncio +async def test_add_a2a_enables_agent_card_caching_for_all_interfaces(): + app = FastAPI() + adapter = _adapter( + AgentInterface( + url="/rpc", + protocol_binding=TransportProtocol.JSONRPC, + ), + AgentInterface( + url="/rest", + protocol_binding=TransportProtocol.HTTP_JSON, + ), + ) + add_a2a( + app, + MagicMock(), + adapter, + use_jwt_middleware=False, + _agent_card_cache_enabled=True, + agent_card_cache_max_age=300, + ) + + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url="http://testserver", + ) as client: + rpc_card = await client.get("/rpc/.well-known/agent-card.json") + rest_card = await client.get("/rest/.well-known/agent-card.json") + + for response in (rpc_card, rest_card): + assert response.headers["cache-control"] == "public, max-age=300" + assert response.headers["etag"] + assert response.headers["last-modified"] + + +@pytest.mark.asyncio +async def test_add_a2a_applies_jwt_middleware_to_registered_routes(): + app = FastAPI() + app.state.agent_configuration = AgentAuthConfiguration() + + @app.get("/health") + async def health(): + return {"status": "ok"} + + adapter = _adapter( + AgentInterface( + url="/rpc", + protocol_binding=TransportProtocol.JSONRPC, + ) + ) + add_a2a(app, MagicMock(), adapter) + + with patch( + "microsoft_agents.hosting.fastapi.jwt_authorization_middleware." + "_authorize_request", + new=AsyncMock( + return_value=HttpResponse( + body={"error": "Authentication required"}, + status_code=401, + ) + ), + ) as authorize: + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url="http://testserver", + ) as client: + health_response = await client.get("/health") + card_response = await client.get("/rpc/.well-known/agent-card.json") + + assert health_response.status_code == 200 + assert health_response.json() == {"status": "ok"} + assert card_response.status_code == 401 + assert card_response.json() == {"error": "Authentication required"} + authorize.assert_awaited_once() + + +def test_add_a2a_rejects_adapter_without_interfaces(): + with pytest.raises(ValueError, match="No agent interfaces found"): + add_a2a(FastAPI(), MagicMock(), _adapter()) + + +def test_add_a2a_rejects_adapter_without_supported_interfaces(): + adapter = _adapter( + AgentInterface( + url="/grpc", + protocol_binding=TransportProtocol.GRPC, + ) + ) + + with pytest.raises(ValueError, match="unsupported protocol bindings"): + add_a2a(FastAPI(), MagicMock(), adapter, use_jwt_middleware=False) + + +def test_add_a2a_ignores_unsupported_interfaces_when_supported_ones_exist(): + app = FastAPI() + adapter = _adapter( + AgentInterface( + url="/grpc", + protocol_binding=TransportProtocol.GRPC, + ), + AgentInterface( + url="/rpc", + protocol_binding=TransportProtocol.JSONRPC, + ), + ) + + add_a2a(app, MagicMock(), adapter, use_jwt_middleware=False) + + paths = {route.path for route in app.routes} + assert "/rpc" in paths + assert "/rpc/.well-known/agent-card.json" in paths + assert "/grpc" not in paths + assert "/grpc/.well-known/agent-card.json" not in paths + + +def test_add_a2a_is_exported_from_package(): + assert exported_add_a2a is add_a2a diff --git a/tests/hosting_a2a/test_blob_task_store.py b/tests/hosting_a2a/test_blob_task_store.py new file mode 100644 index 000000000..782c8b2f0 --- /dev/null +++ b/tests/hosting_a2a/test_blob_task_store.py @@ -0,0 +1,237 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import asyncio +import os +import uuid +from contextlib import asynccontextmanager + +import pytest +from a2a.server.context import ServerCallContext +from a2a.types import ListTasksRequest, Task, TaskState, TaskStatus +from azure.core.exceptions import ResourceNotFoundError +from azure.identity.aio import DefaultAzureCredential +from azure.storage.blob.aio import BlobServiceClient, ContainerClient +from dotenv import load_dotenv +from google.protobuf.message import DecodeError + +from microsoft_agents.hosting.a2a.blob_task_store import BlobTaskStore + +# To enable these tests, run with --run-blob and configure either: +# TEST_BLOB_STORAGE_CONNECTION_STRING or TEST_BLOB_STORAGE_ACCOUNT_URL. + + +def _task( + task_id: str, + *, + context_id: str = "context-1", + state: TaskState = TaskState.TASK_STATE_WORKING, +) -> Task: + return Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=state), + ) + + +@asynccontextmanager +async def _blob_task_store(): + load_dotenv() + connection_string = os.environ.get("TEST_BLOB_STORAGE_CONNECTION_STRING") + credential = None + + if connection_string: + service_client = BlobServiceClient.from_connection_string(connection_string) + else: + account_url = os.environ.get("TEST_BLOB_STORAGE_ACCOUNT_URL") + if not account_url: + pytest.skip( + "Set TEST_BLOB_STORAGE_CONNECTION_STRING or " + "TEST_BLOB_STORAGE_ACCOUNT_URL" + ) + credential = DefaultAzureCredential() + service_client = BlobServiceClient(account_url, credential=credential) + + container_name = f"asdka2atasks{uuid.uuid4().hex}" + container_client = service_client.get_container_client(container_name) + store = BlobTaskStore(container_client) + + try: + yield store, container_client + finally: + try: + await container_client.delete_container() + except ResourceNotFoundError: + pass + await container_client.close() + await service_client.close() + if credential: + await credential.close() + + +async def _upload_blob(container_client: ContainerClient, name: str, data: bytes): + blob_client = await container_client.upload_blob( + name=name, + data=data, + overwrite=True, + ) + await blob_client.close() + + +@pytest.mark.parametrize( + "kwargs", + [ + {}, + {"data_connection_string": "connection"}, + {"container_name": "tasks"}, + ], +) +def test_constructor_rejects_invalid_parameter_combinations(kwargs): + with pytest.raises(ValueError, match="Invalid combination"): + BlobTaskStore(**kwargs) + + +def test_blob_name_uses_a2a_task_namespace_and_url_encoding(): + assert ( + BlobTaskStore._get_blob_name("task/with spaces") == "a2atask/task%2Fwith+spaces" + ) + + +@pytest.mark.blob +class TestBlobTaskStore: + @pytest.mark.asyncio + async def test_save_get_and_overwrite_task(self): + context = ServerCallContext() + + async with _blob_task_store() as (store, container_client): + task = _task("task/with spaces") + + await store.save(task, context) + + saved = await store.get(task.id, context) + assert saved == task + assert saved is not task + + blob_names = [blob.name async for blob in container_client.list_blobs()] + assert blob_names == ["a2atask/task%2Fwith+spaces"] + + task.status.state = TaskState.TASK_STATE_COMPLETED + await store.save(task, context) + + overwritten = await store.get(task.id, context) + assert overwritten == task + assert overwritten.status.state == TaskState.TASK_STATE_COMPLETED + + @pytest.mark.asyncio + async def test_external_blob_change_is_visible(self): + context = ServerCallContext() + external_task = _task("external-task") + + async with _blob_task_store() as (store, container_client): + assert await store.get(external_task.id, context) is None + + await _upload_blob( + container_client, + "a2atask/external-task", + external_task.SerializeToString(), + ) + + assert await store.get(external_task.id, context) == external_task + + @pytest.mark.asyncio + async def test_get_returns_none_for_missing_task(self): + async with _blob_task_store() as (store, _): + assert await store.get("missing-task", ServerCallContext()) is None + + @pytest.mark.asyncio + async def test_get_rejects_corrupted_task_blob(self): + context = ServerCallContext() + + async with _blob_task_store() as (store, container_client): + assert await store.get("corrupted-task", context) is None + await _upload_blob( + container_client, + "a2atask/corrupted-task", + b"not a serialized Task", + ) + + with pytest.raises(DecodeError): + await store.get("corrupted-task", context) + + @pytest.mark.asyncio + async def test_list_filters_by_status_and_context(self): + context = ServerCallContext() + matching = _task("task-1") + tasks = [ + matching, + _task("task-2", context_id="other-context"), + _task("task-3", state=TaskState.TASK_STATE_COMPLETED), + ] + + async with _blob_task_store() as (store, _): + await asyncio.gather(*(store.save(task, context) for task in tasks)) + + response = await store.list( + ListTasksRequest( + status=TaskState.TASK_STATE_WORKING, + context_id="context-1", + ), + context, + ) + + assert list(response.tasks) == [matching] + assert response.next_page_token == "" + + @pytest.mark.asyncio + async def test_list_uses_azure_continuation_tokens(self): + context = ServerCallContext() + tasks = [_task(f"task-{index}") for index in range(3)] + + async with _blob_task_store() as (store, _): + await asyncio.gather(*(store.save(task, context) for task in tasks)) + + first_page = await store.list( + ListTasksRequest(page_size=2), + context, + ) + second_page = await store.list( + ListTasksRequest( + page_size=2, + page_token=first_page.next_page_token, + ), + context, + ) + + assert [task.id for task in first_page.tasks] == [ + "task-0", + "task-1", + ] + assert first_page.next_page_token + assert [task.id for task in second_page.tasks] == ["task-2"] + assert second_page.next_page_token == "" + + @pytest.mark.asyncio + async def test_concurrent_first_operations_share_container_initialization(self): + context = ServerCallContext() + tasks = [_task(f"task-{index}") for index in range(3)] + + async with _blob_task_store() as (store, _): + await asyncio.gather(*(store.save(task, context) for task in tasks)) + saved = await asyncio.gather( + *(store.get(task.id, context) for task in tasks) + ) + + assert saved == tasks + + @pytest.mark.asyncio + async def test_delete_removes_persisted_task(self): + context = ServerCallContext() + task = _task("task/to-delete") + + async with _blob_task_store() as (store, _): + await store.save(task, context) + assert await store.get(task.id, context) == task + + await store.delete(task.id, context) + + assert await store.get(task.id, context) is None diff --git a/tests/hosting_core/test_service_set.py b/tests/hosting_core/test_service_set.py index 1e720091c..ce6e2b01b 100644 --- a/tests/hosting_core/test_service_set.py +++ b/tests/hosting_core/test_service_set.py @@ -31,6 +31,13 @@ def test_get_returns_none_for_missing_service(): assert services.get(Service) is None +def test_get_raises_for_missing_service_when_requested(): + services = _ServiceSet() + + with pytest.raises(KeyError, match="Service.*missing"): + services.get(Service, raise_if_missing=True) + + def test_has_returns_false_for_missing_service(): services = _ServiceSet()