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

83 lines
3.1 KiB
Python

"""Utilities for provider-specific chat message ordering constraints."""
from __future__ import annotations
from typing import Any
from langchain_core.messages import BaseMessage, SystemMessage
_SYSTEM_SEPARATOR = "\n\n---\n\n"
def normalize_chat_payload_messages(payload: dict[str, Any]) -> None:
"""Ensure payload ``system`` messages are represented once at the beginning.
Several OpenAI-compatible gateways reject requests when a transient
``system`` message appears after user/assistant/tool history. LangChain and
DeerFlow middleware can legitimately create those transient messages, so the
provider boundary folds every system payload into a single leading message
while preserving all non-system messages in their original order.
"""
messages = payload.get("messages")
if not isinstance(messages, list):
return
normalized = _normalize_payload_message_list(messages)
if normalized is not messages:
payload["messages"] = normalized
def normalize_langchain_messages(messages: list[BaseMessage]) -> list[BaseMessage]:
"""Return LangChain messages with all system context merged at the front."""
system_messages = [message for message in messages if isinstance(message, SystemMessage)]
if not system_messages:
return messages
first_non_system_index = next((idx for idx, message in enumerate(messages) if not isinstance(message, SystemMessage)), len(messages))
if len(system_messages) == first_non_system_index:
return messages
merged = system_messages[0].model_copy(
update={
"content": _merge_contents([message.content for message in system_messages]),
}
)
return [merged, *(message for message in messages if not isinstance(message, SystemMessage))]
def _normalize_payload_message_list(messages: list[Any]) -> list[Any]:
system_messages = [message for message in messages if isinstance(message, dict) and message.get("role") == "system"]
if not system_messages:
return messages
first_non_system_index = next(
(
idx
for idx, message in enumerate(messages)
if not (isinstance(message, dict) and message.get("role") == "system")
),
len(messages),
)
if len(system_messages) == first_non_system_index:
return messages
merged_system = dict(system_messages[0])
merged_system["content"] = _merge_contents([message.get("content") for message in system_messages])
return [merged_system, *(message for message in messages if not (isinstance(message, dict) and message.get("role") == "system"))]
def _merge_contents(contents: list[Any]) -> Any:
if all(isinstance(content, str) for content in contents):
return _SYSTEM_SEPARATOR.join(content for content in contents if content)
blocks: list[Any] = []
for idx, content in enumerate(contents):
if idx > 0:
blocks.append({"type": "text", "text": _SYSTEM_SEPARATOR})
if isinstance(content, list):
blocks.extend(content)
elif content is not None:
blocks.append({"type": "text", "text": str(content)})
return blocks