deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/models/patched_openai.py
2026-09-07 18:24:55 +08:00

228 lines
10 KiB
Python

"""Patched ChatOpenAI that preserves thought_signature for Gemini thinking models.
When using Gemini with thinking enabled via an OpenAI-compatible gateway (e.g.
Vertex AI, Google AI Studio, or any proxy), the API requires that the
``thought_signature`` field on tool-call objects is echoed back verbatim in
every subsequent request.
The OpenAI-compatible gateway stores the raw tool-call dicts (including
``thought_signature``) in ``additional_kwargs["tool_calls"]``, but standard
``langchain_openai.ChatOpenAI`` only serialises the standard fields (``id``,
``type``, ``function``) into the outgoing payload, silently dropping the
signature. That causes an HTTP 400 ``INVALID_ARGUMENT`` error:
Unable to submit request because function call `<tool>` in the N. content
block is missing a `thought_signature`.
This module fixes the problem by overriding ``_get_request_payload`` to
re-inject thinking metadata back into the outgoing payload for any assistant
message that originally carried it:
- ``thought_signature`` for Gemini-compatible tool calls.
- ``reasoning_content`` for providers (for example DeepSeek) that require it
to be echoed in multi-turn requests when thinking mode is enabled.
"""
from __future__ import annotations
from typing import Any
from langchain_core.language_models import LanguageModelInput
from langchain_core.messages import AIMessage, AIMessageChunk
from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult
from langchain_openai import ChatOpenAI
from deerflow.models.message_ordering import normalize_chat_payload_messages
class PatchedChatOpenAI(ChatOpenAI):
"""ChatOpenAI with ``thought_signature`` preservation for Gemini thinking via OpenAI gateway.
When using Gemini with thinking enabled via an OpenAI-compatible gateway,
the API expects ``thought_signature`` to be present on tool-call objects in
multi-turn conversations. This patched version restores those signatures
from ``AIMessage.additional_kwargs["tool_calls"]`` into the serialised
request payload before it is sent to the API.
Usage in ``config.yaml``::
- name: gemini-2.5-pro-thinking
display_name: Gemini 2.5 Pro (Thinking)
use: deerflow.models.patched_openai:PatchedChatOpenAI
model: google/gemini-2.5-pro-preview
api_key: $GEMINI_API_KEY
base_url: https://<your-openai-compat-gateway>/v1
max_tokens: 16384
supports_thinking: true
supports_vision: true
when_thinking_enabled:
extra_body:
thinking:
type: enabled
"""
def _get_request_payload(
self,
input_: LanguageModelInput,
*,
stop: list[str] | None = None,
**kwargs: Any,
) -> dict:
"""Get request payload with ``thought_signature`` preserved on tool-call objects.
Overrides the parent method to re-inject ``thought_signature`` fields
on tool-call objects that were stored in
``additional_kwargs["tool_calls"]`` by LangChain but dropped during
serialisation.
"""
# Capture the original LangChain messages *before* conversion so we can
# access fields that the serialiser might drop.
original_messages = self._convert_input(input_).to_messages()
# Obtain the base payload from the parent implementation.
payload = super()._get_request_payload(input_, stop=stop, **kwargs)
payload_messages = payload.get("messages", [])
if len(payload_messages) == len(original_messages):
for payload_msg, orig_msg in zip(payload_messages, original_messages):
if payload_msg.get("role") == "assistant" and isinstance(orig_msg, AIMessage):
_restore_tool_call_signatures(payload_msg, orig_msg)
_restore_reasoning_content(payload_msg, orig_msg)
else:
# Fallback: match assistant-role entries positionally against AIMessages.
ai_messages = [m for m in original_messages if isinstance(m, AIMessage)]
assistant_payloads = [(i, m) for i, m in enumerate(payload_messages) if m.get("role") == "assistant"]
for (_, payload_msg), ai_msg in zip(assistant_payloads, ai_messages):
_restore_tool_call_signatures(payload_msg, ai_msg)
_restore_reasoning_content(payload_msg, ai_msg)
normalize_chat_payload_messages(payload)
return payload
def _convert_chunk_to_generation_chunk(
self,
chunk: dict,
default_chunk_class: type,
base_generation_info: dict | None,
) -> ChatGenerationChunk | None:
"""Preserve streamed ``reasoning_content`` in ``additional_kwargs``."""
generation_chunk = super()._convert_chunk_to_generation_chunk(
chunk,
default_chunk_class,
base_generation_info,
)
if generation_chunk is None:
return None
choices = chunk.get("choices", []) or chunk.get("chunk", {}).get("choices", [])
choice = choices[0] if choices else {}
delta = choice.get("delta") if isinstance(choice, dict) else None
reasoning_content = _extract_reasoning_content(delta if isinstance(delta, dict) else None)
if reasoning_content and isinstance(generation_chunk.message, AIMessageChunk):
generation_chunk.message = _append_reasoning_content(
generation_chunk.message,
reasoning_content,
)
return generation_chunk
def _create_chat_result(
self,
response: dict | Any,
generation_info: dict | None = None,
) -> ChatResult:
"""Preserve non-stream ``reasoning_content`` in ``additional_kwargs``."""
result = super()._create_chat_result(response, generation_info)
response_dict = response if isinstance(response, dict) else response.model_dump()
choices = response_dict.get("choices", [])
generations: list[ChatGeneration] = []
for index, generation in enumerate(result.generations):
message = generation.message
choice = choices[index] if index < len(choices) else {}
choice_message = choice.get("message", {}) if isinstance(choice, dict) else {}
reasoning_content = _extract_reasoning_content(choice_message if isinstance(choice_message, dict) else None)
if reasoning_content and isinstance(message, AIMessage):
message = _append_reasoning_content(message, reasoning_content)
generation = ChatGeneration(
message=message,
generation_info=generation.generation_info,
)
generations.append(generation)
return ChatResult(generations=generations, llm_output=result.llm_output)
def _restore_tool_call_signatures(payload_msg: dict, orig_msg: AIMessage) -> None:
"""Re-inject ``thought_signature`` onto tool-call objects in *payload_msg*.
When the Gemini OpenAI-compatible gateway returns a response with function
calls, each tool-call object may carry a ``thought_signature``. LangChain
stores the raw tool-call dicts in ``additional_kwargs["tool_calls"]`` but
only serialises the standard fields (``id``, ``type``, ``function``) into
the outgoing payload, silently dropping the signature.
This function matches raw tool-call entries (by ``id``, falling back to
positional order) and copies the signature back onto the serialised
payload entries.
"""
raw_tool_calls: list[dict] = orig_msg.additional_kwargs.get("tool_calls") or []
payload_tool_calls: list[dict] = payload_msg.get("tool_calls") or []
if not raw_tool_calls or not payload_tool_calls:
return
# Build an id → raw_tc lookup for efficient matching.
raw_by_id: dict[str, dict] = {}
for raw_tc in raw_tool_calls:
tc_id = raw_tc.get("id")
if tc_id:
raw_by_id[tc_id] = raw_tc
for idx, payload_tc in enumerate(payload_tool_calls):
# Try matching by id first, then fall back to positional.
raw_tc = raw_by_id.get(payload_tc.get("id", ""))
if raw_tc is None and idx < len(raw_tool_calls):
raw_tc = raw_tool_calls[idx]
if raw_tc is None:
continue
# The gateway may use either snake_case or camelCase.
sig = raw_tc.get("thought_signature") or raw_tc.get("thoughtSignature")
if sig:
payload_tc["thought_signature"] = sig
def _restore_reasoning_content(payload_msg: dict, orig_msg: AIMessage) -> None:
"""Re-inject ``reasoning_content`` onto assistant payload message.
Some thinking-enabled OpenAI-compatible providers require each historical
assistant turn to include the original ``reasoning_content`` in subsequent
requests. LangChain stores it in ``AIMessage.additional_kwargs`` but does
not always serialize it back into ``messages``.
"""
reasoning_content = orig_msg.additional_kwargs.get("reasoning_content")
if reasoning_content is not None:
payload_msg["reasoning_content"] = reasoning_content
def _extract_reasoning_content(payload: dict | None) -> str | None:
"""Extract reasoning text from provider payload delta/message objects."""
if not isinstance(payload, dict):
return None
reasoning_content = payload.get("reasoning_content")
return reasoning_content if isinstance(reasoning_content, str) and reasoning_content else None
def _append_reasoning_content(message: AIMessage | AIMessageChunk, reasoning_content: str) -> AIMessage | AIMessageChunk:
"""Append streamed reasoning chunks, dedupe complete response reasoning."""
additional_kwargs = dict(message.additional_kwargs)
existing = additional_kwargs.get("reasoning_content")
if isinstance(existing, str) and existing:
if reasoning_content not in existing:
additional_kwargs["reasoning_content"] = f"{existing}{reasoning_content}"
else:
additional_kwargs["reasoning_content"] = reasoning_content
return message.model_copy(update={"additional_kwargs": additional_kwargs})