diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py index 6070c70677..b2bf2c6eb5 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py @@ -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)) diff --git a/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py b/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py index ba741d6c2b..302facefcb 100644 --- a/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py +++ b/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py @@ -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 +