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
11 changes: 11 additions & 0 deletions src/a2a/server/request_handlers/default_request_handler_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -203,6 +203,17 @@ 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')
# Reject terminal tasks before persisting anything (e.g. the push
# notification config below). Previously the terminal-state check
# only happened inside ActiveTask.start(), after set_info() had
# already written the config for an already-completed task.
if task.status.state in TERMINAL_TASK_STATES:
raise InvalidParamsError(
message=(
f'Task {task.id} is in terminal state: '
f'{task.status.state}'
)
)

# Build context to resolve or generate missing IDs
request_context = await self._request_context_builder.build(
Expand Down
62 changes: 48 additions & 14 deletions tests/server/request_handlers/test_default_request_handler_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -926,6 +926,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,
Expand All @@ -939,13 +940,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()
)
Expand All @@ -965,6 +960,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,
Expand All @@ -978,13 +974,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()
):
Expand All @@ -995,6 +985,50 @@ async def test_on_message_send_stream_task_in_terminal_state(terminal_state):
)


@pytest.mark.asyncio
@pytest.mark.parametrize('terminal_state', TERMINAL_TASK_STATES)
async def test_on_message_send_terminal_task_skips_push_config_persistence(
terminal_state,
):
"""Push config must not be persisted for a task already in a terminal state.

Regression test for terminal-state persistence guard: the config was previously written to the
store before the terminal-state check ran inside ActiveTask.start().
"""
state_name = TaskState.Name(terminal_state)
task_id = f'terminal_push_task_{state_name}'
terminal_task = create_sample_task(
task_id=task_id, status_state=terminal_state
)
mock_task_store = AsyncMock(spec=TaskStore)
mock_task_store.get.return_value = terminal_task
push_config_store = AsyncMock(spec=PushNotificationConfigStore)
request_handler = DefaultRequestHandlerV2(
agent_executor=MockAgentExecutor(),
task_store=mock_task_store,
push_config_store=push_config_store,
agent_card=create_default_agent_card(),
)
params = SendMessageRequest(
message=Message(
role=Role.ROLE_USER,
message_id='msg_terminal_push',
parts=[Part(text='hello')],
task_id=task_id,
),
configuration=SendMessageConfiguration(
task_push_notification_config=TaskPushNotificationConfig(
url='http://example.com/cb'
)
),
)
with pytest.raises(InvalidParamsError):
await request_handler.on_message_send(
params, create_server_call_context()
)
push_config_store.set_info.assert_not_awaited()


@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."""
Expand Down
Loading