406 lines
17 KiB
Python
406 lines
17 KiB
Python
from __future__ import annotations
|
||
|
||
import copy
|
||
import hashlib
|
||
import math
|
||
import re
|
||
from collections.abc import Iterable
|
||
from dataclasses import dataclass, field
|
||
from datetime import UTC, datetime
|
||
from typing import Any
|
||
|
||
from server.timeutil import fmt_dt
|
||
|
||
LAYERS = ("fixed", "project", "session", "retrieval")
|
||
|
||
DEFAULT_SYSTEM_PROMPT = (
|
||
"你是 APS 排产智能体。你只拥有提议权:一切写操作必须经门禁确认;"
|
||
"来源与提交语义必须区分,中间状态不得当作已完成事实;"
|
||
"失败必须显式报告,不得用泛化确认掩盖。"
|
||
)
|
||
|
||
DEFAULT_WINDOW = 8
|
||
|
||
_CJK_RE = re.compile(r"[\u4e00-\u9fff]")
|
||
_MESSAGE_ID_FALLBACK = "m{idx}"
|
||
|
||
|
||
def _now_local() -> datetime:
|
||
"""本地无时区时间(与仓库 fmt_dt 语义一致)。"""
|
||
return datetime.now(UTC).astimezone().replace(tzinfo=None)
|
||
|
||
|
||
def estimate_tokens(text: str) -> int:
|
||
"""确定性 token 估算:CJK 字符按 1 token,其余按 4 字符 1 token。
|
||
|
||
仅用于预算控制(可预测、确定性),不等同真实分词器。
|
||
"""
|
||
if not text:
|
||
return 0
|
||
cjk = len(_CJK_RE.findall(text))
|
||
other = len(text) - cjk
|
||
return cjk + math.ceil(other / 4)
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class ContextBudget:
|
||
"""四层 token 预算;各层独立限额,总预算 = 各层之和。"""
|
||
|
||
fixed: int = 1500 # L-固定:系统提示
|
||
project: int = 3000 # L-项目:项目资料摘要
|
||
session: int = 6000 # L-会话:滚动摘要 + 钉住 + 最近 N 轮
|
||
retrieval: int = 4000 # L-检索:RAG 命中 + 引用解析结果
|
||
|
||
def layer_limit(self, layer: str) -> int:
|
||
if layer not in LAYERS:
|
||
raise ValueError(f"unknown layer: {layer!r}")
|
||
return getattr(self, layer)
|
||
|
||
def total(self) -> int:
|
||
return self.fixed + self.project + self.session + self.retrieval
|
||
|
||
@classmethod
|
||
def from_dict(cls, values: dict[str, Any]) -> ContextBudget:
|
||
return cls(**{k: int(v) for k, v in (values or {}).items() if k in LAYERS})
|
||
|
||
|
||
class ContextBudgetExceeded(Exception):
|
||
"""某层上下文超出预算:可预测拒绝(fail closed),携带层/用量/限额。"""
|
||
|
||
def __init__(self, layer: str, required: int, budget: int) -> None:
|
||
super().__init__(f"context layer {layer!r} needs {required} tokens, budget {budget}")
|
||
self.layer = layer
|
||
self.required = required
|
||
self.budget = budget
|
||
|
||
|
||
@dataclass
|
||
class ContextAssembly:
|
||
"""四层组装结果:各层文本可检查,整体可注入 prompt。"""
|
||
|
||
layers: dict[str, str] # layer -> 渲染文本
|
||
usage: dict[str, int] # layer -> 估算 token
|
||
budgets: dict[str, int] # layer -> 预算
|
||
pinned: list[dict[str, Any]] = field(default_factory=list) # 已生效钉住项
|
||
summary: dict[str, Any] | None = None # 当前滚动摘要(含溯源)
|
||
|
||
def layer(self, name: str) -> str:
|
||
if name not in self.layers:
|
||
raise ValueError(f"unknown layer: {name!r}")
|
||
return self.layers[name]
|
||
|
||
def text(self) -> str:
|
||
"""按固定/项目/会话/检索顺序拼接为单一 prompt 文本。"""
|
||
return "\n\n".join(f"[{name}] {self.layers[name]}" for name in LAYERS)
|
||
|
||
def tokens(self) -> int:
|
||
return sum(self.usage.values())
|
||
|
||
|
||
# ---------------- 策略状态(钉住 / 滚动摘要) ----------------
|
||
|
||
def _policies(store_data: dict[str, Any]) -> dict[str, dict[str, Any]]:
|
||
return store_data.setdefault("contextPolicies", {})
|
||
|
||
|
||
def get_policy(store_data: dict[str, Any], session_id: str) -> dict[str, Any]:
|
||
"""读会话上下文策略(钉住项/滚动摘要);不写时返回空策略。"""
|
||
return copy.deepcopy(_policies(store_data).get(session_id) or {})
|
||
|
||
|
||
def _message_id(message: dict[str, Any], idx: int) -> str:
|
||
return str(message.get("id") or _MESSAGE_ID_FALLBACK.format(idx=idx))
|
||
|
||
|
||
def _find_message(store_data: dict[str, Any], session_id: str, message_id: str) -> int:
|
||
messages = (store_data.get("messages") or {}).get(session_id) or []
|
||
for idx, message in enumerate(messages):
|
||
if isinstance(message, dict) and _message_id(message, idx) == message_id:
|
||
return idx
|
||
return -1
|
||
|
||
|
||
def pin_message(store_data: dict[str, Any], session_id: str, message_id: str,
|
||
*, note: str | None = None, by: str | None = None,
|
||
now: datetime | None = None) -> dict[str, Any]:
|
||
"""把关键消息钉进会话上下文:钉住项永不被摘要/窗口淘汰。
|
||
|
||
- 消息必须存在(fail closed:不存在显式 ValueError)
|
||
- 幂等:重复钉同一消息只保留一条
|
||
Returns:
|
||
新钉住项 {messageId, note, pinnedAt, pinnedBy}
|
||
"""
|
||
idx = _find_message(store_data, session_id, message_id)
|
||
if idx < 0:
|
||
raise ValueError(f"message not found: {message_id!r} in session {session_id!r}")
|
||
policies = _policies(store_data)
|
||
policy = policies.setdefault(session_id, {"pinned": [], "summary": None, "summaries": []})
|
||
pinned = policy.setdefault("pinned", [])
|
||
for item in pinned:
|
||
if item.get("messageId") == message_id:
|
||
return copy.deepcopy(item)
|
||
entry = {
|
||
"messageId": message_id,
|
||
"note": note,
|
||
"pinnedAt": fmt_dt(now or _now_local()),
|
||
"pinnedBy": by,
|
||
}
|
||
pinned.append(entry)
|
||
return copy.deepcopy(entry)
|
||
|
||
|
||
def unpin_message(store_data: dict[str, Any], session_id: str, message_id: str) -> bool:
|
||
"""解除钉住;返回是否实际移除。"""
|
||
policies = _policies(store_data)
|
||
policy = policies.get(session_id)
|
||
if not policy:
|
||
return False
|
||
pinned = policy.get("pinned") or []
|
||
kept = [item for item in pinned if item.get("messageId") != message_id]
|
||
if len(kept) == len(pinned):
|
||
return False
|
||
policy["pinned"] = kept
|
||
return True
|
||
|
||
|
||
def scroll_summary(store_data: dict[str, Any], session_id: str, text: str, *,
|
||
source_message_ids: list[str] | tuple[str, ...] | None = None,
|
||
decision_points: Iterable[str] = (),
|
||
open_items: Iterable[str] = (),
|
||
key_numbers: Iterable[str] = (),
|
||
by: str | None = None, now: datetime | None = None,
|
||
summarizer: Any | None = None) -> dict[str, Any]:
|
||
"""滚动压缩:把超出窗口的对话折叠为结构化摘要,可追溯到原消息 id。
|
||
|
||
新摘要继承前一条摘要的溯源链(supersedes),sourceMessageIds 记录
|
||
本次压缩覆盖的原消息 id——「聊了三小时依然记得开头的验收标准」。
|
||
Args:
|
||
summarizer: 压缩生成方标记(如 "llm"/"deterministic",来自
|
||
summarizer.compress_scroll_summary 的 mode)。默认 None 保持纯确定性
|
||
行为(现有调用不变,不写入新字段);传入时在摘要记录中落
|
||
compressedBy 溯源,便于审计压缩来自 LLM 还是确定性回退。
|
||
LLM 实际调用是异步的,在 summarizer.compress_scroll_summary 完成,
|
||
本函数负责把压缩结果按既有溯源链入库。
|
||
Returns:
|
||
新摘要 {summaryId, text, sourceMessageIds, decisionPoints,
|
||
openItems, keyNumbers, createdAt, createdBy, supersedes}
|
||
(传入 summarizer 时另含 compressedBy)
|
||
"""
|
||
if not text or not str(text).strip():
|
||
raise ValueError("scroll summary text must not be empty")
|
||
policies = _policies(store_data)
|
||
policy = policies.setdefault(session_id, {"pinned": [], "summary": None, "summaries": []})
|
||
previous = policy.get("summary")
|
||
generation = int((previous or {}).get("generation") or 0) + 1
|
||
digest = hashlib.sha256(
|
||
f"{session_id}|{generation}|{text}".encode()
|
||
).hexdigest()[:12]
|
||
summary = {
|
||
"summaryId": f"sum_{digest}",
|
||
"text": str(text).strip(),
|
||
"sourceMessageIds": list(source_message_ids or []),
|
||
"decisionPoints": list(decision_points or []),
|
||
"openItems": list(open_items or []),
|
||
"keyNumbers": list(key_numbers or []),
|
||
"createdAt": fmt_dt(now or _now_local()),
|
||
"createdBy": by,
|
||
"supersedes": (previous or {}).get("summaryId"),
|
||
"generation": generation,
|
||
}
|
||
if summarizer is not None:
|
||
summary["compressedBy"] = str(summarizer)
|
||
history = policy.setdefault("summaries", [])
|
||
if previous:
|
||
history.append(copy.deepcopy(previous))
|
||
policy["summary"] = summary
|
||
return copy.deepcopy(summary)
|
||
|
||
|
||
def _pinned_entries(store_data: dict[str, Any], session_id: str) -> list[dict[str, Any]]:
|
||
"""把钉住项解析为可注入内容:消息仍存在 -> 全文;被删 -> 显式 missing。"""
|
||
policy = _policies(store_data).get(session_id) or {}
|
||
messages = (store_data.get("messages") or {}).get(session_id) or []
|
||
by_id = {_message_id(m, idx): m for idx, m in enumerate(messages) if isinstance(m, dict)}
|
||
out: list[dict[str, Any]] = []
|
||
for item in policy.get("pinned") or []:
|
||
message = by_id.get(item.get("messageId"))
|
||
if message is None:
|
||
out.append({**copy.deepcopy(item), "missing": True, "text": None})
|
||
else:
|
||
out.append({
|
||
**copy.deepcopy(item),
|
||
"missing": False,
|
||
"role": message.get("role") or "user",
|
||
"text": str(message.get("text") or message.get("content") or ""),
|
||
})
|
||
return out
|
||
|
||
|
||
def _normalize_message(message: dict[str, Any], idx: int) -> dict[str, Any] | 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": _message_id(message, idx), "role": role, "text": text}
|
||
|
||
|
||
# ---------------- 四层组装 ----------------
|
||
|
||
def _fixed_layer(store_data: dict[str, Any]) -> str:
|
||
prompt = str(store_data.get("systemPrompt") or DEFAULT_SYSTEM_PROMPT)
|
||
return prompt.strip()
|
||
|
||
|
||
def _project_layer(store_data: dict[str, Any], session_id: str) -> str:
|
||
"""L-项目:会话归属项目的资料摘要(scope/共享上下文/基准版本)。"""
|
||
session = None
|
||
for row in store_data.get("sessions") or []:
|
||
if row.get("id") == session_id:
|
||
session = row
|
||
break
|
||
if session is None:
|
||
raise ValueError(f"session not found: {session_id!r}")
|
||
project_id = session.get("projectId") or "__personal__"
|
||
lines = [f"项目 ID:{project_id}"]
|
||
if project_id == "__personal__":
|
||
lines.append("作用域:个人空间(未绑定项目)")
|
||
return "\n".join(lines)
|
||
project = None
|
||
for row in store_data.get("projects") or []:
|
||
if row.get("id") == project_id:
|
||
project = row
|
||
break
|
||
if project is None:
|
||
raise ValueError(f"project not accessible: {project_id!r} for session {session_id!r}")
|
||
lines = [f"项目:{project.get('name') or project_id}",
|
||
f"作用域:{project.get('scopeLabel') or '未设定'}"]
|
||
shared = project.get("sharedContext") or []
|
||
if shared:
|
||
lines.append("共享上下文:" + ";".join(str(item) for item in shared))
|
||
for key in ("paramVersion", "baseVersion", "parametersVersion", "baselineVersion"):
|
||
if project.get(key):
|
||
lines.append(f"{key}:{project[key]}")
|
||
if project.get("archived"):
|
||
lines.append("注意:项目已归档,只读。")
|
||
return "\n".join(lines)
|
||
|
||
|
||
def _session_layer(store_data: dict[str, Any], session_id: str,
|
||
*, window: int = DEFAULT_WINDOW) -> tuple[str, list[dict[str, Any]], dict[str, Any] | None]:
|
||
"""L-会话:滚动摘要 + 钉住消息 + 最近 N 轮对话原文(滑动窗口)。"""
|
||
policy = _policies(store_data).get(session_id) or {}
|
||
summary = policy.get("summary")
|
||
pinned = _pinned_entries(store_data, session_id)
|
||
messages = (store_data.get("messages") or {}).get(session_id) or []
|
||
normalized = [item for item in (
|
||
_normalize_message(m, idx) for idx, m in enumerate(messages) if isinstance(m, dict)
|
||
) if item is not None]
|
||
|
||
parts: list[str] = []
|
||
if summary:
|
||
src = "、".join(str(sid) for sid in (summary.get("sourceMessageIds") or []))
|
||
parts.append("-- 滚动摘要 --")
|
||
parts.append(summary.get("text") or "")
|
||
parts.append(f"(摘要 {summary.get('summaryId')},来源消息:{src})")
|
||
if pinned:
|
||
parts.append("-- 钉住消息 --")
|
||
for item in pinned:
|
||
if item.get("missing"):
|
||
parts.append(f"[钉住 {item['messageId']} 已删除]"
|
||
+ (f"(备注:{item.get('note')})" if item.get("note") else ""))
|
||
else:
|
||
who = "用户" if item.get("role") == "user" else "助手"
|
||
note = f"(备注:{item.get('note')})" if item.get("note") else ""
|
||
parts.append(f"{who}:{item.get('text')}{note}")
|
||
if normalized:
|
||
parts.append("-- 最近对话 --")
|
||
pinned_ids = {item["messageId"] for item in pinned}
|
||
recent = [m for m in normalized[-window:] if m["messageId"] not in pinned_ids]
|
||
for message in recent:
|
||
who = "用户" if message["role"] == "user" else "助手"
|
||
parts.append(f"{who}:{message['text']}")
|
||
return "\n".join(parts), pinned, summary
|
||
|
||
|
||
def _retrieval_layer(store_data: dict[str, Any], query: str | None) -> str:
|
||
"""L-检索:RAG 命中 + 会话引用解析结果(按需注入)。"""
|
||
parts: list[str] = []
|
||
results = store_data.get("ragResults") or store_data.get("retrievalResults")
|
||
if results is None and query and store_data.get("ragAssets"):
|
||
try:
|
||
from server.knowledge.retrieval import hybrid_search
|
||
|
||
units = []
|
||
for asset in store_data.get("ragAssets") or []:
|
||
units.append({
|
||
"assetId": asset.get("assetId"), "title": asset.get("title"),
|
||
"kind": asset.get("kind"), "version": asset.get("version"),
|
||
"tags": asset.get("tags") or [],
|
||
"content": asset.get("content") or "",
|
||
})
|
||
results = hybrid_search(units, query)
|
||
except Exception: # noqa: BLE001 - 检索失败降级为空,绝不阻断上下文组装
|
||
results = []
|
||
if results:
|
||
parts.append("-- RAG 命中 --")
|
||
for hit in results:
|
||
title = hit.get("title") or hit.get("assetId") or "?"
|
||
version = f" v{hit.get('version')}" if hit.get("version") else ""
|
||
snippet = str(hit.get("snippet") or hit.get("content") or "")[:200]
|
||
parts.append(f"[{hit.get('assetId')}] {title}{version}:{snippet}")
|
||
refs = store_data.get("resolvedRefs")
|
||
if refs:
|
||
parts.append("-- 引用解析 --")
|
||
for ref in refs:
|
||
target = ref.get("target_id") or ref.get("targetId") or "?"
|
||
snapshot_at = ref.get("snapshot_at") or ref.get("snapshotAt") or "?"
|
||
parts.append(f"[引用 {ref.get('kind')}:{target} @{snapshot_at}]")
|
||
return "\n".join(parts) if parts else "(无检索命中)"
|
||
|
||
|
||
def assemble_context(store_data: dict[str, Any], session_id: str, query: str | None,
|
||
budget: ContextBudget, *, window: int = DEFAULT_WINDOW) -> ContextAssembly:
|
||
"""四层组装上下文;任何一层超预算 -> ContextBudgetExceeded(可预测拒绝)。
|
||
|
||
Args:
|
||
store_data: 工作区快照(sessions/projects/messages/contextPolicies/ragResults...)
|
||
session_id: 当前会话 ID
|
||
query: 用户问题(用于按需检索;可为 None)
|
||
budget: 四层 token 预算
|
||
window: 会话层最近 N 轮窗口
|
||
Returns:
|
||
ContextAssembly:layers/usage/budgets 可检查,text() 可注入 prompt
|
||
Raises:
|
||
ContextBudgetExceeded: 某层估算 token 超预算(fail closed)
|
||
ValueError: 会话/项目不存在
|
||
"""
|
||
layers: dict[str, str] = {}
|
||
usage: dict[str, int] = {}
|
||
budgets: dict[str, int] = {}
|
||
pinned: list[dict[str, Any]] = []
|
||
summary: dict[str, Any] | None = None
|
||
for layer in LAYERS:
|
||
if layer == "fixed":
|
||
text = _fixed_layer(store_data)
|
||
elif layer == "project":
|
||
text = _project_layer(store_data, session_id)
|
||
elif layer == "session":
|
||
text, pinned, summary = _session_layer(store_data, session_id, window=window)
|
||
else:
|
||
text = _retrieval_layer(store_data, query)
|
||
limit = budget.layer_limit(layer)
|
||
used = estimate_tokens(text)
|
||
if used > limit:
|
||
raise ContextBudgetExceeded(layer=layer, required=used, budget=limit)
|
||
layers[layer] = text
|
||
usage[layer] = used
|
||
budgets[layer] = limit
|
||
return ContextAssembly(
|
||
layers=layers, usage=usage, budgets=budgets,
|
||
pinned=pinned, summary=summary,
|
||
)
|