use crate::ir::{Arena, NodeId};
use crate::normalize::NormalizeError;
use indexmap::IndexMap;
use std::collections::HashSet;
pub fn resolve_refs(
arena: &mut Arena,
defs: &IndexMap<String, NodeId>,
) -> Result<Vec<(NodeId, NodeId)>, NormalizeError> {
let refs: Vec<(NodeId, String)> = arena
.iter()
.filter_map(|(id, node)| {
node.annotations
.r#ref
.as_ref()
.and_then(|v| v.as_str().map(|s| (id, s.to_string())))
})
.collect();
let mut edges = Vec::new();
for (node_id, ref_str) in &refs {
if let Some(target) = resolve_ref_string(ref_str, defs)? {
arena[*node_id].ref_target = Some(target);
edges.push((*node_id, target));
}
}
for (node_id, ref_str) in &refs {
if arena[*node_id].ref_target.is_none() && ref_str.starts_with('#') {
return Err(NormalizeError::UnresolvedRef(ref_str.clone()));
}
}
Ok(edges)
}
fn resolve_ref_string(
ref_str: &str,
defs: &IndexMap<String, NodeId>,
) -> Result<Option<NodeId>, NormalizeError> {
let pointer = ref_str.strip_prefix('#').unwrap_or(ref_str);
let decoded = percent_encoding::percent_decode_str(pointer)
.decode_utf8()
.map_err(|e| {
NormalizeError::ParseError(format!("invalid percent-encoding in $ref: {e}"))
})?;
let decoded = decoded.as_ref();
if let Some(name) = decoded.strip_prefix("/$defs/") {
return Ok(defs.get(name).copied());
}
if let Some(name) = decoded.strip_prefix("/definitions/") {
return Ok(defs.get(name).copied());
}
if ref_str.starts_with("http://") || ref_str.starts_with("https://") || ref_str.starts_with('/')
{
return Ok(None);
}
Ok(None)
}
pub fn transitive_ref_edges(
arena: &Arena,
defs: &IndexMap<String, NodeId>,
) -> Vec<(NodeId, NodeId)> {
let def_ids: HashSet<NodeId> = defs.values().copied().collect();
let mut edges = Vec::new();
for (node_id, node) in arena.iter() {
if let Some(target) = node.ref_target {
let source = if def_ids.contains(&node_id) {
node_id
} else if let Some(parent_id) = node.parent {
parent_id
} else {
node_id
};
edges.push((source, target));
}
}
edges
}
pub fn tarjan_scc(arena: &mut Arena, edges: &[(NodeId, NodeId)]) {
let n = arena.len();
if n == 0 {
return;
}
let mut adj: Vec<Vec<usize>> = vec![Vec::new(); n];
for (u, v) in edges {
adj[u.0 as usize].push(v.0 as usize);
}
let mut index = 0usize;
let mut stack = Vec::new();
let mut on_stack = vec![false; n];
let mut indices = vec![None; n];
let mut lowlinks = vec![0usize; n];
let mut sccs: Vec<Vec<usize>> = Vec::new();
for v in 0..n {
if indices[v].is_none() {
strongconnect(
v,
&adj,
&mut index,
&mut stack,
&mut on_stack,
&mut indices,
&mut lowlinks,
&mut sccs,
);
}
}
for component in &sccs {
let size = component.len();
let has_self_loop = size == 1 && adj[component[0]].contains(&component[0]);
if size > 1 || has_self_loop {
for &node_idx in component {
arena[NodeId(node_idx as u32)].is_cyclic = true;
}
}
}
}
#[allow(clippy::too_many_arguments)]
fn strongconnect(
v: usize,
adj: &[Vec<usize>],
index: &mut usize,
stack: &mut Vec<usize>,
on_stack: &mut [bool],
indices: &mut [Option<usize>],
lowlinks: &mut [usize],
sccs: &mut Vec<Vec<usize>>,
) {
indices[v] = Some(*index);
lowlinks[v] = *index;
*index += 1;
stack.push(v);
on_stack[v] = true;
for &w in &adj[v] {
if indices[w].is_none() {
strongconnect(w, adj, index, stack, on_stack, indices, lowlinks, sccs);
lowlinks[v] = lowlinks[v].min(lowlinks[w]);
} else if on_stack[w] {
lowlinks[v] = lowlinks[v].min(indices[w].unwrap());
}
}
if lowlinks[v] == indices[v].unwrap() {
let mut component = Vec::new();
loop {
let w = stack.pop().unwrap();
on_stack[w] = false;
component.push(w);
if w == v {
break;
}
}
sccs.push(component);
}
}