diff --git a/src/agents/common/middlewares/summary_middleware.py b/src/agents/common/middlewares/summary_middleware.py index e36e8f9a..27052448 100644 --- a/src/agents/common/middlewares/summary_middleware.py +++ b/src/agents/common/middlewares/summary_middleware.py @@ -132,7 +132,8 @@ def _offload_tool_result(msg: ToolMessage, threshold: int, token_counter: TokenC return None # 计算 token 数 - if token_counter([msg]) <= threshold: + msg_tokens = token_counter([msg]) + if msg_tokens <= threshold: return None # 获取工具名称和参数 @@ -450,23 +451,6 @@ class SummaryOffloadMiddleware(AgentMiddleware): return result - def _should_summarize_based_on_reported_tokens(self, messages: list[AnyMessage], threshold: float) -> bool: - """Check if reported token usage from last AIMessage exceeds threshold.""" - last_ai_message = next( - (msg for msg in reversed(messages) if isinstance(msg, AIMessage)), - None, - ) - if ( # noqa: SIM103 - isinstance(last_ai_message, AIMessage) - and last_ai_message.usage_metadata is not None - and (reported_tokens := last_ai_message.usage_metadata.get("total_tokens", -1)) - and reported_tokens >= threshold - and (message_provider := last_ai_message.response_metadata.get("model_provider")) - and message_provider == self.model._get_ls_params().get("ls_provider") # noqa: SLF001 - ): - return True - return False - def _should_summarize(self, messages: list[AnyMessage], total_tokens: int) -> bool: """Determine whether summarization should run for the current token usage.""" if not self._trigger_conditions: @@ -477,8 +461,6 @@ class SummaryOffloadMiddleware(AgentMiddleware): return True if kind == "tokens" and total_tokens >= value: return True - if kind == "tokens" and self._should_summarize_based_on_reported_tokens(messages, value): - return True if kind == "fraction": max_input_tokens = self._get_profile_limits() if max_input_tokens is None: @@ -488,9 +470,6 @@ class SummaryOffloadMiddleware(AgentMiddleware): threshold = 1 if total_tokens >= threshold: return True - - if self._should_summarize_based_on_reported_tokens(messages, threshold): - return True return False def _determine_cutoff_index(self, messages: list[AnyMessage]) -> int: