1use rudb_common::{Error, Result};
34
35use crate::link::Link;
36use crate::rid::{NO_PARENT, PART_ROWS, Rid};
37
38pub const SPARSE_RATIO: u64 = 1000;
44
45pub const STOP_AFTER: u64 = 3;
51
52#[derive(Debug, Clone, Copy, PartialEq, Eq)]
54pub enum Form {
55 Full,
57 Sparse,
59 Dense,
61}
62
63#[derive(Debug, Clone, PartialEq, Eq)]
65pub struct Rids {
66 rows: u64,
67 body: Body,
68}
69
70#[derive(Debug, Clone, PartialEq, Eq)]
71enum Body {
72 Full,
73 Sparse(Vec<Rid>),
74 Dense { words: Vec<u64>, members: u64 },
75}
76
77impl Rids {
78 #[must_use]
80 pub fn full(rows: u64) -> Self {
81 if rows == 0 { Self::none(0) } else { Self { rows, body: Body::Full } }
82 }
83
84 #[must_use]
86 pub fn none(rows: u64) -> Self {
87 Self { rows, body: Body::Sparse(Vec::new()) }
88 }
89
90 pub fn from_sorted(rows: u64, members: Vec<Rid>) -> Result<Self> {
97 let mut previous = None;
98 for &member in &members {
99 if member >= rows || previous.is_some_and(|previous| member <= previous) {
100 return Err(Error::internal(format!(
101 "row {member} is out of order or past the end of a table of {rows} rows"
102 )));
103 }
104 previous = Some(member);
105 }
106 Ok(Self::settle_sparse(rows, members))
107 }
108
109 pub fn from_words(rows: u64, words: Vec<u64>) -> Result<Self> {
115 if count(words.len()) != rows.div_ceil(64) {
116 return Err(Error::internal(format!(
117 "{} words is not a bitmap over {rows} rows",
118 words.len()
119 )));
120 }
121 let tail = rows % 64;
122 if tail != 0 && words.last().is_some_and(|last| last >> tail != 0) {
123 return Err(Error::internal("a bitmap has rows set past the end of its table"));
124 }
125 Ok(Self::settle_dense(rows, words))
126 }
127
128 #[must_use]
130 pub fn rows(&self) -> u64 {
131 self.rows
132 }
133
134 #[must_use]
136 pub fn len(&self) -> u64 {
137 match &self.body {
138 Body::Full => self.rows,
139 Body::Sparse(members) => count(members.len()),
140 Body::Dense { members, .. } => *members,
141 }
142 }
143
144 #[must_use]
146 pub fn is_empty(&self) -> bool {
147 self.len() == 0
148 }
149
150 #[must_use]
152 pub fn is_full(&self) -> bool {
153 matches!(self.body, Body::Full)
154 }
155
156 #[must_use]
158 pub fn form(&self) -> Form {
159 match self.body {
160 Body::Full => Form::Full,
161 Body::Sparse(_) => Form::Sparse,
162 Body::Dense { .. } => Form::Dense,
163 }
164 }
165
166 #[must_use]
168 pub fn bytes(&self) -> usize {
169 match &self.body {
170 Body::Full => 0,
171 Body::Sparse(members) => members.len() * size_of::<Rid>(),
172 Body::Dense { words, .. } => words.len() * size_of::<u64>(),
173 }
174 }
175
176 #[must_use]
178 pub fn contains(&self, rid: Rid) -> bool {
179 if rid >= self.rows {
180 return false;
181 }
182 match &self.body {
183 Body::Full => true,
184 Body::Sparse(members) => members.binary_search(&rid).is_ok(),
185 Body::Dense { words, .. } => bit(words, rid),
186 }
187 }
188
189 #[must_use]
194 pub fn any_between(&self, low: Rid, high: Rid) -> bool {
195 let high = high.min(self.rows.saturating_sub(1));
196 if low > high || self.rows == 0 {
197 return false;
198 }
199 match &self.body {
200 Body::Full => true,
201 Body::Sparse(members) => {
202 let from = members.partition_point(|&member| member < low);
203 members.get(from).is_some_and(|&member| member <= high)
204 }
205 Body::Dense { words, .. } => {
206 let (first, last) = (index(low / 64), index(high / 64));
207 (first..=last).any(|at| {
208 let mut word = words[at];
209 if at == first {
210 word &= u64::MAX << (low % 64);
211 }
212 if at == last {
213 word &= u64::MAX >> (63 - high % 64);
214 }
215 word != 0
216 })
217 }
218 }
219 }
220
221 pub fn iter(&self) -> impl Iterator<Item = Rid> + '_ {
223 let (full, sparse, dense) = match &self.body {
224 Body::Full => (Some(0..self.rows), None, None),
225 Body::Sparse(members) => (None, Some(members.iter().copied()), None),
226 Body::Dense { words, .. } => (None, None, Some(ones(words))),
227 };
228 full.into_iter()
229 .flatten()
230 .chain(sparse.into_iter().flatten())
231 .chain(dense.into_iter().flatten())
232 }
233
234 pub fn intersect(&self, other: &Self) -> Result<Self> {
240 self.same_table(other)?;
241 Ok(match (&self.body, &other.body) {
242 (Body::Full, _) => other.clone(),
243 (_, Body::Full) => self.clone(),
244 (Body::Dense { words: left, .. }, Body::Dense { words: right, .. }) => {
245 let words = left.iter().zip(right).map(|(left, right)| left & right).collect();
246 Self::settle_dense(self.rows, words)
247 }
248 (Body::Sparse(members), _) => Self::settle_sparse(
251 self.rows,
252 members.iter().copied().filter(|&member| other.contains(member)).collect(),
253 ),
254 (_, Body::Sparse(members)) => Self::settle_sparse(
255 self.rows,
256 members.iter().copied().filter(|&member| self.contains(member)).collect(),
257 ),
258 })
259 }
260
261 pub fn union(&self, other: &Self) -> Result<Self> {
267 self.same_table(other)?;
268 if self.is_full() || other.is_full() {
269 return Ok(Self::full(self.rows));
270 }
271 let mut words = self.words();
272 for member in other.iter() {
273 words[index(member / 64)] |= 1 << (member % 64);
274 }
275 Ok(Self::settle_dense(self.rows, words))
276 }
277
278 pub fn forward(&self, link: &Link) -> Result<Pushed> {
288 self.push(link, false)
289 }
290
291 pub fn forward_or_stop(&self, link: &Link) -> Result<Pushed> {
304 self.push(link, true)
305 }
306
307 fn push(&self, link: &Link, stopping: bool) -> Result<Pushed> {
308 if self.rows != link.parents() {
309 return Err(Error::internal(format!(
310 "a set over {} rows pushed through a link whose parent has {}",
311 self.rows,
312 link.parents()
313 )));
314 }
315 let children = link.children();
316 let parts = children.div_ceil(count(PART_ROWS));
317 if self.is_full() && link.linked() == children {
319 return Ok(Pushed { rids: Self::full(children), parts, skipped: 0, stopped: false });
320 }
321 let mut words = vec![0_u64; index(children.div_ceil(64))];
322 let mut parents = vec![NO_PARENT; PART_ROWS];
323 let mut skipped = 0_u64;
324 let mark = children.div_ceil(STOP_AFTER);
328 let mut asked = !stopping;
329 let mut kept = 0_u64;
330 for part in 0..parts {
331 let first = part * count(PART_ROWS);
332 if !asked && first >= mark {
333 asked = true;
334 if kept == first {
335 return Ok(Pushed {
336 rids: Self::full(children),
337 parts,
338 skipped,
339 stopped: true,
340 });
341 }
342 }
343 let reach = match link.part_bounds(index(part)) {
344 Some(Some((low, high))) => self.any_between(low, high),
345 _ => false,
346 };
347 if !reach {
348 skipped += 1;
349 continue;
350 }
351 let run = index((children - first).min(count(PART_ROWS)));
352 link.forward_run(first, &mut parents[..run])?;
353 for (at, &parent) in parents[..run].iter().enumerate() {
354 if parent != NO_PARENT && self.contains(parent) {
355 let child = first + count(at);
356 words[index(child / 64)] |= 1 << (child % 64);
357 kept += 1;
358 }
359 }
360 }
361 Ok(Pushed { rids: Self::settle_dense(children, words), parts, skipped, stopped: false })
362 }
363
364 pub fn backward(&self, link: &Link) -> Result<Self> {
374 if self.rows != link.children() {
375 return Err(Error::internal(format!(
376 "a set over {} rows pushed back through a link whose child has {}",
377 self.rows,
378 link.children()
379 )));
380 }
381 let mut words = vec![0_u64; index(link.parents().div_ceil(64))];
382 let mut parents = vec![NO_PARENT; PART_ROWS];
383 let children = link.children();
384 for part in 0..children.div_ceil(count(PART_ROWS)) {
385 let first = part * count(PART_ROWS);
386 let last = (first + count(PART_ROWS)).min(children) - 1;
387 if !self.any_between(first, last) {
388 continue;
389 }
390 let run = index(last - first + 1);
391 link.forward_run(first, &mut parents[..run])?;
392 for (at, &parent) in parents[..run].iter().enumerate() {
393 if parent != NO_PARENT && self.contains(first + count(at)) {
394 words[index(parent / 64)] |= 1 << (parent % 64);
395 }
396 }
397 }
398 Ok(Self::settle_dense(link.parents(), words))
399 }
400
401 fn same_table(&self, other: &Self) -> Result<()> {
402 if self.rows == other.rows {
403 Ok(())
404 } else {
405 Err(Error::internal(format!(
406 "a set over {} rows combined with one over {}",
407 self.rows, other.rows
408 )))
409 }
410 }
411
412 fn words(&self) -> Vec<u64> {
414 let mut words = vec![0_u64; index(self.rows.div_ceil(64))];
415 match &self.body {
416 Body::Dense { words: held, .. } => words.copy_from_slice(held),
417 _ => {
418 for member in self.iter() {
419 words[index(member / 64)] |= 1 << (member % 64);
420 }
421 }
422 }
423 words
424 }
425
426 fn settle_dense(rows: u64, words: Vec<u64>) -> Self {
428 let members = words.iter().map(|word| u64::from(word.count_ones())).sum::<u64>();
429 match shape(rows, members) {
430 Form::Full => Self::full(rows),
431 Form::Sparse => Self { rows, body: Body::Sparse(ones(&words).collect()) },
432 Form::Dense => Self { rows, body: Body::Dense { words, members } },
433 }
434 }
435
436 fn settle_sparse(rows: u64, members: Vec<Rid>) -> Self {
438 match shape(rows, count(members.len())) {
439 Form::Full => Self::full(rows),
440 Form::Sparse => Self { rows, body: Body::Sparse(members) },
441 Form::Dense => {
442 let mut words = vec![0_u64; index(rows.div_ceil(64))];
443 for member in &members {
444 words[index(member / 64)] |= 1 << (member % 64);
445 }
446 Self { rows, body: Body::Dense { words, members: count(members.len()) } }
447 }
448 }
449 }
450}
451
452#[derive(Debug, Clone, PartialEq, Eq)]
454pub struct Pushed {
455 pub rids: Rids,
457 pub parts: u64,
459 pub skipped: u64,
461 pub stopped: bool,
464}
465
466fn shape(rows: u64, members: u64) -> Form {
468 if rows > 0 && members == rows {
469 Form::Full
470 } else if members == 0 || members.saturating_mul(SPARSE_RATIO) < rows {
471 Form::Sparse
472 } else {
473 Form::Dense
474 }
475}
476
477fn bit(words: &[u64], at: u64) -> bool {
479 words.get(index(at / 64)).is_some_and(|word| word >> (at % 64) & 1 == 1)
480}
481
482fn ones(words: &[u64]) -> impl Iterator<Item = Rid> + '_ {
484 words.iter().enumerate().flat_map(|(at, &word)| {
485 let base = count(at) * 64;
486 let mut rest = word;
487 std::iter::from_fn(move || {
488 if rest == 0 {
489 return None;
490 }
491 let low = u64::from(rest.trailing_zeros());
492 rest &= rest - 1;
493 Some(base + low)
494 })
495 })
496}
497
498fn count(rows: usize) -> u64 {
500 u64::try_from(rows).unwrap_or(u64::MAX)
501}
502
503fn index(rows: u64) -> usize {
505 usize::try_from(rows).unwrap_or(usize::MAX)
506}
507
508#[cfg(test)]
509mod tests {
510 use super::{Form, Rids, SPARSE_RATIO, STOP_AFTER};
511 use crate::link::Link;
512 use crate::rid::{NO_PARENT, PART_ROWS, Rid};
513
514 fn members(rids: &Rids) -> Vec<Rid> {
516 (0..rids.rows()).filter(|&rid| rids.contains(rid)).collect()
517 }
518
519 #[test]
520 fn the_form_follows_the_count_and_not_the_constructor() {
521 let rows = 10 * SPARSE_RATIO;
522 assert_eq!(Rids::from_sorted(rows, vec![1, 2, 3]).expect("sorted").form(), Form::Sparse);
523 let many: Vec<Rid> = (0..rows).step_by(2).collect();
524 assert_eq!(Rids::from_sorted(rows, many).expect("sorted").form(), Form::Dense);
525 let every: Vec<Rid> = (0..rows).collect();
526 assert_eq!(Rids::from_sorted(rows, every).expect("sorted").form(), Form::Full);
527 let mut words = vec![0_u64; usize::try_from(rows.div_ceil(64)).expect("small")];
528 words[0] = 1;
529 assert_eq!(Rids::from_words(rows, words).expect("bitmap").form(), Form::Sparse);
530 }
531
532 #[test]
535 fn the_same_members_are_the_same_set_whichever_way_they_came_in() {
536 let rows = 5000;
537 let members: Vec<Rid> = (0..rows).filter(|rid| rid % 3 == 0).collect();
538 let mut words = vec![0_u64; usize::try_from(rows.div_ceil(64)).expect("small")];
539 for member in &members {
540 words[usize::try_from(member / 64).expect("small")] |= 1 << (member % 64);
541 }
542 let listed = Rids::from_sorted(rows, members).expect("sorted");
543 let mapped = Rids::from_words(rows, words).expect("bitmap");
544 assert_eq!(listed, mapped);
545 }
546
547 #[test]
548 fn a_member_out_of_order_or_past_the_end_is_refused() {
549 assert!(Rids::from_sorted(10, vec![3, 2]).is_err());
550 assert!(Rids::from_sorted(10, vec![3, 3]).is_err());
551 assert!(Rids::from_sorted(10, vec![10]).is_err());
552 assert!(Rids::from_words(10, vec![1 << 10]).is_err(), "a bit past the end");
553 assert!(Rids::from_words(10, vec![0, 0]).is_err(), "a word too many");
554 }
555
556 #[test]
557 fn a_full_set_holds_nothing_and_answers_everything() {
558 let full = Rids::full(1_000_000);
559 assert_eq!(full.bytes(), 0);
560 assert_eq!(full.len(), 1_000_000);
561 assert!(full.contains(999_999));
562 assert!(!full.contains(1_000_000));
563 }
564
565 #[test]
566 fn any_between_looks_only_inside_the_range_in_every_form() {
567 let rows = 4096;
568 for members in [vec![700], (0..rows).filter(|rid| rid % 2 == 0 && *rid != 700).collect()] {
569 let rids = Rids::from_sorted(rows, members.clone()).expect("sorted");
570 for (low, high) in [(0, 63), (64, 699), (699, 701), (700, 700), (1000, 5000)] {
571 let expected = members.iter().any(|&member| (low..=high).contains(&member));
572 assert_eq!(
573 rids.any_between(low, high),
574 expected,
575 "{low}..={high} over {:?}",
576 rids.form()
577 );
578 }
579 }
580 assert!(Rids::full(10).any_between(3, 3));
581 assert!(!Rids::full(10).any_between(10, 20), "past the end is outside the table");
582 }
583
584 #[test]
585 fn intersect_and_union_agree_with_the_slow_answer_across_forms() {
586 let rows = 20_000;
587 let sets = [
588 Rids::none(rows),
589 Rids::from_sorted(rows, vec![5, 700, 19_999]).expect("sorted"),
590 Rids::from_sorted(rows, (0..rows).filter(|rid| rid % 3 == 0).collect())
591 .expect("sorted"),
592 Rids::from_sorted(rows, (0..rows).filter(|rid| rid % 5 == 0).collect())
593 .expect("sorted"),
594 Rids::full(rows),
595 ];
596 for left in &sets {
597 for right in &sets {
598 let both = left.intersect(right).expect("same table");
599 let either = left.union(right).expect("same table");
600 let (left_members, right_members) = (members(left), members(right));
601 let expected_both: Vec<Rid> = left_members
602 .iter()
603 .copied()
604 .filter(|rid| right_members.contains(rid))
605 .collect();
606 let mut expected_either = left_members.clone();
607 expected_either.extend(right_members.iter().copied());
608 expected_either.sort_unstable();
609 expected_either.dedup();
610 assert_eq!(members(&both), expected_both);
611 assert_eq!(members(&either), expected_either);
612 assert_eq!(both.iter().collect::<Vec<_>>(), expected_both, "iteration is in order");
613 }
614 }
615 assert!(Rids::full(3).intersect(&Rids::full(4)).is_err(), "two different tables");
616 }
617
618 fn link(children: u64, parents: u64, parent_of: impl Fn(u64) -> Rid) -> Link {
620 let of: Vec<Rid> = (0..children).map(parent_of).collect();
621 Link::build(&of, parents).expect("a link")
622 }
623
624 #[test]
627 fn a_forward_push_finds_exactly_the_children_that_point_into_the_set() {
628 let parents = 3000;
629 let children = 10 * count(PART_ROWS) + 17;
630 let clustered = link(children, parents, |child| child * parents / children);
631 let scattered = link(children, parents, |child| {
632 if child / count(PART_ROWS) == 4 { NO_PARENT } else { (child * 7919) % parents }
633 });
634 for link in [&clustered, &scattered] {
635 for set in [
636 Rids::none(parents),
637 Rids::from_sorted(parents, vec![0, 1500, 2999]).expect("sorted"),
638 Rids::from_sorted(parents, (0..parents).filter(|p| p % 4 == 1).collect())
639 .expect("sorted"),
640 Rids::full(parents),
641 ] {
642 let pushed = set.forward(link).expect("the same table");
643 let expected: Vec<Rid> = (0..children)
644 .filter(|&child| link.forward(child).is_some_and(|parent| set.contains(parent)))
645 .collect();
646 assert_eq!(
647 members(&pushed.rids),
648 expected,
649 "{:?} through {:?}",
650 set.form(),
651 link.form()
652 );
653 }
654 }
655 }
656
657 fn count(rows: usize) -> u64 {
658 u64::try_from(rows).expect("small")
659 }
660
661 #[test]
664 fn a_push_that_removes_nothing_by_the_third_stops_and_one_that_removes_something_finishes() {
665 let parents = 3000;
666 let children = 4 * STOP_AFTER * count(PART_ROWS);
667 let clustered = link(children, parents, |child| child * parents / children);
668 let every = Rids::full(parents);
669 let all_but_last: Vec<Rid> = (0..parents - 1).collect();
670 let most = Rids::from_sorted(parents, all_but_last).expect("sorted");
671 let stopped = most.forward_or_stop(&clustered).expect("the same table");
672 assert!(stopped.stopped, "nothing was removed in the first third");
673 assert!(stopped.rids.is_full(), "a stopped push keeps every row");
674 assert_eq!(stopped.parts, 4 * STOP_AFTER);
675 assert!(!every.forward_or_stop(&clustered).expect("the same table").stopped);
677
678 let all_but_first: Vec<Rid> = (1..parents).collect();
679 let early = Rids::from_sorted(parents, all_but_first).expect("sorted");
680 let finished = early.forward_or_stop(&clustered).expect("the same table");
681 assert!(!finished.stopped, "the first parent's children were removed before the third");
682 assert_eq!(finished, early.forward(&clustered).expect("the same table"));
683 assert_eq!(
684 finished.rids.len(),
685 children - count((0..children).filter(|child| child * parents / children == 0).count())
686 );
687
688 let orphans =
690 link(children, parents, |child| if child == 5 { NO_PARENT } else { child % parents });
691 assert!(!every.forward_or_stop(&orphans).expect("the same table").stopped);
692 }
693
694 #[test]
697 fn a_clustered_child_skips_every_part_that_points_outside_the_set() {
698 let parents = 1000;
699 let children = 100 * count(PART_ROWS);
700 let link = link(children, parents, |child| child * parents / children);
701 let set = Rids::from_sorted(parents, (100..200).collect()).expect("sorted");
702 let pushed = set.forward(&link).expect("the same table");
703 assert_eq!(pushed.parts, 100);
704 assert!(pushed.skipped >= 88, "only {} of 100 parts were skipped", pushed.skipped);
706 assert_eq!(pushed.rids.len(), children / 10);
707 }
708
709 #[test]
710 fn nothing_in_the_set_skips_every_part_and_everything_skips_the_pass() {
711 let link = link(5000, 100, |child| child % 100);
712 let pushed = Rids::none(100).forward(&link).expect("the same table");
713 assert_eq!((pushed.skipped, pushed.rids.len()), (pushed.parts, 0));
714 let pushed = Rids::full(100).forward(&link).expect("the same table");
715 assert!(pushed.rids.is_full(), "every child matched, so every child is in");
716 assert_eq!(pushed.skipped, 0);
717 }
718
719 #[test]
720 fn a_backward_push_finds_exactly_the_parents_the_set_points_at() {
721 let parents = 500;
722 let children = 7 * count(PART_ROWS) + 3;
723 let clustered = link(children, parents, |child| child * parents / children);
724 let scattered = link(children, parents, |child| {
725 if child % 11 == 0 { NO_PARENT } else { (child * 31) % parents }
726 });
727 for link in [&clustered, &scattered] {
728 let set = Rids::from_sorted(children, (0..children).filter(|c| c % 97 == 3).collect())
729 .expect("sorted");
730 let pushed = set.backward(link).expect("the same table");
731 let mut expected: Vec<Rid> =
732 set.iter().filter_map(|child| link.forward(child)).collect();
733 expected.sort_unstable();
734 expected.dedup();
735 assert_eq!(members(&pushed), expected, "through {:?}", link.form());
736 }
737 }
738
739 #[test]
740 fn a_set_over_the_wrong_table_is_refused_rather_than_pushed() {
741 let link = link(100, 10, |child| child % 10);
742 assert!(Rids::full(11).forward(&link).is_err());
743 assert!(Rids::full(10).backward(&link).is_err());
744 }
745}