mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-07 00:16:57 +08:00
fix(agent): compact oversized final model requests (#6299)
This commit is contained in:
+15
-3
@@ -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.skills import SKILL_TOOL_NAME, SkillsMiddleware
|
||||||
from app.agent.middleware.summarization import (
|
from app.agent.middleware.summarization import (
|
||||||
ContextPreservingSummarizationMiddleware as SummarizationMiddleware,
|
ContextPreservingSummarizationMiddleware as SummarizationMiddleware,
|
||||||
|
FinalRequestCompactionMiddleware,
|
||||||
)
|
)
|
||||||
from app.agent.middleware.subagents import (
|
from app.agent.middleware.subagents import (
|
||||||
SUBAGENT_CONTROL_TOOL_NAME,
|
SUBAGENT_CONTROL_TOOL_NAME,
|
||||||
@@ -1944,6 +1945,12 @@ class MoviePilotAgent:
|
|||||||
if getattr(tool, "name", None) == QUERY_ACTIVITY_LOG_TOOL_NAME
|
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 = [
|
middlewares = [
|
||||||
# 宿主策略必须位于最外层,确保插件覆盖工具基类也不能绕过。
|
# 宿主策略必须位于最外层,确保插件覆盖工具基类也不能绕过。
|
||||||
@@ -1965,9 +1972,7 @@ class MoviePilotAgent:
|
|||||||
# 活动日志依赖记忆上下文,并应在摘要压缩前完成读取与记录。
|
# 活动日志依赖记忆上下文,并应在摘要压缩前完成读取与记录。
|
||||||
*([activity_log_middleware] if activity_log_middleware else []),
|
*([activity_log_middleware] if activity_log_middleware else []),
|
||||||
# 上下文压缩
|
# 上下文压缩
|
||||||
SummarizationMiddleware(
|
summarization_middleware,
|
||||||
model=non_streaming_model, trigger=("fraction", 0.85)
|
|
||||||
),
|
|
||||||
# 错误工具调用修复
|
# 错误工具调用修复
|
||||||
PatchToolCallsMiddleware(),
|
PatchToolCallsMiddleware(),
|
||||||
# 子代理委派
|
# 子代理委派
|
||||||
@@ -1990,6 +1995,13 @@ class MoviePilotAgent:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# 需要在动态 system 与工具筛选完成后按最终输入预算补充压缩。
|
||||||
|
middlewares.append(
|
||||||
|
FinalRequestCompactionMiddleware(
|
||||||
|
summarizer=summarization_middleware,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
# 预算观察器必须位于最内层,才能看到动态 system 和最终筛选后的工具。
|
# 预算观察器必须位于最内层,才能看到动态 system 和最终筛选后的工具。
|
||||||
middlewares.append(
|
middlewares.append(
|
||||||
UsageMiddleware(
|
UsageMiddleware(
|
||||||
|
|||||||
@@ -1,8 +1,33 @@
|
|||||||
"""Agent 会话上下文压缩中间件。"""
|
"""Agent 会话上下文压缩中间件。"""
|
||||||
|
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from importlib import import_module
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
from langchain.agents.middleware import SummarizationMiddleware
|
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 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):
|
class ContextSummarizationError(RuntimeError):
|
||||||
@@ -37,10 +62,16 @@ class ContextPreservingSummarizationMiddleware(SummarizationMiddleware):
|
|||||||
def _create_summary(self, messages_to_summarize: list[AnyMessage]) -> str:
|
def _create_summary(self, messages_to_summarize: list[AnyMessage]) -> str:
|
||||||
"""同步摘要失败时保持原图状态。"""
|
"""同步摘要失败时保持原图状态。"""
|
||||||
formatted_messages = self._prepare_summary_input(messages_to_summarize)
|
formatted_messages = self._prepare_summary_input(messages_to_summarize)
|
||||||
|
summary_model = getattr(self, "_summary_model", self.model)
|
||||||
try:
|
try:
|
||||||
response = self.model.invoke(
|
response = summary_model.invoke(
|
||||||
self.summary_prompt.format(messages=formatted_messages).rstrip(),
|
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:
|
except Exception as err:
|
||||||
raise ContextSummarizationError(self._ERROR_MESSAGE) from 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:
|
async def _acreate_summary(self, messages_to_summarize: list[AnyMessage]) -> str:
|
||||||
"""异步摘要失败时保持原图状态。"""
|
"""异步摘要失败时保持原图状态。"""
|
||||||
formatted_messages = self._prepare_summary_input(messages_to_summarize)
|
formatted_messages = self._prepare_summary_input(messages_to_summarize)
|
||||||
|
summary_model = getattr(self, "_summary_model", self.model)
|
||||||
try:
|
try:
|
||||||
response = await self.model.ainvoke(
|
response = await summary_model.ainvoke(
|
||||||
self.summary_prompt.format(messages=formatted_messages).rstrip(),
|
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:
|
except Exception as err:
|
||||||
raise ContextSummarizationError(self._ERROR_MESSAGE) from err
|
raise ContextSummarizationError(self._ERROR_MESSAGE) from err
|
||||||
return self._require_valid_summary(response.text.strip())
|
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)
|
||||||
|
|||||||
@@ -438,6 +438,7 @@ class AgentBackgroundOutputTest(unittest.IsolatedAsyncioTestCase):
|
|||||||
"memory",
|
"memory",
|
||||||
"summary",
|
"summary",
|
||||||
"patch",
|
"patch",
|
||||||
|
"FinalRequestCompactionMiddleware",
|
||||||
"usage",
|
"usage",
|
||||||
],
|
],
|
||||||
[getattr(item, "name", item) for item in created["middleware"]],
|
[getattr(item, "name", item) for item in created["middleware"]],
|
||||||
@@ -554,6 +555,7 @@ class AgentBackgroundOutputTest(unittest.IsolatedAsyncioTestCase):
|
|||||||
"memory",
|
"memory",
|
||||||
"summary",
|
"summary",
|
||||||
"patch",
|
"patch",
|
||||||
|
"FinalRequestCompactionMiddleware",
|
||||||
"usage",
|
"usage",
|
||||||
],
|
],
|
||||||
[getattr(item, "name", item) for item in created["middleware"]],
|
[getattr(item, "name", item) for item in created["middleware"]],
|
||||||
@@ -759,6 +761,7 @@ class AgentBackgroundOutputTest(unittest.IsolatedAsyncioTestCase):
|
|||||||
"activity",
|
"activity",
|
||||||
"summary",
|
"summary",
|
||||||
"patch",
|
"patch",
|
||||||
|
"FinalRequestCompactionMiddleware",
|
||||||
"usage",
|
"usage",
|
||||||
],
|
],
|
||||||
[getattr(item, "name", item) for item in created["middleware"]],
|
[getattr(item, "name", item) for item in created["middleware"]],
|
||||||
|
|||||||
@@ -491,9 +491,10 @@ async def test_graph_keeps_mcp_first_winner_and_catalogs_all_collisions(
|
|||||||
activity_tool,
|
activity_tool,
|
||||||
subagent_task_tool,
|
subagent_task_tool,
|
||||||
]
|
]
|
||||||
assert captured["middlewares"][-2].name == "selector"
|
assert captured["middlewares"][-3].name == "selector"
|
||||||
else:
|
else:
|
||||||
assert "selection_tools" not in captured
|
assert "selection_tools" not in captured
|
||||||
|
assert captured["middlewares"][-2].name == "FinalRequestCompactionMiddleware"
|
||||||
assert captured["middlewares"][-1].name == "usage"
|
assert captured["middlewares"][-1].name == "usage"
|
||||||
policy_middleware = next(
|
policy_middleware = next(
|
||||||
middleware
|
middleware
|
||||||
|
|||||||
@@ -81,6 +81,7 @@ def test_final_request_can_exceed_window_before_message_fraction_triggers():
|
|||||||
_llm_type="test-chat",
|
_llm_type="test-chat",
|
||||||
profile={"max_input_tokens": 4096},
|
profile={"max_input_tokens": 4096},
|
||||||
)
|
)
|
||||||
|
model.with_retry = lambda: model
|
||||||
summarizer = SummarizationMiddleware(
|
summarizer = SummarizationMiddleware(
|
||||||
model=model,
|
model=model,
|
||||||
trigger=("fraction", 0.85),
|
trigger=("fraction", 0.85),
|
||||||
|
|||||||
@@ -4,7 +4,23 @@ from types import SimpleNamespace
|
|||||||
from unittest.mock import AsyncMock, patch
|
from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
import pytest
|
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
|
import app.agent as agent_module
|
||||||
from app.agent.memory import memory_manager
|
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 (
|
from app.agent.middleware.summarization import (
|
||||||
ContextSummarizationError,
|
ContextSummarizationError,
|
||||||
ContextPreservingSummarizationMiddleware,
|
ContextPreservingSummarizationMiddleware,
|
||||||
|
FinalRequestCompactionMiddleware,
|
||||||
)
|
)
|
||||||
|
from app.agent.middleware.usage import UsageMiddleware
|
||||||
|
|
||||||
|
|
||||||
class _FakeLLM:
|
class _FakeLLM:
|
||||||
@@ -22,6 +40,10 @@ class _FakeLLM:
|
|||||||
self.model = model
|
self.model = model
|
||||||
self.profile = {"max_input_tokens": 64000}
|
self.profile = {"max_input_tokens": 64000}
|
||||||
|
|
||||||
|
def with_retry(self):
|
||||||
|
"""满足新版 LangChain 摘要模型的 Runnable 合同。"""
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
class _FailingSummaryLLM(_FakeLLM):
|
class _FailingSummaryLLM(_FakeLLM):
|
||||||
"""模拟摘要模型暂时不可用。"""
|
"""模拟摘要模型暂时不可用。"""
|
||||||
@@ -47,6 +69,131 @@ class _SuccessfulSummaryLLM(_FakeLLM):
|
|||||||
return AIMessage(content="保留旧事实的摘要")
|
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:
|
class _FailingGraph:
|
||||||
"""模拟上下文压缩阶段失败的 Agent 图。"""
|
"""模拟上下文压缩阶段失败的 Agent 图。"""
|
||||||
|
|
||||||
@@ -278,6 +425,451 @@ def test_summary_success_still_replaces_old_context():
|
|||||||
assert len(update["messages"]) < len(messages)
|
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():
|
def test_unsummarizable_message_requires_new_context_instead_of_retry():
|
||||||
"""确定性不可裁剪的消息应给出可前进路径,而非建议无效重试。"""
|
"""确定性不可裁剪的消息应给出可前进路径,而非建议无效重试。"""
|
||||||
middleware = ContextPreservingSummarizationMiddleware(
|
middleware = ContextPreservingSummarizationMiddleware(
|
||||||
|
|||||||
Reference in New Issue
Block a user