mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-14 16:08:41 +00:00
fix(agents): keep the token budget across goal continuations of a run (#5410)
* fix(agents): keep the token budget across goal continuations of a run TokenBudgetMiddleware cleared its usage in after_agent, and before_agent marks every existing message as seen. A Gateway run re-enters the graph for hidden goal continuations under the same run_id, so each continuation started from zero and could spend another full budget, even after the user turn had already hit the hard stop. Keep the run's usage and warning state across graph entries, as LoopDetectionMiddleware does since #5344. Only the per-message seen map is dropped, and before_agent rebuilds it. Invocations without a run_id in the context still clear everything. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com> * fix(agents): normalize token budget run identity and test cleanup --------- Co-authored-by: Claude Opus 5 <noreply@anthropic.com> Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
This commit is contained in:
parent
bbced51022
commit
96d6ff3aca
@ -109,7 +109,7 @@ Before changing a later authorization phase, read the [authorization RFC](../../
|
||||
hooks topology-stable. Before changing this guard, read
|
||||
[Loop detection lifecycle](../../../../../docs/LOOP_DETECTION.md) for
|
||||
fallback identity, cleanup/LRU/reset, severity ordering, and test invariants.
|
||||
30. **TokenBudgetMiddleware** - *(optional, if `token_budget.enabled`)* Enforces per-run token limits
|
||||
30. **TokenBudgetMiddleware** - `token_budget.enabled`: shares run-ID budgets across continuations; missing/invalid IDs clear invocation state.
|
||||
31. **Custom middlewares** - *(optional)* Any `custom_middlewares` passed to `build_middlewares` are injected here, before config-declared extensions and the terminal-response/safety/clarification tail
|
||||
32. **Configured extension middlewares** - `extensions.middlewares` in `config.yaml` or `extensions_config.json` optionally accepts `module.path:ClassName` strings or `{class, kwargs}` objects. `deerflow.reflection.resolve_class` loads `AgentMiddleware` classes; import, class, and constructor errors fail agent creation. `kwargs` must be JSON-compatible; YAML dates/timestamps become ISO strings. Order: built-ins/custom and loop/token guards → extensions → terminal-response/safety/clarification tail. Subagents share the list before their safety tail; separate lead/subagent lists are unsupported. Trusted operator config only: paths instantiate arbitrary code. Gateway skill/MCP toggles preserve it in raw JSON; adding an API write path requires explicit trust-boundary review.
|
||||
33. **TerminalResponseMiddleware** - When a provider returns an empty terminal `AIMessage` after tool execution, injects a hidden recovery prompt and retries the model once; a second empty response is replaced in checkpoint state by a visible error fallback marked for the run worker, so the run finishes as an error instead of a silent success
|
||||
|
||||
@ -15,6 +15,15 @@ Warning injection uses the deferred pattern:
|
||||
- wrap_model_call injects it as a HumanMessage at the next model call.
|
||||
This preserves AIMessage(tool_calls) → ToolMessage pairing.
|
||||
|
||||
Run scope:
|
||||
Usage and warning state are keyed by ``run_id`` and survive ``after_agent``.
|
||||
A single Gateway run may re-enter the graph for hidden goal continuations,
|
||||
and those continuations share one budget; a later user run gets a new
|
||||
``run_id`` and a fresh budget. Only the per-message ``seen`` map is dropped
|
||||
(``before_agent`` rebuilds it). Invocations without a non-empty string
|
||||
``run_id`` use runtime-local identity and clear their usage/warning state
|
||||
in ``after_agent``.
|
||||
|
||||
Stop-reason surfacing (#3875 Phase 2):
|
||||
The hard stop does NOT raise — it strips tool_calls so the agent loop
|
||||
terminates naturally and produces a final answer. To let the caller (e.g.
|
||||
@ -105,12 +114,16 @@ class TokenBudgetMiddleware(AgentMiddleware[AgentState]):
|
||||
return self._stop_reason.pop(run_id, None)
|
||||
|
||||
@staticmethod
|
||||
def _get_run_id(runtime: Runtime) -> str:
|
||||
def _context_run_id(runtime: Runtime) -> str | None:
|
||||
"""Resolve the explicit identity shared by continuation invocations."""
|
||||
ctx = getattr(runtime, "context", None)
|
||||
if isinstance(ctx, dict) and "run_id" in ctx:
|
||||
return ctx["run_id"]
|
||||
run_id = ctx.get("run_id") if isinstance(ctx, dict) else None
|
||||
return run_id if isinstance(run_id, str) and run_id else None
|
||||
|
||||
@classmethod
|
||||
def _get_run_id(cls, runtime: Runtime) -> str:
|
||||
# Fallback to runtime object ID to prevent collisions across embedded client runs
|
||||
return str(id(runtime))
|
||||
return cls._context_run_id(runtime) or str(id(runtime))
|
||||
|
||||
def _clear_run_state(self, run_id: str) -> None:
|
||||
with self._lock:
|
||||
@ -149,7 +162,15 @@ class TokenBudgetMiddleware(AgentMiddleware[AgentState]):
|
||||
def after_agent(self, state: AgentState, runtime: Runtime) -> None:
|
||||
if not self._config.enabled:
|
||||
return
|
||||
self._clear_run_state(self._get_run_id(runtime))
|
||||
run_id = self._get_run_id(runtime)
|
||||
if self._context_run_id(runtime) is not None:
|
||||
# A Gateway run re-enters the graph for hidden goal continuations
|
||||
# under the same run_id, and they share this run's budget. Keep the
|
||||
# usage and warning state; before_agent rebuilds the seen map.
|
||||
with self._lock:
|
||||
self._seen_messages.pop(run_id, None)
|
||||
return
|
||||
self._clear_run_state(run_id)
|
||||
|
||||
@override
|
||||
async def aafter_agent(self, state: AgentState, runtime: Runtime) -> None:
|
||||
|
||||
@ -1,6 +1,11 @@
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from langchain.agents import create_agent
|
||||
from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel
|
||||
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||||
from langchain_core.tools import tool as as_tool
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
|
||||
from deerflow.agents.middlewares.token_budget_middleware import TokenBudgetMiddleware
|
||||
from deerflow.config.token_budget_config import TokenBudgetConfig
|
||||
@ -80,6 +85,68 @@ class TestTokenBudgetTracking:
|
||||
assert "TOKEN BUDGET WARNING" in mw._pending_warnings["test-run"][0]
|
||||
|
||||
|
||||
class TestTokenBudgetLifecycle:
|
||||
@pytest.mark.parametrize("context", [None, {}, {"run_id": None}, {"run_id": ""}, {"run_id": 0}, {"run_id": []}])
|
||||
@pytest.mark.parametrize("async_hooks", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_or_invalid_run_id_clears_invocation_state(self, context, async_hooks):
|
||||
mw = TokenBudgetMiddleware(TokenBudgetConfig(enabled=True, max_tokens=1000))
|
||||
runtime = _make_runtime()
|
||||
runtime.context = context
|
||||
state = _make_state_with_usage(total=850)
|
||||
if async_hooks:
|
||||
await mw.abefore_agent({"messages": []}, runtime)
|
||||
await mw.aafter_model(state, runtime)
|
||||
else:
|
||||
mw.before_agent({"messages": []}, runtime)
|
||||
mw.after_model(state, runtime)
|
||||
|
||||
# Missing identities are invocation-local, never shared under None or "".
|
||||
key = str(id(runtime))
|
||||
assert mw._get_run_id(runtime) == key
|
||||
assert mw._cumulative_usage[key].total == 850
|
||||
assert mw._warned[key]
|
||||
assert mw._pending_warnings[key]
|
||||
assert mw._seen_messages[key]
|
||||
|
||||
if async_hooks:
|
||||
await mw.aafter_agent(state, runtime)
|
||||
else:
|
||||
mw.after_agent(state, runtime)
|
||||
for values in (mw._cumulative_usage, mw._warned, mw._pending_warnings, mw._seen_messages):
|
||||
assert key not in values
|
||||
|
||||
# Reusing even the same runtime object starts a fresh invocation budget.
|
||||
mw.before_agent(state, runtime)
|
||||
follow_up = _make_state_with_usage(total=200)
|
||||
follow_up["messages"][0].id = "next-msg"
|
||||
assert mw.after_model(follow_up, runtime) is None
|
||||
assert mw._cumulative_usage[key].total == 200
|
||||
assert not mw._warned.get(key)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_valid_run_id_preserves_usage_warnings_and_stop_reason(self):
|
||||
mw = TokenBudgetMiddleware(TokenBudgetConfig(enabled=True, max_tokens=1000))
|
||||
runtime = _make_runtime(run_id="goal-run")
|
||||
mw.after_model(_make_state_with_usage(total=850), runtime)
|
||||
state = _make_state_with_usage(total=1100)
|
||||
assert mw.after_model(state, runtime) is not None
|
||||
await mw.aafter_agent(state, runtime)
|
||||
|
||||
assert "goal-run" not in mw._seen_messages
|
||||
assert mw._cumulative_usage["goal-run"].total == 1100
|
||||
assert mw._warned["goal-run"]
|
||||
assert len(mw._pending_warnings["goal-run"]) == 1
|
||||
assert mw.consume_stop_reason("goal-run") == "token_capped"
|
||||
|
||||
# Continuations may get another Runtime object with the same run identity.
|
||||
continuation = _make_runtime(run_id="goal-run")
|
||||
await mw.abefore_agent(state, continuation)
|
||||
await mw.aafter_model(state, continuation)
|
||||
assert mw._cumulative_usage["goal-run"].total == 1100
|
||||
assert len(mw._pending_warnings["goal-run"]) == 1
|
||||
|
||||
|
||||
class TestTokenBudgetWarning:
|
||||
def test_warn_injected_at_next_model_call(self):
|
||||
config = TokenBudgetConfig(max_tokens=100000, warn_threshold=0.8, enabled=True)
|
||||
@ -196,3 +263,55 @@ class TestIndependentDimensions:
|
||||
|
||||
assert result is not None
|
||||
assert "output token" in result["messages"][0].content
|
||||
|
||||
|
||||
class _ToolCallingFakeModel(FakeMessagesListChatModel):
|
||||
def bind_tools(self, tools, *, tool_choice=None, **kwargs):
|
||||
return self
|
||||
|
||||
|
||||
class TestTokenBudgetAgentGraph:
|
||||
def test_goal_continuation_shares_the_run_budget(self):
|
||||
"""A hidden goal continuation re-enters the graph under the same run_id; it must not get a fresh budget."""
|
||||
executed: list[str] = []
|
||||
|
||||
@as_tool
|
||||
def bash(command: str) -> str:
|
||||
"""Run a fake shell command."""
|
||||
executed.append(command)
|
||||
return "ok"
|
||||
|
||||
def call(command: str, tokens: int = 4000) -> AIMessage:
|
||||
return AIMessage(
|
||||
content="",
|
||||
id=f"ai-{command}",
|
||||
tool_calls=[{"name": "bash", "id": f"call-{command}", "args": {"command": command}}],
|
||||
usage_metadata={"input_tokens": tokens, "output_tokens": 0, "total_tokens": tokens},
|
||||
)
|
||||
|
||||
model = _ToolCallingFakeModel(
|
||||
responses=[
|
||||
call("a"),
|
||||
call("b"),
|
||||
AIMessage(content="first answer", id="ai-answer-1", usage_metadata={"input_tokens": 1000, "output_tokens": 0, "total_tokens": 1000}),
|
||||
call("c"),
|
||||
call("d"),
|
||||
AIMessage(content="second answer", id="ai-answer-2", usage_metadata={"input_tokens": 1000, "output_tokens": 0, "total_tokens": 1000}),
|
||||
]
|
||||
)
|
||||
mw = TokenBudgetMiddleware(TokenBudgetConfig(enabled=True, max_tokens=10_000))
|
||||
graph = create_agent(model=model, tools=[bash], middleware=[mw], checkpointer=InMemorySaver())
|
||||
config = {"configurable": {"thread_id": "goal-thread"}}
|
||||
|
||||
# User turn: 9k of 10k.
|
||||
graph.invoke({"messages": [HumanMessage("research")]}, config=config, context={"thread_id": "goal-thread", "run_id": "run-1"})
|
||||
assert executed == ["a", "b"]
|
||||
|
||||
# Goal continuation in the same run: the next 4k call crosses the cap.
|
||||
result = graph.invoke({"messages": [HumanMessage("keep going")]}, config=config, context={"thread_id": "goal-thread", "run_id": "run-1"})
|
||||
assert executed == ["a", "b"]
|
||||
assert "TOKEN BUDGET EXCEEDED" in result["messages"][-1].content
|
||||
|
||||
# A later user run still starts with a fresh budget.
|
||||
graph.invoke({"messages": [HumanMessage("next question")]}, config=config, context={"thread_id": "goal-thread", "run_id": "run-2"})
|
||||
assert executed == ["a", "b", "d"]
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user