#[cfg(test)]
use super::compat::LinkId;
use super::compat::{EntityRef, Link, LinkStore, StorageContext};
use khive_score::DeterministicScore;
use crate::error::{Result, RetrievalError};
use super::types::Direction;
pub fn get_edge_weight(link: &Link) -> f64 {
link.properties
.as_ref()
.and_then(|props| props.get("weight"))
.and_then(|v| v.as_f64())
.unwrap_or(1.0)
}
pub fn matches_link_type(link: &Link, filter: &Option<Vec<String>>) -> bool {
match filter {
None => true,
Some(types) => types.iter().any(|t| t == &link.relation),
}
}
pub async fn get_neighbors<S: LinkStore>(
store: &S,
ctx: &StorageContext,
entity: &EntityRef,
direction: &Direction,
) -> Result<Vec<Link>> {
let links =
match direction {
Direction::Out => store
.outgoing(ctx, entity)
.await
.map_err(|e| RetrievalError::GraphTraversal(format!("link store error: {e}"))),
Direction::In => store
.incoming(ctx, entity)
.await
.map_err(|e| RetrievalError::GraphTraversal(format!("link store error: {e}"))),
Direction::Both => {
let mut out = store.outgoing(ctx, entity).await.map_err(|e| {
RetrievalError::GraphTraversal(format!("link store error: {e}"))
})?;
let incoming = store.incoming(ctx, entity).await.map_err(|e| {
RetrievalError::GraphTraversal(format!("link store error: {e}"))
})?;
out.extend(incoming);
Ok(out)
}
};
links
}
pub fn proximity_score(depth: usize, max_depth: usize) -> DeterministicScore {
if max_depth == 0 {
return DeterministicScore::from_f64(if depth == 0 { 1.0 } else { 0.0 });
}
let proximity = 1.0 - (depth as f64 / max_depth as f64);
DeterministicScore::from_f64(proximity)
}
pub fn get_neighbor_entity(link: &Link, current: &EntityRef, direction: &Direction) -> EntityRef {
match direction {
Direction::Out => link.target.clone(),
Direction::In => link.source.clone(),
Direction::Both => {
if &link.source == current {
link.target.clone()
} else {
link.source.clone()
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_matches_link_type() {
let link = Link::new(
LinkId::NIL,
EntityRef::External("a".to_string()),
EntityRef::External("b".to_string()),
"contains",
);
assert!(matches_link_type(&link, &None));
assert!(matches_link_type(
&link,
&Some(vec!["contains".to_string()])
));
assert!(!matches_link_type(
&link,
&Some(vec!["references".to_string()])
));
assert!(matches_link_type(
&link,
&Some(vec!["references".to_string(), "contains".to_string()])
));
}
#[test]
fn test_get_edge_weight() {
let link = Link::new(
LinkId::NIL,
EntityRef::External("a".to_string()),
EntityRef::External("b".to_string()),
"test",
);
assert_eq!(get_edge_weight(&link), 1.0);
let link_with_weight = Link::with_properties(
LinkId::NIL,
EntityRef::External("a".to_string()),
EntityRef::External("b".to_string()),
"test",
serde_json::json!({"weight": 2.5}),
);
assert_eq!(get_edge_weight(&link_with_weight), 2.5);
}
#[test]
fn test_get_neighbor_entity() {
let source = EntityRef::External("source".to_string());
let target = EntityRef::External("target".to_string());
let link = Link::new(LinkId::NIL, source.clone(), target.clone(), "test");
assert_eq!(get_neighbor_entity(&link, &source, &Direction::Out), target);
assert_eq!(get_neighbor_entity(&link, &target, &Direction::In), source);
assert_eq!(
get_neighbor_entity(&link, &source, &Direction::Both),
target
);
assert_eq!(
get_neighbor_entity(&link, &target, &Direction::Both),
source
);
}
#[test]
fn test_proximity_score_normal() {
let score = proximity_score(0, 5);
assert!((score.to_f64() - 1.0).abs() < f64::EPSILON);
let score = proximity_score(5, 5);
assert!((score.to_f64() - 0.0).abs() < f64::EPSILON);
let score = proximity_score(2, 4);
assert!((score.to_f64() - 0.5).abs() < f64::EPSILON);
}
#[test]
fn test_proximity_score_max_depth_zero() {
let score = proximity_score(0, 0);
assert!((score.to_f64() - 1.0).abs() < f64::EPSILON);
let score = proximity_score(1, 0);
assert!((score.to_f64() - 0.0).abs() < f64::EPSILON);
}
#[test]
fn test_proximity_score_monotonic() {
let max_depth = 10;
let mut prev_score = f64::MAX;
for depth in 0..=max_depth {
let score = proximity_score(depth, max_depth).to_f64();
assert!(
score <= prev_score,
"Score should be monotonically decreasing"
);
prev_score = score;
}
}
#[test]
fn test_proximity_score_bounded() {
for max_depth in [0, 1, 5, 10, 100] {
for depth in 0..=max_depth {
let score = proximity_score(depth, max_depth).to_f64();
assert!(score >= 0.0, "Score should be >= 0");
assert!(score <= 1.0, "Score should be <= 1");
}
}
}
}