Add support for gcs url for vertex ai embeddings

This commit is contained in:
Sameer Kankute
2026-03-11 10:58:04 +05:30
parent 2c4a495619
commit d25b8e6d00
2 changed files with 118 additions and 4 deletions
@@ -35,6 +35,46 @@ def _is_file_reference(s: str) -> bool:
return isinstance(s, str) and s.startswith("files/")
def _is_gcs_url(s: str) -> bool:
"""Check if string is a GCS URL (gs://...)."""
return isinstance(s, str) and s.startswith("gs://")
def _infer_mime_type_from_gcs_url(gcs_url: str) -> str:
"""
Infer MIME type from GCS URL file extension.
Args:
gcs_url: GCS URL like gs://bucket/path/to/file.png
Returns:
str: Inferred MIME type
Raises:
ValueError: If file extension is not supported
"""
extension_to_mime = {
".png": "image/png",
".jpg": "image/jpeg",
".jpeg": "image/jpeg",
".mp3": "audio/mpeg",
".wav": "audio/wav",
".mp4": "video/mp4",
".mov": "video/quicktime",
".pdf": "application/pdf",
}
gcs_url_lower = gcs_url.lower()
for ext, mime_type in extension_to_mime.items():
if gcs_url_lower.endswith(ext):
return mime_type
raise ValueError(
f"Unable to infer MIME type from GCS URL: {gcs_url}. "
f"Supported extensions: {', '.join(extension_to_mime.keys())}"
)
def _parse_data_url(data_url: str) -> Tuple[str, str]:
"""
Parse a data URL to extract the media type and base64 data.
@@ -76,13 +116,13 @@ def _parse_data_url(data_url: str) -> Tuple[str, str]:
def _is_multimodal_input(input: EmbeddingInput) -> bool:
"""
Check if the input contains multimodal data (data URIs or file references).
Check if the input contains multimodal data (data URIs, file references, or GCS URLs).
Args:
input: EmbeddingInput (str or List[str])
Returns:
bool: True if any element is a data URI or file reference
bool: True if any element is a data URI, file reference, or GCS URL
"""
if isinstance(input, str):
input_list = [input]
@@ -95,6 +135,8 @@ def _is_multimodal_input(input: EmbeddingInput) -> bool:
return True
if _is_file_reference(element):
return True
if _is_gcs_url(element):
return True
return False
@@ -166,15 +208,22 @@ def transform_openai_input_gemini_embed_content(
mime_type, base64_data = _parse_data_url(element)
blob: BlobType = {"mime_type": mime_type, "data": base64_data}
parts.append(PartType(inline_data=blob))
elif _is_gcs_url(element):
mime_type = _infer_mime_type_from_gcs_url(element)
file_data: FileDataType = {
"mime_type": mime_type,
"file_uri": element,
}
parts.append(PartType(file_data=file_data))
elif _is_file_reference(element):
if element not in resolved_files:
raise ValueError(f"File reference {element} not resolved")
file_info = resolved_files[element]
file_data: FileDataType = {
file_data_ref: FileDataType = {
"mime_type": file_info["mime_type"],
"file_uri": file_info["uri"],
}
parts.append(PartType(file_data=file_data))
parts.append(PartType(file_data=file_data_ref))
else:
parts.append(PartType(text=element))
@@ -391,3 +391,68 @@ def test_dimensions_mapped_to_output_dimensionality():
assert result["outputDimensionality"] == 768
assert "dimensions" not in result
def test_is_gcs_url():
"""Test GCS URL detection."""
from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
_is_gcs_url,
)
assert _is_gcs_url("gs://my-bucket/path/to/file.png") is True
assert _is_gcs_url("gs://bucket/image.jpg") is True
assert _is_gcs_url("https://storage.googleapis.com/bucket/file.png") is False
assert _is_gcs_url("files/abc123") is False
assert _is_gcs_url("data:image/png;base64,abc") is False
assert _is_gcs_url("regular text") is False
def test_infer_mime_type_from_gcs_url():
"""Test MIME type inference from GCS URL."""
from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
_infer_mime_type_from_gcs_url,
)
assert _infer_mime_type_from_gcs_url("gs://bucket/image.png") == "image/png"
assert _infer_mime_type_from_gcs_url("gs://bucket/photo.jpg") == "image/jpeg"
assert _infer_mime_type_from_gcs_url("gs://bucket/photo.JPEG") == "image/jpeg"
assert _infer_mime_type_from_gcs_url("gs://bucket/audio.mp3") == "audio/mpeg"
assert _infer_mime_type_from_gcs_url("gs://bucket/audio.wav") == "audio/wav"
assert _infer_mime_type_from_gcs_url("gs://bucket/video.mp4") == "video/mp4"
assert _infer_mime_type_from_gcs_url("gs://bucket/video.mov") == "video/quicktime"
assert _infer_mime_type_from_gcs_url("gs://bucket/doc.pdf") == "application/pdf"
with pytest.raises(ValueError, match="Unable to infer MIME type"):
_infer_mime_type_from_gcs_url("gs://bucket/file.txt")
def test_transform_multimodal_with_gcs_url():
"""Test transformation with GCS URL."""
input_data = [
"Describe this image",
"gs://my-bucket/images/photo.png"
]
result = transform_openai_input_gemini_embed_content(
input=input_data,
model="gemini-embedding-2-preview",
optional_params={},
resolved_files=None,
)
parts = result["content"]["parts"]
assert len(parts) == 2
assert parts[0]["text"] == "Describe this image"
assert parts[1]["file_data"]["mime_type"] == "image/png"
assert parts[1]["file_data"]["file_uri"] == "gs://my-bucket/images/photo.png"
def test_multimodal_input_detection_with_gcs():
"""Test that GCS URLs are detected as multimodal."""
from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
_is_multimodal_input,
)
assert _is_multimodal_input(["text", "gs://bucket/file.png"]) is True
assert _is_multimodal_input("gs://bucket/video.mp4") is True
assert _is_multimodal_input(["just text", "more text"]) is False