fix(gemini): propagate batchEmbedContents usageMetadata for gemini-embedding-2 multimodal usage

This commit is contained in:
Devin AI 2026-07-17 11:11:35 +00:00
parent 4d33964898
commit d33937dee4
4 changed files with 102 additions and 1 deletions

View file

@ -274,6 +274,8 @@ class GoogleBatchEmbeddings(VertexLLM):
model_response=model_response,
_predictions=_predictions,
input=input,
raw_usage_metadata=_json_response.get("usageMetadata"),
resolved_files=resolved_files,
)
async def async_batch_embeddings(
@ -378,4 +380,6 @@ class GoogleBatchEmbeddings(VertexLLM):
model_response=model_response,
_predictions=_predictions,
input=input,
raw_usage_metadata=_json_response.get("usageMetadata"),
resolved_files=resolved_files,
)

View file

@ -371,7 +371,9 @@ def _usage_from_embed_content_response(
prompt_tokens = usage_metadata.get("promptTokenCount", 0)
total_tokens = usage_metadata.get("totalTokenCount") or prompt_tokens
details: Sequence[PromptTokensDetails] = usage_metadata.get("promptTokensDetails") or ()
details: Sequence[PromptTokensDetails] = (
usage_metadata.get("promptTokensDetails") or usage_metadata.get("promptTokenDetails") or ()
)
text_tokens = _tokens_for_modality(details, "TEXT")
audio_tokens = _tokens_for_modality(details, "AUDIO")
video_tokens = _tokens_for_modality(details, "VIDEO")
@ -449,6 +451,8 @@ def process_response(
model_response: EmbeddingResponse,
model: str,
_predictions: VertexAIBatchEmbeddingsResponseObject,
raw_usage_metadata: object = None,
resolved_files: Mapping[str, Mapping[str, str]] | None = None,
) -> EmbeddingResponse:
openai_embeddings: List[Embedding] = []
for idx, embedding in enumerate(_predictions["embeddings"]):
@ -462,6 +466,15 @@ def process_response(
model_response.data = openai_embeddings
model_response.model = model
if _parse_usage_metadata(raw_usage_metadata) is not None:
model_response.usage = _usage_from_embed_content_response(
input=input,
model=model,
raw_usage_metadata=raw_usage_metadata,
resolved_files=resolved_files or {},
)
return model_response
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]

View file

@ -302,6 +302,7 @@ class UsageMetadata(TypedDict, total=False):
toolUsePromptTokenCount: int
toolUsePromptTokensDetails: List[PromptTokensDetails]
promptTokensDetails: List[PromptTokensDetails]
promptTokenDetails: List[PromptTokensDetails]
cacheTokensDetails: List[PromptTokensDetails]
thoughtsTokenCount: int
responseTokensDetails: List[PromptTokensDetails]

View file

@ -282,6 +282,70 @@ class TestProcessResponse:
assert len(result.data) == 1
assert result.usage.prompt_tokens > 0
def test_batch_image_usage_from_metadata_singular_key(self):
"""batchEmbedContents returns usageMetadata with the singular
`promptTokenDetails` key; image-only input must bill 258 tokens / 1 image
instead of the previous prompt_tokens=0."""
predictions: VertexAIBatchEmbeddingsResponseObject = {
"embeddings": [{"values": [0.1, 0.2]}]
}
result = process_response(
input=[IMAGE_DATA_URI],
model_response=EmbeddingResponse(),
model="gemini-embedding-2",
_predictions=predictions,
raw_usage_metadata={
"promptTokenCount": 258,
"promptTokenDetails": [{"modality": "IMAGE", "tokenCount": 258}],
},
)
assert result.usage.prompt_tokens == 258
assert result.usage.total_tokens == 258
assert result.usage.prompt_tokens_details.image_count == 1
prompt_cost, _ = generic_cost_per_token(
model="gemini-embedding-2",
usage=result.usage,
custom_llm_provider="gemini",
)
assert prompt_cost > 0
def test_batch_mixed_text_image_aggregates_usage_from_metadata(self):
"""Mixed text + image over batchEmbedContents aggregates the native
per-modality usageMetadata rather than counting only the text element."""
predictions: VertexAIBatchEmbeddingsResponseObject = {
"embeddings": [{"values": [0.1, 0.2]}, {"values": [0.3, 0.4]}]
}
result = process_response(
input=["hello", IMAGE_DATA_URI],
model_response=EmbeddingResponse(),
model="gemini-embedding-2",
_predictions=predictions,
raw_usage_metadata={
"promptTokenCount": 259,
"promptTokenDetails": [
{"modality": "TEXT", "tokenCount": 1},
{"modality": "IMAGE", "tokenCount": 258},
],
},
)
assert result.usage.prompt_tokens == 259
assert result.usage.prompt_tokens_details.image_count == 1
assert result.usage.prompt_tokens_details.text_tokens == 1
def test_batch_without_metadata_falls_back_to_token_counter(self):
"""No usageMetadata: text-only batch still counts tokens locally."""
predictions: VertexAIBatchEmbeddingsResponseObject = {
"embeddings": [{"values": [0.1, 0.2]}, {"values": [0.3, 0.4]}]
}
result = process_response(
input=["hello", "world"],
model_response=EmbeddingResponse(),
model="gemini-embedding-2",
_predictions=predictions,
)
assert result.usage.prompt_tokens > 0
def test_nested_empty_list_raises(self):
with pytest.raises(ValueError, match="must not be empty"):
transform_openai_input_gemini_content(
@ -333,6 +397,25 @@ class TestProcessEmbedContentResponseUsage:
)
assert prompt_cost > 0
def test_singular_prompt_token_details_key_is_parsed(self):
"""The live embedContent endpoint returns the singular `promptTokenDetails`
key (not `promptTokensDetails`); the per-modality breakdown must still bill."""
response_json = {
"embedding": {"values": [0.1, 0.2, 0.3]},
"usageMetadata": {
"promptTokenCount": 64,
"promptTokenDetails": [{"modality": "AUDIO", "tokenCount": 64}],
},
}
result = process_embed_content_response(
input=["data:audio/mpeg;base64,QUJD"],
model_response=EmbeddingResponse(),
model=self.MODEL,
response_json=response_json,
)
assert result.usage.prompt_tokens == 64
assert result.usage.prompt_tokens_details.audio_tokens == 64
def test_text_modality_detail_populated(self):
response_json = {
"embedding": {"values": [0.1, 0.2]},