fix(batches): keep modality token details in raw vertex batch usage

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-09-15 01:57:50 +00:00
parent 2b32f586c0
commit c0c5044c45
3 changed files with 76 additions and 2 deletions

View file

@ -3,14 +3,14 @@ from collections.abc import Iterable, Iterator, Mapping
from dataclasses import dataclass
from dataclasses import replace as dataclasses_replace
from enum import Enum
from typing import Any, Final, Literal
from typing import Any, Final, Literal, cast
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
from litellm.litellm_core_utils.llm_cost_calc.utils import parse_prompt_tokens_details
from litellm.types.llms.openai import Batch
from litellm.types.utils import ModelInfo, Usage
from litellm.types.utils import ModelInfo, PromptTokensDetailsWrapper, Usage
from litellm.utils import token_counter
@ -310,6 +310,35 @@ def _aggregate_batch_cost_usage_models(
)
def _vertex_prompt_tokens_details(
usage_metadata: Mapping[str, object],
) -> PromptTokensDetailsWrapper | None:
raw_details: Final = usage_metadata.get("promptTokensDetails")
if not isinstance(raw_details, list):
return None
raw_list: Final = cast(list[object], raw_details)
if not all(isinstance(detail, Mapping) for detail in raw_list):
return None
details: Final = tuple(cast(Mapping[str, object], detail) for detail in raw_list)
normalized: Final = tuple(
(modality.upper(), token_count)
for detail in details
if isinstance(modality := detail.get("modality"), str)
and isinstance(token_count := detail.get("tokenCount"), int)
)
if len(normalized) != len(details):
return None
return PromptTokensDetailsWrapper(
text_tokens=sum(token_count for modality, token_count in normalized if modality in ("TEXT", "DOCUMENT")),
audio_tokens=sum(token_count for modality, token_count in normalized if modality == "AUDIO"),
image_tokens=sum(token_count for modality, token_count in normalized if modality == "IMAGE"),
video_tokens=sum(token_count for modality, token_count in normalized if modality == "VIDEO"),
)
def calculate_vertex_ai_batch_cost_and_usage(
vertex_ai_batch_responses: list[dict],
model_name: str | None = None,
@ -356,6 +385,7 @@ def calculate_vertex_ai_batch_cost_and_usage(
prompt_tokens=_prompt,
completion_tokens=_completion,
total_tokens=_total,
prompt_tokens_details=_vertex_prompt_tokens_details(cast(Mapping[str, object], usage_metadata)),
)
try:

View file

@ -695,6 +695,38 @@ def test_vertex_cost_and_usage_aggregation(monkeypatch):
assert result.failed_requests == 0
def test_vertex_batch_usage_preserves_modality_token_details(monkeypatch):
monkeypatch.setitem(
litellm.model_cost,
"vertex_ai/gemini-embedding-2",
{
"input_cost_per_token_batches": 1e-7,
"input_cost_per_audio_token_batches": 3.25e-6,
"input_cost_per_image_token_batches": 2.25e-7,
"input_cost_per_video_token_batches": 6e-6,
},
)
responses = [
{
"response": {
"usageMetadata": {
"promptTokenCount": 84,
"candidatesTokenCount": 0,
"totalTokenCount": 84,
"promptTokensDetails": [
{"modality": "AUDIO", "tokenCount": 64},
{"modality": "TEXT", "tokenCount": 20},
],
}
}
}
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(responses, "gemini-embedding-2")
assert result.prompt_cost == pytest.approx(64 * 3.25e-6 + 20 * 1e-7)
def test_vertex_cost_skips_none_response_body(monkeypatch):
import litellm.cost_calculator as cc

View file

@ -29857,6 +29857,8 @@ export interface components {
input_cost_per_audio_per_second_above_128k_tokens?: number | null;
/** Input Cost Per Audio Token */
input_cost_per_audio_token?: number | null;
/** Input Cost Per Audio Token Batches */
input_cost_per_audio_token_batches?: number | null;
/** Input Cost Per Character */
input_cost_per_character?: number | null;
/** Input Cost Per Character Above 128K Tokens */
@ -29867,6 +29869,8 @@ export interface components {
input_cost_per_image_above_128k_tokens?: number | null;
/** Input Cost Per Image Token */
input_cost_per_image_token?: number | null;
/** Input Cost Per Image Token Batches */
input_cost_per_image_token_batches?: number | null;
/** Input Cost Per Pixel */
input_cost_per_pixel?: number | null;
/** Input Cost Per Query */
@ -29909,6 +29913,8 @@ export interface components {
input_cost_per_video_per_second_above_8s_interval?: number | null;
/** Input Cost Per Video Token */
input_cost_per_video_token?: number | null;
/** Input Cost Per Video Token Batches */
input_cost_per_video_token_batches?: number | null;
/** Itpm */
itpm?: number | null;
/** Keepalive Seconds */
@ -40071,6 +40077,8 @@ export interface components {
input_cost_per_audio_per_second_above_128k_tokens?: number | null;
/** Input Cost Per Audio Token */
input_cost_per_audio_token?: number | null;
/** Input Cost Per Audio Token Batches */
input_cost_per_audio_token_batches?: number | null;
/** Input Cost Per Character */
input_cost_per_character?: number | null;
/** Input Cost Per Character Above 128K Tokens */
@ -40081,6 +40089,8 @@ export interface components {
input_cost_per_image_above_128k_tokens?: number | null;
/** Input Cost Per Image Token */
input_cost_per_image_token?: number | null;
/** Input Cost Per Image Token Batches */
input_cost_per_image_token_batches?: number | null;
/** Input Cost Per Pixel */
input_cost_per_pixel?: number | null;
/** Input Cost Per Query */
@ -40123,6 +40133,8 @@ export interface components {
input_cost_per_video_per_second_above_8s_interval?: number | null;
/** Input Cost Per Video Token */
input_cost_per_video_token?: number | null;
/** Input Cost Per Video Token Batches */
input_cost_per_video_token_batches?: number | null;
/** Itpm */
itpm?: number | null;
/** Keepalive Seconds */