From e40a7ffb04880fd0482b21d6c54a72d33023d37a Mon Sep 17 00:00:00 2001 From: shrey kharbanda Date: Sun, 27 Sep 2026 04:14:46 +0000 Subject: [PATCH] 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 --- litellm/litellm_core_utils/token_counter.py | 72 ++++++++++++++++++- litellm/proxy/proxy_server.py | 4 +- .../code_coverage_tests/recursive_detector.py | 1 + .../litellm_core_utils/test_token_counter.py | 66 +++++++++++++++++ tests/unit/proxy/test_proxy_token_counter.py | 69 ++++++++++++++++++ 5 files changed, 208 insertions(+), 4 deletions(-) diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 5d7956059e4..94757280cda 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -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 "" + return _without_opaque_keys(value, depth + 1) + + +def _without_opaque_keys(value: object, depth: int = 0) -> object: + if depth > DEFAULT_MAX_RECURSE_DEPTH: + return "" + 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) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7a6068f3b6f..4e7f82b6887 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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) diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 659dc438f2d..c53793495ac 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -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. diff --git a/tests/unit/litellm_core_utils/test_token_counter.py b/tests/unit/litellm_core_utils/test_token_counter.py index f7ded4f3fa8..731407a85be 100644 --- a/tests/unit/litellm_core_utils/test_token_counter.py +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -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('""' + "}" * (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": "", "stdout": "ok", "notes": {"data": "two words"}}}' + ) diff --git a/tests/unit/proxy/test_proxy_token_counter.py b/tests/unit/proxy/test_proxy_token_counter.py index b39759eca70..16431c1deda 100644 --- a/tests/unit/proxy/test_proxy_token_counter.py +++ b/tests/unit/proxy/test_proxy_token_counter.py @@ -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=[