use std::any::Any;
use rustc_hash::{FxHashMap, FxHashSet};
use super::{find_column_index, resolve_alias, Expression};
use radixdb_core::{DataType, I64Set, Result, Row, Schema, Value};
#[doc(hidden)]
pub fn has_cross_numeric_physical_variant(target_type: DataType, values: &[Value]) -> bool {
values.iter().any(|value| match target_type {
DataType::Integer => matches!(value, Value::Float(_)) || value.as_decimal_parts().is_some(),
DataType::Float => matches!(value, Value::Integer(_)) || value.as_decimal_parts().is_some(),
DataType::Decimal => matches!(value, Value::Integer(_) | Value::Float(_)),
_ => false,
})
}
#[derive(Debug, Clone)]
enum HashedValues {
None,
Integers(I64Set),
Strings(FxHashSet<String>),
Booleans { has_true: bool, has_false: bool },
Mixed,
}
#[derive(Debug, Clone)]
pub struct InListExpr {
column: String,
values: Vec<Value>,
not: bool,
col_index: Option<usize>,
hashed: HashedValues,
has_null: bool,
aliases: FxHashMap<String, String>,
original_column: Option<String>,
cached_min: Option<Value>,
cached_max: Option<Value>,
}
impl InListExpr {
#[inline]
fn exact_integer_key(value: &Value) -> Option<i64> {
value
.exact_integer_identity()
.filter(|integer| *integer != i64::MIN)
}
#[inline]
fn scalar_values_equal(left: &Value, right: &Value) -> bool {
if left.is_null() || right.is_null() {
return false;
}
match left.compare(right) {
Ok(std::cmp::Ordering::Equal) => true,
Ok(_) => false,
Err(_) => left == right,
}
}
pub fn new(column: impl Into<String>, values: Vec<Value>) -> Self {
let has_null = values.iter().any(|v| v.is_null());
Self {
column: column.into(),
values,
not: false,
col_index: None,
hashed: HashedValues::None,
has_null,
aliases: FxHashMap::default(),
cached_min: None,
cached_max: None,
original_column: None,
}
}
pub fn not_in(column: impl Into<String>, values: Vec<Value>) -> Self {
let has_null = values.iter().any(|v| v.is_null());
Self {
column: column.into(),
values,
not: true,
col_index: None,
hashed: HashedValues::None,
has_null,
aliases: FxHashMap::default(),
original_column: None,
cached_min: None,
cached_max: None,
}
}
pub fn is_not(&self) -> bool {
self.not
}
pub fn values(&self) -> &[Value] {
&self.values
}
pub fn get_values(&self) -> &[Value] {
&self.values
}
fn build_hash_sets(&mut self) {
if self.values.is_empty() {
self.hashed = HashedValues::None;
return;
}
let first_type = self.values.iter().find_map(|v| match v {
Value::Integer(_) => Some("int"),
Value::Float(_) => Some("float"),
Value::Text(_) => Some("text"),
Value::Boolean(_) => Some("bool"),
Value::Null(_) => None,
_ => Some("other"),
});
match first_type {
Some("int") => {
let mut set = I64Set::new();
let mut all_int = true;
for v in &self.values {
match v {
Value::Integer(_) | Value::Float(_) => {
if let Some(integer) = Self::exact_integer_key(v) {
set.insert(integer);
} else {
all_int = false;
break;
}
}
Value::Null(_) => {} _ => {
all_int = false;
break;
}
}
}
if all_int {
self.hashed = HashedValues::Integers(set);
} else {
self.hashed = HashedValues::Mixed;
}
}
Some("text") => {
let mut set = FxHashSet::default();
let mut all_text = true;
for v in &self.values {
match v {
Value::Text(s) => {
set.insert(s.to_string());
}
Value::Null(_) => {}
_ => {
all_text = false;
break;
}
}
}
if all_text {
self.hashed = HashedValues::Strings(set);
} else {
self.hashed = HashedValues::Mixed;
}
}
Some("bool") => {
let mut has_true = false;
let mut has_false = false;
for v in &self.values {
match v {
Value::Boolean(true) => has_true = true,
Value::Boolean(false) => has_false = true,
Value::Null(_) => {}
_ => {}
}
}
self.hashed = HashedValues::Booleans {
has_true,
has_false,
};
}
_ => {
self.hashed = HashedValues::Mixed;
}
}
let mut min_val: Option<Value> = None;
let mut max_val: Option<Value> = None;
let mut comparable = true;
for v in &self.values {
if v.is_null() {
continue;
}
match (&min_val, &max_val) {
(None, _) => {
min_val = Some(v.clone());
max_val = Some(v.clone());
}
(Some(cur_min), Some(cur_max)) => match (v.compare(cur_min), v.compare(cur_max)) {
(Ok(min_order), Ok(max_order)) => {
if min_order == std::cmp::Ordering::Less {
min_val = Some(v.clone());
}
if max_order == std::cmp::Ordering::Greater {
max_val = Some(v.clone());
}
}
_ => {
comparable = false;
break;
}
},
_ => {}
}
}
if comparable {
self.cached_min = min_val;
self.cached_max = max_val;
} else {
self.cached_min = None;
self.cached_max = None;
}
}
#[inline]
fn check_integer(&self, val: i64) -> bool {
match &self.hashed {
HashedValues::Integers(set) if val != i64::MIN => set.contains(val),
_ => self
.values
.iter()
.any(|value| Self::scalar_values_equal(&Value::Integer(val), value)),
}
}
#[inline]
fn check_float(&self, val: f64) -> bool {
self.values
.iter()
.any(|value| Self::scalar_values_equal(&Value::Float(val), value))
}
#[inline]
fn check_string(&self, val: &str) -> bool {
match &self.hashed {
HashedValues::Strings(set) => set.contains(val),
_ => {
for v in &self.values {
if let Some(list_val) = v.as_string() {
if val == list_val {
return true;
}
}
}
false
}
}
}
#[inline]
fn check_boolean(&self, val: bool) -> bool {
match &self.hashed {
HashedValues::Booleans {
has_true,
has_false,
} => {
if val {
*has_true
} else {
*has_false
}
}
_ => {
self.values.iter().any(|v| v.as_boolean() == Some(val))
}
}
}
#[inline]
fn check_scalar(&self, val: &Value) -> bool {
self.values
.iter()
.any(|list_val| Self::scalar_values_equal(val, list_val))
}
}
impl Expression for InListExpr {
fn evaluate(&self, row: &Row) -> Result<bool> {
let col_idx = match self.col_index {
Some(idx) if idx < row.len() => idx,
_ => return Ok(false),
};
let col_value = &row[col_idx];
if col_value.is_null() {
return Ok(false);
}
let found = match col_value {
Value::Integer(val) => self.check_integer(*val),
Value::Float(val) => self.check_float(*val),
Value::Text(val) => self.check_string(val),
Value::Boolean(val) => self.check_boolean(*val),
_ => self.check_scalar(col_value),
};
if found {
Ok(!self.not) } else if self.has_null {
Ok(false)
} else {
Ok(self.not) }
}
fn evaluate_fast(&self, row: &Row) -> bool {
let col_idx = match self.col_index {
Some(idx) if idx < row.len() => idx,
_ => return false, };
let col_value = &row[col_idx];
if col_value.is_null() {
return false;
}
let found = match col_value {
Value::Integer(val) => self.check_integer(*val),
Value::Float(val) => self.check_float(*val),
Value::Text(val) => self.check_string(val),
Value::Boolean(val) => self.check_boolean(*val),
_ => self.check_scalar(col_value),
};
if found {
!self.not } else if self.has_null {
false
} else {
self.not }
}
fn with_aliases(&self, aliases: &FxHashMap<String, String>) -> Box<dyn Expression> {
let resolved = resolve_alias(&self.column, aliases);
let mut expr = self.clone();
if resolved != self.column {
expr.original_column = Some(self.column.clone());
expr.column = resolved.to_string();
}
expr.aliases = aliases.clone();
expr.col_index = None;
expr.hashed = HashedValues::None; Box::new(expr)
}
fn prepare_for_schema(&mut self, schema: &Schema) {
if self.col_index.is_some() {
return;
}
self.col_index = find_column_index(schema, &self.column);
self.build_hash_sets();
}
fn collect_column_indices(&self, out: &mut Vec<usize>) -> bool {
if let Some(idx) = self.col_index {
out.push(idx);
true
} else {
false
}
}
fn is_prepared(&self) -> bool {
self.col_index.is_some()
}
fn get_column_name(&self) -> Option<&str> {
Some(&self.column)
}
fn get_in_list_info(&self) -> Option<(&str, &[Value], bool, bool)> {
Some((&self.column, &self.values, self.not, self.has_null))
}
fn collect_comparisons(&self) -> Vec<(&str, radixdb_core::Operator, &Value)> {
if self.not {
return vec![];
}
match (&self.cached_min, &self.cached_max) {
(Some(min), Some(max)) => {
vec![
(&self.column, radixdb_core::Operator::Gte, min),
(&self.column, radixdb_core::Operator::Lte, max),
]
}
_ => vec![],
}
}
fn can_use_index(&self) -> bool {
true
}
fn is_conjunctive_simple(&self) -> bool {
false
}
fn clone_box(&self) -> Box<dyn Expression> {
Box::new(self.clone())
}
fn is_unknown_due_to_null(&self, row: &Row) -> bool {
let Some(value) = self.col_index.and_then(|index| row.get(index)) else {
return false;
};
if value.is_null() {
return true;
}
self.has_null && !self.check_scalar(value)
}
fn as_any(&self) -> &dyn Any {
self
}
}
#[cfg(test)]
mod tests {
use super::*;
use radixdb_core::{DataType, SchemaBuilder};
fn test_schema() -> Schema {
SchemaBuilder::new("test")
.add_primary_key("id", DataType::Integer)
.add("name", DataType::Text)
.add("status", DataType::Text)
.build()
}
#[test]
fn test_integer_in() {
let schema = test_schema();
let row = Row::from_values(vec![
Value::integer(2),
Value::text("Alice"),
Value::text("active"),
]);
let mut expr = InListExpr::new(
"id",
vec![Value::integer(1), Value::integer(2), Value::integer(3)],
);
expr.prepare_for_schema(&schema);
assert!(expr.evaluate(&row).unwrap());
assert!(expr.evaluate_fast(&row));
let mut expr = InListExpr::new(
"id",
vec![Value::integer(5), Value::integer(6), Value::integer(7)],
);
expr.prepare_for_schema(&schema);
assert!(!expr.evaluate(&row).unwrap());
}
#[test]
fn test_string_in() {
let schema = test_schema();
let row = Row::from_values(vec![
Value::integer(1),
Value::text("Alice"),
Value::text("active"),
]);
let mut expr = InListExpr::new(
"status",
vec![
Value::text("active"),
Value::text("inactive"),
Value::text("pending"),
],
);
expr.prepare_for_schema(&schema);
assert!(expr.evaluate(&row).unwrap());
}
#[test]
fn test_cross_numeric_physical_variant_matrix() {
let decimal = Value::decimal(15, 2, 1);
assert!(!has_cross_numeric_physical_variant(
DataType::Integer,
&[Value::Integer(1)],
));
assert!(has_cross_numeric_physical_variant(
DataType::Integer,
&[Value::Float(1.0)],
));
assert!(has_cross_numeric_physical_variant(
DataType::Integer,
std::slice::from_ref(&decimal),
));
assert!(has_cross_numeric_physical_variant(
DataType::Float,
&[Value::Integer(1)],
));
assert!(!has_cross_numeric_physical_variant(
DataType::Float,
&[Value::Float(1.0)],
));
assert!(has_cross_numeric_physical_variant(
DataType::Decimal,
&[Value::Float(1.0)],
));
assert!(!has_cross_numeric_physical_variant(
DataType::Decimal,
&[decimal],
));
}
#[test]
fn test_not_in() {
let schema = test_schema();
let row = Row::from_values(vec![
Value::integer(4),
Value::text("Alice"),
Value::text("active"),
]);
let mut expr = InListExpr::not_in(
"id",
vec![Value::integer(1), Value::integer(2), Value::integer(3)],
);
expr.prepare_for_schema(&schema);
assert!(expr.evaluate(&row).unwrap());
let mut expr = InListExpr::not_in(
"id",
vec![Value::integer(4), Value::integer(5), Value::integer(6)],
);
expr.prepare_for_schema(&schema);
assert!(!expr.evaluate(&row).unwrap());
}
#[test]
fn test_null_in() {
let schema = test_schema();
let row = Row::from_values(vec![
Value::null(DataType::Integer),
Value::text("Alice"),
Value::text("active"),
]);
let mut expr = InListExpr::new(
"id",
vec![Value::integer(1), Value::integer(2), Value::integer(3)],
);
expr.prepare_for_schema(&schema);
assert!(!expr.evaluate(&row).unwrap());
let mut expr = InListExpr::not_in(
"id",
vec![Value::integer(1), Value::integer(2), Value::integer(3)],
);
expr.prepare_for_schema(&schema);
assert!(!expr.evaluate(&row).unwrap());
}
#[test]
fn test_empty_list() {
let schema = test_schema();
let row = Row::from_values(vec![
Value::integer(1),
Value::text("Alice"),
Value::text("active"),
]);
let mut expr = InListExpr::new("id", vec![]);
expr.prepare_for_schema(&schema);
assert!(!expr.evaluate(&row).unwrap());
let mut expr = InListExpr::not_in("id", vec![]);
expr.prepare_for_schema(&schema);
assert!(expr.evaluate(&row).unwrap());
}
#[test]
fn test_mixed_numeric_types() {
let schema = test_schema();
let row = Row::from_values(vec![
Value::integer(2),
Value::text("Alice"),
Value::text("active"),
]);
let mut expr = InListExpr::new(
"id",
vec![Value::float(1.0), Value::float(2.0), Value::float(3.0)],
);
expr.prepare_for_schema(&schema);
assert!(expr.evaluate(&row).unwrap());
}
#[test]
fn test_in_list_exact_mixed_numeric_identity_boundaries() {
let schema = test_schema();
let two53 = 1_i64 << 53;
let mut exact = InListExpr::new(
"id",
vec![Value::Integer(two53), Value::Float(two53 as f64)],
);
exact.prepare_for_schema(&schema);
let exact_row = Row::from_values(vec![
Value::Integer(two53),
Value::text("Alice"),
Value::text("active"),
]);
assert!(exact.evaluate(&exact_row).unwrap());
assert!(exact.evaluate_fast(&exact_row));
let neighbor_row = Row::from_values(vec![
Value::Integer(two53 + 1),
Value::text("Alice"),
Value::text("active"),
]);
assert!(!exact.evaluate(&neighbor_row).unwrap());
assert!(!exact.evaluate_fast(&neighbor_row));
let mut out_of_range = InListExpr::new(
"id",
vec![Value::Integer(i64::MIN), Value::Float(2_f64.powi(63))],
);
out_of_range.prepare_for_schema(&schema);
let min_row = Row::from_values(vec![
Value::Integer(i64::MIN),
Value::text("Alice"),
Value::text("active"),
]);
assert!(out_of_range.evaluate(&min_row).unwrap());
assert!(out_of_range.evaluate_fast(&min_row));
let max_row = Row::from_values(vec![
Value::Integer(i64::MAX),
Value::text("Alice"),
Value::text("active"),
]);
assert!(!out_of_range.evaluate(&max_row).unwrap());
assert!(!out_of_range.evaluate_fast(&max_row));
let mut fractional = InListExpr::new(
"id",
vec![Value::Integer(0), Value::Float(1.5), Value::Float(-0.0)],
);
fractional.prepare_for_schema(&schema);
let one_row = Row::from_values(vec![
Value::Integer(1),
Value::text("Alice"),
Value::text("active"),
]);
assert!(!fractional.evaluate(&one_row).unwrap());
assert!(!fractional.evaluate_fast(&one_row));
let zero_row = Row::from_values(vec![
Value::Integer(0),
Value::text("Alice"),
Value::text("active"),
]);
assert!(fractional.evaluate(&zero_row).unwrap());
assert!(fractional.evaluate_fast(&zero_row));
let mut nan = InListExpr::new("id", vec![Value::Float(f64::NAN)]);
nan.prepare_for_schema(&schema);
let nan_row = Row::from_values(vec![
Value::Float(f64::NAN),
Value::text("Alice"),
Value::text("active"),
]);
assert!(nan.evaluate(&nan_row).unwrap());
assert!(nan.evaluate_fast(&nan_row));
}
#[test]
fn test_uuid_in_uses_scalar_equality_contract() {
let uuid_a = [0x11; 16];
let uuid_b = [0x22; 16];
let uuid_c = [0x33; 16];
let schema = SchemaBuilder::new("uuid_values")
.add_primary_key("row_id", DataType::Integer)
.add("id", DataType::Uuid)
.build();
let row = Row::from_values(vec![Value::integer(1), Value::uuid(uuid_b)]);
let mut matching = InListExpr::new("id", vec![Value::uuid(uuid_a), Value::uuid(uuid_b)]);
matching.prepare_for_schema(&schema);
assert!(matching.evaluate(&row).unwrap());
assert!(matching.evaluate_fast(&row));
let mut missing = InListExpr::new("id", vec![Value::uuid(uuid_a), Value::uuid(uuid_c)]);
missing.prepare_for_schema(&schema);
assert!(!missing.evaluate(&row).unwrap());
assert!(!missing.evaluate_fast(&row));
let mut duplicate = InListExpr::new(
"id",
vec![
Value::uuid(uuid_b),
Value::uuid(uuid_b),
Value::null(DataType::Uuid),
],
);
duplicate.prepare_for_schema(&schema);
assert!(duplicate.evaluate(&row).unwrap());
assert!(duplicate.evaluate_fast(&row));
}
#[test]
fn test_with_aliases() {
let schema = test_schema();
let row = Row::from_values(vec![
Value::integer(1),
Value::text("Alice"),
Value::text("active"),
]);
let mut aliases = FxHashMap::default();
aliases.insert("i".to_string(), "id".to_string());
let expr = InListExpr::new("i", vec![Value::integer(1), Value::integer(2)]);
let mut aliased = expr.with_aliases(&aliases);
aliased.prepare_for_schema(&schema);
assert!(aliased.evaluate(&row).unwrap());
}
#[test]
fn test_is_not() {
let expr = InListExpr::new("id", vec![Value::integer(1)]);
assert!(!expr.is_not());
let expr = InListExpr::not_in("id", vec![Value::integer(1)]);
assert!(expr.is_not());
}
#[test]
fn test_not_in_with_null_in_list() {
let schema = test_schema();
let row = Row::from_values(vec![
Value::integer(1),
Value::text("Alice"),
Value::text("active"),
]);
let mut expr = InListExpr::not_in(
"id",
vec![Value::integer(2), Value::null(DataType::Integer)],
);
expr.prepare_for_schema(&schema);
assert!(!expr.evaluate(&row).unwrap());
assert!(!expr.evaluate_fast(&row));
let row2 = Row::from_values(vec![
Value::integer(2),
Value::text("Bob"),
Value::text("active"),
]);
assert!(!expr.evaluate(&row2).unwrap());
assert!(!expr.evaluate_fast(&row2));
}
#[test]
fn test_in_with_null_in_list() {
let schema = test_schema();
let row = Row::from_values(vec![
Value::integer(1),
Value::text("Alice"),
Value::text("active"),
]);
let mut expr = InListExpr::new(
"id",
vec![Value::integer(2), Value::null(DataType::Integer)],
);
expr.prepare_for_schema(&schema);
assert!(!expr.evaluate(&row).unwrap());
assert!(!expr.evaluate_fast(&row));
let row2 = Row::from_values(vec![
Value::integer(2),
Value::text("Bob"),
Value::text("active"),
]);
assert!(expr.evaluate(&row2).unwrap());
assert!(expr.evaluate_fast(&row2));
}
#[test]
fn test_not_in_without_null() {
let schema = test_schema();
let row = Row::from_values(vec![
Value::integer(1),
Value::text("Alice"),
Value::text("active"),
]);
let mut expr = InListExpr::not_in("id", vec![Value::integer(2), Value::integer(3)]);
expr.prepare_for_schema(&schema);
assert!(expr.evaluate(&row).unwrap());
assert!(expr.evaluate_fast(&row));
}
#[test]
fn heterogeneous_list_never_exports_partial_bounds() {
let schema = test_schema();
let row = Row::from_values(vec![
Value::integer(1),
Value::text("match"),
Value::text("active"),
]);
let mut expr = InListExpr::new("name", vec![Value::integer(100), Value::text("match")]);
expr.prepare_for_schema(&schema);
assert!(expr.evaluate(&row).unwrap());
assert!(expr.collect_comparisons().is_empty());
}
#[test]
fn in_list_reports_unknown_for_null_input_or_null_tail() {
let schema = test_schema();
let mut expr = InListExpr::new(
"id",
vec![Value::integer(2), Value::null(DataType::Integer)],
);
expr.prepare_for_schema(&schema);
let miss = Row::from_values(vec![
Value::integer(1),
Value::text("Alice"),
Value::text("active"),
]);
let null = Row::from_values(vec![
Value::null(DataType::Integer),
Value::text("Alice"),
Value::text("active"),
]);
assert!(expr.is_unknown_due_to_null(&miss));
assert!(expr.is_unknown_due_to_null(&null));
}
}