diff --git a/src/a2a/server/tasks/task_manager.py b/src/a2a/server/tasks/task_manager.py index c9dfc879f..54b4563ff 100644 --- a/src/a2a/server/tasks/task_manager.py +++ b/src/a2a/server/tasks/task_manager.py @@ -19,6 +19,37 @@ logger = logging.getLogger(__name__) +TERMINAL_TASK_STATES = { + TaskState.TASK_STATE_COMPLETED, + TaskState.TASK_STATE_CANCELED, + TaskState.TASK_STATE_FAILED, + TaskState.TASK_STATE_REJECTED, +} + + +def validate_state_transition( + current_state: TaskState, new_state: TaskState +) -> None: + """Validates a task state transition before it is persisted. + + Terminal states are final: a task in COMPLETED/CANCELED/FAILED/REJECTED + must never move to a different state (including back to SUBMITTED). + Re-persisting the same terminal state is tolerated for idempotency. + + Raises: + InvalidAgentResponseError: If the transition moves a task away from + a terminal state. + """ + if current_state in TERMINAL_TASK_STATES and new_state != current_state: + raise InvalidAgentResponseError( + message=( + f'Illegal state transition from terminal state ' + f'{TaskState.Name(current_state)} to ' + f'{TaskState.Name(new_state)}' + ) + ) + + @trace_function() def append_artifact_to_task(task: Task, event: TaskArtifactUpdateEvent) -> None: """Helper method for updating a Task object with new artifact data from an event. @@ -204,6 +235,7 @@ async def save_task_event( logger.debug( 'Updating task %s status to: %s', task.id, event.status.state ) + validate_state_transition(task.status.state, event.status.state) if task.status.HasField('message'): task.history.append(task.status.message) if event.metadata: diff --git a/tests/server/tasks/test_task_manager.py b/tests/server/tasks/test_task_manager.py index 7daa3ab3d..dff1b0f0c 100644 --- a/tests/server/tasks/test_task_manager.py +++ b/tests/server/tasks/test_task_manager.py @@ -444,3 +444,76 @@ def test_append_artifact_to_task(): match='append=True for nonexistent artifact_id', ): append_artifact_to_task(task, append_event_5) + + +@pytest.mark.parametrize( + ('current', 'new'), + [ + (TaskState.TASK_STATE_COMPLETED, TaskState.TASK_STATE_SUBMITTED), + (TaskState.TASK_STATE_COMPLETED, TaskState.TASK_STATE_WORKING), + (TaskState.TASK_STATE_CANCELED, TaskState.TASK_STATE_SUBMITTED), + (TaskState.TASK_STATE_FAILED, TaskState.TASK_STATE_WORKING), + (TaskState.TASK_STATE_REJECTED, TaskState.TASK_STATE_COMPLETED), + ], +) +def test_validate_state_transition_blocks_terminal_escape(current, new): + """A terminal state must never transition to a different state.""" + from a2a.server.tasks.task_manager import validate_state_transition + + with pytest.raises(InvalidAgentResponseError): + validate_state_transition(current, new) + + +@pytest.mark.parametrize( + ('current', 'new'), + [ + (TaskState.TASK_STATE_COMPLETED, TaskState.TASK_STATE_COMPLETED), + (TaskState.TASK_STATE_SUBMITTED, TaskState.TASK_STATE_WORKING), + (TaskState.TASK_STATE_SUBMITTED, TaskState.TASK_STATE_COMPLETED), + (TaskState.TASK_STATE_WORKING, TaskState.TASK_STATE_COMPLETED), + (TaskState.TASK_STATE_WORKING, TaskState.TASK_STATE_INPUT_REQUIRED), + (TaskState.TASK_STATE_INPUT_REQUIRED, TaskState.TASK_STATE_WORKING), + ], +) +def test_validate_state_transition_allows_legal_transitions(current, new): + """Legal forward transitions (and idempotent repeats) are accepted.""" + from a2a.server.tasks.task_manager import validate_state_transition + + validate_state_transition(current, new) + + +@pytest.mark.asyncio +async def test_save_task_event_blocks_terminal_overwrite(): + """save_task_event must reject a terminal -> non-terminal status update. + + Regression test for BUG-43: the status was blindly overwritten, so a + misbehaving agent could move a completed task back to SUBMITTED. + """ + from a2a.server.tasks import InMemoryTaskStore + from a2a.server.tasks.task_manager import TaskManager + + task_store = InMemoryTaskStore() + task_manager = TaskManager( + task_id=MINIMAL_TASK_ID, + context_id=MINIMAL_CONTEXT_ID, + task_store=task_store, + initial_message=None, + context=TEST_CONTEXT, + ) + + task = create_minimal_task() + task.status.state = TaskState.TASK_STATE_COMPLETED + await task_store.save(task, TEST_CONTEXT) + + illegal_event = TaskStatusUpdateEvent( + task_id=MINIMAL_TASK_ID, + context_id=MINIMAL_CONTEXT_ID, + status=TaskStatus(state=TaskState.TASK_STATE_SUBMITTED), + ) + + with pytest.raises(InvalidAgentResponseError): + await task_manager.save_task_event(illegal_event) + + # The task state must be unchanged. + stored = await task_store.get(MINIMAL_TASK_ID, TEST_CONTEXT) + assert stored.status.state == TaskState.TASK_STATE_COMPLETED