算法: - RouterCache:语义条目写入时预计算向量范数、语义查找单遍完成(消除命中后二次 O(N) 查找)、 相似度=1.0 提前终止;微基准(3000 条目×200 查询):3986ms -> 1685ms,2.37x - TaskGraph.topo_order:O(V²logV) 重排序/成员扫描 -> 邻接表+deque 的 O(V+E) Kahn, 输出顺序契约不变(初始就绪层按插入序、循环依赖按插入序兜底、未知依赖忽略) - RuleClassifier:同分决胜按领域名字典序(与规则表排列无关),次高分 O(n) 扫描 工程卫生: - .mimosa/(扫描器工作目录)加入 .gitignore 并移出索引 - test_review 抽样测试改用内联确定性 LCG,消除 2 个低危(不安全随机数) 测试:新增 11 项(topo 契约 6 + 缓存回归 3 + 分类器 2) pytest 230 passed(基线 219 全绿 + 11)
132 lines
5.0 KiB
Python
132 lines
5.0 KiB
Python
"""黑板(Blackboard)/ 工作记忆:专家系统风格的共享工作区(零依赖)。
|
|
|
|
- TaskNode:子任务节点(DAG 顶点),由 Planner 创建、Router 按拓扑序执行
|
|
- TaskGraph:子任务 DAG,提供拓扑排序与状态查询
|
|
- WorkingMemory:黑板,各知识源(执行器/规则)写入部分解,最后合并为最终答案
|
|
|
|
对齐《可行性调研与落地实现路线报告》第八章:
|
|
"黑板协作:多知识源(领域专家/执行器)通过共享黑板协作,而不是一个模型全包"。
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
from collections import deque
|
|
from dataclasses import dataclass, field
|
|
from typing import Any, Dict, List, Optional
|
|
|
|
|
|
@dataclass
|
|
class TaskNode:
|
|
"""一个子任务节点。"""
|
|
id: str
|
|
kind: str # analyze | design | implement | solve | diagnose | fix
|
|
# | retrieve | conclude | advise | explain | disclaimer | verify
|
|
domain: str
|
|
query: str # 子任务输入(通常为原始查询)
|
|
status: str = "pending" # pending | running | done | failed | skipped
|
|
output: Optional[str] = None
|
|
rule_trace: List[str] = field(default_factory=list)
|
|
deps: List[str] = field(default_factory=list)
|
|
desc: str = ""
|
|
error: Optional[str] = None
|
|
|
|
|
|
class TaskGraph:
|
|
"""子任务 DAG:节点 + 依赖边。"""
|
|
|
|
def __init__(self):
|
|
self._nodes: Dict[str, TaskNode] = {}
|
|
|
|
def add_node(self, node: TaskNode) -> None:
|
|
if node.id in self._nodes:
|
|
raise ValueError(f"节点 id 重复: {node.id}")
|
|
self._nodes[node.id] = node
|
|
|
|
def get(self, node_id: str) -> Optional[TaskNode]:
|
|
return self._nodes.get(node_id)
|
|
|
|
def nodes(self) -> List[TaskNode]:
|
|
return list(self._nodes.values())
|
|
|
|
def topo_order(self) -> List[TaskNode]:
|
|
"""Kahn 拓扑排序:依赖在前;初始就绪层按插入序稳定输出。
|
|
|
|
O(V+E) 实现(邻接表 + deque);循环依赖时按插入序兜底(不崩溃)。
|
|
"""
|
|
insert_pos = {nid: i for i, nid in enumerate(self._nodes)}
|
|
indeg: Dict[str, int] = {nid: 0 for nid in self._nodes}
|
|
dependents: Dict[str, List[str]] = {nid: [] for nid in self._nodes}
|
|
for n in self._nodes.values():
|
|
for d in n.deps:
|
|
if d in indeg: # 未知依赖 id 忽略(与入度统计口径一致)
|
|
indeg[n.id] += 1
|
|
dependents[d].append(n.id)
|
|
ready = deque(sorted((nid for nid, deg in indeg.items() if deg == 0),
|
|
key=insert_pos.__getitem__))
|
|
order_ids: List[str] = []
|
|
while ready:
|
|
nid = ready.popleft()
|
|
order_ids.append(nid)
|
|
for m in dependents[nid]:
|
|
indeg[m] -= 1
|
|
if indeg[m] == 0:
|
|
ready.append(m)
|
|
if len(order_ids) < len(self._nodes):
|
|
# 循环依赖兜底:剩余节点按插入序追加
|
|
placed = set(order_ids)
|
|
order_ids.extend(nid for nid in self._nodes if nid not in placed)
|
|
return [self._nodes[nid] for nid in order_ids]
|
|
|
|
def all_done(self) -> bool:
|
|
return all(n.status == "done" for n in self._nodes.values())
|
|
|
|
def failed(self) -> List[TaskNode]:
|
|
return [n for n in self._nodes.values() if n.status == "failed"]
|
|
|
|
def __len__(self) -> int:
|
|
return len(self._nodes)
|
|
|
|
|
|
class WorkingMemory:
|
|
"""黑板:facts(槽位事实)+ sections(章节部分解)+ trace(推理轨迹)。"""
|
|
|
|
def __init__(self):
|
|
self.facts: Dict[str, Any] = {}
|
|
self.sections: Dict[str, str] = {}
|
|
self.trace: List[str] = []
|
|
|
|
# ---- 事实 ----
|
|
def write_fact(self, key: str, value: Any, rule_id: Optional[str] = None) -> None:
|
|
if key in self.facts:
|
|
self.trace.append(f"overwrite:{key}@{rule_id or '?'}")
|
|
self.facts[key] = value
|
|
if rule_id:
|
|
self.trace.append(f"fact:{key}={str(value)[:40]}@rule:{rule_id}")
|
|
|
|
def get_fact(self, key: str, default: Any = None) -> Any:
|
|
return self.facts.get(key, default)
|
|
|
|
# ---- 章节 ----
|
|
def write_section(self, sid: str, text: str) -> None:
|
|
"""写入章节;同 id 覆盖(记录 trace)。"""
|
|
if sid in self.sections:
|
|
self.trace.append(f"overwrite_section:{sid}")
|
|
self.sections[sid] = text
|
|
|
|
def section(self, sid: str) -> Optional[str]:
|
|
return self.sections.get(sid)
|
|
|
|
def merge(self, order: Optional[List[str]] = None) -> str:
|
|
"""按 order(章节顺序)合并为最终答案;order 为空则按写入顺序。"""
|
|
if order:
|
|
parts = [self.sections[s] for s in order if s in self.sections]
|
|
if parts:
|
|
return "\n\n".join(parts)
|
|
return "\n\n".join(self.sections.values())
|
|
|
|
# ---- 轨迹 ----
|
|
def add_trace(self, item: str) -> None:
|
|
self.trace.append(item)
|
|
|
|
def explain(self) -> List[str]:
|
|
return list(self.trace)
|