Better handle Gemini non-function thought signatures
This commit is contained in:
@@ -209,7 +209,7 @@ class GeminiLLMAdapter(BaseLLMAdapter[GeminiLLMInvocationParams]):
|
|||||||
system_instruction = None
|
system_instruction = None
|
||||||
messages = []
|
messages = []
|
||||||
tool_call_id_to_name_mapping = {}
|
tool_call_id_to_name_mapping = {}
|
||||||
non_fn_thought_signatures = []
|
non_fn_signed_parts = []
|
||||||
|
|
||||||
# Process each message, converting to Google format as needed
|
# Process each message, converting to Google format as needed
|
||||||
for message in universal_context_messages:
|
for message in universal_context_messages:
|
||||||
@@ -237,9 +237,9 @@ class GeminiLLMAdapter(BaseLLMAdapter[GeminiLLMInvocationParams]):
|
|||||||
if (
|
if (
|
||||||
isinstance(message.message, dict)
|
isinstance(message.message, dict)
|
||||||
and message.message.get("type") == "non_fn_thought_signature"
|
and message.message.get("type") == "non_fn_thought_signature"
|
||||||
and (thought_signature := message.message.get("signature"))
|
and (signed_part := message.message.get("signed_part"))
|
||||||
):
|
):
|
||||||
non_fn_thought_signatures.append(thought_signature)
|
non_fn_signed_parts.append(signed_part)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Fall back to assuming that the message is already in Google
|
# Fall back to assuming that the message is already in Google
|
||||||
@@ -269,7 +269,7 @@ class GeminiLLMAdapter(BaseLLMAdapter[GeminiLLMInvocationParams]):
|
|||||||
|
|
||||||
# Apply non-function-call-related thought signatures to the appropriate
|
# Apply non-function-call-related thought signatures to the appropriate
|
||||||
# messages
|
# messages
|
||||||
self._apply_non_function_thought_signatures_to_messages(non_fn_thought_signatures, messages)
|
self._apply_non_function_thought_signatures_to_messages(non_fn_signed_parts, messages)
|
||||||
|
|
||||||
# Check if we only have function-related messages (no regular text)
|
# Check if we only have function-related messages (no regular text)
|
||||||
has_regular_messages = any(
|
has_regular_messages = any(
|
||||||
@@ -476,19 +476,19 @@ class GeminiLLMAdapter(BaseLLMAdapter[GeminiLLMInvocationParams]):
|
|||||||
break
|
break
|
||||||
|
|
||||||
def _apply_non_function_thought_signatures_to_messages(
|
def _apply_non_function_thought_signatures_to_messages(
|
||||||
self, thought_signatures: List[bytes], messages: List[Content]
|
self, signed_parts: List[Part], messages: List[Content]
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Apply non-function-call-related thought signatures to the last part of corresponding non-function-call assistant messages.
|
"""Apply (optional, but recommended) non-function-call-related thought signatures to the last part of corresponding non-function-call assistant messages.
|
||||||
|
|
||||||
Gemini 3 Pro (and, somewhat surprisingly, other models, too, when
|
Gemini 3 Pro (and, somewhat surprisingly, other models, too, when
|
||||||
functions are involved in the conversation) outputs a thought signature
|
functions are involved in the conversation) outputs a thought signature
|
||||||
at the end of assistant responses.
|
at the end of assistant responses.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
thought_signatures: The list of thought signature bytes to apply.
|
signed_parts: A list of signed received Parts containing thought signatures to apply.
|
||||||
messages: List of messages to search through.
|
messages: List of messages to search through.
|
||||||
"""
|
"""
|
||||||
if not thought_signatures:
|
if not signed_parts:
|
||||||
return
|
return
|
||||||
|
|
||||||
# Find all assistant (model) messages that aren't function calls
|
# Find all assistant (model) messages that aren't function calls
|
||||||
@@ -504,17 +504,51 @@ class GeminiLLMAdapter(BaseLLMAdapter[GeminiLLMInvocationParams]):
|
|||||||
if not has_function_call:
|
if not has_function_call:
|
||||||
non_fn_assistant_messages.append(message)
|
non_fn_assistant_messages.append(message)
|
||||||
|
|
||||||
# Warn if counts don't match
|
|
||||||
if len(thought_signatures) != len(non_fn_assistant_messages):
|
|
||||||
logger.warning(
|
|
||||||
f"Thought signature count ({len(thought_signatures)}) doesn't match "
|
|
||||||
f"non-function-call assistant message count ({len(non_fn_assistant_messages)})"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Apply thought signatures to the corresponding assistant messages
|
# Apply thought signatures to the corresponding assistant messages
|
||||||
# Match them in order (oldest to newest)
|
# Match them using content heuristics, maintaining order (messages without signatures are skipped)
|
||||||
for i, thought_signature in enumerate(thought_signatures):
|
message_start_index = 0 # Track where to start searching for the next match
|
||||||
if i < len(non_fn_assistant_messages):
|
for signed_part in signed_parts:
|
||||||
|
thought_signature = getattr(signed_part, "thought_signature", None)
|
||||||
|
if not thought_signature:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Search through remaining non-function assistant messages for a match
|
||||||
|
for i in range(message_start_index, len(non_fn_assistant_messages)):
|
||||||
message = non_fn_assistant_messages[i]
|
message = non_fn_assistant_messages[i]
|
||||||
if message.parts:
|
if not message.parts:
|
||||||
message.parts[-1].thought_signature = thought_signature
|
continue
|
||||||
|
|
||||||
|
last_part = message.parts[-1]
|
||||||
|
matched = False
|
||||||
|
|
||||||
|
# Check if signed part has text and last message part text has the same text or
|
||||||
|
# - is a prefix of that text (in case spoken text was truncated due to interruption)
|
||||||
|
# - is prefixed by that text (in case signed part was not the end of the assistant response...
|
||||||
|
# which is NOT supposed to happen, according to Google's docs, but seems to, for long responses...)
|
||||||
|
if hasattr(signed_part, "text") and signed_part.text:
|
||||||
|
if hasattr(last_part, "text") and last_part.text:
|
||||||
|
# Normalize whitespace for comparison
|
||||||
|
signed_text = " ".join(signed_part.text.split())
|
||||||
|
last_text = " ".join(last_part.text.split())
|
||||||
|
if (
|
||||||
|
last_text == signed_text
|
||||||
|
or signed_text.startswith(last_text)
|
||||||
|
or last_text.startswith(signed_text)
|
||||||
|
):
|
||||||
|
last_part.thought_signature = thought_signature
|
||||||
|
matched = True
|
||||||
|
|
||||||
|
# Check if signed part has inline_data and last message part has matching inline_data
|
||||||
|
elif hasattr(signed_part, "inline_data") and signed_part.inline_data:
|
||||||
|
if (
|
||||||
|
hasattr(last_part, "inline_data")
|
||||||
|
and last_part.inline_data
|
||||||
|
and last_part.inline_data.data == signed_part.inline_data.data
|
||||||
|
):
|
||||||
|
last_part.thought_signature = thought_signature
|
||||||
|
matched = True
|
||||||
|
|
||||||
|
# If we found a match, update start index and stop searching for this signed part
|
||||||
|
if matched:
|
||||||
|
message_start_index = i + 1
|
||||||
|
break
|
||||||
|
|||||||
@@ -1018,7 +1018,7 @@ class GoogleLLMService(LLMService):
|
|||||||
self.get_llm_adapter().create_llm_specific_message(
|
self.get_llm_adapter().create_llm_specific_message(
|
||||||
{
|
{
|
||||||
"type": "non_fn_thought_signature",
|
"type": "non_fn_thought_signature",
|
||||||
"signature": part.thought_signature,
|
"signed_part": part,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
|
|||||||
Reference in New Issue
Block a user