use std::collections::HashSet;
use rudb_common::{LogicalType, Result, Value};
use rudb_vector::{Data, Form, Validity, Vector};
use crate::shape::{first, identity, nulls_of, single};
#[derive(Debug)]
pub struct Members {
held: Held,
has_null: bool,
negated: bool,
}
#[derive(Debug)]
enum Held {
Whole(HashSet<i128>),
Text(HashSet<String>),
}
impl Members {
#[must_use]
pub fn of(values: &[Value], negated: bool) -> Option<Self> {
if values.len() < 2 {
return None;
}
let mut whole: HashSet<i128> = HashSet::new();
let mut text: HashSet<String> = HashSet::new();
let mut has_null = false;
let mut kind: Option<std::mem::Discriminant<Value>> = None;
for value in values {
if matches!(value, Value::Null) {
has_null = true;
continue;
}
let held = std::mem::discriminant(value);
if *kind.get_or_insert(held) != held {
return None;
}
match value {
Value::Varchar(held) => {
text.insert(held.clone());
}
other => {
whole.insert(number(other)?);
}
}
}
let held = if text.is_empty() {
if whole.is_empty() {
return None;
}
Held::Whole(whole)
} else {
Held::Text(text)
};
Some(Self { held, has_null, negated })
}
#[must_use]
pub fn len(&self) -> usize {
match &self.held {
Held::Whole(set) => set.len(),
Held::Text(set) => set.len(),
}
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
fn number(value: &Value) -> Option<i128> {
match *value {
Value::TinyInt(held) => Some(i128::from(held)),
Value::SmallInt(held) => Some(i128::from(held)),
Value::Integer(held) | Value::Date(held) => Some(i128::from(held)),
Value::BigInt(held) | Value::Time(held) | Value::Timestamp(held) => Some(i128::from(held)),
Value::HugeInt(held) => Some(held),
Value::UTinyInt(held) => Some(i128::from(held)),
Value::USmallInt(held) => Some(i128::from(held)),
Value::UInteger(held) => Some(i128::from(held)),
Value::UBigInt(held) => Some(i128::from(held)),
_ => None,
}
}
pub fn in_set(input: &Vector, members: &Members, returns: &LogicalType) -> Result<Vector> {
let rows = input.len();
let base = nulls_of(input);
match input.form() {
Form::Flat => match input.data() {
Some(data) => look(data, identity, members, &base, rows, returns),
None => row_at_a_time(input, members, &base, rows, returns),
},
Form::Dictionary | Form::Rle => {
let Some((codes, values)) = input.positions() else {
return row_at_a_time(input, members, &base, rows, returns);
};
let Some(data) = values.data().filter(|_| codes.len() >= rows) else {
return row_at_a_time(input, members, &base, rows, returns);
};
let at = move |index: usize| codes[index] as usize;
look(data, at, members, &base, rows, returns)
}
Form::Constant => {
let Some(value) = input.constant_value() else {
return row_at_a_time(input, members, &base, rows, returns);
};
let Some(held) = single(input.logical_type(), value) else {
return row_at_a_time(input, members, &base, rows, returns);
};
match held.data() {
Some(data) => look(data, first, members, &base, rows, returns),
None => row_at_a_time(input, members, &base, rows, returns),
}
}
_ => row_at_a_time(input, members, &base, rows, returns),
}
}
fn look<A: Fn(usize) -> usize>(
data: &Data,
at: A,
members: &Members,
base: &Validity,
rows: usize,
returns: &LogicalType,
) -> Result<Vector> {
match (&members.held, data) {
(Held::Text(set), Data::Varlen(column)) => answer(rows, base, members, returns, |index| {
column.get(at(index)).is_some_and(|text| set.contains(text))
}),
(Held::Whole(set), Data::Int8(held)) => {
answer(rows, base, members, returns, |index| holds(set, held.as_slice(), at(index)))
}
(Held::Whole(set), Data::Int16(held)) => {
answer(rows, base, members, returns, |index| holds(set, held.as_slice(), at(index)))
}
(Held::Whole(set), Data::Int32(held)) => {
answer(rows, base, members, returns, |index| holds(set, held.as_slice(), at(index)))
}
(Held::Whole(set), Data::Int64(held)) => {
answer(rows, base, members, returns, |index| holds(set, held.as_slice(), at(index)))
}
(Held::Whole(set), Data::Int128(held)) => {
answer(rows, base, members, returns, |index| holds(set, held.as_slice(), at(index)))
}
(Held::Whole(set), Data::UInt8(held)) => {
answer(rows, base, members, returns, |index| holds(set, held.as_slice(), at(index)))
}
(Held::Whole(set), Data::UInt16(held)) => {
answer(rows, base, members, returns, |index| holds(set, held.as_slice(), at(index)))
}
(Held::Whole(set), Data::UInt32(held)) => {
answer(rows, base, members, returns, |index| holds(set, held.as_slice(), at(index)))
}
(Held::Whole(set), Data::UInt64(held)) => {
answer(rows, base, members, returns, |index| holds(set, held.as_slice(), at(index)))
}
_ => Err(rudb_common::Error::internal(format!(
"an IN list over a column this kernel does not read, which is {returns}"
))),
}
}
fn holds<T: Copy>(set: &HashSet<i128>, values: &[T], index: usize) -> bool
where
i128: From<T>,
{
values.get(index).is_some_and(|&held| set.contains(&i128::from(held)))
}
fn answer(
rows: usize,
base: &Validity,
members: &Members,
returns: &LogicalType,
found: impl Fn(usize) -> bool,
) -> Result<Vector> {
let mut out = vec![false; rows];
let mut live = vec![false; rows];
for index in 0..rows {
if !base.is_valid(index) {
continue;
}
let hit = found(index);
live[index] = hit || !members.has_null;
out[index] = hit != members.negated;
}
let validity = Validity::from_run(&live).normalize(rows);
Ok(Vector::flat(returns.clone(), Data::Bool(out.into()))?.with_validity(validity))
}
fn row_at_a_time(
input: &Vector,
members: &Members,
base: &Validity,
rows: usize,
returns: &LogicalType,
) -> Result<Vector> {
let held: Vec<Value> = (0..rows).map(|index| input.value_at(index)).collect();
answer(rows, base, members, returns, |index| match (&members.held, &held[index]) {
(Held::Text(set), Value::Varchar(text)) => set.contains(text.as_str()),
(Held::Whole(set), value) => number(value).is_some_and(|held| set.contains(&held)),
_ => false,
})
}
#[cfg(test)]
mod tests {
use rudb_common::{LogicalType, Value};
use rudb_vector::Vector;
use super::{Members, in_set};
fn over(input: &Vector, list: &[Value], negated: bool) -> Vec<Value> {
let members = Members::of(list, negated).expect("this list folds");
let answer = in_set(input, &members, &LogicalType::Boolean).expect("the lookup runs");
(0..input.len()).map(|row| answer.value_at(row)).collect()
}
fn numbers() -> Vector {
Vector::from_values(
LogicalType::Integer,
&[Value::Integer(1), Value::Integer(7), Value::Null, Value::Integer(3)],
)
.expect("four integers")
}
#[test]
fn a_row_in_the_list_is_true_and_a_row_outside_it_is_false() {
assert_eq!(
over(&numbers(), &[Value::Integer(1), Value::Integer(3)], false),
[Value::Boolean(true), Value::Boolean(false), Value::Null, Value::Boolean(true)]
);
}
#[test]
fn a_not_in_is_the_same_lookup_read_the_other_way() {
assert_eq!(
over(&numbers(), &[Value::Integer(1), Value::Integer(3)], true),
[Value::Boolean(false), Value::Boolean(true), Value::Null, Value::Boolean(false)]
);
}
#[test]
fn a_miss_against_a_list_with_a_null_in_it_is_null() {
let list = [Value::Integer(1), Value::Null, Value::Integer(3)];
assert_eq!(
over(&numbers(), &list, false),
[Value::Boolean(true), Value::Null, Value::Null, Value::Boolean(true)]
);
assert_eq!(
over(&numbers(), &list, true),
[Value::Boolean(false), Value::Null, Value::Null, Value::Boolean(false)]
);
}
#[test]
fn a_dictionary_column_is_read_through_its_codes() {
let values = Vector::from_values(
LogicalType::Varchar,
&[Value::Varchar("a".into()), Value::Varchar("b".into()), Value::Null],
)
.expect("builds");
let text = Vector::dictionary(vec![0, 2, 1, 0], values).expect("codes are in range");
let list = [Value::Varchar("a".into()), Value::Varchar("c".into())];
assert_eq!(
over(&text, &list, false),
[Value::Boolean(true), Value::Null, Value::Boolean(false), Value::Boolean(true)]
);
}
#[test]
fn a_constant_column_answers_every_row_the_same() {
let held = Vector::constant(LogicalType::Integer, Value::Integer(3), 3);
let list = [Value::Integer(1), Value::Integer(3)];
assert_eq!(over(&held, &list, false), vec![Value::Boolean(true); 3]);
}
#[test]
fn a_run_length_column_reads_the_same_as_the_flat_one_it_stands_for() {
let flat = Vector::from_values(
LogicalType::Integer,
&[Value::Integer(1), Value::Integer(1), Value::Integer(7), Value::Integer(7)],
)
.expect("four integers");
let runs = flat.clone().run_encoded().expect("two runs");
let list = [Value::Integer(1), Value::Integer(3)];
assert_eq!(over(&runs, &list, false), over(&flat, &list, false));
}
#[test]
fn a_list_of_one_is_left_alone_because_a_comparison_is_already_that() {
assert!(Members::of(&[Value::Integer(1)], false).is_none());
}
#[test]
fn a_list_of_floats_does_not_fold() {
assert!(Members::of(&[Value::Double(1.0), Value::Double(2.0)], false).is_none());
}
#[test]
fn a_list_of_two_kinds_does_not_fold() {
let mixed = [Value::Integer(1), Value::Varchar("a".into())];
assert!(Members::of(&mixed, false).is_none());
}
#[test]
fn a_list_of_nothing_but_nulls_does_not_fold() {
assert!(Members::of(&[Value::Null, Value::Null], false).is_none());
}
#[test]
fn a_list_says_how_many_distinct_values_it_holds() {
let list = [Value::Integer(1), Value::Integer(1), Value::Integer(2), Value::Null];
let members = Members::of(&list, false).expect("this list folds");
assert_eq!(members.len(), 2);
assert!(!members.is_empty());
}
}