mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-14 16:08:41 +00:00
* fix(subagents): isolate background execution IDs * fix(subagents): preserve correlation scope and isolate usage * fix(subagents): make usage attribution idempotent
370 lines
13 KiB
Python
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}"
|