use std::collections::HashSet;
use crate::store::{Store, StoreError};
use crate::{Edge, EdgeKind, Node};
const DIM: usize = 256;
const EMBED_REF: &str = "embedding:hash/v1";
#[derive(Debug, Clone, Copy)]
pub struct InferenceConfig {
pub min_confidence: f64,
pub top_k: usize,
}
impl Default for InferenceConfig {
fn default() -> Self {
Self {
min_confidence: 0.4,
top_k: 5,
}
}
}
type Embedding = [f32; DIM];
fn fnv1a(s: &str) -> u64 {
let mut hash: u64 = 0xcbf2_9ce4_8422_2325;
for b in s.bytes() {
hash ^= u64::from(b);
hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
}
hash
}
fn tokens(text: &str) -> Vec<String> {
let lower = text.to_lowercase();
let mut out = Vec::new();
for word in lower
.split(|c: char| !c.is_alphanumeric())
.filter(|w| !w.is_empty())
{
out.push(word.to_owned());
let chars: Vec<char> = word.chars().collect();
if chars.len() >= 3 {
for w in chars.windows(3) {
out.push(w.iter().collect());
}
}
}
out
}
#[must_use]
pub fn embed(text: &str) -> Embedding {
let mut v = [0f32; DIM];
for tok in tokens(text) {
let h = fnv1a(&tok);
let idx = usize::try_from(h % DIM as u64).unwrap_or(0);
let sign = if (h >> 63) & 1 == 1 { -1.0 } else { 1.0 };
v[idx] += sign;
}
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for x in &mut v {
*x /= norm;
}
}
v
}
#[must_use]
pub fn similarity(a: &Embedding, b: &Embedding) -> f64 {
let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
f64::from(dot).clamp(0.0, 1.0)
}
fn node_text(node: &Node) -> String {
let mut text = node.name.clone();
if let Some(path) = &node.path
&& let Some(stem) = std::path::Path::new(path)
.file_stem()
.and_then(|s| s.to_str())
{
text.push(' ');
text.push_str(stem);
}
text
}
pub fn infer_edges(store: &Store, config: InferenceConfig) -> Result<Vec<Edge>, StoreError> {
let keys = store.all_keys()?;
let mut nodes: Vec<(String, Embedding)> = Vec::with_capacity(keys.len());
for key in &keys {
if let Some(node) = store.get_node(key)? {
nodes.push((node.key.clone(), embed(&node_text(&node))));
}
}
let mut existing_owned: Vec<(String, String)> = Vec::new();
for (key, _) in &nodes {
for edge in store.edges_from(key)? {
existing_owned.push((edge.src, edge.dst));
}
}
let existing: HashSet<(&str, &str)> = existing_owned
.iter()
.map(|(s, d)| (s.as_str(), d.as_str()))
.collect();
let connected = |a: &str, b: &str| existing.contains(&(a, b)) || existing.contains(&(b, a));
let mut edges = Vec::new();
for (i, (src, src_vec)) in nodes.iter().enumerate() {
let mut candidates: Vec<(f64, &str)> = Vec::new();
for (j, (dst, dst_vec)) in nodes.iter().enumerate() {
if i == j || connected(src, dst) {
continue;
}
let sim = similarity(src_vec, dst_vec);
if sim >= config.min_confidence {
candidates.push((sim, dst.as_str()));
}
}
candidates.sort_by(|a, b| {
b.0.partial_cmp(&a.0)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.1.cmp(b.1))
});
for (sim, dst) in candidates.into_iter().take(config.top_k) {
let mut edge = Edge::inferred(src.clone(), dst.to_owned(), EdgeKind::Related, sim);
edge.src_ref = Some(EMBED_REF.to_owned());
edges.push(edge);
}
}
Ok(edges)
}
#[cfg(test)]
mod tests {
use super::{InferenceConfig, embed, infer_edges, similarity};
use crate::{EdgeKind, FactSet, Node, NodeKind, Provenance, Store};
#[test]
fn embedding_is_deterministic_and_unit_length() {
let a = embed("Store::apply_factset");
let b = embed("Store::apply_factset");
let bits = |v: &[f32; super::DIM]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
assert_eq!(bits(&a), bits(&b), "embedding must be deterministic");
let norm: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-5, "unit length, got {norm}");
}
#[test]
fn similar_names_score_higher_than_unrelated() {
let from = embed("edges_from");
let to = embed("edges_to");
let far = embed("cloudflare deployment pipeline");
let near = similarity(&from, &to);
let distant = similarity(&from, &far);
assert!(
near > distant,
"near {near} should exceed distant {distant}"
);
assert!(near > 0.3, "related names should share features: {near}");
}
#[test]
fn similarity_is_in_range() {
let a = embed("anything at all");
assert!((0.0..=1.0).contains(&similarity(&a, &a)));
assert!((0.0..=1.0).contains(&similarity(&a, &embed(""))));
}
#[test]
fn infers_confident_related_edges_and_skips_known_facts() {
let mut store = Store::open_in_memory().expect("store");
let facts = FactSet::new()
.with_node(Node::new(
"sym:rust:a.rs#edges_from",
NodeKind::Fn,
"edges_from",
))
.with_node(Node::new(
"sym:rust:a.rs#edges_to",
NodeKind::Fn,
"edges_to",
))
.with_node(Node::new(
"sym:rust:a.rs#edges_by_provenance",
NodeKind::Fn,
"edges_by_provenance",
))
.with_node(Node::new("sym:rust:a.rs#unrelated", NodeKind::Fn, "quokka"))
.with_edge(crate::Edge::derived(
"sym:rust:a.rs#edges_from",
"sym:rust:a.rs#edges_to",
EdgeKind::Calls,
));
store.apply_factset(&facts).expect("apply");
let inferred = infer_edges(&store, InferenceConfig::default()).expect("infer");
assert!(
!inferred.is_empty(),
"should infer at least one edge among the unconnected similar fns",
);
for e in &inferred {
assert_eq!(e.provenance, Provenance::Inferred);
assert_eq!(e.kind, EdgeKind::Related);
let c = e.confidence.expect("confidence present");
assert!((0.0..=1.0).contains(&c));
assert!(e.is_valid());
assert!(
!(e.src == "sym:rust:a.rs#edges_from" && e.dst == "sym:rust:a.rs#edges_to"),
"must not re-suggest an existing edge",
);
}
assert!(
inferred
.iter()
.all(|e| e.dst != "sym:rust:a.rs#unrelated" && e.src != "sym:rust:a.rs#unrelated"),
"unrelated node must not be inferred-linked",
);
let mut s2 = store;
s2.apply_factset(&FactSet {
nodes: vec![],
edges: inferred,
})
.expect("inferred edges satisfy store invariants");
assert!(
!s2.edges_by_provenance(Provenance::Inferred)
.expect("q")
.is_empty()
);
}
#[test]
fn re_inferring_is_authoritative_after_clearing() {
let mut store = Store::open_in_memory().expect("store");
let facts = FactSet::new()
.with_node(Node::new(
"sym:rust:a.rs#handle_read",
NodeKind::Fn,
"handle_read",
))
.with_node(Node::new(
"sym:rust:a.rs#handle_write",
NodeKind::Fn,
"handle_write",
))
.with_node(Node::new(
"sym:rust:a.rs#handler_pool",
NodeKind::Fn,
"handler_pool",
));
store.apply_factset(&facts).expect("apply");
let apply = |store: &mut Store, min: f64| {
let edges = infer_edges(
store,
InferenceConfig {
min_confidence: min,
top_k: 5,
},
)
.expect("infer");
store
.apply_factset(&FactSet {
nodes: vec![],
edges,
})
.expect("apply inferred");
};
apply(&mut store, 0.3);
let loose = store
.edges_by_provenance(Provenance::Inferred)
.expect("q")
.len();
assert!(loose > 0);
let removed = store
.delete_edges_by_provenance(Provenance::Inferred)
.expect("delete");
assert_eq!(
usize::try_from(removed).unwrap(),
loose,
"delete removes exactly the inferred edges"
);
assert!(
store
.edges_by_provenance(Provenance::Inferred)
.expect("q")
.is_empty()
);
apply(&mut store, 0.9);
let strict = store
.edges_by_provenance(Provenance::Inferred)
.expect("q")
.len();
assert!(
strict <= loose,
"stricter re-run must not accumulate: {strict} vs {loose}"
);
}
#[test]
fn top_k_bounds_edges_per_source() {
let mut store = Store::open_in_memory().expect("store");
let mut facts = FactSet::new();
for i in 0..10 {
facts = facts.with_node(Node::new(
format!("sym:rust:a.rs#handler{i}"),
NodeKind::Fn,
format!("handler{i}"),
));
}
store.apply_factset(&facts).expect("apply");
let cfg = InferenceConfig {
min_confidence: 0.3,
top_k: 2,
};
let inferred = infer_edges(&store, cfg).expect("infer");
for key in store.all_keys().expect("keys") {
let from_key = inferred.iter().filter(|e| e.src == key).count();
assert!(
from_key <= 2,
"top_k=2 bound exceeded for {key}: {from_key}"
);
}
}
}