feat: allow users to include artifacts from artifact_service in A2A events using provided interceptor
PiperOrigin-RevId: 897143688
This commit is contained in:
committed by
Copybara-Service
parent
dcc485b23e
commit
e63d991be8
@@ -260,17 +260,15 @@ class A2aAgentExecutor(AgentExecutor):
|
||||
context.context_id,
|
||||
self._config.gen_ai_part_converter,
|
||||
):
|
||||
a2a_event = await execute_after_event_interceptors(
|
||||
a2a_events = await execute_after_event_interceptors(
|
||||
a2a_event,
|
||||
executor_context,
|
||||
adk_event,
|
||||
self._config.execute_interceptors,
|
||||
)
|
||||
if a2a_event is None:
|
||||
continue
|
||||
|
||||
task_result_aggregator.process_event(a2a_event)
|
||||
await event_queue.enqueue_event(a2a_event)
|
||||
for e in a2a_events:
|
||||
task_result_aggregator.process_event(e)
|
||||
await event_queue.enqueue_event(e)
|
||||
|
||||
# publish the task result event - this is final
|
||||
if (
|
||||
|
||||
@@ -49,6 +49,7 @@ from ..converters.utils import _get_adk_metadata_key
|
||||
from ..experimental import a2a_experimental
|
||||
from .config import A2aAgentExecutorConfig
|
||||
from .executor_context import ExecutorContext
|
||||
from .interceptors.include_artifacts_in_a2a_event import include_artifacts_in_a2a_event_interceptor
|
||||
from .utils import execute_after_agent_interceptors
|
||||
from .utils import execute_after_event_interceptors
|
||||
from .utils import execute_before_agent_interceptors
|
||||
@@ -221,15 +222,14 @@ class _A2aAgentExecutor(AgentExecutor):
|
||||
self._config.gen_ai_part_converter,
|
||||
):
|
||||
a2a_event.metadata = self._get_invocation_metadata(executor_context)
|
||||
a2a_event = await execute_after_event_interceptors(
|
||||
a2a_events = await execute_after_event_interceptors(
|
||||
a2a_event,
|
||||
executor_context,
|
||||
adk_event,
|
||||
self._config.execute_interceptors,
|
||||
)
|
||||
if not a2a_event:
|
||||
continue
|
||||
await event_queue.enqueue_event(a2a_event)
|
||||
for e in a2a_events:
|
||||
await event_queue.enqueue_event(e)
|
||||
|
||||
if error_event:
|
||||
final_event = error_event
|
||||
|
||||
@@ -57,7 +57,7 @@ class ExecuteInterceptor:
|
||||
after_event: Optional[
|
||||
Callable[
|
||||
[ExecutorContext, A2AEvent, Event],
|
||||
Awaitable[Union[A2AEvent, None]],
|
||||
Awaitable[Union[A2AEvent, list[A2AEvent], None]],
|
||||
]
|
||||
] = None
|
||||
"""Hook executed after an ADK event is converted to an A2A event.
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
# Copyright 2026 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from .include_artifacts_in_a2a_event import include_artifacts_in_a2a_event_interceptor
|
||||
|
||||
__all__ = [
|
||||
"include_artifacts_in_a2a_event_interceptor",
|
||||
]
|
||||
@@ -0,0 +1,73 @@
|
||||
# Copyright 2026 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Union
|
||||
|
||||
from a2a.server.events import Event as A2AEvent
|
||||
from a2a.types import Artifact
|
||||
from a2a.types import TaskArtifactUpdateEvent
|
||||
from a2a.types import TaskStatusUpdateEvent
|
||||
from google.adk.a2a.executor.config import ExecuteInterceptor
|
||||
from google.adk.a2a.executor.config import ExecutorContext
|
||||
|
||||
from ....events.event import Event
|
||||
from ...converters.part_converter import convert_genai_part_to_a2a_part
|
||||
|
||||
|
||||
async def _after_agent(
|
||||
ctx: ExecutorContext, a2a_event: A2AEvent, adk_event: Event
|
||||
) -> Union[A2AEvent, list[A2AEvent]]:
|
||||
"""After agent interceptor that includes artifacts in A2A events."""
|
||||
if isinstance(a2a_event, (TaskStatusUpdateEvent, TaskArtifactUpdateEvent)):
|
||||
artifact_service = ctx.runner.artifact_service
|
||||
if artifact_service and adk_event.actions.artifact_delta:
|
||||
new_events = []
|
||||
for filename, version in adk_event.actions.artifact_delta.items():
|
||||
genai_part = await artifact_service.load_artifact(
|
||||
app_name=ctx.app_name,
|
||||
user_id=ctx.user_id,
|
||||
session_id=ctx.session_id,
|
||||
filename=filename,
|
||||
version=version,
|
||||
)
|
||||
if genai_part:
|
||||
a2a_part = convert_genai_part_to_a2a_part(genai_part)
|
||||
if a2a_part:
|
||||
a2a_artifact = Artifact(
|
||||
artifact_id=f"{filename}_{version}",
|
||||
name=filename,
|
||||
parts=[a2a_part],
|
||||
)
|
||||
new_event = TaskArtifactUpdateEvent(
|
||||
task_id=a2a_event.task_id,
|
||||
context_id=a2a_event.context_id,
|
||||
artifact=a2a_artifact,
|
||||
metadata=a2a_event.metadata,
|
||||
append=False,
|
||||
last_chunk=True,
|
||||
)
|
||||
new_events.append(new_event)
|
||||
|
||||
adk_event.actions.artifact_delta = {}
|
||||
|
||||
if new_events:
|
||||
return [a2a_event] + new_events
|
||||
|
||||
return a2a_event
|
||||
|
||||
|
||||
include_artifacts_in_a2a_event_interceptor = ExecuteInterceptor(
|
||||
after_event=_after_agent
|
||||
)
|
||||
@@ -41,16 +41,24 @@ async def execute_after_event_interceptors(
|
||||
executor_context: ExecutorContext,
|
||||
adk_event: Event,
|
||||
execute_interceptors: Optional[list[ExecuteInterceptor]],
|
||||
) -> Optional[A2AEvent]:
|
||||
) -> list[A2AEvent]:
|
||||
events = [a2a_event]
|
||||
if execute_interceptors:
|
||||
for interceptor in execute_interceptors:
|
||||
if interceptor.after_event:
|
||||
a2a_event = await interceptor.after_event(
|
||||
executor_context, a2a_event, adk_event
|
||||
)
|
||||
if a2a_event is None:
|
||||
return None
|
||||
return a2a_event
|
||||
next_events = []
|
||||
for e in events:
|
||||
res = await interceptor.after_event(executor_context, e, adk_event)
|
||||
if res is None:
|
||||
continue
|
||||
if isinstance(res, list):
|
||||
next_events.extend(res)
|
||||
else:
|
||||
next_events.append(res)
|
||||
events = next_events
|
||||
if not events:
|
||||
return []
|
||||
return events
|
||||
|
||||
|
||||
async def execute_after_agent_interceptors(
|
||||
|
||||
@@ -14,6 +14,7 @@
|
||||
|
||||
"""A2A Server for integration tests."""
|
||||
|
||||
from unittest.mock import AsyncMock
|
||||
from unittest.mock import Mock
|
||||
|
||||
from a2a.server.apps.jsonrpc.fastapi_app import A2AFastAPIApplication
|
||||
@@ -23,9 +24,12 @@ from a2a.types import AgentCapabilities
|
||||
from a2a.types import AgentCard
|
||||
from a2a.types import AgentSkill
|
||||
from google.adk.a2a.executor.a2a_agent_executor import A2aAgentExecutor
|
||||
from google.adk.a2a.executor.config import A2aAgentExecutorConfig
|
||||
from google.adk.a2a.executor.interceptors.include_artifacts_in_a2a_event import include_artifacts_in_a2a_event_interceptor
|
||||
from google.adk.agents.base_agent import BaseAgent
|
||||
from google.adk.runners import Runner
|
||||
from google.adk.sessions.in_memory_session_service import InMemorySessionService
|
||||
from google.genai import types
|
||||
|
||||
|
||||
class FakeRunner(Runner):
|
||||
@@ -43,6 +47,12 @@ class FakeRunner(Runner):
|
||||
)
|
||||
self.run_async_fn = run_async_fn
|
||||
|
||||
mock_artifact_service = Mock()
|
||||
mock_artifact_service.load_artifact = AsyncMock(
|
||||
return_value=types.Part(text="artifact content")
|
||||
)
|
||||
self.artifact_service = mock_artifact_service
|
||||
|
||||
async def run_async(self, **kwargs):
|
||||
async for event in self.run_async_fn(**kwargs):
|
||||
yield event
|
||||
@@ -63,18 +73,21 @@ agent_card = AgentCard(
|
||||
)
|
||||
|
||||
|
||||
def create_server_app(run_async_fn):
|
||||
def create_server_app(
|
||||
run_async_fn=None, config: A2aAgentExecutorConfig | None = None
|
||||
):
|
||||
"""Creates an A2A FastAPI application with a mocked runner.
|
||||
|
||||
Args:
|
||||
run_async_fn: A generator function that takes **kwargs and yields Event
|
||||
objects.
|
||||
include_artifacts: Whether to include artifacts in A2A events.
|
||||
|
||||
Returns:
|
||||
A FastAPI application instance.
|
||||
"""
|
||||
runner = FakeRunner(run_async_fn)
|
||||
executor = A2aAgentExecutor(runner=runner)
|
||||
executor = A2aAgentExecutor(runner=runner, config=config)
|
||||
task_store = InMemoryTaskStore()
|
||||
handler = DefaultRequestHandler(
|
||||
agent_executor=executor, task_store=task_store
|
||||
|
||||
@@ -19,6 +19,8 @@ from a2a.types import Part as A2APart
|
||||
from a2a.types import Task
|
||||
from a2a.types import TaskState
|
||||
from a2a.types import TextPart
|
||||
from google.adk.a2a.executor.config import A2aAgentExecutorConfig
|
||||
from google.adk.a2a.executor.interceptors.include_artifacts_in_a2a_event import include_artifacts_in_a2a_event_interceptor
|
||||
from google.adk.agents.remote_a2a_agent import A2A_METADATA_PREFIX
|
||||
from google.adk.events.event import Event
|
||||
from google.adk.events.event_actions import EventActions
|
||||
@@ -586,3 +588,51 @@ async def test_user_follow_up():
|
||||
)
|
||||
|
||||
assert last_event is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_include_artifacts_in_a2a_event():
|
||||
"""Test that artifacts are included in A2A events when the interceptor is enabled."""
|
||||
|
||||
async def mock_run_async(**kwargs):
|
||||
yield Event(
|
||||
actions=EventActions(artifact_delta={"artifact1": 1, "artifact2": 1}),
|
||||
author="agent",
|
||||
content=types.Content(
|
||||
parts=[types.Part(text="Here are the artifacts")]
|
||||
),
|
||||
)
|
||||
|
||||
config = A2aAgentExecutorConfig(
|
||||
execute_interceptors=[include_artifacts_in_a2a_event_interceptor]
|
||||
)
|
||||
built_app = create_server_app(mock_run_async, config=config)
|
||||
|
||||
a2a_client = create_a2a_client(built_app, streaming=False)
|
||||
|
||||
request = A2AMessage(
|
||||
message_id="test_message_id",
|
||||
parts=[A2APart(root=TextPart(text="Hi"))],
|
||||
role="user",
|
||||
)
|
||||
|
||||
events = []
|
||||
async for event in a2a_client.send_message(request=request):
|
||||
events.append(event)
|
||||
|
||||
assert len(events) == 1
|
||||
|
||||
task = events[0][0]
|
||||
assert isinstance(task, Task)
|
||||
assert task.artifacts is not None
|
||||
assert len(task.artifacts) == 3
|
||||
|
||||
assert task.artifacts[0].parts[0].root.text == "Here are the artifacts"
|
||||
|
||||
assert task.artifacts[1].artifact_id == "artifact1_1"
|
||||
assert task.artifacts[1].name == "artifact1"
|
||||
assert task.artifacts[1].parts[0].root.text == "artifact content"
|
||||
|
||||
assert task.artifacts[2].artifact_id == "artifact2_1"
|
||||
assert task.artifacts[2].name == "artifact2"
|
||||
assert task.artifacts[2].parts[0].root.text == "artifact content"
|
||||
|
||||
Reference in New Issue
Block a user