diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index ba59e3fa997..712092f4759 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1,4 +1,5 @@ import base64 +import binascii import copy import hashlib import json @@ -65,6 +66,7 @@ BAD_MESSAGE_ERROR_STR: Final = "Invalid Message " # Separator used to embed Gemini thought signatures in tool call IDs # See: https://ai.google.dev/gemini-api/docs/thought-signatures THOUGHT_SIGNATURE_SEPARATOR: Final = "__thought__" +THOUGHT_SIGNATURE_CHECKSUM_SEPARATOR: Final = "__checksum__" # used to interweave user messages, to ensure user/assistant alternating DEFAULT_USER_CONTINUE_MESSAGE: Final = { @@ -1133,7 +1135,7 @@ def _gemini_tool_call_invoke_helper( args=arguments_dict, ) if tool_call_id: - clean_id: Final = tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0] + clean_id, _ = _decode_tool_call_id_with_signature(tool_call_id) if clean_id: function_call["id"] = clean_id return function_call @@ -1149,15 +1151,48 @@ def _encode_tool_call_id_with_signature(tool_call_id: str, thought_signature: st Returns: Tool call ID with embedded signature if present, otherwise original ID - Format: call___thought__ + Format: call___checksum____thought__ See: https://ai.google.dev/gemini-api/docs/thought-signatures """ if thought_signature: - return f"{tool_call_id}{THOUGHT_SIGNATURE_SEPARATOR}{thought_signature}" + checksum: Final = hashlib.sha256(thought_signature.encode("utf-8")).hexdigest() + return ( + f"{tool_call_id}{THOUGHT_SIGNATURE_CHECKSUM_SEPARATOR}{checksum}" + f"{THOUGHT_SIGNATURE_SEPARATOR}{thought_signature}" + ) return tool_call_id +def _is_valid_thought_signature(signature: str) -> bool: + """Check strict standard-alphabet base64 while allowing omitted padding.""" + if not signature: + return False + try: + padding = "=" * (-len(signature) % 4) + base64.b64decode(signature + padding, validate=True) + except (binascii.Error, ValueError): + return False + return True + + +def _decode_tool_call_id_with_signature(tool_call_id: str) -> tuple[str, str | None]: + payload, signature_separator, signature = tool_call_id.partition(THOUGHT_SIGNATURE_SEPARATOR) + if not signature_separator: + return tool_call_id, None + + base_id, checksum_separator, checksum = payload.rpartition(THOUGHT_SIGNATURE_CHECKSUM_SEPARATOR) + if not checksum_separator: + return payload, None + if not _is_valid_thought_signature(signature): + return base_id, None + + expected_checksum: Final = hashlib.sha256(signature.encode("utf-8")).hexdigest() + if checksum != expected_checksum: + return base_id, None + return base_id, signature + + def _get_thought_signature_from_tool(tool: dict) -> str | None: """Extract thought signature from tool call's provider_specific_fields. @@ -1190,11 +1225,9 @@ def _get_thought_signature_from_tool(tool: dict) -> str | None: return signature # Check if thought signature is embedded in tool call ID tool_call_id: Final = tool.get("id") - if tool_call_id and THOUGHT_SIGNATURE_SEPARATOR in tool_call_id: - parts: Final = tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1) - if len(parts) == 2: - _, signature = parts - return signature + if tool_call_id and isinstance(tool_call_id, str): + _, signature = _decode_tool_call_id_with_signature(tool_call_id) + return signature return None @@ -1453,7 +1486,7 @@ def convert_to_gemini_tool_call_result( if forward_function_call_id: raw_tool_call_id: Final = message.get("tool_call_id") if raw_tool_call_id and isinstance(raw_tool_call_id, str): - stripped_id: Final = raw_tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0] + stripped_id, _ = _decode_tool_call_id_with_signature(raw_tool_call_id) if stripped_id: gemini_call_id = stripped_id diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py b/tests/test_litellm/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py index 208cba519f3..964f5e0e673 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py @@ -9,13 +9,19 @@ Note: Embedding signatures in tool call IDs is a beta feature that requires enable_preview_features=True to be enabled. """ +import base64 + import pytest import litellm from litellm.litellm_core_utils.prompt_templates.factory import ( + THOUGHT_SIGNATURE_CHECKSUM_SEPARATOR, THOUGHT_SIGNATURE_SEPARATOR, + _decode_tool_call_id_with_signature, _encode_tool_call_id_with_signature, + _get_dummy_thought_signature, _get_thought_signature_from_tool, + _is_valid_thought_signature, convert_to_gemini_tool_call_invoke, ) from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( @@ -32,6 +38,7 @@ def test_encode_decode_tool_call_id_with_signature(): # Test encoding encoded_id = _encode_tool_call_id_with_signature(base_id, test_signature) assert THOUGHT_SIGNATURE_SEPARATOR in encoded_id + assert THOUGHT_SIGNATURE_CHECKSUM_SEPARATOR in encoded_id assert encoded_id.startswith(base_id) # Test decoding using factory function with realistic tool call structure @@ -48,7 +55,7 @@ def test_encode_decode_tool_call_id_with_signature(): assert extracted_signature == test_signature # Verify base ID is preserved - decoded_base_id = encoded_id.split(THOUGHT_SIGNATURE_SEPARATOR)[0] + decoded_base_id, _ = _decode_tool_call_id_with_signature(encoded_id) assert decoded_base_id == base_id @@ -94,10 +101,7 @@ def test_tool_call_id_includes_signature_in_response(enable_preview_features): 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 - ) + 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 @@ -234,9 +238,7 @@ def test_openai_client_e2e_flow(enable_preview_features): ], } # Step 4: LiteLLM converts back to Gemini format, extracting signature - gemini_parts_converted = convert_to_gemini_tool_call_invoke( - openai_assistant_message - ) + 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 @@ -247,7 +249,7 @@ def test_openai_client_e2e_flow(enable_preview_features): @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" + signature1 = base64.b64encode(b"signature_for_first_call").decode("ascii") # Only first call has signature (Gemini behavior for parallel calls) gemini_parts = [ @@ -271,10 +273,7 @@ def test_parallel_tool_calls_with_signatures(enable_preview_features): 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 - ) + 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"] @@ -285,3 +284,131 @@ def test_parallel_tool_calls_with_signatures(enable_preview_features): 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 + + +REAL_SIGNATURE = ( + "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdT" + "tfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed" + "0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7C" + "vykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" +) + + +@pytest.mark.parametrize( + "mangled_signature", + [ + "AY89a1/_57b05e78dc", + "AY89a1/_ee781c9832", + "AY89a1_S3YvIpUCcBTFSgDfesRLDnA_775ff49bcd", + ], +) +def test_is_valid_thought_signature_rejects_client_normalized_values(mangled_signature): + assert _is_valid_thought_signature(REAL_SIGNATURE) is True + assert _is_valid_thought_signature(mangled_signature) is False + + +def test_is_valid_thought_signature_tolerates_missing_padding(): + encoded = base64.b64encode(b"hello").decode("ascii").rstrip("=") + assert "=" not in encoded + assert _is_valid_thought_signature(encoded) is True + + +def test_get_thought_signature_rejects_decodable_tampering(): + encoded_id = _encode_tool_call_id_with_signature("call_abc123", REAL_SIGNATURE) + checksum = encoded_id.split(THOUGHT_SIGNATURE_CHECKSUM_SEPARATOR, 1)[1].split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0] + decodable_but_wrong = base64.b64encode(b"not-the-real-signature").decode("ascii") + tampered_id = ( + f"call_abc123{THOUGHT_SIGNATURE_CHECKSUM_SEPARATOR}{checksum}{THOUGHT_SIGNATURE_SEPARATOR}{decodable_but_wrong}" + ) + + assert _is_valid_thought_signature(decodable_but_wrong) is True + assert _get_thought_signature_from_tool({"id": tampered_id, "type": "function"}) is None + + +def test_get_thought_signature_rejects_unsigned_embedded_signature(): + unsigned_id = f"call_abc123{THOUGHT_SIGNATURE_SEPARATOR}{REAL_SIGNATURE}" + + assert _get_thought_signature_from_tool({"id": unsigned_id, "type": "function"}) is None + + +def test_get_thought_signature_drops_client_mangled_id_suffix(): + """When the segment after ``__thought__`` isn't valid base64, the extractor + must return ``None`` so the caller can either fall back to the dummy + skip-validator signature (Gemini 3+) or drop the signature entirely + (Gemini 2.x), instead of forwarding a corrupted value that Vertex would + reject. Regression test for issue #37849.""" + mangled_id = _encode_tool_call_id_with_signature("call_2156408", "AY89a1/_57b05e78dc") + tool = {"id": mangled_id, "type": "function"} + + assert _get_thought_signature_from_tool(tool) is None + + +def test_get_thought_signature_still_prefers_provider_fields_even_when_id_mangled(): + """A valid signature in ``provider_specific_fields`` must win over the + mangled tail so we don't downgrade a good signal.""" + mangled_id = f"call_abc{THOUGHT_SIGNATURE_SEPARATOR}not_base64!!" + tool = { + "id": mangled_id, + "type": "function", + "function": {"name": "get_temperature", "arguments": '{"location": "Paris"}'}, + "provider_specific_fields": {"thought_signature": REAL_SIGNATURE}, + } + + assert _get_thought_signature_from_tool(tool) == REAL_SIGNATURE + + +def test_convert_to_gemini_uses_dummy_signature_when_client_mangles_id_on_gemini_3(): + """On Gemini 3+ a mangled id must degrade to the documented skip-validator + dummy signature rather than forwarding the corrupted bytes.""" + mangled_id = f"call_2156408{THOUGHT_SIGNATURE_SEPARATOR}AY89a1/_57b05e78dc" + assistant_message = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": mangled_id, + "type": "function", + "function": { + "name": "get_temperature", + "arguments": '{"location": "Paris"}', + }, + } + ], + } + + parts = convert_to_gemini_tool_call_invoke( + assistant_message, + model="vertex_ai/gemini-3.1-pro-preview", + ) + + assert len(parts) == 1 + assert parts[0].get("thoughtSignature") == _get_dummy_thought_signature() + + +def test_convert_to_gemini_drops_signature_when_client_mangles_id_on_gemini_2(): + """On older Gemini models the placeholder fallback doesn't apply, so a + mangled id must simply drop the signature (rather than forward a value + Vertex will 400 on).""" + mangled_id = f"call_2158562{THOUGHT_SIGNATURE_SEPARATOR}AY89a1/_ee781c9832" + assistant_message = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": mangled_id, + "type": "function", + "function": { + "name": "get_temperature", + "arguments": '{"location": "Paris"}', + }, + } + ], + } + + parts = convert_to_gemini_tool_call_invoke( + assistant_message, + model="vertex_ai/gemini-2.5-pro", + ) + + assert len(parts) == 1 + assert "thoughtSignature" not in parts[0]