mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-10 05:58:36 +00:00
fix(todo-middleware): call super in wrap_model_call to restore write_todos prompt injection (#4714) (#4735)
* fix(todo-middleware): call super in wrap_model_call to restore write_todos prompt injection (#4714) * test(todo-middleware): add async no-reminder system-prompt passthrough test (review feedback)
This commit is contained in:
parent
21e2cfd719
commit
2bb230b334
@ -337,7 +337,11 @@ class TodoMiddleware(TodoListMiddleware):
|
|||||||
request: ModelRequest,
|
request: ModelRequest,
|
||||||
handler: Callable[[ModelRequest], ModelResponse],
|
handler: Callable[[ModelRequest], ModelResponse],
|
||||||
) -> ModelCallResult:
|
) -> ModelCallResult:
|
||||||
return handler(self._augment_request(request))
|
# The base class appends the `write_todos` system prompt to the request;
|
||||||
|
# without calling it the model is never told about the todo list feature.
|
||||||
|
# Augment with pending completion reminders on the request that already
|
||||||
|
# carries the injected system prompt.
|
||||||
|
return super().wrap_model_call(request, lambda req: handler(self._augment_request(req)))
|
||||||
|
|
||||||
@override
|
@override
|
||||||
async def awrap_model_call(
|
async def awrap_model_call(
|
||||||
@ -345,7 +349,11 @@ class TodoMiddleware(TodoListMiddleware):
|
|||||||
request: ModelRequest,
|
request: ModelRequest,
|
||||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||||
) -> ModelCallResult:
|
) -> ModelCallResult:
|
||||||
return await handler(self._augment_request(request))
|
# See wrap_model_call: preserve the base class system-prompt injection.
|
||||||
|
async def augmented_handler(req: ModelRequest) -> ModelResponse:
|
||||||
|
return await handler(self._augment_request(req))
|
||||||
|
|
||||||
|
return await super().awrap_model_call(request, augmented_handler)
|
||||||
|
|
||||||
@override
|
@override
|
||||||
def after_agent(self, state: ThreadState, runtime: Runtime) -> dict[str, Any] | None:
|
def after_agent(self, state: ThreadState, runtime: Runtime) -> dict[str, Any] | None:
|
||||||
|
|||||||
@ -2,11 +2,12 @@
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
from langchain.agents import create_agent
|
from langchain.agents import create_agent
|
||||||
|
from langchain.agents.middleware.types import ModelRequest
|
||||||
from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel
|
from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel
|
||||||
from langchain_core.messages import AIMessage, HumanMessage
|
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage
|
||||||
from pydantic import PrivateAttr
|
from pydantic import PrivateAttr
|
||||||
|
|
||||||
from deerflow.agents.middlewares.todo_middleware import (
|
from deerflow.agents.middlewares.todo_middleware import (
|
||||||
@ -59,6 +60,15 @@ def _make_runtime_for(thread_id: str, run_id: str):
|
|||||||
return runtime
|
return runtime
|
||||||
|
|
||||||
|
|
||||||
|
def _make_model_request(messages: list[Any], *, runtime=None) -> ModelRequest:
|
||||||
|
return ModelRequest(
|
||||||
|
model=object(),
|
||||||
|
messages=list(messages),
|
||||||
|
state={"messages": list(messages)},
|
||||||
|
runtime=runtime,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _sample_todos():
|
def _sample_todos():
|
||||||
return [
|
return [
|
||||||
{"status": "completed", "content": "Set up project"},
|
{"status": "completed", "content": "Set up project"},
|
||||||
@ -342,21 +352,24 @@ class TestAfterModel:
|
|||||||
assert result["jump_to"] == "model"
|
assert result["jump_to"] == "model"
|
||||||
assert "messages" not in result
|
assert "messages" not in result
|
||||||
|
|
||||||
request = MagicMock()
|
request = _make_model_request(state["messages"], runtime=runtime)
|
||||||
request.runtime = runtime
|
seen: list[ModelRequest] = []
|
||||||
request.messages = state["messages"]
|
|
||||||
request.override.return_value = "patched-request"
|
def handler(model_request: ModelRequest):
|
||||||
handler = MagicMock(return_value="response")
|
seen.append(model_request)
|
||||||
|
return "response"
|
||||||
|
|
||||||
assert mw.wrap_model_call(request, handler) == "response"
|
assert mw.wrap_model_call(request, handler) == "response"
|
||||||
request.override.assert_called_once()
|
assert len(seen) == 1
|
||||||
reminder = request.override.call_args.kwargs["messages"][-1]
|
sent = seen[0]
|
||||||
|
assert sent.system_message is not None
|
||||||
|
assert "write_todos" in sent.system_message.text
|
||||||
|
reminder = sent.messages[-1]
|
||||||
assert isinstance(reminder, HumanMessage)
|
assert isinstance(reminder, HumanMessage)
|
||||||
assert reminder.name == "todo_completion_reminder"
|
assert reminder.name == "todo_completion_reminder"
|
||||||
assert reminder.additional_kwargs["hide_from_ui"] is True
|
assert reminder.additional_kwargs["hide_from_ui"] is True
|
||||||
assert "Step 2" in reminder.content
|
assert "Step 2" in reminder.content
|
||||||
assert "Step 3" in reminder.content
|
assert "Step 3" in reminder.content
|
||||||
handler.assert_called_once_with("patched-request")
|
|
||||||
|
|
||||||
def test_reminder_lists_only_incomplete_items(self):
|
def test_reminder_lists_only_incomplete_items(self):
|
||||||
mw = TodoMiddleware()
|
mw = TodoMiddleware()
|
||||||
@ -368,12 +381,15 @@ class TestAfterModel:
|
|||||||
result = mw.after_model(state, runtime)
|
result = mw.after_model(state, runtime)
|
||||||
assert result is not None
|
assert result is not None
|
||||||
|
|
||||||
request = MagicMock()
|
request = _make_model_request(state["messages"], runtime=runtime)
|
||||||
request.runtime = runtime
|
seen: list[ModelRequest] = []
|
||||||
request.messages = state["messages"]
|
|
||||||
request.override.return_value = "patched-request"
|
def handler(model_request: ModelRequest):
|
||||||
mw.wrap_model_call(request, MagicMock(return_value="response"))
|
seen.append(model_request)
|
||||||
content = request.override.call_args.kwargs["messages"][-1].content
|
return "response"
|
||||||
|
|
||||||
|
mw.wrap_model_call(request, handler)
|
||||||
|
content = seen[0].messages[-1].content
|
||||||
assert "Step 1" not in content # completed — should not appear
|
assert "Step 1" not in content # completed — should not appear
|
||||||
assert "Step 2" in content
|
assert "Step 2" in content
|
||||||
assert "Step 3" in content
|
assert "Step 3" in content
|
||||||
@ -453,16 +469,26 @@ class TestAafterModel:
|
|||||||
|
|
||||||
|
|
||||||
class TestWrapModelCall:
|
class TestWrapModelCall:
|
||||||
def test_no_pending_reminder_passthrough(self):
|
def test_no_pending_reminder_still_injects_todo_system_prompt(self):
|
||||||
|
"""Regression for bytedance/deer-flow#4714: the base class system-prompt
|
||||||
|
injection must survive the override, otherwise the model never learns
|
||||||
|
about `write_todos` and no todo list is ever produced."""
|
||||||
mw = TodoMiddleware()
|
mw = TodoMiddleware()
|
||||||
request = MagicMock()
|
runtime = _make_runtime()
|
||||||
request.runtime = _make_runtime()
|
request = _make_model_request([HumanMessage(content="hi")], runtime=runtime)
|
||||||
request.messages = [HumanMessage(content="hi")]
|
seen: list[ModelRequest] = []
|
||||||
handler = MagicMock(return_value="response")
|
|
||||||
|
def handler(model_request: ModelRequest):
|
||||||
|
seen.append(model_request)
|
||||||
|
return "response"
|
||||||
|
|
||||||
assert mw.wrap_model_call(request, handler) == "response"
|
assert mw.wrap_model_call(request, handler) == "response"
|
||||||
request.override.assert_not_called()
|
assert len(seen) == 1
|
||||||
handler.assert_called_once_with(request)
|
sent = seen[0]
|
||||||
|
assert sent.system_message is not None
|
||||||
|
assert "write_todos" in sent.system_message.text
|
||||||
|
# No pending reminder — messages must not be augmented.
|
||||||
|
assert sent.messages == [HumanMessage(content="hi")]
|
||||||
|
|
||||||
def test_pending_reminder_is_injected_once(self):
|
def test_pending_reminder_is_injected_once(self):
|
||||||
mw = TodoMiddleware()
|
mw = TodoMiddleware()
|
||||||
@ -473,22 +499,28 @@ class TestWrapModelCall:
|
|||||||
}
|
}
|
||||||
mw.after_model(state, runtime)
|
mw.after_model(state, runtime)
|
||||||
|
|
||||||
request = MagicMock()
|
request = _make_model_request(state["messages"], runtime=runtime)
|
||||||
request.runtime = runtime
|
seen: list[ModelRequest] = []
|
||||||
request.messages = state["messages"]
|
|
||||||
request.override.return_value = "patched-request"
|
def handler(model_request: ModelRequest):
|
||||||
handler = MagicMock(return_value="response")
|
seen.append(model_request)
|
||||||
|
return "response"
|
||||||
|
|
||||||
assert mw.wrap_model_call(request, handler) == "response"
|
assert mw.wrap_model_call(request, handler) == "response"
|
||||||
injected_messages = request.override.call_args.kwargs["messages"]
|
assert len(seen) == 1
|
||||||
assert injected_messages[-1].name == "todo_completion_reminder"
|
sent = seen[0]
|
||||||
|
assert sent.system_message is not None
|
||||||
|
assert "write_todos" in sent.system_message.text
|
||||||
|
injected = sent.messages[-1]
|
||||||
|
assert isinstance(injected, HumanMessage)
|
||||||
|
assert injected.name == "todo_completion_reminder"
|
||||||
|
|
||||||
request.override.reset_mock()
|
# Second call: the reminder was drained, so only the system prompt is
|
||||||
handler.reset_mock()
|
# injected and the original messages pass through untouched.
|
||||||
handler.return_value = "second-response"
|
seen.clear()
|
||||||
assert mw.wrap_model_call(request, handler) == "second-response"
|
assert mw.wrap_model_call(request, handler) == "response"
|
||||||
request.override.assert_not_called()
|
assert len(seen) == 1
|
||||||
handler.assert_called_once_with(request)
|
assert seen[0].messages == state["messages"]
|
||||||
|
|
||||||
|
|
||||||
class TestTodoMiddlewareAgentGraphIntegration:
|
class TestTodoMiddlewareAgentGraphIntegration:
|
||||||
@ -528,6 +560,13 @@ class TestTodoMiddlewareAgentGraphIntegration:
|
|||||||
|
|
||||||
assert result["todos"] == [{"content": "Step 1", "status": "pending"}]
|
assert result["todos"] == [{"content": "Step 1", "status": "pending"}]
|
||||||
|
|
||||||
|
# Regression for bytedance/deer-flow#4714: the model request must carry
|
||||||
|
# the `write_todos` system prompt (injected by the base class), otherwise
|
||||||
|
# the model never produces a todo list.
|
||||||
|
first_model_call = model.seen_messages[0]
|
||||||
|
assert isinstance(first_model_call[0], SystemMessage)
|
||||||
|
assert "write_todos" in first_model_call[0].text
|
||||||
|
|
||||||
def test_completion_reminder_is_transient_in_real_agent_graph(self):
|
def test_completion_reminder_is_transient_in_real_agent_graph(self):
|
||||||
mw = TodoMiddleware()
|
mw = TodoMiddleware()
|
||||||
model = _CapturingFakeMessagesListChatModel(
|
model = _CapturingFakeMessagesListChatModel(
|
||||||
@ -642,6 +681,27 @@ class TestRunScopedReminderCleanup:
|
|||||||
|
|
||||||
|
|
||||||
class TestAwrapModelCall:
|
class TestAwrapModelCall:
|
||||||
|
def test_async_no_pending_reminder_still_injects_todo_system_prompt(self):
|
||||||
|
"""Mirror of the sync no-reminder test: the async path must also keep the
|
||||||
|
base class system-prompt injection while leaving messages untouched when
|
||||||
|
there are no pending completion reminders."""
|
||||||
|
mw = TodoMiddleware()
|
||||||
|
runtime = _make_runtime()
|
||||||
|
request = _make_model_request([HumanMessage(content="hi")], runtime=runtime)
|
||||||
|
seen: list[ModelRequest] = []
|
||||||
|
|
||||||
|
async def handler(model_request: ModelRequest):
|
||||||
|
seen.append(model_request)
|
||||||
|
return "response"
|
||||||
|
|
||||||
|
result = asyncio.run(mw.awrap_model_call(request, handler))
|
||||||
|
assert result == "response"
|
||||||
|
assert len(seen) == 1
|
||||||
|
sent = seen[0]
|
||||||
|
assert sent.system_message is not None
|
||||||
|
assert "write_todos" in sent.system_message.text
|
||||||
|
assert sent.messages == [HumanMessage(content="hi")]
|
||||||
|
|
||||||
def test_async_pending_reminder_is_injected(self):
|
def test_async_pending_reminder_is_injected(self):
|
||||||
mw = TodoMiddleware()
|
mw = TodoMiddleware()
|
||||||
runtime = _make_runtime()
|
runtime = _make_runtime()
|
||||||
@ -651,14 +711,20 @@ class TestAwrapModelCall:
|
|||||||
}
|
}
|
||||||
mw.after_model(state, runtime)
|
mw.after_model(state, runtime)
|
||||||
|
|
||||||
request = MagicMock()
|
request = _make_model_request(state["messages"], runtime=runtime)
|
||||||
request.runtime = runtime
|
seen: list[ModelRequest] = []
|
||||||
request.messages = state["messages"]
|
|
||||||
request.override.return_value = "patched-request"
|
async def handler(model_request: ModelRequest):
|
||||||
handler = AsyncMock(return_value="response")
|
seen.append(model_request)
|
||||||
|
return "response"
|
||||||
|
|
||||||
result = asyncio.run(mw.awrap_model_call(request, handler))
|
result = asyncio.run(mw.awrap_model_call(request, handler))
|
||||||
assert result == "response"
|
assert result == "response"
|
||||||
injected_messages = request.override.call_args.kwargs["messages"]
|
assert len(seen) == 1
|
||||||
assert injected_messages[-1].name == "todo_completion_reminder"
|
sent = seen[0]
|
||||||
handler.assert_awaited_once_with("patched-request")
|
assert sent.system_message is not None
|
||||||
|
assert "write_todos" in sent.system_message.text
|
||||||
|
injected = sent.messages[-1]
|
||||||
|
assert isinstance(injected, HumanMessage)
|
||||||
|
assert injected.name == "todo_completion_reminder"
|
||||||
|
assert injected.additional_kwargs["hide_from_ui"] is True
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user