mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 7e8baac8f6 into 02f61c9c42
This commit is contained in:
commit
454430a331
2 changed files with 434 additions and 15 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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