Skip to main content

p_memory/
graph_search.rs

1//! 图谱搜索:基于 petgraph 的**内存态图**。
2//!
3//! 存储仍是 SQLite(权威);本模块按需把 `relations` 投影表读进内存建图,
4//! 调 petgraph 现成算法,用完即弃。算法一律用 petgraph 实现,不自己造。
5
6use crate::graph::{Entity, GraphStore};
7use crate::storage::{self, filter_sql};
8use crate::types::{ReadFilter, RecordKind, WriteReceipt};
9use crate::{Error, Result};
10use petgraph::algo::{astar, connected_components, tarjan_scc};
11use petgraph::graph::{DiGraph, Graph, NodeIndex};
12use rusqlite::{params, params_from_iter, Connection};
13use std::collections::{BTreeMap, HashMap, HashSet};
14
15/// 别名等价关系的内置谓词:建图时并查集缩点,把互为别名的实体折成同一超节点。
16const SAME_AS: &str = "sys:same_as";
17
18/// 查询期并查集:只服务 `sys:same_as` 的缩点,用完即弃,不落盘。
19#[derive(Default)]
20struct UnionFind { parent: HashMap<i64, i64> }
21
22impl UnionFind {
23    fn find(&mut self, x: i64) -> i64 {
24        if !self.parent.contains_key(&x) { self.parent.insert(x, x); return x; }
25        let mut root = x;
26        while self.parent[&root] != root { root = self.parent[&root]; }
27        let mut cursor = x;
28        while self.parent[&cursor] != root { let next = self.parent[&cursor]; self.parent.insert(cursor, root); cursor = next; }
29        root
30    }
31    fn union(&mut self, a: i64, b: i64) {
32        let (ra, rb) = (self.find(a), self.find(b));
33        if ra != rb { self.parent.insert(ra, rb); }
34    }
35}
36
37/// 从 `relations` 构建的内存态图。节点权重 = 实体 record_id,边权重 = 谓词文本。
38///
39/// 图是**有向**的(petgraph `Graph` 默认 `Directed`,`neighbors` 只返回出边):
40/// 反向查询(如由「父亲」反查「子女」)不会自动连通,靠 `predicate_rules` 在 `snapshot`
41/// 里补出对称/逆关系的虚拟边来打通,物理表不落双向边。
42///
43/// 节点 id 是**缩点后的代表 id**:互为 `sys:same_as` 的实体被并查集折叠成同一节点,
44/// `canonical` 保存「原始 id -> 代表 id」的映射,查询入口先做一次换算。
45pub struct GraphView {
46    graph: Graph<i64, String>,
47    index: HashMap<i64, NodeIndex>,
48    canonical: HashMap<i64, i64>,
49}
50
51/// 读一次快照:`filter` 范围内的实体 id,以及有向边 (subject, object, predicate)。
52///
53/// 边里已按 `predicate_rules` 补入内存虚拟边——对称谓词补反向同谓词边、有逆谓词的补对偶边。
54fn snapshot(conn: &Connection, filter: &ReadFilter) -> Result<(Vec<i64>, Vec<(i64, i64, String)>)> {
55    let (cond, values) = filter_sql(filter, &[RecordKind::Entity], false)?;
56    let mut stmt = conn.prepare(&format!(
57        "SELECT e.record_id FROM entities e JOIN records r ON r.id=e.record_id WHERE {cond}"
58    ))?;
59    let nodes = stmt
60        .query_map(params_from_iter(values), |row| row.get::<_, i64>(0))?
61        .collect::<rusqlite::Result<Vec<_>>>()?;
62
63    let (cond, values) = filter_sql(filter, &[RecordKind::Relation], false)?;
64    let mut stmt = conn.prepare(&format!(
65        "SELECT rel.subject_id, rel.object_id, s.text, pr.is_symmetric, inv.text \
66         FROM relations rel JOIN records r ON r.id=rel.record_id \
67         JOIN strings s ON s.id=rel.predicate_id \
68         LEFT JOIN predicate_rules pr ON pr.predicate_id=rel.predicate_id \
69         LEFT JOIN strings inv ON inv.id=pr.inverse_predicate_id WHERE {cond}"
70    ))?;
71    let mut edges: Vec<(i64, i64, String)> = Vec::new();
72    let mut rows = stmt.query(params_from_iter(values))?;
73    while let Some(row) = rows.next()? {
74        let subject: i64 = row.get(0)?;
75        let object: i64 = row.get(1)?;
76        let predicate: String = row.get(2)?;
77        let symmetric: Option<i64> = row.get(3)?;
78        let inverse: Option<String> = row.get(4)?;
79        edges.push((subject, object, predicate.clone()));
80        if symmetric == Some(1) {
81            edges.push((object, subject, predicate));
82        } else if let Some(inverse) = inverse {
83            edges.push((object, subject, inverse));
84        }
85    }
86    Ok((nodes, edges))
87}
88
89/// 按 `sys:same_as` 求每个 id 的代表 id(含 nodes 与所有边端点)。
90fn canonical_ids(nodes: &[i64], edges: &[(i64, i64, String)]) -> HashMap<i64, i64> {
91    let mut uf = UnionFind::default();
92    for (s, o, pred) in edges {
93        if pred == SAME_AS { uf.union(*s, *o); }
94    }
95    let mut canonical: HashMap<i64, i64> = HashMap::new();
96    for id in nodes.iter().copied().chain(edges.iter().flat_map(|(s, o, _)| [*s, *o])) {
97        let representative = uf.find(id);
98        canonical.insert(id, representative);
99    }
100    canonical
101}
102
103impl GraphStore {
104    /// 把 `filter` 范围(namespace/scope/tag)内的实体铺成节点、关系连成边,建一张有向图。
105    ///
106    /// 节点来自 entity 记录(含没有关系的孤立实体,「谁没连进主图」才看得出来),
107    /// 边来自 relation 记录。范围外的记录不进图——多命名空间不会串。
108    pub fn build_graph(&self, filter: &ReadFilter) -> Result<GraphView> {
109        let state = self.0.read()?;
110        let (nodes, edges) = snapshot(state.conn(), filter)?;
111        let canonical = canonical_ids(&nodes, &edges);
112        let mut graph = Graph::<i64, String>::new();
113        let mut index: HashMap<i64, NodeIndex> = HashMap::new();
114        for raw in nodes.iter().copied().chain(edges.iter().flat_map(|(s, o, _)| [*s, *o])) {
115            let representative = canonical[&raw];
116            if !index.contains_key(&representative) {
117                let i = graph.add_node(representative);
118                index.insert(representative, i);
119            }
120        }
121        for (s, o, predicate) in edges {
122            if predicate == SAME_AS { continue; }
123            graph.add_edge(index[&canonical[&s]], index[&canonical[&o]], predicate);
124        }
125        Ok(GraphView { graph, index, canonical })
126    }
127
128    /// 强连通分量(有向,`tarjan_scc`):互相可达的实体环,如「互为对手/同伙」的闭环。
129    /// 只返回大小 > 1 的分量,孤点自环被滤掉。
130    pub fn strongly_connected(&self, filter: &ReadFilter) -> Result<Vec<Vec<i64>>> {
131        let (nodes, edges) = {
132            let state = self.0.read()?;
133            snapshot(state.conn(), filter)?
134        };
135        let canonical = canonical_ids(&nodes, &edges);
136        let mut graph = DiGraph::<i64, ()>::new();
137        let mut index: HashMap<i64, NodeIndex> = HashMap::new();
138        for raw in nodes.iter().copied().chain(edges.iter().flat_map(|(s, o, _)| [*s, *o])) {
139            let representative = canonical[&raw];
140            if !index.contains_key(&representative) {
141                let i = graph.add_node(representative);
142                index.insert(representative, i);
143            }
144        }
145        for (s, o, _pred) in edges {
146            let (si, oi) = (index[&canonical[&s]], index[&canonical[&o]]);
147            if si != oi { graph.add_edge(si, oi, ()); }
148        }
149        Ok(tarjan_scc(&graph)
150            .into_iter()
151            .filter(|c| c.len() > 1)
152            .map(|c| c.into_iter().map(|n| graph[n]).collect())
153            .collect())
154    }
155
156    /// 从 `root` 出发 `depth` 跳内的实体(带 name / aliases / attributes)。
157    pub fn ego(&self, root: i64, depth: usize, filter: &ReadFilter, limit: usize) -> Result<Vec<Entity>> {
158        let ids = self.build_graph(filter)?.ego_ids(root, depth, limit);
159        let state = self.0.read()?;
160        // 一次批量取回:逐条 `get` 会为每个实体各跑一遍过滤与装配。
161        let mut loaded: BTreeMap<i64, Entity> = storage::load_many(state.conn(), &ids, filter)?;
162        Ok(ids.iter().filter_map(|id| loaded.remove(id)).collect())
163    }
164
165    /// `from` → `to` 桥接路径上的实体(含两端)。不连通返回 None。
166    pub fn path(&self, from: i64, to: i64, filter: &ReadFilter) -> Result<Option<Vec<Entity>>> {
167        let Some(ids) = self.build_graph(filter)?.path_ids(from, to) else {
168            return Ok(None);
169        };
170        let state = self.0.read()?;
171        let mut loaded: BTreeMap<i64, Entity> = storage::load_many(state.conn(), &ids, filter)?;
172        Ok(Some(ids.iter().filter_map(|id| loaded.remove(id)).collect()))
173    }
174
175    /// 登记谓词元规则:`symmetric` 声明对称谓词(反向即自身),`inverse` 声明逆谓词
176    /// (反向补一条对偶边,如 `父亲` 的逆是 `子女`)。二者互斥,规则只影响建图时的内存补边。
177    pub fn set_predicate_rule(&self, predicate: &str, inverse: Option<&str>, symmetric: bool) -> Result<WriteReceipt<()>> {
178        storage::validate_identity("predicate", predicate)?;
179        if symmetric && inverse.is_some() {
180            return Err(Error::Validation("a symmetric predicate must not declare an inverse".into()));
181        }
182        if let Some(inverse) = inverse { storage::validate_identity("inverse predicate", inverse)?; }
183        self.0.mutate_meta(|tx| {
184            let predicate_id = storage::term_id(tx, predicate)?;
185            let inverse_id = match inverse { Some(text) => Some(storage::term_id(tx, text)?), None => None };
186            tx.execute("INSERT INTO predicate_rules(predicate_id,inverse_predicate_id,is_symmetric) VALUES (?1,?2,?3)
187                ON CONFLICT(predicate_id) DO UPDATE SET inverse_predicate_id=excluded.inverse_predicate_id,is_symmetric=excluded.is_symmetric",
188                params![predicate_id, inverse_id, i64::from(symmetric)])?;
189            Ok(())
190        })
191    }
192
193    /// 撤销一条谓词元规则(对称或逆谓词)。内置的 `sys:same_as` 同样可以撤。
194    /// 返回是否命中;没登记过的谓词不报错。
195    pub fn delete_predicate_rule(&self, predicate: &str) -> Result<WriteReceipt<bool>> {
196        storage::validate_identity("predicate", predicate)?;
197        self.0.mutate_meta(|tx| {
198            let removed = tx.execute("DELETE FROM predicate_rules WHERE predicate_id=(SELECT id FROM strings WHERE text=?1)",
199                [crate::text::normalized_tag(predicate)])?;
200            Ok(removed > 0)
201        })
202    }
203}
204
205impl GraphView {
206    pub fn node_count(&self) -> usize {
207        self.graph.node_count()
208    }
209
210    pub fn edge_count(&self) -> usize {
211        self.graph.edge_count()
212    }
213
214    /// 从 `root` 出发 `depth` 跳内的实体 id(不含 root)。`root` 会先折算成缩点后的代表 id。
215    pub fn ego_ids(&self, root: i64, depth: usize, limit: usize) -> Vec<i64> {
216        let root = self.canonical.get(&root).copied().unwrap_or(root);
217        let start = match self.index.get(&root) {
218            Some(&i) => i,
219            None => return Vec::new(),
220        };
221        let mut seen: HashSet<NodeIndex> = HashSet::new();
222        seen.insert(start);
223        let mut frontier = vec![start];
224        let mut out: Vec<i64> = Vec::new();
225        for _ in 0..depth {
226            let mut next = Vec::new();
227            for n in &frontier {
228                for nb in self.graph.neighbors(*n) {
229                    if seen.insert(nb) {
230                        out.push(self.graph[nb]);
231                        next.push(nb);
232                        if out.len() >= limit {
233                            return out;
234                        }
235                    }
236                }
237            }
238            if next.is_empty() {
239                break;
240            }
241            frontier = next;
242        }
243        out
244    }
245
246    /// `from` 到 `to` 的最短路径(节点 id 列表,含两端)。用 petgraph 的 `astar`。
247    /// 两端会先折算成缩点后的代表 id。
248    pub fn path_ids(&self, from: i64, to: i64) -> Option<Vec<i64>> {
249        let from = self.canonical.get(&from).copied().unwrap_or(from);
250        let to = self.canonical.get(&to).copied().unwrap_or(to);
251        let start = self.index.get(&from)?;
252        let goal = self.index.get(&to)?;
253        let (_cost, path) = astar(&self.graph, *start, |n| n == *goal, |_e| 1i32, |_n| 0i32)?;
254        Some(path.into_iter().map(|n| self.graph[n]).collect())
255    }
256
257    /// 连通分量数量(无向)。用 petgraph 的 `connected_components`。
258    pub fn component_count(&self) -> usize {
259        connected_components(&self.graph)
260    }
261}
262
263#[cfg(test)]
264mod tests {
265    use crate::graph::{EntityInput, GraphBatch, RelationInput};
266    use crate::types::{ReadFilter, RecordInput};
267    use crate::KnowledgeBase;
268    use std::collections::BTreeMap;
269
270    fn ent(name: &str) -> EntityInput {
271        EntityInput {
272            record: RecordInput::default(),
273            name: name.into(),
274            entity_type: "person".into(),
275            aliases: Vec::new(),
276            attributes: BTreeMap::new(),
277            summary: String::new(),
278        }
279    }
280
281    fn rel(s: i64, p: &str, o: i64) -> RelationInput {
282        RelationInput {
283            record: RecordInput::default(),
284            subject_id: s,
285            predicate: p.into(),
286            object_id: o,
287            confidence: 1.0,
288            reason: String::new(),
289        }
290    }
291
292    #[test]
293    fn ego_path_components() {
294        let dir = tempfile::tempdir().unwrap();
295        let kb = KnowledgeBase::open(dir.path()).unwrap();
296        let ents = kb
297            .graph()
298            .apply_batch(&GraphBatch { entities: vec![ent("A"), ent("B"), ent("C")], ..Default::default() })
299            .unwrap()
300            .value
301            .entities;
302        let (a, b, c) = (ents[0].header.id, ents[1].header.id, ents[2].header.id);
303        kb.graph()
304            .apply_batch(&GraphBatch { relations: vec![rel(a, "knows", b), rel(b, "knows", c)], ..Default::default() })
305            .unwrap();
306
307        let view = kb.graph().build_graph(&ReadFilter::default()).unwrap();
308        assert_eq!(view.node_count(), 3);
309        assert_eq!(view.edge_count(), 2);
310        assert_eq!(view.component_count(), 1);
311
312        // 邻域:A 的一跳只有 B,两跳兜住 B、C
313        assert_eq!(view.ego_ids(a, 1, 100), vec![b]);
314        let mut ego = view.ego_ids(a, 2, 100);
315        ego.sort();
316        let mut expect = vec![b, c];
317        expect.sort();
318        assert_eq!(ego, expect);
319
320        // 桥接:A→C 的最短路径
321        assert_eq!(view.path_ids(a, c), Some(vec![a, b, c]));
322
323        // Entity 级查询:返回的实体对象自带 name 字段
324        let names: Vec<String> = kb
325            .graph()
326            .ego(a, 2, &ReadFilter::default(), 100)
327            .unwrap()
328            .into_iter()
329            .map(|e| e.name)
330            .collect();
331        assert!(names.contains(&"B".to_string()) && names.contains(&"C".to_string()));
332        let hop: Vec<String> = kb
333            .graph()
334            .path(a, c, &ReadFilter::default())
335            .unwrap()
336            .unwrap()
337            .into_iter()
338            .map(|e| e.name)
339            .collect();
340        assert_eq!(hop, vec!["A", "B", "C"]);
341
342        // 强连通:补 B→A 形成环,SCC 找到互相可达的 {A,B}
343        kb.graph()
344            .apply_batch(&GraphBatch { relations: vec![rel(b, "knows", a)], ..Default::default() })
345            .unwrap();
346        let scc = kb.graph().strongly_connected(&ReadFilter::default()).unwrap();
347        assert_eq!(scc.len(), 1);
348        let mut ring = scc[0].clone();
349        ring.sort();
350        let mut exp = vec![a, b];
351        exp.sort();
352        assert_eq!(ring, exp);
353
354        // 孤立实体独立成块
355        kb.graph()
356            .apply_batch(&GraphBatch { entities: vec![ent("D")], ..Default::default() })
357            .unwrap();
358        let view = kb.graph().build_graph(&ReadFilter::default()).unwrap();
359        assert_eq!(view.node_count(), 4);
360        assert_eq!(view.component_count(), 2);
361    }
362
363    #[test]
364    fn filter_isolates_namespaces() {
365        let dir = tempfile::tempdir().unwrap();
366        let kb = KnowledgeBase::open(dir.path()).unwrap();
367
368        let mut ea = ent("A");
369        ea.record.namespace = "a".into();
370        let mut eb = ent("B");
371        eb.record.namespace = "a".into();
372        let mut ec = ent("C");
373        ec.record.namespace = "b".into();
374        let got = kb
375            .graph()
376            .apply_batch(&GraphBatch { entities: vec![ea, eb, ec], ..Default::default() })
377            .unwrap()
378            .value
379            .entities;
380        let (a, b) = (got[0].header.id, got[1].header.id);
381        let mut r = rel(a, "knows", b);
382        r.record.namespace = "a".into();
383        kb.graph()
384            .apply_batch(&GraphBatch { relations: vec![r], ..Default::default() })
385            .unwrap();
386
387        let fa = ReadFilter { namespace: "a".into(), scopes: vec!["public".into()], tags: vec![], note_ids: vec![] };
388        let fb = ReadFilter { namespace: "b".into(), scopes: vec!["public".into()], tags: vec![], note_ids: vec![] };
389        let va = kb.graph().build_graph(&fa).unwrap();
390        assert_eq!((va.node_count(), va.edge_count()), (2, 1));
391        let vb = kb.graph().build_graph(&fb).unwrap();
392        assert_eq!((vb.node_count(), vb.edge_count()), (1, 0));
393    }
394
395    #[test]
396    fn virtual_inverse_and_symmetric_edges_enable_reverse_queries() {
397        let dir = tempfile::tempdir().unwrap();
398        let kb = KnowledgeBase::open(dir.path()).unwrap();
399        let filter = ReadFilter::default();
400        let ents = kb
401            .graph()
402            .apply_batch(&GraphBatch { entities: vec![ent("张伟"), ent("张父"), ent("甲"), ent("乙")], ..Default::default() })
403            .unwrap()
404            .value
405            .entities;
406        let (son, father, jia, yi) = (ents[0].header.id, ents[1].header.id, ents[2].header.id, ents[3].header.id);
407        // 物理只写单向边:张伟 --父亲--> 张父;甲 --同事--> 乙
408        kb.graph()
409            .apply_batch(&GraphBatch { relations: vec![rel(son, "父亲", father), rel(jia, "同事", yi)], ..Default::default() })
410            .unwrap();
411
412        // 未登记规则:反向查询天然不通(有向图)
413        let view = kb.graph().build_graph(&filter).unwrap();
414        assert!(view.ego_ids(father, 1, 100).is_empty());
415        assert_eq!(view.path_ids(father, son), None);
416
417        // 登记逆谓词:反向补出「子女」边,反向 ego/path 打通
418        kb.graph().set_predicate_rule("父亲", Some("子女"), false).unwrap();
419        let view = kb.graph().build_graph(&filter).unwrap();
420        assert_eq!(view.ego_ids(father, 1, 100), vec![son]);
421        assert_eq!(view.path_ids(father, son), Some(vec![father, son]));
422
423        // 登记对称谓词:反向即自身
424        kb.graph().set_predicate_rule("同事", None, true).unwrap();
425        let view = kb.graph().build_graph(&filter).unwrap();
426        assert_eq!(view.ego_ids(yi, 1, 100), vec![jia]);
427
428        // 对称谓词声明逆谓词属于非法组合
429        assert!(kb.graph().set_predicate_rule("同事", Some("同事"), true).is_err());
430    }
431
432    #[test]
433    fn same_as_contracts_alias_nodes() {
434        let dir = tempfile::tempdir().unwrap();
435        let kb = KnowledgeBase::open(dir.path()).unwrap();
436        let filter = ReadFilter::default();
437        let ents = kb
438            .graph()
439            .apply_batch(&GraphBatch { entities: vec![ent("乙"), ent("乙先生"), ent("导师")], ..Default::default() })
440            .unwrap()
441            .value
442            .entities;
443        let (lin, alias, mentor) = (ents[0].header.id, ents[1].header.id, ents[2].header.id);
444        // 别名等价用 sys:same_as 显式关系表达,绝不物理合并实体
445        kb.graph()
446            .apply_batch(&GraphBatch { relations: vec![rel(lin, "sys:same_as", alias), rel(alias, "导师", mentor)], ..Default::default() })
447            .unwrap();
448
449        let view = kb.graph().build_graph(&filter).unwrap();
450        // 三个实体折叠成两个节点:别名与主实体是同一超节点
451        assert_eq!(view.node_count(), 2);
452        assert_eq!(view.edge_count(), 1);
453        // 别名不再打断邻域:乙与乙先生的一跳都直接到导师
454        assert_eq!(view.ego_ids(lin, 1, 100), vec![mentor]);
455        assert_eq!(view.ego_ids(alias, 1, 100), vec![mentor]);
456        // 路径里不再夹着别名中间节点:只剩两个节点
457        let path = view.path_ids(lin, mentor).unwrap();
458        assert_eq!(path.len(), 2);
459        assert!(path.contains(&mentor));
460    }
461}