mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(vertex): bill image inputs at the image rate when usage lacks modality details
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
6c9fe65608
commit
e5845c17ff
3 changed files with 113 additions and 1 deletions
|
|
@ -268,6 +268,7 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
model_response=model_response,
|
||||
model=model,
|
||||
response_json=_json_response,
|
||||
resolved_files=resolved_files,
|
||||
)
|
||||
else:
|
||||
_predictions: Final = VertexAIBatchEmbeddingsResponseObject(**_json_response)
|
||||
|
|
@ -371,6 +372,7 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
model_response=model_response,
|
||||
model=model,
|
||||
response_json=_json_response,
|
||||
resolved_files=resolved_files,
|
||||
)
|
||||
else:
|
||||
_predictions: Final = VertexAIBatchEmbeddingsResponseObject(**_json_response)
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ Transformation logic from OpenAI /v1/embeddings format to Google AI Studio /batc
|
|||
Why separate file? Make it easy to see how transformation works
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
|
@ -297,6 +297,7 @@ def transform_openai_input_gemini_embed_content(
|
|||
return request_body
|
||||
|
||||
|
||||
_IMAGE_MIME_TYPES: Final = frozenset({"image/png", "image/jpeg"})
|
||||
_usage_metadata_adapter: Final = TypeAdapter(UsageMetadata)
|
||||
|
||||
|
||||
|
|
@ -309,6 +310,40 @@ def _parse_usage_metadata(raw_usage_metadata: object) -> UsageMetadata | None:
|
|||
return None
|
||||
|
||||
|
||||
def _flatten_input(input: GeminiEmbeddingInput) -> tuple[str, ...]:
|
||||
if isinstance(input, str):
|
||||
return (input,)
|
||||
return tuple(sub for element in input for sub in (element if isinstance(element, list) else [element]))
|
||||
|
||||
|
||||
def _is_image_element(
|
||||
element: str,
|
||||
resolved_files: Mapping[str, Mapping[str, str]],
|
||||
) -> bool:
|
||||
if element.startswith("data:") and ";base64," in element:
|
||||
try:
|
||||
mime_type, _ = _parse_data_url(element)
|
||||
except ValueError:
|
||||
return False
|
||||
return mime_type in _IMAGE_MIME_TYPES
|
||||
if _is_gcs_url(element):
|
||||
try:
|
||||
return _infer_mime_type_from_gcs_url(element) in _IMAGE_MIME_TYPES
|
||||
except ValueError:
|
||||
return False
|
||||
if _is_file_reference(element):
|
||||
file_info: Final = resolved_files.get(element)
|
||||
return file_info is not None and file_info.get("mime_type") in _IMAGE_MIME_TYPES
|
||||
return False
|
||||
|
||||
|
||||
def _count_input_images(
|
||||
input: GeminiEmbeddingInput,
|
||||
resolved_files: Mapping[str, Mapping[str, str]],
|
||||
) -> int:
|
||||
return sum(1 for element in _flatten_input(input) if _is_image_element(element, resolved_files))
|
||||
|
||||
|
||||
def _tokens_for_modality(details: Sequence[PromptTokensDetails], modality: str) -> int:
|
||||
return sum(detail["tokenCount"] for detail in details if detail["modality"] == modality)
|
||||
|
||||
|
|
@ -325,6 +360,7 @@ def _usage_from_embed_content_response(
|
|||
input: GeminiEmbeddingInput,
|
||||
model: str,
|
||||
raw_usage_metadata: object,
|
||||
resolved_files: Mapping[str, Mapping[str, str]],
|
||||
) -> Usage:
|
||||
usage_metadata: Final = _parse_usage_metadata(raw_usage_metadata)
|
||||
if usage_metadata is None:
|
||||
|
|
@ -334,6 +370,17 @@ def _usage_from_embed_content_response(
|
|||
total_tokens: Final = usage_metadata.get("totalTokenCount") or prompt_tokens
|
||||
|
||||
details: Final[Sequence[PromptTokensDetails]] = usage_metadata.get("promptTokensDetails") or ()
|
||||
if not details:
|
||||
image_tokens: Final = prompt_tokens if _count_input_images(input, resolved_files) else 0
|
||||
return Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
total_tokens=total_tokens,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=0,
|
||||
image_tokens=image_tokens,
|
||||
),
|
||||
)
|
||||
|
||||
text_tokens: Final = _tokens_for_modality(details, "TEXT")
|
||||
audio_tokens: Final = _tokens_for_modality(details, "AUDIO")
|
||||
image_tokens: Final = _tokens_for_modality(details, "IMAGE")
|
||||
|
|
@ -356,6 +403,7 @@ def process_embed_content_response(
|
|||
model_response: EmbeddingResponse,
|
||||
model: str,
|
||||
response_json: dict,
|
||||
resolved_files: Mapping[str, Mapping[str, str]] | None = None,
|
||||
) -> EmbeddingResponse:
|
||||
"""
|
||||
Process Gemini embedContent response (single embedding for multimodal input).
|
||||
|
|
@ -365,6 +413,7 @@ def process_embed_content_response(
|
|||
model_response: EmbeddingResponse to populate
|
||||
model: Model name
|
||||
response_json: Raw JSON response from embedContent endpoint
|
||||
resolved_files: Mapping of file references to resolved metadata
|
||||
|
||||
Returns:
|
||||
EmbeddingResponse with single embedding
|
||||
|
|
@ -386,6 +435,7 @@ def process_embed_content_response(
|
|||
input=input,
|
||||
model=model,
|
||||
raw_usage_metadata=response_json.get("usageMetadata"),
|
||||
resolved_files=resolved_files or {},
|
||||
)
|
||||
|
||||
return model_response
|
||||
|
|
|
|||
|
|
@ -411,6 +411,12 @@ class TestProcessEmbedContentResponseUsage:
|
|||
model_response=EmbeddingResponse(),
|
||||
model=self.MODEL,
|
||||
response_json=response_json,
|
||||
resolved_files={
|
||||
"files/img123": {
|
||||
"mime_type": "image/png",
|
||||
"uri": "https://example.com/img123",
|
||||
}
|
||||
},
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.image_tokens == 258
|
||||
assert result.usage.prompt_tokens_details.text_tokens == 0
|
||||
|
|
@ -437,6 +443,12 @@ class TestProcessEmbedContentResponseUsage:
|
|||
model_response=EmbeddingResponse(),
|
||||
model=self.MODEL,
|
||||
response_json=response_json,
|
||||
resolved_files={
|
||||
"files/clip1": {
|
||||
"mime_type": "audio/mpeg",
|
||||
"uri": "https://example.com/clip1",
|
||||
}
|
||||
},
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.audio_tokens == 64
|
||||
assert result.usage.prompt_tokens_details.image_tokens == 0
|
||||
|
|
@ -477,3 +489,51 @@ class TestProcessEmbedContentResponseUsage:
|
|||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
assert prompt_cost == pytest.approx(516 * 1.2e-5 + 64 * 6.5e-6)
|
||||
|
||||
def test_image_without_modality_details_uses_image_rate(self):
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1]},
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 258,
|
||||
"totalTokenCount": 258,
|
||||
},
|
||||
}
|
||||
result = process_embed_content_response(
|
||||
input=IMAGE_DATA_URI,
|
||||
model_response=EmbeddingResponse(),
|
||||
model=self.MODEL,
|
||||
response_json=response_json,
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.image_tokens == 258
|
||||
assert result.usage.prompt_tokens_details.text_tokens == 0
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model=self.MODEL,
|
||||
usage=result.usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
assert prompt_cost == pytest.approx(258 * 4.5e-7)
|
||||
|
||||
def test_text_without_modality_details_uses_text_rate(self):
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1]},
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 12,
|
||||
"totalTokenCount": 12,
|
||||
},
|
||||
}
|
||||
result = process_embed_content_response(
|
||||
input="a short caption",
|
||||
model_response=EmbeddingResponse(),
|
||||
model=self.MODEL,
|
||||
response_json=response_json,
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.text_tokens == 0
|
||||
assert result.usage.prompt_tokens_details.image_tokens == 0
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model=self.MODEL,
|
||||
usage=result.usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
assert prompt_cost == pytest.approx(12 * 2e-7)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue