fix(gemini): harden count_tokens translation for anthropic, openai and native shapes

- detect anthropic-shaped input and route through the anthropic->openai adapter,
  with openai/responses-flat/native tools normalized instead of 400ing
- count server-side anthropic blocks and hosted tools (textified mass,
  codeExecution/urlContext) instead of silently dropping them
- merge thoughtSignature duplicates into their thought parts
- honor per-model system message support (gemini-1.5 folds into contents)
- wrap handler/client failures into typed litellm errors so callers fall back
- forward system/tools on the google-endpoints count route
This commit is contained in:
shrey kharbanda 2026-09-24 02:56:09 +00:00
parent cdd3d2f930
commit 61994a6704
11 changed files with 818 additions and 238 deletions

View file

@ -496,58 +496,69 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
import copy
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload
from litellm.llms.gemini.count_tokens.transformation import (
build_count_tokens_payload,
normalize_count_tokens_tools,
)
if contents is None and not messages:
return None
deployment = deployment or {}
count_tokens_params_request: Final = copy.deepcopy(deployment.get("litellm_params", {}))
payload: Final = (
build_count_tokens_payload(model=model_to_use, messages=messages, system=system, tools=tools)
if contents is None
else None
)
system_instruction: Final = payload.system_instruction if payload is not None else system
gemini_tools: Final = payload.tools if payload is not None else tools
count_tokens_params: Final = {
"model": model_to_use,
"contents": payload.contents if payload is not None else contents,
**(
{"system_instruction": system_instruction} # mutable-ok: kwargs dict for acount_tokens
if system_instruction is not None
else {} # mutable-ok: kwargs dict for acount_tokens
),
**(
{"tools": gemini_tools} # mutable-ok: kwargs dict for acount_tokens
if gemini_tools is not None
else {} # mutable-ok: kwargs dict for acount_tokens
),
}
count_tokens_params_request.update(count_tokens_params)
try:
payload: Final = (
build_count_tokens_payload(model=model_to_use, messages=messages, system=system, tools=tools)
if contents is None
else None
)
system_instruction: Final = (
payload.system_instruction
if payload is not None
else (
{"parts": [{"text": system}]} # mutable-ok: SystemInstructions wire shape
if isinstance(system, str)
else system
)
)
gemini_tools: Final = payload.tools if payload is not None else normalize_count_tokens_tools(tools)
count_tokens_params: Final = { # mutable-ok: kwargs dict for acount_tokens
"model": model_to_use,
"contents": payload.contents if payload is not None else contents,
**(
{"system_instruction": system_instruction} # mutable-ok: kwargs dict for acount_tokens
if system_instruction is not None
else {} # mutable-ok: kwargs dict for acount_tokens
),
**(
{"tools": gemini_tools} # mutable-ok: kwargs dict for acount_tokens
if gemini_tools is not None
else {} # mutable-ok: kwargs dict for acount_tokens
),
}
count_tokens_params_request.update(count_tokens_params)
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
client=client,
**count_tokens_params_request,
)
except (litellm.APIError, litellm.APIConnectionError) as e:
if result is not None:
return TokenCountResponse(
total_tokens=result.get("totalTokens", 0),
request_model=request_model,
model_used=model_to_use,
tokenizer_type="gemini_api",
original_response=result,
)
return None
except Exception as e:
# provider counting is best-effort: translation, credential, and request
# failures all degrade to the proxy's local-tokenizer fallback
return TokenCountResponse(
total_tokens=0,
request_model=request_model,
model_used=model_to_use,
tokenizer_type="gemini_api",
error=True,
error_message=e.message,
status_code=e.status_code,
error_message=getattr(e, "message", None) or str(e),
status_code=getattr(e, "status_code", None) or 500,
)
if result is not None:
return TokenCountResponse(
total_tokens=result.get("totalTokens", 0),
request_model=request_model,
model_used=model_to_use,
tokenizer_type=result.get("tokenizer_used", ""),
original_response=result,
)
return None

View file

