From 228de7bd6780d3444a4a99f63b1243d4cfc5a632 Mon Sep 17 00:00:00 2001 From: InfinityPacer <160988576+InfinityPacer@users.noreply.github.com> Date: Thu, 13 Aug 2026 17:26:13 +0800 Subject: [PATCH] fix(agent): compact oversized final model requests (#6299) --- app/agent/__init__.py | 18 +- app/agent/middleware/summarization.py | 371 +++++++++++- tests/test_agent_background_output.py | 3 + tests/test_agent_graph_cache.py | 3 +- tests/test_agent_request_budget.py | 1 + tests/test_agent_summarization_streaming.py | 594 +++++++++++++++++++- 6 files changed, 980 insertions(+), 10 deletions(-) diff --git a/app/agent/__init__.py b/app/agent/__init__.py index 3f284e3c1..717e41c28 100644 --- a/app/agent/__init__.py +++ b/app/agent/__init__.py @@ -42,6 +42,7 @@ from app.agent.middleware.runtime_config import RuntimeConfigMiddleware from app.agent.middleware.skills import SKILL_TOOL_NAME, SkillsMiddleware from app.agent.middleware.summarization import ( ContextPreservingSummarizationMiddleware as SummarizationMiddleware, + FinalRequestCompactionMiddleware, ) from app.agent.middleware.subagents import ( SUBAGENT_CONTROL_TOOL_NAME, @@ -1944,6 +1945,12 @@ class MoviePilotAgent: if getattr(tool, "name", None) == QUERY_ACTIVITY_LOG_TOOL_NAME ) + summarization_middleware = SummarizationMiddleware( + model=non_streaming_model, + trigger=("fraction", 0.85), + keep=("messages", 20), + ) + # 中间件 middlewares = [ # 宿主策略必须位于最外层,确保插件覆盖工具基类也不能绕过。 @@ -1965,9 +1972,7 @@ class MoviePilotAgent: # 活动日志依赖记忆上下文,并应在摘要压缩前完成读取与记录。 *([activity_log_middleware] if activity_log_middleware else []), # 上下文压缩 - SummarizationMiddleware( - model=non_streaming_model, trigger=("fraction", 0.85) - ), + summarization_middleware, # 错误工具调用修复 PatchToolCallsMiddleware(), # 子代理委派 @@ -1990,6 +1995,13 @@ class MoviePilotAgent: ) ) + # 需要在动态 system 与工具筛选完成后按最终输入预算补充压缩。 + middlewares.append( + FinalRequestCompactionMiddleware( + summarizer=summarization_middleware, + ) + ) + # 预算观察器必须位于最内层,才能看到动态 system 和最终筛选后的工具。 middlewares.append( UsageMiddleware( diff --git a/app/agent/middleware/summarization.py b/app/agent/middleware/summarization.py index a3d2cd70d..995f54c21 100644 --- a/app/agent/middleware/summarization.py +++ b/app/agent/middleware/summarization.py @@ -1,8 +1,33 @@ """Agent 会话上下文压缩中间件。""" +from collections.abc import Awaitable, Callable +from importlib import import_module +from typing import Any + from langchain.agents.middleware import SummarizationMiddleware -from langchain_core.messages import AnyMessage +from langchain.agents.middleware.types import ( + AgentMiddleware, + ExtendedModelResponse, + ModelRequest, + ModelResponse, +) +from langchain_core.messages import AnyMessage, HumanMessage, RemoveMessage, ToolMessage from langchain_core.messages.utils import get_buffer_string +from langgraph.graph.message import REMOVE_ALL_MESSAGES +from langgraph.types import Command + +from app.agent.middleware.usage import UsageMiddleware +from app.log import logger + +try: + _internal_call_metadata = import_module( + "langchain.agents.middleware.internal_call_transformer" + ).internal_call_metadata +except ImportError: + + def _internal_call_metadata() -> dict[str, Any]: + """旧版 LangChain 没有内部模型调用的流式过滤标记。""" + return {} class ContextSummarizationError(RuntimeError): @@ -37,10 +62,16 @@ class ContextPreservingSummarizationMiddleware(SummarizationMiddleware): def _create_summary(self, messages_to_summarize: list[AnyMessage]) -> str: """同步摘要失败时保持原图状态。""" formatted_messages = self._prepare_summary_input(messages_to_summarize) + summary_model = getattr(self, "_summary_model", self.model) try: - response = self.model.invoke( + response = summary_model.invoke( self.summary_prompt.format(messages=formatted_messages).rstrip(), - config={"metadata": {"lc_source": "summarization"}}, + config={ + "metadata": { + "lc_source": "summarization", + **_internal_call_metadata(), + } + }, ) except Exception as err: raise ContextSummarizationError(self._ERROR_MESSAGE) from err @@ -49,11 +80,341 @@ class ContextPreservingSummarizationMiddleware(SummarizationMiddleware): async def _acreate_summary(self, messages_to_summarize: list[AnyMessage]) -> str: """异步摘要失败时保持原图状态。""" formatted_messages = self._prepare_summary_input(messages_to_summarize) + summary_model = getattr(self, "_summary_model", self.model) try: - response = await self.model.ainvoke( + response = await summary_model.ainvoke( self.summary_prompt.format(messages=formatted_messages).rstrip(), - config={"metadata": {"lc_source": "summarization"}}, + config={ + "metadata": { + "lc_source": "summarization", + **_internal_call_metadata(), + } + }, ) except Exception as err: raise ContextSummarizationError(self._ERROR_MESSAGE) from err return self._require_valid_summary(response.text.strip()) + + def partition_for_token_limit( + self, messages: list[AnyMessage], token_limit: int + ) -> tuple[list[AnyMessage], list[AnyMessage]] | None: + """按 token 上限拆分历史,并保持 LangChain 的工具调用事务边界。""" + self._ensure_message_ids(messages) + if self.token_counter(messages) <= token_limit: + return None + + left, right = 0, len(messages) + cutoff_candidate = len(messages) + while left < right: + midpoint = (left + right) // 2 + if self._partial_token_counter(messages[midpoint:]) <= token_limit: + cutoff_candidate = midpoint + right = midpoint + else: + left = midpoint + 1 + + if cutoff_candidate >= len(messages): + cutoff_candidate = len(messages) + cutoff_index = self._find_safe_cutoff_point(messages, cutoff_candidate) + if cutoff_index <= 0: + return None + return self._partition_messages(messages, cutoff_index) + + def build_summary_messages(self, summary: str) -> list[AnyMessage]: + """将摘要转换为 LangChain 约定的可识别历史消息。""" + return self._build_new_messages(summary) + + def ensure_message_ids(self, messages: list[AnyMessage]) -> None: + """为压缩后消息补齐 LangGraph reducer 所需的稳定 ID。""" + self._ensure_message_ids(messages) + + def create_summary(self, messages_to_summarize: list[AnyMessage]) -> str: + """通过 MoviePilot 的失败保护合同生成同步摘要。""" + return self._create_summary(messages_to_summarize) + + async def acreate_summary(self, messages_to_summarize: list[AnyMessage]) -> str: + """通过 MoviePilot 的失败保护合同生成异步摘要。""" + return await self._acreate_summary(messages_to_summarize) + + +class FinalRequestCompactionMiddleware(AgentMiddleware): + """按最终模型请求预算压缩历史,并在模型成功后原子提交新状态。""" + + _COMPACTION_ANCHOR_KEY = "moviepilot_compaction_anchor_id" + _UNCOMPRESSIBLE_REQUEST = ( + "最终模型请求压缩后仍超出上下文窗口,原有上下文已保留," + "请减少启用工具或切换更大上下文模型" + ) + + def __init__( + self, + *, + summarizer: ContextPreservingSummarizationMiddleware, + trigger_fraction: float = 0.85, + keep_fraction: float = 0.10, + ) -> None: + self.summarizer = summarizer + self.trigger_fraction = trigger_fraction + self.keep_fraction = keep_fraction + + def _should_compact(self, budget: dict[str, Any]) -> bool: + """以最终请求实际模型窗口判断是否需要压缩。""" + estimated_tokens = budget.get("estimated_input_tokens") + context_window = budget.get("context_window_tokens") + return ( + isinstance(estimated_tokens, int) + and isinstance(context_window, int) + and estimated_tokens >= context_window * self.trigger_fraction + ) + + def _compaction_partition( + self, request: ModelRequest + ) -> tuple[list[AnyMessage], list[AnyMessage]] | None: + """最终输入达到阈值时,拆分需要摘要和需要原样保留的消息。""" + messages = list(request.messages) + try: + budget = UsageMiddleware.estimate_request(request) + except Exception as error: + logger.debug( + "最终模型请求预算评估失败,继续原请求: error_type=%s", + type(error).__name__, + ) + return None + + context_window = budget.get("context_window_tokens") + if self._should_skip_after_current_turn_compaction(messages, budget): + return None + if not self._should_compact(budget) or not isinstance(context_window, int): + return None + try: + partition = self.summarizer.partition_for_token_limit( + messages, + max(1, int(context_window * self.keep_fraction)), + ) + except Exception as error: + logger.debug( + "最终请求历史拆分失败,继续原请求: error_type=%s", + type(error).__name__, + ) + if budget["estimated_input_tokens"] > context_window: + raise ContextSummarizationError( + self._UNCOMPRESSIBLE_REQUEST + ) from error + return None + if partition is None: + if budget["estimated_input_tokens"] > context_window: + raise ContextSummarizationError(self._UNCOMPRESSIBLE_REQUEST) + return None + messages_to_summarize, preserved_messages = partition + if all( + message.additional_kwargs.get("lc_source") == "summarization" + for message in messages_to_summarize + ): + if budget["estimated_input_tokens"] > context_window: + raise ContextSummarizationError(self._UNCOMPRESSIBLE_REQUEST) + return None + return messages_to_summarize, preserved_messages + + @classmethod + def _should_skip_after_current_turn_compaction( + cls, messages: list[AnyMessage], budget: dict[str, Any] + ) -> bool: + """同轮只在新工具结果已使请求超窗时再次压缩。""" + for message in reversed(messages): + anchor_id = message.additional_kwargs.get(cls._COMPACTION_ANCHOR_KEY) + if not isinstance(anchor_id, str): + continue + anchor_index = next( + ( + index + for index, candidate in enumerate(messages) + if candidate.id == anchor_id + ), + None, + ) + if anchor_index is None: + return False + messages_after_anchor = messages[anchor_index + 1 :] + if any( + isinstance(candidate, HumanMessage) + for candidate in messages_after_anchor + ): + return False + if any(isinstance(candidate, ToolMessage) for candidate in messages_after_anchor): + estimated_tokens = budget.get("estimated_input_tokens") + context_window = budget.get("context_window_tokens") + return not ( + isinstance(estimated_tokens, int) + and isinstance(context_window, int) + and estimated_tokens > context_window + ) + return True + return False + + def _build_compacted_messages( + self, summary: str, preserved_messages: list[AnyMessage] + ) -> list[AnyMessage]: + """构造摘要与近期历史,并记录本轮压缩输入边界。""" + summary_messages = self.summarizer.build_summary_messages(summary) + if summary_messages: + self.summarizer.ensure_message_ids(summary_messages) + anchor_id = ( + preserved_messages[-1].id + if preserved_messages + else summary_messages[0].id + ) + first_summary = summary_messages[0] + summary_messages[0] = first_summary.model_copy( + update={ + "additional_kwargs": { + **first_summary.additional_kwargs, + self._COMPACTION_ANCHOR_KEY: anchor_id, + } + } + ) + return [*summary_messages, *preserved_messages] + + def _validate_or_repartition( + self, + request: ModelRequest, + messages_to_summarize: list[AnyMessage], + preserved_messages: list[AnyMessage], + summary: str, + ) -> tuple[list[AnyMessage], tuple[list[AnyMessage], list[AnyMessage]] | None]: + """复核压缩后的最终预算,并计算一次更小的近期历史分区。""" + compacted_messages = self._build_compacted_messages(summary, preserved_messages) + compacted_budget = UsageMiddleware.estimate_request( + request.override(messages=compacted_messages) + ) + context_window = compacted_budget.get("context_window_tokens") + estimated_tokens = compacted_budget.get("estimated_input_tokens") + if not isinstance(context_window, int) or not isinstance(estimated_tokens, int): + return compacted_messages, None + + target_tokens = max(1, int(context_window * self.trigger_fraction)) + if estimated_tokens <= target_tokens: + return compacted_messages, None + + summary_messages = self.summarizer.build_summary_messages(summary) + fixed_summary_budget = UsageMiddleware.estimate_request( + request.override(messages=summary_messages) + ) + available_recent_tokens = ( + target_tokens - fixed_summary_budget["estimated_input_tokens"] + ) + repartition = self.summarizer.partition_for_token_limit( + list(request.messages), max(1, available_recent_tokens) + ) + if ( + available_recent_tokens <= 0 + or repartition is None + or len(repartition[0]) <= len(messages_to_summarize) + ): + if estimated_tokens > context_window: + raise ContextSummarizationError(self._UNCOMPRESSIBLE_REQUEST) + return compacted_messages, None + return compacted_messages, repartition + + def _require_within_window( + self, request: ModelRequest, compacted_messages: list[AnyMessage] + ) -> None: + """禁止把已知仍超过主模型窗口的请求发送给 provider。""" + budget = UsageMiddleware.estimate_request( + request.override(messages=compacted_messages) + ) + estimated_tokens = budget.get("estimated_input_tokens") + context_window = budget.get("context_window_tokens") + if ( + isinstance(estimated_tokens, int) + and isinstance(context_window, int) + and estimated_tokens > context_window + ): + raise ContextSummarizationError(self._UNCOMPRESSIBLE_REQUEST) + + def _prepare_messages(self, request: ModelRequest) -> list[AnyMessage] | None: + """同步生成摘要与需要原样保留的近期消息。""" + partition = self._compaction_partition(request) + if partition is None: + return None + messages_to_summarize, preserved_messages = partition + summary = self.summarizer.create_summary(messages_to_summarize) + compacted_messages, repartition = self._validate_or_repartition( + request, + messages_to_summarize, + preserved_messages, + summary, + ) + if repartition is not None: + messages_to_summarize, preserved_messages = repartition + summary = self.summarizer.create_summary(messages_to_summarize) + compacted_messages = self._build_compacted_messages( + summary, preserved_messages + ) + self._require_within_window(request, compacted_messages) + return compacted_messages + + async def _aprepare_messages( + self, request: ModelRequest + ) -> list[AnyMessage] | None: + """异步生成摘要与需要原样保留的近期消息。""" + partition = self._compaction_partition(request) + if partition is None: + return None + messages_to_summarize, preserved_messages = partition + summary = await self.summarizer.acreate_summary(messages_to_summarize) + compacted_messages, repartition = self._validate_or_repartition( + request, + messages_to_summarize, + preserved_messages, + summary, + ) + if repartition is not None: + messages_to_summarize, preserved_messages = repartition + summary = await self.summarizer.acreate_summary(messages_to_summarize) + compacted_messages = self._build_compacted_messages( + summary, preserved_messages + ) + self._require_within_window(request, compacted_messages) + return compacted_messages + + @staticmethod + def _with_state_update( + response: ModelResponse, compacted_messages: list[AnyMessage] + ) -> ExtendedModelResponse: + """主模型成功后一次性替换历史,同时保留本次模型结果。""" + return ExtendedModelResponse( + model_response=response, + command=Command( + update={ + "messages": [ + RemoveMessage(id=REMOVE_ALL_MESSAGES), + *compacted_messages, + *response.result, + ] + } + ), + ) + + def wrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + ) -> ModelResponse | ExtendedModelResponse: + """同步压缩最终请求;模型失败时不提交摘要状态。""" + compacted_messages = self._prepare_messages(request) + if compacted_messages is None: + return handler(request) + response = handler(request.override(messages=compacted_messages)) + return self._with_state_update(response, compacted_messages) + + async def awrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + ) -> ModelResponse | ExtendedModelResponse: + """异步压缩最终请求;模型失败时不提交摘要状态。""" + compacted_messages = await self._aprepare_messages(request) + if compacted_messages is None: + return await handler(request) + response = await handler(request.override(messages=compacted_messages)) + return self._with_state_update(response, compacted_messages) diff --git a/tests/test_agent_background_output.py b/tests/test_agent_background_output.py index a4f4a1809..d38868144 100644 --- a/tests/test_agent_background_output.py +++ b/tests/test_agent_background_output.py @@ -438,6 +438,7 @@ class AgentBackgroundOutputTest(unittest.IsolatedAsyncioTestCase): "memory", "summary", "patch", + "FinalRequestCompactionMiddleware", "usage", ], [getattr(item, "name", item) for item in created["middleware"]], @@ -554,6 +555,7 @@ class AgentBackgroundOutputTest(unittest.IsolatedAsyncioTestCase): "memory", "summary", "patch", + "FinalRequestCompactionMiddleware", "usage", ], [getattr(item, "name", item) for item in created["middleware"]], @@ -759,6 +761,7 @@ class AgentBackgroundOutputTest(unittest.IsolatedAsyncioTestCase): "activity", "summary", "patch", + "FinalRequestCompactionMiddleware", "usage", ], [getattr(item, "name", item) for item in created["middleware"]], diff --git a/tests/test_agent_graph_cache.py b/tests/test_agent_graph_cache.py index d4b003169..2d174f767 100644 --- a/tests/test_agent_graph_cache.py +++ b/tests/test_agent_graph_cache.py @@ -491,9 +491,10 @@ async def test_graph_keeps_mcp_first_winner_and_catalogs_all_collisions( activity_tool, subagent_task_tool, ] - assert captured["middlewares"][-2].name == "selector" + assert captured["middlewares"][-3].name == "selector" else: assert "selection_tools" not in captured + assert captured["middlewares"][-2].name == "FinalRequestCompactionMiddleware" assert captured["middlewares"][-1].name == "usage" policy_middleware = next( middleware diff --git a/tests/test_agent_request_budget.py b/tests/test_agent_request_budget.py index 1a9124543..47b5b1465 100644 --- a/tests/test_agent_request_budget.py +++ b/tests/test_agent_request_budget.py @@ -81,6 +81,7 @@ def test_final_request_can_exceed_window_before_message_fraction_triggers(): _llm_type="test-chat", profile={"max_input_tokens": 4096}, ) + model.with_retry = lambda: model summarizer = SummarizationMiddleware( model=model, trigger=("fraction", 0.85), diff --git a/tests/test_agent_summarization_streaming.py b/tests/test_agent_summarization_streaming.py index c38c3e46f..21fc67e00 100644 --- a/tests/test_agent_summarization_streaming.py +++ b/tests/test_agent_summarization_streaming.py @@ -4,7 +4,23 @@ from types import SimpleNamespace from unittest.mock import AsyncMock, patch import pytest -from langchain_core.messages import AIMessage, HumanMessage +from langchain.agents import create_agent +from langchain.agents.middleware import AgentMiddleware +from langchain.agents.middleware.types import ( + ExtendedModelResponse, + ModelRequest, + ModelResponse, +) +from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel +from langchain_core.messages import ( + AIMessage, + AnyMessage, + HumanMessage, + SystemMessage, +) +from langchain_core.tools import tool +from langgraph.checkpoint.memory import InMemorySaver +from pydantic import Field import app.agent as agent_module from app.agent.memory import memory_manager @@ -12,7 +28,9 @@ from app.agent.middleware.runtime_config import RuntimeConfigMiddleware from app.agent.middleware.summarization import ( ContextSummarizationError, ContextPreservingSummarizationMiddleware, + FinalRequestCompactionMiddleware, ) +from app.agent.middleware.usage import UsageMiddleware class _FakeLLM: @@ -22,6 +40,10 @@ class _FakeLLM: self.model = model self.profile = {"max_input_tokens": 64000} + def with_retry(self): + """满足新版 LangChain 摘要模型的 Runnable 合同。""" + return self + class _FailingSummaryLLM(_FakeLLM): """模拟摘要模型暂时不可用。""" @@ -47,6 +69,131 @@ class _SuccessfulSummaryLLM(_FakeLLM): return AIMessage(content="保留旧事实的摘要") +class _CountingSummaryLLM(_SuccessfulSummaryLLM): + """记录真实 Agent 图触发的摘要次数。""" + + def __init__(self, model: str): + super().__init__(model) + self.calls = 0 + + async def ainvoke(self, *_args, **_kwargs): + """记录异步摘要请求。""" + self.calls += 1 + return AIMessage(content="保留旧事实的摘要") + + def invoke(self, *_args, **_kwargs): + """记录同步摘要请求。""" + self.calls += 1 + return AIMessage(content="保留旧事实的摘要") + + +class _LongSummaryLLM(_SuccessfulSummaryLLM): + """返回本身无法装入主模型窗口的摘要。""" + + async def ainvoke(self, *_args, **_kwargs): + """返回异常超长摘要。""" + return AIMessage(content="异常冗长摘要 " * 2000) + + def invoke(self, *_args, **_kwargs): + """返回异常超长摘要。""" + return AIMessage(content="异常冗长摘要 " * 2000) + + +class _MetadataRecordingSummaryLLM(_SuccessfulSummaryLLM): + """记录摘要内部模型调用使用的 metadata。""" + + def __init__(self, model: str): + super().__init__(model) + self.configs = [] + + async def ainvoke(self, *_args, **kwargs): + """记录异步摘要调用配置。""" + self.configs.append(kwargs.get("config")) + return AIMessage(content="保留旧事实的摘要") + + def invoke(self, *_args, **kwargs): + """记录同步摘要调用配置。""" + self.configs.append(kwargs.get("config")) + return AIMessage(content="保留旧事实的摘要") + +class _RecordingChatModel(FakeMessagesListChatModel): + """记录主模型实际收到的最终请求。""" + + seen_messages: list[list[AnyMessage]] = Field(default_factory=list) + + def bind_tools(self, _tools, **_kwargs): + """保留测试模型,同时满足工具绑定契约。""" + return self + + def _generate(self, messages, *args, **kwargs): + """记录包含 system 的最终消息序列。""" + self.seen_messages.append(list(messages)) + return super()._generate(messages, *args, **kwargs) + + +class _FailingMainModel(_RecordingChatModel): + """模拟最终请求进入主模型后失败。""" + + def _generate(self, messages, *_args, **_kwargs): + """记录请求后终止模型调用。""" + self.seen_messages.append(list(messages)) + raise TimeoutError("main provider unavailable") + + +class _DynamicSystemMiddleware(AgentMiddleware): + """模拟运行时中间件追加的大型 system prompt。""" + + async def awrap_model_call(self, request, handler): + """在最终压缩器之前补充动态 system prompt。""" + return await handler( + request.override( + system_message=SystemMessage(content="动态系统约束 " * 250) + ) + ) + + +def _final_request(*, messages, system_message=None, tools=None) -> ModelRequest: + """构造包含最终系统提示词和工具目录的模型请求。""" + return ModelRequest( + model=SimpleNamespace( + model="small-model", + profile={"max_input_tokens": 2048}, + ), + messages=list(messages), + system_message=system_message, + tools=list(tools or []), + state={"messages": list(messages)}, + runtime=None, + ) + + +def _real_compaction_graph(*, model, summarizer, tools=None, checkpointer=None): + """构造包含真实 LangChain 状态归并路径的最小 Agent 图。""" + summary_middleware = ContextPreservingSummarizationMiddleware( + model=summarizer, + trigger=("fraction", 0.85), + keep=("messages", 20), + ) + return create_agent( + model=model, + tools=list(tools or []), + middleware=[ + summary_middleware, + _DynamicSystemMiddleware(), + FinalRequestCompactionMiddleware(summarizer=summary_middleware), + ], + checkpointer=checkpointer, + ) + + +def _oversized_final_request_history() -> list[HumanMessage]: + """生成历史本身未达阈值、叠加动态 system 后超阈值的消息。""" + return [ + HumanMessage(content=(f"必须保留的历史事实 {index} " * 80)) + for index in range(6) + ] + + class _FailingGraph: """模拟上下文压缩阶段失败的 Agent 图。""" @@ -278,6 +425,451 @@ def test_summary_success_still_replaces_old_context(): assert len(update["messages"]) < len(messages) +def test_final_request_compaction_includes_dynamic_system_and_tools(): + """最终 system 和工具预算达到阈值时,应在同轮压缩历史。""" + summarizer = ContextPreservingSummarizationMiddleware( + model=_SuccessfulSummaryLLM("summary"), + trigger=("fraction", 0.85), + keep=("fraction", 0.10), + ) + middleware = FinalRequestCompactionMiddleware(summarizer=summarizer) + messages = [ + HumanMessage(content=f"必须保留的历史事实 {index} " * 120) + for index in range(6) + ] + request = _final_request( + messages=messages, + system_message=SystemMessage(content="动态系统约束 " * 40), + tools=[ + { + "type": "function", + "function": { + "name": "large_tool", + "description": "工具业务说明 " * 100, + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + ) + received = [] + + async def _handler(compacted_request): + received.append(compacted_request) + return ModelResponse(result=[AIMessage(content="继续完成")]) + + result = asyncio.run(middleware.awrap_model_call(request, _handler)) + + assert isinstance(result, ExtendedModelResponse) + assert len(received) == 1 + assert len(received[0].messages) < len(messages) + assert "保留旧事实的摘要" in received[0].messages[0].content + compacted_budget = UsageMiddleware.estimate_request(received[0]) + assert compacted_budget["estimated_input_ratio"] <= 0.85 + assert result.command is not None + assert "保留旧事实的摘要" in result.command.update["messages"][1].content + assert result.command.update["messages"][-1].content == "继续完成" + + +def test_final_request_compaction_preserves_history_when_summary_fails(): + """动态压缩失败时中止本轮,不提交摘要或调用主模型。""" + summarizer = ContextPreservingSummarizationMiddleware( + model=_FailingSummaryLLM("summary"), + trigger=("fraction", 0.85), + keep=("fraction", 0.10), + ) + middleware = FinalRequestCompactionMiddleware(summarizer=summarizer) + messages = [ + HumanMessage(content=f"必须保留的历史事实 {index} " * 120) + for index in range(6) + ] + request = _final_request( + messages=messages, + system_message=SystemMessage(content="动态系统约束 " * 80), + tools=[ + { + "type": "function", + "function": { + "name": "large_tool", + "description": "工具业务说明 " * 300, + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + ) + received = [] + + async def _handler(original_request): + received.append(original_request) + return ModelResponse(result=[AIMessage(content="不应执行")]) + + with pytest.raises(ContextSummarizationError, match="会话上下文压缩失败"): + asyncio.run(middleware.awrap_model_call(request, _handler)) + + assert received == [] + + +@pytest.mark.parametrize("failure", ["fixed-overhead", "long-summary"]) +def test_final_request_compaction_rejects_known_overflow_before_main_model(failure): + """压缩后仍已知超窗时不得把请求发送给主模型。""" + summary_model = ( + _LongSummaryLLM("summary") + if failure == "long-summary" + else _SuccessfulSummaryLLM("summary") + ) + summarizer = ContextPreservingSummarizationMiddleware( + model=summary_model, + trigger=("fraction", 0.85), + keep=("fraction", 0.10), + ) + middleware = FinalRequestCompactionMiddleware(summarizer=summarizer) + request = _final_request( + messages=_oversized_final_request_history(), + system_message=SystemMessage( + content=( + "不可缩减的系统约束 " * 1500 + if failure == "fixed-overhead" + else "动态系统约束 " * 250 + ) + ), + ) + received = [] + + async def _handler(compacted_request): + received.append(compacted_request) + return ModelResponse(result=[AIMessage(content="不应执行")]) + + with pytest.raises(ContextSummarizationError, match="仍超出上下文窗口"): + asyncio.run(middleware.awrap_model_call(request, _handler)) + + assert received == [] + + +def test_uncompactable_request_below_window_still_calls_main_model(): + """主动压缩线不是硬拒绝线,窗口内请求应保持可用。""" + summarizer = ContextPreservingSummarizationMiddleware( + model=_SuccessfulSummaryLLM("summary"), + trigger=("fraction", 0.85), + keep=("fraction", 0.10), + ) + middleware = FinalRequestCompactionMiddleware(summarizer=summarizer) + request = _final_request( + messages=[HumanMessage(content="最新问题")], + system_message=SystemMessage(content="不可缩减的系统约束 " * 750), + ) + budget = UsageMiddleware.estimate_request(request) + assert 0.85 <= budget["estimated_input_ratio"] <= 1 + received = [] + + async def _handler(original_request): + received.append(original_request) + return ModelResponse(result=[AIMessage(content="继续完成")]) + + result = asyncio.run(middleware.awrap_model_call(request, _handler)) + + assert isinstance(result, ModelResponse) + assert received == [request] + + +def test_summary_calls_include_version_compatible_internal_metadata(): + """摘要调用应合并当前 LangChain 提供的内部流式过滤标记。""" + summary_model = _MetadataRecordingSummaryLLM("summary") + middleware = ContextPreservingSummarizationMiddleware( + model=summary_model, + trigger=("messages", 2), + ) + messages = [HumanMessage(content="旧消息"), HumanMessage(content="新消息")] + + with patch( + "app.agent.middleware.summarization._internal_call_metadata", + return_value={"lc_internal_call": "process-marker"}, + ): + middleware.create_summary(messages) + asyncio.run(middleware.acreate_summary(messages)) + + assert [config["metadata"] for config in summary_model.configs] == [ + {"lc_source": "summarization", "lc_internal_call": "process-marker"}, + {"lc_source": "summarization", "lc_internal_call": "process-marker"}, + ] + + +@pytest.mark.parametrize("execution", ["ainvoke", "astream"]) +def test_real_agent_commits_final_request_compaction(execution): + """真实图在普通和流式执行中应提交相同的压缩后最终状态。""" + summarizer = _CountingSummaryLLM("summary") + model = _RecordingChatModel( + responses=[AIMessage(content="继续完成")], + profile={"max_input_tokens": 2048}, + ) + checkpointer = InMemorySaver() + graph = _real_compaction_graph( + model=model, + summarizer=summarizer, + checkpointer=checkpointer, + ) + messages = _oversized_final_request_history() + config = {"configurable": {"thread_id": f"compaction-{execution}"}} + + async def _run(): + if execution == "ainvoke": + result = await graph.ainvoke({"messages": messages}, config=config) + return result, (await graph.aget_state(config)).values + final_state = None + async for state in graph.astream( + {"messages": messages}, config=config, stream_mode="values" + ): + final_state = state + return final_state, (await graph.aget_state(config)).values + + result, persisted = asyncio.run(_run()) + + assert result is not None + assert [message.content for message in result["messages"]] == [ + message.content for message in persisted["messages"] + ] + assert summarizer.calls == 1 + assert len(model.seen_messages) == 1 + assert isinstance(model.seen_messages[0][0], SystemMessage) + assert "保留旧事实的摘要" in model.seen_messages[0][1].content + assert "保留旧事实的摘要" in result["messages"][0].content + assert result["messages"][-1].content == "继续完成" + assert len(result["messages"]) < len(messages) + + +def test_real_agent_does_not_compact_request_below_threshold(): + """最终请求低于主模型阈值时不得调用摘要模型。""" + summarizer = _CountingSummaryLLM("summary") + model = _RecordingChatModel( + responses=[AIMessage(content="直接完成")], + profile={"max_input_tokens": 8192}, + ) + graph = _real_compaction_graph(model=model, summarizer=summarizer) + messages = _oversized_final_request_history() + + result = asyncio.run(graph.ainvoke({"messages": messages})) + + assert summarizer.calls == 0 + assert len(model.seen_messages[0]) == len(messages) + 1 + assert [message.content for message in result["messages"][:-1]] == [ + message.content for message in messages + ] + + +def test_real_agent_executes_compacted_tool_call_once(): + """压缩不得重试主模型或重复执行工具事务。""" + calls = [] + + @tool + def record_value(value: str) -> str: + """记录工具调用次数。""" + calls.append(value) + return value + + summarizer = _CountingSummaryLLM("summary") + model = _RecordingChatModel( + responses=[ + AIMessage( + content="", + tool_calls=[ + {"name": "record_value", "args": {"value": "once"}, "id": "call-1"} + ], + ), + AIMessage(content="工具完成"), + ], + profile={"max_input_tokens": 2048}, + ) + graph = _real_compaction_graph( + model=model, + summarizer=summarizer, + tools=[record_value], + ) + + result = asyncio.run( + graph.ainvoke({"messages": _oversized_final_request_history()}) + ) + + assert calls == ["once"] + assert summarizer.calls == 1 + assert len(model.seen_messages) == 2 + assert result["messages"][-1].content == "工具完成" + + +def test_real_agent_does_not_recompact_small_tool_result_during_same_loop(): + """小工具结果不会让同一轮请求重新压缩。""" + calls = [] + + @tool + def record_value(value: str) -> str: + """记录工具调用次数。""" + calls.append(value) + return value + + summarizer = _CountingSummaryLLM("summary") + model = _RecordingChatModel( + responses=[ + AIMessage( + content="", + tool_calls=[ + {"name": "record_value", "args": {"value": "once"}, "id": "call-1"} + ], + ), + AIMessage(content="工具完成"), + ], + profile={"max_input_tokens": 2048}, + ) + graph = _real_compaction_graph( + model=model, + summarizer=summarizer, + tools=[record_value], + ) + + result = asyncio.run( + graph.ainvoke({"messages": _oversized_final_request_history()}) + ) + + assert calls == ["once"] + assert summarizer.calls == 1 + assert len(model.seen_messages) == 2 + assert result["messages"][-1].content == "工具完成" + + +def test_large_tool_result_allows_same_turn_recompaction(): + """新工具结果使请求超窗时,同轮 anchor 不得阻止再次压缩。""" + calls = [] + + @tool + def large_result() -> str: + """返回足以再次耗尽主模型窗口的工具结果。""" + calls.append("once") + return "超长工具结果 " * 2000 + + summarizer = _CountingSummaryLLM("summary") + model = _RecordingChatModel( + responses=[ + AIMessage( + content="", + tool_calls=[{"name": "large_result", "args": {}, "id": "large-call"}], + ), + AIMessage(content="工具完成"), + ], + profile={"max_input_tokens": 2048}, + ) + graph = _real_compaction_graph( + model=model, + summarizer=summarizer, + tools=[large_result], + ) + + result = asyncio.run( + graph.ainvoke({"messages": _oversized_final_request_history()}) + ) + + assert calls == ["once"] + assert summarizer.calls == 2 + assert len(model.seen_messages) == 2 + assert "超长工具结果" not in str(model.seen_messages[1]) + assert result["messages"][-1].content == "工具完成" + + +def test_real_agent_can_compact_again_after_new_user_message(): + """同轮保护不得阻止后续用户轮次继续滚动压缩。""" + summarizer = _CountingSummaryLLM("summary") + model = _RecordingChatModel( + responses=[AIMessage(content="第一轮完成"), AIMessage(content="第二轮完成")], + profile={"max_input_tokens": 1024}, + ) + checkpointer = InMemorySaver() + graph = _real_compaction_graph( + model=model, + summarizer=summarizer, + checkpointer=checkpointer, + ) + config = {"configurable": {"thread_id": "compaction-next-turn"}} + + async def _run(): + await graph.ainvoke( + {"messages": _oversized_final_request_history()}, + config=config, + ) + return await graph.ainvoke( + {"messages": [HumanMessage(content="新的用户问题 " * 300)]}, + config=config, + ) + + result = asyncio.run(_run()) + + assert summarizer.calls == 2 + assert len(model.seen_messages) == 2 + assert result["messages"][-1].content == "第二轮完成" + + +def test_real_agent_does_not_resummarize_existing_summary_only(): + """可移除历史只有既有摘要时,不应反复摘要同一内容。""" + summarizer = _CountingSummaryLLM("summary") + model = _RecordingChatModel( + responses=[AIMessage(content="继续完成")], + profile={"max_input_tokens": 2048}, + ) + graph = _real_compaction_graph(model=model, summarizer=summarizer) + messages = [ + HumanMessage( + content="已有摘要 " * 400, + additional_kwargs={"lc_source": "summarization"}, + ), + HumanMessage(content="最新问题"), + ] + + result = asyncio.run(graph.ainvoke({"messages": messages})) + + assert summarizer.calls == 0 + assert "已有摘要" in model.seen_messages[0][1].content + assert result["messages"][-1].content == "继续完成" + + +@pytest.mark.parametrize("failure", ["summary", "main"]) +def test_real_agent_failure_does_not_commit_compacted_history(failure): + """摘要或主模型失败时,checkpoint 只保留原始历史。""" + checkpointer = InMemorySaver() + summarizer = ( + _FailingSummaryLLM("summary") + if failure == "summary" + else _CountingSummaryLLM("summary") + ) + model = ( + _FailingMainModel( + responses=[AIMessage(content="不会返回")], + profile={"max_input_tokens": 2048}, + ) + if failure == "main" + else _RecordingChatModel( + responses=[AIMessage(content="不会调用")], + profile={"max_input_tokens": 2048}, + ) + ) + graph = _real_compaction_graph( + model=model, + summarizer=summarizer, + checkpointer=checkpointer, + ) + messages = _oversized_final_request_history() + config = {"configurable": {"thread_id": f"compaction-{failure}"}} + + async def _run(): + with pytest.raises((ContextSummarizationError, TimeoutError)): + await graph.ainvoke({"messages": messages}, config=config) + return await graph.aget_state(config) + + snapshot = asyncio.run(_run()) + + assert [message.content for message in snapshot.values["messages"]] == [ + message.content for message in messages + ] + if failure == "summary": + assert model.seen_messages == [] + else: + assert summarizer.calls == 1 + assert len(model.seen_messages) == 1 + + def test_unsummarizable_message_requires_new_context_instead_of_retry(): """确定性不可裁剪的消息应给出可前进路径,而非建议无效重试。""" middleware = ContextPreservingSummarizationMiddleware(