diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 87f8fd3946e..397bc0a35a2 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -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: diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index 768ea332677..8b04d7af70a 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -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 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 8aa05cf8c7c..73245806bb2 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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 */