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
33 changes: 33 additions & 0 deletions src/google/adk/artifacts/artifact_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,6 +137,39 @@ def is_artifact_ref(artifact: types.Part) -> bool:
)


def _new_rewind_tombstone() -> types.Part:
"""Returns the part a session rewind saves to remove an artifact.

A rewind cannot delete an artifact that did not exist at the rewind point,
because the versions saved after that point stay in its history. It saves
this part as a new version instead, and every artifact service loads a
version holding it as absent.

An empty ``application/octet-stream`` payload is the marker, so an artifact
saved with exactly that content also loads as absent. A narrower marker would
avoid that, but sessions rewound before now already store this one.

Returns:
A new marker part, safe for the caller to store or modify.
"""
return types.Part(
inline_data=types.Blob(mime_type="application/octet-stream", data=b"")
)


def _is_rewind_tombstone(artifact: types.Part) -> bool:
"""Checks if an artifact part is the marker a session rewind saves.

Args:
artifact: The artifact part to check.

Returns:
True if the part equals the marker from ``_new_rewind_tombstone``, False
otherwise.
"""
return artifact == _new_rewind_tombstone()


def validate_artifact_reference_scope(
*,
app_name: str,
Expand Down
5 changes: 4 additions & 1 deletion src/google/adk/artifacts/file_artifact_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -658,13 +658,16 @@ def _load_artifact_sync(
"Binary artifact %s missing at %s", filename, content_path
)
return None
return types.Part(
artifact = types.Part(
inline_data=types.Blob(
mime_type=mime_type,
data=data,
display_name=metadata.display_name if metadata else None,
)
)
if artifact_util._is_rewind_tombstone(artifact):
return None
return artifact

text = _read_text_if_present(content_path)
if text is None:
Expand Down
5 changes: 4 additions & 1 deletion src/google/adk/artifacts/gcs_artifact_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -458,9 +458,12 @@ def _load_artifact(
display_name=display_name,
)
)
return types.Part.from_bytes(
artifact = types.Part.from_bytes(
data=artifact_bytes, mime_type=blob.content_type
)
if artifact_util._is_rewind_tombstone(artifact):
return None
return artifact

def _list_artifact_keys(
self, app_name: str, user_id: str, session_id: Optional[str]
Expand Down
21 changes: 3 additions & 18 deletions src/google/adk/artifacts/in_memory_artifact_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,23 +48,6 @@ class _ArtifactEntry:
artifact_version: ArtifactVersion


# Runner._compute_artifact_delta_for_rewind marks an artifact as
# inaccessible by saving exactly this part. Match it exactly rather than
# treating every empty payload as absent, so a caller that saves a
# legitimately empty artifact can read it back.
#
# Notes:
# 1. A caller that saves an empty artifact with mime type exactly
# application/octet-stream will still read back None. That collision is
# inherent to using content shape as a tombstone; narrowing the match
# shrinks the hole from every empty artifact to one specific mime type.
# 2. This tombstone convention is in-memory only; other artifact services
# (such as GcsArtifactService) do not perform this empty-payload check.
_REWIND_TOMBSTONE = types.Part(
inline_data=types.Blob(mime_type="application/octet-stream", data=b"")
)


class InMemoryArtifactService(BaseArtifactService, BaseModel):
"""An in-memory implementation of the artifact service.

Expand Down Expand Up @@ -244,7 +227,9 @@ async def _load_artifact(
remaining_depth=remaining_depth - 1,
)

if artifact_data == types.Part() or artifact_data == _REWIND_TOMBSTONE:
if artifact_data == types.Part() or artifact_util._is_rewind_tombstone(
artifact_data
):
return None
return artifact_data.model_copy(deep=True)

Expand Down
13 changes: 5 additions & 8 deletions src/google/adk/sessions/_rewind_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -91,6 +91,9 @@ async def compute_artifact_delta_for_rewind(
if not artifact_service:
return {}

# Imported here so that importing sessions does not load artifacts.
from ..artifacts import artifact_util

versions_at_rewind_point: dict[str, int] = {}
for i in range(rewind_event_index):
event = session.events[i]
Expand All @@ -115,9 +118,7 @@ async def compute_artifact_delta_for_rewind(
artifact: types.Part
if vt is None:
# Artifact did not exist at rewind point. Mark it as inaccessible.
artifact = types.Part(
inline_data=types.Blob(mime_type="application/octet-stream", data=b"")
)
artifact = artifact_util._new_rewind_tombstone()
else:
# Artifact version changed after rewind point. Restore to version at
# rewind point by loading the actual data via the artifact service.
Expand All @@ -136,11 +137,7 @@ async def compute_artifact_delta_for_rewind(
vt,
session.id,
)
artifact = types.Part(
inline_data=types.Blob(
mime_type="application/octet-stream", data=b""
)
)
artifact = artifact_util._new_rewind_tombstone()
else:
artifact = loaded_artifact
await artifact_service.save_artifact(
Expand Down
80 changes: 80 additions & 0 deletions tests/unittests/artifacts/test_artifact_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
import threading
from types import SimpleNamespace
from typing import Any
from typing import Callable
from typing import Optional
from typing import Union
from unittest import mock
Expand All @@ -37,6 +38,7 @@
from google.adk.artifacts import file_artifact_service
from google.adk.artifacts import gcs_artifact_service
from google.adk.artifacts.base_artifact_service import ArtifactVersion
from google.adk.artifacts.base_artifact_service import BaseArtifactService
from google.adk.artifacts.base_artifact_service import ensure_part
from google.adk.artifacts.file_artifact_service import FileArtifactService
from google.adk.artifacts.gcs_artifact_service import GcsArtifactService
Expand Down Expand Up @@ -3568,6 +3570,84 @@ async def test_save_load_empty_bytes_artifact(
assert loaded.inline_data.data == b""


# The part Runner.rewind_async saves to remove an artifact that did not exist
# at the rewind point. Sessions rewound by earlier releases store this exact
# part, so it is spelled out here rather than taken from the code under test.
_REWIND_TOMBSTONE = types.Part(
inline_data=types.Blob(mime_type="application/octet-stream", data=b"")
)


@pytest.mark.asyncio
@pytest.mark.parametrize(
"service_type",
[
ArtifactServiceType.IN_MEMORY,
ArtifactServiceType.GCS,
ArtifactServiceType.FILE,
],
)
@pytest.mark.parametrize("version", [None, 1])
async def test_artifact_removed_by_rewind_loads_as_none(
service_type: ArtifactServiceType,
artifact_service_factory: Callable[
[ArtifactServiceType], BaseArtifactService
],
version: Optional[int],
) -> None:
"""A version holding the rewind tombstone loads as absent."""
artifact_service = artifact_service_factory(service_type)
scope = dict(app_name="app0", user_id="user0", session_id="123")
await artifact_service.save_artifact(
**scope, filename="report.txt", artifact=types.Part.from_text(text="v0")
)
await artifact_service.save_artifact(
**scope,
filename="report.txt",
artifact=_REWIND_TOMBSTONE,
)

loaded = await artifact_service.load_artifact(
**scope, filename="report.txt", version=version
)

assert loaded is None


@pytest.mark.asyncio
@pytest.mark.parametrize(
"service_type",
[
ArtifactServiceType.IN_MEMORY,
ArtifactServiceType.GCS,
ArtifactServiceType.FILE,
],
)
async def test_version_before_rewind_tombstone_still_loads(
service_type: ArtifactServiceType,
artifact_service_factory: Callable[
[ArtifactServiceType], BaseArtifactService
],
) -> None:
"""Versions saved before the rewind tombstone keep their content."""
artifact_service = artifact_service_factory(service_type)
scope = dict(app_name="app0", user_id="user0", session_id="123")
await artifact_service.save_artifact(
**scope, filename="report.txt", artifact=types.Part.from_text(text="v0")
)
await artifact_service.save_artifact(
**scope,
filename="report.txt",
artifact=_REWIND_TOMBSTONE,
)

loaded = await artifact_service.load_artifact(
**scope, filename="report.txt", version=0
)

assert loaded == types.Part.from_text(text="v0")


def _write_tampered_metadata(
root: Path,
*,
Expand Down
57 changes: 57 additions & 0 deletions tests/unittests/runners/test_runner_rewind.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,12 +14,15 @@

"""Tests for runner.rewind_async."""

from pathlib import Path
from typing import Any
from typing import Optional
from typing import Union

from google.adk.agents.base_agent import BaseAgent
from google.adk.artifacts.base_artifact_service import BaseArtifactService
from google.adk.artifacts.base_artifact_service import ensure_part
from google.adk.artifacts.file_artifact_service import FileArtifactService
from google.adk.artifacts.in_memory_artifact_service import InMemoryArtifactService
from google.adk.events.event import Event
from google.adk.events.event import EventActions
Expand Down Expand Up @@ -398,3 +401,57 @@ async def _compute_state_delta_for_rewind(
)
last_event = rewound_session.events[-1]
assert last_event.actions.state_delta == {"overridden": True}


@pytest.mark.asyncio
@pytest.mark.parametrize("storage", ["in_memory", "file"])
async def test_rewind_removes_artifact_created_after_rewind_point(
storage: str, tmp_path: Path
) -> None:
"""An artifact first saved after the rewind point no longer loads."""
artifact_service: BaseArtifactService = (
FileArtifactService(root_dir=tmp_path)
if storage == "file"
else InMemoryArtifactService()
)
runner = Runner(
app_name="test_app",
agent=BaseAgent(name="test_agent"),
session_service=InMemorySessionService(),
artifact_service=artifact_service,
)
session = await runner.session_service.create_session(
app_name="test_app", user_id="u1", session_id="s1"
)
await runner.session_service.append_event(
session=session, event=Event(invocation_id="inv1", author="user")
)
await artifact_service.save_artifact(
app_name="test_app",
user_id="u1",
session_id="s1",
filename="report.txt",
artifact=types.Part.from_text(text="draft"),
)
await runner.session_service.append_event(
session=session,
event=Event(
invocation_id="inv2",
author="agent",
actions=EventActions(artifact_delta={"report.txt": 0}),
),
)

await runner.rewind_async(
user_id="u1", session_id="s1", rewind_before_invocation_id="inv2"
)

assert (
await artifact_service.load_artifact(
app_name="test_app",
user_id="u1",
session_id="s1",
filename="report.txt",
)
is None
)
Loading