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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 32 additions & 0 deletions src/a2a/server/tasks/task_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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:
Expand Down
73 changes: 73 additions & 0 deletions tests/server/tasks/test_task_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading