mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-08-14 02:05:13 +08:00
fix(agent): preserve context during request compaction (#6300)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user