diff --git a/app/agent/__init__.py b/app/agent/__init__.py index 717e41c28..95df59048 100644 --- a/app/agent/__init__.py +++ b/app/agent/__init__.py @@ -1969,10 +1969,8 @@ class MoviePilotAgent: RuntimeConfigMiddleware(), # 记忆管理 MemoryMiddleware(memory_dir=str(agent_runtime_manager.memory_dir)), - # 活动日志依赖记忆上下文,并应在摘要压缩前完成读取与记录。 + # 活动日志依赖记忆上下文,并应在最终请求压缩前完成读取与记录。 *([activity_log_middleware] if activity_log_middleware else []), - # 上下文压缩 - summarization_middleware, # 错误工具调用修复 PatchToolCallsMiddleware(), # 子代理委派 @@ -1995,7 +1993,7 @@ class MoviePilotAgent: ) ) - # 需要在动态 system 与工具筛选完成后按最终输入预算补充压缩。 + # 所有压缩都在最终请求边界完成,避免主模型失败前写入摘要状态。 middlewares.append( FinalRequestCompactionMiddleware( summarizer=summarization_middleware, diff --git a/app/agent/middleware/summarization.py b/app/agent/middleware/summarization.py index 995f54c21..3d727b39a 100644 --- a/app/agent/middleware/summarization.py +++ b/app/agent/middleware/summarization.py @@ -96,12 +96,26 @@ class ContextPreservingSummarizationMiddleware(SummarizationMiddleware): return self._require_valid_summary(response.text.strip()) def partition_for_token_limit( - self, messages: list[AnyMessage], token_limit: int + self, + messages: list[AnyMessage], + token_limit: int, + *, + force: bool = False, + minimum_cutoff: int = 1, + strict_token_limit: bool = False, ) -> tuple[list[AnyMessage], list[AnyMessage]] | None: """按 token 上限拆分历史,并保持 LangChain 的工具调用事务边界。""" self._ensure_message_ids(messages) if self.token_counter(messages) <= token_limit: - return None + return ( + self._minimum_safe_partition( + messages, + minimum_cutoff=minimum_cutoff, + token_limit=token_limit if strict_token_limit else None, + ) + if force + else None + ) left, right = 0, len(messages) cutoff_candidate = len(messages) @@ -116,10 +130,89 @@ class ContextPreservingSummarizationMiddleware(SummarizationMiddleware): if cutoff_candidate >= len(messages): cutoff_candidate = len(messages) cutoff_index = self._find_safe_cutoff_point(messages, cutoff_candidate) + if ( + cutoff_index <= 0 + or cutoff_index >= len(messages) + or cutoff_index < minimum_cutoff + or not self._contains_unsummarized_message(messages[:cutoff_index]) + or ( + strict_token_limit + and self._partial_token_counter(messages[cutoff_index:]) > token_limit + ) + ): + if cutoff_candidate >= len(messages): + if strict_token_limit: + return None + return self._latest_safe_partition( + messages, + minimum_cutoff=minimum_cutoff, + ) + return self._minimum_safe_partition( + messages, + minimum_cutoff=max(minimum_cutoff, cutoff_candidate), + token_limit=token_limit if strict_token_limit else None, + ) + return self._partition_messages(messages, cutoff_index) + + def partition_for_retention( + self, messages: list[AnyMessage] + ) -> tuple[list[AnyMessage], list[AnyMessage]] | None: + """按摘要器既有触发和保留策略拆分历史。""" + self._ensure_message_ids(messages) + total_tokens = self.token_counter(messages) + if not self._should_summarize(messages, total_tokens): + return None + cutoff_index = self._determine_cutoff_index(messages) if cutoff_index <= 0: return None return self._partition_messages(messages, cutoff_index) + def _minimum_safe_partition( + self, + messages: list[AnyMessage], + *, + minimum_cutoff: int = 1, + token_limit: int | None = None, + ) -> tuple[list[AnyMessage], list[AnyMessage]] | None: + """至少摘要一段旧历史,同时保留最新完整消息事务。""" + for candidate in range(max(1, minimum_cutoff), len(messages)): + cutoff_index = self._find_safe_cutoff_point(messages, candidate) + if ( + minimum_cutoff <= cutoff_index < len(messages) + and self._contains_unsummarized_message(messages[:cutoff_index]) + and ( + token_limit is None + or self._partial_token_counter(messages[cutoff_index:]) + <= token_limit + ) + ): + return self._partition_messages(messages, cutoff_index) + return None + + def _latest_safe_partition( + self, + messages: list[AnyMessage], + *, + minimum_cutoff: int, + ) -> tuple[list[AnyMessage], list[AnyMessage]] | None: + """保留无法满足软预算时的最新完整消息事务。""" + for candidate in range(len(messages) - 1, minimum_cutoff - 1, -1): + cutoff_index = self._find_safe_cutoff_point(messages, candidate) + if ( + minimum_cutoff <= cutoff_index < len(messages) + and self._contains_unsummarized_message(messages[:cutoff_index]) + ): + return self._partition_messages(messages, cutoff_index) + return None + + @staticmethod + def _contains_unsummarized_message(messages: list[AnyMessage]) -> bool: + """确认待摘要段包含可推进上下文的原始消息。""" + return any( + message.additional_kwargs.get("lc_source") != "summarization" + for message in messages + ) + def build_summary_messages(self, summary: str) -> list[AnyMessage]: """将摘要转换为 LangChain 约定的可识别历史消息。""" return self._build_new_messages(summary) @@ -187,10 +280,13 @@ class FinalRequestCompactionMiddleware(AgentMiddleware): 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)), - ) + partition = self.summarizer.partition_for_retention(messages) + if partition is None: + partition = self.summarizer.partition_for_token_limit( + messages, + max(1, int(context_window * self.keep_fraction)), + force=budget["estimated_input_tokens"] > context_window, + ) except Exception as error: logger.debug( "最终请求历史拆分失败,继续原请求: error_type=%s", @@ -291,8 +387,7 @@ class FinalRequestCompactionMiddleware(AgentMiddleware): 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: + if estimated_tokens <= context_window: return compacted_messages, None summary_messages = self.summarizer.build_summary_messages(summary) @@ -300,19 +395,22 @@ class FinalRequestCompactionMiddleware(AgentMiddleware): request.override(messages=summary_messages) ) available_recent_tokens = ( - target_tokens - fixed_summary_budget["estimated_input_tokens"] + context_window - fixed_summary_budget["estimated_input_tokens"] ) + if available_recent_tokens <= 0: + raise ContextSummarizationError(self._UNCOMPRESSIBLE_REQUEST) repartition = self.summarizer.partition_for_token_limit( - list(request.messages), max(1, available_recent_tokens) + list(request.messages), + available_recent_tokens, + force=True, + minimum_cutoff=len(messages_to_summarize) + 1, + strict_token_limit=True, ) if ( - available_recent_tokens <= 0 - or repartition is None + 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 + raise ContextSummarizationError(self._UNCOMPRESSIBLE_REQUEST) return compacted_messages, repartition def _require_within_window( diff --git a/tests/test_agent_background_output.py b/tests/test_agent_background_output.py index d38868144..f77142a72 100644 --- a/tests/test_agent_background_output.py +++ b/tests/test_agent_background_output.py @@ -436,7 +436,6 @@ class AgentBackgroundOutputTest(unittest.IsolatedAsyncioTestCase): "jobs", "runtime", "memory", - "summary", "patch", "FinalRequestCompactionMiddleware", "usage", @@ -553,7 +552,6 @@ class AgentBackgroundOutputTest(unittest.IsolatedAsyncioTestCase): "jobs", "runtime", "memory", - "summary", "patch", "FinalRequestCompactionMiddleware", "usage", @@ -759,7 +757,6 @@ class AgentBackgroundOutputTest(unittest.IsolatedAsyncioTestCase): "runtime", "memory", "activity", - "summary", "patch", "FinalRequestCompactionMiddleware", "usage", diff --git a/tests/test_agent_summarization_streaming.py b/tests/test_agent_summarization_streaming.py index 21fc67e00..3f26378ca 100644 --- a/tests/test_agent_summarization_streaming.py +++ b/tests/test_agent_summarization_streaming.py @@ -17,6 +17,7 @@ from langchain_core.messages import ( AnyMessage, HumanMessage, SystemMessage, + ToolMessage, ) from langchain_core.tools import tool from langgraph.checkpoint.memory import InMemorySaver @@ -178,7 +179,6 @@ def _real_compaction_graph(*, model, summarizer, tools=None, checkpointer=None): model=model, tools=list(tools or []), middleware=[ - summary_middleware, _DynamicSystemMiddleware(), FinalRequestCompactionMiddleware(summarizer=summary_middleware), ], @@ -232,14 +232,18 @@ def test_streaming_agent_uses_non_streaming_llm_for_summary(): ): asyncio.run(agent._create_agent(streaming=True)) - summary_middleware = next( + compaction_middleware = next( middleware for middleware in captured["middleware"] - if isinstance(middleware, ContextPreservingSummarizationMiddleware) + if isinstance(middleware, FinalRequestCompactionMiddleware) ) assert captured["model"] is main_llm - assert summary_middleware.model is non_streaming_llm + assert compaction_middleware.summarizer.model is non_streaming_llm + assert not any( + isinstance(middleware, ContextPreservingSummarizationMiddleware) + for middleware in captured["middleware"] + ) def test_streaming_agent_uses_non_streaming_llm_for_model_middlewares(): @@ -355,14 +359,18 @@ def test_non_streaming_agent_reuses_main_llm_for_summary(): ): asyncio.run(agent._create_agent(streaming=False)) - summary_middleware = next( + compaction_middleware = next( middleware for middleware in captured["middleware"] - if isinstance(middleware, ContextPreservingSummarizationMiddleware) + if isinstance(middleware, FinalRequestCompactionMiddleware) ) assert captured["model"] is main_llm - assert summary_middleware.model is main_llm + assert compaction_middleware.summarizer.model is main_llm + assert not any( + isinstance(middleware, ContextPreservingSummarizationMiddleware) + for middleware in captured["middleware"] + ) def test_summary_failure_does_not_replace_existing_context(): @@ -390,11 +398,12 @@ def test_summary_failure_does_not_replace_existing_context(): ): asyncio.run(agent._create_agent(streaming=False)) - summary_middleware = next( + compaction_middleware = next( middleware for middleware in captured["middleware"] - if isinstance(middleware, ContextPreservingSummarizationMiddleware) + if isinstance(middleware, FinalRequestCompactionMiddleware) ) + summary_middleware = compaction_middleware.summarizer messages = [ HumanMessage(content=f"必须保留的旧上下文 {index} " * 200) for index in range(160) @@ -570,6 +579,222 @@ def test_uncompactable_request_below_window_still_calls_main_model(): assert received == [request] +def test_overflow_with_small_history_compacts_to_hard_window(): + """历史低于常规保留量时,超窗请求仍应尝试压缩而不是直接拒绝。""" + summarizer = ContextPreservingSummarizationMiddleware( + model=_SuccessfulSummaryLLM("summary"), + trigger=("fraction", 0.85), + keep=("messages", 20), + ) + middleware = FinalRequestCompactionMiddleware(summarizer=summarizer) + messages = [ + HumanMessage(content=f"近期历史 {index} " * 5) for index in range(4) + ] + request = _final_request( + messages=messages, + system_message=SystemMessage(content="固定系统约束 " * 1150), + ) + original_budget = UsageMiddleware.estimate_request(request) + fixed_budget = UsageMiddleware.estimate_request(request.override(messages=[])) + assert fixed_budget["estimated_input_ratio"] < 1 + assert summarizer.token_counter(messages) < 2048 * 0.10 + assert original_budget["estimated_input_ratio"] > 1 + 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 + compacted_budget = UsageMiddleware.estimate_request(received[0]) + assert 0.85 < compacted_budget["estimated_input_ratio"] <= 1 + + +def test_compaction_uses_hard_window_when_soft_target_is_unreachable(): + """固定开销超过软线时,应缩小近期历史直到完整窗口可承载。""" + summarizer = ContextPreservingSummarizationMiddleware( + model=_SuccessfulSummaryLLM("summary"), + trigger=("fraction", 0.85), + keep=("messages", 20), + ) + middleware = FinalRequestCompactionMiddleware(summarizer=summarizer) + messages = [ + HumanMessage(content=f"需要择量保留的历史 {index} " * 20) + for index in range(24) + ] + request = _final_request( + messages=messages, + system_message=SystemMessage(content="固定系统约束 " * 1100), + ) + fixed_budget = UsageMiddleware.estimate_request(request.override(messages=[])) + assert 0.85 < fixed_budget["estimated_input_ratio"] < 1 + assert UsageMiddleware.estimate_request(request)["estimated_input_ratio"] > 1 + 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 + compacted_budget = UsageMiddleware.estimate_request(received[0]) + assert 0.85 < compacted_budget["estimated_input_ratio"] <= 1 + + +def test_forced_compaction_advances_beyond_existing_summary(): + """超窗时应继续压缩旧事实,不能因首条已有摘要而误判无进展。""" + summarizer = ContextPreservingSummarizationMiddleware( + model=_SuccessfulSummaryLLM("summary"), + trigger=("fraction", 0.85), + keep=("messages", 20), + ) + middleware = FinalRequestCompactionMiddleware(summarizer=summarizer) + messages = [ + AIMessage( + content="已有摘要 " * 570, + additional_kwargs={"lc_source": "summarization"}, + ), + HumanMessage(content="仍可继续压缩的旧事实 " * 92), + HumanMessage(content="最新问题"), + ] + request = _final_request( + messages=messages, + system_message=SystemMessage(content="固定系统约束 " * 700), + ) + fixed_budget = UsageMiddleware.estimate_request(request.override(messages=[])) + assert fixed_budget["estimated_input_ratio"] < 1 + assert UsageMiddleware.estimate_request(request)["estimated_input_ratio"] > 1 + partition = summarizer.partition_for_token_limit( + messages, + int(2048 * 0.10), + force=True, + ) + assert partition is not None + assert len(partition[0]) == 2 + 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 received[0].messages[-1].id == messages[-1].id + assert UsageMiddleware.estimate_request(received[0])["estimated_input_ratio"] <= 1 + + +def test_forced_compaction_keeps_only_current_message_instead_of_summarizing_it(): + """唯一当前消息不可被强制摘要,无法装入窗口时应保留历史并拒绝。""" + summarizer = ContextPreservingSummarizationMiddleware( + model=_SuccessfulSummaryLLM("summary"), + trigger=("fraction", 0.85), + keep=("messages", 20), + ) + middleware = FinalRequestCompactionMiddleware(summarizer=summarizer) + latest_message = HumanMessage(content="唯一且必须保留的当前问题 " * 40) + request = _final_request( + messages=[latest_message], + system_message=SystemMessage(content="固定系统约束 " * 1100), + ) + fixed_budget = UsageMiddleware.estimate_request(request.override(messages=[])) + assert fixed_budget["estimated_input_ratio"] < 1 + assert UsageMiddleware.estimate_request(request)["estimated_input_ratio"] > 1 + 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 == [] + assert request.messages == [latest_message] + + +def test_forced_compaction_rejects_single_long_current_message(): + """当前消息超过保留预算时也不可被整体摘要。""" + summarizer = ContextPreservingSummarizationMiddleware( + model=_SuccessfulSummaryLLM("summary"), + trigger=("fraction", 0.85), + keep=("messages", 20), + ) + middleware = FinalRequestCompactionMiddleware(summarizer=summarizer) + latest_message = HumanMessage(content="当前用户唯一问题 " * 830) + request = _final_request( + messages=[latest_message], + system_message=SystemMessage(content="固定系统约束 " * 100), + ) + assert UsageMiddleware.estimate_request(request)["estimated_input_ratio"] > 1 + assert summarizer.partition_for_token_limit( + [latest_message], + int(2048 * 0.10), + force=True, + ) is None + 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 == [] + assert request.messages == [latest_message] + + +def test_repartition_advances_past_complete_old_tool_transaction(): + """二次分区应整体摘要旧工具事务,而不是停在相同安全边界。""" + summarizer = ContextPreservingSummarizationMiddleware( + model=_SuccessfulSummaryLLM("summary"), + trigger=("fraction", 0.85), + keep=("messages", 20), + ) + middleware = FinalRequestCompactionMiddleware(summarizer=summarizer) + latest_message = HumanMessage(content="新的用户问题") + messages = [ + HumanMessage(content="更早历史 " * 30), + AIMessage( + content="", + tool_calls=[ + { + "name": "old_tool", + "args": {"text": "工具参数 " * 500}, + "id": "old-call", + } + ], + ), + ToolMessage(content="旧工具结果", tool_call_id="old-call"), + latest_message, + ] + request = _final_request( + messages=messages, + system_message=SystemMessage(content="固定系统约束 " * 780), + ) + assert UsageMiddleware.estimate_request(request)["estimated_input_ratio"] > 1 + 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 received[0].messages[-1].id == latest_message.id + assert not any(isinstance(message, ToolMessage) for message in received[0].messages) + assert UsageMiddleware.estimate_request(received[0])["estimated_input_ratio"] <= 1 + + def test_summary_calls_include_version_compatible_internal_metadata(): """摘要调用应合并当前 LangChain 提供的内部流式过滤标记。""" summary_model = _MetadataRecordingSummaryLLM("summary") @@ -740,7 +965,7 @@ def test_large_tool_result_allows_same_turn_recompaction(): def large_result() -> str: """返回足以再次耗尽主模型窗口的工具结果。""" calls.append("once") - return "超长工具结果 " * 2000 + return "超长工具结果 " * 800 summarizer = _CountingSummaryLLM("summary") model = _RecordingChatModel( @@ -766,7 +991,10 @@ def test_large_tool_result_allows_same_turn_recompaction(): assert calls == ["once"] assert summarizer.calls == 2 assert len(model.seen_messages) == 2 - assert "超长工具结果" not in str(model.seen_messages[1]) + second_request = model.seen_messages[1] + assert isinstance(second_request[-2], AIMessage) + assert isinstance(second_request[-1], ToolMessage) + assert second_request[-1].tool_call_id == second_request[-2].tool_calls[0]["id"] assert result["messages"][-1].content == "工具完成" @@ -870,6 +1098,82 @@ def test_real_agent_failure_does_not_commit_compacted_history(failure): assert len(model.seen_messages) == 1 +def test_history_triggered_compaction_waits_for_main_model_success(): + """历史本身触发压缩时,主模型失败也不得提前提交摘要状态。""" + checkpointer = InMemorySaver() + summarizer = _CountingSummaryLLM("summary") + summarizer.profile = {"max_input_tokens": 2048} + model = _FailingMainModel( + responses=[AIMessage(content="不会返回")], + profile={"max_input_tokens": 2048}, + ) + graph = _real_compaction_graph( + model=model, + summarizer=summarizer, + checkpointer=checkpointer, + ) + messages = [ + HumanMessage(content=f"历史直接触发压缩 {index} " * 40) + for index in range(30) + ] + summary_middleware = ContextPreservingSummarizationMiddleware( + model=summarizer, + trigger=("fraction", 0.85), + keep=("messages", 20), + ) + assert summary_middleware.token_counter(messages) >= 2048 * 0.85 + config = {"configurable": {"thread_id": "history-trigger-main-failure"}} + + async def _run(): + with pytest.raises(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 + ] + assert summarizer.calls >= 1 + assert len(model.seen_messages) == 1 + + +def test_history_triggered_compaction_commits_after_main_model_success(): + """历史本身触发压缩时,摘要与模型结果应在成功后一次性提交。""" + checkpointer = InMemorySaver() + summarizer = _CountingSummaryLLM("summary") + summarizer.profile = {"max_input_tokens": 2048} + model = _RecordingChatModel( + responses=[AIMessage(content="继续完成")], + profile={"max_input_tokens": 2048}, + ) + graph = _real_compaction_graph( + model=model, + summarizer=summarizer, + checkpointer=checkpointer, + ) + messages = [ + HumanMessage(content=f"历史直接触发压缩 {index} " * 40) + for index in range(30) + ] + config = {"configurable": {"thread_id": "history-trigger-main-success"}} + + async def _run(): + result = await graph.ainvoke({"messages": messages}, config=config) + return result, await graph.aget_state(config) + + result, snapshot = asyncio.run(_run()) + + assert [message.content for message in result["messages"]] == [ + message.content for message in snapshot.values["messages"] + ] + assert "保留旧事实的摘要" in result["messages"][0].content + assert result["messages"][-1].content == "继续完成" + assert len(result["messages"]) < len(messages) + assert summarizer.calls >= 1 + assert len(model.seen_messages) == 1 + + def test_unsummarizable_message_requires_new_context_instead_of_retry(): """确定性不可裁剪的消息应给出可前进路径,而非建议无效重试。""" middleware = ContextPreservingSummarizationMiddleware( diff --git a/tests/test_agent_tool_policy.py b/tests/test_agent_tool_policy.py index fecc2256e..8882b38ee 100644 --- a/tests/test_agent_tool_policy.py +++ b/tests/test_agent_tool_policy.py @@ -4,7 +4,6 @@ from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest -from langchain.agents.middleware import SummarizationMiddleware from langchain_core.messages import ToolMessage from pydantic import BaseModel, Field @@ -12,6 +11,7 @@ import app.agent as agent_module from app.agent.middleware.activity_log import ActivityLogMiddleware from app.agent.middleware.memory import MemoryMiddleware from app.agent.middleware.policy import AgentPolicyMiddleware +from app.agent.middleware.summarization import FinalRequestCompactionMiddleware from app.agent.policy import ( DEFAULT_TOOL_POLICY_ORCHESTRATOR, DEFAULT_TOOL_POLICY_REGISTRY, @@ -838,12 +838,12 @@ def test_main_agent_preserves_activity_log_middleware_order() -> None: for index, middleware in enumerate(middlewares) if isinstance(middleware, ActivityLogMiddleware) ) - summary_index = next( + compaction_index = next( index for index, middleware in enumerate(middlewares) - if isinstance(middleware, SummarizationMiddleware) + if isinstance(middleware, FinalRequestCompactionMiddleware) ) assert policy_index == 0 assert activity_index == memory_index + 1 - assert summary_index == activity_index + 1 + assert compaction_index > activity_index