83 lines
3.1 KiB
Python
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
|