memoryweave/rust/src/graph_prune.rs

255 lines
8.8 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// 知识图谱修剪 — 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(())
}
}