255 lines
8.8 KiB
Rust
255 lines
8.8 KiB
Rust
// 知识图谱修剪 — 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<PruneStats, Box<dyn std::error::Error>> {
|
||
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<usize, rusqlite::Error> {
|
||
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<usize, rusqlite::Error> {
|
||
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<usize, rusqlite::Error> {
|
||
// 查找 (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<String> = 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<String, f64> = node_ids.iter()
|
||
.map(|id| (id.clone(), 1.0 / n))
|
||
.collect();
|
||
|
||
// 构建出边映射
|
||
let mut out_edges: HashMap<String, Vec<(String, f64)>> = 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<String, f64> = 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(())
|
||
}
|
||
}
|