274 lines
9.8 KiB
Python
274 lines
9.8 KiB
Python
"""Tests for PositionArtifactCapMiddleware."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from types import SimpleNamespace
|
|
|
|
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
|
|
|
from deerflow.agents.middlewares.position_artifact_cap_middleware import (
|
|
PositionArtifactCapMiddleware,
|
|
_reminder_for,
|
|
_successful_outputs_writes,
|
|
_successful_presents,
|
|
)
|
|
|
|
|
|
def _runtime(cap: int | None = 1) -> SimpleNamespace:
|
|
context: dict = {"thread_id": "t-pos"}
|
|
if cap is not None:
|
|
context["position_artifact_cap"] = cap
|
|
return SimpleNamespace(context=context, config={"configurable": dict(context)})
|
|
|
|
|
|
def _write_call(call_id: str, path: str = "/mnt/user-data/outputs/report.md") -> dict:
|
|
return {
|
|
"id": call_id,
|
|
"name": "write_file",
|
|
"args": {"path": path, "content": "# hello", "description": "deliver"},
|
|
}
|
|
|
|
|
|
def _present_call(call_id: str, path: str = "/mnt/user-data/outputs/report.md") -> dict:
|
|
return {
|
|
"id": call_id,
|
|
"name": "present_files",
|
|
"args": {"filepaths": [path]},
|
|
}
|
|
|
|
|
|
def test_successful_outputs_writes_counts_ok_in_current_turn_only() -> None:
|
|
messages = [
|
|
HumanMessage(content="上一轮"),
|
|
AIMessage(content="", tool_calls=[_write_call("old")]),
|
|
ToolMessage(content="OK", tool_call_id="old", name="write_file"),
|
|
HumanMessage(content="本轮请交付"),
|
|
AIMessage(content="", tool_calls=[_write_call("new")]),
|
|
ToolMessage(content="OK", tool_call_id="new", name="write_file"),
|
|
]
|
|
assert _successful_outputs_writes(messages, turn_start=3) == 1
|
|
|
|
|
|
def test_after_model_truncates_second_outputs_write_in_same_response() -> None:
|
|
mw = PositionArtifactCapMiddleware()
|
|
state = {
|
|
"messages": [
|
|
HumanMessage(content="请交付一份报告"),
|
|
AIMessage(
|
|
content="正在写入",
|
|
tool_calls=[
|
|
_write_call("w1"),
|
|
_write_call("w2", "/mnt/user-data/outputs/report-copy.md"),
|
|
_present_call("p1"),
|
|
],
|
|
),
|
|
]
|
|
}
|
|
update = mw.after_model(state, _runtime(1))
|
|
assert update is not None
|
|
tool_calls = update["messages"][0].tool_calls
|
|
assert [tc["id"] for tc in tool_calls] == ["w1", "p1"]
|
|
|
|
|
|
def test_after_model_strips_all_writes_once_cap_already_met() -> None:
|
|
mw = PositionArtifactCapMiddleware()
|
|
state = {
|
|
"messages": [
|
|
HumanMessage(content="请交付"),
|
|
AIMessage(content="", tool_calls=[_write_call("done")]),
|
|
ToolMessage(content="OK", tool_call_id="done", name="write_file"),
|
|
AIMessage(content="再写一份", tool_calls=[_write_call("again")]),
|
|
]
|
|
}
|
|
update = mw.after_model(state, _runtime(1))
|
|
assert update is not None
|
|
assert update["messages"][0].tool_calls == []
|
|
assert update.get("jump_to") == "end"
|
|
assert "产物上限" not in str(update["messages"][0].content)
|
|
assert "tool_calls" not in (update["messages"][0].additional_kwargs or {})
|
|
|
|
|
|
def test_after_model_noop_without_cap() -> None:
|
|
mw = PositionArtifactCapMiddleware()
|
|
state = {
|
|
"messages": [
|
|
HumanMessage(content="请交付"),
|
|
AIMessage(content="", tool_calls=[_write_call("w1"), _write_call("w2")]),
|
|
]
|
|
}
|
|
assert mw.after_model(state, _runtime(None)) is None
|
|
|
|
|
|
def test_wrap_tool_call_rejects_after_successful_delivery() -> None:
|
|
mw = PositionArtifactCapMiddleware()
|
|
state = {
|
|
"messages": [
|
|
HumanMessage(content="请交付"),
|
|
AIMessage(content="", tool_calls=[_write_call("done")]),
|
|
ToolMessage(content="OK", tool_call_id="done", name="write_file"),
|
|
]
|
|
}
|
|
request = SimpleNamespace(
|
|
tool_call=_write_call("again"),
|
|
runtime=_runtime(1),
|
|
state=state,
|
|
)
|
|
called = {"n": 0}
|
|
|
|
def handler(_request):
|
|
called["n"] += 1
|
|
return ToolMessage(content="OK", tool_call_id="again", name="write_file")
|
|
|
|
result = mw.wrap_tool_call(request, handler)
|
|
assert called["n"] == 0
|
|
assert isinstance(result, ToolMessage)
|
|
assert str(result.content).startswith("Error:")
|
|
|
|
|
|
def test_wrap_tool_call_allows_first_write() -> None:
|
|
mw = PositionArtifactCapMiddleware()
|
|
state = {"messages": [HumanMessage(content="请交付")]}
|
|
request = SimpleNamespace(
|
|
tool_call=_write_call("w1"),
|
|
runtime=_runtime(1),
|
|
state=state,
|
|
)
|
|
|
|
def handler(_request):
|
|
return ToolMessage(content="OK", tool_call_id="w1", name="write_file")
|
|
|
|
result = mw.wrap_tool_call(request, handler)
|
|
assert result.content == "OK"
|
|
|
|
|
|
def test_wrap_tool_call_rejects_parallel_second_write_in_same_batch() -> None:
|
|
"""Two write_file in one AIMessage can run concurrently; order in the list is the budget."""
|
|
mw = PositionArtifactCapMiddleware()
|
|
state = {
|
|
"messages": [
|
|
HumanMessage(content="请交付"),
|
|
AIMessage(content="", tool_calls=[_write_call("w1"), _write_call("w2")]),
|
|
]
|
|
}
|
|
first = SimpleNamespace(tool_call=_write_call("w1"), runtime=_runtime(1), state=state)
|
|
second = SimpleNamespace(tool_call=_write_call("w2"), runtime=_runtime(1), state=state)
|
|
|
|
def handler(request):
|
|
return ToolMessage(content="OK", tool_call_id=request.tool_call["id"], name="write_file")
|
|
|
|
assert mw.wrap_tool_call(first, handler).content == "OK"
|
|
rejected = mw.wrap_tool_call(second, handler)
|
|
assert str(rejected.content).startswith("Error:")
|
|
|
|
|
|
def test_after_model_clears_provider_tool_call_metadata() -> None:
|
|
mw = PositionArtifactCapMiddleware()
|
|
extra = {
|
|
"tool_calls": [
|
|
{"id": "w1", "type": "function", "function": {"name": "write_file"}},
|
|
{"id": "w2", "type": "function", "function": {"name": "write_file"}},
|
|
]
|
|
}
|
|
ai = AIMessage(content="写两份", tool_calls=[_write_call("w1"), _write_call("w2")])
|
|
ai.additional_kwargs = extra
|
|
state = {"messages": [HumanMessage(content="请交付"), ai]}
|
|
update = mw.after_model(state, _runtime(1))
|
|
assert update is not None
|
|
kept = update["messages"][0]
|
|
assert [tc["id"] for tc in kept.tool_calls] == ["w1"]
|
|
assert [item["id"] for item in kept.additional_kwargs.get("tool_calls", [])] == ["w1"]
|
|
|
|
|
|
def test_action_plan_style_cap_two_allows_second_write() -> None:
|
|
"""Action-plan is exempt in the UI (no flag); if a cap of 2 were set it must allow both."""
|
|
mw = PositionArtifactCapMiddleware()
|
|
state = {
|
|
"messages": [
|
|
HumanMessage(content="生成行动规划"),
|
|
AIMessage(content="", tool_calls=[_write_call("md", "/mnt/user-data/outputs/行动规划报告.md")]),
|
|
ToolMessage(content="OK", tool_call_id="md", name="write_file"),
|
|
AIMessage(
|
|
content="",
|
|
tool_calls=[_write_call("json", "/mnt/user-data/outputs/action-plan-subtasks.json")],
|
|
),
|
|
]
|
|
}
|
|
assert mw.after_model(state, _runtime(2)) is None
|
|
|
|
|
|
def test_after_model_truncates_second_present_files_in_same_response() -> None:
|
|
mw = PositionArtifactCapMiddleware()
|
|
state = {
|
|
"messages": [
|
|
HumanMessage(content="请交付"),
|
|
AIMessage(
|
|
content="",
|
|
tool_calls=[_write_call("w1"), _present_call("p1"), _present_call("p2")],
|
|
),
|
|
]
|
|
}
|
|
update = mw.after_model(state, _runtime(1))
|
|
assert update is not None
|
|
assert [tc["id"] for tc in update["messages"][0].tool_calls] == ["w1", "p1"]
|
|
|
|
|
|
def test_after_model_strips_second_present_after_first_succeeded() -> None:
|
|
"""Screenshot bug: write+present already done, weak model presents the same file again."""
|
|
mw = PositionArtifactCapMiddleware()
|
|
state = {
|
|
"messages": [
|
|
HumanMessage(content="请评估"),
|
|
AIMessage(content="", tool_calls=[_write_call("w1"), _present_call("p1")]),
|
|
ToolMessage(content="OK", tool_call_id="w1", name="write_file"),
|
|
ToolMessage(content="Successfully presented files", tool_call_id="p1", name="present_files"),
|
|
AIMessage(content="评估报告已交付", tool_calls=[_present_call("p2")]),
|
|
]
|
|
}
|
|
update = mw.after_model(state, _runtime(1))
|
|
assert update is not None
|
|
assert update["messages"][0].tool_calls == []
|
|
assert update.get("jump_to") == "end"
|
|
|
|
|
|
def test_wrap_tool_call_rejects_second_present_files() -> None:
|
|
mw = PositionArtifactCapMiddleware()
|
|
state = {
|
|
"messages": [
|
|
HumanMessage(content="请交付"),
|
|
AIMessage(content="", tool_calls=[_present_call("p1")]),
|
|
ToolMessage(content="Successfully presented files", tool_call_id="p1", name="present_files"),
|
|
]
|
|
}
|
|
request = SimpleNamespace(
|
|
tool_call=_present_call("p2"),
|
|
runtime=_runtime(1),
|
|
state=state,
|
|
)
|
|
called = {"n": 0}
|
|
|
|
def handler(_request):
|
|
called["n"] += 1
|
|
return ToolMessage(content="Successfully presented files", tool_call_id="p2", name="present_files")
|
|
|
|
result = mw.wrap_tool_call(request, handler)
|
|
assert called["n"] == 0
|
|
assert isinstance(result, ToolMessage)
|
|
assert "已经呈现" in str(result.content)
|
|
|
|
|
|
def test_reminder_asks_present_only_when_not_yet_presented() -> None:
|
|
assert "present_files" in (_reminder_for(writes_done=True, presents_done=False) or "")
|
|
done = _reminder_for(writes_done=True, presents_done=True) or ""
|
|
assert "已经 present_files" in done
|
|
assert "请立即 present_files" not in done
|
|
assert _reminder_for(writes_done=False, presents_done=False) is None
|
|
|
|
|
|
def test_successful_presents_counts_tool_result() -> None:
|
|
messages = [
|
|
HumanMessage(content="请交付"),
|
|
AIMessage(content="", tool_calls=[_present_call("p1")]),
|
|
ToolMessage(content="Successfully presented files", tool_call_id="p1", name="present_files"),
|
|
]
|
|
assert _successful_presents(messages, turn_start=0) == 1
|