Skip to content
Merged
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
27 changes: 20 additions & 7 deletions src/a2a/server/agent_execution/active_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@
from __future__ import annotations

import asyncio
import contextvars
import logging
import uuid

Expand Down Expand Up @@ -468,9 +469,9 @@ def __init__( # noqa: PLR0913
self._reference_count = 0

# Queue for incoming requests
self._request_queue: AsyncQueue[tuple[RequestContext, uuid.UUID]] = (
create_async_queue()
)
self._request_queue: AsyncQueue[
tuple[RequestContext, uuid.UUID, contextvars.Context]
] = create_async_queue()

@property
def task_id(self) -> str:
Expand All @@ -480,9 +481,16 @@ def task_id(self) -> str:
async def enqueue_request(
self, request_context: RequestContext
) -> uuid.UUID:
"""Enqueues a request for the active task to process."""
"""Enqueues a request for the active task to process.

The caller's contextvars are captured so the producer runs
`AgentExecutor.execute` in the context of the request that sent the
message, not the request that started the producer.
"""
request_id = uuid.uuid4()
await self._request_queue.put((request_context, request_id))
await self._request_queue.put(
(request_context, request_id, contextvars.copy_context())
)
return request_id

async def start(
Expand Down Expand Up @@ -573,6 +581,7 @@ async def _run_producer(self) -> None:
(
request_context,
request_id,
sender_context,
) = await self._request_queue.get()
await self._request_lock.acquire()
# TODO: Should we create task manager every time?
Expand Down Expand Up @@ -607,8 +616,12 @@ async def _run_producer(self) -> None:
_RequestStarted(request_id, request_context),
)
)
await self._agent_executor.execute(
request_context, self._event_queue_agent
# Awaiting the child task propagates producer cancellation.
await sender_context.run(
asyncio.create_task,
self._agent_executor.execute(
request_context, self._event_queue_agent
),
)
logger.debug(
'Producer[%s]: Execution finished successfully',
Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import asyncio
import contextvars
import logging
import time
import uuid
Expand Down Expand Up @@ -2443,3 +2444,87 @@ async def test_cancel_of_input_required_task_cannot_be_undone():
assert stored.status.state == TaskState.TASK_STATE_CANCELED
assert agent.execute_calls == 1
await handler.aclose()


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


class _InputRequiredThenCompleteAgent(AgentExecutor):
"""Asks for input on the first message of a task and completes on the next."""

def __init__(self) -> None:
self.seen_tags: list[str] = []
self.task: Task | None = None

async def execute(
self, context: RequestContext, event_queue: EventQueue
) -> None:
self.seen_tags.append(_REQUEST_TAG.get())
if context.current_task:
updater = TaskUpdater(
event_queue, context.task_id, context.context_id
)
await updater.complete()
return
self.task = new_task_from_user_message(context.message)
await event_queue.enqueue_event(self.task)
updater = TaskUpdater(event_queue, self.task.id, self.task.context_id)
await updater.update_status(TaskState.TASK_STATE_INPUT_REQUIRED)

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


@pytest.mark.timeout(10)
@pytest.mark.asyncio
@pytest.mark.parametrize('streaming', [False, True])
async def test_follow_up_message_runs_in_sender_contextvars(streaming):
"""A follow-up message on a live task must see its own request's contextvars."""
agent = _InputRequiredThenCompleteAgent()
handler = DefaultRequestHandlerV2(
agent_executor=agent,
task_store=InMemoryTaskStore(),
agent_card=create_default_agent_card(),
)

async def send(tag: str, message: Message) -> None:
_REQUEST_TAG.set(tag)
params = SendMessageRequest(message=message)
if streaming:
async for _ in handler.on_message_send_stream(
params, create_server_call_context()
):
pass
else:
await handler.on_message_send(params, create_server_call_context())

await asyncio.create_task(
send(
'request-1',
Message(
message_id='msg-1',
role=Role.ROLE_USER,
parts=[Part(text='book a flight')],
),
)
)
assert agent.task is not None
await asyncio.create_task(
send(
'request-2',
Message(
message_id='msg-2',
role=Role.ROLE_USER,
parts=[Part(text='Friday')],
task_id=agent.task.id,
context_id=agent.task.context_id,
),
)
)

assert agent.seen_tags == ['request-1', 'request-2']
await handler.aclose()
Loading