diff --git a/src/a2a/server/agent_execution/active_task.py b/src/a2a/server/agent_execution/active_task.py index ea1955fdd..6d3c49967 100644 --- a/src/a2a/server/agent_execution/active_task.py +++ b/src/a2a/server/agent_execution/active_task.py @@ -391,6 +391,8 @@ def __init__( # `_request_lock` protects parallel request processing. self._request_lock = asyncio.Lock() + # `_snapshot_lock` serializes cache reads with request-boundary refreshes. + self._snapshot_lock = asyncio.Lock() # _task_created is set when initial version of task is stored in DB. self._task_created = asyncio.Event() @@ -419,6 +421,14 @@ def task_id(self) -> str: """The ID of the task.""" return self._task_id + @staticmethod + def _raise_if_task_terminal(task: Task) -> None: + """Rejects operations that would restart a terminal task.""" + if task.status.state in TERMINAL_TASK_STATES: + raise InvalidParamsError( + message=f'Task {task.id} is in terminal state: {task.status.state}' + ) + async def enqueue_request( self, request_context: RequestContext ) -> uuid.UUID: @@ -427,6 +437,24 @@ async def enqueue_request( await self._request_queue.put((request_context, request_id)) return request_id + async def refresh_task_if_idle( + self, call_context: ServerCallContext + ) -> None: + """Drops an idle task snapshot before a new request boundary. + + An ``ActiveTask`` remains registered while a task waits for + human-in-the-loop input, + so another process may persist a newer task snapshot in the meantime. + The request lock, rather than the subscriber count, defines whether + execution is idle. A subscriber may detach before background artifact + streaming finishes, while passive subscribers may remain after the + previous request is fully persisted. + """ + async with self._snapshot_lock, self._lock: + if not self._request_lock.locked(): + self._task_manager._call_context = call_context + self._task_manager._current_task = None + async def start( self, call_context: ServerCallContext, @@ -468,10 +496,7 @@ async def start( if task: self._task_created.set() - if task.status.state in TERMINAL_TASK_STATES: - raise InvalidParamsError( - message=f'Task {task.id} is in terminal state: {task.status.state}' - ) + self._raise_if_task_terminal(task) elif not create_task_if_missing: raise TaskNotFoundError @@ -510,6 +535,7 @@ async def _run_producer(self) -> None: """ logger.debug('Producer[%s]: Started', self._task_id) request_context = None + task_missing_at_boundary = False try: while True: ( @@ -517,13 +543,26 @@ async def _run_producer(self) -> None: request_id, ) = await self._request_queue.get() await self._request_lock.acquire() - # TODO: Should we create task manager every time? - self._task_manager._call_context = request_context.call_context - - request_context.current_task = ( - await self._task_manager.get_task() - ) + # Order task loading after any idle refresh that began before + # this request acquired `_request_lock`. Later refreshes observe + # the held request lock and leave the streaming snapshot intact. + async with self._snapshot_lock: + self._task_manager._call_context = ( + request_context.call_context + ) + # This is the request boundary for queued requests that + # were discovered while the previous request was active. + self._task_manager._current_task = None + request_context.current_task = ( + await self._task_manager.get_task() + ) + if ( + request_context.current_task is None + and self._task_created.is_set() + ): + task_missing_at_boundary = True + raise TaskNotFoundError(f'Task {self._task_id} not found') logger.debug( 'Producer[%s]: Executing agent task %s', self._task_id, @@ -562,7 +601,7 @@ async def _run_producer(self) -> None: ) # Persist the failure directly instead of relying on the closing # event queue to carry a final status update. - if request_context: + if request_context and not task_missing_at_boundary: task = await self._task_manager.ensure_task_id( self._task_id, request_context.context_id or '', @@ -643,6 +682,7 @@ async def subscribe( self._task_id, ) task = await self.get_task() + self._raise_if_task_terminal(task) yield task while True: @@ -715,9 +755,13 @@ async def cancel(self, call_context: ServerCallContext) -> Task: logger.debug('Cancel[%s]: Cancelling task', self._task_id) # TODO: Conflicts with call_context on the pending request. - self._task_manager._call_context = call_context - - task = await self._task_manager.get_task() + async with self._snapshot_lock: + self._task_manager._call_context = call_context + task = await self._task_manager.get_task() + if task is None and self._task_created.is_set(): + raise TaskNotFoundError(f'Task {self._task_id} not found') + if task is not None and task.status.state in TERMINAL_TASK_STATES: + return task request_context = RequestContext( call_context=call_context, task_id=self._task_id, @@ -843,7 +887,8 @@ async def _mark_task_as_failed(self, exception: Exception) -> Task | None: async def get_task(self) -> Task: """Get task from db.""" await self._task_created.wait() - task = await self._task_manager.get_task() + async with self._snapshot_lock: + task = await self._task_manager.get_task() if not task: raise RuntimeError('Task should have been created') return task diff --git a/src/a2a/server/agent_execution/active_task_registry.py b/src/a2a/server/agent_execution/active_task_registry.py index ab7d6a11c..b7d2c33e8 100644 --- a/src/a2a/server/agent_execution/active_task_registry.py +++ b/src/a2a/server/agent_execution/active_task_registry.py @@ -46,28 +46,38 @@ async def get_or_create( initial_message: Message | None = None, ) -> ActiveTask: """Retrieves an existing ActiveTask or creates a new one.""" - async with self._lock: - if self._closed: - raise RuntimeError('ActiveTaskRegistry is closed') - if task_id in self._active_tasks: - return self._active_tasks[task_id] - - task_manager = TaskManager( - task_id=task_id, - context_id=context_id, - task_store=self._task_store, - initial_message=initial_message, - context=call_context, - ) - - active_task = ActiveTask( - agent_executor=self._agent_executor, - task_id=task_id, - task_manager=task_manager, - push_sender=self._push_sender, - on_cleanup=self._on_active_task_cleanup, - ) - self._active_tasks[task_id] = active_task + while True: + async with self._lock: + if self._closed: + raise RuntimeError('ActiveTaskRegistry is closed') + active_task = self._active_tasks.get(task_id) + if active_task is None: + task_manager = TaskManager( + task_id=task_id, + context_id=context_id, + task_store=self._task_store, + initial_message=initial_message, + context=call_context, + ) + + active_task = ActiveTask( + agent_executor=self._agent_executor, + task_id=task_id, + task_manager=task_manager, + push_sender=self._push_sender, + on_cleanup=self._on_active_task_cleanup, + ) + self._active_tasks[task_id] = active_task + break + + # A refresh can wait behind task-store I/O, so do not hold the + # global registry lock while synchronizing this individual task. + await active_task.refresh_task_if_idle(call_context) + async with self._lock: + if self._closed: + raise RuntimeError('ActiveTaskRegistry is closed') + if self._active_tasks.get(task_id) is active_task: + return active_task await active_task.start( call_context=call_context, diff --git a/src/a2a/server/request_handlers/default_request_handler_v2.py b/src/a2a/server/request_handlers/default_request_handler_v2.py index 872a3bfa2..862cdbc81 100644 --- a/src/a2a/server/request_handlers/default_request_handler_v2.py +++ b/src/a2a/server/request_handlers/default_request_handler_v2.py @@ -164,6 +164,13 @@ async def on_cancel_task( # noqa: D102 context: ServerCallContext, ) -> Task | None: task_id = params.id + task = await self.task_store.get(task_id, context) + if not task: + raise TaskNotFoundError(f'Task {task_id} not found') + if task.status.state in TERMINAL_TASK_STATES: + raise TaskNotCancelableError( + message=f'Task cannot be canceled - current state: {task.status.state}' + ) try: active_task = await self._active_task_registry.get_or_create( @@ -203,6 +210,10 @@ async def _setup_active_task( task = await self.task_store.get(original_task_id, call_context) if not task: raise TaskNotFoundError(f'Task {original_task_id} not found') + if task.status.state in TERMINAL_TASK_STATES: + raise InvalidParamsError( + message=f'Task {task.id} is in terminal state: {task.status.state}' + ) # Build context to resolve or generate missing IDs request_context = await self._request_context_builder.build( diff --git a/tests/server/agent_execution/test_active_task.py b/tests/server/agent_execution/test_active_task.py index 1be233ee1..8d1050e23 100644 --- a/tests/server/agent_execution/test_active_task.py +++ b/tests/server/agent_execution/test_active_task.py @@ -27,6 +27,13 @@ logger = logging.getLogger(__name__) +def _working_task() -> Task: + return Task( + id='test-task-id', + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + + class TestActiveTask: """Tests for the ActiveTask class.""" @@ -105,10 +112,9 @@ async def execute_mock(req, q): agent_executor.execute = AsyncMock(side_effect=execute_mock) agent_executor.cancel = AsyncMock() task_manager.get_task.side_effect = [ - Task( - id='test-task-id', - status=TaskStatus(state=TaskState.TASK_STATE_WORKING), - ) + _working_task(), + _working_task(), + _working_task(), ] + [ Task( id='test-task-id', @@ -153,10 +159,8 @@ async def execute_mock(req, q): agent_executor.execute = AsyncMock(side_effect=execute_mock) task_manager.get_task.side_effect = [ - Task( - id='test-task-id', - status=TaskStatus(state=TaskState.TASK_STATE_WORKING), - ) + _working_task(), + _working_task(), ] + [task_obj] * 10 await active_task.start( @@ -200,10 +204,8 @@ async def execute_mock(req, q): agent_executor.execute = AsyncMock(side_effect=execute_mock) task_manager.get_task.side_effect = [ - Task( - id='test-task-id', - status=TaskStatus(state=TaskState.TASK_STATE_WORKING), - ) + _working_task(), + _working_task(), ] + [task_obj] * 10 await active_task.start( @@ -266,10 +268,8 @@ async def execute_mock(req, q): agent_executor.execute = AsyncMock(side_effect=execute_mock) task_manager.get_task.side_effect = [ - Task( - id='test-task-id', - status=TaskStatus(state=TaskState.TASK_STATE_WORKING), - ) + _working_task(), + _working_task(), ] + [task_obj] * 10 await active_task.start( @@ -380,10 +380,8 @@ async def execute_mock(req, q): agent_executor.execute = AsyncMock(side_effect=execute_mock) task_manager.get_task.side_effect = [ - Task( - id='test-task-id', - status=TaskStatus(state=TaskState.TASK_STATE_WORKING), - ) + _working_task(), + _working_task(), ] + [task_obj] * 10 await active_task.start( @@ -517,10 +515,8 @@ async def execute_mock(req, q): agent_executor.execute = AsyncMock(side_effect=execute_mock) task_manager.get_task.side_effect = [ - Task( - id='test-task-id', - status=TaskStatus(state=TaskState.TASK_STATE_WORKING), - ) + _working_task(), + _working_task(), ] + [task_obj] * 10 await active_task.start( @@ -662,10 +658,8 @@ async def execute_mock(req, q): agent_executor.execute = AsyncMock(side_effect=execute_mock) task_manager.get_task.side_effect = [ - Task( - id='test-task-id', - status=TaskStatus(state=TaskState.TASK_STATE_WORKING), - ) + _working_task(), + _working_task(), ] + [task_obj] * 10 await active_task.start( diff --git a/tests/server/agent_execution/test_active_task_registry.py b/tests/server/agent_execution/test_active_task_registry.py index 16d9c8797..9085d3336 100644 --- a/tests/server/agent_execution/test_active_task_registry.py +++ b/tests/server/agent_execution/test_active_task_registry.py @@ -10,7 +10,20 @@ from a2a.server.agent_execution.context import RequestContext from a2a.server.context import ServerCallContext from a2a.server.events.event_queue_v2 import EventQueue -from a2a.server.tasks import InMemoryTaskStore +from a2a.server.tasks import InMemoryTaskStore, TaskStore +from a2a.types.a2a_pb2 import ( + Artifact, + Message, + Part, + Role, + SendMessageRequest, + Task, + TaskArtifactUpdateEvent, + TaskState, + TaskStatus, + TaskStatusUpdateEvent, +) +from a2a.utils.errors import InvalidParamsError, TaskNotFoundError class _SlowExecutor(AgentExecutor): @@ -27,6 +40,161 @@ async def cancel( return None +class _RecordingInputRequiredExecutor(AgentExecutor): + """Records the task snapshot used for a resumed request.""" + + def __init__(self) -> None: + self.seen_tasks: list[Task] = [] + + async def execute( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + if context.current_task is not None: + task = Task() + task.CopyFrom(context.current_task) + self.seen_tasks.append(task) + + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id=context.task_id or '', + context_id=context.context_id or '', + status=TaskStatus( + state=TaskState.TASK_STATE_INPUT_REQUIRED, + ), + ) + ) + + async def cancel( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + return None + + +class _PausedArtifactExecutor(AgentExecutor): + """Keeps an artifact stream open after publishing its first chunk.""" + + def __init__(self) -> None: + self.release_append = asyncio.Event() + + async def execute( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + task_id = context.task_id or '' + context_id = context.context_id or '' + await event_queue.enqueue_event( + TaskArtifactUpdateEvent( + task_id=task_id, + context_id=context_id, + artifact=Artifact( + artifact_id='streamed-artifact', + parts=[Part(text='first chunk')], + ), + ) + ) + await self.release_append.wait() + await event_queue.enqueue_event( + TaskArtifactUpdateEvent( + task_id=task_id, + context_id=context_id, + artifact=Artifact( + artifact_id='streamed-artifact', + parts=[Part(text='second chunk')], + ), + append=True, + ) + ) + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id=task_id, + context_id=context_id, + status=TaskStatus( + state=TaskState.TASK_STATE_INPUT_REQUIRED, + ), + ) + ) + + async def cancel( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + return None + + +class _QueuedResumeExecutor(AgentExecutor): + """Records snapshots for a request queued behind a streaming request.""" + + def __init__(self) -> None: + self.calls = 0 + self.first_started = asyncio.Event() + self.release_first = asyncio.Event() + self.second_started = asyncio.Event() + self.seen_tasks: list[Task] = [] + + async def execute( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + self.calls += 1 + if context.current_task is not None: + task = Task() + task.CopyFrom(context.current_task) + self.seen_tasks.append(task) + + task_id = context.task_id or '' + context_id = context.context_id or '' + if self.calls == 1: + await event_queue.enqueue_event( + TaskArtifactUpdateEvent( + task_id=task_id, + context_id=context_id, + artifact=Artifact( + artifact_id='first-request-artifact', + parts=[Part(text='first request')], + ), + ) + ) + self.first_started.set() + await self.release_first.wait() + else: + self.second_started.set() + + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id=task_id, + context_id=context_id, + status=TaskStatus( + state=TaskState.TASK_STATE_INPUT_REQUIRED, + ), + ) + ) + + async def cancel( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + return None + + +def _request_context( + task_id: str, + context_id: str, + call_context: ServerCallContext, + message_id: str, + text: str, +) -> RequestContext: + return RequestContext( + call_context=call_context, + request=SendMessageRequest( + message=Message( + task_id=task_id, + context_id=context_id, + message_id=message_id, + role=Role.ROLE_USER, + parts=[Part(text=text)], + ) + ), + task_id=task_id, + context_id=context_id, + ) + + def _make_registry() -> ActiveTaskRegistry: return ActiveTaskRegistry( agent_executor=_SlowExecutor(), @@ -105,3 +273,697 @@ async def test_aclose_logs_and_swallows_task_errors(caplog): await registry.aclose() assert 'Error draining active task' in caplog.text + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +@pytest.mark.parametrize( + 'keep_passive_subscriber', + [False, True], + ids=['no-subscriber', 'passive-subscriber'], +) +async def test_reused_idle_active_task_refreshes_shared_store_snapshot( + keep_passive_subscriber: bool, +): + """A resumed request uses data persisted by another registry instance.""" + task_id = 'shared-task' + context_id = 'shared-context' + call_context = ServerCallContext() + task_store = InMemoryTaskStore() + await task_store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus( + state=TaskState.TASK_STATE_INPUT_REQUIRED, + message=Message( + message_id='initial-status', + role=Role.ROLE_AGENT, + parts=[Part(text='initial status')], + ), + ), + ), + call_context, + ) + + replica_a_executor = _RecordingInputRequiredExecutor() + registry_a = ActiveTaskRegistry(replica_a_executor, task_store) + registry_b = ActiveTaskRegistry(_SlowExecutor(), task_store) + passive_stream = None + + try: + active_a = await registry_a.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + active_b = await registry_b.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + if keep_passive_subscriber: + passive_stream = active_a.subscribe(include_initial_task=True) + initial_task = await anext(passive_stream) + assert isinstance(initial_task, Task) + + task_from_b = await active_b.get_task() + task_from_b.artifacts.append( + Artifact( + artifact_id='replica-b-artifact', + parts=[Part(text='persisted by replica B')], + ) + ) + task_from_b.history.append( + Message( + message_id='replica-a-resume', + role=Role.ROLE_USER, + parts=[Part(text='resume from replica A')], + ) + ) + task_from_b.status.CopyFrom( + TaskStatus( + state=TaskState.TASK_STATE_INPUT_REQUIRED, + message=Message( + message_id='replica-b-status', + role=Role.ROLE_AGENT, + parts=[Part(text='replica B needs input')], + ), + ) + ) + await active_b._task_manager.save_task_event(task_from_b) + + active_a = await registry_a.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + + request_context = _request_context( + task_id, + context_id, + call_context, + 'replica-a-resume', + 'resume from replica A', + ) + events = [ + event async for event in active_a.subscribe(request=request_context) + ] + + assert events + assert replica_a_executor.seen_tasks + assert ( + replica_a_executor.seen_tasks[0].artifacts[0].artifact_id + == 'replica-b-artifact' + ) + assert ( + replica_a_executor.seen_tasks[0].status.message.message_id + == 'replica-b-status' + ) + + persisted_task = await task_store.get(task_id, call_context) + assert persisted_task is not None + assert persisted_task.artifacts[0].artifact_id == 'replica-b-artifact' + assert {message.message_id for message in persisted_task.history} >= { + 'replica-b-status', + 'replica-a-resume', + } + assert ( + sum( + message.message_id == 'replica-a-resume' + for message in persisted_task.history + ) + == 1 + ) + finally: + if passive_stream is not None: + await passive_stream.aclose() + await registry_a.aclose() + await registry_b.aclose() + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_reused_task_rejects_terminal_passive_subscription(): + """A passive subscriber rejects a terminal snapshot from another replica.""" + task_id = 'terminal-task' + context_id = 'terminal-context' + call_context = ServerCallContext() + task_store = InMemoryTaskStore() + await task_store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + call_context, + ) + executor = _RecordingInputRequiredExecutor() + registry = ActiveTaskRegistry(executor, task_store) + + try: + active = await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + await task_store.save( + Task( + id=task_id, + context_id=context_id, + artifacts=[ + Artifact( + artifact_id='completed-artifact', + parts=[Part(text='completed elsewhere')], + ) + ], + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ), + call_context, + ) + await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + + with pytest.raises(InvalidParamsError, match='terminal state'): + await anext(active.subscribe(include_initial_task=True)) + + assert not executor.seen_tasks + persisted_task = await task_store.get(task_id, call_context) + assert persisted_task is not None + assert persisted_task.status.state == TaskState.TASK_STATE_COMPLETED + assert persisted_task.artifacts[0].artifact_id == 'completed-artifact' + finally: + await registry.aclose() + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +@pytest.mark.parametrize( + 'stored_state', + [ + TaskState.TASK_STATE_CANCELED, + None, + ], + ids=['terminal', 'missing'], +) +async def test_reused_task_does_not_cancel_unavailable_snapshot( + stored_state: TaskState | None, +): + """Cancellation honors terminal and missing authoritative snapshots.""" + task_id = 'terminal-cancel-task' + context_id = 'terminal-cancel-context' + call_context = ServerCallContext() + task_store = InMemoryTaskStore() + await task_store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + call_context, + ) + executor = AsyncMock(spec=AgentExecutor) + + async def cancel(context: RequestContext, event_queue: EventQueue) -> None: + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id=context.task_id or '', + context_id=context.context_id or '', + status=TaskStatus(state=TaskState.TASK_STATE_CANCELED), + ) + ) + + executor.cancel.side_effect = cancel + registry = ActiveTaskRegistry(executor, task_store) + + try: + active = await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + if stored_state is None: + await task_store.delete(task_id, call_context) + else: + await task_store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=stored_state), + ), + call_context, + ) + await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + + if stored_state is None: + with pytest.raises(TaskNotFoundError): + await active.cancel(call_context) + else: + result = await active.cancel(call_context) + assert result.status.state == stored_state + + executor.cancel.assert_not_awaited() + persisted_task = await task_store.get(task_id, call_context) + if stored_state is None: + assert persisted_task is None + else: + assert persisted_task is not None + assert persisted_task.status.state == stored_state + finally: + await registry.aclose() + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_reused_task_does_not_recreate_missing_snapshot(): + """A task removed from the authoritative store is not resurrected.""" + task_id = 'removed-task' + context_id = 'removed-context' + call_context = ServerCallContext() + task_store = InMemoryTaskStore() + await task_store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + call_context, + ) + executor = _RecordingInputRequiredExecutor() + registry = ActiveTaskRegistry(executor, task_store) + + try: + active = await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + await task_store.delete(task_id, call_context) + await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + + with pytest.raises(TaskNotFoundError): + async for _ in active.subscribe( + request=_request_context( + task_id, + context_id, + call_context, + 'late-resume', + 'must not recreate', + ) + ): + pass + + assert not executor.seen_tasks + assert await task_store.get(task_id, call_context) is None + finally: + await registry.aclose() + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_reuse_preserves_snapshot_during_background_artifact_stream(): + """A request boundary cannot refresh an artifact stream still in flight.""" + task_id = 'streaming-task' + context_id = 'streaming-context' + call_context = ServerCallContext() + task_store = InMemoryTaskStore() + await task_store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ), + call_context, + ) + + executor = _PausedArtifactExecutor() + registry = ActiveTaskRegistry(executor, task_store) + active = await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + request_context = _request_context( + task_id, + context_id, + call_context, + 'streaming-request', + 'start streaming', + ) + stream = active.subscribe(request=request_context) + + try: + first_event = await anext(stream) + assert isinstance(first_event, TaskArtifactUpdateEvent) + await stream.aclose() + assert active._reference_count == 1 + assert active._request_lock.locked() + + # Model a stale concurrent replica replacing the persisted snapshot + # while this replica still owns an open artifact stream. + await task_store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ), + call_context, + ) + reused = await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + assert reused is active + + executor.release_append.set() + request_finished = asyncio.create_task(active._request_lock.acquire()) + consumer_finished = asyncio.create_task(active._is_finished.wait()) + done, pending = await asyncio.wait( + {request_finished, consumer_finished}, + timeout=1, + return_when=asyncio.FIRST_COMPLETED, + ) + assert done + if request_finished in done: + request_finished.result() + active._request_lock.release() + for task in pending: + task.cancel() + await asyncio.gather(*pending, return_exceptions=True) + + persisted_task = await task_store.get(task_id, call_context) + assert persisted_task is not None + assert not active._is_finished.is_set() + assert [part.text for part in persisted_task.artifacts[0].parts] == [ + 'first chunk', + 'second chunk', + ] + finally: + executor.release_append.set() + await stream.aclose() + await registry.aclose() + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_idle_refresh_wins_over_inflight_passive_snapshot_read(): + """A stale passive read cannot repopulate the cache after invalidation.""" + task_id = 'racing-read-task' + context_id = 'racing-read-context' + call_context = ServerCallContext() + latest_task = Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ) + passive_get_started = asyncio.Event() + release_passive_get = asyncio.Event() + get_count = 0 + + async def get_task( + requested_task_id: str, context: ServerCallContext + ) -> Task: + nonlocal get_count + assert requested_task_id == task_id + assert context is call_context + get_count += 1 + task = Task() + task.CopyFrom(latest_task) + if get_count == 2: + passive_get_started.set() + await release_passive_get.wait() + return task + + async def save_task(task: Task, context: ServerCallContext) -> None: + nonlocal latest_task + assert context is call_context + latest_task = Task() + latest_task.CopyFrom(task) + + task_store = AsyncMock(spec=TaskStore) + task_store.get.side_effect = get_task + task_store.save.side_effect = save_task + registry = ActiveTaskRegistry(_SlowExecutor(), task_store) + active = await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + passive_stream = active.subscribe(include_initial_task=True) + passive_initial_task = asyncio.create_task(anext(passive_stream)) + + try: + await asyncio.wait_for(passive_get_started.wait(), timeout=1) + await task_store.save( + Task( + id=task_id, + context_id=context_id, + artifacts=[ + Artifact( + artifact_id='newer-artifact', + parts=[Part(text='newer replica state')], + ) + ], + status=TaskStatus( + state=TaskState.TASK_STATE_INPUT_REQUIRED, + message=Message( + message_id='newer-status', + role=Role.ROLE_AGENT, + parts=[Part(text='newer input request')], + ), + ), + ), + call_context, + ) + reuse_task = asyncio.create_task( + registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + ) + for _ in range(3): + await asyncio.sleep(0) + release_passive_get.set() + initial_task, reused = await asyncio.gather( + passive_initial_task, reuse_task + ) + assert isinstance(initial_task, Task) + assert reused is active + + refreshed_task = await active.get_task() + assert refreshed_task.artifacts[0].artifact_id == 'newer-artifact' + assert refreshed_task.status.message.message_id == 'newer-status' + finally: + release_passive_get.set() + passive_initial_task.cancel() + await asyncio.gather(passive_initial_task, return_exceptions=True) + await passive_stream.aclose() + await registry.aclose() + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_aclose_not_blocked_by_pending_idle_refresh(): + """Registry shutdown does not wait for a caller's blocked snapshot read.""" + task_id = 'blocked-refresh-task' + context_id = 'blocked-refresh-context' + call_context = ServerCallContext() + task = Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ) + passive_get_started = asyncio.Event() + release_passive_get = asyncio.Event() + get_count = 0 + + async def get_task( + requested_task_id: str, context: ServerCallContext + ) -> Task: + nonlocal get_count + assert requested_task_id == task_id + assert context is call_context + get_count += 1 + if get_count > 1: + passive_get_started.set() + await release_passive_get.wait() + result = Task() + result.CopyFrom(task) + return result + + task_store = AsyncMock(spec=TaskStore) + task_store.get.side_effect = get_task + registry = ActiveTaskRegistry(_RecordingInputRequiredExecutor(), task_store) + active = await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + passive_stream = active.subscribe(include_initial_task=True) + passive_initial_task = asyncio.create_task(anext(passive_stream)) + reuse_task = None + + try: + await asyncio.wait_for(passive_get_started.wait(), timeout=1) + reuse_task = asyncio.create_task( + registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + ) + + for _ in range(3): + await asyncio.sleep(0) + assert not reuse_task.done() + await asyncio.wait_for(registry.aclose(), timeout=1) + finally: + release_passive_get.set() + passive_initial_task.cancel() + await asyncio.gather(passive_initial_task, return_exceptions=True) + if reuse_task is not None: + result = await asyncio.gather(reuse_task, return_exceptions=True) + assert isinstance(result[0], RuntimeError) + await passive_stream.aclose() + await registry.aclose() + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_queued_request_refreshes_snapshot_when_previous_request_finishes(): + """A queued request re-reads the store when it begins processing.""" + task_id = 'queued-resume-task' + context_id = 'queued-resume-context' + call_context = ServerCallContext() + latest_task = Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + second_get_started = asyncio.Event() + release_second_get = asyncio.Event() + get_count = 0 + + async def get_task( + requested_task_id: str, context: ServerCallContext + ) -> Task: + nonlocal get_count + assert requested_task_id == task_id + assert context is call_context + get_count += 1 + if get_count == 3: + second_get_started.set() + await release_second_get.wait() + task = Task() + task.CopyFrom(latest_task) + return task + + async def save_task(task: Task, context: ServerCallContext) -> None: + nonlocal latest_task + assert context is call_context + latest_task = Task() + latest_task.CopyFrom(task) + + task_store = AsyncMock(spec=TaskStore) + task_store.get.side_effect = get_task + task_store.save.side_effect = save_task + executor = _QueuedResumeExecutor() + registry = ActiveTaskRegistry(executor, task_store) + active = await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + + try: + await active.enqueue_request( + _request_context( + task_id, + context_id, + call_context, + 'first-request', + 'first-request', + ) + ) + await asyncio.wait_for(executor.first_started.wait(), timeout=1) + assert active._request_lock.locked() + + # The registry sees the second request while the first request is + # active, so the actual producer boundary must perform the refresh. + await registry.get_or_create( + task_id, + call_context=call_context, + create_task_if_missing=False, + ) + await active.enqueue_request( + _request_context( + task_id, + context_id, + call_context, + 'second-request', + 'second-request', + ) + ) + executor.release_first.set() + await asyncio.wait_for(second_get_started.wait(), timeout=1) + + await task_store.save( + Task( + id=task_id, + context_id=context_id, + artifacts=[ + Artifact( + artifact_id='queued-newer-artifact', + parts=[Part(text='queued newer state')], + ) + ], + status=TaskStatus( + state=TaskState.TASK_STATE_INPUT_REQUIRED, + message=Message( + message_id='queued-newer-status', + role=Role.ROLE_AGENT, + parts=[Part(text='queued newer input')], + ), + ), + ), + call_context, + ) + release_second_get.set() + await asyncio.wait_for(executor.second_started.wait(), timeout=1) + + assert len(executor.seen_tasks) >= 2 + assert ( + executor.seen_tasks[1].artifacts[0].artifact_id + == 'queued-newer-artifact' + ) + assert ( + executor.seen_tasks[1].status.message.message_id + == 'queued-newer-status' + ) + finally: + executor.release_first.set() + release_second_get.set() + await registry.aclose() diff --git a/tests/server/request_handlers/test_default_request_handler_v2.py b/tests/server/request_handlers/test_default_request_handler_v2.py index b276fb77a..c293237ba 100644 --- a/tests/server/request_handlers/test_default_request_handler_v2.py +++ b/tests/server/request_handlers/test_default_request_handler_v2.py @@ -37,6 +37,7 @@ InvalidAgentResponseError, InvalidParamsError, PushNotificationNotSupportedError, + TaskNotCancelableError, TaskNotFoundError, ) from a2a.types.a2a_pb2 import ( @@ -926,6 +927,7 @@ async def test_on_message_send_task_in_terminal_state(terminal_state): task_id=task_id, status_state=terminal_state ) mock_task_store = AsyncMock(spec=TaskStore) + mock_task_store.get.return_value = terminal_task request_handler = DefaultRequestHandlerV2( agent_executor=MockAgentExecutor(), task_store=mock_task_store, @@ -939,13 +941,7 @@ async def test_on_message_send_task_in_terminal_state(terminal_state): task_id=task_id, ) ) - with ( - patch( - 'a2a.server.request_handlers.default_request_handler.TaskManager.get_task', - return_value=terminal_task, - ), - pytest.raises(InvalidParamsError) as exc_info, - ): + with pytest.raises(InvalidParamsError) as exc_info: await request_handler.on_message_send( params, create_server_call_context() ) @@ -965,6 +961,7 @@ async def test_on_message_send_stream_task_in_terminal_state(terminal_state): task_id=task_id, status_state=terminal_state ) mock_task_store = AsyncMock(spec=TaskStore) + mock_task_store.get.return_value = terminal_task request_handler = DefaultRequestHandlerV2( agent_executor=MockAgentExecutor(), task_store=mock_task_store, @@ -978,13 +975,7 @@ async def test_on_message_send_stream_task_in_terminal_state(terminal_state): task_id=task_id, ) ) - with ( - patch( - 'a2a.server.request_handlers.default_request_handler.TaskManager.get_task', - return_value=terminal_task, - ), - pytest.raises(InvalidParamsError) as exc_info, - ): + with pytest.raises(InvalidParamsError) as exc_info: async for _ in request_handler.on_message_send_stream( params, create_server_call_context() ): @@ -995,6 +986,68 @@ async def test_on_message_send_stream_task_in_terminal_state(terminal_state): ) +@pytest.mark.asyncio +async def test_reused_task_admission_rejects_external_terminal_state(): + """Send and cancel read the store even when the registry has the task.""" + task_id = 'externally-completed-task' + context_id = 'externally-completed-context' + context = create_server_call_context() + task_store = InMemoryTaskStore() + await task_store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + context, + ) + executor = AsyncMock(spec=AgentExecutor) + request_handler = DefaultRequestHandlerV2( + agent_executor=executor, + task_store=task_store, + agent_card=create_default_agent_card(), + ) + + try: + await request_handler._active_task_registry.get_or_create( + task_id, + call_context=context, + create_task_if_missing=False, + ) + await task_store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ), + context, + ) + + with pytest.raises(InvalidParamsError, match='terminal state'): + await request_handler.on_message_send( + SendMessageRequest( + message=Message( + task_id=task_id, + context_id=context_id, + message_id='late-resume', + role=Role.ROLE_USER, + parts=[Part(text='must not resume')], + ) + ), + context, + ) + + with pytest.raises(TaskNotCancelableError): + await request_handler.on_cancel_task( + CancelTaskRequest(id=task_id), context + ) + + executor.execute.assert_not_awaited() + executor.cancel.assert_not_awaited() + finally: + await request_handler.aclose() + + @pytest.mark.asyncio async def test_on_message_send_task_id_provided_but_task_not_found(): """Test on_message_send when task_id is provided but task doesn't exist."""