fix(token_counter): count Gemini function_declarations tools (#43417)

* fix(token_counter): count Gemini function_declarations tools

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(token_counter): skip non-dict tools when formatting definitions

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-26 22:28:31 -07:00 • committed by GitHub
parent 4274bdda44
commit 8e6d99d74a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 126 additions and 33 deletions

View file

@ -951,7 +951,7 @@ def _count_content_list(
)
def _format_function_definitions(tools):
def _format_function_definitions(tools: Sequence[object]) -> str:
"""Formats tool definitions in the format that OpenAI appears to use.
Based on https://github.com/forestwanglin/openai-java/blob/main/jtokkit/src/main/java/xyz/felh/openai/jtokkit/utils/TikTokenUtils.java
"""
@ -959,41 +959,57 @@ def _format_function_definitions(tools):
lines.append("namespace functions {")
lines.append("")
for tool in tools:
if not isinstance(tool, dict):
if not isinstance(tool, Mapping):
continue
function = tool.get("function")
if not isinstance(function, dict):
# Anthropic tool shape → OpenAI function dict for token counting.
params = tool.get("input_schema") or tool.get("parameters") or {}
if not isinstance(params, dict):
params = {}
function = {
"name": tool.get("name"),
"description": tool.get("description"),
"parameters": params,
}
function_name = function.get("name")
if not function_name:
# Skip malformed tools missing a name to avoid emitting
# ``type None = ...`` which would produce inaccurate token counts.
continue
if function_description := function.get("description"):
lines.append(f"// {function_description}")
parameters = function.get("parameters") or {}
if not isinstance(parameters, dict):
parameters = {}
properties = parameters.get("properties")
if properties and properties.keys():
lines.append(f"type {function_name} = (_: {{")
lines.append(_format_object_parameters(parameters, 0))
lines.append("}) => any;")
else:
lines.append(f"type {function_name} = () => any;")
lines.append("")
for function in _function_definitions_for_tool(cast(Mapping[str, object], tool)):
lines.extend(_format_single_function_definition(function))
lines.append("} // namespace functions")
return "\n".join(lines)
def _function_definitions_for_tool(tool: Mapping[str, object]) -> Iterable[Mapping[str, object]]:
function: Final = tool.get("function")
if isinstance(function, Mapping):
yield function
return
declarations: Final = tool.get("function_declarations") or tool.get("functionDeclarations")
if isinstance(declarations, list):
for declaration in declarations:
if isinstance(declaration, Mapping):
yield declaration
return
parameters: Final = tool.get("input_schema") or tool.get("parameters") or {}
normalized_parameters: Final = parameters if isinstance(parameters, Mapping) else {}
yield {
"name": tool.get("name"),
"description": tool.get("description"),
"parameters": normalized_parameters,
}
def _format_single_function_definition(function: Mapping[str, object]) -> tuple[str, ...]:
function_name: Final = function.get("name")
if not function_name:
return ()
function_description: Final = function.get("description")
parameters_value: Final = function.get("parameters") or {}
parameters: Final = parameters_value if isinstance(parameters_value, Mapping) else {}
properties: Final = parameters.get("properties")
if isinstance(properties, Mapping) and properties:
return (
*((f"// {function_description}",) if function_description else ()),
f"type {function_name} = (_: {{",
_format_object_parameters(parameters, 0),
"}) => any;",
"",
)
return (
*((f"// {function_description}",) if function_description else ()),
f"type {function_name} = () => any;",
"",
)
def _format_object_parameters(parameters, indent):
properties: Final = parameters.get("properties")
if not properties:

View file

@ -442,6 +442,78 @@ def test_token_counter_with_tools(message_count_pair):
), f"Expected {expected_tokens} tokens, got {counted_tokens}."
def test_token_counter_counts_gemini_function_declarations():
openai_tools: Final = [
{
"type": "function",
"function": {
"name": "lookup_weather",
"description": "Find current weather conditions for a location",
"parameters": {
"type": "object",
"properties": {
"location": {"type": "string", "description": "City and region"},
"units": {"type": "string", "enum": ["celsius", "fahrenheit"]},
},
"required": ["location"],
},
},
}
]
gemini_tools: Final = litellm.utils.get_optional_params(
model="gemini-2.5-pro",
custom_llm_provider="gemini",
tools=openai_tools,
)["tools"]
camel_case_tools: Final = [{"functionDeclarations": gemini_tools[0]["function_declarations"]}]
openai_tokens: Final = token_counter_new(
model="gemini-2.5-pro",
messages=[{"role": "user", "content": "What's the weather?"}],
tools=openai_tools,
)
gemini_tokens: Final = token_counter_new(
model="gemini-2.5-pro",
messages=[{"role": "user", "content": "What's the weather?"}],
tools=gemini_tools,
)
camel_case_tokens: Final = token_counter_new(
model="gemini-2.5-pro",
messages=[{"role": "user", "content": "What's the weather?"}],
tools=camel_case_tools,
)
assert openai_tokens == gemini_tokens == camel_case_tokens
def test_token_counter_skips_non_mapping_tools():
openai_tool: Final = {
"type": "function",
"function": {
"name": "lookup_weather",
"description": "Find current weather conditions for a location",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string", "description": "City and region"}},
"required": ["location"],
},
},
}
messages: Final = [{"role": "user", "content": "What's the weather?"}]
valid_tokens: Final = token_counter_new(
model="gemini-2.5-pro",
messages=messages,
tools=[openai_tool],
)
mixed_tokens: Final = token_counter_new(
model="gemini-2.5-pro",
messages=messages,
tools=["bad", None, openai_tool],
)
assert mixed_tokens == valid_tokens
class NeedsToleranceUpdateError(Exception):
"""Custom exception to mark tests that have improved"""

View file

@ -1,4 +1,4 @@
from typing import List
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -1530,7 +1530,7 @@ class TestContextCachingEndpoints:
]
all_messages = short_cached_messages + non_cached_messages
large_tools = [
openai_large_tools: Final = [
{
"type": "function",
"function": {
@ -1548,6 +1548,11 @@ class TestContextCachingEndpoints:
}
for i in range(12)
]
large_tools: Final = litellm.utils.get_optional_params(
model="gemini-1.5-pro",
custom_llm_provider="gemini",
tools=openai_large_tools,
)["tools"]
optional_params = {
**self.sample_optional_params,