use crate::tree::Node;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum PoolError {
Empty,
TooMany,
}
#[derive(Debug, Default)]
pub(crate) struct CategoryPool {
ids: Vec<u32>,
}
impl CategoryPool {
pub(crate) fn with_capacity(capacity: usize) -> Self {
CategoryPool {
ids: Vec::with_capacity(capacity),
}
}
pub(crate) fn push_split(
&mut self,
node: &mut Node,
ids: impl IntoIterator<Item = u32>,
) -> Result<(), PoolError> {
let begin = self.ids.len();
self.ids.extend(ids);
let end = self.ids.len();
let range = u32::try_from(begin).ok().zip(u32::try_from(end).ok());
let error = match range {
_ if end == begin => PoolError::Empty,
Some((begin, end)) => {
node.cat_begin = begin;
node.cat_end = end;
return Ok(());
}
None => PoolError::TooMany,
};
self.ids.truncate(begin);
Err(error)
}
pub(crate) fn finish(self) -> Vec<u32> {
self.ids
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn splits_get_consecutive_ranges_and_failures_leave_no_trace() {
let mut pool = CategoryPool::default();
let mut a = Node::leaf(0.0, 0.0);
let mut b = Node::leaf(0.0, 0.0);
pool.push_split(&mut a, [3, 1]).unwrap();
assert_eq!(pool.push_split(&mut b, []), Err(PoolError::Empty));
assert_eq!((b.cat_begin, b.cat_end), (0, 0));
pool.push_split(&mut b, [7]).unwrap();
assert_eq!(
(a.cat_begin, a.cat_end, b.cat_begin, b.cat_end),
(0, 2, 2, 3)
);
assert_eq!(pool.finish(), [3, 1, 7]);
}
}