85 lines
3.6 KiB
Python
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
|