From 8b8121b91785c370a164990364c9049575a2b13b Mon Sep 17 00:00:00 2001 From: KirschQAQ <114209152+KirschBluteX@users.noreply.github.com> Date: Fri, 14 Aug 2026 22:20:44 -0700 Subject: [PATCH 1/4] fix(server): refresh idle ActiveTask snapshots on reuse --- src/a2a/server/agent_execution/active_task.py | 16 ++ .../agent_execution/active_task_registry.py | 6 +- .../test_active_task_registry.py | 148 ++++++++++++++++++ 3 files changed, 168 insertions(+), 2 deletions(-) diff --git a/src/a2a/server/agent_execution/active_task.py b/src/a2a/server/agent_execution/active_task.py index ea1955fdd..d2b8ed682 100644 --- a/src/a2a/server/agent_execution/active_task.py +++ b/src/a2a/server/agent_execution/active_task.py @@ -427,6 +427,22 @@ 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 HITL input, + so another process may persist a newer task snapshot in the meantime. + Refreshing only while the producer/consumer reference is idle avoids + replacing the task object during an active streaming request, where + artifact chunks may still be appended to the in-memory snapshot. + """ + async with self._lock: + if self._reference_count <= 1: + self._task_manager._call_context = call_context + self._task_manager._current_task = None + async def start( self, call_context: ServerCallContext, diff --git a/src/a2a/server/agent_execution/active_task_registry.py b/src/a2a/server/agent_execution/active_task_registry.py index ab7d6a11c..c4fd1bae9 100644 --- a/src/a2a/server/agent_execution/active_task_registry.py +++ b/src/a2a/server/agent_execution/active_task_registry.py @@ -49,8 +49,10 @@ async def get_or_create( 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] + active_task = self._active_tasks.get(task_id) + if active_task is not None: + await active_task.refresh_task_if_idle(call_context) + return active_task task_manager = TaskManager( task_id=task_id, diff --git a/tests/server/agent_execution/test_active_task_registry.py b/tests/server/agent_execution/test_active_task_registry.py index 16d9c8797..e6f581ce8 100644 --- a/tests/server/agent_execution/test_active_task_registry.py +++ b/tests/server/agent_execution/test_active_task_registry.py @@ -11,6 +11,17 @@ from a2a.server.context import ServerCallContext from a2a.server.events.event_queue_v2 import EventQueue from a2a.server.tasks import InMemoryTaskStore +from a2a.types.a2a_pb2 import ( + Artifact, + Message, + Part, + Role, + SendMessageRequest, + Task, + TaskState, + TaskStatus, + TaskStatusUpdateEvent, +) class _SlowExecutor(AgentExecutor): @@ -27,6 +38,36 @@ 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 + + def _make_registry() -> ActiveTaskRegistry: return ActiveTaskRegistry( agent_executor=_SlowExecutor(), @@ -105,3 +146,110 @@ 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 +async def test_reused_idle_active_task_refreshes_shared_store_snapshot(): + """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) + + 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, + ) + + 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.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, + ) + + resume_request = SendMessageRequest( + message=Message( + task_id=task_id, + context_id=context_id, + message_id='replica-a-resume', + role=Role.ROLE_USER, + parts=[Part(text='resume from replica A')], + ) + ) + request_context = RequestContext( + call_context=call_context, + request=resume_request, + task_id=task_id, + context_id=context_id, + ) + 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', + } + finally: + await registry_a.aclose() + await registry_b.aclose() From 5c466c054ec14619ceb61b4eaa183c2039d4343a Mon Sep 17 00:00:00 2001 From: KirschQAQ <114209152+KirschBluteX@users.noreply.github.com> Date: Fri, 14 Aug 2026 22:52:53 -0700 Subject: [PATCH 2/4] fix(ci): spell out human-in-the-loop in docs --- src/a2a/server/agent_execution/active_task.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/a2a/server/agent_execution/active_task.py b/src/a2a/server/agent_execution/active_task.py index d2b8ed682..ee59838c0 100644 --- a/src/a2a/server/agent_execution/active_task.py +++ b/src/a2a/server/agent_execution/active_task.py @@ -432,7 +432,8 @@ async def refresh_task_if_idle( ) -> None: """Drops an idle task snapshot before a new request boundary. - An ``ActiveTask`` remains registered while a task waits for HITL input, + 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. Refreshing only while the producer/consumer reference is idle avoids replacing the task object during an active streaming request, where From 0561ff8ddc3b6dddee9d324d6d0976bc68f47f50 Mon Sep 17 00:00:00 2001 From: KirschQAQ <114209152+KirschBluteX@users.noreply.github.com> Date: Sat, 15 Aug 2026 11:22:38 -0700 Subject: [PATCH 3/4] fix(server): harden ActiveTask snapshot refresh boundaries --- src/a2a/server/agent_execution/active_task.py | 68 +- .../agent_execution/active_task_registry.py | 56 +- .../default_request_handler_v2.py | 11 + .../agent_execution/test_active_task.py | 50 +- .../test_active_task_registry.py | 746 +++++++++++++++++- .../test_default_request_handler_v2.py | 81 +- 6 files changed, 910 insertions(+), 102 deletions(-) diff --git a/src/a2a/server/agent_execution/active_task.py b/src/a2a/server/agent_execution/active_task.py index ee59838c0..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: @@ -435,12 +445,13 @@ async def refresh_task_if_idle( 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. - Refreshing only while the producer/consumer reference is idle avoids - replacing the task object during an active streaming request, where - artifact chunks may still be appended to the in-memory snapshot. + 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._lock: - if self._reference_count <= 1: + 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 @@ -485,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 @@ -527,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: ( @@ -534,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, @@ -579,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 '', @@ -660,6 +682,7 @@ async def subscribe( self._task_id, ) task = await self.get_task() + self._raise_if_task_terminal(task) yield task while True: @@ -732,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, @@ -860,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 c4fd1bae9..b7d2c33e8 100644 --- a/src/a2a/server/agent_execution/active_task_registry.py +++ b/src/a2a/server/agent_execution/active_task_registry.py @@ -46,30 +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') - active_task = self._active_tasks.get(task_id) - if active_task is not None: - await active_task.refresh_task_if_idle(call_context) - return active_task - - 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 e6f581ce8..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,7 @@ 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, @@ -18,10 +18,12 @@ Role, SendMessageRequest, Task, + TaskArtifactUpdateEvent, TaskState, TaskStatus, TaskStatusUpdateEvent, ) +from a2a.utils.errors import InvalidParamsError, TaskNotFoundError class _SlowExecutor(AgentExecutor): @@ -68,6 +70,131 @@ async def cancel( 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(), @@ -150,7 +277,14 @@ async def test_aclose_logs_and_swallows_task_errors(caplog): @pytest.mark.timeout(5) @pytest.mark.asyncio -async def test_reused_idle_active_task_refreshes_shared_store_snapshot(): +@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' @@ -175,6 +309,7 @@ async def test_reused_idle_active_task_refreshes_shared_store_snapshot(): 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( @@ -187,6 +322,10 @@ async def test_reused_idle_active_task_refreshes_shared_store_snapshot(): 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( @@ -195,6 +334,13 @@ async def test_reused_idle_active_task_refreshes_shared_store_snapshot(): 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, @@ -213,20 +359,12 @@ async def test_reused_idle_active_task_refreshes_shared_store_snapshot(): create_task_if_missing=False, ) - resume_request = SendMessageRequest( - message=Message( - task_id=task_id, - context_id=context_id, - message_id='replica-a-resume', - role=Role.ROLE_USER, - parts=[Part(text='resume from replica A')], - ) - ) - request_context = RequestContext( - call_context=call_context, - request=resume_request, - task_id=task_id, - context_id=context_id, + 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) @@ -250,6 +388,582 @@ async def test_reused_idle_active_task_refreshes_shared_store_snapshot(): '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.""" From dc88fbccc8da0f682007ab274f91e67114a76e29 Mon Sep 17 00:00:00 2001 From: KirschQAQ <114209152+KirschBluteX@users.noreply.github.com> Date: Mon, 17 Aug 2026 08:09:10 -0700 Subject: [PATCH 4/4] fix(server): propagate request-boundary store failures --- src/a2a/server/agent_execution/active_task.py | 44 ++++++++----- .../test_default_request_handler_v2.py | 61 +++++++++++++++++++ 2 files changed, 89 insertions(+), 16 deletions(-) diff --git a/src/a2a/server/agent_execution/active_task.py b/src/a2a/server/agent_execution/active_task.py index 6d3c49967..851406e36 100644 --- a/src/a2a/server/agent_execution/active_task.py +++ b/src/a2a/server/agent_execution/active_task.py @@ -140,16 +140,22 @@ async def run(self) -> None: logger.exception('Consumer[%s]: Failed', self.active_task._task_id) updated_task = None - task = await self.active_task._task_manager.get_task() - if task and task.status.state not in TERMINAL_TASK_STATES: - handled_event = TaskStatusUpdateEvent( - task_id=task.id, - context_id=task.context_id, - status=TaskStatus( - state=TaskState.TASK_STATE_FAILED, - ), + try: + task = await self.active_task._task_manager.get_task() + if task and task.status.state not in TERMINAL_TASK_STATES: + handled_event = TaskStatusUpdateEvent( + task_id=task.id, + context_id=task.context_id, + status=TaskStatus( + state=TaskState.TASK_STATE_FAILED, + ), + ) + updated_task = await self._handle_task_event(handled_event) + except Exception: + logger.exception( + 'Consumer[%s]: Failed to persist task failure', + self.active_task._task_id, ) - updated_task = await self._handle_task_event(handled_event) await self._enqueue_to_subscribers(cast('Event', e), updated_task) @@ -601,15 +607,21 @@ 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 and not task_missing_at_boundary: - task = await self._task_manager.ensure_task_id( + try: + 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 '', + ) + if task.status.state not in TERMINAL_TASK_STATES: + task.status.state = TaskState.TASK_STATE_FAILED + await self._task_manager.save_task_event(task) + self._task_created.set() + except Exception: + logger.exception( + 'Producer[%s]: Failed to persist task failure', self._task_id, - request_context.context_id or '', ) - if task.status.state not in TERMINAL_TASK_STATES: - task.status.state = TaskState.TASK_STATE_FAILED - await self._task_manager.save_task_event(task) - self._task_created.set() await self._event_queue_agent.enqueue_event(cast('Event', e)) finally: 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 c293237ba..c271a6e6d 100644 --- a/tests/server/request_handlers/test_default_request_handler_v2.py +++ b/tests/server/request_handlers/test_default_request_handler_v2.py @@ -1223,6 +1223,67 @@ async def test_on_message_send_limit_history(): assert task.history is not None and len(task.history) > 1 +@pytest.mark.asyncio +@pytest.mark.timeout(5) +@pytest.mark.parametrize('consumer_persistence_fails', [False, True]) +async def test_on_message_send_propagates_request_boundary_store_failure( + consumer_persistence_fails: bool, +) -> None: + task_id = 'request-boundary-store-failure' + context_id = 'request-boundary-context' + context = create_server_call_context() + stored_task = create_sample_task( + task_id=task_id, + context_id=context_id, + status_state=TaskState.TASK_STATE_INPUT_REQUIRED, + ) + task_store = AsyncMock(spec=TaskStore) + task_store.get.return_value = stored_task + agent_executor = AsyncMock(spec=AgentExecutor) + request_handler = DefaultRequestHandlerV2( + agent_executor=agent_executor, + task_store=task_store, + agent_card=create_default_agent_card(), + ) + + try: + await request_handler._active_task_registry.get_or_create( + task_id, + context_id=context_id, + call_context=context, + create_task_if_missing=False, + ) + task_store.get.reset_mock() + get_results: list[Task | OSError] = [ + stored_task, + OSError('request-boundary read failed'), + OSError('failure-state read failed'), + ] + if consumer_persistence_fails: + get_results.append(OSError('consumer failure-state read failed')) + get_results.append(stored_task) + task_store.get.side_effect = get_results + + with pytest.raises(OSError, match='request-boundary read failed'): + await request_handler.on_message_send( + SendMessageRequest( + message=Message( + task_id=task_id, + context_id=context_id, + message_id='resume-after-store-failure', + role=Role.ROLE_USER, + parts=[Part(text='resume')], + ) + ), + context, + ) + + assert task_store.get.await_count == 4 + agent_executor.execute.assert_not_awaited() + finally: + await request_handler.aclose() + + @pytest.mark.asyncio async def test_on_message_send_early_producer_exception_marks_task_failed_and_preserves_originating_message(): task_store = InMemoryTaskStore()