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

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