mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
4274bdda44
commit
8e6d99d74a
3 changed files with 126 additions and 33 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue