fix: resolve file refs in nested lists, only flag truly multimodal nested inputs

This commit is contained in:
Chesars 2026-03-22 01:37:18 -03:00
parent 912f08b61d
commit 11e4ec00f7
3 changed files with 26 additions and 15 deletions

View file

@ -215,18 +215,22 @@ class GoogleBatchEmbeddings(VertexLLM):
)
else:
input_list = [input] if isinstance(input, str) else input
has_file_refs = any(
_is_file_reference(e) for e in input_list if isinstance(e, str)
)
flat_elements = [
e
for item in input_list
for e in (item if isinstance(item, list) else [item])
if isinstance(e, str)
]
has_file_refs = any(_is_file_reference(e) for e in flat_elements)
if has_file_refs and not api_key:
raise ValueError(
"An API key is required to resolve Gemini file references (files/...). "
"Pass api_key= or set GEMINI_API_KEY."
)
resolved_files = {}
if api_key and _is_multimodal_input(input):
if api_key and has_file_refs:
resolved_files = self._resolve_file_references(
input=input, api_key=api_key, sync_handler=sync_handler
input=flat_elements, api_key=api_key, sync_handler=sync_handler
)
request_data = transform_openai_input_gemini_content(
input=input,
@ -320,18 +324,22 @@ class GoogleBatchEmbeddings(VertexLLM):
)
else:
input_list = [input] if isinstance(input, str) else input
has_file_refs = any(
_is_file_reference(e) for e in input_list if isinstance(e, str)
)
flat_elements = [
e
for item in input_list
for e in (item if isinstance(item, list) else [item])
if isinstance(e, str)
]
has_file_refs = any(_is_file_reference(e) for e in flat_elements)
if has_file_refs and not api_key:
raise ValueError(
"An API key is required to resolve Gemini file references (files/...). "
"Pass api_key= or set GEMINI_API_KEY."
)
resolved_files = {}
if api_key and _is_multimodal_input(input):
if api_key and has_file_refs:
resolved_files = await self._async_resolve_file_references(
input=input, api_key=api_key, async_handler=async_handler
input=flat_elements, api_key=api_key, async_handler=async_handler
)
data = transform_openai_input_gemini_content(
input=input,

View file

@ -130,8 +130,9 @@ def _is_multimodal_input(input: EmbeddingInput) -> bool:
for element in input:
if isinstance(element, list):
return True
if isinstance(element, str) and _is_multimodal_element(element):
if any(_is_multimodal_element(sub) for sub in element if isinstance(sub, str)):
return True
elif isinstance(element, str) and _is_multimodal_element(element):
return True
return False
@ -350,7 +351,8 @@ def process_response(
model_response.data = openai_embeddings
model_response.model = model
if _is_multimodal_input(input):
has_nested = isinstance(input, list) and any(isinstance(e, list) for e in input)
if _is_multimodal_input(input) or has_nested:
input_list = input if isinstance(input, list) else [input]
text_elements = []
for e in input_list:

View file

@ -44,8 +44,9 @@ class TestIsMultimodalInput:
def test_mixed_text_and_image(self):
assert _is_multimodal_input(["hello", IMAGE_DATA_URI]) is True
def test_nested_list_is_multimodal(self):
assert _is_multimodal_input([["text_a", "text_b"]]) is True
def test_nested_text_only_is_not_multimodal(self):
"""Nested list with only text is not multimodal."""
assert _is_multimodal_input([["text_a", "text_b"]]) is False
def test_nested_list_with_image_is_multimodal(self):
assert _is_multimodal_input([["a red shoe", IMAGE_DATA_URI]]) is True