use crate::canonical::ir::{drop_subsumed, BoundNumber, NumberLeaf, Side};
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct NumberLeaves {
leaves: Vec<NumberLeaf>,
canonical: bool,
}
impl Default for NumberLeaves {
fn default() -> Self {
Self {
leaves: Vec::new(),
canonical: true,
}
}
}
impl NumberLeaves {
pub(crate) fn insert(&mut self, leaf: NumberLeaf) {
self.leaves.push(leaf);
self.canonical = false;
}
fn canonicalize(&mut self) {
if self.canonical {
return;
}
let was_empty = self.leaves.is_empty();
self.leaves = merge(std::mem::take(&mut self.leaves));
drop_subsumed(&mut self.leaves, |outer, inner| {
covers(outer, inner) && outer.multiple_of.divide_all(&inner.multiple_of)
});
self.canonical = true;
debug_assert_eq!(
self.leaves.is_empty(),
was_empty,
"merging emptied the leaves"
);
}
pub(crate) fn clear(&mut self) {
self.leaves.clear();
self.canonical = true;
}
pub(crate) fn retain(&mut self, keep: impl FnMut(&NumberLeaf) -> bool) {
self.canonicalize();
self.leaves.retain(keep);
}
pub(crate) fn is_empty(&self) -> bool {
self.leaves.is_empty()
}
pub(crate) fn as_slice(&mut self) -> &[NumberLeaf] {
self.canonicalize();
&self.leaves
}
}
impl IntoIterator for NumberLeaves {
type Item = NumberLeaf;
type IntoIter = std::vec::IntoIter<NumberLeaf>;
fn into_iter(mut self) -> Self::IntoIter {
self.canonicalize();
self.leaves.into_iter()
}
}
fn merge(mut leaves: Vec<NumberLeaf>) -> Vec<NumberLeaf> {
if leaves.len() < 2 {
return leaves;
}
leaves.sort_by(|left, right| {
left.multiple_of.cmp(&right.multiple_of).then_with(|| {
match (&left.minimum, &right.minimum) {
(Some(left), Some(right)) if left.to_number() == right.to_number() => {
right.is_inclusive().cmp(&left.is_inclusive())
}
(left, right) => left.cmp(right),
}
})
});
let mut merged: Vec<NumberLeaf> = Vec::with_capacity(leaves.len());
for leaf in leaves {
match merged.last_mut() {
Some(last) if last.multiple_of == leaf.multiple_of && reaches(last, &leaf) => {
*last = hull(std::mem::take(last), leaf);
}
_ => merged.push(leaf),
}
}
merged
}
fn covers(outer: &NumberLeaf, inner: &NumberLeaf) -> bool {
let wider = |outer: &BoundNumber, inner: &BoundNumber, side| {
!outer.is_tighter_than(inner, side) || inner.is_tighter_than(outer, side)
};
let minimum = match (&outer.minimum, &inner.minimum) {
(None, _) => true,
(Some(_), None) => false,
(Some(outer), Some(inner)) => wider(outer, inner, Side::Lower),
};
let maximum = match (&outer.maximum, &inner.maximum) {
(None, _) => true,
(Some(_), None) => false,
(Some(outer), Some(inner)) => wider(outer, inner, Side::Upper),
};
minimum && maximum
}
fn reaches(last: &NumberLeaf, next: &NumberLeaf) -> bool {
let (Some(end), Some(start)) = (&last.maximum, &next.minimum) else {
return true;
};
if end.to_number() == start.to_number() {
return end.is_inclusive() || start.is_inclusive();
}
end.admits(&start.to_number(), Side::Upper)
}
fn hull(last: NumberLeaf, next: NumberLeaf) -> NumberLeaf {
debug_assert!(
!(last.minimum.is_some() && next.minimum.is_none()),
"an unbounded minimum sorted after a bounded one"
);
debug_assert!(
!matches!((&last.minimum, &next.minimum), (Some(left), Some(right))
if left.is_tighter_than(right, Side::Lower)
&& !right.is_tighter_than(left, Side::Lower)),
"the tighter minimum sorted first"
);
debug_assert_eq!(
last.multiple_of, next.multiple_of,
"folding intervals under different divisors"
);
NumberLeaf {
multiple_of: last.multiple_of,
minimum: last.minimum,
maximum: match (last.maximum, next.maximum) {
(Some(left), Some(right)) => Some(if left.is_tighter_than(&right, Side::Upper) {
right
} else {
left
}),
_ => None,
},
}
}