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
14 changes: 11 additions & 3 deletions src/a2a/server/request_handlers/default_request_handler_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,33 +20,34 @@
from a2a.server.request_handlers.request_handler import (
RequestHandler,
validate,
validate_request_params,
)
from a2a.types.a2a_pb2 import (
AgentCard,
CancelTaskRequest,
DeleteTaskPushNotificationConfigRequest,
GetExtendedAgentCardRequest,
GetTaskPushNotificationConfigRequest,
GetTaskRequest,
ListTaskPushNotificationConfigsRequest,
ListTaskPushNotificationConfigsResponse,
ListTasksRequest,
ListTasksResponse,
Message,
SendMessageRequest,
SubscribeToTaskRequest,
Task,
TaskPushNotificationConfig,
TaskState,
)
from a2a.utils.errors import (
ExtendedAgentCardNotConfiguredError,
InternalError,
InvalidParamsError,
PushNotificationNotSupportedError,
TaskNotCancelableError,
TaskNotFoundError,
)

Check notice on line 50 in src/a2a/server/request_handlers/default_request_handler_v2.py

View workflow job for this annotation

GitHub Actions / Lint Code Base

Copy/pasted code

see src/a2a/server/request_handlers/default_request_handler.py (33-60)
from a2a.utils.task import (
apply_history_length,
validate_history_length,
Expand Down Expand Up @@ -298,9 +299,16 @@
):
self._validate_task_id_match(task_id, event.id)
result = event
# DO break here as it's "return_immediately".
# AgentExecutor will continue to run in the background.
break
# A FAILED task may be followed by a producer exception. Keep
# the task as the fallback result, but let the subscription
# surface that exception or finish the current request.
if (
params.configuration.return_immediately
or event.status.state != TaskState.TASK_STATE_FAILED
):
# AgentExecutor will continue to run in the background
# when return_immediately is set.
break

if isinstance(event, Message):
result = event
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -381,6 +381,33 @@ async def cancel(self, context: RequestContext, event_queue: EventQueue):
pass


class FailedStatusAgentExecutor(AgentExecutor):
async def execute(
self, context: RequestContext, event_queue: EventQueue
) -> None:
assert context.message is not None
task = new_task_from_user_message(context.message)
await event_queue.enqueue_event(task)
task_updater = TaskUpdater(event_queue, task.id, task.context_id)
await task_updater.update_status(TaskState.TASK_STATE_FAILED)

async def cancel(
self, context: RequestContext, event_queue: EventQueue
) -> None:
pass


class FailedStatusThenRaisesAgentExecutor(FailedStatusAgentExecutor):
def __init__(self) -> None:
self.exception = RuntimeError('late producer failure')

async def execute(
self, context: RequestContext, event_queue: EventQueue
) -> None:
await super().execute(context, event_queue)
raise self.exception


async def send_message_with_early_failure(
request_handler: DefaultRequestHandlerV2,
params: SendMessageRequest,
Expand Down Expand Up @@ -1297,6 +1324,54 @@ async def save_task_and_signal_terminal_state(self, task):
assert stored_task.status.state == terminal_state


@pytest.mark.asyncio
async def test_on_message_send_failed_task_does_not_hide_producer_exception() -> (
None
):
agent_executor = FailedStatusThenRaisesAgentExecutor()
request_handler = DefaultRequestHandlerV2(
agent_executor=agent_executor,
task_store=InMemoryTaskStore(),
agent_card=create_default_agent_card(),
)
params = SendMessageRequest(
message=Message(
role=Role.ROLE_USER,
message_id='msg_failed_then_raised',
parts=[Part(text='Hi')],
)
)

with pytest.raises(RuntimeError, match='late producer failure') as exc_info:
await request_handler.on_message_send(
params, create_server_call_context()
)
assert exc_info.value is agent_executor.exception


@pytest.mark.asyncio
async def test_on_message_send_returns_agent_declared_failed_task() -> None:
request_handler = DefaultRequestHandlerV2(
agent_executor=FailedStatusAgentExecutor(),
task_store=InMemoryTaskStore(),
agent_card=create_default_agent_card(),
)
params = SendMessageRequest(
message=Message(
role=Role.ROLE_USER,
message_id='msg_declared_failure',
parts=[Part(text='Hi')],
)
)

result = await request_handler.on_message_send(
params, create_server_call_context()
)

assert isinstance(result, Task)
assert result.status.state == TaskState.TASK_STATE_FAILED


@pytest.mark.asyncio
async def test_on_message_send_early_producer_exception_preserves_originating_message():
task_store = InMemoryTaskStore()
Expand Down
Loading