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