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

85 lines
3.6 KiB
Python

import pytest
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage
from deerflow.agents.middlewares.writing_setup_round_cap_middleware import WritingSetupRoundCapMiddleware
from deerflow.models.message_ordering import normalize_chat_payload_messages, normalize_langchain_messages
from deerflow.models.patched_openai import PatchedChatOpenAI
def _multi_turn_payload_with_late_systems(turns: int):
messages = [{"role": "system", "content": "base system"}]
for idx in range(turns):
messages.append({"role": "user", "content": f"user {idx}"})
messages.append({"role": "assistant", "content": f"assistant {idx}"})
if idx >= 9 and (idx + 1) % 10 == 0:
messages.append({"role": "system", "content": f"late summary after {idx + 1} turns"})
messages.append({"role": "system", "content": "tail forced instruction"})
return {"messages": messages}
@pytest.mark.parametrize("turns", [12, 25, 50])
def test_normalize_chat_payload_messages_merges_late_systems_after_many_turns(turns):
payload = _multi_turn_payload_with_late_systems(turns)
normalize_chat_payload_messages(payload)
roles = [message["role"] for message in payload["messages"]]
assert roles[0] == "system"
assert roles.count("system") == 1
assert roles[1:5] == ["user", "assistant", "user", "assistant"]
assert "base system" in payload["messages"][0]["content"]
assert "late summary after 10 turns" in payload["messages"][0]["content"]
assert "tail forced instruction" in payload["messages"][0]["content"]
def test_normalize_langchain_messages_moves_late_systems_before_history():
messages = [SystemMessage(content="base system")]
for idx in range(12):
messages.append(HumanMessage(content=f"user {idx}"))
messages.append(AIMessage(content=f"assistant {idx}"))
if idx == 10:
messages.append(SystemMessage(content="late compaction summary"))
normalized = normalize_langchain_messages(messages)
assert isinstance(normalized[0], SystemMessage)
assert sum(isinstance(message, SystemMessage) for message in normalized) == 1
assert normalized[0].content == "base system\n\n---\n\nlate compaction summary"
assert isinstance(normalized[1], HumanMessage)
assert normalized[1].content == "user 0"
@pytest.mark.parametrize("turns", [12, 25, 50])
def test_patched_openai_payload_has_only_leading_system_after_many_turns(turns):
model = PatchedChatOpenAI(
model="test-model",
api_key="test-key",
base_url="http://127.0.0.1:9/v1",
)
messages = [SystemMessage(content="base system")]
for idx in range(turns):
messages.append(HumanMessage(content=f"user {idx}"))
messages.append(AIMessage(content=f"assistant {idx}"))
if idx >= 9 and (idx + 1) % 10 == 0:
messages.append(SystemMessage(content=f"late system after {idx + 1} turns"))
messages.append(SystemMessage(content="tail system"))
payload = model._get_request_payload(messages)
roles = [message["role"] for message in payload["messages"]]
assert roles[0] == "system"
assert roles.count("system") == 1
assert "base system" in payload["messages"][0]["content"]
assert "late system after 10 turns" in payload["messages"][0]["content"]
assert "tail system" in payload["messages"][0]["content"]
def test_writing_round_cap_reminder_is_not_late_system_message():
middleware = WritingSetupRoundCapMiddleware(max_rounds=3)
reminder = middleware._reminder()
assert isinstance(reminder, HumanMessage)
assert not isinstance(reminder, SystemMessage)
assert reminder.additional_kwargs["hide_from_ui"] is True