aps-agent/server/agent_core/summarizer.py

285 lines
12 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# ============================================================
# 滚动摘要 · LLM 压缩生成方(moduleId: core-summarizer, 可重生 ✅)
# 矩阵 64 行剩余项:滚动摘要的 LLM 压缩生成方接线 + 越权上下文测试。
# - compress_scroll_summary:固定中文提示词 -> ModelProvider.chat_text 压缩;
# provider 不可用/超时/异常/输出为空/超预算 -> 确定性回退
# (deterministic_compress,复用 estimate_tokens 预算口径截断)。
# - maybe_roll_session_summary:每轮对话后的预算门控接线(fail closed)。
# 绝不编造:提示词限定只基于给定消息;任何异常输出都回退确定性截断;
# 压缩结果经 scroll_summary 入库,sourceMessageIds/supersedes 可溯源。
# ============================================================
from __future__ import annotations
import inspect
import re
from dataclasses import dataclass, field
from datetime import datetime
from typing import Any
from server.agent_core.context_policy import (
DEFAULT_WINDOW,
ContextBudget,
ContextBudgetExceeded,
assemble_context,
estimate_tokens,
get_policy,
scroll_summary,
)
# 固定中文提示词:只允许压缩给定消息,禁止编造。
SYSTEM_PROMPT = (
"你是一名 APS 排产助手的长会话压缩器。你的唯一任务:把给定的对话消息压缩成"
"一段结构化中文摘要。\n"
"铁律:\n"
"1. 只能使用给定对话里出现过的内容——事实、数字、决策、未决事项、验收标准;\n"
"2. 禁止编造、补全或推断给定对话中不存在的信息;\n"
"3. 输出纯文本摘要正文,不要输出 JSON、markdown 代码块、消息编号或任何解释;\n"
"4. 必须保留:用户的验收标准与硬性要求、已做决策及理由、未决事项、关键数字、"
"当前进度;\n"
"5. 长度不得超过给定预算(中文 1 字 ≈ 1 token 估算),宁可保守删减细节,不可超长。"
)
@dataclass(frozen=True)
class SummaryCompression:
"""压缩结果:摘要文本 + 被压缩消息 id 列表 + 生成方标记与回退原因。"""
text: str
compressed_message_ids: list[str] = field(default_factory=list)
mode: str = "deterministic" # "llm" | "deterministic"
reason: str | None = None # 仅回退时记录原因(审计/测试可检查)
def _normalize_message(message: Any, idx: int) -> dict[str, Any] | None:
"""与 context_policy._normalize_message 同口径:只保留 user/agent 且有文本的消息。"""
if not isinstance(message, dict):
return None
role = str(message.get("role") or "")
if role in ("assistant", "bot", "ai"):
role = "agent"
if role not in ("user", "agent"):
return None
text = str(message.get("text") or message.get("content") or "").strip()
if not text:
return None
return {
"messageId": str(message.get("messageId") or message.get("id") or f"m{idx}"),
"role": role,
"text": text,
}
def _normalize_messages(messages: list[Any]) -> list[dict[str, Any]]:
return [
item for item in (
_normalize_message(message, idx)
for idx, message in enumerate(messages)
)
if item is not None
]
def build_compression_prompt(messages: list[dict[str, Any]], budget_tokens: int) -> str:
"""固定中文提示词模板:把给定消息原样呈现给 LLM,禁止引用外部信息。"""
lines = [
f"预算:摘要估算 token 不得超过 {budget_tokens}。",
"",
"需要压缩的对话(按时间顺序):",
"",
]
for idx, message in enumerate(messages, start=1):
who = "用户" if message.get("role") == "user" else "助手"
lines.append(f"[消息{idx}] {who}:{message.get('text')}")
lines.extend(["", "请只输出压缩后的摘要正文:"])
return "\n".join(lines)
def deterministic_compress(messages: list[Any], budget_tokens: int) -> SummaryCompression:
"""确定性回退:复用 estimate_tokens 预算口径,从最旧开始保留直到超预算。
只使用给定消息内容(绝不引入外部信息);输出可溯源到消息 id。
"""
if budget_tokens <= 0:
raise ValueError("budget_tokens must be positive")
normalized = _normalize_messages(messages)
if not normalized:
raise ValueError("no compressible messages")
message_ids = [m["messageId"] for m in normalized]
lines = [f"{'用户' if m['role'] == 'user' else '助手'}:{m['text']}" for m in normalized]
kept: list[str] = []
for line in lines:
candidate = "\n".join([*kept, line])
if estimate_tokens(candidate) > budget_tokens:
break
kept.append(line)
text = "\n".join(kept).strip()
if not text:
# 单条都放不下:对最新一条做预算内强制截断,保证非空回退。
text = lines[-1]
while estimate_tokens(text) > budget_tokens and len(text) > 8:
text = text[:-8].rstrip()
return SummaryCompression(
text=text,
compressed_message_ids=message_ids,
mode="deterministic",
reason="budget-truncation",
)
_FENCE_RE = re.compile(r"```(?:json|text)?\s*([\s\S]*?)```", re.IGNORECASE)
def _strip_fence(text: str) -> str:
"""容错:剥离 markdown 代码块围栏(部分国产端点会在文本外套 ```)。"""
match = _FENCE_RE.search(text)
return (match.group(1) if match else text).strip()
def _chat_accepts_timeout(provider: Any) -> bool:
try:
return "timeout" in inspect.signature(provider.chat_text).parameters
except (TypeError, ValueError):
return True # 拿不到签名时按完整签名处理
async def _provider_chat(provider: Any, system: str, user: str, timeout: float) -> str | None:
if _chat_accepts_timeout(provider):
return await provider.chat_text(system, user, timeout=timeout)
return await provider.chat_text(system, user)
async def compress_scroll_summary(
messages: list[dict[str, Any]], *,
provider: Any, budget_tokens: int = 600, timeout: float = 45.0,
) -> SummaryCompression:
"""LLM 压缩生成方:固定中文提示词 + provider.chat_text(ModelProvider 接口)。
- 只允许压缩给定消息内容:提示词限定 + 输出校验 + 异常回退,绝不编造。
- provider 不可用(未启用/返回 None)/ 超时 / 异常 / 输出为空 / 输出超预算
-> 确定性回退 deterministic_compress(复用 estimate_tokens 预算截断)。
Returns:
SummaryCompression{text, compressedMessageIds, mode, reason}
Raises:
ValueError: 无可压缩消息或 budget_tokens <= 0(fail closed)
"""
normalized = _normalize_messages(messages)
if not normalized:
raise ValueError("compress_scroll_summary: no compressible messages")
if budget_tokens <= 0:
raise ValueError("compress_scroll_summary: budget_tokens must be positive")
message_ids = [m["messageId"] for m in normalized]
def fallback(reason: str) -> SummaryCompression:
compressed = deterministic_compress(normalized, budget_tokens)
return SummaryCompression(
text=compressed.text,
compressed_message_ids=message_ids,
mode="deterministic",
reason=reason,
)
raw: str | None = None
reason: str | None = None
try:
raw = await _provider_chat(
provider,
SYSTEM_PROMPT,
build_compression_prompt(normalized, budget_tokens),
timeout,
)
except Exception as exc: # noqa: BLE001 - provider 异常一律回退,绝不阻断
reason = f"provider-error:{type(exc).__name__}"
if raw is None:
return fallback(reason or "provider-unavailable")
text = _strip_fence(str(raw))
if not text:
return fallback("llm-empty-output")
if estimate_tokens(text) > budget_tokens:
return fallback("llm-over-budget")
return SummaryCompression(
text=text,
compressed_message_ids=message_ids,
mode="llm",
reason=None,
)
async def maybe_roll_session_summary(
store_data: dict[str, Any], session_id: str, *,
provider: Any | None = None,
budget: ContextBudget | None = None,
window: int = DEFAULT_WINDOW,
auto_summarize: bool | None = None,
by: str | None = None,
now: datetime | None = None,
) -> dict[str, Any] | None:
"""每轮对话后的滚动摘要钩子:预算门控 + LLM 压缩(fail closed)。
触发条件(预算门控):
- 会话原始消息(未钉住)估算 token 超过会话层预算(长会话),或
- 当前会话层组装已超预算(layer == "session")。
满足时,把「滑出窗口且未被既有摘要覆盖」的消息交给 compress_scroll_summary
(LLM 压缩;异常/超时/输出不合格 -> 确定性回退),再经 scroll_summary 入库
(sourceMessageIds + supersedes 溯源链,compressedBy 标记生成方)。
滚动后会话层仍超预算 -> ContextBudgetExceeded(可预测拒绝,fail closed)。
auto_summarize: None -> 读会话策略 contextPolicies[session].autoSummarize,
缺省 True(长会话默认开启);False -> 永不滚动。
Returns:
新摘要 dict(发生滚动)或 None(未触发/被禁用)。
Raises:
ContextBudgetExceeded: 滚动后仍超预算(可预测拒绝)
ValueError: 会话/项目不可访问(越权 fail closed)
"""
budget = budget or ContextBudget()
session_limit = budget.layer_limit("session")
policy = get_policy(store_data, session_id)
enabled = policy.get("autoSummarize", True) if auto_summarize is None else auto_summarize
if not enabled:
return None
normalized = _normalize_messages((store_data.get("messages") or {}).get(session_id) or [])
pinned_ids = {str(item.get("messageId")) for item in policy.get("pinned") or []}
raw_history = [m for m in normalized if m["messageId"] not in pinned_ids]
if not raw_history:
return None
raw_tokens = estimate_tokens("\n".join(m["text"] for m in raw_history))
layer_over = False
try:
assemble_context(store_data, session_id, None, budget, window=window)
except ContextBudgetExceeded as exc:
if exc.layer != "session":
raise
layer_over = True
if raw_tokens <= session_limit and not layer_over:
return None
covered: set[str] = set()
for entry in [policy.get("summary"), *(policy.get("summaries") or [])]:
if isinstance(entry, dict):
covered.update(str(sid) for sid in (entry.get("sourceMessageIds") or []))
recent_ids = {m["messageId"] for m in normalized[-window:]}
eligible = [
m for m in raw_history
if m["messageId"] not in recent_ids and m["messageId"] not in covered
]
if not eligible:
if layer_over:
# 无可压缩消息但会话层仍超预算:可预测拒绝(fail closed)。
assemble_context(store_data, session_id, None, budget, window=window)
return None
summary_budget = max(1, session_limit // 2)
compression = await compress_scroll_summary(
eligible, provider=provider, budget_tokens=summary_budget,
)
new_summary = scroll_summary(
store_data, session_id, compression.text,
source_message_ids=compression.compressed_message_ids,
by=by or "planner", now=now, summarizer=compression.mode,
)
# 滚动后验证:会话层仍超预算 -> ContextBudgetExceeded(可预测拒绝)。
assemble_context(store_data, session_id, None, budget, window=window)
return new_summary