Skip to content

Commit 43115b2

Browse files
blueeye-040DeanChensj
authored andcommitted
feat: support file_data URI references in GcsArtifactService
Merge #5322 `GcsArtifactService` previously raised `NotImplementedError` when saving a `types.Part` with `file_data`. This PR adds full `file_data` URI support to match the existing `InMemoryArtifactService` behavior. ### Link to Issue or Description of Change - Closes: #5230 **Problem:** `GcsArtifactService._save_artifact` raised `NotImplementedError` for `file_data` parts, while `InMemoryArtifactService` already handled them. This made `GcsArtifactService` unusable for any workflow that passes file URI references (e.g. `gs://` or `artifact://` URIs) as artifacts. **Solution:** Store `file_data` URI references as zero-byte GCS blobs with a `file_uri` custom metadata key — no content to upload, just the URI pointer. On load, check for that metadata key before downloading bytes and restore the original `types.Part(file_data=...)`. Also switched `_load_artifact` to use `bucket.get_blob()` (which populates metadata) instead of `bucket.blob()` (which does not). `artifact://` URIs are validated via `artifact_util.parse_artifact_uri` before saving. ### Testing Plan **Unit Tests:** - [x] I have added or updated unit tests for my change. - [x] All unit tests pass locally. Added 4 new tests in `tests/unittests/artifacts/test_artifact_service.py`: - `test_gcs_save_artifact_with_external_gcs_uri` — roundtrip save/load with a `gs://` URI - `test_gcs_save_artifact_with_artifact_ref_uri` — roundtrip save/load with an `artifact://` URI - `test_gcs_save_artifact_file_data_without_mime_type` — `file_data` with no mime_type - `test_gcs_save_artifact_file_data_missing_uri_raises` — empty `file_uri` raises `InputValidationError` Co-authored-by: Shangjie Chen <deanchen@google.com> COPYBARA_INTEGRATE_REVIEW=#5322 from blueeye-040:feat/gcs-file-data-support 99a932d PiperOrigin-RevId: 934629799
1 parent 9ecbaed commit 43115b2

3 files changed

Lines changed: 253 additions & 4 deletions

File tree

src/google/adk/artifacts/gcs_artifact_service.py

Lines changed: 57 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@
3232
from google.genai import types
3333
from typing_extensions import override
3434

35+
from . import artifact_util
3536
from ..errors.input_validation_error import InputValidationError
3637
from .base_artifact_service import ArtifactVersion
3738
from .base_artifact_service import BaseArtifactService
@@ -41,6 +42,8 @@
4142

4243
_GCS_DISPLAY_NAME_METADATA_KEY = "adkDisplayName"
4344
_GCS_IS_TEXT_METADATA_KEY = "adkIsText"
45+
_GCS_FILE_URI_METADATA_KEY = "adkFileUri"
46+
_GCS_FILE_MIME_TYPE_METADATA_KEY = "adkFileMimeType"
4447

4548

