This commit is contained in:
PINYO PATTANAWASANPORN 2026-10-05 07:10:19 -04:00 • committed by GitHub
commit 454430a331
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 434 additions and 15 deletions

View file

@ -4,6 +4,8 @@ Transformation logic from OpenAI format to Gemini format.
Why separate file? Make it easy to see how transformation works
"""
import base64
import binascii
import json
import os
import re
@ -12,8 +14,6 @@ from typing import TYPE_CHECKING, Any, Final, Literal, cast
from urllib.parse import quote
import httpx
from pydantic import BaseModel
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.asyncify import asyncify
@ -56,6 +56,7 @@ from litellm.types.llms.vertex_ai import (
Tools,
)
from litellm.types.utils import GenericImageParsingChunk, LlmProviders
from pydantic import BaseModel
from ..common_utils import (
_check_text_in_content,
@ -179,7 +180,7 @@ def _apply_gemini_metadata(
part: PartType,
model: str | None,
media_resolution_enum: dict[str, str] | None,
video_metadata: Mapping[str, object] | None,
video_metadata: dict[str, Any] | None,
) -> PartType:
"""
Apply media_resolution and video_metadata parameters to a Gemini part.
@ -637,6 +638,23 @@ def check_if_part_exists_in_parts(parts: list[PartType], part: PartType, exclude
return False
def _get_valid_base64_thought_signature(signature: object) -> str | None:
if not isinstance(signature, str):
return None
stripped: Final = signature.strip()
if not stripped:
return None
padded: Final = stripped + "=" * (-len(stripped) % 4)
try:
decoded: Final = base64.b64decode(padded.encode("utf-8"), altchars=b"-_", validate=True)
except (binascii.Error, ValueError):
return None
return padded if decoded else None
def _collect_tool_call_thought_signatures(
assistant_msg: ChatCompletionAssistantMessage,
) -> frozenset[str]:
@ -682,8 +700,9 @@ def _collect_tool_call_thought_signatures(
continue
for key in ("thought_signature", "response_thought_signature"):
invocation_signature = invocation.get(key)
if isinstance(invocation_signature, str) and invocation_signature:
signatures += (invocation_signature,)
valid_sig = _get_valid_base64_thought_signature(invocation_signature)
if valid_sig:
signatures += (valid_sig,)
return frozenset(signatures)
@ -889,19 +908,23 @@ def _gemini_convert_messages_with_history(
if block["type"] == "thinking":
block_thinking_str = block.get("thinking")
block_signature = block.get("signature")
if block_thinking_str is not None and block_signature is not None:
valid_block_sig = _get_valid_base64_thought_signature(block_signature)
if block_thinking_str is not None:
sig_kwargs: dict[str, Any] = (
{"thoughtSignature": valid_block_sig} if valid_block_sig is not None else {}
)
try:
assistant_content.append(
PartType(
thoughtSignature=block_signature,
**sig_kwargs,
**json.loads(block_thinking_str),
)
)
except Exception:
assistant_content.append(
PartType(
thoughtSignature=block_signature,
text=block_thinking_str,
**sig_kwargs,
)
)
if _message_content is not None and isinstance(_message_content, list):
@ -927,17 +950,18 @@ def _gemini_convert_messages_with_history(
# reasoning token count on gemini-3 and newer models
tool_call_signatures = _collect_tool_call_thought_signatures(assistant_msg)
if (
thought_signatures
and isinstance(thought_signatures, list)
and len(thought_signatures) > 0
and thought_signatures[0] not in tool_call_signatures
):
valid_text_signature = (
_get_valid_base64_thought_signature(thought_signatures[0])
if (thought_signatures and isinstance(thought_signatures, list) and len(thought_signatures) > 0)
else None
)
if valid_text_signature and valid_text_signature not in tool_call_signatures:
# Use the first signature for the text part (Gemini expects one signature per part)
assistant_content.append(
PartType(
text=assistant_text,
thoughtSignature=thought_signatures[0],
thoughtSignature=valid_text_signature,
)
)
else:

View file

@ -0,0 +1,395 @@
"""
Tests for embedding thought signatures in tool call IDs for OpenAI client compatibility.
When using OpenAI clients (instead of LiteLLM SDK), provider_specific_fields are not preserved.
This test suite validates that thought signatures can be embedded in tool call IDs and extracted
when converting back to Gemini format.
Note: Embedding signatures in tool call IDs is a beta feature that requires
enable_preview_features=True to be enabled.
"""
import base64
import litellm
import pytest
from litellm.litellm_core_utils.prompt_templates.factory import (
THOUGHT_SIGNATURE_SEPARATOR,
_encode_tool_call_id_with_signature,
_get_thought_signature_from_tool,
convert_to_gemini_tool_call_invoke,
)
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
from litellm.types.llms.vertex_ai import HttpxPartType
def test_encode_decode_tool_call_id_with_signature():
"""Test that thought signatures can be encoded in and decoded from tool call IDs"""
base_id = "call_abc123"
test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5"
# Test encoding
encoded_id = _encode_tool_call_id_with_signature(base_id, test_signature)
assert THOUGHT_SIGNATURE_SEPARATOR in encoded_id
assert encoded_id.startswith(base_id)
# Test decoding using factory function with realistic tool call structure
tool = {
"id": encoded_id,
"type": "function",
"function": {
"name": "get_current_temperature",
"arguments": '{"location": "Paris"}',
},
}
extracted_signature = _get_thought_signature_from_tool(tool)
assert extracted_signature == test_signature
# Verify base ID is preserved
decoded_base_id = encoded_id.split(THOUGHT_SIGNATURE_SEPARATOR)[0]
assert decoded_base_id == base_id
def test_encode_tool_call_id_without_signature():
"""Test that IDs without signatures are returned unchanged"""
base_id = "call_abc123def456"
# Encode without signature
encoded_id = _encode_tool_call_id_with_signature(base_id, None)
assert encoded_id == base_id
assert THOUGHT_SIGNATURE_SEPARATOR not in encoded_id
# Decode ID without signature using factory function
tool_obj = {"id": base_id, "type": "function"}
decoded_signature = _get_thought_signature_from_tool(tool_obj)
assert decoded_signature is None
@pytest.mark.parametrize("enable_preview_features", [True, False])
def test_tool_call_id_includes_signature_in_response(enable_preview_features):
"""Test that tool call IDs in responses include embedded thought signatures only when preview features are enabled"""
test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5"
parts_with_signature = [
HttpxPartType(
functionCall={
"name": "get_current_temperature",
"args": {"location": "Paris"},
},
thoughtSignature=test_signature,
)
]
function, tools, _ = VertexGeminiConfig._transform_parts(
parts=parts_with_signature,
cumulative_tool_call_idx=0,
is_function_call=False,
)
# Verify tool call exists
assert tools is not None
assert len(tools) == 1
tool_call_id = tools[0]["id"]
# Verify signature is always in provider_specific_fields
assert tools[0].get("provider_specific_fields", {}).get("thought_signature") == test_signature
# When preview features enabled, signature should be embedded in ID
assert THOUGHT_SIGNATURE_SEPARATOR in tool_call_id
# Verify we can decode it using the factory function
tool_obj = {"id": tool_call_id, "type": "function"}
decoded_sig = _get_thought_signature_from_tool(tool_obj)
assert decoded_sig == test_signature
def test_get_thought_signature_backward_compatibility():
"""Test that provider_specific_fields still works (backward compatibility)"""
test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5"
# Test with provider_specific_fields (LiteLLM SDK scenario)
tool = {
"id": "call_abc123",
"type": "function",
"function": {
"name": "get_current_temperature",
"arguments": '{"location": "Paris"}',
},
"provider_specific_fields": {"thought_signature": test_signature},
}
extracted_signature = _get_thought_signature_from_tool(tool)
assert extracted_signature == test_signature
def test_get_thought_signature_prioritizes_provider_fields():
"""Test that provider_specific_fields takes priority over tool call ID"""
signature_in_fields = "signature_from_fields"
signature_in_id = "signature_from_id"
encoded_id = _encode_tool_call_id_with_signature("call_abc123", signature_in_id)
tool = {
"id": encoded_id,
"type": "function",
"function": {
"name": "get_current_temperature",
"arguments": '{"location": "Paris"}',
},
"provider_specific_fields": {"thought_signature": signature_in_fields},
}
extracted_signature = _get_thought_signature_from_tool(tool)
# Should prioritize provider_specific_fields
assert extracted_signature == signature_in_fields
def test_convert_to_gemini_with_embedded_signature():
"""Test that convert_to_gemini_tool_call_invoke extracts signatures from tool call IDs"""
test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5"
# Create tool call ID with embedded signature (as OpenAI client would send)
base_id = "call_abc123"
encoded_id = _encode_tool_call_id_with_signature(base_id, test_signature)
# Assistant message as sent by OpenAI client (no provider_specific_fields)
assistant_message = {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": encoded_id, # ID has signature embedded
"type": "function",
"function": {
"name": "get_current_temperature",
"arguments": '{"location": "Paris"}',
},
}
],
}
gemini_parts = convert_to_gemini_tool_call_invoke(assistant_message)
# Verify thought signature is extracted and sent to Gemini
assert len(gemini_parts) == 1
assert "function_call" in gemini_parts[0]
assert "thoughtSignature" in gemini_parts[0]
assert gemini_parts[0]["thoughtSignature"] == test_signature
@pytest.mark.parametrize("enable_preview_features", [True, False])
def test_openai_client_e2e_flow(enable_preview_features):
"""
End-to-end test simulating OpenAI client usage:
1. LiteLLM receives response from Gemini with thought signature
2. LiteLLM embeds signature in tool call ID (if preview features enabled)
3. OpenAI client sends message back with same tool call ID
4. LiteLLM extracts signature from ID/provider_specific_fields and sends to Gemini
"""
test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5"
# Step 1: Gemini returns function call with thought signature
gemini_parts = [
HttpxPartType(
functionCall={
"name": "get_current_temperature",
"args": {"location": "Paris"},
},
thoughtSignature=test_signature,
)
]
# Step 2: LiteLLM transforms to OpenAI format
function, tools, _ = VertexGeminiConfig._transform_parts(
parts=gemini_parts,
cumulative_tool_call_idx=0,
is_function_call=False,
)
assert tools is not None
assert len(tools) == 1
tool_call_id = tools[0]["id"]
assert THOUGHT_SIGNATURE_SEPARATOR in tool_call_id
# Step 3: OpenAI client sends back assistant message
# For the disabled case, we simulate that the client might have provider_specific_fields
# or we use the embedded ID if preview features were enabled
openai_assistant_message = {
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": tool_call_id, # Preserved from response (with embedded signature)
"type": "function",
"function": {
"name": "get_current_temperature",
"arguments": '{"location": "Paris"}',
},
}
],
}
# Step 4: LiteLLM converts back to Gemini format, extracting signature
gemini_parts_converted = convert_to_gemini_tool_call_invoke(openai_assistant_message)
# Verify signature is preserved through the round trip
assert len(gemini_parts_converted) == 1
assert "thoughtSignature" in gemini_parts_converted[0]
assert gemini_parts_converted[0]["thoughtSignature"] == test_signature
@pytest.mark.parametrize("enable_preview_features", [True, False])
def test_parallel_tool_calls_with_signatures(enable_preview_features):
"""Test that parallel tool calls preserve signatures correctly"""
signature1 = "signature_for_first_call"
# Only first call has signature (Gemini behavior for parallel calls)
gemini_parts = [
HttpxPartType(
functionCall={"name": "get_temperature", "args": {"location": "Paris"}},
thoughtSignature=signature1,
),
HttpxPartType(
functionCall={"name": "get_temperature", "args": {"location": "London"}},
# No signature for second parallel call
),
]
function, tools, _ = VertexGeminiConfig._transform_parts(
parts=gemini_parts,
cumulative_tool_call_idx=0,
is_function_call=False,
)
assert tools is not None
assert len(tools) == 2
# First tool call should have signature in provider_specific_fields
assert tools[0].get("provider_specific_fields", {}).get("thought_signature") == signature1
# When preview features enabled, first tool call has signature in ID
assert THOUGHT_SIGNATURE_SEPARATOR in tools[0]["id"]
sig1 = _get_thought_signature_from_tool({"id": tools[0]["id"], "type": "function"})
assert sig1 == signature1
# Second tool call has no signature in ID (regardless of flag)
assert THOUGHT_SIGNATURE_SEPARATOR not in tools[1]["id"]
sig2 = _get_thought_signature_from_tool({"id": tools[1]["id"], "type": "function"})
assert sig2 is None
def test_get_valid_base64_thought_signature_helper():
from litellm.llms.vertex_ai.gemini.transformation import (
_get_valid_base64_thought_signature,
)
valid_sig = base64.b64encode(b"valid_thought_signature_content").decode("utf-8")
assert _get_valid_base64_thought_signature(valid_sig) == valid_sig
assert _get_valid_base64_thought_signature(f" {valid_sig} \n") == valid_sig
urlsafe_sig = base64.urlsafe_b64encode(b"valid_thought_signature_content").decode("utf-8")
assert _get_valid_base64_thought_signature(urlsafe_sig) == urlsafe_sig
unpadded_sig = valid_sig.rstrip("=")
assert _get_valid_base64_thought_signature(unpadded_sig) == valid_sig
assert _get_valid_base64_thought_signature(None) is None
assert _get_valid_base64_thought_signature("") is None
assert _get_valid_base64_thought_signature(" ") is None
assert _get_valid_base64_thought_signature("not_base64_content!@#") is None
assert _get_valid_base64_thought_signature("abc") == "abc="
assert _get_valid_base64_thought_signature("abcde") is None
assert _get_valid_base64_thought_signature("====") is None
assert _get_valid_base64_thought_signature(12345) is None
def test_gemini_transformation_omits_malformed_thought_signature_in_replayed_history():
from litellm.llms.vertex_ai.gemini.transformation import (
_gemini_convert_messages_with_history,
)
def _get_part_val(part: object, key: str) -> object:
if isinstance(part, dict):
return part.get(key)
return getattr(part, key, None)
valid_sig = base64.b64encode(b"gemini_replayed_thought_sig").decode("utf-8")
malformed_sig = "malformed_base64_thought_sig_!@#"
# 1. Malformed signature in provider_specific_fields["thought_signatures"]
messages_malformed = [
{"role": "user", "content": "Analyze this data."},
{
"role": "assistant",
"content": "Analysis complete.",
"provider_specific_fields": {"thought_signatures": [malformed_sig]},
},
]
contents = _gemini_convert_messages_with_history(messages=messages_malformed)
model_parts = [part for content in contents if content["role"] == "model" for part in content["parts"]]
assert len(model_parts) >= 1
assert _get_part_val(model_parts[0], "text") == "Analysis complete."
assert _get_part_val(model_parts[0], "thoughtSignature") is None
# 2. Valid signature in provider_specific_fields["thought_signatures"]
messages_valid = [
{"role": "user", "content": "Analyze this data."},
{
"role": "assistant",
"content": "Analysis complete.",
"provider_specific_fields": {"thought_signatures": [valid_sig]},
},
]
contents_valid = _gemini_convert_messages_with_history(messages=messages_valid)
model_parts_valid = [part for content in contents_valid if content["role"] == "model" for part in content["parts"]]
assert len(model_parts_valid) >= 1
assert _get_part_val(model_parts_valid[0], "thoughtSignature") == valid_sig
# 3. Malformed signature in thinking_blocks
messages_thinking_malformed = [
{"role": "user", "content": "Plan next steps."},
{
"role": "assistant",
"content": "Here is the plan.",
"thinking_blocks": [
{
"type": "thinking",
"thinking": "Step 1: Check constraints.",
"signature": malformed_sig,
}
],
},
]
contents_thinking = _gemini_convert_messages_with_history(messages=messages_thinking_malformed)
thinking_parts = [part for content in contents_thinking if content["role"] == "model" for part in content["parts"]]
assert len(thinking_parts) >= 1
assert _get_part_val(thinking_parts[0], "thoughtSignature") is None
# 4. Valid signature in thinking_blocks
messages_thinking_valid = [
{"role": "user", "content": "Plan next steps."},
{
"role": "assistant",
"content": "Here is the plan.",
"thinking_blocks": [
{
"type": "thinking",
"thinking": "Step 1: Check constraints.",
"signature": valid_sig,
}
],
},
]
contents_thinking_valid = _gemini_convert_messages_with_history(messages=messages_thinking_valid)
thinking_parts_valid = [
part for content in contents_thinking_valid if content["role"] == "model" for part in content["parts"]
]
assert len(thinking_parts_valid) >= 1
assert _get_part_val(thinking_parts_valid[0], "thoughtSignature") == valid_sig