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:
Nan Gao 2026-08-11 11:41:31 +08:00 committed by GitHub
parent 21e2cfd719
commit 2bb230b334
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 120 additions and 46 deletions

View File

@ -337,7 +337,11 @@ class TodoMiddleware(TodoListMiddleware):
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> 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
async def awrap_model_call(
@ -345,7 +349,11 @@ class TodoMiddleware(TodoListMiddleware):
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> 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
def after_agent(self, state: ThreadState, runtime: Runtime) -> dict[str, Any] | None:

View File

@ -2,11 +2,12 @@
import asyncio
from typing import Any
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import MagicMock
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.messages import AIMessage, HumanMessage
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage
from pydantic import PrivateAttr
from deerflow.agents.middlewares.todo_middleware import (
@ -59,6 +60,15 @@ def _make_runtime_for(thread_id: str, run_id: str):
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():
return [
{"status": "completed", "content": "Set up project"},
@ -342,21 +352,24 @@ class TestAfterModel:
assert result["jump_to"] == "model"
assert "messages" not in result
request = MagicMock()
request.runtime = runtime
request.messages = state["messages"]
request.override.return_value = "patched-request"
handler = MagicMock(return_value="response")
request = _make_model_request(state["messages"], runtime=runtime)
seen: list[ModelRequest] = []
def handler(model_request: ModelRequest):
seen.append(model_request)
return "response"
assert mw.wrap_model_call(request, handler) == "response"
request.override.assert_called_once()
reminder = request.override.call_args.kwargs["messages"][-1]
assert len(seen) == 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 reminder.name == "todo_completion_reminder"
assert reminder.additional_kwargs["hide_from_ui"] is True
assert "Step 2" in reminder.content
assert "Step 3" in reminder.content
handler.assert_called_once_with("patched-request")
def test_reminder_lists_only_incomplete_items(self):
mw = TodoMiddleware()
@ -368,12 +381,15 @@ class TestAfterModel:
result = mw.after_model(state, runtime)
assert result is not None
request = MagicMock()
request.runtime = runtime
request.messages = state["messages"]
request.override.return_value = "patched-request"
mw.wrap_model_call(request, MagicMock(return_value="response"))
content = request.override.call_args.kwargs["messages"][-1].content
request = _make_model_request(state["messages"], runtime=runtime)
seen: list[ModelRequest] = []
def handler(model_request: ModelRequest):
seen.append(model_request)
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 2" in content
assert "Step 3" in content
@ -453,16 +469,26 @@ class TestAafterModel:
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()
request = MagicMock()
request.runtime = _make_runtime()
request.messages = [HumanMessage(content="hi")]
handler = MagicMock(return_value="response")
runtime = _make_runtime()
request = _make_model_request([HumanMessage(content="hi")], runtime=runtime)
seen: list[ModelRequest] = []
def handler(model_request: ModelRequest):
seen.append(model_request)
return "response"
assert mw.wrap_model_call(request, handler) == "response"
request.override.assert_not_called()
handler.assert_called_once_with(request)
assert len(seen) == 1
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):
mw = TodoMiddleware()
@ -473,22 +499,28 @@ class TestWrapModelCall:
}
mw.after_model(state, runtime)
request = MagicMock()
request.runtime = runtime
request.messages = state["messages"]
request.override.return_value = "patched-request"
handler = MagicMock(return_value="response")
request = _make_model_request(state["messages"], runtime=runtime)
seen: list[ModelRequest] = []
def handler(model_request: ModelRequest):
seen.append(model_request)
return "response"
assert mw.wrap_model_call(request, handler) == "response"
injected_messages = request.override.call_args.kwargs["messages"]
assert injected_messages[-1].name == "todo_completion_reminder"
assert len(seen) == 1
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()
handler.reset_mock()
handler.return_value = "second-response"
assert mw.wrap_model_call(request, handler) == "second-response"
request.override.assert_not_called()
handler.assert_called_once_with(request)
# Second call: the reminder was drained, so only the system prompt is
# injected and the original messages pass through untouched.
seen.clear()
assert mw.wrap_model_call(request, handler) == "response"
assert len(seen) == 1
assert seen[0].messages == state["messages"]
class TestTodoMiddlewareAgentGraphIntegration:
@ -528,6 +560,13 @@ class TestTodoMiddlewareAgentGraphIntegration:
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):
mw = TodoMiddleware()
model = _CapturingFakeMessagesListChatModel(
@ -642,6 +681,27 @@ class TestRunScopedReminderCleanup:
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):
mw = TodoMiddleware()
runtime = _make_runtime()
@ -651,14 +711,20 @@ class TestAwrapModelCall:
}
mw.after_model(state, runtime)
request = MagicMock()
request.runtime = runtime
request.messages = state["messages"]
request.override.return_value = "patched-request"
handler = AsyncMock(return_value="response")
request = _make_model_request(state["messages"], 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"
injected_messages = request.override.call_args.kwargs["messages"]
assert injected_messages[-1].name == "todo_completion_reminder"
handler.assert_awaited_once_with("patched-request")
assert len(seen) == 1
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"
assert injected.additional_kwargs["hide_from_ui"] is True