diff --git a/litellm/llms/gemini/count_tokens/transformation.py b/litellm/llms/gemini/count_tokens/transformation.py index 5a046544901..c4098adb823 100644 --- a/litellm/llms/gemini/count_tokens/transformation.py +++ b/litellm/llms/gemini/count_tokens/transformation.py @@ -8,6 +8,7 @@ and would drop tool_calls, tool messages, and web search tool declarations. """ import json +import re from collections.abc import Mapping, Sequence from dataclasses import dataclass from typing import Final, cast @@ -51,14 +52,14 @@ _ANTHROPIC_PART_TYPES: Final = frozenset( } ) -_ANTHROPIC_TOOL_TYPE_PREFIXES: Final = ( - "web_search_", - "web_fetch_", - "code_execution_", - "computer_", - "text_editor_", - "bash_", - "mcp_", +# Anthropic hosted tools are date-versioned (web_search_20250305). OpenAI +# names like web_search_preview or computer_use must not match, so prefixes +# are excluded unless they carry a date suffix. +_ANTHROPIC_TOOL_TYPE_NAMES: Final = frozenset( + {"web_search", "web_fetch", "code_execution", "computer", "text_editor", "bash", "mcp_toolset"} +) +_ANTHROPIC_TOOL_TYPE_RE: Final = re.compile( + r"^(web_search|web_fetch|code_execution|computer|text_editor|bash|mcp_toolset)_\d{8}$" ) # Server-side Anthropic content blocks the anthropic->openai adapter drops, so @@ -126,7 +127,9 @@ def _has_anthropic_shape( if "input_schema" in tool: return True tool_type = tool.get("type") - if isinstance(tool_type, str) and tool_type.startswith(_ANTHROPIC_TOOL_TYPE_PREFIXES): + if isinstance(tool_type, str) and ( + tool_type in _ANTHROPIC_TOOL_TYPE_NAMES or _ANTHROPIC_TOOL_TYPE_RE.match(tool_type) + ): return True return False @@ -423,3 +426,61 @@ def build_count_tokens_payload( if _has_anthropic_shape(system=system, tools=tools, messages=messages): return _build_anthropic_payload(model=model, messages=messages, system=system, tools=tools) return _build_openai_payload(model=model, messages=messages, system=system, tools=tools) + + +# Matches real inlineData blobs; a short or non-base64 `data` field (tool args, +# function responses) stays text so the fallback count keeps its mass. +_BASE64_BLOB_RE: Final = re.compile(r"[A-Za-z0-9+/=]{16,}") + + +def _elide_data_key(obj: dict[str, object]) -> dict[str, object]: + """json.loads object_hook that replaces base64 blobs (inlineData.data) + so serialized parts stay a sane size for the local tokenizer.""" + return { # mutable-ok: object_hook contract returns a rebuilt object per JSON node + key: ("" if key == "data" and isinstance(value, str) and _BASE64_BLOB_RE.fullmatch(value) else value) + for key, value in obj.items() + } + + +def _serialize_part(part: object) -> str: + return json.dumps(json.loads(json.dumps(part, default=str), object_hook=_elide_data_key), default=str) + + +def _part_to_text(part: object) -> str: + if isinstance(part, Mapping) and isinstance(part.get("text"), str): + return part["text"] + return _serialize_part(part) + + +def _content_parts(content: Mapping[str, object]) -> tuple[object, ...]: + parts: Final = content.get("parts") + if isinstance(parts, list): + return tuple(parts) + return (content,) + + +def gemini_contents_as_chat_messages(contents: object) -> tuple[Mapping[str, object], ...] | None: + """Approximate gemini contents as chat messages for the local fallback + tokenizer. Text parts count as text; other parts count as their JSON + frame with base64 blobs elided.""" + if contents is None: + return None + if isinstance(contents, list): + messages: Final = tuple( + { # mutable-ok: transient chat-shaped message for the local tokenizer + "role": "assistant" if content.get("role") == "model" else "user", + "content": "\n".join(_part_to_text(part) for part in _content_parts(content)), + } + for content in contents + if isinstance(content, Mapping) + ) + counted: Final = tuple(message for message in messages if message["content"]) + if counted: + return counted + fallback: Final[tuple[Mapping[str, object], ...]] = ( + { # mutable-ok: transient chat-shaped message for the local tokenizer + "role": "user", + "content": _serialize_part(contents), + }, + ) + return fallback diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index b81b311f9de..87992d8da07 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -13333,6 +13333,7 @@ async def run_thread( # ) # async def get_available_routes(user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth)): from litellm.llms.base_llm.base_utils import BaseTokenCounter +from litellm.llms.gemini.count_tokens.transformation import gemini_contents_as_chat_messages from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.model_repository import ModelRepository @@ -13459,58 +13460,6 @@ def _system_message(system: object) -> ChatCompletionSystemMessage | None: return message -def _elide_data_key(obj: dict[str, object]) -> dict[str, object]: - """json.loads object_hook that replaces base64 blobs (inlineData.data) - so serialized parts stay a sane size for the local tokenizer.""" - return { # mutable-ok: object_hook contract returns a rebuilt object per JSON node - key: ("" if key == "data" and isinstance(value, str) else value) for key, value in obj.items() - } - - -def _serialize_part(part: object) -> str: - return json.dumps(json.loads(json.dumps(part, default=str), object_hook=_elide_data_key), default=str) - - -def _part_to_text(part: object) -> str: - if isinstance(part, Mapping) and isinstance(part.get("text"), str): - return part["text"] - return _serialize_part(part) - - -def _content_parts(content: Mapping[str, object]) -> tuple[object, ...]: - parts: Final = content.get("parts") - if isinstance(parts, list): - return tuple(parts) - return (content,) - - -def _contents_as_messages(contents: object) -> tuple[Mapping[str, object], ...] | None: - """Approximate gemini contents as chat messages for the local fallback - tokenizer. Text parts count as text; other parts count as their JSON - frame with base64 blobs elided.""" - if contents is None: - return None - if isinstance(contents, list): - messages: Final = tuple( - { # mutable-ok: transient chat-shaped message for the local tokenizer - "role": "assistant" if content.get("role") == "model" else "user", - "content": "\n".join(_part_to_text(part) for part in _content_parts(content)), - } - for content in contents - if isinstance(content, Mapping) - ) - counted: Final = tuple(message for message in messages if message["content"]) - if counted: - return counted - fallback: Final[tuple[Mapping[str, object], ...]] = ( - { # mutable-ok: transient chat-shaped message for the local tokenizer - "role": "user", - "content": _serialize_part(contents), - }, - ) - return fallback - - @router.post( "/utils/token_counter", tags=["llm utils"], @@ -13614,7 +13563,7 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False) 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 if messages is not None else _contents_as_messages(contents), + messages if messages is not None else gemini_contents_as_chat_messages(contents), ) counted_messages: Final = ( typed_messages if typed_messages is None or system_message is None else (system_message, *typed_messages) diff --git a/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py b/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py index 13d186f3021..254e4b96858 100644 --- a/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py +++ b/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py @@ -119,6 +119,39 @@ def test_build_count_tokens_payload_maps_openai_web_search_tool(): assert payload.tools == [{"googleSearch": {}}] +def test_build_count_tokens_payload_routes_openai_tool_types_to_openai_path(): + """Regression: web_search_preview and computer_use are OpenAI tool types; + they must not trip the anthropic shape detector (which drops tool_calls).""" + payload = build_count_tokens_payload( + model="gemini-2.5-flash", + messages=[ + {"role": "user", "content": "check it"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "get_weather", "arguments": '{"city":"Paris"}'}, + } + ], + }, + {"role": "tool", "content": "sunny", "tool_call_id": "call_1"}, + ], + system=None, + tools=[ + {"type": "web_search_preview"}, + {"type": "computer_use", "display_width": 1024, "display_height": 768}, + ], + ) + + function_call = payload.contents[1]["parts"][0].get("function_call") + assert function_call == {"name": "get_weather", "args": {"city": "Paris"}} + function_response = payload.contents[2]["parts"][0].get("function_response") + assert function_response["name"] == "get_weather" + + def test_build_count_tokens_payload_wraps_responses_api_tool(): payload = build_count_tokens_payload( model="gemini-2.5-flash", diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_utils.py b/tests/test_litellm/proxy/proxy_server/test_routes_utils.py index fd67e97427a..07e8610fc5e 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_utils.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_utils.py @@ -412,9 +412,9 @@ def test_token_counter_contents_only_request_counts_text_parts(client, auth_as, monkeypatch.setattr(proxy_server, "llm_router", None) monkeypatch.setattr(litellm, "disable_token_counter", False, raising=False) contents = [ - {"role": "user", "parts": [{"text": "hello world"}, {"inline_data": {"data": "abc"}}]}, + {"role": "user", "parts": [{"text": "hello world"}, {"inline_data": {"data": "QUJDREVGR0hJSktMTU5PUFFSUw=="}}]}, {"role": "model", "parts": [{"text": "hi there"}]}, - {"role": "user", "parts": [{"inline_data": {"data": "abc"}}]}, + {"role": "user", "parts": [{"inline_data": {"data": "QUJDREVGR0hJSktMTU5PUFFSUw=="}}]}, ] with auth_as(): @@ -437,8 +437,11 @@ def test_token_counter_media_only_contents_falls_back_instead_of_500(client, aut monkeypatch.setattr(proxy_server, "llm_router", None) monkeypatch.setattr(litellm, "disable_token_counter", False, raising=False) contents = [ - {"role": "user", "parts": [{"inline_data": {"mime_type": "image/png", "data": "aGVsbG8="}}]}, - {"role": "model", "parts": [{"function_call": {"name": "get_weather", "args": {"city": "sf"}}}]}, + {"role": "user", "parts": [{"inline_data": {"mime_type": "image/png", "data": "QUJDREVGR0hJSktMTU5PUFFSUw=="}}]}, + { + "role": "model", + "parts": [{"function_call": {"name": "get_weather", "args": {"city": "sf", "data": "daily notes"}}}], + }, ] with auth_as(): @@ -449,6 +452,9 @@ def test_token_counter_media_only_contents_falls_back_instead_of_500(client, aut model="claude-fable-5", messages=[ {"role": "user", "content": '{"inline_data": {"mime_type": "image/png", "data": ""}}'}, - {"role": "assistant", "content": '{"function_call": {"name": "get_weather", "args": {"city": "sf"}}}'}, + { + "role": "assistant", + "content": '{"function_call": {"name": "get_weather", "args": {"city": "sf", "data": "daily notes"}}}', + }, ], )