mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
test(vertex_ai): add thought signature Base64 validation tests (#42201)
This commit is contained in:
parent
de7073812d
commit
7e8baac8f6
1 changed files with 395 additions and 0 deletions
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue