Skip to content

Commit df55309

Browse files
baba9811copybara-github
authored andcommitted
fix: preserve empty text and bytes in memory
Merge #7072 PiperOrigin-RevId: 981274973
1 parent b1b0adc commit df55309

3 files changed

Lines changed: 65 additions & 11 deletions

File tree

‎src/google/adk/artifacts/in_memory_artifact_service.py‎

Lines changed: 18 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,23 @@ class _ArtifactEntry:
4747
artifact_version: ArtifactVersion
4848

4949

50+
# Runner._compute_artifact_delta_for_rewind marks an artifact as
51+
# inaccessible by saving exactly this part. Match it exactly rather than
52+
# treating every empty payload as absent, so a caller that saves a
53+
# legitimately empty artifact can read it back.
54+
#
55+
# Notes:
56+
# 1. A caller that saves an empty artifact with mime type exactly
57+
# application/octet-stream will still read back None. That collision is
58+
# inherent to using content shape as a tombstone; narrowing the match
59+
# shrinks the hole from every empty artifact to one specific mime type.
60+
# 2. This tombstone convention is in-memory only; other artifact services
61+
# (such as GcsArtifactService) do not perform this empty-payload check.
62+
_REWIND_TOMBSTONE = types.Part(
63+
inline_data=types.Blob(mime_type="application/octet-stream", data=b"")
64+
)
65+
66+
5067
class InMemoryArtifactService(BaseArtifactService, BaseModel):
5168
"""An in-memory implementation of the artifact service.
5269
@@ -208,11 +225,7 @@ async def load_artifact(
208225
version=parsed_uri.version,
209226
)
210227

211-
if (
212-
artifact_data == types.Part()
213-
or artifact_data == types.Part(text="")
214-
or (artifact_data.inline_data and not artifact_data.inline_data.data)
215-
):
228+
if artifact_data == types.Part() or artifact_data == _REWIND_TOMBSTONE:
216229
return None
217230
return artifact_data
218231

‎tests/unittests/artifacts/test_artifact_service.py‎

Lines changed: 43 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2954,12 +2954,15 @@ async def test_save_load_text_artifact(
29542954
@pytest.mark.parametrize(
29552955
"service_type",
29562956
[
2957+
ArtifactServiceType.IN_MEMORY,
29572958
ArtifactServiceType.GCS,
29582959
ArtifactServiceType.FILE,
29592960
],
29602961
)
2962+
@pytest.mark.parametrize("filename", ["empty.txt", "user:empty.txt"])
2963+
@pytest.mark.parametrize("version", [None, 0])
29612964
async def test_save_load_empty_text_artifact(
2962-
service_type, artifact_service_factory
2965+
service_type, artifact_service_factory, filename, version
29632966
):
29642967
"""Tests that empty text artifacts survive round-trip save/load."""
29652968
artifact_service = artifact_service_factory(service_type)
@@ -2969,20 +2972,57 @@ async def test_save_load_empty_text_artifact(
29692972
app_name="app0",
29702973
user_id="user0",
29712974
session_id="123",
2972-
filename="empty.txt",
2975+
filename=filename,
29732976
artifact=artifact,
29742977
)
29752978
loaded = await artifact_service.load_artifact(
29762979
app_name="app0",
29772980
user_id="user0",
29782981
session_id="123",
2979-
filename="empty.txt",
2982+
filename=filename,
2983+
version=version,
29802984
)
29812985
assert loaded is not None
29822986
assert loaded.text == ""
29832987
assert loaded.inline_data is None
29842988

29852989

2990+
@pytest.mark.parametrize(
2991+
"service_type",
2992+
[
2993+
ArtifactServiceType.IN_MEMORY,
2994+
ArtifactServiceType.GCS,
2995+
ArtifactServiceType.FILE,
2996+
],
2997+
)
2998+
@pytest.mark.parametrize("filename", ["empty.bin", "user:empty.bin"])
2999+
@pytest.mark.parametrize("version", [None, 0])
3000+
async def test_save_load_empty_bytes_artifact(
3001+
service_type, artifact_service_factory, filename, version
3002+
):
3003+
"""Tests that empty bytes artifacts survive round-trip save/load."""
3004+
artifact_service = artifact_service_factory(service_type)
3005+
artifact = types.Part.from_bytes(data=b"", mime_type="application/pdf")
3006+
3007+
await artifact_service.save_artifact(
3008+
app_name="app0",
3009+
user_id="user0",
3010+
session_id="123",
3011+
filename=filename,
3012+
artifact=artifact,
3013+
)
3014+
loaded = await artifact_service.load_artifact(
3015+
app_name="app0",
3016+
user_id="user0",
3017+
session_id="123",
3018+
filename=filename,
3019+
version=version,
3020+
)
3021+
assert loaded is not None
3022+
assert loaded.inline_data is not None
3023+
assert loaded.inline_data.data == b""
3024+
3025+
29863026
def _write_tampered_metadata(
29873027
root: Path,
29883028
*,

‎tests/unittests/runners/test_runner_rewind.py‎

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -76,7 +76,8 @@ def setup_method(self):
7676
)
7777

7878
@pytest.mark.asyncio
79-
async def test_rewind_async_with_state_and_artifacts(self):
79+
@pytest.mark.parametrize("initial_text", ["f1v0", ""])
80+
async def test_rewind_async_with_state_and_artifacts(self, initial_text):
8081
"""Tests rewind_async rewinds state and artifacts."""
8182
runner = self.runner
8283
user_id = "test_user"
@@ -93,7 +94,7 @@ async def test_rewind_async_with_state_and_artifacts(self):
9394
user_id=user_id,
9495
session_id=session_id,
9596
filename="f1",
96-
artifact=types.Part.from_text(text="f1v0"),
97+
artifact=types.Part.from_text(text=initial_text),
9798
)
9899
event1 = Event(
99100
invocation_id="invocation1",
@@ -177,7 +178,7 @@ async def test_rewind_async_with_state_and_artifacts(self):
177178
user_id=user_id,
178179
session_id=session_id,
179180
filename="f1",
180-
) == types.Part.from_text(text="f1v0")
181+
) == types.Part.from_text(text=initial_text)
181182
# f2 should not exist
182183
assert (
183184
await runner.artifact_service.load_artifact(

0 commit comments

Comments
 (0)