use std::collections::{HashMap, VecDeque};
use super::compat::{EntityRef, Link, LinkStore, StorageContext};
use crate::error::{Result, RetrievalError};
use super::types::{PathNode, MAX_TRAVERSAL_DEPTH};
pub async fn find_shortest_path<S: LinkStore>(
store: &S,
ctx: &StorageContext,
from: EntityRef,
to: EntityRef,
max_depth: usize,
) -> Result<Option<Vec<PathNode>>> {
let max_depth = max_depth.min(MAX_TRAVERSAL_DEPTH);
if from == to {
return Ok(Some(vec![PathNode::start(from)]));
}
let mut forward_visited: HashMap<EntityRef, (usize, Option<EntityRef>, Option<Link>)> =
HashMap::new();
let mut forward_queue: VecDeque<EntityRef> = VecDeque::new();
forward_visited.insert(from.clone(), (0, None, None));
forward_queue.push_back(from.clone());
let mut backward_visited: HashMap<EntityRef, (usize, Option<EntityRef>, Option<Link>)> =
HashMap::new();
let mut backward_queue: VecDeque<EntityRef> = VecDeque::new();
backward_visited.insert(to.clone(), (0, None, None));
backward_queue.push_back(to.clone());
let mut best_meeting: Option<(EntityRef, usize)> = None; let mut current_depth = 0;
while !forward_queue.is_empty() || !backward_queue.is_empty() {
if current_depth > max_depth {
break;
}
let forward_level_size = forward_queue.len();
for _ in 0..forward_level_size {
if let Some(current) = forward_queue.pop_front() {
let outgoing = store.outgoing(ctx, ¤t).await.map_err(|e| {
RetrievalError::GraphTraversal(format!("link store error: {e}"))
})?;
for link in outgoing {
let neighbor = link.target.clone();
if !forward_visited.contains_key(&neighbor) {
let fwd_dist = current_depth + 1;
forward_visited.insert(
neighbor.clone(),
(fwd_dist, Some(current.clone()), Some(link)),
);
forward_queue.push_back(neighbor.clone());
if let Some((bwd_dist, _, _)) = backward_visited.get(&neighbor) {
let total = fwd_dist + bwd_dist;
if best_meeting.as_ref().is_none_or(|&(_, best)| total < best) {
best_meeting = Some((neighbor, total));
}
}
}
}
}
}
if best_meeting.is_some() {
break;
}
let backward_level_size = backward_queue.len();
for _ in 0..backward_level_size {
if let Some(current) = backward_queue.pop_front() {
let incoming = store.incoming(ctx, ¤t).await.map_err(|e| {
RetrievalError::GraphTraversal(format!("link store error: {e}"))
})?;
for link in incoming {
let neighbor = link.source.clone();
if !backward_visited.contains_key(&neighbor) {
let bwd_dist = current_depth + 1;
backward_visited.insert(
neighbor.clone(),
(bwd_dist, Some(current.clone()), Some(link)),
);
backward_queue.push_back(neighbor.clone());
if let Some((fwd_dist, _, _)) = forward_visited.get(&neighbor) {
let total = fwd_dist + bwd_dist;
if best_meeting.as_ref().is_none_or(|&(_, best)| total < best) {
best_meeting = Some((neighbor, total));
}
}
}
}
}
}
if best_meeting.is_some() {
break;
}
current_depth += 1;
}
match best_meeting {
Some((mid, _total_dist)) => {
let path = reconstruct_path(&forward_visited, &backward_visited, &mid);
Ok(Some(path))
}
None => Ok(None),
}
}
fn reconstruct_path(
forward_visited: &HashMap<EntityRef, (usize, Option<EntityRef>, Option<Link>)>,
backward_visited: &HashMap<EntityRef, (usize, Option<EntityRef>, Option<Link>)>,
meeting_point: &EntityRef,
) -> Vec<PathNode> {
let mut forward_entities: Vec<EntityRef> = Vec::new();
let mut forward_links: Vec<Option<Link>> = Vec::new();
let mut current = meeting_point.clone();
while let Some((_, parent, link)) = forward_visited.get(¤t) {
forward_entities.push(current.clone());
forward_links.push(link.clone());
match parent {
Some(p) => current = p.clone(),
None => break,
}
}
forward_entities.reverse();
forward_links.reverse();
let mut backward_entities: Vec<EntityRef> = Vec::new();
let mut backward_links: Vec<Option<Link>> = Vec::new();
if let Some((_, Some(child), link)) = backward_visited.get(meeting_point) {
backward_links.push(link.clone());
current = child.clone();
while let Some((_, next_child, link)) = backward_visited.get(¤t) {
backward_entities.push(current.clone());
match next_child {
Some(nc) => {
backward_links.push(link.clone());
current = nc.clone();
}
None => break,
}
}
if backward_entities.last().map_or(true, |e| e != ¤t) {
backward_entities.push(current.clone());
}
}
let mut path: Vec<PathNode> = Vec::new();
for (i, entity) in forward_entities.iter().enumerate() {
let link = if i == 0 {
None } else {
forward_links.get(i).cloned().flatten()
};
path.push(PathNode {
entity_id: entity.clone(),
depth: i,
via_link: link,
path_weight: i as f64,
});
}
let base_depth = path.len();
for (i, entity) in backward_entities.iter().enumerate() {
let link = backward_links.get(i).cloned().flatten();
path.push(PathNode {
entity_id: entity.clone(),
depth: base_depth + i,
via_link: link,
path_weight: (base_depth + i) as f64,
});
}
path
}