#[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,
) -> Option<EntityRef> {
match direction {
Direction::Out if &link.source == current => Some(link.target.clone()),
Direction::In if &link.target == current => Some(link.source.clone()),
Direction::Both if &link.source == current => Some(link.target.clone()),
Direction::Both if &link.target == current => Some(link.source.clone()),
_ => None,
}
}
#[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),
Some(target.clone())
);
assert_eq!(
get_neighbor_entity(&link, &target, &Direction::In),
Some(source.clone())
);
assert_eq!(
get_neighbor_entity(&link, &source, &Direction::Both),
Some(target.clone())
);
assert_eq!(
get_neighbor_entity(&link, &target, &Direction::Both),
Some(source.clone())
);
}
#[test]
fn test_get_neighbor_entity_unrelated_node_returns_none() {
let source = EntityRef::External("source".to_string());
let target = EntityRef::External("target".to_string());
let unrelated = EntityRef::External("unrelated".to_string());
let link = Link::new(LinkId::NIL, source.clone(), target.clone(), "test");
assert_eq!(
get_neighbor_entity(&link, &unrelated, &Direction::Out),
None,
"Out direction: current must be source; unrelated node must return None"
);
assert_eq!(
get_neighbor_entity(&link, &unrelated, &Direction::In),
None,
"In direction: current must be target; unrelated node must return None"
);
assert_eq!(
get_neighbor_entity(&link, &unrelated, &Direction::Both),
None,
"Both direction: unrelated node must return None"
);
}
#[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");
}
}
}
}