4649
class GcsArtifactService(BaseArtifactService):
@@ -243,9 +246,27 @@ def _save_artifact(
243246
content_type="text/plain",
244247
)
245248
elif artifact.file_data:
246-
raise NotImplementedError(
247-
"Saving artifact with file_data is not supported yet in"
248-
" GcsArtifactService."
249+
file_data = artifact.file_data
250+
assert file_data is not None
251+
file_uri = file_data.file_uri
252+
if not file_uri:
253+
raise InputValidationError("Artifact file_data must have a file_uri.")
254+
if artifact_util.is_artifact_ref(artifact):
255+
if not artifact_util.parse_artifact_uri(file_uri):
256+
raise InputValidationError(
257+
f"Invalid artifact reference URI: {file_uri}"
258+
)
259+
# Store the URI and mime_type (if any) as blob metadata; no content to upload.
260+
metadata = {
261+
**(blob.metadata or {}),
262+
_GCS_FILE_URI_METADATA_KEY: file_uri,
263+
}
264+
if file_data.mime_type:
265+
metadata[_GCS_FILE_MIME_TYPE_METADATA_KEY] = file_data.mime_type
266+
blob.metadata = metadata
267+
blob.upload_from_string(
268+
b"",
269+
content_type=file_data.mime_type or None,
249270
)
250271
else:
251272
raise InputValidationError(
@@ -280,6 +301,39 @@ def _load_artifact(
280301
if not blob:
281302
return None
282303

304+
# If the artifact was saved as a file_data URI reference, restore or resolve it.
305+
file_uri = None
306+
if blob.metadata:
307+
file_uri = blob.metadata.get(
308+
_GCS_FILE_URI_METADATA_KEY
309+
) or blob.metadata.get("file_uri")
310+
311+
if file_uri:
312+
if file_uri.startswith("artifact://"):
313+
parsed_uri = artifact_util.parse_artifact_uri(file_uri)
314+
if not parsed_uri:
315+
raise InputValidationError(
316+
f"Invalid artifact reference URI: {file_uri}"
317+
)
318+
return self._load_artifact(
319+
app_name=parsed_uri.app_name,
320+
user_id=parsed_uri.user_id,
321+
session_id=parsed_uri.session_id,
322+
filename=parsed_uri.filename,
323+
version=parsed_uri.version,
324+
)
325+
mime_type = None
326+
if blob.metadata:
327+
mime_type = blob.metadata.get(_GCS_FILE_MIME_TYPE_METADATA_KEY)
328+
if mime_type is None:
329+
mime_type = blob.content_type or None
330+
return types.Part(
331+
file_data=types.FileData(
332+
file_uri=file_uri,
333+
mime_type=mime_type,
334+
)
335+
)
336+
283337
artifact_bytes = blob.download_as_bytes()
284338
if blob.metadata and blob.metadata.get(_GCS_IS_TEXT_METADATA_KEY) == "true":
285339
return types.Part(text=artifact_bytes.decode("utf-8"))

src/google/adk/errors/input_validation_error.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
class InputValidationError(ValueError):
1919
"""Represents an error raised when user input fails validation."""
2020

21-
def __init__(self, message="Invalid input."):
21+
def __init__(self, message: str = "Invalid input.") -> None:
2222
"""Initializes the InputValidationError exception.
2323
2424
Args:

tests/unittests/artifacts/test_artifact_service.py

Lines changed: 195 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -938,6 +938,201 @@ def test_converts_text_dict(self):
938938
assert result.text == "hello world"
939939

940940

941+
# ---------------------------------------------------------------------------
942+
# GCS file_data (URI reference) tests
943+
# ---------------------------------------------------------------------------
944+
945+
946+
@pytest.mark.asyncio # type: ignore[untyped-decorator]
947+
async def test_gcs_save_artifact_with_external_gcs_uri() -> None:
948+
"""GcsArtifactService saves and loads a gs:// file_data URI reference."""
949+
service = mock_gcs_artifact_service() # type: ignore[no-untyped-call]
950+
artifact = types.Part(
951+
file_data=types.FileData(
952+
file_uri="gs://my-bucket/report.pdf",
953+
mime_type="application/pdf",
954+
)
955+
)
956+
957+
version = await service.save_artifact(
958+
app_name="app",
959+
user_id="user1",
960+
session_id="sess1",
961+
filename="report.pdf",
962+
artifact=artifact,
963+
)
964+
assert version == 0
965+
966+
loaded = await service.load_artifact(
967+
app_name="app",
968+
user_id="user1",
969+
session_id="sess1",
970+
filename="report.pdf",
971+
)
972+
assert loaded is not None
973+
assert loaded.file_data is not None
974+
assert loaded.file_data.file_uri == "gs://my-bucket/report.pdf"
975+
assert loaded.file_data.mime_type == "application/pdf"
976+
977+
978+
@pytest.mark.asyncio # type: ignore[untyped-decorator]
979+
async def test_gcs_save_artifact_with_artifact_ref_uri() -> None:
980+
"""GcsArtifactService saves and recursively loads an internal artifact:// URI reference."""
981+
service = mock_gcs_artifact_service() # type: ignore[no-untyped-call]
982+
983+
# Save the referenced (source) artifact first.
984+
source_artifact = types.Part(text="source content")
985+
await service.save_artifact(
986+
app_name="app",
987+
user_id="user1",
988+
session_id="sess1",
989+
filename="source.txt",
990+
artifact=source_artifact,
991+
)
992+
993+
artifact_ref_uri = "artifact://apps/app/users/user1/sessions/sess1/artifacts/source.txt/versions/0"
994+
artifact = types.Part(
995+
file_data=types.FileData(
996+
file_uri=artifact_ref_uri,
997+
mime_type="text/plain",
998+
)
999+
)
1000+
1001+
version = await service.save_artifact(
1002+
app_name="app",
1003+
user_id="user1",
1004+
session_id="sess1",
1005+
filename="ref.txt",
1006+
artifact=artifact,
1007+
)
1008+
assert version == 0
1009+
1010+
loaded = await service.load_artifact(
1011+
app_name="app",
1012+
user_id="user1",
1013+
session_id="sess1",
1014+
filename="ref.txt",
1015+
)
1016+
assert loaded is not None
1017+
assert loaded.text == "source content"
1018+
1019+
1020+
@pytest.mark.asyncio # type: ignore[untyped-decorator]
1021+
async def test_gcs_save_artifact_file_data_without_mime_type() -> None:
1022+
"""GcsArtifactService handles file_data with no mime_type."""
1023+
service = mock_gcs_artifact_service() # type: ignore[no-untyped-call]
1024+
artifact = types.Part(
1025+
file_data=types.FileData(file_uri="gs://my-bucket/data.bin")
1026+
)
1027+
1028+
version = await service.save_artifact(
1029+
app_name="app",
1030+
user_id="user1",
1031+
session_id="sess1",
1032+
filename="data.bin",
1033+
artifact=artifact,
1034+
)
1035+
assert version == 0
1036+
1037+
loaded = await service.load_artifact(
1038+
app_name="app",
1039+
user_id="user1",
1040+
session_id="sess1",
1041+
filename="data.bin",
1042+
)
1043+
assert loaded is not None
1044+
assert loaded.file_data is not None
1045+
assert loaded.file_data.file_uri == "gs://my-bucket/data.bin"
1046+
1047+
1048+
@pytest.mark.asyncio # type: ignore[untyped-decorator]
1049+
async def test_gcs_save_artifact_file_data_missing_uri_raises() -> None:
1050+
"""GcsArtifactService raises InputValidationError when file_uri is empty."""
1051+
service = mock_gcs_artifact_service() # type: ignore[no-untyped-call]
1052+
artifact = types.Part(file_data=types.FileData(file_uri=""))
1053+
1054+
with pytest.raises(InputValidationError):
1055+
await service.save_artifact(
1056+
app_name="app",
1057+
user_id="user1",
1058+
session_id="sess1",
1059+
filename="empty.bin",
1060+
artifact=artifact,
1061+
)
1062+
1063+
1064+
@pytest.mark.asyncio # type: ignore[untyped-decorator]
1065+
async def test_gcs_save_artifact_file_data_invalid_uri_raises() -> None:
1066+
"""GcsArtifactService raises InputValidationError when file_uri is an invalid artifact:// URI template."""
1067+
service = mock_gcs_artifact_service() # type: ignore[no-untyped-call]
1068+
artifact = types.Part(
1069+
file_data=types.FileData(
1070+
file_uri="artifact://apps/app/invalid",
1071+
mime_type="text/plain",
1072+
)
1073+
)
1074+
1075+
with pytest.raises(InputValidationError):
1076+
await service.save_artifact(
1077+
app_name="app",
1078+
user_id="user1",
1079+
session_id="sess1",
1080+
filename="invalid_ref.txt",
1081+
artifact=artifact,
1082+
)
1083+
1084+
1085+
@pytest.mark.asyncio # type: ignore[untyped-decorator]
1086+
async def test_gcs_save_artifact_metadata_namespacing_and_mime() -> None:
1087+
"""GcsArtifactService saves file_data using namespaced metadata keys."""
1088+
service = mock_gcs_artifact_service() # type: ignore[no-untyped-call]
1089+
artifact = types.Part(
1090+
file_data=types.FileData(
1091+
file_uri="gs://my-bucket/report.pdf",
1092+
mime_type="application/pdf",
1093+
)
1094+
)
1095+
1096+
await service.save_artifact(
1097+
app_name="app",
1098+
user_id="user1",
1099+
session_id="sess1",
1100+
filename="report.pdf",
1101+
artifact=artifact,
1102+
)
1103+
1104+
blob_name = service._get_blob_name("app", "user1", "report.pdf", 0, "sess1")
1105+
blob = service.bucket.get_blob(blob_name)
1106+
assert blob is not None
1107+
assert blob.metadata.get("adkFileUri") == "gs://my-bucket/report.pdf"
1108+
assert blob.metadata.get("adkFileMimeType") == "application/pdf"
1109+
assert "file_uri" not in blob.metadata
1110+
1111+
1112+
@pytest.mark.asyncio # type: ignore[untyped-decorator]
1113+
async def test_gcs_load_artifact_file_data_fallback_compatibility() -> None:
1114+
"""GcsArtifactService loads file_data with old file_uri metadata key for backward compatibility."""
1115+
service = mock_gcs_artifact_service() # type: ignore[no-untyped-call]
1116+
blob_name = service._get_blob_name(
1117+
"app", "user1", "old_report.pdf", 0, "sess1"
1118+
)
1119+
blob = service.bucket.blob(blob_name)
1120+
# Manually setup metadata with old key
1121+
blob.metadata = {"file_uri": "gs://my-bucket/old_report.pdf"}
1122+
blob.upload_from_string(b"", content_type="application/pdf")
1123+
1124+
loaded = await service.load_artifact(
1125+
app_name="app",
1126+
user_id="user1",
1127+
session_id="sess1",
1128+
filename="old_report.pdf",
1129+
)
1130+
assert loaded is not None
1131+
assert loaded.file_data is not None
1132+
assert loaded.file_data.file_uri == "gs://my-bucket/old_report.pdf"
1133+
assert loaded.file_data.mime_type == "application/pdf"
1134+
1135+
9411136
@pytest.mark.asyncio
9421137
@pytest.mark.parametrize(
9431138
"service_type",

0 commit comments

Comments
 (0)