use yo_common::num::DIGITS_MAX;
use crate::Elements;
use crate::set::{Limits, Needle, Set};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Plan {
Probe,
Accumulate,
}
pub fn inter<F>(sets: &[&Set], limit: usize, f: F) -> usize
where
F: FnMut(&[u8]),
{
inter_with(Plan::Probe, sets, limit, f)
}
pub fn inter_with<F>(how: Plan, sets: &[&Set], limit: usize, f: F) -> usize
where
F: FnMut(&[u8]),
{
if sets.is_empty() || sets.iter().any(|s| s.is_empty()) {
return 0;
}
match how {
Plan::Probe => inter_probe(sets, limit, f),
Plan::Accumulate => inter_accumulate(sets, limit, f),
}
}
fn inter_probe<F>(sets: &[&Set], limit: usize, mut f: F) -> usize
where
F: FnMut(&[u8]),
{
let mut order: Vec<usize> = (0..sets.len()).collect();
order.sort_unstable_by_key(|&i| sets[i].len());
let (&first, rest) = order.split_first().expect("not empty");
let mut digits = [0u8; DIGITS_MAX];
let mut found = 0usize;
for m in sets[first].iter() {
let needle = Needle::of(m, &mut digits);
if rest.iter().all(|&i| sets[i].has(&needle)) {
f(needle.bytes());
found += 1;
if limit != 0 && found == limit {
break;
}
}
}
found
}
fn inter_accumulate<F>(sets: &[&Set], limit: usize, mut f: F) -> usize
where
F: FnMut(&[u8]),
{
let mut order: Vec<usize> = (0..sets.len()).collect();
order.sort_unstable_by_key(|&i| sets[i].len());
let (&first, rest) = order.split_first().expect("not empty");
let mut digits = [0u8; DIGITS_MAX];
let mut seen = Elements::<u32>::with_capacity(sets[first].len());
for m in sets[first].iter() {
seen.insert(text(m, &mut digits), 1)
.expect("no larger than its source");
}
for &i in rest {
for m in sets[i].iter() {
if let Some(count) = seen.get_mut(text(m, &mut digits)) {
*count += 1;
}
}
}
let k = sets.len() as u32;
let mut found = 0usize;
for m in sets[first].iter() {
let name = text(m, &mut digits);
if seen.get(name) == Some(&k) {
f(name);
found += 1;
if limit != 0 && found == limit {
break;
}
}
}
found
}
#[inline]
fn text<'a>(m: crate::set::Member<'a>, digits: &'a mut [u8; DIGITS_MAX]) -> &'a [u8] {
match m {
crate::set::Member::Str(s) => s,
crate::set::Member::Int(n) => yo_common::num::i64_digits(digits, n),
}
}
pub fn union<F>(sets: &[&Set], mut f: F) -> usize
where
F: FnMut(&[u8]),
{
let biggest = sets.iter().map(|s| s.len()).max().unwrap_or(0);
let mut digits = [0u8; DIGITS_MAX];
let mut seen = Elements::<()>::with_capacity(biggest);
let mut found = 0usize;
for s in sets {
for m in s.iter() {
let name = text(m, &mut digits);
if seen.insert(name, ()).is_ok_and(|was| was.is_none()) {
f(name);
found += 1;
}
}
}
found
}
pub fn diff<F>(sets: &[&Set], mut f: F) -> usize
where
F: FnMut(&[u8]),
{
let Some((first, rest)) = sets.split_first() else {
return 0;
};
let mut order: Vec<usize> = (0..rest.len()).collect();
order.sort_unstable_by_key(|&i| rest[i].len());
let mut digits = [0u8; DIGITS_MAX];
let mut found = 0usize;
for m in first.iter() {
let needle = Needle::of(m, &mut digits);
if !order.iter().any(|&i| rest[i].has(&needle)) {
f(needle.bytes());
found += 1;
}
}
found
}
pub fn collect(
upper: usize,
limits: &Limits,
run: impl FnOnce(&mut dyn FnMut(&[u8])),
) -> Option<Set> {
let mut out: Option<Set> = None;
run(&mut |name| match &mut out {
Some(s) => {
s.add(name, limits);
}
None => {
let mut s = Set::with_hint(name, upper, limits);
s.add(name, limits);
out = Some(s);
}
});
out
}
#[cfg(test)]
mod tests {
use super::*;
use crate::set::Encoding;
fn set(members: &[&str]) -> Set {
of(members.iter().map(|m| m.as_bytes()))
}
fn of<'a>(members: impl IntoIterator<Item = &'a [u8]>) -> Set {
let mut s = Set::new();
for m in members {
s.add(m, &Limits::DEFAULT);
}
s
}
fn banded(members: &[&str], limits: &Limits) -> Set {
let mut s = Set::new();
for m in members {
s.add(m.as_bytes(), limits);
}
s
}
const AS_INTSET: Limits = Limits {
max_intset_entries: usize::MAX,
max_listpack_entries: usize::MAX,
max_listpack_value: usize::MAX,
};
const AS_LISTPACK: Limits = Limits {
max_intset_entries: 0,
max_listpack_entries: usize::MAX,
max_listpack_value: usize::MAX,
};
const AS_TABLE: Limits = Limits {
max_intset_entries: 0,
max_listpack_entries: 0,
max_listpack_value: 0,
};
fn run<F>(op: F) -> Vec<String>
where
F: FnOnce(&mut dyn FnMut(&[u8])) -> usize,
{
let mut got = Vec::new();
let n = op(&mut |m| got.push(String::from_utf8_lossy(m).into_owned()));
assert_eq!(n, got.len(), "the count and the members disagree");
got
}
#[test]
fn an_intersection_is_what_they_all_have() {
let a = set(&["a", "b", "c", "d"]);
let b = set(&["b", "c", "d", "e"]);
let c = set(&["c", "d", "e", "f"]);
let got = run(|f| inter(&[&a, &b, &c], 0, f));
assert_eq!(got, vec!["c", "d"]);
}
#[test]
fn an_intersection_of_one_set_is_that_set() {
let a = set(&["x", "y"]);
assert_eq!(run(|f| inter(&[&a], 0, f)), vec!["x", "y"]);
}
#[test]
fn an_empty_set_anywhere_empties_the_intersection() {
let a = set(&["a", "b"]);
let empty = set(&[]);
assert_eq!(run(|f| inter(&[&a, &empty], 0, f)), Vec::<String>::new());
assert_eq!(run(|f| inter(&[&empty, &a], 0, f)), Vec::<String>::new());
assert_eq!(run(|f| inter(&[], 0, f)), Vec::<String>::new());
}
#[test]
fn a_limit_stops_the_intersection_early() {
let a = set(&["a", "b", "c", "d", "e"]);
let b = set(&["a", "b", "c", "d", "e"]);
assert_eq!(run(|f| inter(&[&a, &b], 2, f)), vec!["a", "b"]);
assert_eq!(run(|f| inter(&[&a, &b], 99, f)).len(), 5);
assert_eq!(run(|f| inter(&[&a, &b], 0, f)).len(), 5, "zero is no limit");
}
#[test]
fn both_plans_give_the_same_answer_in_the_same_order() {
let sets: Vec<Set> = (0..9)
.map(|s| {
let members: Vec<String> = (0..200)
.filter(|i| i % (s + 2) != 1)
.map(|i| format!("m{i}"))
.collect();
set(&members.iter().map(String::as_str).collect::<Vec<_>>())
})
.collect();
let refs: Vec<&Set> = sets.iter().collect();
let probed = run(|f| inter_with(Plan::Probe, &refs, 0, f));
let piled = run(|f| inter_with(Plan::Accumulate, &refs, 0, f));
assert_eq!(probed, piled);
assert!(!probed.is_empty(), "the fixture should overlap");
assert_eq!(
run(|f| inter(&refs, 0, f)),
probed,
"and so does the chooser"
);
}
#[test]
fn a_union_has_everything_once() {
let a = set(&["a", "b"]);
let b = set(&["b", "c"]);
let c = set(&["c", "d"]);
assert_eq!(run(|f| union(&[&a, &b, &c], f)), vec!["a", "b", "c", "d"]);
assert_eq!(run(|f| union(&[], f)), Vec::<String>::new());
}
#[test]
fn a_difference_takes_the_others_out_of_the_first() {
let a = set(&["a", "b", "c", "d"]);
let b = set(&["b"]);
let c = set(&["d", "e"]);
assert_eq!(run(|f| diff(&[&a, &b, &c], f)), vec!["a", "c"]);
assert_eq!(run(|f| diff(&[&a], f)), vec!["a", "b", "c", "d"]);
assert_eq!(run(|f| diff(&[], f)), Vec::<String>::new());
}
#[test]
fn the_plans_agree_where_every_set_holds_everything() {
let members: Vec<String> = (0..100).map(|i| format!("m{i}")).collect();
let names: Vec<&str> = members.iter().map(String::as_str).collect();
let sets: Vec<Set> = (0..10).map(|_| set(&names)).collect();
let refs: Vec<&Set> = sets.iter().collect();
let probed = run(|f| inter_with(Plan::Probe, &refs, 0, f));
assert_eq!(probed, members, "everything is in all ten");
assert_eq!(run(|f| inter_with(Plan::Accumulate, &refs, 0, f)), probed);
assert_eq!(run(|f| inter(&refs, 0, f)), probed);
}
#[test]
fn a_store_form_builds_a_set_of_the_result() {
let a = set(&["a", "b", "c"]);
let b = set(&["b", "c", "d"]);
let out = collect(a.len().min(b.len()), &Limits::DEFAULT, |f| {
inter(&[&a, &b], 0, f);
})
.expect("two members is a set");
assert_eq!(out.len(), 2);
assert!(out.contains(b"b") && out.contains(b"c"));
assert!(!out.contains(b"a"));
}
#[test]
fn a_store_form_of_nothing_is_nothing() {
let a = set(&["a"]);
let b = set(&["b"]);
assert!(
collect(1, &Limits::DEFAULT, |f| {
inter(&[&a, &b], 0, f);
})
.is_none()
);
}
#[test]
fn a_store_form_keeps_the_representation_its_members_deserve() {
let a = set(&["1", "2", "3"]);
let b = set(&["2", "3", "4"]);
assert_eq!(a.encoding(), Encoding::Intset);
let out = collect(3, &Limits::DEFAULT, |f| {
inter(&[&a, &b], 0, f);
})
.expect("two members");
assert_eq!(out.encoding(), Encoding::Intset);
assert!(out.contains(b"2") && out.contains(b"3"));
let c = set(&["x"]);
let out = collect(4, &Limits::DEFAULT, |f| {
union(&[&a, &c], f);
})
.expect("four members");
assert_ne!(out.encoding(), Encoding::Intset);
assert!(out.contains(b"1") && out.contains(b"x"));
}
#[test]
fn the_three_representations_intersect_each_other() {
let names = ["1", "2", "3", "4"];
let others = ["3", "4", "5", "6"];
let bands = [
("intset", AS_INTSET),
("listpack", AS_LISTPACK),
("table", AS_TABLE),
];
for (ln, left) in bands {
for (rn, right) in bands {
let a = banded(&names, &left);
let b = banded(&others, &right);
let mut got = run(|f| inter(&[&a, &b], 0, f));
got.sort();
assert_eq!(got, ["3", "4"], "{ln} against {rn}");
let mut got = run(|f| union(&[&a, &b], f));
got.sort();
assert_eq!(got, ["1", "2", "3", "4", "5", "6"], "{ln} with {rn}");
let mut got = run(|f| diff(&[&a, &b], f));
got.sort();
assert_eq!(got, ["1", "2"], "{ln} without {rn}");
}
}
}
#[test]
fn a_number_and_its_untidy_spelling_stay_two_members() {
let a = banded(&["42", "042", "-0"], &AS_LISTPACK);
let b = banded(&["42"], &AS_INTSET);
assert_eq!(run(|f| inter(&[&a, &b], 0, f)), vec!["42"]);
let mut got = run(|f| diff(&[&a, &b], f));
got.sort();
assert_eq!(got, ["-0", "042"]);
let mut got = run(|f| union(&[&a, &b], f));
got.sort();
assert_eq!(
got,
["-0", "042", "42"],
"and the union does not merge them"
);
}
#[test]
fn members_that_are_not_text_work_the_same() {
let a = of([&b"\x00\xff"[..], b"\xc3\x28", b""]);
let b = of([&b"\xc3\x28"[..], b""]);
let mut got: Vec<Vec<u8>> = Vec::new();
let n = inter(&[&a, &b], 0, |m| got.push(m.to_vec()));
assert_eq!(n, 2);
assert_eq!(got, vec![b"\xc3\x28".to_vec(), b"".to_vec()]);
}
}