fix: account for video duration in token counter

This commit is contained in:
pragnyanramtha 2026-05-18 04:45:55 +00:00
parent 031acb9d38
commit aac41538cd
9 changed files with 504 additions and 25 deletions

View file

@ -2,6 +2,7 @@
## Helper utilities for token counting
import base64
import io
import math
import struct
from typing import (
Any,
@ -44,6 +45,47 @@ from litellm.types.llms.openai import (
)
from litellm.types.utils import Message, SelectTokenizerResponse
DEFAULT_VIDEO_TOKEN_COUNT_PER_SECOND = 263
DEFAULT_AUDIO_TOKEN_COUNT_PER_SECOND = 32
def messages_contain_video_url(messages: Any) -> bool:
if not isinstance(messages, list):
return False
for message in messages:
if not isinstance(message, dict):
continue
content = message.get("content")
if not isinstance(content, list):
continue
if any(
isinstance(content_block, dict) and content_block.get("type") == "video_url"
for content_block in content
):
return True
return False
def get_token_count_for_limit_enforcement(
input_tokens: int,
messages: Any,
token_limit: Optional[Union[int, float]],
) -> int:
if not messages_contain_video_url(messages):
return input_tokens
if (
token_limit is None
or isinstance(token_limit, bool)
or not isinstance(token_limit, (int, float))
or not math.isfinite(token_limit)
):
return input_tokens
# Video metadata in the request is client-provided. For admission gates,
# reserve the full finite limit rather than trusting understated duration/fps.
return max(input_tokens, math.ceil(token_limit))
def get_modified_max_tokens(
model: str,
@ -623,6 +665,155 @@ def _count_image_tokens(
)
def _coerce_float(value: Any) -> Optional[float]:
if isinstance(value, bool):
return None
if isinstance(value, (int, float)):
return float(value)
if isinstance(value, str):
number = value.strip()
if number.endswith("s"):
number = number[:-1]
try:
return float(number)
except ValueError:
return None
return None
def _coerce_duration_seconds(value: Any) -> Optional[float]:
if isinstance(value, dict):
seconds = _coerce_float(value.get("seconds"))
nanos = _coerce_float(value.get("nanos"))
if seconds is None:
return None
return seconds + ((nanos or 0) / 1_000_000_000)
return _coerce_float(value)
def _get_video_value(
video_url: Mapping[str, Any],
video_metadata: Mapping[str, Any],
keys: Tuple[str, ...],
) -> Any:
for key in keys:
if key in video_url:
return video_url[key]
if key in video_metadata:
return video_metadata[key]
return None
def _coerce_bool(value: Any) -> Optional[bool]:
if isinstance(value, bool):
return value
if isinstance(value, str):
lowered_value = value.strip().lower()
if lowered_value in ("true", "1", "yes"):
return True
if lowered_value in ("false", "0", "no"):
return False
return None
def _get_video_duration_seconds(
video_url: Mapping[str, Any],
video_metadata: Mapping[str, Any],
) -> Optional[float]:
duration_seconds = _coerce_duration_seconds(
_get_video_value(
video_url,
video_metadata,
("duration_seconds", "duration", "seconds"),
)
)
if duration_seconds is not None:
return duration_seconds
start_offset = _coerce_duration_seconds(
_get_video_value(
video_url,
video_metadata,
("start_offset", "startOffset"),
)
)
end_offset = _coerce_duration_seconds(
_get_video_value(
video_url,
video_metadata,
("end_offset", "endOffset"),
)
)
if start_offset is not None and end_offset is not None:
return end_offset - start_offset
return None
def _count_video_tokens(video_url: Any) -> int:
"""
Count tokens for a video_url content block without tokenizing the URL/base64 bytes.
"""
video_metadata: Mapping[str, Any] = {}
if isinstance(video_url, dict):
url = video_url.get("url")
if not url:
raise ValueError("Missing required key 'url' in video_url dict.")
metadata = video_url.get("video_metadata")
if isinstance(metadata, Mapping):
video_metadata = metadata
elif isinstance(video_url, str):
if not video_url.strip():
raise ValueError("Empty video_url string is not valid.")
else:
raise ValueError(
f"Invalid video_url type: {type(video_url).__name__}. "
"Expected str or dict with 'url' field."
)
# If callers do not provide duration metadata, avoid inspecting/fetching media
# and use a one-second minimum rather than counting URL or base64 text.
duration_seconds = 1.0
if isinstance(video_url, dict):
parsed_duration_seconds = _get_video_duration_seconds(video_url, video_metadata)
if parsed_duration_seconds is not None:
duration_seconds = parsed_duration_seconds
if duration_seconds < 0:
raise ValueError("video_url duration must be non-negative.")
fps = 1.0
if isinstance(video_url, dict):
parsed_fps = _coerce_duration_seconds(
_get_video_value(video_url, video_metadata, ("fps",))
)
if parsed_fps is not None:
fps = parsed_fps
if fps < 0:
raise ValueError("video_url fps must be non-negative.")
has_audio = True
if isinstance(video_url, dict):
parsed_has_audio = _coerce_bool(
_get_video_value(
video_url,
video_metadata,
("has_audio", "contains_audio", "audio"),
)
)
if parsed_has_audio is not None:
has_audio = parsed_has_audio
video_tokens = math.ceil(
duration_seconds * fps * DEFAULT_VIDEO_TOKEN_COUNT_PER_SECOND
)
audio_tokens = (
math.ceil(duration_seconds * DEFAULT_AUDIO_TOKEN_COUNT_PER_SECOND)
if has_audio
else 0
)
return video_tokens + audio_tokens
def _validate_anthropic_content(content: Mapping[str, Any]) -> type:
"""
Validate and determine which Anthropic TypedDict applies.
@ -724,7 +915,7 @@ def _count_content_list(
image_url, use_default_image_token_count
)
elif c["type"] == "video_url":
num_tokens += DEFAULT_IMAGE_TOKEN_COUNT
num_tokens += _count_video_tokens(c.get("video_url"))
elif c["type"] in ("tool_use", "tool_result"):
num_tokens += _count_anthropic_content(
c,

View file

@ -9,6 +9,9 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching import DualCache
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.litellm_core_utils.token_counter import (
get_token_count_for_limit_enforcement,
)
from litellm.proxy._types import (
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
@ -932,12 +935,18 @@ def _estimate_input_tokens(
) -> Optional[int]:
try:
if "messages" in request_body:
return litellm.token_counter(
messages = request_body.get("messages") or []
input_tokens = litellm.token_counter(
model=model,
messages=request_body.get("messages") or [],
messages=messages,
tools=request_body.get("tools"),
tool_choice=request_body.get("tool_choice"),
)
return get_token_count_for_limit_enforcement(
input_tokens=input_tokens,
messages=messages,
token_limit=_to_int(model_info.get("max_input_tokens")),
)
if "prompt" in request_body:
return _count_text_tokens(model=model, text=request_body.get("prompt"))
if "input" in request_body:

View file

@ -8,6 +8,9 @@ from litellm import ModelResponse, token_counter, verbose_logger
from litellm._logging import verbose_router_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.token_counter import (
get_token_count_for_limit_enforcement,
)
class LowestCostLoggingHandler(CustomLogger):
@ -260,6 +263,11 @@ class LowestCostLoggingHandler(CustomLogger):
or _deployment.get("model_info", {}).get("rpm", None)
or float("inf")
)
input_tokens_for_tpm = get_token_count_for_limit_enforcement(
input_tokens=input_tokens,
messages=messages,
token_limit=_deployment_tpm,
)
item_litellm_model_name = _deployment.get("litellm_params", {}).get("model")
item_litellm_model_cost_map = litellm.model_cost.get(
item_litellm_model_name, {}
@ -314,7 +322,7 @@ class LowestCostLoggingHandler(CustomLogger):
# -------------- #
if (
item_tpm + input_tokens > _deployment_tpm
item_tpm + input_tokens_for_tpm > _deployment_tpm
or item_rpm + 1 > _deployment_rpm
): # if user passed in tpm / rpm in the model_list
continue

View file

@ -10,6 +10,9 @@ from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import safe_divide_seconds
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
from litellm.litellm_core_utils.token_counter import (
get_token_count_for_limit_enforcement,
)
from litellm.types.utils import LiteLLMPydanticObjectBase
if TYPE_CHECKING:
@ -485,6 +488,11 @@ class LowestLatencyLoggingHandler(CustomLogger):
or _deployment.get("model_info", {}).get("rpm", None)
or float("inf")
)
input_tokens_for_tpm = get_token_count_for_limit_enforcement(
input_tokens=input_tokens,
messages=messages,
token_limit=_deployment_tpm,
)
item_latency = item_map.get("latency", [])
item_ttft_latency = item_map.get("time_to_first_token", [])
item_rpm = item_map.get(precise_minute, {}).get("rpm", 0)
@ -524,7 +532,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
# -------------- #
if (
item_tpm + input_tokens > _deployment_tpm
item_tpm + input_tokens_for_tpm > _deployment_tpm
or item_rpm + 1 > _deployment_rpm
): # if user passed in tpm / rpm in the model_list
continue

View file

@ -8,6 +8,9 @@ from litellm import token_counter
from litellm._logging import verbose_router_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.token_counter import (
get_token_count_for_limit_enforcement,
)
from litellm.types.utils import LiteLLMPydanticObjectBase
from litellm.utils import print_verbose
@ -225,6 +228,11 @@ class LowestTPMLoggingHandler(CustomLogger):
_deployment_tpm = _deployment.get("model_info", {}).get("tpm")
if _deployment_tpm is None:
_deployment_tpm = float("inf")
input_tokens_for_tpm = get_token_count_for_limit_enforcement(
input_tokens=input_tokens,
messages=messages,
token_limit=_deployment_tpm,
)
_deployment_rpm = None
if _deployment_rpm is None:
@ -236,7 +244,7 @@ class LowestTPMLoggingHandler(CustomLogger):
if _deployment_rpm is None:
_deployment_rpm = float("inf")
if item_tpm + input_tokens > _deployment_tpm:
if item_tpm + input_tokens_for_tpm > _deployment_tpm:
continue
elif (rpm_dict is not None and item in rpm_dict) and (
rpm_dict[item] + 1 >= _deployment_rpm

View file

@ -11,6 +11,9 @@ from litellm._logging import verbose_logger, verbose_router_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
from litellm.litellm_core_utils.token_counter import (
get_token_count_for_limit_enforcement,
)
from litellm.types.router import RouterErrors
from litellm.types.utils import LiteLLMPydanticObjectBase, StandardLoggingPayload
from litellm.utils import get_utc_datetime, print_verbose
@ -330,6 +333,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
all_deployments: Dict,
input_tokens: int,
rpm_dict: Dict,
messages: Optional[List[Dict[str, str]]] = None,
):
lowest_tpm = float("inf")
potential_deployments = [] # if multiple deployments have the same low value
@ -355,6 +359,11 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
_deployment_tpm = _deployment.get("model_info", {}).get("tpm")
if _deployment_tpm is None:
_deployment_tpm = float("inf")
input_tokens_for_tpm = get_token_count_for_limit_enforcement(
input_tokens=input_tokens,
messages=messages,
token_limit=_deployment_tpm,
)
_deployment_rpm = None
if _deployment_rpm is None:
@ -365,7 +374,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
_deployment_rpm = _deployment.get("model_info", {}).get("rpm")
if _deployment_rpm is None:
_deployment_rpm = float("inf")
if item_tpm + input_tokens > _deployment_tpm:
if item_tpm + input_tokens_for_tpm > _deployment_tpm:
continue
elif (
(rpm_dict is not None and item in rpm_dict)
@ -433,6 +442,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
all_deployments=all_deployments,
input_tokens=input_tokens,
rpm_dict=rpm_dict,
messages=messages,
)
print_verbose("returning picked lowest tpm/rpm deployment.")

View file

@ -725,6 +725,50 @@ def test_return_potential_deployments():
assert len(potential_deployments) == 1
def test_return_potential_deployments_uses_full_tpm_for_video_url():
test_cache = DualCache()
lowest_tpm_logger = LowestTPMLoggingHandler(router_cache=test_cache)
deployment_id = "video-deployment"
potential_deployments = lowest_tpm_logger._return_potential_deployments(
healthy_deployments=[
{
"model_name": "model-test",
"litellm_params": {
"model": "openai/gpt-4o",
},
"model_info": {
"id": deployment_id,
"tpm": 100,
},
},
],
all_deployments={
f"{deployment_id}:tpm:02-17": 1,
},
input_tokens=10,
rpm_dict={},
messages=[
{
"role": "user",
"content": [
{
"type": "video_url",
"video_url": {
"url": "https://example.com/long-video.mp4",
"duration_seconds": 1,
"fps": 0,
"has_audio": False,
},
}
],
}
],
)
assert potential_deployments == []
@pytest.mark.asyncio
async def test_tpm_rpm_routing_model_name_checks():
deployment = {

View file

@ -16,7 +16,10 @@ from unittest.mock import AsyncMock, MagicMock, patch
import litellm
from litellm import create_pretrained_tokenizer, decode, encode, get_modified_max_tokens
from litellm import token_counter as token_counter_old
from litellm.constants import DEFAULT_IMAGE_TOKEN_COUNT
from litellm.litellm_core_utils.token_counter import (
get_token_count_for_limit_enforcement,
messages_contain_video_url,
)
from litellm.litellm_core_utils.token_counter import token_counter as token_counter_new
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.messages_with_counts import (
@ -944,9 +947,6 @@ def test_token_counter_with_video_url():
assert (
tokens_str > 0
), f"Expected positive token count for string video_url, got {tokens_str}"
assert (
tokens_str > DEFAULT_IMAGE_TOKEN_COUNT
), f"Expected default video token budget, got {tokens_str}"
messages_base64_str = [
{
@ -959,40 +959,200 @@ def test_token_counter_with_video_url():
],
}
]
assert (
token_counter(model="gpt-4o", messages=messages_base64_str) == tokens_str
)
assert token_counter(model="gpt-4o", messages=messages_base64_str) == tokens_str
messages_empty_url = [
messages_one_second_without_audio = [
{
"role": "user",
"content": [{"type": "video_url", "video_url": ""}],
"content": [
{
"type": "video_url",
"video_url": {
"url": "https://example.com/video.mp4",
"duration_seconds": 1,
"has_audio": False,
},
}
],
}
]
messages_none_url = [
messages_one_second_with_audio = [
{
"role": "user",
"content": [{"type": "video_url", "video_url": None}],
"content": [
{
"type": "video_url",
"video_url": {
"url": "https://example.com/video.mp4",
"duration_seconds": 1,
"has_audio": True,
},
}
],
}
]
messages_none_nested_url = [
messages_two_seconds_with_audio = [
{
"role": "user",
"content": [{"type": "video_url", "video_url": {"url": None}}],
"content": [
{
"type": "video_url",
"video_url": {
"url": "https://example.com/video.mp4",
"duration_seconds": 2,
"has_audio": True,
},
}
],
}
]
tokens_empty_url = token_counter(model="gpt-4o", messages=messages_empty_url)
assert tokens_empty_url > DEFAULT_IMAGE_TOKEN_COUNT
tokens_one_second_without_audio = token_counter(
model="gpt-4o", messages=messages_one_second_without_audio
)
tokens_one_second_with_audio = token_counter(
model="gpt-4o", messages=messages_one_second_with_audio
)
tokens_two_seconds_with_audio = token_counter(
model="gpt-4o", messages=messages_two_seconds_with_audio
)
expected_audio_tokens_per_second = 32
expected_video_tokens_per_second = 263
assert (
token_counter(model="gpt-4o", messages=messages_none_url) == tokens_empty_url
tokens_one_second_with_audio - tokens_one_second_without_audio
== expected_audio_tokens_per_second
)
assert (
token_counter(model="gpt-4o", messages=messages_none_nested_url)
== tokens_empty_url
tokens_two_seconds_with_audio - tokens_one_second_with_audio
== expected_video_tokens_per_second + expected_audio_tokens_per_second
)
def _count_video_url_for_test(video_url):
return token_counter(
model="gpt-4o",
messages=[
{
"role": "user",
"content": [{"type": "video_url", "video_url": video_url}],
}
],
)
def test_token_counter_video_url_metadata_shapes():
zero_second_no_audio_tokens = _count_video_url_for_test(
{
"url": "https://example.com/video.mp4",
"duration_seconds": 0,
"has_audio": False,
}
)
metadata_tokens = _count_video_url_for_test(
{
"url": "https://example.com/video.mp4",
"video_metadata": {
"duration": "2s",
"fps": "2",
"audio": "no",
},
}
)
assert metadata_tokens - zero_second_no_audio_tokens == 2 * 2 * 263
offset_tokens = _count_video_url_for_test(
{
"url": "https://example.com/video.mp4",
"video_metadata": {
"startOffset": {"seconds": 1, "nanos": 0},
"endOffset": {"seconds": 4, "nanos": 0},
"contains_audio": "yes",
},
}
)
assert offset_tokens - zero_second_no_audio_tokens == 3 * (263 + 32)
invalid_duration_tokens = _count_video_url_for_test(
{
"url": "https://example.com/video.mp4",
"duration_seconds": "not-a-duration",
"has_audio": False,
}
)
one_second_no_audio_tokens = _count_video_url_for_test(
{
"url": "https://example.com/video.mp4",
"duration_seconds": 1,
"has_audio": False,
}
)
assert invalid_duration_tokens == one_second_no_audio_tokens
def test_video_url_token_count_for_limit_enforcement_uses_full_limit():
messages = [
{
"role": "user",
"content": [
{
"type": "video_url",
"video_url": {
"url": "https://example.com/long-video.mp4",
"duration_seconds": 1,
"fps": 0,
"has_audio": False,
},
}
],
}
]
input_tokens = token_counter(model="gpt-4o", messages=messages)
assert messages_contain_video_url(messages) is True
assert (
get_token_count_for_limit_enforcement(
input_tokens=input_tokens,
messages=messages,
token_limit=128_000,
)
== 128_000
)
text_messages = [{"role": "user", "content": "hello"}]
assert (
get_token_count_for_limit_enforcement(
input_tokens=5,
messages=text_messages,
token_limit=128_000,
)
== 5
)
@pytest.mark.parametrize(
"video_url,error",
[
("", "Empty video_url string is not valid"),
(None, "Invalid video_url type"),
({"url": ""}, "Missing required key 'url'"),
(
{"url": "https://example.com/video.mp4", "duration_seconds": -1},
"duration must be non-negative",
),
(
{"url": "https://example.com/video.mp4", "fps": -1},
"fps must be non-negative",
),
],
)
def test_token_counter_invalid_video_url_metadata(video_url, error):
with pytest.raises(ValueError) as exc_info:
_count_video_url_for_test(video_url)
assert error in str(exc_info.value)
def test_token_counter_invalid_content_type_lists_video_url():
messages = [
{

View file

@ -699,6 +699,47 @@ async def test_should_clamp_reservation_to_model_ceiling_when_caller_overrequest
await release_budget_reservation(reservation)
def test_should_reserve_max_input_tokens_for_video_url_budget_estimate():
request_body = {
"model": "gpt-4o-mini",
"messages": [
{
"role": "user",
"content": [
{
"type": "video_url",
"video_url": {
"url": "https://example.com/long-video.mp4",
"duration_seconds": 1,
"fps": 0,
"has_audio": False,
},
}
],
}
],
"max_tokens": 0,
}
input_cost_per_token = 1e-6
max_input_tokens = 128_000
with patch(
"litellm.proxy.spend_tracking.budget_reservation._get_model_cost_info",
return_value={
"input_cost_per_token": input_cost_per_token,
"max_input_tokens": max_input_tokens,
},
):
estimated = estimate_request_max_cost(
request_body=request_body,
route="/chat/completions",
llm_router=None,
)
assert estimated == pytest.approx(max_input_tokens * input_cost_per_token)
@pytest.mark.asyncio
async def test_should_reserve_image_generation_cost_per_image(
spend_counter_state,