From 61994a6704bb38658dd6c83a7b5d70b698557562 Mon Sep 17 00:00:00 2001 From: shrey kharbanda Date: Thu, 24 Sep 2026 02:56:09 +0000 Subject: [PATCH] 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 --- litellm/llms/gemini/common_utils.py | 85 ++-- litellm/llms/gemini/count_tokens/handler.py | 78 ++-- .../gemini/count_tokens/transformation.py | 432 ++++++++++++++++-- litellm/proxy/google_endpoints/endpoints.py | 2 + .../llms/gemini/count_tokens/test_handler.py | 35 ++ .../count_tokens/test_transformation.py | 221 +++++++++ .../llms/gemini/test_cost_calculator.py | 24 +- .../llms/gemini/test_gemini_client_setup.py | 9 +- .../llms/gemini/test_gemini_common_utils.py | 90 ++-- ..._gemini_image_generation_transformation.py | 20 +- .../llms/gemini/test_gemini_tts.py | 60 +-- 11 files changed, 818 insertions(+), 238 deletions(-) diff --git a/litellm/llms/gemini/common_utils.py b/litellm/llms/gemini/common_utils.py index dd905d5a72a..d2e1f0ad294 100644 --- a/litellm/llms/gemini/common_utils.py +++ b/litellm/llms/gemini/common_utils.py @@ -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 diff --git a/litellm/llms/gemini/count_tokens/handler.py b/litellm/llms/gemini/count_tokens/handler.py index 142bafc89f2..3894a6f4837 100644 --- a/litellm/llms/gemini/count_tokens/handler.py +++ b/litellm/llms/gemini/count_tokens/handler.py @@ -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 diff --git a/litellm/llms/gemini/count_tokens/transformation.py b/litellm/llms/gemini/count_tokens/transformation.py index 8e2c4a66ee2..5a046544901 100644 --- a/litellm/llms/gemini/count_tokens/transformation.py +++ b/litellm/llms/gemini/count_tokens/transformation.py @@ -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) diff --git a/litellm/proxy/google_endpoints/endpoints.py b/litellm/proxy/google_endpoints/endpoints.py index 9d39430cfa9..00276fa06d8 100644 --- a/litellm/proxy/google_endpoints/endpoints.py +++ b/litellm/proxy/google_endpoints/endpoints.py @@ -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 diff --git a/tests/test_litellm/llms/gemini/count_tokens/test_handler.py b/tests/test_litellm/llms/gemini/count_tokens/test_handler.py index 3a3068fba4b..c245e02141e 100644 --- a/tests/test_litellm/llms/gemini/count_tokens/test_handler.py +++ b/tests/test_litellm/llms/gemini/count_tokens/test_handler.py @@ -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, + ) diff --git a/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py b/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py index 829f8d49e29..13d186f3021 100644 --- a/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py +++ b/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py @@ -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"}]} + ] diff --git a/tests/test_litellm/llms/gemini/test_cost_calculator.py b/tests/test_litellm/llms/gemini/test_cost_calculator.py index b2633c6091b..8acba0e5cdc 100644 --- a/tests/test_litellm/llms/gemini/test_cost_calculator.py +++ b/tests/test_litellm/llms/gemini/test_cost_calculator.py @@ -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="")) diff --git a/tests/test_litellm/llms/gemini/test_gemini_client_setup.py b/tests/test_litellm/llms/gemini/test_gemini_client_setup.py index 48b010aca48..7edc7357759 100644 --- a/tests/test_litellm/llms/gemini/test_gemini_client_setup.py +++ b/tests/test_litellm/llms/gemini/test_gemini_client_setup.py @@ -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}" diff --git a/tests/test_litellm/llms/gemini/test_gemini_common_utils.py b/tests/test_litellm/llms/gemini/test_gemini_common_utils.py index 5262adbb5f4..062bf207357 100644 --- a/tests/test_litellm/llms/gemini/test_gemini_common_utils.py +++ b/tests/test_litellm/llms/gemini/test_gemini_common_utils.py @@ -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 diff --git a/tests/test_litellm/llms/gemini/test_gemini_image_generation_transformation.py b/tests/test_litellm/llms/gemini/test_gemini_image_generation_transformation.py index d9509856759..6025fddc07d 100644 --- a/tests/test_litellm/llms/gemini/test_gemini_image_generation_transformation.py +++ b/tests/test_litellm/llms/gemini/test_gemini_image_generation_transformation.py @@ -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": { diff --git a/tests/test_litellm/llms/gemini/test_gemini_tts.py b/tests/test_litellm/llms/gemini/test_gemini_tts.py index 4893825373a..19b992226f1 100644 --- a/tests/test_litellm/llms/gemini/test_gemini_tts.py +++ b/tests/test_litellm/llms/gemini/test_gemini_tts.py @@ -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"]