From 9accd522c1a5009b43a960624d0bc5f55b2d728c Mon Sep 17 00:00:00 2001 From: pragnyanramtha Date: Tue, 19 May 2026 01:50:38 +0000 Subject: [PATCH] fix: keep video spend fallback conservative --- litellm/cost_calculator.py | 175 +++++++++++++++++- litellm/litellm_core_utils/litellm_logging.py | 1 + tests/local_testing/test_completion_cost.py | 107 +++++++++++ .../test_litellm_logging.py | 65 +++++++ 4 files changed, 346 insertions(+), 2 deletions(-) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 2257861aff6..e0b5f60a68a 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -1,9 +1,10 @@ # What is this? ## File for 'response_cost' calculation in Logging import logging +import math import time from functools import lru_cache -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, List, Literal, Mapping, Optional, Tuple, Union, cast from httpx import Response from pydantic import BaseModel @@ -31,6 +32,10 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import ( get_billable_input_tokens, select_cost_metric_for_model, ) +from litellm.litellm_core_utils.token_counter import ( + get_token_count_for_limit_enforcement, + messages_contain_video_url, +) from litellm.llms.anthropic.cost_calculation import ( cost_per_token as anthropic_cost_per_token, ) @@ -165,6 +170,13 @@ _SEARCH_CALL_TYPES = frozenset( } ) +_CHAT_COMPLETION_CALL_TYPES = frozenset( + { + CallTypes.completion.value, + CallTypes.acompletion.value, + } +) + _AREALTIME_CALL_TYPE = CallTypes.arealtime.value _MCP_CALL_TYPE = CallTypes.call_mcp_tool.value @@ -285,6 +297,136 @@ def _transcription_usage_has_token_details( return (prompt_tokens_val > 0) or (completion_tokens_val > 0) +def _is_positive_finite_number(value: Any) -> bool: + return ( + not isinstance(value, bool) + and isinstance(value, (int, float)) + and math.isfinite(value) + and value > 0 + ) + + +def _get_metadata_model_infos( + litellm_logging_obj: Optional[LitellmLoggingObject], +) -> List[Mapping[str, Any]]: + litellm_params = getattr(litellm_logging_obj, "litellm_params", None) + if not isinstance(litellm_params, dict): + return [] + + model_infos: List[Mapping[str, Any]] = [] + for metadata_key in ("litellm_metadata", "metadata"): + metadata = litellm_params.get(metadata_key, {}) or {} + if not isinstance(metadata, dict): + continue + model_info = metadata.get("model_info", {}) or {} + if isinstance(model_info, Mapping): + model_infos.append(model_info) + return model_infos + + +def _get_max_input_tokens_for_cost_fallback( + model: Optional[str], + custom_llm_provider: Optional[str], + litellm_logging_obj: Optional[LitellmLoggingObject], +) -> Optional[Union[int, float]]: + metadata_model_infos = _get_metadata_model_infos( + litellm_logging_obj=litellm_logging_obj + ) + for model_info in metadata_model_infos: + for token_limit_key in ("max_input_tokens", "max_tokens"): + token_limit = model_info.get(token_limit_key) + if _is_positive_finite_number(token_limit): + return cast(Union[int, float], token_limit) + + if model is None: + return None + + provider = custom_llm_provider + model_names_to_try = [model] + if "/" in model: + provider_from_model, model_without_provider = model.split("/", 1) + provider = provider or provider_from_model + model_names_to_try.append(model_without_provider) + + for model_name in model_names_to_try: + try: + model_info = litellm.get_model_info( + model=model_name, custom_llm_provider=provider + ) + except Exception: + continue + for token_limit_key in ("max_input_tokens", "max_tokens"): + token_limit = model_info.get(token_limit_key) + if _is_positive_finite_number(token_limit): + return cast(Union[int, float], token_limit) + return None + + +def _usage_has_token_counts(usage_object: Optional[Usage]) -> bool: + if usage_object is None: + return False + + for attr in ("prompt_tokens", "completion_tokens", "total_tokens"): + if _is_positive_finite_number(getattr(usage_object, attr, 0)): + return True + + prompt_details = getattr(usage_object, "prompt_tokens_details", None) + if prompt_details is not None: + for attr in ("audio_tokens", "cached_tokens", "text_tokens"): + if _is_positive_finite_number(getattr(prompt_details, attr, 0)): + return True + + completion_details = getattr(usage_object, "completion_tokens_details", None) + if completion_details is not None: + for attr in ("audio_tokens", "reasoning_tokens", "text_tokens"): + if _is_positive_finite_number(getattr(completion_details, attr, 0)): + return True + + return False + + +def _usage_with_conservative_prompt_tokens( + usage_object: Optional[Usage], + prompt_tokens: int, + completion_tokens: int, +) -> Usage: + if usage_object is not None: + usage_dict = usage_object.model_dump() + else: + usage_dict = {} + + usage_dict["prompt_tokens"] = prompt_tokens + usage_dict["completion_tokens"] = completion_tokens + usage_dict["total_tokens"] = max( + int(usage_dict.get("total_tokens") or 0), + prompt_tokens + completion_tokens, + ) + return Usage(**usage_dict) + + +def _get_conservative_video_prompt_tokens_for_cost_fallback( + *, + prompt_tokens: int, + messages: List, + model: Optional[str], + custom_llm_provider: Optional[str], + litellm_logging_obj: Optional[LitellmLoggingObject], +) -> int: + if not messages_contain_video_url(messages): + return prompt_tokens + + token_limit = _get_max_input_tokens_for_cost_fallback( + model=model, + custom_llm_provider=custom_llm_provider, + litellm_logging_obj=litellm_logging_obj, + ) + return get_token_count_for_limit_enforcement( + input_tokens=prompt_tokens, + messages=messages, + token_limit=token_limit, + ) + + def cost_per_token( # noqa: PLR0915 model: str = "", prompt_tokens: int = 0, @@ -1094,7 +1236,7 @@ def completion_cost( # noqa: PLR0915 completion_response=None, model: Optional[str] = None, prompt="", - messages: List = [], + messages: Optional[List] = None, completion="", total_time: Optional[float] = 0.0, # used for replicate, sagemaker call_type: Optional[CallTypesLiteral] = None, @@ -1147,6 +1289,7 @@ def completion_cost( # noqa: PLR0915 - For un-mapped Replicate models, the cost is calculated based on the total time used for the request. """ try: + messages = messages or [] call_type = _infer_call_type(call_type, completion_response) or "completion" if ( @@ -1525,6 +1668,32 @@ def completion_cost( # noqa: PLR0915 return MCPCostCalculator.calculate_mcp_tool_call_cost( litellm_logging_obj=litellm_logging_obj ) + + if ( + call_type in _CHAT_COMPLETION_CALL_TYPES + and not _usage_has_token_counts(cost_per_token_usage_object) + ): + # Video metadata is client-provided. When provider usage is + # missing, keep spend reconciliation conservative. + conservative_prompt_tokens = ( + _get_conservative_video_prompt_tokens_for_cost_fallback( + prompt_tokens=prompt_tokens or 0, + messages=messages, + model=model, + custom_llm_provider=custom_llm_provider, + litellm_logging_obj=litellm_logging_obj, + ) + ) + if conservative_prompt_tokens != (prompt_tokens or 0): + prompt_tokens = conservative_prompt_tokens + cost_per_token_usage_object = ( + _usage_with_conservative_prompt_tokens( + usage_object=cost_per_token_usage_object, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens or 0, + ) + ) + # Calculate cost based on prompt_tokens, completion_tokens if ( "togethercomputer" in model @@ -1804,6 +1973,7 @@ def response_cost_calculator( cache_hit: Optional[bool] = None, base_model: Optional[str] = None, custom_pricing: Optional[bool] = None, + messages: Optional[List] = None, prompt: str = "", standard_built_in_tools_params: Optional[StandardBuiltInToolsParams] = None, litellm_model_name: Optional[str] = None, @@ -1838,6 +2008,7 @@ def response_cost_calculator( optional_params=optional_params, custom_pricing=custom_pricing, base_model=base_model, + messages=messages, prompt=prompt, standard_built_in_tools_params=standard_built_in_tools_params, litellm_model_name=litellm_model_name, diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 876f1b167db..0de6777d4f6 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -1543,6 +1543,7 @@ class Logging(LiteLLMLoggingBaseClass): "call_type": self.call_type, "optional_params": self.optional_params, "custom_pricing": custom_pricing, + "messages": self.messages or [], "prompt": prompt, "standard_built_in_tools_params": self.standard_built_in_tools_params, "router_model_id": router_model_id, diff --git a/tests/local_testing/test_completion_cost.py b/tests/local_testing/test_completion_cost.py index cf0c645615d..cadafc34035 100644 --- a/tests/local_testing/test_completion_cost.py +++ b/tests/local_testing/test_completion_cost.py @@ -982,6 +982,113 @@ def test_completion_cost_azure_common_deployment_name(): assert "azure/gpt-4" == mock_client.call_args.kwargs["base_model"] +def test_completion_cost_uses_conservative_video_fallback_without_usage(): + model = "openai/test-video-cost-fallback" + input_cost_per_token = 0.25 + max_input_tokens = 8 + litellm.register_model( + model_cost={ + model: { + "input_cost_per_token": input_cost_per_token, + "output_cost_per_token": 0.0, + "max_tokens": max_input_tokens, + "max_input_tokens": max_input_tokens, + "max_output_tokens": 4, + "litellm_provider": "openai", + "mode": "chat", + } + } + ) + messages = [ + { + "role": "user", + "content": [ + { + "type": "video_url", + "video_url": { + "url": "https://example.com/video.mp4", + "video_metadata": { + "duration_seconds": 0, + "fps": 0, + "has_audio": False, + }, + }, + } + ], + } + ] + + try: + cost = completion_cost( + completion_response={"model": model, "usage": {}}, + model=model, + messages=messages, + custom_llm_provider="openai", + ) + finally: + litellm.model_cost.pop(model, None) + + assert cost == pytest.approx(max_input_tokens * input_cost_per_token) + + +def test_completion_cost_uses_provider_video_usage_when_present(): + model = "openai/test-video-provider-usage" + input_cost_per_token = 0.25 + output_cost_per_token = 0.5 + litellm.register_model( + model_cost={ + model: { + "input_cost_per_token": input_cost_per_token, + "output_cost_per_token": output_cost_per_token, + "max_tokens": 128, + "max_input_tokens": 128, + "max_output_tokens": 16, + "litellm_provider": "openai", + "mode": "chat", + } + } + ) + messages = [ + { + "role": "user", + "content": [ + { + "type": "video_url", + "video_url": { + "url": "https://example.com/video.mp4", + "video_metadata": { + "duration_seconds": 0, + "fps": 0, + "has_audio": False, + }, + }, + } + ], + } + ] + + try: + cost = completion_cost( + completion_response={ + "model": model, + "usage": { + "prompt_tokens": 2, + "completion_tokens": 3, + "total_tokens": 5, + }, + }, + model=model, + messages=messages, + custom_llm_provider="openai", + ) + finally: + litellm.model_cost.pop(model, None) + + assert cost == pytest.approx( + (2 * input_cost_per_token) + (3 * output_cost_per_token) + ) + + @pytest.mark.parametrize( "model, custom_llm_provider", [ diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index c6961477a58..0ec62ca3208 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -269,6 +269,71 @@ def test_response_cost_calculator_uses_router_model_id_from_litellm_metadata(): litellm.model_cost.pop(custom_model_id, None) +def test_response_cost_calculator_passes_messages_for_video_cost_fallback(): + import litellm + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + + model = "openai/test-video-logging-cost-fallback" + input_cost_per_token = 0.125 + max_input_tokens = 16 + litellm.register_model( + model_cost={ + model: { + "input_cost_per_token": input_cost_per_token, + "output_cost_per_token": 0.0, + "max_tokens": max_input_tokens, + "max_input_tokens": max_input_tokens, + "max_output_tokens": 4, + "litellm_provider": "openai", + "mode": "chat", + } + } + ) + messages = [ + { + "role": "user", + "content": [ + { + "type": "video_url", + "video_url": { + "url": "https://example.com/video.mp4", + "video_metadata": { + "duration_seconds": 0, + "fps": 0, + "has_audio": False, + }, + }, + } + ], + } + ] + + try: + logging_obj = LiteLLMLoggingObj( + model=model, + messages=messages, + stream=False, + call_type="completion", + start_time=time.time(), + litellm_call_id="test-video-fallback", + function_id="test-fn", + ) + logging_obj.update_environment_variables( + model=model, + user="", + optional_params={}, + litellm_params={"api_base": ""}, + ) + + cost = logging_obj._response_cost_calculator( + result={"model": model, "usage": {}} + ) + + assert cost == pytest.approx(max_input_tokens * input_cost_per_token) + finally: + litellm.model_cost.pop(model, None) + + class TestGetRouterModelId: """Tests for the get_router_model_id helper method."""