"""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