// 知识图谱修剪 — SQLite 操作 (rusqlite) // 深度整合步骤 2:孤立节点/低权重边/冗余边合并/PageRank 更新 use rusqlite::{Connection, params}; use serde::{Deserialize, Serialize}; use std::collections::{HashMap, HashSet}; use std::path::Path; /// 修剪统计 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct PruneStats { pub isolated_nodes_removed: usize, pub low_weight_edges_removed: usize, pub redundant_edges_merged: usize, pub nodes_before: usize, pub nodes_after: usize, pub edges_before: usize, pub edges_after: usize, } /// 图谱修剪器 pub struct GraphPruner { db_path: String, /// 孤立节点阈值: 14天无关联边 isolated_days: u32, /// 低权重阈值 low_weight_threshold: f64, /// 最多 BFS 跳数 max_hops: usize, } impl GraphPruner { pub fn new(db_path: &str) -> Self { Self { db_path: db_path.to_string(), isolated_days: 14, low_weight_threshold: 0.15, max_hops: 3, } } /// 执行全部修剪流程 pub fn prune(&self) -> Result> { let conn = Connection::open(&self.db_path)?; // 统计修剪前 let nodes_before: usize = conn.query_row("SELECT COUNT(*) FROM graph_nodes", [], |r| r.get(0))?; let edges_before: usize = conn.query_row("SELECT COUNT(*) FROM graph_edges", [], |r| r.get(0))?; // Step 1: 删除孤立节点 (14天无关联边) let isolated_removed = self.remove_isolated_nodes(&conn)?; // Step 2: 删除低权重边 (weight < 0.15) let low_weight_removed = self.remove_low_weight_edges(&conn)?; // Step 3: 合并冗余边 (同 source→target 的多条边 → 取 weight 加权平均) let redundant_merged = self.merge_redundant_edges(&conn)?; // Step 4: 更新 PageRank self.update_pagerank(&conn)?; // 统计修剪后 let nodes_after: usize = conn.query_row("SELECT COUNT(*) FROM graph_nodes", [], |r| r.get(0))?; let edges_after: usize = conn.query_row("SELECT COUNT(*) FROM graph_edges", [], |r| r.get(0))?; let stats = PruneStats { isolated_nodes_removed: isolated_removed, low_weight_edges_removed: low_weight_removed, redundant_edges_merged: redundant_merged, nodes_before, nodes_after, edges_before, edges_after, }; eprintln!( "[graph_prune] nodes {}→{}, edges {}→{} (isolated={}, low_wt={}, merged={})", nodes_before, nodes_after, edges_before, edges_after, isolated_removed, low_weight_removed, redundant_merged, ); conn.close().ok(); Ok(stats) } /// 删除孤立节点 fn remove_isolated_nodes(&self, conn: &Connection) -> Result { let cutoff = format!("-{} days", self.isolated_days); let deleted = conn.execute( "DELETE FROM graph_nodes WHERE id IN ( SELECT n.id FROM graph_nodes n LEFT JOIN graph_edges e ON n.id = e.source OR n.id = e.target WHERE e.id IS NULL AND n.created_at < datetime('now', '-14 days') )", [], )?; Ok(deleted) } /// 删除低权重边 fn remove_low_weight_edges(&self, conn: &Connection) -> Result { let deleted = conn.execute( "DELETE FROM graph_edges WHERE weight < ?1", params![self.low_weight_threshold], )?; Ok(deleted) } /// 合并冗余边 fn merge_redundant_edges(&self, conn: &Connection) -> Result { // 查找 (source, target, relation) 相同的冗余边 let mut stmt = conn.prepare( "SELECT source, target, relation, COUNT(*) as cnt, SUM(weight) as total_weight, SUM(evidence_count) as total_evidence FROM graph_edges GROUP BY source, target, relation HAVING cnt > 1" )?; let to_merge: Vec<(String, String, String, i64, f64, i64)> = stmt.query_map( [], |row| Ok(( row.get::<_, String>(0)?, row.get::<_, String>(1)?, row.get::<_, String>(2)?, row.get::<_, i64>(3)?, row.get::<_, f64>(4)?, row.get::<_, i64>(5)?, )) )?.filter_map(|r| r.ok()).collect(); let mut merged = 0usize; for (src, tgt, rel, cnt, total_wt, total_ev) in to_merge { // 删除所有冗余边 let deleted = conn.execute( "DELETE FROM graph_edges WHERE source=?1 AND target=?2 AND relation=?3", params![src, tgt, rel], )?; // 插入合并后单边(加权平均) let avg_weight = total_wt / cnt as f64; let avg_evidence = total_ev / cnt; conn.execute( "INSERT INTO graph_edges (id, source, target, relation, weight, evidence_count, namespace, created_at) VALUES (?1, ?2, ?3, ?4, ?5, ?6, 'shared', datetime('now'))", params![ format!("merged_{}_{}", src, tgt), src, tgt, rel, avg_weight, avg_evidence, ], )?; merged += deleted as usize - 1; // -1 因为已插入新边 } Ok(merged) } /// 幂等添加列(迁移用) fn ensure_column(&self, conn: &Connection, table: &str, column: &str, def: &str) -> Result<(), rusqlite::Error> { let check_sql = format!("SELECT 1 FROM pragma_table_info('{}') WHERE name='{}'", table, column); let exists: bool = conn.query_row(&check_sql, [], |_| Ok(true)).unwrap_or(false); if !exists { let alter_sql = format!("ALTER TABLE {} ADD COLUMN {} {}", table, column, def); conn.execute(&alter_sql, [])?; eprintln!("[graph_prune] migrated: added {} column to {} (def: {})", column, table, def); } Ok(()) } /// 简单 PageRank 迭代更新 fn update_pagerank(&self, conn: &Connection) -> Result<(), rusqlite::Error> { // 确保 pagerank 列存在(可能由 Go migrate 创建,也可能没有) self.ensure_column(conn, "graph_nodes", "pagerank", "REAL DEFAULT 1.0")?; let damping = 0.85; let iterations = 20; // 获取所有节点 let mut node_ids: Vec = Vec::new(); let mut stmt = conn.prepare("SELECT id FROM graph_nodes")?; let rows = stmt.query_map([], |row| row.get::<_, String>(0))?; for row in rows { node_ids.push(row?); } if node_ids.is_empty() { return Ok(()); } let n = node_ids.len() as f64; let base = (1.0 - damping) / n; let mut ranks: HashMap = node_ids.iter() .map(|id| (id.clone(), 1.0 / n)) .collect(); // 构建出边映射 let mut out_edges: HashMap> = HashMap::new(); for node in &node_ids { out_edges.insert(node.clone(), Vec::new()); } let mut edge_stmt = conn.prepare( "SELECT source, target, weight FROM graph_edges" )?; let edge_rows = edge_stmt.query_map([], |row| Ok(( row.get::<_, String>(0)?, row.get::<_, String>(1)?, row.get::<_, f64>(2)?, )))?; for edge in edge_rows { if let Ok((src, tgt, wt)) = edge { out_edges.entry(src).or_insert_with(Vec::new).push((tgt, wt)); } } // PageRank 迭代 for _ in 0..iterations { let mut new_ranks: HashMap = HashMap::new(); for node in &node_ids { let mut rank = base; for (other_node, edges) in &out_edges { for (tgt, wt) in edges { if tgt == node { let total_out_wt: f64 = edges.iter().map(|(_, w)| w).sum(); if total_out_wt > 0.0 { rank += damping * ranks.get(other_node).unwrap_or(&0.0) * wt / total_out_wt; } } } } new_ranks.insert(node.clone(), rank); } ranks = new_ranks; } // 写入数据库 for (node_id, rank) in &ranks { conn.execute( "UPDATE graph_nodes SET pagerank = ?1 WHERE id = ?2", params![rank, node_id], )?; } eprintln!("[graph_prune] PageRank updated for {} nodes ({} iterations)", n as usize, iterations); Ok(()) } }