fix(proxy): count web search and other unknown blocks in the local estimate

Blocks the local tokenizer does not know, such as server_tool_use and
web_search_tool_result, are counted as the text of their payload instead of
crashing the count. Ids, signatures, cache_control and encrypted content are
dropped first, and inline base64 data is elided, so none of it inflates the
estimate
This commit is contained in:
shrey kharbanda 2026-09-27 04:14:46 +00:00
parent 613caa5e39
commit e40a7ffb04
No known key found for this signature in database
5 changed files with 208 additions and 4 deletions

View file

@ -2,6 +2,8 @@
## Helper utilities for token counting
import base64
import io
import json
import re
import struct
from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
from typing import Final, Literal, cast
@ -19,6 +21,7 @@ from litellm.constants import (
DEFAULT_IMAGE_HEIGHT,
DEFAULT_IMAGE_TOKEN_COUNT,
DEFAULT_IMAGE_WIDTH,
DEFAULT_MAX_RECURSE_DEPTH,
MAX_IMAGE_URL_DOWNLOAD_SIZE_MB,
MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES,
MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES,
@ -866,6 +869,20 @@ def _count_anthropic_content(
return tokens
LOCALLY_COUNTABLE_BLOCK_TYPES: Final = (
"text",
"image_url",
"image",
"document",
"file",
"tool_use",
"tool_result",
"thinking",
"redacted_thinking",
"tool_reference",
)
def _count_content_list(
count_function: TokenCounterFunction,
content_list: str
@ -938,9 +955,7 @@ def _count_content_list(
content_type = c.get("type", type(c).__name__) if isinstance(c, dict) else type(c).__name__
raise ValueError(
f"Invalid content item type: {content_type}. "
f"Expected str or dict with 'type' field "
f"(text, image_url, image, document, file, tool_use, tool_result, thinking, redacted_thinking, "
f"tool_reference)."
f"Expected str or dict with 'type' field ({', '.join(LOCALLY_COUNTABLE_BLOCK_TYPES)})."
)
return num_tokens
except Exception as e:
@ -1033,3 +1048,54 @@ def _format_type(props, indent):
else:
# This is a guess, as an empty string doesn't yield the expected token count
return "any"
_INLINE_DATA_BASE64_RE: Final = re.compile(r"[A-Za-z0-9+/=_-]{16,}")
_OPAQUE_BLOCK_KEYS: Final = frozenset(
{"id", "tool_use_id", "cache_control", "signature", "encrypted_content", "encrypted_index"}
)
def _countable_value(key: object, value: object, depth: int) -> object:
if key == "data" and isinstance(value, str) and _INLINE_DATA_BASE64_RE.fullmatch(value):
return "<binary>"
return _without_opaque_keys(value, depth + 1)
def _without_opaque_keys(value: object, depth: int = 0) -> object:
if depth > DEFAULT_MAX_RECURSE_DEPTH:
return "<truncated>"
if isinstance(value, Mapping):
return { # mutable-ok: json.dumps input
key: _countable_value(key, item, depth) for key, item in value.items() if key not in _OPAQUE_BLOCK_KEYS
}
if isinstance(value, (list, tuple)):
return [_without_opaque_keys(item, depth + 1) for item in value] # mutable-ok: json.dumps input
return value
def _countable_leaf_block(block: object) -> object:
if not isinstance(block, Mapping) or block.get("type") in LOCALLY_COUNTABLE_BLOCK_TYPES:
return block
return {"type": "text", "text": json.dumps(_without_opaque_keys(block), default=str)}
def _countable_block(block: object) -> object:
if isinstance(block, Mapping) and block.get("type") == "tool_result" and isinstance(block.get("content"), list):
return {**block, "content": [_countable_leaf_block(item) for item in block["content"]]}
return _countable_leaf_block(block)
def _countable_message(message: object) -> object:
if not isinstance(message, Mapping) or not isinstance(message.get("content"), list):
return message
return {
**message,
"content": [_countable_block(block) for block in message["content"]],
}
def messages_with_uncountable_blocks_as_text(messages: Sequence[object]) -> tuple[object, ...]:
return tuple(_countable_message(message) for message in messages)

View file

@ -13564,6 +13564,7 @@ async def run_thread(
# dependencies=[Depends(user_api_key_auth)],
# )
# async def get_available_routes(user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth)):
from litellm.litellm_core_utils.token_counter import messages_with_uncountable_blocks_as_text
from litellm.llms.base_llm.base_utils import BaseTokenCounter
from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient
from litellm.repositories.config_repository import ConfigRepository
@ -13794,7 +13795,8 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False)
tokenizer_used: Final = str(_tokenizer_used["type"])
system_message: Final = _system_message(system)
typed_messages: Final = cast( # cast-ok: request messages are raw chat-shaped dicts that token_counter normalizes
Sequence[AllMessageValues] | None, messages
Sequence[AllMessageValues] | None,
None if messages is None else messages_with_uncountable_blocks_as_text(messages),
)
counted_messages: Final = (
typed_messages if typed_messages is None or system_message is None else (system_message, *typed_messages)

View file

@ -38,6 +38,7 @@ IGNORE_FUNCTIONS = [
"_mask_sequence", # max depth set.
"_delete_nested_value_custom", # max depth set (bounded by number of path segments).
"filter_exceptions_from_params", # max depth set (default 20) to prevent infinite recursion.
"_without_opaque_keys", # max depth set (DEFAULT_MAX_RECURSE_DEPTH).
"__getattr__", # lazy loading pattern in litellm/__init__.py with proper caching to prevent infinite recursion.
"_validate_inheritance_chain", # max depth set (default 100) to prevent infinite recursion in policy inheritance validation.
"_basic_json_schema_validate", # max depth set.

View file

@ -28,6 +28,7 @@ from litellm import token_counter as token_counter_old
import litellm.constants
from litellm.constants import TOKEN_COUNTER_MAX_CONCURRENT_COUNTS
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.litellm_core_utils.token_counter import (
_encoding_count,
_get_exact_count_function,
@ -35,6 +36,7 @@ from litellm.litellm_core_utils.token_counter import (
_get_tiktoken_count_function,
calculate_img_tokens,
high_detail_image_token_upper_bound,
messages_with_uncountable_blocks_as_text,
offload_token_count,
)
from litellm.litellm_core_utils.token_counter import token_counter as token_counter_new
@ -1561,3 +1563,67 @@ def test_token_counter_uses_the_tokenizer_of_each_model_family_and_of_a_custom_t
"custom": expected["Xenova/llama-3-tokenizer"],
"requested": sorted(served),
}
class _Unprintable:
def __str__(self) -> str:
raise AssertionError("opaque block values must be dropped before they are serialized")
def test_uncountable_block_drops_opaque_values_without_serializing_them():
(message,) = messages_with_uncountable_blocks_as_text(
[
{
"role": "assistant",
"content": [
{
"type": "web_search_tool_result",
"tool_use_id": _Unprintable(),
"content": [{"type": "web_search_result", "title": "Paris", "encrypted_content": _Unprintable()}],
}
],
}
]
)
assert message["content"] == [
{
"type": "text",
"text": '{"type": "web_search_tool_result", "content": [{"type": "web_search_result", "title": "Paris"}]}',
}
]
def test_uncountable_block_nesting_past_the_depth_limit_is_truncated():
nested: object = "leaf"
for _ in range(DEFAULT_MAX_RECURSE_DEPTH + 5):
nested = {"child": nested}
(message,) = messages_with_uncountable_blocks_as_text(
[{"role": "assistant", "content": [{"type": "server_tool_use", "input": nested}]}]
)
text = message["content"][0]["text"]
assert text.endswith('"<truncated>"' + "}" * (DEFAULT_MAX_RECURSE_DEPTH + 1))
assert "leaf" not in text
def test_uncountable_block_elides_inline_base64_data_but_keeps_plain_text_data():
(message,) = messages_with_uncountable_blocks_as_text(
[
{
"role": "assistant",
"content": [
{
"type": "code_execution_tool_result",
"content": {"data": "iVBORw0KGgo" * 20, "stdout": "ok", "notes": {"data": "two words"}},
}
],
}
]
)
assert message["content"][0]["text"] == (
'{"type": "code_execution_tool_result", '
'"content": {"data": "<binary>", "stdout": "ok", "notes": {"data": "two words"}}}'
)

View file

@ -1247,6 +1247,75 @@ async def test_anthropic_endpoint_429_rate_limit_error_format():
proxy_server.token_counter = original_token_counter
def _server_tool_history(stdout: str, encrypted_content: str) -> list[dict[str, object]]:
return [
{"role": "user", "content": "weather in Paris?"},
{
"role": "assistant",
"content": [
{"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "paris"}},
{
"type": "web_search_tool_result",
"tool_use_id": "srvtoolu_1",
"content": [
{
"type": "web_search_result",
"url": "https://example.com/paris",
"title": "Paris weather",
"encrypted_content": encrypted_content,
}
],
},
{
"type": "bash_code_execution_tool_result",
"tool_use_id": "srvtoolu_2",
"content": {"type": "bash_code_execution_result", "stdout": stdout, "stderr": "", "return_code": 0},
},
{
"type": "text_editor_code_execution_tool_result",
"tool_use_id": "srvtoolu_3",
"content": {"type": "text_editor_code_execution_view_result", "content": "notes"},
},
{"type": "tool_use", "id": "toolu_1", "name": "lookup", "input": {}},
],
},
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "toolu_1",
"content": [
{
"type": "search_result",
"source": "https://example.com",
"title": "t",
"content": [{"type": "text", "text": "18C"}],
}
],
}
],
},
]
@pytest.mark.asyncio
async def test_local_token_count_estimates_server_tool_history_without_counting_ciphertext(monkeypatch):
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", None)
async def count(stdout: str, encrypted_content: str) -> int:
result = await token_counter(
request=TokenCountRequest(model="gpt-4o", messages=_server_tool_history(stdout, encrypted_content))
)
return result.total_tokens
baseline = await count("18C", "RW5jcnlwdGVk")
assert baseline > 0
assert await count("18C", "RW5jcnlwdGVk" * 2000) == baseline
assert await count("18C and sunny for the rest of the week", "RW5jcnlwdGVk") > baseline
def _gemini_router() -> Router:
return Router(
model_list=[