diff --git a/backend/packages/harness/deerflow/agents/middlewares/token_budget_middleware.py b/backend/packages/harness/deerflow/agents/middlewares/token_budget_middleware.py index a9fe9ae4d..181fa1148 100644 --- a/backend/packages/harness/deerflow/agents/middlewares/token_budget_middleware.py +++ b/backend/packages/harness/deerflow/agents/middlewares/token_budget_middleware.py @@ -96,31 +96,33 @@ class TokenBudgetMiddleware(AgentMiddleware[AgentState]): return str(id(runtime)) def _clear_run_state(self, run_id: str) -> None: - self._warned.pop(run_id, None) - self._pending_warnings.pop(run_id, None) - self._seen_messages.pop(run_id, None) - self._cumulative_usage.pop(run_id, None) + with self._lock: + self._warned.pop(run_id, None) + self._pending_warnings.pop(run_id, None) + self._seen_messages.pop(run_id, None) + self._cumulative_usage.pop(run_id, None) @override def before_agent(self, state: AgentState, runtime: Runtime) -> None: if not self._config.enabled: return - # Mark all old messages from previous runss as 'seen' so they don't count toward THIS run's budget + # Mark all old messages from previous runs as 'seen' so they don't count toward THIS run's budget messages = state.get("messages", []) if not messages: return run_id = self._get_run_id(runtime) - seen = self._seen_messages.setdefault(run_id, {}) - self._cumulative_usage.setdefault(run_id, TokenUsage()) + with self._lock: + seen = self._seen_messages.setdefault(run_id, {}) + self._cumulative_usage.setdefault(run_id, TokenUsage()) - for msg in messages: - if isinstance(msg, AIMessage) and msg.id and hasattr(msg, "usage_metadata"): - usage = msg.usage_metadata or {} - input_tokens = usage.get("input_tokens", 0) - output_tokens = usage.get("output_tokens", 0) - seen[msg.id] = (input_tokens, output_tokens) + for msg in messages: + if isinstance(msg, AIMessage) and msg.id and hasattr(msg, "usage_metadata"): + usage = msg.usage_metadata or {} + input_tokens = usage.get("input_tokens", 0) + output_tokens = usage.get("output_tokens", 0) + seen[msg.id] = (input_tokens, output_tokens) @override async def abefore_agent(self, state: AgentState, runtime: Runtime) -> None: @@ -182,68 +184,69 @@ class TokenBudgetMiddleware(AgentMiddleware[AgentState]): run_id = self._get_run_id(runtime) - seen = self._seen_messages.setdefault(run_id, {}) - usage_accum = self._cumulative_usage.setdefault(run_id, TokenUsage()) + with self._lock: + seen = self._seen_messages.setdefault(run_id, {}) + usage_accum = self._cumulative_usage.setdefault(run_id, TokenUsage()) - for msg in messages: - if isinstance(msg, AIMessage) and msg.id and hasattr(msg, "usage_metadata"): - usage = msg.usage_metadata or {} + for msg in messages: + if isinstance(msg, AIMessage) and msg.id and hasattr(msg, "usage_metadata"): + usage = msg.usage_metadata or {} - input_tokens = usage.get("input_tokens", 0) - output_tokens = usage.get("output_tokens", 0) + input_tokens = usage.get("input_tokens", 0) + output_tokens = usage.get("output_tokens", 0) - # Check what previously recorded for this exact message - prev_input, prev_output = seen.get(msg.id, (0, 0)) + # Check what previously recorded for this exact message + prev_input, prev_output = seen.get(msg.id, (0, 0)) - # Calculate if any new tokens were added (handles retroactive subagent tokens) - diff_input = max(0, input_tokens - prev_input) - diff_output = max(0, output_tokens - prev_output) + # Calculate if any new tokens were added (handles retroactive subagent tokens) + diff_input = max(0, input_tokens - prev_input) + diff_output = max(0, output_tokens - prev_output) - if diff_input > 0 or diff_output > 0: - usage_accum.input += diff_input - usage_accum.output += diff_output - usage_accum.total += diff_input + diff_output - seen[msg.id] = (input_tokens, output_tokens) + if diff_input > 0 or diff_output > 0: + usage_accum.input += diff_input + usage_accum.output += diff_output + usage_accum.total += diff_input + diff_output + seen[msg.id] = (input_tokens, output_tokens) + + if usage_accum.total <= 0: + return None + + fractions = [("total", usage_accum.total, self._config.max_tokens)] + if self._config.max_input_tokens: + fractions.append(("input", usage_accum.input, self._config.max_input_tokens)) + if self._config.max_output_tokens: + fractions.append(("output", usage_accum.output, self._config.max_output_tokens)) + + highest_fraction = 0.0 + trigger_reason = "" + trigger_used = 0 + trigger_budget = 0 + + for reason, used, limit in fractions: + frac = used / limit + if frac > highest_fraction: + highest_fraction = frac + trigger_reason = reason + trigger_used = used + trigger_budget = limit + + if highest_fraction >= self._config.hard_stop_threshold: + logger.warning("Token budget hard stop triggered for run %s: %s limit exceeded", run_id, trigger_reason) + stop_text = _BUDGET_EXCEEDED_MSG.format(reason=trigger_reason, used=trigger_used, budget=trigger_budget) + return self._build_hard_stop_update(last_msg, stop_text) + + if highest_fraction >= self._config.warn_threshold and not self._warned.get(run_id, False): + self._warned[run_id] = True + percent = highest_fraction * 100 + warn_text = _BUDGET_WARNING_MSG.format(reason=trigger_reason, used=trigger_used, budget=trigger_budget, percent=percent) + logger.info("Token budget warning triggered for run %s: %s limit at %.1f%%", run_id, trigger_reason, percent) + # queue warning for wrap_model_call + warnings = self._pending_warnings.setdefault(run_id, []) + warnings.append(warn_text) + return None - if usage_accum.total <= 0: return None - fractions = [("total", usage_accum.total, self._config.max_tokens)] - if self._config.max_input_tokens: - fractions.append(("input", usage_accum.input, self._config.max_input_tokens)) - if self._config.max_output_tokens: - fractions.append(("output", usage_accum.output, self._config.max_output_tokens)) - - highest_fraction = 0.0 - trigger_reason = "" - trigger_used = 0 - trigger_budget = 0 - - for reason, used, limit in fractions: - frac = used / limit - if frac > highest_fraction: - highest_fraction = frac - trigger_reason = reason - trigger_used = used - trigger_budget = limit - - if highest_fraction >= self._config.hard_stop_threshold: - logger.warning("Token budget hard stop triggered for run %s: %s limit exceeded", run_id, trigger_reason) - stop_text = _BUDGET_EXCEEDED_MSG.format(reason=trigger_reason, used=trigger_used, budget=trigger_budget) - return self._build_hard_stop_update(last_msg, stop_text) - - if highest_fraction >= self._config.warn_threshold and not self._warned.get(run_id, False): - self._warned[run_id] = True - percent = highest_fraction * 100 - warn_text = _BUDGET_WARNING_MSG.format(reason=trigger_reason, used=trigger_used, budget=trigger_budget, percent=percent) - logger.info("Token budget warning triggered for run %s: %s limit at %.1f%%", run_id, trigger_reason, percent) - # queue warning for wrap_model_call - warnings = self._pending_warnings.setdefault(run_id, []) - warnings.append(warn_text) - return None - - return None - @override def after_model(self, state: AgentState, runtime: Runtime) -> dict | None: return self._apply(state, runtime) @@ -257,7 +260,8 @@ class TokenBudgetMiddleware(AgentMiddleware[AgentState]): return [] run_id = self._get_run_id(runtime) - warnings = self._pending_warnings.pop(run_id, None) + with self._lock: + warnings = self._pending_warnings.pop(run_id, None) return warnings or [] def _inject_warnings(self, request: ModelRequest, warnings: list[str]) -> ModelRequest: diff --git a/frontend/src/content/zh/introduction/index.mdx b/frontend/src/content/zh/introduction/index.mdx index 8b6708d93..23392af2e 100644 --- a/frontend/src/content/zh/introduction/index.mdx +++ b/frontend/src/content/zh/introduction/index.mdx @@ -2,7 +2,6 @@ title: 简介 description: DeerFlow 是一个运行时 Harness,用于构建支持技能、记忆、工具和沙箱执行的长时序 Agent 系统。 --- -a import { Cards } from "nextra/components"; # 简介