use crate::datatypes::Value;
use crate::graph::algorithms::hnsw::HnswParams;
use crate::graph::algorithms::vector::DistanceMetric;
use crate::graph::dir_graph::DirGraph;
use crate::graph::schema::EmbeddingStore;
use crate::graph::storage::GraphRead;
use petgraph::graph::NodeIndex;
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct EmbeddingIngestReport {
pub embeddings_stored: usize,
pub dimension: usize,
pub skipped: usize,
pub store_created: bool,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct VectorIndexReport {
pub indexed: usize,
pub metric: String,
pub m: usize,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct EmbeddingStoreInfo {
pub node_type: String,
pub text_column: String,
pub store_name: String,
pub dimension: usize,
pub count: usize,
pub metric: String,
}
pub fn store_name(text_column: &str) -> String {
format!("{}_emb", text_column)
}
pub fn store_key(node_type: &str, text_column: &str) -> (String, String) {
(node_type.to_string(), store_name(text_column))
}
pub fn text_column_of(store: &str) -> Option<&str> {
store.strip_suffix("_emb")
}
pub fn embedded_text_columns<'a>(graph: &'a DirGraph, node_types: &[&str]) -> Vec<&'a str> {
let mut columns: Vec<&str> = graph
.embeddings
.keys()
.filter(|(stored_type, _)| node_types.contains(&stored_type.as_str()))
.map(|(_, name)| text_column_of(name).unwrap_or(name))
.collect();
columns.sort_unstable();
columns.dedup();
columns
}
pub fn unknown_column_hint(
graph: &DirGraph,
node_types: &[&str],
text_column: &str,
caller: &str,
) -> String {
if let Some(stripped) = text_column_of(text_column) {
if node_types.iter().any(|node_type| {
graph
.embedding_store(node_type, &store_name(stripped))
.is_some()
}) {
return format!(
" Did you mean '{stripped}'? {caller} takes the text column; \
'{text_column}' is the embedding store's own name."
);
}
}
let columns = embedded_text_columns(graph, node_types);
let suggestion = crate::graph::mutation::validation::did_you_mean(text_column, &columns);
if !suggestion.is_empty() {
return suggestion;
}
if columns.is_empty() {
String::new()
} else {
format!(" Embedded text columns: {}.", columns.join(", "))
}
}
pub fn list_embeddings(graph: &DirGraph) -> Vec<EmbeddingStoreInfo> {
graph
.embeddings
.iter()
.map(|((node_type, name), store)| EmbeddingStoreInfo {
node_type: node_type.clone(),
text_column: text_column_of(name).unwrap_or(name).to_string(),
store_name: name.clone(),
dimension: store.dimension,
count: store.len(),
metric: store.metric.as_deref().unwrap_or("cosine").to_string(),
})
.collect()
}
pub fn set_embeddings<I, V>(
graph: &mut DirGraph,
node_type: &str,
text_column: &str,
metric: Option<&str>,
entries: I,
) -> Result<EmbeddingIngestReport, String>
where
I: IntoIterator<Item = (Value, V)>,
V: AsRef<[f32]>,
{
let key = store_key(node_type, text_column);
let prepared = prepare(graph, node_type, text_column, None, entries)?;
let Some(dim) = prepared.dimension else {
return Ok(EmbeddingIngestReport {
skipped: prepared.skipped,
..Default::default()
});
};
let mut store = match metric {
Some(m) => EmbeddingStore::with_metric(dim, m),
None => EmbeddingStore::new(dim),
};
store.data.reserve(prepared.entries.len() * dim);
for (node_idx, vector) in &prepared.entries {
store.set_embedding(node_idx.index(), vector.as_ref());
}
let embeddings_stored = store.len();
graph.embeddings.insert(key, store);
graph.bump_version();
Ok(EmbeddingIngestReport {
embeddings_stored,
dimension: dim,
skipped: prepared.skipped,
store_created: true,
})
}
pub fn add_embeddings<I, V>(
graph: &mut DirGraph,
node_type: &str,
text_column: &str,
metric: Option<&str>,
entries: I,
) -> Result<EmbeddingIngestReport, String>
where
I: IntoIterator<Item = (Value, V)>,
V: AsRef<[f32]>,
{
let key = store_key(node_type, text_column);
let existing_dim = graph.embeddings.get(&key).map(|s| s.dimension);
let store_existed = existing_dim.is_some();
let prepared = prepare(graph, node_type, text_column, existing_dim, entries)?;
let Some(dim) = prepared.dimension else {
return Ok(EmbeddingIngestReport {
skipped: prepared.skipped,
..Default::default()
});
};
let store = graph.embeddings.entry(key).or_insert_with(|| match metric {
Some(m) => EmbeddingStore::with_metric(dim, m),
None => EmbeddingStore::new(dim),
});
for (node_idx, vector) in &prepared.entries {
store.set_embedding(node_idx.index(), vector.as_ref());
}
let embeddings_stored = store.len();
graph.bump_version();
Ok(EmbeddingIngestReport {
embeddings_stored,
dimension: dim,
skipped: prepared.skipped,
store_created: !store_existed,
})
}
pub fn build_vector_index(
graph: &mut DirGraph,
node_type: &str,
text_column: &str,
m: Option<usize>,
ef_construction: Option<usize>,
ef_search: Option<usize>,
metric: Option<&str>,
) -> Result<VectorIndexReport, String> {
let key = store_key(node_type, text_column);
let metric_name = match metric {
Some(m) => m.to_string(),
None => graph
.embeddings
.get(&key)
.and_then(|s| s.metric.clone())
.unwrap_or_else(|| "cosine".to_string()),
};
let distance = match metric_name.as_str() {
"cosine" => DistanceMetric::Cosine,
"dot_product" => DistanceMetric::DotProduct,
"euclidean" => DistanceMetric::Euclidean,
"poincare" => {
return Err(
"build_vector_index: the 'poincare' metric is not supported by HNSW; \
Poincaré search stays on the exact (brute-force) path."
.to_string(),
)
}
other => {
return Err(format!(
"Unknown metric '{}'. Use 'cosine', 'dot_product', or 'euclidean'.",
other
))
}
};
let defaults = HnswParams::default();
let params = HnswParams {
m: m.unwrap_or(defaults.m).max(2),
ef_construction: ef_construction.unwrap_or(defaults.ef_construction).max(1),
ef_search: ef_search.unwrap_or(defaults.ef_search).max(1),
};
if !graph.embeddings.contains_key(&key) {
let hint = unknown_column_hint(graph, &[node_type], text_column, "build_vector_index()");
return Err(format!(
"No embedding store '{}.{}' to index.{} Call set_embeddings()/embed_texts() first.",
node_type,
store_name(text_column),
hint
));
}
let store = graph
.embeddings
.get_mut(&key)
.expect("store presence checked immediately above");
let indexed = store.len();
let seed = 0x9E37_79B9_7F4A_7C15 ^ (indexed as u64);
store.build_index(distance, params, seed)?;
Ok(VectorIndexReport {
indexed,
metric: metric_name,
m: params.m,
})
}
struct Prepared<V> {
entries: Vec<(NodeIndex, V)>,
dimension: Option<usize>,
skipped: usize,
}
fn prepare<I, V>(
graph: &mut DirGraph,
node_type: &str,
text_column: &str,
constraint: Option<usize>,
entries: I,
) -> Result<Prepared<V>, String>
where
I: IntoIterator<Item = (Value, V)>,
V: AsRef<[f32]>,
{
let _arena_guard = graph.graph.begin_query();
if !graph.type_indices.contains_key(node_type) {
return Err(format!(
"Node type '{}' does not exist in the graph",
node_type
));
}
let mut incoming = entries.into_iter().peekable();
let non_empty = incoming.peek().is_some();
if non_empty {
resolve_source_column(graph, node_type, text_column)?;
}
graph.build_id_index(node_type);
let mut resolved: Vec<(NodeIndex, V)> = Vec::new();
let mut skipped = 0usize;
let mut dimension = constraint;
for (id, vector) in incoming {
let Some(node_idx) = graph.lookup_by_id(node_type, &id) else {
skipped += 1;
continue;
};
let len = vector.as_ref().len();
match dimension {
None => dimension = Some(len),
Some(d) if len != d => {
return Err(match constraint {
Some(_) => format!(
"Inconsistent embedding dimension: store has {} but got {}",
d, len
),
None => format!(
"Inconsistent embedding dimensions: expected {} but got {}",
d, len
),
})
}
Some(_) => {}
}
resolved.push((node_idx, vector));
}
if resolved.is_empty() {
dimension = None;
}
Ok(Prepared {
entries: resolved,
dimension,
skipped,
})
}
pub fn resolve_source_column<'a>(
graph: &'a DirGraph,
node_type: &str,
text_column: &'a str,
) -> Result<&'a str, String> {
let resolved = graph.resolve_alias(node_type, text_column);
if matches!(resolved, "id" | "title") {
return Ok(resolved);
}
let present = graph
.type_indices
.get(node_type)
.map(|indices| {
indices.iter().any(|idx| {
graph
.graph
.node_view(idx)
.map(|n| n.has_property(resolved))
.unwrap_or(false)
})
})
.unwrap_or(false);
if present {
return Ok(resolved);
}
if crate::graph::schema::soft_alias_fallback(resolved).is_some() {
return Ok(resolved);
}
Err(format!(
"Source column '{}' not found on any '{}' node. \
set_embeddings() expects the text column name \
(e.g. 'summary'), not the embedding store name.",
text_column, node_type
))
}
#[cfg(test)]
#[path = "embeddings_tests.rs"]
mod tests;