1use roaring::RoaringBitmap;
11
12use crate::csr::{Csr, CsrError};
13use crate::ordinal::{Ordinal, to_usize};
14
15#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
17pub enum ClosureError {
18 #[error("the is-a hierarchy has a cycle through {} node(s), for example {first}", .members.len())]
20 Cycle {
21 first: Ordinal,
23 members: Vec<Ordinal>,
25 },
26 #[error(transparent)]
28 Csr(#[from] CsrError),
29}
30
31#[derive(Debug, Clone, PartialEq, Eq)]
33pub struct Closure {
34 ancestors: Vec<RoaringBitmap>,
35 descendants: Vec<RoaringBitmap>,
36}
37
38static EMPTY: std::sync::LazyLock<RoaringBitmap> = std::sync::LazyLock::new(RoaringBitmap::new);
39
40impl Closure {
41 pub fn compute(is_a: &Csr) -> Result<Self, ClosureError> {
47 let parents = is_a;
48 let children = is_a.transpose()?;
49 let order = topological_order(parents, &children)?;
50 let nodes = to_usize(parents.nodes());
51 Ok(Self {
52 ancestors: sweep(parents, order.iter(), nodes),
53 descendants: sweep(&children, order.iter().rev(), nodes),
54 })
55 }
56
57 #[must_use]
59 pub fn from_parts(ancestors: Vec<RoaringBitmap>, descendants: Vec<RoaringBitmap>) -> Self {
60 Self {
61 ancestors,
62 descendants,
63 }
64 }
65
66 #[must_use]
68 pub fn nodes(&self) -> u32 {
69 u32::try_from(self.ancestors.len()).unwrap_or(u32::MAX)
70 }
71
72 #[must_use]
74 pub fn ancestors(&self, node: Ordinal) -> &RoaringBitmap {
75 self.ancestors
76 .get(node.as_usize())
77 .unwrap_or_else(|| &*EMPTY)
78 }
79
80 #[must_use]
82 pub fn descendants(&self, node: Ordinal) -> &RoaringBitmap {
83 self.descendants
84 .get(node.as_usize())
85 .unwrap_or_else(|| &*EMPTY)
86 }
87
88 #[must_use]
90 pub fn descendants_or_self(&self, node: Ordinal) -> RoaringBitmap {
91 let mut set = self.descendants(node).clone();
92 set.insert(node.index());
93 set
94 }
95
96 #[must_use]
98 pub fn ancestors_or_self(&self, node: Ordinal) -> RoaringBitmap {
99 let mut set = self.ancestors(node).clone();
100 set.insert(node.index());
101 set
102 }
103
104 #[must_use]
106 pub fn is_ancestor(&self, ancestor: Ordinal, node: Ordinal) -> bool {
107 self.ancestors(node).contains(ancestor.index())
108 }
109
110 #[must_use]
112 pub fn ancestor_sets(&self) -> &[RoaringBitmap] {
113 &self.ancestors
114 }
115
116 #[must_use]
118 pub fn descendant_sets(&self) -> &[RoaringBitmap] {
119 &self.descendants
120 }
121}
122
123fn sweep<'a>(
129 edges: &Csr,
130 order: impl Iterator<Item = &'a Ordinal>,
131 nodes: usize,
132) -> Vec<RoaringBitmap> {
133 let mut sets: Vec<RoaringBitmap> = vec![RoaringBitmap::new(); nodes];
134 for node in order {
135 let mut set = RoaringBitmap::new();
136 for neighbour in edges.neighbours(*node) {
137 set.insert(*neighbour);
138 if let Some(reached) = sets.get(to_usize(*neighbour)) {
139 set |= reached;
140 }
141 }
142 if let Some(slot) = sets.get_mut(node.as_usize()) {
143 *slot = set;
144 }
145 }
146 sets
147}
148
149fn topological_order(parents: &Csr, children: &Csr) -> Result<Vec<Ordinal>, ClosureError> {
151 let nodes = parents.nodes();
152 let mut remaining: Vec<u32> = (0..nodes)
153 .map(|n| u32::try_from(parents.neighbours(Ordinal::new(n)).len()).unwrap_or(u32::MAX))
154 .collect();
155 let mut ready: Vec<Ordinal> = (0..nodes)
156 .filter(|n| remaining.get(to_usize(*n)) == Some(&0))
157 .map(Ordinal::new)
158 .collect();
159 let mut order = Vec::with_capacity(to_usize(nodes));
160 while let Some(node) = ready.pop() {
161 order.push(node);
162 for child in children.neighbours(node) {
163 if let Some(count) = remaining.get_mut(to_usize(*child)) {
164 *count = count.saturating_sub(1);
165 if *count == 0 {
166 ready.push(Ordinal::new(*child));
167 }
168 }
169 }
170 }
171 if order.len() != to_usize(nodes) {
172 let members: Vec<Ordinal> = (0..nodes)
173 .filter(|n| remaining.get(to_usize(*n)).is_some_and(|c| *c > 0))
174 .map(Ordinal::new)
175 .collect();
176 let first = members.first().copied().unwrap_or(Ordinal::new(0));
177 return Err(ClosureError::Cycle { first, members });
178 }
179 Ok(order)
180}
181
182#[cfg(test)]
183mod tests {
184 use super::{Closure, ClosureError};
185 use crate::csr::Csr;
186 use crate::ordinal::Ordinal;
187
188 fn o(i: u32) -> Ordinal {
189 Ordinal::new(i)
190 }
191
192 fn diamond() -> Closure {
194 let is_a = Csr::build(4, [(o(1), o(0)), (o(2), o(0)), (o(3), o(1)), (o(3), o(2))])
195 .expect("builds");
196 Closure::compute(&is_a).expect("acyclic")
197 }
198
199 #[test]
200 fn ancestors_and_descendants_are_transitive_and_inverse() {
201 let closure = diamond();
202 assert_eq!(
203 closure.ancestors(o(3)).iter().collect::<Vec<_>>(),
204 vec![0, 1, 2]
205 );
206 assert_eq!(
207 closure.descendants(o(0)).iter().collect::<Vec<_>>(),
208 vec![1, 2, 3]
209 );
210 assert_eq!(
211 closure.descendants(o(1)).iter().collect::<Vec<_>>(),
212 vec![3]
213 );
214 assert!(closure.ancestors(o(0)).is_empty());
215 assert!(closure.descendants(o(3)).is_empty());
216 assert!(closure.is_ancestor(o(0), o(3)));
217 assert!(!closure.is_ancestor(o(3), o(0)));
218 assert!(closure.descendants_or_self(o(3)).contains(3));
219 assert_eq!(closure.ancestors_or_self(o(3)).len(), 4);
220 }
221
222 #[test]
223 fn a_cycle_is_refused() {
224 let is_a = Csr::build(3, [(o(0), o(1)), (o(1), o(2)), (o(2), o(0))]).expect("builds");
225 match Closure::compute(&is_a) {
226 Err(ClosureError::Cycle { members, .. }) => assert_eq!(members.len(), 3),
227 other => panic!("expected a cycle, got {other:?}"),
228 }
229 }
230
231 #[test]
232 fn unknown_nodes_have_empty_sets() {
233 let closure = diamond();
234 assert!(closure.ancestors(o(99)).is_empty());
235 assert!(closure.descendants(o(99)).is_empty());
236 }
237}