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,
|
|||
|
|
)
|