Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 40 additions & 1 deletion src/bedrock_agentcore/runtime/a2a.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@
extraction, health checks, and Docker host detection.
"""

import asyncio
import contextvars
import logging
import uuid
from importlib import import_module
Expand Down Expand Up @@ -34,6 +36,8 @@
A2A_CONTRACT_PORT = 9000
A2A_PORT_ENV = "A2A_PORT"

_CONTEXTVARS_STATE_KEY = "_bedrock_agentcore_contextvars"


def _check_a2a_sdk() -> None:
"""Raise ImportError with install instructions if a2a-sdk is missing."""
Expand Down Expand Up @@ -247,6 +251,37 @@ def build(self, request: Any) -> Any:
pass


class _ContextvarsSnapshotRequestContextBuilder:
"""Stores a copy of the sending request's contextvars on each RequestContext.

a2a-sdk 1.x runs every message of a task in one producer task created by the
first request, so the executor would otherwise see that request's contextvars.
Remove with ``_PerMessageContextvarsExecutor`` once a2a-python#1316 is fixed.
"""

def __init__(self, inner: Any) -> None:
self._inner = inner

async def build(self, *args: Any, **kwargs: Any) -> Any:
request_context = await self._inner.build(*args, **kwargs)
request_context.call_context.state[_CONTEXTVARS_STATE_KEY] = contextvars.copy_context()
return request_context


class _PerMessageContextvarsExecutor:
"""Runs the wrapped executor in the contextvars of the request that sent the message."""

def __init__(self, inner: Any) -> None:
self._inner = inner

async def execute(self, context: Any, event_queue: Any) -> None:
snapshot = context.call_context.state[_CONTEXTVARS_STATE_KEY]
await snapshot.run(asyncio.create_task, self._inner.execute(context, event_queue))

async def cancel(self, context: Any, event_queue: Any) -> None:
await self._inner.cancel(context, event_queue)


def build_a2a_app(
executor: Any,
agent_card: Any = None,
Expand Down Expand Up @@ -300,12 +335,16 @@ def build_a2a_app(

request_handler_type: Any = DefaultRequestHandler
if is_a2a_v1:
from a2a.server.agent_execution import SimpleRequestContextBuilder
from a2a.server.routes import create_agent_card_routes, create_jsonrpc_routes

http_handler = request_handler_type(
agent_executor=executor,
agent_executor=_PerMessageContextvarsExecutor(executor),
task_store=task_store,
agent_card=agent_card,
request_context_builder=_ContextvarsSnapshotRequestContextBuilder(
SimpleRequestContextBuilder(should_populate_referred_tasks=False, task_store=task_store)
),
)
routes = create_agent_card_routes(agent_card)
routes.extend(
Expand Down
82 changes: 72 additions & 10 deletions tests/bedrock_agentcore/runtime/test_a2a.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,42 @@ async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None
pass


_REQUEST_TAG: contextvars.ContextVar[str] = contextvars.ContextVar("request_tag", default="unset")


class _RequestTagMiddleware:
def __init__(self, app):
self.app = app

async def __call__(self, scope, receive, send):
token = _REQUEST_TAG.set(dict(scope.get("headers", [])).get(b"x-request-tag", b"unset").decode())
try:
await self.app(scope, receive, send)
finally:
_REQUEST_TAG.reset(token)


class _InputRequiredExecutor(AgentExecutor):
"""Asks for input on the first message of a task and completes on the follow-up."""

def __init__(self):
self.seen = []

async def execute(self, context: RequestContext, event_queue: EventQueue) -> None:
token = BedrockAgentCoreContext.get_workload_access_token()
self.seen.append((context.task_id, context.context_id, _REQUEST_TAG.get(), token))
task = context.current_task or _new_task(context.message)
updater = TaskUpdater(event_queue, task.id, task.context_id)
if context.current_task:
await updater.complete()
else:
await event_queue.enqueue_event(task)
await updater.requires_input()

async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None:
pass


def _make_agent_card() -> AgentCard:
card_kwargs = {
"name": "test-agent",
Expand Down Expand Up @@ -116,22 +152,23 @@ def _jsonrpc_request(method: str, params: dict | None = None) -> dict:
return body


def _send_message_params(text: str = "hello") -> dict:
def _send_message_params(text: str = "hello", task_id: str | None = None, context_id: str | None = None) -> dict:
if not IS_A2A_V1:
return {
"message": {
"message_id": str(uuid.uuid4()),
"role": "user",
"parts": [{"kind": "text", "text": text}],
}
message = {
"message_id": str(uuid.uuid4()),
"role": "user",
"parts": [{"kind": "text", "text": text}],
}
return {
"message": {
else:
message = {
"messageId": str(uuid.uuid4()),
"role": "ROLE_USER",
"parts": [{"text": text}],
}
}
if task_id:
message["taskId"] = task_id
message["contextId"] = context_id
return {"message": message}


class TestBuildA2AApp:
Expand Down Expand Up @@ -194,6 +231,31 @@ def test_message_send_executes_and_returns_completed_task(self):
assert task["status"]["state"] == _completed_state()
assert task["artifacts"][0]["parts"][0]["text"] == "echo: unit-test"

@pytest.mark.parametrize(
"method",
[_method("SendMessage", "message/send"), _method("SendStreamingMessage", "message/stream")],
)
def test_follow_up_message_runs_in_its_own_request_contextvars(self, method):
"""Regression for #690: a2a-sdk 1.x reuses the first request's producer task for every message of a task."""
executor = _InputRequiredExecutor()
app = build_a2a_app(executor, _make_agent_card())
app.add_middleware(_RequestTagMiddleware)

def send(tag, **ids):
resp = client.post(
"/",
json=_jsonrpc_request(method, _send_message_params(tag, **ids)),
headers={"x-request-tag": tag, "WorkloadAccessToken": f"wat-{tag}"},
)
assert resp.status_code == 200

with _test_client(app) as client:
send("turn-1")
task_id, context_id = executor.seen[0][:2]
send("turn-2", task_id=task_id, context_id=context_id)

assert [row[2:] for row in executor.seen] == [("turn-1", "wat-turn-1"), ("turn-2", "wat-turn-2")]

@pytest.mark.skipif(not IS_A2A_V1, reason="v0.3 compatibility routes are provided by a2a-sdk v1")
def test_v03_message_send_remains_compatible(self):
app = build_a2a_app(_EchoExecutor(), _make_agent_card())
Expand Down
Loading