@ -124,47 +124,45 @@ class GoogleAIStudioTokenCounter:
Exception: For any other unexpected errors
"""
# Prepare headers
headers, url = await self.validate_environment(
api_key=api_key,
api_base=api_base,
headers={},
model=model,
litellm_params=kwargs,
)
# Prepare request body - clean up contents to remove unsupported fields
cleaned_contents: Final = self._clean_contents_for_gemini_api(contents)
request_body: Final = (
{"contents": cleaned_contents} # mutable-ok: httpx json body takes a plain dict
if system_instruction is None and tools is None
else { # mutable-ok: httpx json body takes a plain dict
"generateContentRequest": { # mutable-ok: httpx json body takes a plain dict
"model": f"models/{model}",
"contents": cleaned_contents,
**(
{ # mutable-ok: httpx json body takes a plain dict
"systemInstruction": system_instruction,
}
if system_instruction is not None
else {} # mutable-ok: httpx json body takes a plain dict
),
**(
{ # mutable-ok: httpx json body takes a plain dict
"tools": tools,
}
if tools is not None
else {} # mutable-ok: httpx json body takes a plain dict
),
}
}
)
async_httpx_client: Final = client or get_async_httpx_client(
llm_provider=LlmProviders.GEMINI,
)
try:
headers, url = await self.validate_environment(
api_key=api_key,
api_base=api_base,
headers={}, # mutable-ok: validate_environment merges into this dict
model=model,
litellm_params=kwargs,
)
cleaned_contents: Final = self._clean_contents_for_gemini_api(contents)
request_body: Final = (
{"contents": cleaned_contents} # mutable-ok: httpx json body takes a plain dict
if system_instruction is None and tools is None
else { # mutable-ok: httpx json body takes a plain dict
"generateContentRequest": { # mutable-ok: httpx json body takes a plain dict
"model": f"models/{model}",
"contents": cleaned_contents,
**(
{ # mutable-ok: httpx json body takes a plain dict
"systemInstruction": system_instruction,
}
if system_instruction is not None
else {} # mutable-ok: httpx json body takes a plain dict
),
**(
{ # mutable-ok: httpx json body takes a plain dict
"tools": tools,
}
if tools is not None
else {} # mutable-ok: httpx json body takes a plain dict
),
}
}
)
async_httpx_client: Final = client or get_async_httpx_client(
llm_provider=LlmProviders.GEMINI,
)
response: Final = await async_httpx_client.post(url=url, headers=headers, json=request_body)
# Check for HTTP errors

View file

@ -1,10 +1,18 @@
"""Translate an Anthropic /v1/messages/count_tokens request into a Gemini
countTokens payload (contents + systemInstruction + tools)."""
"""Translate a token-count request into a Gemini countTokens payload
(contents + systemInstruction + tools).
Callers send Anthropic Messages shapes (/v1/messages/count_tokens) or
already-OpenAI shapes (/v1/responses/input_tokens, /utils/token_counter).
OpenAI input skips the Anthropic adapter, which only reads Anthropic fields
and would drop tool_calls, tool messages, and web search tool declarations.
"""
import json
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Final, cast
import litellm
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
LiteLLMAnthropicMessagesAdapter,
)
@ -15,6 +23,7 @@ from litellm.llms.vertex_ai.gemini.transformation import (
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig
from litellm.types.llms.anthropic import AnthropicMessagesRequest
from litellm.types.llms.vertex_ai import ContentType, SystemInstructions, Tools
from litellm.types.utils import AllMessageValues
@dataclass(frozen=True, slots=True)
@ -24,44 +33,393 @@ class GeminiCountTokensPayload:
tools: list[Tools] | None
_ANTHROPIC_PART_TYPES: Final = frozenset(
{
"tool_use",
"tool_result",
"thinking",
"redacted_thinking",
"image",
"document",
"server_tool_use",
"web_search_tool_result",
"web_fetch_tool_result",
"code_execution_tool_result",
"mcp_tool_use",
"mcp_tool_result",
"container_upload",
}
)
_ANTHROPIC_TOOL_TYPE_PREFIXES: Final = (
"web_search_",
"web_fetch_",
"code_execution_",
"computer_",
"text_editor_",
"bash_",
"mcp_",
)
# Server-side Anthropic content blocks the anthropic->openai adapter drops, so
# they would silently count as ~1 token each. They are flattened to text
# (closest token mass to the serialized block Anthropic itself bills).
_SERVER_SIDE_PART_TYPES: Final = frozenset(
{
"server_tool_use",
"web_search_tool_result",
"web_fetch_tool_result",
"code_execution_tool_result",
"code_execution_result",
"mcp_tool_use",
"mcp_tool_result",
"container_upload",
"redacted_thinking",
}
)
# Anthropic hosted tools with a native Gemini equivalent: mapping preserves
# roughly the hosted-tool token overhead instead of collapsing them to a
# name-only function declaration.
_ANTHROPIC_HOSTED_TOOL_TYPES: Final = (
("code_execution", "codeExecution"),
("web_fetch", "urlContext"),
)
_GEMINI_TOOL_KEYS: Final = frozenset(
{
"function_declarations",
"functionDeclarations",
"googleSearch",
"google_search",
"urlContext",
"url_context",
"codeExecution",
"code_execution",
"enterpriseWebSearch",
"googleSearchRetrieval",
"retrieval",
"computerUse",
"computer_use",
}
)
def _has_anthropic_shape(
system: object | None,
tools: Sequence[Mapping[str, object]] | None,
messages: Sequence[Mapping[str, object]] | None,
) -> bool:
# A top-level system given as a list of blocks only exists in the Anthropic API
if isinstance(system, list):
return True
for message in messages or ():
if not isinstance(message, Mapping):
continue
content = message.get("content")
if isinstance(content, list):
for part in content:
if isinstance(part, Mapping) and part.get("type") in _ANTHROPIC_PART_TYPES:
return True
for tool in tools or ():
if isinstance(tool, Mapping):
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):
return True
return False
def _is_web_search_tool(tool: Mapping[str, object]) -> bool:
tool_type = tool.get("type")
return isinstance(tool_type, str) and tool_type.startswith("web_search")
def _is_gemini_tool_shape(tool: Mapping[str, object]) -> bool:
return any(key in tool for key in _GEMINI_TOOL_KEYS)
def _hosted_tool_type(tool: Mapping[str, object]) -> str | None:
tool_type = tool.get("type")
if not isinstance(tool_type, str):
return None
for prefix, gemini_name in _ANTHROPIC_HOSTED_TOOL_TYPES:
if tool_type.startswith(prefix):
return gemini_name
return None
def _normalize_openai_tool(tool: Mapping[str, object]) -> dict[str, object]:
if "function" in tool:
return dict(tool) # mutable-ok: _map_function takes plain tool dicts
if "input_schema" in tool:
return { # mutable-ok: normalized tool dict for _map_function (anthropic function shape reaching the openai path)
"type": "function",
"function": { # mutable-ok: normalized tool dict for _map_function
"name": tool.get("name"),
"description": tool.get("description"),
"parameters": tool.get("input_schema"),
},
}
if tool.get("type") == "function" and "name" in tool:
return { # mutable-ok: normalized tool dict for _map_function (responses-api flat shape)
"type": "function",
"function": { # mutable-ok: normalized tool dict for _map_function
key: tool[key] for key in ("name", "description", "parameters", "strict") if key in tool
},
}
return dict(tool) # mutable-ok: _map_function takes plain tool dicts
def _apply_mixed_tool_drop_rule(merged: Sequence[Tools]) -> list[Tools] | None:
if not merged:
return None
optional_params: Final = { # mutable-ok: shared Vertex drop-rule mutates the tools list in place
"tools": list(merged) # mutable-ok: shared Vertex drop-rule mutates the tools list in place
}
VertexGeminiConfig._drop_search_tools_mixed_with_functions(optional_params)
kept: Final = optional_params["tools"]
return kept or None
def _map_to_gemini_tools(
openai_tools: Sequence[Mapping[str, object]],
web_search_options: object | None,
) -> list[Tools] | None:
merged: Final = (
VertexGeminiConfig()._map_function(
value=[dict(tool) for tool in openai_tools], # mutable-ok: _map_function takes plain tool dicts
optional_params={}, # mutable-ok: _map_function signature takes a dict
)
if openai_tools
else [] # mutable-ok: merged with the mapped tools list below
) + (
[VertexGeminiConfig()._map_web_search_options({})] # mutable-ok: merged tools list for the drop-rule
if web_search_options is not None
else [] # mutable-ok: merged with the mapped tools list
)
return _apply_mixed_tool_drop_rule(merged)
def normalize_count_tokens_tools(
tools: Sequence[Mapping[str, object]] | None,
) -> list[Tools] | None:
"""Tools arriving alongside native Gemini contents may be in Gemini,
OpenAI, Responses-API, or Anthropic shape. Gemini-shaped entries pass
through; the rest are normalized so mixed-shape callers do not 400."""
if not tools:
return None
gemini_shaped: Final = tuple(tool for tool in tools if _is_gemini_tool_shape(tool))
rest: Final = tuple(tool for tool in tools if not _is_gemini_tool_shape(tool))
mapped: Final = _map_to_gemini_tools(
openai_tools=tuple(_normalize_openai_tool(tool) for tool in rest if not _is_web_search_tool(tool)),
web_search_options={} # mutable-ok: truthy marker for _map_web_search_options
if any(_is_web_search_tool(tool) for tool in rest)
else None,
)
merged: Final = (
[
cast( # cast-ok: plain dicts for the Tools wire shape
Tools,
dict(tool), # mutable-ok: plain dicts for the Tools wire shape
)
for tool in gemini_shaped
]
+ list(mapped or []) # mutable-ok: concat with mapped tools
)
return _apply_mixed_tool_drop_rule(merged)
def _dedupe_thought_signature_parts(contents: Sequence[ContentType]) -> list[ContentType]:
"""The openai->gemini converter emits both a {thought: true, text} part
and a duplicate {thoughtSignature, text} part for one signed thinking
block; the signature belongs as a field on the thought part, not a second
part, or every reasoning turn is counted twice."""
def _merge(parts: Sequence[object]) -> Sequence[object]:
sig_by_text: Final = { # mutable-ok: text->signature lookup
part.get("text"): part.get("thoughtSignature")
for part in parts
if isinstance(part, Mapping)
and part.get("thoughtSignature") is not None
and part.get("thought") is not True
}
if not sig_by_text:
return parts
thought_texts: Final = frozenset(
part.get("text") for part in parts if isinstance(part, Mapping) and part.get("thought") is True
)
return tuple(
(
{ # mutable-ok: rebuilt part with merged signature
**dict(part), # mutable-ok: rebuilt part with merged signature
"thoughtSignature": sig_by_text[part.get("text")],
}
if isinstance(part, Mapping)
and part.get("thought") is True
and sig_by_text.get(part.get("text")) is not None
else part
)
for part in parts
if not (
isinstance(part, Mapping)
and part.get("thoughtSignature") is not None
and part.get("thought") is not True
and part.get("text") in thought_texts
)
)
return [ # mutable-ok: rebuilt contents list
cast( # cast-ok: same ContentType shape with deduped parts
ContentType,
{ # mutable-ok: rebuilt content with deduped parts
**dict(content), # mutable-ok: rebuilt content with deduped parts
"parts": _merge(content.get("parts", [])), # mutable-ok: default parts list for the merge
},
)
if isinstance(content.get("parts"), list)
else content
for content in contents
]
def _textify_server_side_blocks(
messages: Sequence[Mapping[str, object]],
) -> tuple[Mapping[str, object], ...]:
"""Flatten server-side Anthropic content blocks to text so the
anthropic->openai adapter (which drops them) still counts their mass."""
def _textify(content: object) -> object:
if not isinstance(content, list):
return content
return [ # mutable-ok: rebuilt content list
(
{ # mutable-ok: textified block for counting
"type": "text",
"text": json.dumps(dict(block), ensure_ascii=False), # mutable-ok: plain dict copy for the dump
}
if isinstance(block, Mapping) and block.get("type") in _SERVER_SIDE_PART_TYPES
else block
)
for block in content
] # mutable-ok: rebuilt content list
return tuple(
(
{**dict(message), "content": _textify(message.get("content"))} # mutable-ok: rebuilt message
if isinstance(message.get("content"), list)
else message
)
for message in messages
)
def _payload_from_openai_parts(
model: str,
messages: Sequence[object],
tools: Sequence[Mapping[str, object]],
web_search_options: object | None,
) -> GeminiCountTokensPayload:
system_instruction, remaining_messages = _transform_system_message(
supports_system_message=litellm.supports_system_messages(model=model, custom_llm_provider="gemini"),
messages=cast( # cast-ok: chat-shaped message dicts accepted by the helper
"list[AllMessageValues]",
list(messages), # mutable-ok: helper contract takes a list
),
)
contents: Final = _dedupe_thought_signature_parts(
_gemini_convert_messages_with_history(
messages=remaining_messages,
model=model,
custom_llm_provider="gemini",
)
)
return GeminiCountTokensPayload(
contents=contents,
system_instruction=system_instruction,
tools=_map_to_gemini_tools(openai_tools=tools, web_search_options=web_search_options),
)
def _build_anthropic_payload(
model: str,
messages: Sequence[Mapping[str, object]],
system: object | None,
tools: Sequence[Mapping[str, object]] | None,
) -> GeminiCountTokensPayload:
hosted_tools: Final = tuple(tool for tool in tools or () if _hosted_tool_type(tool) is not None)
adapter_tools: Final = tuple(tool for tool in tools or () if _hosted_tool_type(tool) is None)
anthropic_request: Final[AnthropicMessagesRequest] = cast( # cast-ok: adapter reads only the keys supplied
AnthropicMessagesRequest,
{ # mutable-ok: transient request dict for the anthropic adapter
"model": model,
"messages": list( # mutable-ok: adapter contract takes a list of messages
_textify_server_side_blocks(messages)
),
**({"system": system} if system else {}), # mutable-ok: transient request dict for the anthropic adapter
**(
{"tools": list(adapter_tools)}
if adapter_tools
else {} # mutable-ok: transient request dict for the anthropic adapter
),
},
)
openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
anthropic_request, custom_llm_provider="gemini"
)
openai_tools: Final = openai_request.get("tools")
payload: Final = _payload_from_openai_parts(
model=model,
messages=openai_request["messages"],
tools=tuple(dict(tool) for tool in openai_tools) # mutable-ok: plain dicts for the normalizer
if openai_tools
else (),
web_search_options=openai_request.get("web_search_options"),
)
if not hosted_tools:
return payload
merged_tools: Final = _apply_mixed_tool_drop_rule(
list(payload.tools or []) # mutable-ok: merged tools for the drop-rule
+ [ # mutable-ok: merged tools list for the drop-rule
cast(Tools, {gemini_name: {}}) # cast-ok: hosted tool wire shape # mutable-ok: hosted tool wire shape
for gemini_name in {_hosted_tool_type(tool) for tool in hosted_tools} # mutable-ok: set dedupe
]
)
return GeminiCountTokensPayload(
contents=payload.contents,
system_instruction=payload.system_instruction,
tools=merged_tools,
)
def _build_openai_payload(
model: str,
messages: Sequence[Mapping[str, object]],
system: object | None,
tools: Sequence[Mapping[str, object]] | None,
) -> GeminiCountTokensPayload:
openai_messages: Final = (
[{"role": "system", "content": system}, *messages] # mutable-ok: transient list for the message converter
if system is not None
else list(messages) # mutable-ok: transient list for the message converter
)
return _payload_from_openai_parts(
model=model,
messages=openai_messages,
tools=tuple(_normalize_openai_tool(tool) for tool in tools or () if not _is_web_search_tool(tool)),
web_search_options={} # mutable-ok: truthy marker for _map_web_search_options
if any(_is_web_search_tool(tool) for tool in tools or ())
else None,
)
def build_count_tokens_payload(
model: str,
messages: Sequence[Mapping[str, object]],
system: object | None,
tools: Sequence[Mapping[str, object]] | None,
) -> GeminiCountTokensPayload:
anthropic_request: Final[AnthropicMessagesRequest] = cast( # cast-ok: adapter reads only the keys supplied
AnthropicMessagesRequest,
{ # mutable-ok: transient request dict for the anthropic adapter
"model": model,
"messages": list(messages), # mutable-ok: adapter contract takes a list of messages
**({"system": system} if system else {}), # mutable-ok: transient request dict for the anthropic adapter
**({"tools": list(tools)} if tools else {}), # mutable-ok: transient request dict for the anthropic adapter
},
)
openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
anthropic_request, custom_llm_provider="gemini"
)
system_instruction, remaining_messages = _transform_system_message(
supports_system_message=True,
messages=list(openai_request["messages"]), # mutable-ok: helper pops the leading system message
)
contents: Final = _gemini_convert_messages_with_history(
messages=remaining_messages,
model=model,
custom_llm_provider="gemini",
)
openai_tools: Final = openai_request.get("tools")
gemini_tools: Final = (
VertexGeminiConfig()._map_function(
value=[dict(tool) for tool in openai_tools], # mutable-ok: _map_function takes plain tool dicts
optional_params={}, # mutable-ok: _map_function signature takes a dict
)
if openai_tools
else None
)
return GeminiCountTokensPayload(
contents=contents,
system_instruction=system_instruction,
tools=gemini_tools,
)
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)

View file

@ -181,6 +181,8 @@ async def google_count_tokens(request: Request, model_name: str):
model=model_name,
contents=contents,
messages=messages, # compatibility when use openai-like endpoint
tools=data.get("tools"),
system=data.get("systemInstruction"),
)
# Call the internal token counter function with direct request flag set to False

View file

@ -58,3 +58,38 @@ async def test_acount_tokens_keeps_contents_body_without_system_or_tools():
body = json.loads(recorded[-1].content)
assert body == {"contents": [{"role": "user", "parts": [{"text": "hi"}]}]}
@pytest.mark.asyncio
async def test_acount_tokens_wraps_unexpected_error_in_api_error():
import litellm
def _handler(request: httpx.Request) -> httpx.Response:
raise RuntimeError("transport exploded")
client = httpx.AsyncClient(transport=httpx.MockTransport(_handler))
with pytest.raises(litellm.APIError) as excinfo:
await GoogleAIStudioTokenCounter().acount_tokens(
model="gemini-2.5-flash",
contents=[{"role": "user", "parts": [{"text": "hello"}]}],
api_key="test-key",
client=client,
)
assert excinfo.value.status_code == 500
@pytest.mark.asyncio
async def test_acount_tokens_wraps_malformed_contents_error_in_api_error():
import litellm
client = httpx.AsyncClient(transport=httpx.MockTransport(lambda request: httpx.Response(200, json={})))
with pytest.raises(litellm.APIError):
await GoogleAIStudioTokenCounter().acount_tokens(
model="gemini-2.5-flash",
contents=5, # pyright: ignore[reportArgumentType] # malformed caller input exercises the error boundary
api_key="test-key",
client=client,
)

View file

@ -1,3 +1,5 @@
import json
from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload
@ -60,3 +62,222 @@ def test_build_count_tokens_payload_passes_openai_tools_through():
assert payload.tools is not None
assert payload.tools[0]["function_declarations"][0]["name"] == "get_weather"
def test_build_count_tokens_payload_keeps_openai_tool_calls_and_results():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[
{"role": "system", "content": "be helpful"},
{"role": "user", "content": "what's the weather?"},
{
"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=None,
)
assert payload.system_instruction is not None
assert payload.system_instruction["parts"][0].get("text") == "be helpful"
assert payload.contents[0]["parts"][0].get("text") == "what's the weather?"
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"
assert function_response["response"] == {"content": "sunny"}
def test_build_count_tokens_payload_maps_anthropic_web_search_tool():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[{"role": "user", "content": "hi"}],
system=None,
tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 3}],
)
assert payload.tools == [{"googleSearch": {}}]
def test_build_count_tokens_payload_maps_openai_web_search_tool():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[{"role": "user", "content": "hi"}],
system=None,
tools=[{"type": "web_search_preview"}],
)
assert payload.tools == [{"googleSearch": {}}]
def test_build_count_tokens_payload_wraps_responses_api_tool():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[{"role": "user", "content": "hi"}],
system=None,
tools=[
{
"type": "function",
"name": "get_weather",
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}},
}
],
)
assert payload.tools is not None
function_declaration = payload.tools[0]["function_declarations"][0]
assert function_declaration["name"] == "get_weather"
assert function_declaration["parameters"] == {
"type": "object",
"properties": {"city": {"type": "string"}},
}
def test_build_count_tokens_payload_drops_search_tool_when_mixed_with_functions():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[{"role": "user", "content": "hi"}],
system=None,
tools=[
{"type": "web_search_20250305", "name": "web_search"},
{"name": "get_weather", "input_schema": {"type": "object"}},
],
)
assert payload.tools == [{"function_declarations": [{"name": "get_weather", "parameters": {"type": "object"}}]}]
def test_build_count_tokens_payload_merges_thought_signature_into_one_part():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": [
{
"type": "thinking",
"thinking": "let me reason about this",
"signature": "sig123",
},
{"type": "text", "text": "answer"},
],
},
],
system=None,
tools=None,
)
assert payload.contents[1]["parts"] == (
{"thought": True, "text": "let me reason about this", "thoughtSignature": "sig123"},
{"text": "answer"},
)
def test_build_count_tokens_payload_counts_server_side_blocks_as_text():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[
{
"role": "assistant",
"content": [{"type": "server_tool_use", "id": "s1", "name": "web_search", "input": {"query": "q"}}],
},
{
"role": "user",
"content": [
{
"type": "web_search_tool_result",
"tool_use_id": "s1",
"content": [{"type": "web_search_result", "url": "u", "title": "t"}],
}
],
},
],
system=None,
tools=None,
)
assert payload.contents[0]["parts"][0].get("text") == json.dumps(
{"type": "server_tool_use", "id": "s1", "name": "web_search", "input": {"query": "q"}},
ensure_ascii=False,
)
result_text = payload.contents[1]["parts"][0].get("text")
assert isinstance(result_text, str)
assert "web_search_tool_result" in result_text
assert "tool_use_id" in result_text
def test_build_count_tokens_payload_maps_anthropic_hosted_tools_to_native_gemini_tools():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[{"role": "user", "content": "hi"}],
system=None,
tools=[{"type": "code_execution_20250522", "name": "code_execution"}],
)
assert payload.tools == [{"codeExecution": {}}]
def test_build_count_tokens_payload_maps_web_fetch_tool_to_url_context():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[{"role": "user", "content": "hi"}],
system=None,
tools=[{"type": "web_fetch_20250910", "name": "web_fetch"}],
)
assert payload.tools == [{"urlContext": {}}]
def test_build_count_tokens_payload_drops_url_context_when_mixed_with_functions():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[{"role": "user", "content": "hi"}],
system=None,
tools=[
{"type": "web_fetch_20250910", "name": "web_fetch"},
{"name": "get_weather", "input_schema": {"type": "object"}},
],
)
assert payload.tools == [{"function_declarations": [{"name": "get_weather", "parameters": {"type": "object"}}]}]
def test_build_count_tokens_payload_folds_system_into_contents_for_models_without_system_support():
payload = build_count_tokens_payload(
model="gemini-1.5-flash",
messages=[{"role": "user", "content": "hi"}],
system="be nice",
tools=None,
)
assert payload.system_instruction is None
assert [part.get("text") for part in payload.contents[0]["parts"]] == ["be nice", "hi"]
def test_normalize_count_tokens_tools_handles_each_tool_shape():
from litellm.llms.gemini.count_tokens.transformation import normalize_count_tokens_tools
assert normalize_count_tokens_tools(None) is None
assert normalize_count_tokens_tools([{"function_declarations": [{"name": "g"}]}]) == [
{"function_declarations": [{"name": "g"}]}
]
assert normalize_count_tokens_tools([{"googleSearch": {}}]) == [{"googleSearch": {}}]
assert normalize_count_tokens_tools(
[{"type": "function", "function": {"name": "f", "parameters": {"type": "object"}}}]
) == [{"function_declarations": [{"name": "f", "parameters": {"type": "object"}}]}]
assert normalize_count_tokens_tools([{"name": "f", "input_schema": {"type": "object"}}]) == [
{"function_declarations": [{"name": "f", "parameters": {"type": "object"}}]}
]
assert normalize_count_tokens_tools([{"googleSearch": {}}, {"type": "function", "function": {"name": "f"}}]) == [
{"function_declarations": [{"name": "f"}]}
]

View file

@ -200,14 +200,6 @@ def test_maps_no_usage_details():
assert cost_per_google_maps_grounding_request(usage=usage, model_info=model_info) == 0.0
def _image_response_with_web_search(web_search_requests):
usage = ImageUsage(
input_tokens=20,
@ -223,10 +215,6 @@ def _image_response_with_web_search(web_search_requests):
return ImageResponse(data=[ImageObject(b64_json="img1")], usage=usage)
@pytest.mark.parametrize(
"traffic_type, expected_service_tier",
[
@ -242,9 +230,7 @@ def _image_response_with_web_search(web_search_requests):
("SOMETHING_UNKNOWN", None),
],
)
def test_map_traffic_type_to_service_tier(
traffic_type: str | None, expected_service_tier: str | None
):
def test_map_traffic_type_to_service_tier(traffic_type: str | None, expected_service_tier: str | None):
"""
Gemini/Vertex usageMetadata.trafficType maps to the LiteLLM service_tier
that selects flex/priority cost keys. ON_DEMAND_FLEX (Vertex's flex opt-in
@ -252,9 +238,7 @@ def test_map_traffic_type_to_service_tier(
"""
from litellm.cost_calculator import _map_traffic_type_to_service_tier
assert (
_map_traffic_type_to_service_tier(traffic_type) == expected_service_tier
)
assert _map_traffic_type_to_service_tier(traffic_type) == expected_service_tier
# Alias targets are the `modelVersion` returned by
@ -267,9 +251,7 @@ def test_map_traffic_type_to_service_tier(
("gemini/gemini-pro-latest", "gemini/gemini-3.1-pro-preview"),
],
)
def test_latest_aliases_cost_the_same_as_their_current_target(
monkeypatch, alias, target
):
def test_latest_aliases_cost_the_same_as_their_current_target(monkeypatch, alias, target):
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))

View file

@ -28,7 +28,7 @@ def test_gemini_completion_no_api_key():
del os.environ[key]
# Test without mock_response to ensure actual API key validation
with pytest.raises(Exception, match='in _complete_vertex_ai_beta') as exc_info:
with pytest.raises(Exception, match="in _complete_vertex_ai_beta") as exc_info:
completion(
model="gemini/gemini-1.5-flash",
messages=[{"role": "user", "content": "Test message"}],
@ -60,7 +60,7 @@ def test_gemini_completion_no_api_key_with_mock():
with patch("litellm.get_secret") as mock_get_secret:
mock_get_secret.return_value = None
with pytest.raises(Exception, match='in _complete_vertex_ai_beta') as exc_info:
with pytest.raises(Exception, match="in _complete_vertex_ai_beta") as exc_info:
completion(
model="gemini/gemini-1.5-flash",
messages=[{"role": "user", "content": "Test message"}],
@ -95,7 +95,4 @@ def test_gemini_completion_both_env_vars(monkeypatch, api_key_env):
messages=[{"role": "user", "content": f"Test with {api_key_env}"}],
mock_response=f"Mocked response using {api_key_env}",
)
assert (
response["choices"][0]["message"]["content"]
== f"Mocked response using {api_key_env}"
)
assert response["choices"][0]["message"]["content"] == f"Mocked response using {api_key_env}"

View file

@ -37,18 +37,10 @@ class TestGeminiModelInfo:
# Test edge cases where model names end with characters from "models/"
# These would be incorrectly processed if using strip("models/") instead of replace("models/", "")
models = [
{
"name": "models/gemini-1.5-pro"
}, # ends with 'o' - would become "gemini-1.5-pr" with strip()
{
"name": "models/test-model"
}, # ends with 'l' - would become "gemini/test-mode" with strip()
{
"name": "models/custom-models"
}, # ends with 's' - would become "gemini/custom-model" with strip()
{
"name": "models/demo"
}, # ends with 'o' - would become "gemini/dem" with strip()
{"name": "models/gemini-1.5-pro"}, # ends with 'o' - would become "gemini-1.5-pr" with strip()
{"name": "models/test-model"}, # ends with 'l' - would become "gemini/test-mode" with strip()
{"name": "models/custom-models"}, # ends with 's' - would become "gemini/custom-model" with strip()
{"name": "models/demo"}, # ends with 'o' - would become "gemini/dem" with strip()
]
result = gemini_model_info.process_model_name(models)
@ -99,16 +91,10 @@ class TestGoogleAIStudioTokenCounter:
token_counter = GoogleAIStudioTokenCounter()
# Test with gemini provider - should return True
assert (
token_counter.should_use_token_counting_api(LlmProviders.GEMINI.value)
is True
)
assert token_counter.should_use_token_counting_api(LlmProviders.GEMINI.value) is True
# Test with other providers - should return False
assert (
token_counter.should_use_token_counting_api(LlmProviders.OPENAI.value)
is False
)
assert token_counter.should_use_token_counting_api(LlmProviders.OPENAI.value) is False
assert token_counter.should_use_token_counting_api("anthropic") is False
assert token_counter.should_use_token_counting_api("vertex_ai") is False
@ -158,9 +144,7 @@ class TestGoogleAIStudioTokenCounter:
assert result.original_response == mock_response
# Verify the mock was called correctly
mock_acount_tokens.assert_called_once_with(
model=model_to_use, contents=contents, client=None
)
mock_acount_tokens.assert_called_once_with(model=model_to_use, contents=contents, client=None)
@pytest.mark.asyncio
async def test_count_tokens_translates_anthropic_messages_system_and_tools(self):
@ -268,6 +252,51 @@ class TestGoogleAIStudioTokenCounter:
assert result.total_tokens == 0
assert result.error_message is not None
@pytest.mark.asyncio
async def test_count_tokens_translation_error_falls_back(self):
"""A crash translating bad message shapes must surface as an error
TokenCountResponse so the proxy falls back instead of 500ing."""
token_counter = GoogleAIStudioTokenCounter()
result = await token_counter.count_tokens(
model_to_use="gemini-2.5-flash",
messages=[{"role": "tool", "content": "orphaned result", "tool_call_id": "missing-call"}],
contents=None,
deployment={"litellm_params": {"api_key": "test-key"}},
request_model="gemini/gemini-2.5-flash",
)
assert result is not None
assert result.error is True
assert result.status_code == 500
assert result.total_tokens == 0
assert result.error_message is not None
@pytest.mark.asyncio
async def test_count_tokens_unexpected_handler_error_returns_error_response(self):
"""A non-litellm exception escaping the handler must still surface as an
error TokenCountResponse so the proxy can fall back."""
token_counter = GoogleAIStudioTokenCounter()
with patch(
"litellm.llms.gemini.count_tokens.handler.GoogleAIStudioTokenCounter.acount_tokens",
new_callable=AsyncMock,
) as mock_acount_tokens:
mock_acount_tokens.side_effect = RuntimeError("unexpected failure")
result = await token_counter.count_tokens(
model_to_use="gemini-2.5-flash",
messages=[{"role": "user", "content": "hello"}],
contents=None,
deployment=None,
request_model="gemini/gemini-2.5-flash",
)
assert result is not None
assert result.error is True
assert result.status_code == 500
assert "unexpected failure" in (result.error_message or "")
@pytest.mark.asyncio
async def test_count_tokens_returns_none_without_contents_or_messages(self):
token_counter = GoogleAIStudioTokenCounter()
@ -297,9 +326,7 @@ class TestGoogleAIStudioTokenCounter:
"functionResponse": {
"id": "read_many_files-1757526647518-730a691aac11c", # This should be removed
"name": "read_many_files",
"response": {
"output": "No files matching the criteria were found or all were skipped."
},
"response": {"output": "No files matching the criteria were found or all were skipped."},
}
}
],
@ -308,9 +335,7 @@ class TestGoogleAIStudioTokenCounter:
]
# Clean the contents
cleaned_contents = token_counter._clean_contents_for_gemini_api(
contents_with_id
)
cleaned_contents = token_counter._clean_contents_for_gemini_api(contents_with_id)
# Verify the 'id' field was removed
function_response = cleaned_contents[1]["parts"][0]["functionResponse"]
@ -319,8 +344,7 @@ class TestGoogleAIStudioTokenCounter:
assert "response" in function_response
assert function_response["name"] == "read_many_files"
assert (
function_response["response"]["output"]
== "No files matching the criteria were found or all were skipped."
function_response["response"]["output"] == "No files matching the criteria were found or all were skipped."
)
def test_clean_contents_for_gemini_api_preserves_other_fields(self):
@ -336,9 +360,7 @@ class TestGoogleAIStudioTokenCounter:
]
# Clean the contents
cleaned_contents = token_counter._clean_contents_for_gemini_api(
contents_without_function_response
)
cleaned_contents = token_counter._clean_contents_for_gemini_api(contents_without_function_response)
# Verify the contents are unchanged
assert cleaned_contents == contents_without_function_response

View file

@ -189,9 +189,7 @@ def test_gemini_image_generation_usage_includes_chat_token_details():
assert usage["output_tokens_details"]["text_tokens"] == 596
assert usage["output_tokens_details"]["image_tokens"] == 1120
logging_usage = StandardLoggingPayloadSetup.get_usage_as_dict(
response_obj=result.model_dump()
)
logging_usage = StandardLoggingPayloadSetup.get_usage_as_dict(response_obj=result.model_dump())
assert logging_usage["completion_tokens_details"]["text_tokens"] == 596
assert logging_usage["completion_tokens_details"]["image_tokens"] == 1120
@ -263,18 +261,14 @@ def test_gemini_image_generation_preserves_tool_config_side_effect():
config = GoogleImageGenConfig()
mapped = config.map_openai_params(
non_default_params={
"tools": [{"googleMaps": {"latitude": 37.7, "longitude": -122.4}}]
},
non_default_params={"tools": [{"googleMaps": {"latitude": 37.7, "longitude": -122.4}}]},
optional_params={},
model="gemini-3.1-flash-image-preview",
drop_params=False,
)
assert mapped["tools"] == [{"googleMaps": {}}]
assert mapped["toolConfig"] == {
"retrievalConfig": {"latLng": {"latitude": 37.7, "longitude": -122.4}}
}
assert mapped["toolConfig"] == {"retrievalConfig": {"latLng": {"latitude": 37.7, "longitude": -122.4}}}
request = config.transform_image_generation_request(
model="gemini-3.1-flash-image-preview",
@ -285,9 +279,7 @@ def test_gemini_image_generation_preserves_tool_config_side_effect():
)
assert request["tools"] == [{"googleMaps": {}}]
assert request["toolConfig"] == {
"retrievalConfig": {"latLng": {"latitude": 37.7, "longitude": -122.4}}
}
assert request["toolConfig"] == {"retrievalConfig": {"latLng": {"latitude": 37.7, "longitude": -122.4}}}
def test_gemini_image_generation_usage_without_output_details_treats_output_as_image():
@ -351,9 +343,7 @@ def test_gemini_image_generation_response_tracks_web_search_requests():
}
]
},
"groundingMetadata": {
"webSearchQueries": ["latest iphone", "iphone colors"]
},
"groundingMetadata": {"webSearchQueries": ["latest iphone", "iphone colors"]},
}
],
"usageMetadata": {

View file

@ -19,9 +19,7 @@ class TestGeminiTTSTransformation:
config = GoogleAIStudioGeminiConfig()
# Test TTS models (both preview and non-preview versions)
assert (
config.is_model_gemini_audio_model("gemini-2.5-flash-preview-tts") == True
)
assert config.is_model_gemini_audio_model("gemini-2.5-flash-preview-tts") == True
assert config.is_model_gemini_audio_model("gemini-2.5-pro-preview-tts") == True
assert config.is_model_gemini_audio_model("gemini-2.5-flash-tts") == True
assert config.is_model_gemini_audio_model("gemini-2.5-pro-tts") == True
@ -66,10 +64,7 @@ class TestGeminiTTSTransformation:
assert "speechConfig" in result
assert "voiceConfig" in result["speechConfig"]
assert "prebuiltVoiceConfig" in result["speechConfig"]["voiceConfig"]
assert (
result["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"]
== "Kore"
)
assert result["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore"
# Check response modalities
assert "responseModalities" in result
@ -78,9 +73,7 @@ class TestGeminiTTSTransformation:
def test_gemini_tts_audio_parameter_mapping_with_language_code(self):
config = GoogleAIStudioGeminiConfig()
non_default_params = {
"audio": {"voice": "Kore", "format": "pcm16", "language_code": "en-US"}
}
non_default_params = {"audio": {"voice": "Kore", "format": "pcm16", "language_code": "en-US"}}
optional_params = {}
result = config.map_openai_params(
@ -92,17 +85,12 @@ class TestGeminiTTSTransformation:
assert "speechConfig" in result
assert result["speechConfig"]["languageCode"] == "en-US"
assert (
result["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"]
== "Kore"
)
assert result["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore"
def test_map_audio_params_language_code(self):
config = GoogleAIStudioGeminiConfig()
result = config._map_audio_params(
{"voice": "Kore", "format": "pcm16", "language_code": "de-DE"}
)
result = config._map_audio_params({"voice": "Kore", "format": "pcm16", "language_code": "de-DE"})
assert result["languageCode"] == "de-DE"
assert result["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore"
@ -198,9 +186,7 @@ class TestGeminiTTSTransformation:
}
optional_params = {}
with pytest.raises(
ValueError, match="Unsupported audio format for Gemini TTS models"
):
with pytest.raises(ValueError, match="Unsupported audio format for Gemini TTS models"):
config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
@ -257,9 +243,7 @@ class TestGeminiTTSSpeechConfigInRequestBody:
("gemini-2.5-pro-tts", "vertex_ai"),
],
)
def test_speechconfig_in_generation_config_transform_request_body(
self, model, custom_llm_provider
):
def test_speechconfig_in_generation_config_transform_request_body(self, model, custom_llm_provider):
"""Test that speechConfig is included in generationConfig after _transform_request_body()"""
from litellm.llms.vertex_ai.gemini.transformation import (
_transform_request_body,
@ -267,9 +251,7 @@ class TestGeminiTTSSpeechConfigInRequestBody:
# Simulate optional_params after map_openai_params() has run
optional_params = {
"speechConfig": {
"voiceConfig": {"prebuiltVoiceConfig": {"voiceName": "Kore"}}
},
"speechConfig": {"voiceConfig": {"prebuiltVoiceConfig": {"voiceName": "Kore"}}},
"responseModalities": ["AUDIO"],
}
@ -292,12 +274,7 @@ class TestGeminiTTSSpeechConfigInRequestBody:
f"speechConfig was filtered out of generationConfig for model={model}, provider={custom_llm_provider}. "
"Ensure speechConfig is in the GenerationConfig TypedDict."
)
assert (
generation_config["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"][
"voiceName"
]
== "Kore"
)
assert generation_config["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore"
@pytest.mark.parametrize(
"model,custom_llm_provider",
@ -351,18 +328,12 @@ class TestGeminiTTSSpeechConfigInRequestBody:
f"speechConfig was filtered out during _transform_request_body() for model={model}, provider={custom_llm_provider}. "
"This breaks Gemini TTS - speechConfig must be in GenerationConfig TypedDict."
)
assert (
generation_config["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"][
"voiceName"
]
== "Puck"
)
assert generation_config["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Puck"
# Also verify responseModalities is present
assert "responseModalities" in generation_config
assert "AUDIO" in generation_config["responseModalities"]
@pytest.mark.parametrize(
"model,custom_llm_provider",
[
@ -381,9 +352,7 @@ class TestGeminiTTSSpeechConfigInRequestBody:
config = VertexGeminiConfig()
non_default_params = {
"audio": {"voice": "Puck", "format": "pcm16", "language_code": "pt-BR"}
}
non_default_params = {"audio": {"voice": "Puck", "format": "pcm16", "language_code": "pt-BR"}}
optional_params = {}
mapped_params = config.map_openai_params(
@ -406,12 +375,7 @@ class TestGeminiTTSSpeechConfigInRequestBody:
generation_config = request_body["generationConfig"]
assert generation_config["speechConfig"]["languageCode"] == "pt-BR"
assert (
generation_config["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"][
"voiceName"
]
== "Puck"
)
assert generation_config["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Puck"
assert "AUDIO" in generation_config["responseModalities"]