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
8 changes: 6 additions & 2 deletions src/google/adk/agents/base_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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,
Expand Down
68 changes: 68 additions & 0 deletions tests/unittests/agents/test_base_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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'])
Expand Down
Loading