mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
cdd3d2f930
commit
61994a6704
11 changed files with 818 additions and 238 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"}]}
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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=""))
|
||||
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue