diff --git a/src/google/adk/agents/base_agent.py b/src/google/adk/agents/base_agent.py index 651ede37ca..a9a3d4bb41 100644 --- a/src/google/adk/agents/base_agent.py +++ b/src/google/adk/agents/base_agent.py @@ -567,7 +567,8 @@ async def _handle_agent_callbacks( end_invocation_on_content: Whether returned content ends the invocation. Returns: - Optional[Event]: an event if a callback provides content or changed state. + Optional[Event]: an event if a callback provides content or changed state + or artifacts. """ callback_context = CallbackContext(ctx) @@ -599,7 +600,10 @@ async def _handle_agent_callbacks( ctx.end_invocation = True return ret_event - if callback_context.state.has_delta(): + if ( + callback_context.state.has_delta() + or callback_context._event_actions.artifact_delta + ): return Event( invocation_id=ctx.invocation_id, author=self.name, diff --git a/tests/unittests/agents/test_base_agent.py b/tests/unittests/agents/test_base_agent.py index 7a8e0ba09a..63f43507e1 100644 --- a/tests/unittests/agents/test_base_agent.py +++ b/tests/unittests/agents/test_base_agent.py @@ -32,6 +32,7 @@ from google.adk.agents.invocation_context import InvocationContext from google.adk.agents.llm_agent import LlmAgent from google.adk.apps.app import ResumabilityConfig +from google.adk.artifacts.in_memory_artifact_service import InMemoryArtifactService from google.adk.events.event import Event from google.adk.plugins.base_plugin import BasePlugin from google.adk.plugins.plugin_manager import PluginManager @@ -720,6 +721,73 @@ async def test_run_async_with_async_after_agent_callback_append_reply( ) +async def _before_agent_callback_save_artifact( + callback_context: CallbackContext, +) -> None: + await callback_context.save_artifact( + 'before_agent.txt', types.Part.from_text(text='saved before agent') + ) + + +async def _after_agent_callback_save_artifact( + callback_context: CallbackContext, +) -> None: + await callback_context.save_artifact( + 'after_agent.txt', types.Part.from_text(text='saved after agent') + ) + + +@pytest.mark.asyncio +async def test_run_async_before_agent_callback_records_artifact_delta( + request: pytest.FixtureRequest, +): + # Arrange + agent = _TestingAgent( + name=f'{request.function.__name__}_test_agent', + before_agent_callback=_before_agent_callback_save_artifact, + ) + parent_ctx = await _create_parent_invocation_context( + request.function.__name__, agent + ) + parent_ctx.artifact_service = InMemoryArtifactService() + + # Act + events = [e async for e in agent.run_async(parent_ctx)] + + # Assert + # The first event records the artifact saved by before_agent_callback, the + # second event is the regular agent response. + assert len(events) == 2 + assert events[0].author == agent.name + assert events[0].actions.artifact_delta == {'before_agent.txt': 0} + assert events[1].content.parts[0].text == 'Hello, world!' + + +@pytest.mark.asyncio +async def test_run_async_after_agent_callback_records_artifact_delta( + request: pytest.FixtureRequest, +): + # Arrange + agent = _TestingAgent( + name=f'{request.function.__name__}_test_agent', + after_agent_callback=_after_agent_callback_save_artifact, + ) + parent_ctx = await _create_parent_invocation_context( + request.function.__name__, agent + ) + parent_ctx.artifact_service = InMemoryArtifactService() + + # Act + events = [e async for e in agent.run_async(parent_ctx)] + + # Assert + # The first event is the regular agent response, the second event records + # the artifact saved by after_agent_callback. + assert len(events) == 2 + assert events[1].author == agent.name + assert events[1].actions.artifact_delta == {'after_agent.txt': 0} + + @pytest.mark.asyncio @pytest.mark.parametrize('method_name', ['run_async', 'run_live']) @pytest.mark.parametrize('exit_type', ['cancelled', 'generator_exit'])