Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b2fa8c3c81 | ||
|
|
e9cfb29b75 | ||
|
|
8d77c8c0c0 |
@@ -35,3 +35,6 @@ Thumbs.db
|
||||
config/model_pool.json
|
||||
agent_runs/
|
||||
agent_workspace/
|
||||
|
||||
# 安全扫描器工作目录(不入库)
|
||||
.mimosa/
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
# 分支:v2-coding-agent — 端云协同编程智能体(第二代)
|
||||
|
||||
> **快照点**:`747d85c`(v2 核心 + v3 Web 应用化 + v4 模型池与工具智能体全部完成、测试全绿时点)。
|
||||
> 历史路标,冻结不再演进;集成主线见 `master`。
|
||||
|
||||
## 这一代是什么
|
||||
|
||||
**命题**:《基于端云协同的编程智能体系统设计与实现》——大模型(API)任务分析/决策/终审 +
|
||||
小模型(本地 llama.cpp)实现/自验证 +「交流文本」结构化共享工作区 + 人工检验队列。
|
||||
|
||||
- **v2 核心**:`Workspace` 交流文本协议(schema/锚点/rollup/双渲染)→ `ArchitectClient`
|
||||
(brief/decide/final_review,JSON 约束)→ `WorkerLoop`(工具化实现+接地验证+自修≤2)
|
||||
→ `CollaborativePipeline`(快路径/协作循环/双护栏熔断/终审)+ 运维层(llama-server 进程管理、
|
||||
三档硬件模板)+ `ReviewQueue` 人工检验 + 打包分发
|
||||
- **v3 Web 应用化**:异步任务 + SSE 实时协作可视化 + Vue3 SPA(对话/协作过程/检验/指标)+ 网关安全加固
|
||||
- **v4 增补**:多价位模型池(local/budget/premium 角色指派)+ 工具智能体(harness 级工具/工作区选择/
|
||||
两级智能体/审批流/token 级流式)
|
||||
- v1 保留为 legacy(`POST /chat/legacy`),离线降级可用
|
||||
|
||||
## 基线
|
||||
|
||||
测试 281 项全绿(T31 时点);E1 token 经济学实测:交流文本较全量上下文降 ~61%(含缓存计费)。
|
||||
|
||||
## 文档
|
||||
|
||||
`实现方案_v2_端云协同编程智能体系统.md`、`实现方案_v3_Web应用化.md`、
|
||||
`实现方案_v4_模型池与工具智能体.md`、`毕业设计_进度记录.md`
|
||||
|
||||
## 与其他分支的关系
|
||||
|
||||
- 第一代(规则路由)以 legacy 形式包含在本快照内
|
||||
- 第三代(校园缓存代理层)在本快照之后的 master 上演进 → 见 `campus-cache-proxy` 分支
|
||||
@@ -211,7 +211,7 @@ class LlamaManager:
|
||||
args.extend(extra_args)
|
||||
|
||||
LOG_FILE.parent.mkdir(parents=True, exist_ok=True)
|
||||
log_f = open(LOG_FILE, "w", encoding="utf-8", buffering=1)
|
||||
log_f = LOG_FILE.open("w", encoding="utf-8", buffering=1)
|
||||
|
||||
try:
|
||||
self._proc = subprocess.Popen(
|
||||
|
||||
+48
-34
@@ -7,6 +7,11 @@
|
||||
高频语义命中会提升为 O(1) 的精确缓存条目。
|
||||
|
||||
只缓存"未升级"的结果(升级路径每次都走大模型,不缓存,避免陈旧)。
|
||||
|
||||
性能设计(2026-09 优化):
|
||||
- 每条语义缓存条目在写入时预计算并缓存向量范数,查询时免重复计算(原来每对比较都重算)
|
||||
- 语义查找单遍完成:扫描即跟踪最优条目与命中计数,命中后不再二次线性查找
|
||||
- 相似度达到 1.0(完全相同查询)时提前终止扫描(余弦相似度上界,不可能更优)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -29,18 +34,6 @@ def _ngrams(text: str, n: int = 3) -> List[str]:
|
||||
return [cleaned[i:i + n] for i in range(len(cleaned) - n + 1)]
|
||||
|
||||
|
||||
def _cosine(vec_a: Dict[str, float], vec_b: Dict[str, float]) -> float:
|
||||
if not vec_a or not vec_b:
|
||||
return 0.0
|
||||
common = set(vec_a) & set(vec_b)
|
||||
dot = sum(vec_a[k] * vec_b[k] for k in common)
|
||||
na = sum(v * v for v in vec_a.values()) ** 0.5
|
||||
nb = sum(v * v for v in vec_b.values()) ** 0.5
|
||||
if na == 0 or nb == 0:
|
||||
return 0.0
|
||||
return dot / (na * nb)
|
||||
|
||||
|
||||
def _tf_vector(grams: List[str]) -> Dict[str, float]:
|
||||
vec: Dict[str, float] = {}
|
||||
for g in grams:
|
||||
@@ -48,6 +41,17 @@ def _tf_vector(grams: List[str]) -> Dict[str, float]:
|
||||
return vec
|
||||
|
||||
|
||||
def _norm(vec: Dict[str, float]) -> float:
|
||||
return sum(v * v for v in vec.values()) ** 0.5
|
||||
|
||||
|
||||
def _dot(vec_a: Dict[str, float], vec_b: Dict[str, float]) -> float:
|
||||
"""点积:遍历较小的一方,另一侧用 get 兜底。"""
|
||||
if len(vec_a) > len(vec_b):
|
||||
vec_a, vec_b = vec_b, vec_a
|
||||
return sum(v * vec_b.get(k, 0.0) for k, v in vec_a.items())
|
||||
|
||||
|
||||
class RouterCache:
|
||||
"""L1 精确缓存 + L2 语义缓存。"""
|
||||
|
||||
@@ -61,6 +65,7 @@ class RouterCache:
|
||||
self._exact: Dict[str, CacheEntry] = {}
|
||||
self._semantic: List[Tuple[str, CacheEntry]] = [] # (query, entry)
|
||||
self._sem_vecs: Dict[str, Dict[str, float]] = {}
|
||||
self._sem_norms: Dict[str, float] = {} # 预计算范数,避免查询期重算
|
||||
self.hits = {"exact": 0, "semantic": 0}
|
||||
self.misses = 0
|
||||
|
||||
@@ -74,36 +79,41 @@ class RouterCache:
|
||||
|
||||
if self.semantic_enabled:
|
||||
q_vec = _tf_vector(_ngrams(query))
|
||||
q_norm = _norm(q_vec)
|
||||
best_sim = 0.0
|
||||
best_query: Optional[str] = None
|
||||
best_result: Optional[Dict[str, Any]] = None
|
||||
for q, e in self._semantic:
|
||||
sim = _cosine(q_vec, self._sem_vecs.get(q, {}))
|
||||
if sim > best_sim:
|
||||
best_sim = sim
|
||||
best_query = q
|
||||
best_result = e.result
|
||||
if best_query is not None and best_sim >= self.similarity_threshold:
|
||||
best_idx = -1
|
||||
if q_norm > 0.0:
|
||||
# 单遍扫描:同时跟踪最优相似度与条目位置
|
||||
for i, (q, _e) in enumerate(self._semantic):
|
||||
n_q = self._sem_norms.get(q, 0.0)
|
||||
if n_q <= 0.0:
|
||||
continue
|
||||
sim = _dot(q_vec, self._sem_vecs.get(q, {})) / (q_norm * n_q)
|
||||
if sim > best_sim:
|
||||
best_sim = sim
|
||||
best_idx = i
|
||||
if sim >= 1.0:
|
||||
break # 余弦相似度上界:完全相同查询,提前终止
|
||||
if best_idx >= 0 and best_sim >= self.similarity_threshold:
|
||||
best_q, best_entry = self._semantic[best_idx]
|
||||
# 完全相同查询(相似度=1.0)计为 exact 命中
|
||||
is_exact = best_sim >= 0.999
|
||||
level = "exact" if is_exact else "semantic"
|
||||
self.hits[level] += 1
|
||||
self._semantic_hit(best_query)
|
||||
return (level, best_result)
|
||||
self._bump_semantic(best_idx, best_q, best_entry)
|
||||
return (level, best_entry.result)
|
||||
|
||||
self.misses += 1
|
||||
return None
|
||||
|
||||
def _semantic_hit(self, query: str):
|
||||
"""语义命中:累计命中次数,达到阈值提升为精确缓存。"""
|
||||
for i, (q, e) in enumerate(self._semantic):
|
||||
if q == query:
|
||||
e.hits += 1
|
||||
if e.hits >= self.promote_frequency:
|
||||
self._exact[query] = e
|
||||
self._semantic.pop(i)
|
||||
self._sem_vecs.pop(query, None)
|
||||
break
|
||||
def _bump_semantic(self, idx: int, query: str, entry: CacheEntry):
|
||||
"""语义命中:累计命中次数,达到阈值提升为精确缓存(O(1),无需二次查找)。"""
|
||||
entry.hits += 1
|
||||
if entry.hits >= self.promote_frequency:
|
||||
self._exact[query] = entry
|
||||
self._semantic.pop(idx)
|
||||
self._sem_vecs.pop(query, None)
|
||||
self._sem_norms.pop(query, None)
|
||||
|
||||
# ---- 写入 ----
|
||||
def put(self, query: str, result: Dict[str, Any]):
|
||||
@@ -114,8 +124,11 @@ class RouterCache:
|
||||
if len(self._semantic) >= self.max_semantic:
|
||||
old_q, _ = self._semantic.pop(0)
|
||||
self._sem_vecs.pop(old_q, None)
|
||||
self._sem_norms.pop(old_q, None)
|
||||
self._semantic.append((query, entry))
|
||||
self._sem_vecs[query] = _tf_vector(_ngrams(query))
|
||||
vec = _tf_vector(_ngrams(query))
|
||||
self._sem_vecs[query] = vec
|
||||
self._sem_norms[query] = _norm(vec)
|
||||
else:
|
||||
self._exact[query] = entry
|
||||
if len(self._exact) > self.max_exact:
|
||||
@@ -137,5 +150,6 @@ class RouterCache:
|
||||
self._exact.clear()
|
||||
self._semantic.clear()
|
||||
self._sem_vecs.clear()
|
||||
self._sem_norms.clear()
|
||||
self.hits = {"exact": 0, "semantic": 0}
|
||||
self.misses = 0
|
||||
|
||||
@@ -186,7 +186,8 @@ class RuleClassifier(BaseClassifier):
|
||||
matched_rules=[],
|
||||
)
|
||||
|
||||
best_domain = max(raw, key=raw.get)
|
||||
# 同分决胜:按领域名字典序,保证与规则表排列顺序无关的确定性
|
||||
best_domain = max(sorted(raw), key=lambda d: raw[d])
|
||||
best_score = raw[best_domain]
|
||||
confidence = 1.0 - math.exp(-best_score)
|
||||
|
||||
@@ -196,7 +197,7 @@ class RuleClassifier(BaseClassifier):
|
||||
|
||||
# 与次高分的差距影响置信度(区分度)
|
||||
if len(raw) > 1:
|
||||
second = sorted(raw.values(), reverse=True)[1]
|
||||
second = max(v for d, v in raw.items() if d != best_domain)
|
||||
if second > 0.7 * best_score:
|
||||
confidence *= 0.85
|
||||
|
||||
|
||||
@@ -0,0 +1,370 @@
|
||||
"""执行器体系(L0 默认专家 + NodeExecutor 后端抽象)。
|
||||
|
||||
设计对齐"专家系统风格"(《可行性调研与落地实现路线报告》第八章):
|
||||
- 输出 = 结构化模板填充(回显查询、知识库事实、领域结构),不追求自然语言流畅度
|
||||
- 确定性:同输入 → 同输出(无采样随机)
|
||||
- 最小参数:零模型参数;L2 模式下同一节点可改由本地小模型执行(Router 按配置切换)
|
||||
|
||||
kind(子任务动作类型)与模板对应:
|
||||
analyze 需求/条件分析 | design 方案设计 | implement 代码实现 | solve 数学求解
|
||||
diagnose 错误定位 | fix 修复方案 | retrieve 知识检索 | conclude 结论
|
||||
advise 一般建议 | explain 展开解释 | disclaimer 免责/警示 | verify 自检
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from .experts import Expert, extract_content_terms
|
||||
from .knowledge import KnowledgeBase
|
||||
from .memory import TaskNode, WorkingMemory
|
||||
from .models import ExpertResponse
|
||||
|
||||
# 各领域"分析"步骤的目标描述
|
||||
_GOALS = {
|
||||
"code": "输出可运行的代码实现",
|
||||
"math": "得到问题的解并给出推导",
|
||||
"legal": "给出法律结论与依据",
|
||||
"medical": "给出科普性建议",
|
||||
"finance": "给出理财/金融建议与风险提示",
|
||||
"life": "给出实用生活建议",
|
||||
"education": "给出学习/行动方案",
|
||||
"general": "给出结构化说明",
|
||||
}
|
||||
|
||||
# 各领域"约束/边界"提示
|
||||
_CONSTRAINTS = {
|
||||
"code": "边界条件(空输入、极端值);复杂度目标",
|
||||
"math": "定义域、无解/多解情况、特殊值",
|
||||
"legal": "以现行有效法律为准,个案需咨询律师",
|
||||
"medical": "个体差异;非诊断,请遵医嘱",
|
||||
"finance": "市场有风险,投资需谨慎;不构成投资建议",
|
||||
"life": "结合个人实际情况,安全第一",
|
||||
"education": "结合个人基础与目标,循序渐进",
|
||||
"general": "围绕核心问题,避免无关展开",
|
||||
}
|
||||
|
||||
# 各领域"验证"清单
|
||||
_VERIFY_CHECKS = {
|
||||
"code": ["输入输出覆盖", "边界条件", "复杂度合理", "可运行性"],
|
||||
"math": ["中间步骤正确", "结果代入验证", "边界/特殊值", "单位与符号"],
|
||||
"legal": ["法条依据充分", "事实对应", "免责提示", "结论可执行"],
|
||||
"medical": ["建议有依据", "警示信号明确", "免责提示", "不构成诊断"],
|
||||
"finance": ["风险提示完整", "数据/规则准确", "免责提示", "建议可执行"],
|
||||
"life": ["建议实用", "安全提示", "贴合场景"],
|
||||
"education": ["方案可执行", "目标可衡量", "符合个人基础"],
|
||||
"general": ["要点覆盖", "逻辑连贯", "无事实错误"],
|
||||
}
|
||||
|
||||
|
||||
def _kw(query: str, n: int = 6) -> str:
|
||||
terms = extract_content_terms(query)
|
||||
return "、".join(terms[:n]) if terms else "该主题"
|
||||
|
||||
|
||||
class RuleExecutor(Expert):
|
||||
"""规则执行器:实现 Expert 接口;L0 模式的默认领域执行器。"""
|
||||
|
||||
name = "rule-executor"
|
||||
|
||||
def __init__(self, name: str = "rule-executor", domain: str = "general",
|
||||
kb: Optional[KnowledgeBase] = None):
|
||||
self.name = name
|
||||
self.domain = domain
|
||||
self.kb = kb
|
||||
|
||||
async def generate(self, query: str, difficulty: str,
|
||||
memory: Optional[WorkingMemory] = None,
|
||||
node: Optional[TaskNode] = None) -> ExpertResponse:
|
||||
"""按节点 kind 生成确定性输出。兼容 Expert 基类签名(后两参可选)。"""
|
||||
kind = node.kind if node is not None else "explain"
|
||||
domain = node.domain if node is not None else self.domain
|
||||
text = self._template(kind, domain, query, difficulty, memory)
|
||||
tokens = max(8, int(len(text) / 2.2))
|
||||
return ExpertResponse(
|
||||
text=text,
|
||||
model_used=f"rule:{domain}:{kind}",
|
||||
latency_ms=0.0,
|
||||
tokens=tokens,
|
||||
cost_est=0.0, # 零参数执行器无推理成本
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
def _template(self, kind: str, domain: str, query: str, difficulty: str,
|
||||
memory: Optional[WorkingMemory]) -> str:
|
||||
facts: Dict[str, Any] = memory.facts if memory else {}
|
||||
goal = _GOALS.get(domain, _GOALS["general"])
|
||||
constraints = _CONSTRAINTS.get(domain, _CONSTRAINTS["general"])
|
||||
kw = _kw(query)
|
||||
|
||||
if kind == "analyze":
|
||||
return (
|
||||
f"【{domain} 分析】\n"
|
||||
f"- 任务:{query}\n"
|
||||
f"- 关键要素:{kw}\n"
|
||||
f"- 目标:{goal}\n"
|
||||
f"- 约束/边界:{constraints}\n"
|
||||
f"- 难度评估:{difficulty}"
|
||||
)
|
||||
if kind == "design":
|
||||
return (
|
||||
f"【{domain} 方案设计】\n"
|
||||
f"针对「{query}」的设计思路:\n"
|
||||
f"1. 明确核心目标与验收标准\n"
|
||||
f"2. 选择合适的方法/数据结构(依据:{kw})\n"
|
||||
f"3. 拆解实现步骤并标注复杂度\n"
|
||||
f"4. 预留边界处理与异常路径\n"
|
||||
f"5. 设计自测用例(正常/边界/异常)"
|
||||
)
|
||||
if kind == "implement":
|
||||
return (
|
||||
f"【{domain} 实现】\n"
|
||||
f"```python\n"
|
||||
f"def solve() -> None:\n"
|
||||
f" # 关键点:{kw}\n"
|
||||
f" # 1. 校验输入与边界条件\n"
|
||||
f" # 2. 核心逻辑(依据 design 步骤)\n"
|
||||
f" # 3. 输出结果\n"
|
||||
f" pass\n"
|
||||
f"```\n"
|
||||
f"要点:{kw};复杂度与边界说明见 design/verify 步骤。"
|
||||
)
|
||||
if kind == "solve":
|
||||
return (
|
||||
f"【{domain} 求解】\n"
|
||||
f"题目:{query}\n"
|
||||
f"步骤:\n"
|
||||
f"1. 提取已知条件({kw})\n"
|
||||
f"2. 选择方法:代数变形/公式代入/逐步推导\n"
|
||||
f"3. 求解并化简中间结果\n"
|
||||
f"4. 检查特殊值与边界\n"
|
||||
f"结论:在标准假设下可得到闭合形式解;完整推导见正式解答。"
|
||||
)
|
||||
if kind == "diagnose":
|
||||
return (
|
||||
f"【{domain} 诊断】\n"
|
||||
f"错误现象:{query}\n"
|
||||
f"排查步骤:\n"
|
||||
f"1. 复现并定位出错行\n"
|
||||
f"2. 检查变量类型与取值(重点:{kw})\n"
|
||||
f"3. 核对函数签名、作用域与返回值\n"
|
||||
f"4. 打印中间变量验证假设\n"
|
||||
f"5. 用最小样例隔离问题"
|
||||
)
|
||||
if kind == "fix":
|
||||
return (
|
||||
f"【{domain} 修复方案】\n"
|
||||
f"针对「{query}」:\n"
|
||||
f"1. 根因:见 diagnose 步骤\n"
|
||||
f"2. 修复:调整类型/增加空值判断/修正逻辑分支\n"
|
||||
f"```python\n"
|
||||
f"def fixed() -> None:\n"
|
||||
f" # 修复点:{kw}\n"
|
||||
f" pass\n"
|
||||
f"```\n"
|
||||
f"3. 回归:补充对应单测后重跑"
|
||||
)
|
||||
if kind == "retrieve":
|
||||
return self._retrieve(domain, query, memory)
|
||||
if kind == "conclude":
|
||||
return (
|
||||
f"【{domain} 结论】\n"
|
||||
f"综合「{query}」:\n"
|
||||
f"1. 事实梳理:{kw}\n"
|
||||
f"2. 适用规则/依据(见 retrieve 步骤)\n"
|
||||
f"3. 结论:在所述前提下,按上述规则处理\n"
|
||||
f"4. 注意事项:个案差异,必要时咨询专业人士"
|
||||
)
|
||||
if kind == "advise":
|
||||
return (
|
||||
f"【{domain} 建议】\n"
|
||||
f"关于「{query}」的一般性建议:\n"
|
||||
f"1. 基础注意事项({kw})\n"
|
||||
f"2. 可操作建议:分步执行并观察效果\n"
|
||||
f"3. 警示信号:出现下列情况应及时就医(见 warning 步骤)"
|
||||
)
|
||||
if kind == "explain":
|
||||
if domain == "code":
|
||||
return (
|
||||
f"【code 代码讲解】\n"
|
||||
f"代码/片段:{query}\n"
|
||||
f"讲解结构:\n"
|
||||
f"1. 整体目的:这段代码要解决什么问题({kw})\n"
|
||||
f"2. 执行流程:按行/按函数梳理数据流与调用链\n"
|
||||
f"3. 关键点:数据结构、边界处理、异常路径\n"
|
||||
f"4. 可改进点:命名/复杂度/可读性建议"
|
||||
)
|
||||
return (
|
||||
f"【{domain} 说明】\n"
|
||||
f"主题:{query}\n"
|
||||
f"1. 背景与定义\n"
|
||||
f"2. 核心要点:{kw}\n"
|
||||
f"3. 分类/维度/机制\n"
|
||||
f"4. 实际应用与注意事项\n"
|
||||
f"如需更深入分析,可补充上下文。"
|
||||
)
|
||||
if kind == "disclaimer":
|
||||
if domain == "legal":
|
||||
return (
|
||||
"⚠️ 提示:以上为一般性法律分析,不构成正式法律意见;"
|
||||
"个案请咨询执业律师。"
|
||||
)
|
||||
if domain == "medical":
|
||||
return (
|
||||
"⚠️ 提示:以上内容仅供健康科普,不能替代医生诊断;"
|
||||
"如有不适请及时就医。"
|
||||
)
|
||||
if domain == "finance":
|
||||
return (
|
||||
"⚠️ 提示:以上为一般性金融科普,不构成投资建议;"
|
||||
"投资有风险,决策前请结合自身情况并咨询专业人士。"
|
||||
)
|
||||
return ""
|
||||
if kind == "verify":
|
||||
checks = _VERIFY_CHECKS.get(domain, _VERIFY_CHECKS["general"])
|
||||
items = "\n".join(f"- {c}" for c in checks)
|
||||
return f"【{domain} 自检】\n{items}"
|
||||
if kind == "refactor":
|
||||
return (
|
||||
f"【code 重构方案】\n"
|
||||
f"针对「{query}」:\n"
|
||||
f"1. 现状问题:重复代码/长函数/命名不清/耦合({kw})\n"
|
||||
f"2. 重构手法:提取函数、消除魔法数字、引入类或模块、统一命名\n"
|
||||
f"3. 目标结构:单一职责、清晰分层、可测试性\n"
|
||||
f"4. 验证:重构前后行为等价(跑通全部测试)"
|
||||
)
|
||||
if kind == "testcase":
|
||||
return (
|
||||
f"【code 测试用例】\n"
|
||||
f"针对「{query}」设计测试:\n"
|
||||
f"```python\n"
|
||||
f"def test_xxx():\n"
|
||||
f" # 正常路径:{kw}\n"
|
||||
f" pass\n\n"
|
||||
f"def test_edge():\n"
|
||||
f" # 边界:空输入/极值/None\n"
|
||||
f" pass\n\n"
|
||||
f"def test_error():\n"
|
||||
f" # 异常路径:非法参数\n"
|
||||
f" pass\n"
|
||||
f"```\n"
|
||||
f"覆盖策略:正常 + 边界 + 异常三组,断言明确"
|
||||
)
|
||||
if kind == "complexity":
|
||||
return (
|
||||
f"【code 复杂度分析】\n"
|
||||
f"针对「{query}」:\n"
|
||||
f"1. 时间复杂度:核心循环/递归层数 → 平均与最坏情况({kw})\n"
|
||||
f"2. 空间复杂度:辅助数据结构占用\n"
|
||||
f"3. 优化建议:若可接受,给出降复杂度的替代思路"
|
||||
)
|
||||
if kind == "optimize":
|
||||
return (
|
||||
f"【math 最优化求解】\n"
|
||||
f"问题:{query}\n"
|
||||
f"步骤:\n"
|
||||
f"1. 建立目标函数与约束({kw})\n"
|
||||
f"2. 求导/配方/不等式法找候选极值点\n"
|
||||
f"3. 比较候选值并与边界比较\n"
|
||||
f"4. 结论:给出最大值/最小值及取到条件"
|
||||
)
|
||||
if kind == "draft":
|
||||
return (
|
||||
f"【写作初稿】\n"
|
||||
f"主题:{query}\n"
|
||||
f"结构:\n"
|
||||
f"1. 开头:点明主题与背景({kw})\n"
|
||||
f"2. 主体:分点展开,每点配一个例子或依据\n"
|
||||
f"3. 结尾:总结观点 + 行动建议\n"
|
||||
f"(初稿完成,待 polish 步骤润色)"
|
||||
)
|
||||
if kind == "polish":
|
||||
return (
|
||||
f"【写作润色】\n"
|
||||
f"基于初稿检查:\n"
|
||||
f"1. 语法与错别字\n"
|
||||
f"2. 逻辑衔接与段落过渡\n"
|
||||
f"3. 语气统一(正式/亲切)与受众匹配\n"
|
||||
f"4. 长度控制与重点突出({kw})"
|
||||
)
|
||||
# 未知 kind 兜底
|
||||
return f"(规则执行器)「{query}」:{kw}"
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
def _retrieve(self, domain: str, query: str,
|
||||
memory: Optional[WorkingMemory]) -> str:
|
||||
"""知识检索:从知识库事实表取命中的条目;无命中则给出查阅建议。"""
|
||||
if self.kb is None:
|
||||
return (
|
||||
f"【{domain} 知识检索】\n"
|
||||
f"未配置知识库,建议查阅权威资料({_kw(query)})。"
|
||||
)
|
||||
facts = self.kb.facts(domain)
|
||||
hits = [f for f in facts if any(k in query for k in f.get("keywords", []))]
|
||||
if hits:
|
||||
lines = [f"- {f['statement']}" for f in hits]
|
||||
return f"【{domain} 知识检索】\n" + "\n".join(lines)
|
||||
return (
|
||||
f"【{domain} 知识检索】\n"
|
||||
f"未命中知识库条目;建议以现行有效法规/最新指南为准,"
|
||||
f"并结合个案情况分析({_kw(query)})。"
|
||||
)
|
||||
|
||||
|
||||
# ===============================================================
|
||||
# NodeExecutor:子任务执行后端抽象(T1:整体项目部分拆解·先行实现)
|
||||
#
|
||||
# Router._execute_node 不再内联 if-else 分支,而是依赖 NodeExecutor 接口:
|
||||
# - RuleNodeExecutor :L0 规则执行器(零参数、确定性)
|
||||
# - ModelNodeExecutor:L2 专家池小模型(≤8B,按需加载)
|
||||
# - 未来可加:多路采样执行器、API 执行器、组内模型执行器……
|
||||
# 工厂按配置选择后端,新增后端无需改动 Router。
|
||||
# ===============================================================
|
||||
|
||||
|
||||
class NodeExecutor:
|
||||
"""子任务执行后端抽象接口。"""
|
||||
|
||||
name: str = "node-executor"
|
||||
|
||||
async def execute(self, node: TaskNode, domain: str, difficulty: str,
|
||||
memory: WorkingMemory) -> ExpertResponse:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class RuleNodeExecutor(NodeExecutor):
|
||||
"""L0:规则执行器后端(零参数、确定性、零成本)。"""
|
||||
|
||||
name = "rule"
|
||||
|
||||
def __init__(self, kb: Optional[KnowledgeBase] = None):
|
||||
self._rule = RuleExecutor("rule-executor", "general", kb=kb)
|
||||
|
||||
async def execute(self, node: TaskNode, domain: str, difficulty: str,
|
||||
memory: WorkingMemory) -> ExpertResponse:
|
||||
return await self._rule.generate(node.query, difficulty, memory, node)
|
||||
|
||||
|
||||
class ModelNodeExecutor(NodeExecutor):
|
||||
"""L2:专家池小模型后端(≤8B;组内模型按需加载,用完即卸载由推理服务管理)。"""
|
||||
|
||||
name = "model"
|
||||
|
||||
def __init__(self, experts: Dict[str, Expert]):
|
||||
self._experts = experts
|
||||
|
||||
async def execute(self, node: TaskNode, domain: str, difficulty: str,
|
||||
memory: WorkingMemory) -> ExpertResponse:
|
||||
expert = self._experts.get(node.domain) or self._experts.get("general")
|
||||
return await expert.generate(node.query, difficulty)
|
||||
|
||||
|
||||
def build_node_executor(backend: str, kb: Optional[KnowledgeBase] = None,
|
||||
experts: Optional[Dict[str, Expert]] = None) -> NodeExecutor:
|
||||
"""按配置选择子任务执行后端。"""
|
||||
if backend == "rule":
|
||||
return RuleNodeExecutor(kb=kb)
|
||||
if backend in ("hf", "api", "model"):
|
||||
if not experts:
|
||||
raise ValueError("ModelNodeExecutor 需要专家池(experts)")
|
||||
return ModelNodeExecutor(experts)
|
||||
raise ValueError(f"未知执行后端: {backend}(支持 rule | hf | api | model)")
|
||||
@@ -0,0 +1,93 @@
|
||||
"""前向链推理机:知识库规则驱动的工作记忆演化(专家系统推理核心,零依赖)。
|
||||
|
||||
流程(经典前向链 forward chaining):
|
||||
1. 初始化黑板:写入领域/难度/置信度等事实
|
||||
2. 循环:在领域内匹配规则(未触发过的)→ 按优先级执行
|
||||
- 命中即记录轨迹 rule:<id>@<priority>
|
||||
- 规则带 output 模板 → 渲染后写入黑板章节(部分解)
|
||||
- 规则带 actions → 执行动作(写事实/写章节)
|
||||
3. 终止:无新规则可触发 / 达到步数上限(防死循环)
|
||||
|
||||
确定性保证:规则匹配基于子串包含,无随机性;同输入 → 同轨迹。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from .knowledge import KnowledgeBase, Rule
|
||||
from .memory import WorkingMemory
|
||||
|
||||
|
||||
def render_template(template: str, query: str, facts: Dict[str, Any]) -> str:
|
||||
"""渲染输出模板:替换 {query} 与 {facts.<key>} 占位符;缺失以 [未提供] 占位,不抛异常。"""
|
||||
out = template.replace("{query}", query)
|
||||
for key, value in facts.items():
|
||||
out = out.replace(f"{{facts.{key}}}", str(value))
|
||||
# 剩余占位符兜底
|
||||
while "{" in out and "}" in out:
|
||||
start = out.find("{")
|
||||
end = out.find("}", start)
|
||||
if end == -1:
|
||||
break
|
||||
out = out[:start] + "[未提供]" + out[end + 1:]
|
||||
return out
|
||||
|
||||
|
||||
class InferenceEngine:
|
||||
"""前向链推理机。"""
|
||||
|
||||
def __init__(self, kb: KnowledgeBase, max_steps: int = 20):
|
||||
self.kb = kb
|
||||
self.max_steps = max_steps
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
def initialize(self, query: str, domain: str, difficulty: str,
|
||||
confidence: float, memory: WorkingMemory) -> None:
|
||||
"""把分类结果写入黑板(事实初始化)。"""
|
||||
memory.write_fact("query", query)
|
||||
memory.write_fact("domain", domain)
|
||||
memory.write_fact("difficulty", difficulty)
|
||||
memory.write_fact("confidence", round(confidence, 4))
|
||||
memory.add_trace(f"init:domain={domain},difficulty={difficulty},conf={confidence:.2f}")
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
def run(self, query: str, domain: str, memory: WorkingMemory,
|
||||
max_steps: Optional[int] = None) -> List[str]:
|
||||
"""前向链主循环。返回触发规则 id 列表(按触发顺序)。"""
|
||||
steps = max_steps or self.max_steps
|
||||
fired: List[str] = []
|
||||
for _ in range(steps):
|
||||
rules = self.kb.match(query, domain=domain)
|
||||
# 选第一个"未触发过"的规则
|
||||
target: Optional[Rule] = None
|
||||
for r in rules:
|
||||
if r.id not in fired:
|
||||
target = r
|
||||
break
|
||||
if target is None:
|
||||
break # 无新规则可触发 → 终止
|
||||
fired.append(target.id)
|
||||
self._fire(target, query, memory)
|
||||
return fired
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
def _fire(self, rule: Rule, query: str, memory: WorkingMemory) -> None:
|
||||
"""执行一条规则:记录轨迹 + 写事实 + 产出章节。"""
|
||||
memory.add_trace(f"rule:{rule.id}@{rule.priority}")
|
||||
# 规则动作
|
||||
for action in rule.actions:
|
||||
self._apply_action(action, rule, query, memory)
|
||||
# 规则输出模板 → 章节
|
||||
if rule.output:
|
||||
text = render_template(rule.output, query, memory.facts)
|
||||
memory.write_section(rule.id, text)
|
||||
|
||||
def _apply_action(self, action: str, rule: Rule, query: str,
|
||||
memory: WorkingMemory) -> None:
|
||||
"""动作格式:write_fact:key=value(value 支持 {query} 占位)。"""
|
||||
if action.startswith("write_fact:"):
|
||||
kv = action[len("write_fact:"):]
|
||||
key, _, value = kv.partition("=")
|
||||
value = value.replace("{query}", query)
|
||||
memory.write_fact(key.strip(), value.strip(), rule_id=rule.id)
|
||||
# 其他动作类型暂不实现(保留扩展位)
|
||||
@@ -0,0 +1,490 @@
|
||||
"""知识库:专家系统风格的规则与知识表示(零依赖,纯标准库)。
|
||||
|
||||
设计原则(对齐《可行性调研与落地实现路线报告》第八章"专家系统内核"):
|
||||
- 领域知识显式化:写在规则文件里(config/knowledge/<domain>.yaml),不藏在模型参数中
|
||||
- 确定性:规则匹配 = 子串包含(大小写不敏感),同输入同输出
|
||||
- 可解释:每次命中都记录规则 id,形成推理轨迹
|
||||
- 最小参数:L0 模式零模型参数,规则即知识
|
||||
|
||||
规则文件格式(YAML;若 pyyaml 不可用,可提供同名 .json):
|
||||
domain: code
|
||||
rules:
|
||||
- id: code-sort
|
||||
priority: 90 # 越大越先触发
|
||||
patterns: ["排序", "sort"] # 任一子串命中即触发
|
||||
template: code-implement # 可选:Planner 任务模板 id
|
||||
output: | # 可选:输出模板({query} 等占位符)
|
||||
(规则输出)...
|
||||
facts: # 领域事实表(Judge 校验 / retrieve 执行器用)
|
||||
- id: legal-nc
|
||||
keywords: ["竞业"]
|
||||
statement: "竞业限制期限不得超过二年"
|
||||
|
||||
任务模板(config/knowledge/tasks.yaml):
|
||||
task_templates:
|
||||
code-implement:
|
||||
steps:
|
||||
- {id: analyze, kind: analyze, domain: code}
|
||||
- {id: design, kind: design, domain: code, deps: [analyze]}
|
||||
|
||||
加载顺序:内置默认规则(代码内兜底)→ 文件规则按 id 合并覆盖。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
DEFAULT_RULES_DIR = Path(__file__).resolve().parent.parent / "config" / "knowledge"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Rule:
|
||||
"""一条领域规则。"""
|
||||
id: str
|
||||
domain: str
|
||||
priority: int = 50
|
||||
patterns: List[str] = field(default_factory=list)
|
||||
template: Optional[str] = None # 引用的任务模板 id
|
||||
output: Optional[str] = None # 输出模板
|
||||
actions: List[str] = field(default_factory=list) # 保留字段:动作扩展
|
||||
subdomain: Optional[str] = None # 二级子领域(如 investing/labor/calculus)
|
||||
subdomain2: Optional[str] = None # 三级子领域(如 fund/overtime/sorting)
|
||||
|
||||
def matches(self, text: str) -> bool:
|
||||
"""任一 pattern 是 text 的子串即命中(大小写不敏感)。"""
|
||||
if not self.patterns:
|
||||
return False
|
||||
q = text.lower()
|
||||
return any(p.lower() in q for p in self.patterns)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# 三级子领域映射(rule_id -> subdomain2)
|
||||
# 集中维护:新增规则时在此加一行即可完成三级细化标注
|
||||
# ---------------------------------------------------------------
|
||||
SUBDOMAIN2_MAP: Dict[str, str] = {
|
||||
# ---- code ----
|
||||
"code-sort": "sorting",
|
||||
"code-debug": "error-analysis",
|
||||
"code-algorithm": "algorithm-general",
|
||||
"code-refactor": "code-quality",
|
||||
"code-database": "sql",
|
||||
"code-explain": "code-reading",
|
||||
"code-test": "unit-test",
|
||||
"code-web": "web-dev",
|
||||
"code-implement-general": "implementation",
|
||||
"code-git-knowledge": "git",
|
||||
"code-docker-knowledge": "container",
|
||||
"code-python-knowledge": "python-env",
|
||||
# ---- math ----
|
||||
"math-equation": "equation",
|
||||
"math-calculus": "calculus",
|
||||
"math-algebra": "algebra",
|
||||
"math-geometry": "geometry",
|
||||
"math-proof": "proof",
|
||||
"math-probability": "probability",
|
||||
"math-number-theory": "number-theory",
|
||||
"math-trigonometry": "trigonometry",
|
||||
"math-optimization": "optimization",
|
||||
"math-general": "math-general",
|
||||
# ---- legal ----
|
||||
"legal-contract": "contract",
|
||||
"legal-labor": "labor",
|
||||
"legal-ip": "intellectual-property",
|
||||
"legal-housing": "housing",
|
||||
"legal-marriage": "family-law",
|
||||
"legal-tax": "tax",
|
||||
"legal-consumer": "consumer-rights",
|
||||
"legal-litigation": "litigation",
|
||||
"legal-compliance": "compliance",
|
||||
"legal-general": "legal-general",
|
||||
# ---- medical ----
|
||||
"medical-hypertension": "hypertension",
|
||||
"medical-drug": "medication",
|
||||
"medical-common": "common-illness",
|
||||
"medical-chronic": "chronic-disease",
|
||||
"medical-digestive": "digestive",
|
||||
"medical-nutrition": "nutrition",
|
||||
"medical-mental": "mental-health",
|
||||
"medical-firstaid": "first-aid",
|
||||
"medical-pediatrics": "pediatrics",
|
||||
"medical-general": "medical-general",
|
||||
# ---- finance ----
|
||||
"finance-investing": "investing",
|
||||
"finance-saving": "saving",
|
||||
"finance-loan": "loan",
|
||||
"finance-insurance": "insurance",
|
||||
"finance-credit-card": "credit",
|
||||
"finance-personal-budget": "budgeting",
|
||||
"finance-general": "finance-general",
|
||||
# ---- life ----
|
||||
"life-food": "cooking",
|
||||
"life-travel": "travel",
|
||||
"life-home": "home",
|
||||
"life-pet": "pet",
|
||||
"life-fitness": "fitness",
|
||||
"life-weather": "weather",
|
||||
"life-general": "life-general",
|
||||
# ---- education ----
|
||||
"edu-study-method": "study-method",
|
||||
"edu-exam": "exam",
|
||||
"edu-language": "language",
|
||||
"edu-course": "course",
|
||||
"edu-career": "career",
|
||||
"edu-general": "education-general",
|
||||
# ---- general ----
|
||||
"general-explain": "explain",
|
||||
"general-writing": "writing",
|
||||
"general-compare": "compare",
|
||||
"general-translate": "translate",
|
||||
"general-knowledge": "explain",
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# 二级子领域映射(rule_id -> subdomain)
|
||||
# 三级 subdomain2 的父级类别;与 SUBDOMAIN2_MAP 按 rule_id 对齐维护。
|
||||
# ---------------------------------------------------------------
|
||||
SUBDOMAIN_MAP: Dict[str, str] = {
|
||||
# ---- code ----
|
||||
"code-sort": "algorithm",
|
||||
"code-debug": "debugging",
|
||||
"code-algorithm": "algorithm",
|
||||
"code-refactor": "quality",
|
||||
"code-database": "data",
|
||||
"code-explain": "reading",
|
||||
"code-test": "quality",
|
||||
"code-web": "web",
|
||||
"code-implement-general": "implementation",
|
||||
"code-git-knowledge": "tooling",
|
||||
"code-docker-knowledge": "tooling",
|
||||
"code-python-knowledge": "tooling",
|
||||
# ---- math ----
|
||||
"math-equation": "algebra",
|
||||
"math-calculus": "analysis",
|
||||
"math-algebra": "algebra",
|
||||
"math-geometry": "geometry",
|
||||
"math-proof": "proof",
|
||||
"math-probability": "probability",
|
||||
"math-number-theory": "number-theory",
|
||||
"math-trigonometry": "trigonometry",
|
||||
"math-optimization": "optimization",
|
||||
"math-general": "general",
|
||||
# ---- legal ----
|
||||
"legal-contract": "contract",
|
||||
"legal-labor": "labor",
|
||||
"legal-ip": "ip",
|
||||
"legal-housing": "civil",
|
||||
"legal-marriage": "civil",
|
||||
"legal-tax": "tax",
|
||||
"legal-consumer": "consumer",
|
||||
"legal-litigation": "procedure",
|
||||
"legal-compliance": "compliance",
|
||||
"legal-general": "general",
|
||||
# ---- medical ----
|
||||
"medical-hypertension": "chronic",
|
||||
"medical-drug": "medication",
|
||||
"medical-common": "common",
|
||||
"medical-chronic": "chronic",
|
||||
"medical-digestive": "common",
|
||||
"medical-nutrition": "nutrition",
|
||||
"medical-mental": "mental",
|
||||
"medical-firstaid": "emergency",
|
||||
"medical-pediatrics": "pediatrics",
|
||||
"medical-general": "general",
|
||||
# ---- finance ----
|
||||
"finance-investing": "investing",
|
||||
"finance-saving": "personal-finance",
|
||||
"finance-loan": "credit",
|
||||
"finance-insurance": "insurance",
|
||||
"finance-credit-card": "credit",
|
||||
"finance-personal-budget": "personal-finance",
|
||||
"finance-general": "general",
|
||||
# ---- life ----
|
||||
"life-food": "daily",
|
||||
"life-travel": "daily",
|
||||
"life-home": "daily",
|
||||
"life-pet": "daily",
|
||||
"life-fitness": "health",
|
||||
"life-weather": "daily",
|
||||
"life-general": "general",
|
||||
# ---- education ----
|
||||
"edu-study-method": "learning",
|
||||
"edu-exam": "learning",
|
||||
"edu-language": "language",
|
||||
"edu-course": "learning",
|
||||
"edu-career": "development",
|
||||
"edu-general": "general",
|
||||
# ---- general ----
|
||||
"general-explain": "explanation",
|
||||
"general-writing": "writing",
|
||||
"general-compare": "analysis",
|
||||
"general-translate": "language",
|
||||
"general-knowledge": "explanation",
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# 内置默认规则(兜底:即使规则文件缺失/损坏,系统仍可运行)
|
||||
# ---------------------------------------------------------------
|
||||
BUILTIN_RULES: List[Dict[str, Any]] = [
|
||||
# ---- code ----
|
||||
{"id": "code-sort", "domain": "code", "priority": 90,
|
||||
"patterns": ["排序", "快速排序", "排序算法", "sort", "quicksort"],
|
||||
"template": "code-implement"},
|
||||
{"id": "code-debug", "domain": "code", "priority": 85,
|
||||
"patterns": ["报错", "错误", "调试", "bug", "debug", "typeerror", "异常", "报 TypeError"],
|
||||
"template": "code-debug"},
|
||||
{"id": "code-implement-general", "domain": "code", "priority": 50,
|
||||
"patterns": ["实现", "编写", "写一个", "函数", "代码", "编程", "用 python", "用 java",
|
||||
"用 javascript", "sql", "接口", "算法"],
|
||||
"template": "code-implement"},
|
||||
# ---- math ----
|
||||
{"id": "math-equation", "domain": "math", "priority": 90,
|
||||
"patterns": ["方程", "求解", "求根", "solve", "equation", "解方程"],
|
||||
"template": "math-solve"},
|
||||
{"id": "math-calculus", "domain": "math", "priority": 85,
|
||||
"patterns": ["积分", "导数", "微积分", "求导", "integral", "derivative", "∫"],
|
||||
"template": "math-solve"},
|
||||
{"id": "math-general", "domain": "math", "priority": 50,
|
||||
"patterns": ["数学", "证明", "定理", "概率", "统计", "计算", "等于", "math", "不等式"],
|
||||
"template": "math-solve"},
|
||||
# ---- legal ----
|
||||
{"id": "legal-contract", "domain": "legal", "priority": 90,
|
||||
"patterns": ["合同", "条款", "违约", "离职", "竞业", "劳动", "contract", "clause", "赔偿"],
|
||||
"template": "legal-advice"},
|
||||
{"id": "legal-ip", "domain": "legal", "priority": 85,
|
||||
"patterns": ["专利", "版权", "商标", "知识产权", "patent", "copyright", "trademark"],
|
||||
"template": "legal-advice"},
|
||||
{"id": "legal-general", "domain": "legal", "priority": 50,
|
||||
"patterns": ["法律", "合规", "诉讼", "仲裁", "法条", "law", "legal", "法规"],
|
||||
"template": "legal-advice"},
|
||||
# ---- medical ----
|
||||
{"id": "medical-hypertension", "domain": "medical", "priority": 90,
|
||||
"patterns": ["高血压", "hypertension", "血压"],
|
||||
"template": "medical-advice"},
|
||||
{"id": "medical-drug", "domain": "medical", "priority": 85,
|
||||
"patterns": ["药物", "吃药", "剂量", "副作用", "退烧药", "降压药", "dosage", "prescription"],
|
||||
"template": "medical-advice"},
|
||||
{"id": "medical-general", "domain": "medical", "priority": 50,
|
||||
"patterns": ["医疗", "症状", "诊断", "治疗", "感冒", "发烧", "糖尿病", "医生", "患者",
|
||||
"体检", "疫苗", "medical", "symptom", "disease"],
|
||||
"template": "medical-advice"},
|
||||
# ---- finance ----
|
||||
{"id": "finance-investing", "domain": "finance", "priority": 90,
|
||||
"patterns": ["基金", "定投", "收益率", "股票", "投资", "炒股", "证券", "invest", "stock"]},
|
||||
{"id": "finance-saving", "domain": "finance", "priority": 85,
|
||||
"patterns": ["存款", "储蓄", "利息", "零钱通", "余额宝", "saving"]},
|
||||
{"id": "finance-loan", "domain": "finance", "priority": 80,
|
||||
"patterns": ["贷款", "房贷", "借款", "按揭", "loan"]},
|
||||
{"id": "finance-insurance", "domain": "finance", "priority": 75,
|
||||
"patterns": ["保险", "理赔", "保单", "投保", "insurance"]},
|
||||
{"id": "finance-credit-card", "domain": "finance", "priority": 70,
|
||||
"patterns": ["信用卡", "花呗", "白条", "credit card"]},
|
||||
{"id": "finance-personal-budget", "domain": "finance", "priority": 60,
|
||||
"patterns": ["预算", "记账", "开销", "省钱", "budget"]},
|
||||
{"id": "finance-general", "domain": "finance", "priority": 50,
|
||||
"patterns": ["金融", "财务", "外汇", "汇率", "finance"]},
|
||||
# ---- life ----
|
||||
{"id": "life-food", "domain": "life", "priority": 90,
|
||||
"patterns": ["做饭", "做菜", "菜谱", "食谱", "烹饪", "cooking"]},
|
||||
{"id": "life-travel", "domain": "life", "priority": 85,
|
||||
"patterns": ["旅游", "旅行", "攻略", "景点", "签证", "travel"]},
|
||||
{"id": "life-home", "domain": "life", "priority": 80,
|
||||
"patterns": ["装修", "租房", "家电", "清洁", "搬家", "home"]},
|
||||
{"id": "life-pet", "domain": "life", "priority": 75,
|
||||
"patterns": ["宠物", "养猫", "养狗", "撸猫", "pet"]},
|
||||
{"id": "life-fitness", "domain": "life", "priority": 70,
|
||||
"patterns": ["健身", "减肥", "跑步", "锻炼", "fitness"]},
|
||||
{"id": "life-weather", "domain": "life", "priority": 65,
|
||||
"patterns": ["天气", "下雨", "台风", "降温", "weather"]},
|
||||
{"id": "life-general", "domain": "life", "priority": 50,
|
||||
"patterns": ["生活", "日常", "家居", "life"]},
|
||||
# ---- education ----
|
||||
{"id": "edu-study-method", "domain": "education", "priority": 90,
|
||||
"patterns": ["学习方法", "记忆", "做笔记", "笔记法", "专注力"]},
|
||||
{"id": "edu-exam", "domain": "education", "priority": 85,
|
||||
"patterns": ["考试", "考研", "复习", "真题", "四六级", "exam"]},
|
||||
{"id": "edu-language", "domain": "education", "priority": 80,
|
||||
"patterns": ["英语", "单词", "口语", "语法", "english"]},
|
||||
{"id": "edu-course", "domain": "education", "priority": 75,
|
||||
"patterns": ["课程", "网课", "慕课", "选修", "course"]},
|
||||
{"id": "edu-career", "domain": "education", "priority": 70,
|
||||
"patterns": ["职业规划", "求职", "面试", "简历", "校招", "career"]},
|
||||
{"id": "edu-general", "domain": "education", "priority": 50,
|
||||
"patterns": ["教育", "大学", "专业选择", "education"]},
|
||||
# ---- general ----
|
||||
{"id": "general-explain", "domain": "general", "priority": 30,
|
||||
"patterns": ["总结", "介绍", "解释", "为什么", "优缺点", "是什么", "翻译", "邮件",
|
||||
"summarize", "explain", "what is", "写一封"],
|
||||
"template": "general-explain"},
|
||||
]
|
||||
|
||||
# 内置默认任务模板(兜底)
|
||||
BUILTIN_TASKS: Dict[str, Dict[str, Any]] = {
|
||||
"code-implement": {"steps": [
|
||||
{"id": "analyze", "kind": "analyze", "domain": "code", "desc": "需求与约束分析"},
|
||||
{"id": "design", "kind": "design", "domain": "code", "deps": ["analyze"], "desc": "算法与数据结构设计"},
|
||||
{"id": "implement", "kind": "implement", "domain": "code", "deps": ["design"], "desc": "实现代码"},
|
||||
{"id": "verify", "kind": "verify", "domain": "code", "deps": ["implement"], "desc": "自测校验"},
|
||||
]},
|
||||
"code-debug": {"steps": [
|
||||
{"id": "analyze", "kind": "analyze", "domain": "code", "desc": "错误现象与复现分析"},
|
||||
{"id": "diagnose", "kind": "diagnose", "domain": "code", "deps": ["analyze"], "desc": "定位错误根因"},
|
||||
{"id": "fix", "kind": "fix", "domain": "code", "deps": ["diagnose"], "desc": "给出修复方案"},
|
||||
{"id": "verify", "kind": "verify", "domain": "code", "deps": ["fix"], "desc": "修复后验证"},
|
||||
]},
|
||||
"math-solve": {"steps": [
|
||||
{"id": "conditions", "kind": "analyze", "domain": "math", "desc": "明确已知条件与目标"},
|
||||
{"id": "solve", "kind": "solve", "domain": "math", "deps": ["conditions"], "desc": "选择方法并求解"},
|
||||
{"id": "verify", "kind": "verify", "domain": "math", "deps": ["solve"], "desc": "检查边界与验证"},
|
||||
]},
|
||||
"legal-advice": {"steps": [
|
||||
{"id": "facts", "kind": "analyze", "domain": "legal", "desc": "梳理事实与法律问题"},
|
||||
{"id": "retrieve", "kind": "retrieve", "domain": "legal", "deps": ["facts"], "desc": "检索适用法规"},
|
||||
{"id": "conclude", "kind": "conclude", "domain": "legal", "deps": ["retrieve"], "desc": "给出法律意见"},
|
||||
{"id": "disclaimer", "kind": "disclaimer", "domain": "legal", "deps": ["conclude"], "desc": "免责提示"},
|
||||
]},
|
||||
"medical-advice": {"steps": [
|
||||
{"id": "symptoms", "kind": "analyze", "domain": "medical", "desc": "梳理症状与背景"},
|
||||
{"id": "advise", "kind": "advise", "domain": "medical", "deps": ["symptoms"], "desc": "给出一般建议"},
|
||||
{"id": "warning", "kind": "disclaimer", "domain": "medical", "deps": ["advise"], "desc": "就医警示"},
|
||||
]},
|
||||
"general-explain": {"steps": [
|
||||
{"id": "outline", "kind": "analyze", "domain": "general", "desc": "梳理主题要点"},
|
||||
{"id": "explain", "kind": "explain", "domain": "general", "deps": ["outline"], "desc": "展开解释"},
|
||||
{"id": "conclude", "kind": "conclude", "domain": "general", "deps": ["explain"], "desc": "总结"},
|
||||
]},
|
||||
}
|
||||
|
||||
# 内置默认事实表(兜底)
|
||||
BUILTIN_FACTS: Dict[str, List[Dict[str, Any]]] = {
|
||||
"legal": [
|
||||
{"id": "legal-noncompete", "keywords": ["竞业", "离职", "同业"],
|
||||
"statement": "竞业限制期限不得超过二年,且用人单位应在限制期内按月给予经济补偿"},
|
||||
{"id": "legal-renew-compensation", "keywords": ["不续签", "经济补偿", "劳动合同"],
|
||||
"statement": "劳动合同期满用人单位不续签的,通常应支付经济补偿(每满一年一个月工资)"},
|
||||
],
|
||||
"medical": [
|
||||
{"id": "medical-hypertension-diet", "keywords": ["高血压", "饮食"],
|
||||
"statement": "高血压患者应低盐低脂饮食、控制体重、规律运动、戒烟限酒,并在医生指导下用药"},
|
||||
{"id": "medical-fever-drug", "keywords": ["发烧", "退烧"],
|
||||
"statement": "体温超过 38.5℃ 可在药师指导下使用退烧药;持续发热或出现严重症状应及时就医"},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def _try_load_yaml(path: Path) -> Optional[Dict[str, Any]]:
|
||||
try:
|
||||
import yaml # type: ignore
|
||||
except ImportError:
|
||||
return None
|
||||
try:
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
data = yaml.safe_load(f)
|
||||
return data if isinstance(data, dict) else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _try_load_json(path: Path) -> Optional[Dict[str, Any]]:
|
||||
json_path = path.with_suffix(".json")
|
||||
if not json_path.exists():
|
||||
return None
|
||||
try:
|
||||
with open(json_path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
return data if isinstance(data, dict) else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
class KnowledgeBase:
|
||||
"""知识库:加载规则文件,提供规则匹配、任务模板、事实表查询。"""
|
||||
|
||||
def __init__(self, rules_dir: Optional[str | Path] = None):
|
||||
self.rules_dir = Path(rules_dir) if rules_dir else DEFAULT_RULES_DIR
|
||||
self._rules: Dict[str, Rule] = {}
|
||||
self._tasks: Dict[str, Dict[str, Any]] = {}
|
||||
self._facts: Dict[str, List[Dict[str, Any]]] = {}
|
||||
self.load()
|
||||
|
||||
# ---- 加载 ----
|
||||
def load(self) -> None:
|
||||
"""内置默认 + 规则文件合并(文件规则按 id 覆盖内置)。"""
|
||||
self._rules = {}
|
||||
self._tasks = dict(BUILTIN_TASKS)
|
||||
for item in BUILTIN_RULES:
|
||||
self._register_rule(item)
|
||||
self._facts = {d: [dict(f) for f in facts] for d, facts in BUILTIN_FACTS.items()}
|
||||
|
||||
if self.rules_dir.is_dir():
|
||||
for f in sorted(self.rules_dir.glob("*.yaml")):
|
||||
data = _try_load_yaml(f)
|
||||
if data is not None:
|
||||
self._load_file_data(f, data)
|
||||
for f in sorted(self.rules_dir.glob("*.json")):
|
||||
if f.name not in {p.name for p in self.rules_dir.glob("*.yaml")}:
|
||||
data = _try_load_json(f)
|
||||
if data is not None:
|
||||
self._load_file_data(f, data)
|
||||
|
||||
def _load_file_data(self, path: Path, data: Dict[str, Any]) -> None:
|
||||
name = path.stem
|
||||
if name == "tasks":
|
||||
for tid, tpl in (data.get("task_templates") or {}).items():
|
||||
if isinstance(tpl, dict) and isinstance(tpl.get("steps"), list):
|
||||
self._tasks[tid] = tpl
|
||||
return
|
||||
domain = data.get("domain", name)
|
||||
for item in data.get("rules") or []:
|
||||
if isinstance(item, dict) and item.get("id"):
|
||||
self._register_rule({**item, "domain": domain})
|
||||
for fact in data.get("facts") or []:
|
||||
if isinstance(fact, dict) and fact.get("id"):
|
||||
self._facts.setdefault(domain, []).append(fact)
|
||||
|
||||
def _register_rule(self, item: Dict[str, Any]) -> None:
|
||||
rule = Rule(
|
||||
id=str(item["id"]),
|
||||
domain=str(item.get("domain", "general")),
|
||||
priority=int(item.get("priority", 50)),
|
||||
patterns=[str(p) for p in item.get("patterns", [])],
|
||||
template=item.get("template"),
|
||||
output=item.get("output"),
|
||||
actions=[str(a) for a in item.get("actions", [])],
|
||||
subdomain=item.get("subdomain") or SUBDOMAIN_MAP.get(str(item["id"])),
|
||||
subdomain2=item.get("subdomain2") or SUBDOMAIN2_MAP.get(str(item["id"])),
|
||||
)
|
||||
self._rules[rule.id] = rule
|
||||
|
||||
# ---- 查询 ----
|
||||
def match(self, text: str, domain: Optional[str] = None) -> List[Rule]:
|
||||
"""返回命中的规则,按优先级降序。domain 为空则全领域匹配。"""
|
||||
hits = []
|
||||
for rule in self._rules.values():
|
||||
if domain is not None and rule.domain != domain:
|
||||
continue
|
||||
if rule.matches(text):
|
||||
hits.append(rule)
|
||||
hits.sort(key=lambda r: r.priority, reverse=True)
|
||||
return hits
|
||||
|
||||
def rule(self, rule_id: str) -> Optional[Rule]:
|
||||
return self._rules.get(rule_id)
|
||||
|
||||
def rules_count(self) -> int:
|
||||
return len(self._rules)
|
||||
|
||||
def task_template(self, tid: str) -> Optional[Dict[str, Any]]:
|
||||
return self._tasks.get(tid)
|
||||
|
||||
def task_ids(self) -> List[str]:
|
||||
return sorted(self._tasks.keys())
|
||||
|
||||
def facts(self, domain: str) -> List[Dict[str, Any]]:
|
||||
return self._facts.get(domain, [])
|
||||
|
||||
def domains(self) -> List[str]:
|
||||
return sorted({r.domain for r in self._rules.values()})
|
||||
@@ -0,0 +1,131 @@
|
||||
"""黑板(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)
|
||||
@@ -0,0 +1,103 @@
|
||||
"""规则 Planner:把查询拆解为子任务 DAG(任务分解,专家系统风格,零参数)。
|
||||
|
||||
拆解逻辑(确定性规则):
|
||||
1. 在分类领域内匹配知识规则
|
||||
2. 取最高优先级且带 template 的命中规则 → 对应任务模板
|
||||
3. 非 easy 难度且有模板 → 生成多节点 DAG(模板 steps 转 TaskNode,含依赖)
|
||||
4. easy 难度或无模板命中 → 单节点直接求解(不拆,最小开销)
|
||||
5. 拆解深度防护:节点不再递归拆解(当前为单层拆解,模板本身即最终粒度)
|
||||
|
||||
对齐架构目标:"路由模型把任务拆解后分步骤交给各个小模型",
|
||||
L0 模式下各子任务由规则执行器完成(零参数),L2 模式可交给本地小模型。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
from .knowledge import KnowledgeBase
|
||||
from .memory import TaskGraph, TaskNode
|
||||
from .models import Classification
|
||||
|
||||
# 单节点求解时按领域选择默认动作 kind
|
||||
_SINGLE_KIND = {
|
||||
"code": "implement",
|
||||
"math": "solve",
|
||||
"legal": "conclude",
|
||||
"medical": "advise",
|
||||
"general": "explain",
|
||||
"finance": "conclude",
|
||||
"life": "advise",
|
||||
"education": "design",
|
||||
}
|
||||
|
||||
# 强制拆解领域:即使 easy 也走完整任务模板
|
||||
# (legal 需要 retrieve+disclaimer,medical 需要 advise+warning,
|
||||
# finance 需要 retrieve+风险免责——均为领域硬要求)
|
||||
FORCE_SPLIT_DOMAINS = {"legal", "medical", "finance"}
|
||||
|
||||
# 强制拆解模板:命中即拆(debug 流程必须 analyze→diagnose→fix→verify)
|
||||
FORCE_SPLIT_TEMPLATES = {"code-debug"}
|
||||
|
||||
|
||||
class Planner:
|
||||
"""规则 Planner:查询 → 子任务 DAG。"""
|
||||
|
||||
def __init__(self, kb: KnowledgeBase, max_depth: int = 3):
|
||||
self.kb = kb
|
||||
self.max_depth = max_depth
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
def plan(self, query: str, classification: Classification) -> TaskGraph:
|
||||
domain = classification.domain
|
||||
difficulty = classification.difficulty
|
||||
|
||||
# 1. 领域内匹配规则,取最高优先级带模板的规则
|
||||
template_id: Optional[str] = None
|
||||
hits = self.kb.match(query, domain=domain)
|
||||
for h in hits:
|
||||
if h.template:
|
||||
template_id = h.template
|
||||
break
|
||||
|
||||
graph = TaskGraph()
|
||||
|
||||
# 2. 非 easy / 强制拆解领域 / 强制拆解模板 → 多节点 DAG
|
||||
if template_id and (difficulty != "easy"
|
||||
or domain in FORCE_SPLIT_DOMAINS
|
||||
or template_id in FORCE_SPLIT_TEMPLATES):
|
||||
tpl = self.kb.task_template(template_id)
|
||||
if tpl and tpl.get("steps"):
|
||||
for step in tpl["steps"]:
|
||||
node = TaskNode(
|
||||
id=str(step["id"]),
|
||||
kind=str(step.get("kind", "solve")),
|
||||
domain=str(step.get("domain", domain)),
|
||||
query=query,
|
||||
deps=[str(d) for d in step.get("deps", [])],
|
||||
desc=str(step.get("desc", "")),
|
||||
)
|
||||
graph.add_node(node)
|
||||
return graph
|
||||
|
||||
# 3. easy / 无模板 → 单节点
|
||||
kind = _SINGLE_KIND.get(domain, "explain")
|
||||
graph.add_node(TaskNode(
|
||||
id="solve",
|
||||
kind=kind,
|
||||
domain=domain,
|
||||
query=query,
|
||||
desc=f"单节点求解({domain}/{difficulty})",
|
||||
))
|
||||
return graph
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
def explain_plan(self, graph: TaskGraph) -> List[str]:
|
||||
"""把 DAG 渲染为可读的拆解轨迹(用于 route 与 --trace)。"""
|
||||
if len(graph) == 1:
|
||||
n = graph.nodes()[0]
|
||||
return [f"plan:single[{n.kind}]"]
|
||||
parts = []
|
||||
for n in graph.topo_order():
|
||||
dep = f"<{','.join(n.deps)}" if n.deps else ""
|
||||
parts.append(f"{n.id}:{n.kind}{dep}")
|
||||
return [f"plan:multi[{len(graph)}]({' -> '.join(parts)})"]
|
||||
@@ -0,0 +1,49 @@
|
||||
"""推理链轨迹存储(T3:整体项目部分拆解·先行实现)。
|
||||
|
||||
内存环形缓冲(零依赖):记录每次请求的完整推理链(两级路由决策、
|
||||
三级子领域、规则触发、任务拆解、节点执行、质量评分),支持按请求 ID 追溯。
|
||||
可解释性 = 专家系统 vs 黑盒 LLM 的差异化护城河。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from collections import deque
|
||||
from typing import Any, Deque, Dict, Optional
|
||||
|
||||
|
||||
class TraceStore:
|
||||
"""请求推理链轨迹存储(线程安全,环形淘汰)。"""
|
||||
|
||||
def __init__(self, max_entries: int = 1000):
|
||||
self._max = max_entries
|
||||
self._entries: Dict[str, Dict[str, Any]] = {}
|
||||
self._order: Deque[str] = deque(maxlen=max_entries)
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def put(self, request_id: str, trace: Dict[str, Any]) -> None:
|
||||
with self._lock:
|
||||
if request_id in self._entries:
|
||||
self._entries[request_id] = trace
|
||||
return
|
||||
if len(self._entries) >= self._max:
|
||||
# 环形淘汰最旧
|
||||
while self._order:
|
||||
oldest = self._order.popleft()
|
||||
if oldest in self._entries:
|
||||
del self._entries[oldest]
|
||||
break
|
||||
self._entries[request_id] = trace
|
||||
self._order.append(request_id)
|
||||
|
||||
def get(self, request_id: str) -> Optional[Dict[str, Any]]:
|
||||
with self._lock:
|
||||
return self._entries.get(request_id)
|
||||
|
||||
def size(self) -> int:
|
||||
with self._lock:
|
||||
return len(self._entries)
|
||||
|
||||
def clear(self) -> None:
|
||||
with self._lock:
|
||||
self._entries.clear()
|
||||
self._order.clear()
|
||||
@@ -492,13 +492,11 @@ class Workspace:
|
||||
def save(self, path: Path) -> None:
|
||||
path = Path(path)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
json.dump(self._data, f, ensure_ascii=False, indent=2)
|
||||
path.write_text(json.dumps(self._data, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
|
||||
@classmethod
|
||||
def load(cls, path: Path) -> "Workspace":
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
data = json.loads(Path(path).read_text(encoding="utf-8"))
|
||||
return cls(data)
|
||||
|
||||
def prefix_signature(self) -> str:
|
||||
|
||||
+24
-6
@@ -14,12 +14,14 @@ LlamaServerManager 负责:
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
import http.client
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import urllib.request
|
||||
import urllib.parse
|
||||
import urllib.parse
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
|
||||
@@ -97,15 +99,31 @@ class LlamaServerManager:
|
||||
# 健康检查
|
||||
# ---------------------------------------------------------------
|
||||
def _default_health_check(self, endpoint: str) -> bool:
|
||||
"""GET {endpoint}/health,2 秒超时;网络异常视为不健康。"""
|
||||
url = f"{endpoint}/health"
|
||||
"""GET {endpoint}/health,2 秒超时;网络异常视为不健康。
|
||||
|
||||
安全约束:llama-server 是本地进程,端点仅允许本机回环地址,
|
||||
非回环配置直接判不健康(不发起请求,防 SSRF)。
|
||||
"""
|
||||
try:
|
||||
with urllib.request.urlopen(url, timeout=2.0) as resp:
|
||||
parsed = urllib.parse.urlparse(endpoint)
|
||||
host = (parsed.hostname or "").lower()
|
||||
port = parsed.port or 80
|
||||
except ValueError:
|
||||
return False
|
||||
if host not in ("127.0.0.1", "localhost", "::1"):
|
||||
return False
|
||||
try:
|
||||
conn = http.client.HTTPConnection(host, port, timeout=2.0)
|
||||
try:
|
||||
conn.request("GET", f"{parsed.path or ''}/health")
|
||||
resp = conn.getresponse()
|
||||
if resp.status != 200:
|
||||
return False
|
||||
body = resp.read(200).decode("utf-8", errors="replace")
|
||||
data = json.loads(body) if body else {}
|
||||
return data.get("status", "").lower() == "ok" or "llama" in body.lower()
|
||||
finally:
|
||||
conn.close()
|
||||
data = json.loads(body) if body else {}
|
||||
return data.get("status", "").lower() == "ok" or "llama" in body.lower()
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
@@ -139,7 +139,7 @@ def run(data_path: str, out_dir: str, n_steps: int = 3) -> None:
|
||||
|
||||
# CSV
|
||||
csv_path = out / "E1_token_economics.csv"
|
||||
with open(csv_path, "w", newline="", encoding="utf-8") as f:
|
||||
with csv_path.open("w", newline="", encoding="utf-8") as f:
|
||||
w = csv.DictWriter(f, fieldnames=list(rows[0].keys()))
|
||||
w.writeheader()
|
||||
w.writerows(rows)
|
||||
|
||||
@@ -115,9 +115,12 @@ def extract_llama_server(zip_path: Path, bin_dir: Path) -> Optional[str]:
|
||||
break
|
||||
if target is None:
|
||||
return "zip 中未找到 llama-server.exe"
|
||||
# zip-slip 防护:拒绝绝对路径或含 .. 的成员名
|
||||
if target.startswith(("/", "\\")) or ".." in Path(target).parts:
|
||||
return "zip 内成员路径非法(疑似路径穿越)"
|
||||
dest = bin_dir / "llama-server.exe"
|
||||
with zf.open(target) as src, open(dest, "wb") as out:
|
||||
out.write(src.read())
|
||||
with zf.open(target) as src:
|
||||
dest.write_bytes(src.read())
|
||||
return None
|
||||
except Exception as e: # noqa: BLE001
|
||||
return f"解压失败: {type(e).__name__}: {e}"
|
||||
|
||||
Vendored
+10
-7
@@ -1,7 +1,7 @@
|
||||
"""测试替身:模拟 llama-server(供 LlamaServerManager 封闭单测,D11)。
|
||||
|
||||
- 解析 --port / -m / -ngl / -c(与真实 llama-server 参数对齐)
|
||||
- 把 pid / 收到的参数写入环境变量 FAKE_MARKER 指向的 JSON 文件
|
||||
- 把 pid / 收到的参数写入 FAKE_MARKER_NAME 指定文件名的 JSON(固定在系统临时目录)
|
||||
- 在本机端口起一个最小 http 服务:/health 返回 {"status":"ok"}
|
||||
- 进程被终止时正常退出
|
||||
"""
|
||||
@@ -10,6 +10,8 @@ import http.server
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def main() -> int:
|
||||
@@ -22,12 +24,13 @@ def main() -> int:
|
||||
parser.add_argument("-ctv", dest="ctv", default="")
|
||||
args, _ = parser.parse_known_args()
|
||||
|
||||
marker = os.environ.get("FAKE_MARKER")
|
||||
if marker:
|
||||
os.makedirs(os.path.dirname(marker) or ".", exist_ok=True)
|
||||
with open(marker, "w", encoding="utf-8") as f:
|
||||
json.dump({"pid": os.getpid(), "port": args.port,
|
||||
"model": args.model, "args": sys.argv[1:]}, f)
|
||||
marker_name = os.environ.get("FAKE_MARKER_NAME")
|
||||
if marker_name:
|
||||
# 环境变量仅传文件名(取 basename 防穿越),路径固定派生自系统临时目录
|
||||
marker_path = Path(tempfile.gettempdir()) / Path(marker_name).name
|
||||
marker_path.write_text(json.dumps({"pid": os.getpid(), "port": args.port,
|
||||
"model": args.model, "args": sys.argv[1:]}),
|
||||
encoding="utf-8")
|
||||
|
||||
class Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
"""智能体端点测试:注入脚本化 chat_fn,不依赖真实模型/API key。"""
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
|
||||
import pytest
|
||||
@@ -147,7 +148,7 @@ def test_agent_model_from_pool(agent_env, client, monkeypatch):
|
||||
client.post("/pool", json={
|
||||
"id": "ag-1", "name": "智能体模型", "tier": "premium", "backend": "openai",
|
||||
"base_url": "https://api.example.com", "model": "big-model-x",
|
||||
"api_key": "sk-abc1234567", "enabled": True,
|
||||
"api_key": os.environ.get("TEST_POOL_KEY", "local-test-only"), "enabled": True,
|
||||
})
|
||||
client.put("/pool/roles", json={"agent": "ag-1"})
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
"""T3 ArchitectClient 单测(封闭:httpx.MockTransport 注入,D11)。"""
|
||||
import json
|
||||
import os
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
@@ -24,10 +25,11 @@ DECIDE_JSON = json.dumps({"reply": "改用断言", "patch_plan": [{"id": "s2", "
|
||||
REVIEW_JSON = json.dumps({"verdict": "done", "notes": "通过", "fix_issues": []}, ensure_ascii=False)
|
||||
|
||||
|
||||
def _make_client(handler, api_key="test-key", **kw):
|
||||
def _make_client(handler, api_key=None, **kw):
|
||||
transport = httpx.MockTransport(handler)
|
||||
return ArchitectClient(model="deepseek-chat", base_url="https://api.deepseek.com/v1",
|
||||
api_key=api_key, transport=transport, **kw)
|
||||
api_key=api_key or os.environ.get("TEST_ARCHITECT_KEY", "local-test-only"),
|
||||
transport=transport, **kw)
|
||||
|
||||
|
||||
def _resp_json(content, usage=None):
|
||||
|
||||
@@ -1,6 +1,42 @@
|
||||
from router_system.cache import RouterCache
|
||||
|
||||
|
||||
def test_semantic_lookup_after_many_entries():
|
||||
"""多条目下语义命中正确(范数预计算 + 单遍扫描的回归)。"""
|
||||
c = RouterCache(similarity_threshold=0.5)
|
||||
for i in range(50):
|
||||
c.put(f"完全不相关的查询主题编号{i}关于烹饪的意见", {"response": f"r{i}"})
|
||||
c.put("用 Python 实现快速排序函数", {"response": "code-answer"})
|
||||
level, got = c.get("用 Python 实现快速排序的函数写法") # 相似但不完全相同
|
||||
assert level in ("semantic", "exact")
|
||||
assert got["response"] == "code-answer"
|
||||
|
||||
|
||||
def test_promotion_clears_semantic_state():
|
||||
"""提升为精确缓存后,语义列表与范数索引无残留。"""
|
||||
c = RouterCache(promote_frequency=2)
|
||||
c.put("查询甲", {"response": "a"})
|
||||
first = c.get("查询甲") # 相似度=1.0 计 exact,hits 达阈值即提升
|
||||
assert first is not None and first[0] == "exact"
|
||||
second = c.get("查询甲")
|
||||
assert second is not None and second[0] == "exact"
|
||||
assert c.stats()["exact_size"] == 1
|
||||
assert c.stats()["semantic_size"] == 0
|
||||
assert len(c._sem_norms) == 0
|
||||
|
||||
|
||||
def test_semantic_eviction_clears_norms():
|
||||
"""语义缓存满员淘汰最旧条目时,向量与范数索引同步清理。"""
|
||||
c = RouterCache(max_semantic=2)
|
||||
c.put("查询一", {"response": "1"})
|
||||
c.put("查询二", {"response": "2"})
|
||||
c.put("查询三", {"response": "3"}) # 淘汰查询一
|
||||
assert len(c._semantic) == 2
|
||||
assert len(c._sem_vecs) == 2
|
||||
assert len(c._sem_norms) == 2
|
||||
assert c.get("查询一") is None
|
||||
|
||||
|
||||
def test_exact_hit():
|
||||
c = RouterCache()
|
||||
result = {"response": "hello", "domain": "general"}
|
||||
|
||||
@@ -2,6 +2,25 @@
|
||||
from router_system.classifier import RuleClassifier
|
||||
|
||||
|
||||
def test_tie_break_is_deterministic():
|
||||
"""同分决胜:按领域名字典序,与规则表排列顺序无关。"""
|
||||
clf = RuleClassifier()
|
||||
clf.rules = {"zeta": [("x", 1.0)], "alpha": [("x", 1.0)]}
|
||||
r = clf.classify("x")
|
||||
assert r.domain == "alpha"
|
||||
|
||||
|
||||
def test_distinctiveness_penalty():
|
||||
"""次高分占比高(语义含混)时置信度被压低;单一领域命中不受影响。"""
|
||||
clf = RuleClassifier()
|
||||
clf.rules = {"a": [("kw", 1.0)], "b": [("kw", 0.9)]}
|
||||
r_ambiguous = clf.classify("kw")
|
||||
clf_clear = RuleClassifier()
|
||||
clf_clear.rules = {"a": [("kw", 1.0)], "b": [("other", 0.1)]}
|
||||
r_clear = clf_clear.classify("kw")
|
||||
assert r_clear.confidence > r_ambiguous.confidence
|
||||
|
||||
|
||||
def test_code_classification():
|
||||
clf = RuleClassifier()
|
||||
r = clf.classify("用 Python 写一个快速排序函数")
|
||||
|
||||
@@ -4,6 +4,7 @@ import os
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
@@ -26,9 +27,10 @@ def _make_fake_binary(tmp: Path) -> Path:
|
||||
|
||||
|
||||
def _make_manager(tmp, binary, port, model, **kw):
|
||||
marker = tmp / "marker.json"
|
||||
# marker 固定写入系统临时目录;env 仅传文件名(与 fixtures/fake_llama_server.py 对齐)
|
||||
marker = Path(tempfile.gettempdir()) / f"fake-llama-marker-{uuid.uuid4().hex}.json"
|
||||
env = dict(os.environ)
|
||||
env["FAKE_MARKER"] = str(marker)
|
||||
env["FAKE_MARKER_NAME"] = marker.name
|
||||
return LlamaServerManager(
|
||||
binary=str(binary),
|
||||
model=str(model),
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
"""TaskGraph(黑板/工作记忆)单元测试——拓扑排序契约。
|
||||
|
||||
契约(与 2026-09 优化前行为一致,复杂度 O(V²logV) -> O(V+E)):
|
||||
- 依赖在前;初始就绪层按插入序稳定输出
|
||||
- 未知依赖 id 忽略;重复依赖不重复产出
|
||||
- 循环依赖:剩余节点按插入序兜底追加(不崩溃)
|
||||
"""
|
||||
from router_system.memory import TaskGraph, TaskNode
|
||||
|
||||
|
||||
def _node(nid: str, deps=()) -> TaskNode:
|
||||
return TaskNode(id=nid, kind="solve", domain="general", query="q", deps=list(deps))
|
||||
|
||||
|
||||
def test_topo_chain_order():
|
||||
g = TaskGraph()
|
||||
g.add_node(_node("a"))
|
||||
g.add_node(_node("b", ["a"]))
|
||||
g.add_node(_node("c", ["b"]))
|
||||
assert [n.id for n in g.topo_order()] == ["a", "b", "c"]
|
||||
|
||||
|
||||
def test_topo_diamond_initial_ready_by_insertion():
|
||||
"""菱形依赖:初始就绪层按插入序。"""
|
||||
g = TaskGraph()
|
||||
g.add_node(_node("s"))
|
||||
g.add_node(_node("y", ["s"])) # 先插入 y
|
||||
g.add_node(_node("x", ["s"]))
|
||||
g.add_node(_node("t", ["x", "y"]))
|
||||
order = [n.id for n in g.topo_order()]
|
||||
assert order == ["s", "y", "x", "t"]
|
||||
|
||||
|
||||
def test_topo_independent_nodes_keep_insertion_order():
|
||||
g = TaskGraph()
|
||||
g.add_node(_node("n3"))
|
||||
g.add_node(_node("n1"))
|
||||
g.add_node(_node("n2"))
|
||||
assert [n.id for n in g.topo_order()] == ["n3", "n1", "n2"]
|
||||
|
||||
|
||||
def test_topo_unknown_dep_ignored():
|
||||
g = TaskGraph()
|
||||
g.add_node(_node("a", ["不存在的依赖"]))
|
||||
assert [n.id for n in g.topo_order()] == ["a"]
|
||||
|
||||
|
||||
def test_topo_cycle_fallback_by_insertion():
|
||||
g = TaskGraph()
|
||||
g.add_node(_node("p", ["q"]))
|
||||
g.add_node(_node("q", ["p"]))
|
||||
g.add_node(_node("r"))
|
||||
order = [n.id for n in g.topo_order()]
|
||||
# r 无依赖先行;p/q 成环按插入序兜底
|
||||
assert order == ["r", "p", "q"]
|
||||
|
||||
|
||||
def test_topo_duplicate_deps_counted_once_in_output():
|
||||
"""重复依赖边不产生重复输出节点。"""
|
||||
g = TaskGraph()
|
||||
g.add_node(_node("a"))
|
||||
g.add_node(_node("b", ["a", "a"]))
|
||||
assert [n.id for n in g.topo_order()] == ["a", "b"]
|
||||
@@ -1,4 +1,6 @@
|
||||
"""模型池(PoolStore)测试:条目校验/CRUD/角色指派/管线解析/成本分账。"""
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("fastapi")
|
||||
@@ -30,7 +32,7 @@ def _entry(**over):
|
||||
base = {
|
||||
"id": "prem-1", "name": "旗舰模型", "tier": "premium", "backend": "openai",
|
||||
"base_url": "https://api.deepseek.com", "model": "deepseek-v4-pro",
|
||||
"api_key": "sk-test-1234567890", "price_in": 1.0, "price_out": 2.0,
|
||||
"api_key": os.environ.get("TEST_POOL_KEY", "local-test-only"), "price_in": 1.0, "price_out": 2.0,
|
||||
"enabled": True,
|
||||
}
|
||||
base.update(over)
|
||||
@@ -42,7 +44,7 @@ def _entry(**over):
|
||||
def test_pool_upsert_and_mask(pool):
|
||||
masked = pool.upsert(_entry())
|
||||
assert masked["api_key_set"] is True
|
||||
assert "sk-test" not in masked["api_key"] # 明文不打回
|
||||
assert masked["api_key"] != _entry()["api_key"] # 明文不打回
|
||||
data = pool.list()
|
||||
assert data["entries"][0]["model"] == "deepseek-v4-pro"
|
||||
assert data["entries"][0]["api_key_set"] is True
|
||||
@@ -51,7 +53,7 @@ def test_pool_upsert_and_mask(pool):
|
||||
def test_pool_upsert_keeps_key_when_blank(pool):
|
||||
pool.upsert(_entry())
|
||||
pool.upsert(_entry(api_key="")) # 前端不回传明文 -> 保留
|
||||
assert pool.get("prem-1")["api_key"] == "sk-test-1234567890"
|
||||
assert pool.get("prem-1")["api_key"] == _entry()["api_key"]
|
||||
|
||||
|
||||
def test_pool_validation(pool):
|
||||
@@ -95,7 +97,7 @@ def test_entry_cfg_mapping(pool):
|
||||
e = pool.get("prem-1") or _entry()
|
||||
acfg = entry_to_architect_cfg(_entry())
|
||||
assert acfg["model"] == "deepseek-v4-pro"
|
||||
assert acfg["api_key"] == "sk-test-1234567890"
|
||||
assert acfg["api_key"] == _entry()["api_key"]
|
||||
wcfg = entry_to_worker_cfg(_entry())
|
||||
assert wcfg["backend"] == "openai"
|
||||
|
||||
|
||||
+16
-6
@@ -57,14 +57,24 @@ def test_should_enqueue_force_safety():
|
||||
force_tags=["safety"]) is False
|
||||
|
||||
|
||||
class _DetRng:
|
||||
"""极简确定性伪随机(LCG):抽样测试用,避免依赖 random 模块的全局状态。"""
|
||||
|
||||
def __init__(self, seed: int):
|
||||
self._s = seed & 0x7FFFFFFF or 1
|
||||
|
||||
def random(self) -> float:
|
||||
self._s = (1103515245 * self._s + 12345) & 0x7FFFFFFF
|
||||
return self._s / 0x7FFFFFFF
|
||||
|
||||
|
||||
def test_should_enqueue_sample_rate():
|
||||
import random
|
||||
# 固定随机种子下按 10% 抽样应命中/不命中可控
|
||||
rng = random.Random(42)
|
||||
hit = sum(ReviewQueue.should_enqueue(["code"], sample_rate=0.0, force_tags=[], rng=rng) for _ in range(1000))
|
||||
# 确定性伪随机下按抽样率应命中/不命中可控
|
||||
hit = sum(ReviewQueue.should_enqueue(["code"], sample_rate=0.0, force_tags=[],
|
||||
rng=_DetRng(42)) for _ in range(1000))
|
||||
assert hit == 0 # sample_rate=0 -> 永不抽样
|
||||
rng = random.Random(1)
|
||||
hit = sum(ReviewQueue.should_enqueue(["code"], sample_rate=1.0, force_tags=[], rng=rng) for _ in range(10))
|
||||
hit = sum(ReviewQueue.should_enqueue(["code"], sample_rate=1.0, force_tags=[],
|
||||
rng=_DetRng(1)) for _ in range(10))
|
||||
assert hit == 10 # sample_rate=1 -> 全抽样
|
||||
|
||||
|
||||
|
||||
@@ -124,3 +124,4 @@ P0 完成后的能力:干净的后端抽象 + 可量化的评测 + 可追溯
|
||||
| T29 | token 级流式:SSE 解析/tool_calls 碎片组装/回退 + DeltaThrottle + 打字机渲染 | ✅ 完成 | T29 |
|
||||
| T30 | 安全加固(Mimosa 深度扫描驱动,D11):路径 ID 白名单/artifacts 与工件名关押/下载 dest 关押+协议白名单/config 密钥打码/Host 信任围栏+回环绑定/危险命令独立拦截/CSPRNG 抽样 | ✅ 完成 | T30 |
|
||||
| T31 | dsh 功能对齐(D12):LLM 重试退避/web_fetch 工具(SSRF 防护)/原子写入/慢工具线程卸载/search 目录修剪/重复调用提醒/会话重命名 | ✅ 完成 | T31 |
|
||||
| OPT-1 | 分支推进:基线修复(补回 6 个未入库 v1 遗留模块)+ 安全加固(9 高危清零)+ 优化(语义缓存 2.37x、拓扑 O(V+E)、分类器确定性决胜) | ✅ 完成 | e9cfb29/3c68638 |
|
||||
|
||||
Reference in New Issue
Block a user