use std::collections::{HashMap, HashSet};
use anyhow::Result;
use leiden_rs::{GraphDataBuilder, Leiden, LeidenConfig, QualityType};
use petgraph::visit::{EdgeRef, IntoEdgeReferences, IntoNodeReferences};
use crate::model::*;
const FEATURE_LEIDEN_SEED: u64 = 42;
const FEATURE_RESOLUTION: f64 = 0.4;
const STRUCTURE_WEIGHT: f64 = 0.5;
const SEMANTIC_WEIGHT: f64 = 0.5;
pub trait Embedder: Send + Sync {
fn embed(&self, text: &str) -> Result<Vec<f32>>;
fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>>;
fn cosine_similarity(&self, a: &[f32], b: &[f32]) -> f64;
}
pub fn detect_features(
graph: &KnowledgeGraph,
embedder: Option<&dyn Embedder>,
) -> Result<Vec<Feature>> {
let funcs: Vec<NodeId> = graph
.graph
.node_references()
.filter(|(_, n)| n.kind == NodeKind::Function)
.map(|(id, _)| id)
.collect();
if funcs.is_empty() {
return Ok(Vec::new());
}
let func_set: HashSet<NodeId> = funcs.iter().copied().collect();
let mut entity_to_file: HashMap<NodeId, NodeId> = HashMap::new();
for edge in graph.graph.edge_references() {
let kind = graph.graph.edge_weight(edge.id()).map(|e| e.kind.clone());
if kind == Some(EdgeKind::Contains) {
entity_to_file.insert(edge.target(), edge.source());
}
}
let mut neighbors: HashMap<NodeId, HashSet<NodeId>> = HashMap::new();
let mut cross_edges: Vec<(NodeId, NodeId)> = Vec::new();
for edge in graph.graph.edge_references() {
let e = graph
.graph
.edge_weight(edge.id())
.expect("边权重必然存在");
if e.kind != EdgeKind::Calls {
continue;
}
let (s, t) = (edge.source(), edge.target());
if !func_set.contains(&s) || !func_set.contains(&t) {
continue;
}
if entity_to_file.get(&s) == entity_to_file.get(&t) {
continue; }
neighbors.entry(s).or_default().insert(t);
neighbors.entry(t).or_default().insert(s);
cross_edges.push((s, t));
}
if cross_edges.is_empty() {
return Ok(Vec::new());
}
let embeddings: Option<HashMap<NodeId, Vec<f32>>> = if let Some(emb) = embedder {
let mut involved: Vec<NodeId> = cross_edges
.iter()
.flat_map(|(s, t)| [*s, *t])
.collect();
involved.sort();
involved.dedup();
let texts: Vec<String> = involved
.iter()
.map(|nid| {
let n = graph.graph.node_weight(*nid).expect("实体节点必然存在");
format!(
"{} {:?} {}",
n.name,
n.kind,
n.signature.as_deref().unwrap_or("")
)
})
.collect();
match emb.embed_batch(&texts) {
Ok(vecs) => Some(involved.into_iter().zip(vecs).collect()),
Err(e) => {
tracing::warn!("特征聚类 embedding 失败,降级为纯结构聚类: {e}");
None
}
}
} else {
None
};
let compact: HashMap<NodeId, usize> = funcs
.iter()
.enumerate()
.map(|(i, &n)| (n, i))
.collect();
let mut weights: HashMap<(usize, usize), f64> = HashMap::new();
for (s, t) in &cross_edges {
let structural = 0.5 + 0.5 * jaccard(neighbors.get(s), neighbors.get(t));
let semantic = match (&embeddings, embedder) {
(Some(em), Some(emb)) => match (em.get(s), em.get(t)) {
(Some(a), Some(b)) => emb.cosine_similarity(a, b),
_ => 0.0,
},
_ => 0.0,
};
let weight = STRUCTURE_WEIGHT * structural + SEMANTIC_WEIGHT * semantic;
let (si, ti) = (compact[s], compact[t]);
*weights.entry((si, ti)).or_insert(0.0) += weight;
}
let mut builder = GraphDataBuilder::new(funcs.len()).directed();
for ((s, t), w) in &weights {
builder
.add_edge(*s, *t, *w)
.expect("边权重均为有限非负数(Embedder 契约:余弦相似度必须有限)");
}
let data = builder.build().expect("图数据构造失败");
let config = LeidenConfig {
quality: QualityType::CPM,
resolution: FEATURE_RESOLUTION,
seed: Some(FEATURE_LEIDEN_SEED),
..Default::default()
};
let result = Leiden::new(config)
.run(&data)
.expect("Leiden 特征聚类失败");
let membership = result.partition.as_slice();
let mut groups: HashMap<usize, Vec<NodeId>> = HashMap::new();
for (i, &comm) in membership.iter().enumerate() {
groups.entry(comm).or_default().push(funcs[i]);
}
let mut features: Vec<Vec<NodeId>> = groups
.into_values()
.map(|mut node_ids| {
node_ids.sort_by_key(|nid| {
graph
.graph
.node_weight(*nid)
.map(|n| n.name.clone())
.unwrap_or_default()
});
node_ids
})
.collect();
features.sort_by_key(|node_ids| {
node_ids
.first()
.map(|nid| {
graph
.graph
.node_weight(*nid)
.map(|n| n.name.clone())
.unwrap_or_default()
})
.unwrap_or_default()
});
Ok(features
.into_iter()
.enumerate()
.map(|(idx, node_ids)| Feature {
name: format!("feature_{idx}"),
node_ids,
description: None,
})
.collect())
}
fn jaccard(a: Option<&HashSet<NodeId>>, b: Option<&HashSet<NodeId>>) -> f64 {
match (a, b) {
(Some(a), Some(b)) => {
let union: HashSet<NodeId> = a.union(b).copied().collect();
if union.is_empty() {
0.0
} else {
a.intersection(b).count() as f64 / union.len() as f64
}
}
_ => 0.0,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_graph() -> KnowledgeGraph {
let mut kg = KnowledgeGraph::default();
let g = &mut kg.graph;
let add_file =
|g: &mut petgraph::stable_graph::StableDiGraph<CodeNode, CodeEdge>,
path: &str|
-> (NodeId, NodeId) {
let nid = g.add_node(CodeNode {
id: NodeId::new(g.node_count()),
kind: NodeKind::File,
name: path.into(),
file_path: Some(path.into()),
line_range: None,
doc_comment: None,
signature: None, visibility: None,
module_path: vec!["src".into()],
});
let eid = g.add_node(CodeNode {
id: NodeId::new(g.node_count()),
kind: NodeKind::Function,
name: format!("fn_{}", path.replace(['/', '.'], "_")),
file_path: Some(path.into()),
line_range: None,
doc_comment: None,
signature: None, visibility: None,
module_path: Vec::new(),
});
g.add_edge(
nid,
eid,
CodeEdge {
id: EdgeId::new(g.edge_count()),
kind: EdgeKind::Contains,
source: nid,
target: eid,
weight: 1.0,
location: None,
},
);
(nid, eid)
};
let (_fa_file, fa) = add_file(g, "src/a.rs");
let (_fb_file, fb) = add_file(g, "src/b.rs");
let (_fc_file, fc) = add_file(g, "src/c.rs");
let (_fd_file, fd) = add_file(g, "src/d.rs");
for (s, t) in [(fa, fb), (fc, fd)] {
g.add_edge(
s,
t,
CodeEdge {
id: EdgeId::new(g.edge_count()),
kind: EdgeKind::Calls,
source: s,
target: t,
weight: 0.7,
location: None,
},
);
}
kg
}
#[test]
fn test_detect_features_basic() {
let kg = make_graph();
let features = detect_features(&kg, None).unwrap();
assert!(features.len() >= 2, "应检出至少 2 个特征: {:?}", features);
let names: Vec<String> = features
.iter()
.flat_map(|f| {
f.node_ids
.iter()
.map(|nid| kg.graph.node_weight(*nid).unwrap().name.clone())
.collect::<Vec<_>>()
})
.collect();
assert!(names.contains(&"fn_src_a_rs".to_string()));
assert!(names.contains(&"fn_src_b_rs".to_string()));
}
#[test]
fn test_detect_features_empty_graph() {
let kg = KnowledgeGraph::default();
let features = detect_features(&kg, None).unwrap();
assert!(features.is_empty());
}
#[test]
fn test_detect_features_no_cross_file_calls() {
let mut kg = KnowledgeGraph::default();
let g = &mut kg.graph;
let f = g.add_node(CodeNode {
id: NodeId::new(0),
kind: NodeKind::File,
name: "src/a.rs".into(),
file_path: Some("src/a.rs".into()),
line_range: None,
doc_comment: None,
signature: None, visibility: None,
module_path: vec!["src".into()],
});
let e1 = g.add_node(CodeNode {
id: NodeId::new(1),
kind: NodeKind::Function,
name: "f1".into(),
file_path: Some("src/a.rs".into()),
line_range: None,
doc_comment: None,
signature: None, visibility: None,
module_path: Vec::new(),
});
let e2 = g.add_node(CodeNode {
id: NodeId::new(2),
kind: NodeKind::Function,
name: "f2".into(),
file_path: Some("src/a.rs".into()),
line_range: None,
doc_comment: None,
signature: None, visibility: None,
module_path: Vec::new(),
});
for e in [e1, e2] {
g.add_edge(
f,
e,
CodeEdge {
id: EdgeId::new(g.edge_count()),
kind: EdgeKind::Contains,
source: f,
target: e,
weight: 1.0,
location: None,
},
);
}
g.add_edge(
e1,
e2,
CodeEdge {
id: EdgeId::new(g.edge_count()),
kind: EdgeKind::Calls,
source: e1,
target: e2,
weight: 0.7,
location: None,
},
);
let features = detect_features(&kg, None).unwrap();
assert!(features.is_empty(), "同文件调用不构成特征");
}
}