use crate::datatypes::values::Value;
use crate::graph::embedding_hints::{missing_column_hint, missing_store_error, Surface};
use crate::graph::embedding_inventory::EmbeddingEntity;
use crate::graph::schema::DirGraph;
use crate::graph::storage::GraphRead;
use super::carry::{key_value, RelationshipKeys};
use super::ingest::{resolve_address, RelationshipVector};
use super::vector_index::{query_edge_embedding_stores, EdgeVectorQueryOptions};
use super::{drop_edge_embedding_store, edge_store_key};
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub struct RelationshipSearchHit {
pub relationship_type: String,
pub source_type: String,
pub source_id: Value,
pub target_type: String,
pub target_id: Value,
pub key: Option<Value>,
pub score: f64,
}
#[derive(Debug, Clone, Default, PartialEq)]
#[non_exhaustive]
pub struct RelationshipSearchOptions {
pub top_k: usize,
pub exact: bool,
pub metric: Option<String>,
}
impl RelationshipSearchOptions {
pub fn new(top_k: usize) -> Self {
Self {
top_k,
exact: false,
metric: None,
}
}
pub fn with_exact(mut self, exact: bool) -> Self {
self.exact = exact;
self
}
pub fn with_metric(mut self, metric: Option<&str>) -> Self {
self.metric = metric.map(str::to_owned);
self
}
}
pub fn search_relationship_embeddings(
graph: &DirGraph,
types: Option<&[String]>,
text_column: &str,
query: &[f32],
options: &RelationshipSearchOptions,
keys: &RelationshipKeys,
) -> Result<Vec<RelationshipSearchHit>, String> {
let types: Vec<String> = match types {
Some([]) => {
return Err(format!(
"types is empty; name at least one relationship type, or omit types to rank \
every '{text_column}' store"
))
}
Some(types) => {
let mut types = types.to_vec();
types.sort();
types.dedup();
types
}
None => {
let mut types: Vec<String> = graph
.edge_embeddings
.keys()
.filter(|(_, store)| {
crate::graph::embeddings::text_column_of(store) == Some(text_column)
})
.map(|(rel_type, _)| rel_type.clone())
.collect();
if types.is_empty() {
return Err(format!(
"No relationship embedding store for text column '{text_column}'.{}",
missing_column_hint(
graph,
EmbeddingEntity::Relationship,
text_column,
Surface::Method
)
));
}
types.sort();
types
}
};
let hits = query_edge_embedding_stores(
graph,
&types,
text_column,
query,
EdgeVectorQueryOptions {
top_k: options.top_k,
exact: options.exact,
metric: options.metric.clone(),
},
Surface::Method,
)?;
let _guard = graph.graph.begin_query();
Ok(hits
.into_iter()
.filter_map(|hit| {
let (source, target) = graph.graph.edge_endpoints(hit.edge)?;
let source_view = graph.graph.node_view(source)?;
let target_view = graph.graph.node_view(target)?;
let key = keys
.get(&hit.rel_type)
.and_then(|property| key_value(graph, hit.edge, property));
Some(RelationshipSearchHit {
source_type: source_view.node_type_str(&graph.interner).to_string(),
source_id: source_view.id().into_owned(),
target_type: target_view.node_type_str(&graph.interner).to_string(),
target_id: target_view.id().into_owned(),
relationship_type: hit.rel_type,
key,
score: hit.score,
})
})
.collect())
}
pub fn relationship_embedding(
graph: &DirGraph,
relationship_type: &str,
text_column: &str,
address: &RelationshipVector,
keys: &RelationshipKeys,
) -> Result<Option<Vec<f32>>, String> {
let store = graph
.edge_embeddings
.get(&edge_store_key(relationship_type, text_column))
.ok_or_else(|| {
missing_store_error(
graph,
EmbeddingEntity::Relationship,
relationship_type,
text_column,
Surface::Method,
)
})?;
let edge = resolve_address(graph, relationship_type, address, keys)?;
Ok(store.get(edge).map(<[f32]>::to_vec))
}
pub fn relationship_embedding_dim(
graph: &DirGraph,
relationship_type: &str,
text_column: &str,
) -> Option<usize> {
graph
.edge_embeddings
.get(&edge_store_key(relationship_type, text_column))
.map(|store| store.dimension())
}
pub fn remove_relationship_embeddings(
graph: &mut DirGraph,
relationship_type: &str,
text_column: &str,
) -> Result<(), String> {
if drop_edge_embedding_store(graph, relationship_type, text_column)? {
return Ok(());
}
Err(missing_store_error(
graph,
EmbeddingEntity::Relationship,
relationship_type,
text_column,
Surface::Method,
))
}