diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 5d7956059e4..cdd2d0654be 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -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: diff --git a/tests/unit/litellm_core_utils/test_token_counter.py b/tests/unit/litellm_core_utils/test_token_counter.py index f7ded4f3fa8..c71b1496bdd 100644 --- a/tests/unit/litellm_core_utils/test_token_counter.py +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -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""" diff --git a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 283ed3710d0..67d78d6030e 100644 --- a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -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,