deer-flow/backend/tests/test_token_usage_middleware.py
ChiHaYa 88252e9b31
fix(subagents): isolate background tasks from reused tool call IDs (#4758)
* fix(subagents): isolate background execution IDs

* fix(subagents): preserve correlation scope and isolate usage

* fix(subagents): make usage attribution idempotent
2026-08-12 09:25:05 +08:00

370 lines
13 KiB
Python

"""Tests for TokenUsageMiddleware attribution annotations."""
import logging
from unittest.mock import MagicMock
from langchain_core.messages import AIMessage, ToolMessage
from langgraph.graph.message import add_messages
from deerflow.agents.middlewares.token_usage_middleware import (
TOKEN_USAGE_ATTRIBUTION_KEY,
TokenUsageMiddleware,
_build_todo_actions,
)
from deerflow.subagents.status_contract import SUBAGENT_TOKEN_USAGE_KEY
def _make_runtime():
runtime = MagicMock()
runtime.context = {"thread_id": "test-thread"}
return runtime
class TestTokenUsageMiddleware:
def test_logs_cache_token_details(self, caplog):
middleware = TokenUsageMiddleware()
message = AIMessage(
content="Here is the final answer.",
usage_metadata={
"input_tokens": 350,
"output_tokens": 240,
"total_tokens": 590,
"input_token_details": {
"audio": 10,
"cache_creation": 200,
"cache_read": 100,
},
"output_token_details": {
"audio": 10,
"reasoning": 200,
},
},
)
with caplog.at_level(
logging.INFO,
logger="deerflow.agents.middlewares.token_usage_middleware",
):
result = middleware.after_model({"messages": [message]}, _make_runtime())
assert result is not None
assert "LLM token usage: input=350 output=240 total=590" in caplog.text
assert "input_token_details={'audio': 10, 'cache_creation': 200, 'cache_read': 100}" in caplog.text
assert "output_token_details={'audio': 10, 'reasoning': 200}" in caplog.text
def test_logs_basic_tokens_when_no_detail_fields_in_usage_metadata(self, caplog):
"""When usage_metadata has only totals (no input_token_details), log just the counts."""
middleware = TokenUsageMiddleware()
message = AIMessage(
content="Here is the final answer.",
usage_metadata={
"input_tokens": 350,
"output_tokens": 240,
"total_tokens": 590,
},
)
with caplog.at_level(
logging.INFO,
logger="deerflow.agents.middlewares.token_usage_middleware",
):
result = middleware.after_model({"messages": [message]}, _make_runtime())
assert result is not None
assert "LLM token usage: input=350 output=240 total=590" in caplog.text
assert "input_token_details" not in caplog.text
def test_no_log_when_usage_metadata_is_missing(self, caplog):
"""When usage_metadata is absent, no token usage line is logged."""
middleware = TokenUsageMiddleware()
message = AIMessage(
content="Here is the final answer.",
response_metadata={
"usage": {
"input_tokens": 350,
"output_tokens": 240,
"total_tokens": 590,
}
},
)
with caplog.at_level(
logging.INFO,
logger="deerflow.agents.middlewares.token_usage_middleware",
):
result = middleware.after_model({"messages": [message]}, _make_runtime())
assert result is not None
assert "LLM token usage" not in caplog.text
def test_annotates_todo_updates_with_structured_actions(self):
middleware = TokenUsageMiddleware()
message = AIMessage(
content="",
tool_calls=[
{
"id": "write_todos:1",
"name": "write_todos",
"args": {
"todos": [
{"content": "Inspect streaming path", "status": "completed"},
{"content": "Design token attribution schema", "status": "in_progress"},
]
},
}
],
usage_metadata={"input_tokens": 100, "output_tokens": 20, "total_tokens": 120},
)
state = {
"messages": [message],
"todos": [
{"content": "Inspect streaming path", "status": "in_progress"},
{"content": "Design token attribution schema", "status": "pending"},
],
}
result = middleware.after_model(state, _make_runtime())
assert result is not None
updated_message = result["messages"][0]
attribution = updated_message.additional_kwargs[TOKEN_USAGE_ATTRIBUTION_KEY]
assert attribution["kind"] == "tool_batch"
assert attribution["shared_attribution"] is True
assert attribution["tool_call_ids"] == ["write_todos:1"]
assert attribution["actions"] == [
{
"kind": "todo_complete",
"content": "Inspect streaming path",
"tool_call_id": "write_todos:1",
},
{
"kind": "todo_start",
"content": "Design token attribution schema",
"tool_call_id": "write_todos:1",
},
]
def test_annotates_subagent_and_search_steps(self):
middleware = TokenUsageMiddleware()
message = AIMessage(
content="",
tool_calls=[
{
"id": "task:1",
"name": "task",
"args": {
"description": "spec-coder patch message grouping",
"subagent_type": "general-purpose",
},
},
{
"id": "web_search:1",
"name": "web_search",
"args": {"query": "LangGraph useStream messages tuple"},
},
],
)
result = middleware.after_model({"messages": [message]}, _make_runtime())
assert result is not None
attribution = result["messages"][0].additional_kwargs[TOKEN_USAGE_ATTRIBUTION_KEY]
assert attribution["kind"] == "tool_batch"
assert attribution["shared_attribution"] is True
assert attribution["actions"] == [
{
"kind": "subagent",
"description": "spec-coder patch message grouping",
"subagent_type": "general-purpose",
"tool_call_id": "task:1",
},
{
"kind": "search",
"tool_name": "web_search",
"query": "LangGraph useStream messages tuple",
"tool_call_id": "web_search:1",
},
]
def test_marks_final_answer_when_no_tools(self):
middleware = TokenUsageMiddleware()
message = AIMessage(content="Here is the final answer.")
result = middleware.after_model({"messages": [message]}, _make_runtime())
assert result is not None
attribution = result["messages"][0].additional_kwargs[TOKEN_USAGE_ATTRIBUTION_KEY]
assert attribution["kind"] == "final_answer"
assert attribution["shared_attribution"] is False
assert attribution["actions"] == []
def test_annotates_removed_todos(self):
middleware = TokenUsageMiddleware()
message = AIMessage(
content="",
tool_calls=[
{
"id": "write_todos:remove",
"name": "write_todos",
"args": {
"todos": [],
},
}
],
)
result = middleware.after_model(
{
"messages": [message],
"todos": [
{"content": "Archive obsolete plan", "status": "pending"},
],
},
_make_runtime(),
)
assert result is not None
attribution = result["messages"][0].additional_kwargs[TOKEN_USAGE_ATTRIBUTION_KEY]
assert attribution["kind"] == "todo_update"
assert attribution["shared_attribution"] is False
assert attribution["actions"] == [
{
"kind": "todo_remove",
"content": "Archive obsolete plan",
"tool_call_id": "write_todos:remove",
}
]
def test_merges_subagent_usage_by_message_position_when_ai_message_ids_are_missing(self):
middleware = TokenUsageMiddleware()
first_dispatch = AIMessage(
content="",
tool_calls=[{"id": "task:first", "name": "task", "args": {}}],
)
second_dispatch = AIMessage(
content="",
tool_calls=[
{"id": "task:second-a", "name": "task", "args": {}},
{"id": "task:second-b", "name": "task", "args": {}},
],
)
messages = [
first_dispatch,
ToolMessage(content="first", tool_call_id="task:first"),
second_dispatch,
ToolMessage(
content="second-a",
tool_call_id="task:second-a",
additional_kwargs={SUBAGENT_TOKEN_USAGE_KEY: {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}},
),
ToolMessage(
content="second-b",
tool_call_id="task:second-b",
additional_kwargs={SUBAGENT_TOKEN_USAGE_KEY: {"input_tokens": 20, "output_tokens": 7, "total_tokens": 27}},
),
AIMessage(content="done"),
]
result = middleware.after_model({"messages": messages}, _make_runtime())
assert result is not None
usage_updates = [message for message in result["messages"] if getattr(message, "usage_metadata", None)]
assert len(usage_updates) == 1
updated = usage_updates[0]
assert updated.tool_calls == second_dispatch.tool_calls
assert updated.usage_metadata == {
"input_tokens": 30,
"output_tokens": 12,
"total_tokens": 42,
}
def test_reused_tool_call_id_keeps_usage_scoped_to_each_run_history(self):
middleware = TokenUsageMiddleware()
tool_call_id = "reused-provider-tool-call-id"
def apply_usage(usage):
dispatch = AIMessage(
content="",
tool_calls=[{"id": tool_call_id, "name": "task", "args": {}}],
)
messages = [
dispatch,
ToolMessage(
content="task result",
tool_call_id=tool_call_id,
additional_kwargs={SUBAGENT_TOKEN_USAGE_KEY: usage},
),
AIMessage(content="done"),
]
result = middleware.after_model({"messages": messages}, _make_runtime())
assert result is not None
return next(message.usage_metadata for message in result["messages"] if getattr(message, "usage_metadata", None))
assert apply_usage({"input_tokens": 10, "output_tokens": 2, "total_tokens": 12}) == {
"input_tokens": 10,
"output_tokens": 2,
"total_tokens": 12,
}
assert apply_usage({"input_tokens": 90, "output_tokens": 8, "total_tokens": 98}) == {
"input_tokens": 90,
"output_tokens": 8,
"total_tokens": 98,
}
def test_subagent_usage_attribution_is_idempotent_when_state_is_reprocessed(self):
middleware = TokenUsageMiddleware()
dispatch = AIMessage(
id="dispatch-message",
content="",
tool_calls=[{"id": "task:replayed", "name": "task", "args": {}}],
)
tool_result = ToolMessage(
id="tool-message",
content="task result",
tool_call_id="task:replayed",
additional_kwargs={
SUBAGENT_TOKEN_USAGE_KEY: {
"input_tokens": 10,
"output_tokens": 2,
"total_tokens": 12,
}
},
)
final = AIMessage(id="final-message", content="done")
messages = [dispatch, tool_result, final]
first_update = middleware.after_model({"messages": messages}, _make_runtime())
assert first_update is not None
checkpoint_messages = add_messages(messages, first_update["messages"])
updated_dispatch = next(message for message in checkpoint_messages if message.id == dispatch.id)
assert updated_dispatch.usage_metadata == {
"input_tokens": 10,
"output_tokens": 2,
"total_tokens": 12,
}
second_update = middleware.after_model({"messages": checkpoint_messages}, _make_runtime())
assert second_update is None
class TestBuildTodoActions:
def test_duplicate_content_emits_todo_remove(self):
"""When next_todos has duplicate content entries that exhaust previous_by_content,
the positional fallback must not consume an unrelated previous todo as matched.
The unrelated previous entry should still produce a todo_remove action."""
previous = [
{"content": "A", "status": "pending"},
{"content": "B", "status": "pending"},
]
next_todos = [
{"content": "A", "status": "in_progress"},
{"content": "A", "status": "completed"},
]
actions = _build_todo_actions(previous, next_todos)
assert any(a.get("kind") == "todo_remove" and a.get("content") == "B" for a in actions), f"Expected todo_remove for B but got: {actions}"