use std::borrow::Cow;
use std::cmp::Ordering;
use arrow_array::{Array, RecordBatch, cast::AsArray, types::*};
use arrow_schema::DataType;
use crate::backend::types::{CmpOp, FilterExpr, Scalar, StringMatchKind};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum EvalError {
UnknownColumn(String),
Unsupported(String),
}
impl std::fmt::Display for EvalError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
EvalError::UnknownColumn(c) => write!(f, "no such column: {c}"),
EvalError::Unsupported(why) => write!(f, "cannot evaluate filter: {why}"),
}
}
}
impl std::error::Error for EvalError {}
#[derive(Debug, Clone, PartialEq)]
pub enum Cell<'a> {
Null,
Bool(bool),
Int(i64),
UInt(u64),
Float(f64),
Str(Cow<'a, str>),
List(Vec<Cell<'a>>),
}
pub trait RowAccessor {
fn column(&self, name: &str) -> Result<Cell<'_>, EvalError>;
}
pub struct ArrowRow<'a> {
batch: &'a RecordBatch,
row: usize,
}
impl<'a> ArrowRow<'a> {
pub fn new(batch: &'a RecordBatch, row: usize) -> Self {
Self { batch, row }
}
}
impl RowAccessor for ArrowRow<'_> {
fn column(&self, name: &str) -> Result<Cell<'_>, EvalError> {
let idx = self
.batch
.schema()
.index_of(name)
.map_err(|_| EvalError::UnknownColumn(name.to_string()))?;
cell_at(self.batch.column(idx).as_ref(), self.row)
}
}
fn cell_at(array: &dyn Array, row: usize) -> Result<Cell<'_>, EvalError> {
if array.is_null(row) {
return Ok(Cell::Null);
}
Ok(match array.data_type() {
DataType::Boolean => Cell::Bool(array.as_boolean().value(row)),
DataType::Int8 => Cell::Int(array.as_primitive::<Int8Type>().value(row) as i64),
DataType::Int16 => Cell::Int(array.as_primitive::<Int16Type>().value(row) as i64),
DataType::Int32 => Cell::Int(array.as_primitive::<Int32Type>().value(row) as i64),
DataType::Int64 => Cell::Int(array.as_primitive::<Int64Type>().value(row)),
DataType::UInt8 => Cell::UInt(array.as_primitive::<UInt8Type>().value(row) as u64),
DataType::UInt16 => Cell::UInt(array.as_primitive::<UInt16Type>().value(row) as u64),
DataType::UInt32 => Cell::UInt(array.as_primitive::<UInt32Type>().value(row) as u64),
DataType::UInt64 => Cell::UInt(array.as_primitive::<UInt64Type>().value(row)),
DataType::Float32 => Cell::Float(array.as_primitive::<Float32Type>().value(row) as f64),
DataType::Float64 => Cell::Float(array.as_primitive::<Float64Type>().value(row)),
DataType::Utf8 => Cell::Str(Cow::Borrowed(array.as_string::<i32>().value(row))),
DataType::LargeUtf8 => Cell::Str(Cow::Borrowed(array.as_string::<i64>().value(row))),
DataType::List(_) => {
let inner = array.as_list::<i32>().value(row);
Cell::List(collect_list(inner.as_ref())?)
}
DataType::LargeList(_) => {
let inner = array.as_list::<i64>().value(row);
Cell::List(collect_list(inner.as_ref())?)
}
other => {
return Err(EvalError::Unsupported(format!(
"no native evaluation for column type {other}"
)));
}
})
}
fn collect_list(inner: &dyn Array) -> Result<Vec<Cell<'static>>, EvalError> {
(0..inner.len())
.map(|i| {
Ok(match cell_at(inner, i)? {
Cell::Null => Cell::Null,
Cell::Bool(b) => Cell::Bool(b),
Cell::Int(v) => Cell::Int(v),
Cell::UInt(v) => Cell::UInt(v),
Cell::Float(v) => Cell::Float(v),
Cell::Str(s) => Cell::Str(Cow::Owned(s.into_owned())),
Cell::List(_) => {
return Err(EvalError::Unsupported("nested lists".to_string()));
}
})
})
.collect()
}
impl FilterExpr {
pub fn eval(&self, row: &dyn RowAccessor) -> Result<Option<bool>, EvalError> {
match self {
FilterExpr::Literal(b) => Ok(Some(*b)),
FilterExpr::And(parts) => {
let mut unknown = false;
for p in parts {
match p.eval(row)? {
Some(false) => return Ok(Some(false)),
Some(true) => {}
None => unknown = true,
}
}
Ok(if unknown { None } else { Some(true) })
}
FilterExpr::Or(parts) => {
let mut unknown = false;
for p in parts {
match p.eval(row)? {
Some(true) => return Ok(Some(true)),
Some(false) => {}
None => unknown = true,
}
}
Ok(if unknown { None } else { Some(false) })
}
FilterExpr::Not(inner) => Ok(inner.eval(row)?.map(|b| !b)),
FilterExpr::Compare { column, op, value } => compare(&row.column(column)?, *op, value),
FilterExpr::In { column, values } => {
let cell = row.column(column)?;
let mut unknown = false;
for v in values {
match compare(&cell, CmpOp::Eq, v)? {
Some(true) => return Ok(Some(true)),
Some(false) => {}
None => unknown = true,
}
}
Ok(if unknown { None } else { Some(false) })
}
FilterExpr::ArrayContains { column, value } => {
let cell = row.column(column)?;
list_contains(&cell, value)
}
FilterExpr::StringMatch {
column,
kind,
pattern,
} => match row.column(column)? {
Cell::Null => Ok(None),
Cell::Str(s) => Ok(Some(match kind {
StringMatchKind::Contains => s.contains(pattern.as_str()),
StringMatchKind::StartsWith => s.starts_with(pattern.as_str()),
StringMatchKind::EndsWith => s.ends_with(pattern.as_str()),
})),
other => Err(EvalError::Unsupported(format!(
"string match on a non-string cell {other:?}"
))),
},
FilterExpr::IsNull(column) => Ok(Some(matches!(row.column(column)?, Cell::Null))),
FilterExpr::IsNotNull(column) => Ok(Some(!matches!(row.column(column)?, Cell::Null))),
FilterExpr::Raw(s) => Err(EvalError::Unsupported(format!(
"raw backend predicate {s:?} has no native evaluation"
))),
}
}
}
fn compare(cell: &Cell<'_>, op: CmpOp, value: &Scalar) -> Result<Option<bool>, EvalError> {
if matches!(cell, Cell::Null) || matches!(value, Scalar::Null) {
return Ok(None);
}
Ok(order(cell, value)?.map(|o| match op {
CmpOp::Eq => o == Ordering::Equal,
CmpOp::NotEq => o != Ordering::Equal,
CmpOp::Lt => o == Ordering::Less,
CmpOp::LtEq => o != Ordering::Greater,
CmpOp::Gt => o == Ordering::Greater,
CmpOp::GtEq => o != Ordering::Less,
}))
}
fn list_contains(cell: &Cell<'_>, value: &Scalar) -> Result<Option<bool>, EvalError> {
let Cell::List(items) = cell else {
if matches!(cell, Cell::Null) {
return Ok(None);
}
return Err(EvalError::Unsupported(
"array_contains on a non-list column".to_string(),
));
};
if matches!(value, Scalar::Null) {
return Ok(None);
}
let mut unknown = false;
for item in items {
match compare(item, CmpOp::Eq, value)? {
Some(true) => return Ok(Some(true)),
Some(false) => {}
None => unknown = true,
}
}
Ok(if unknown { None } else { Some(false) })
}
#[derive(Debug, Clone, Copy)]
enum Num {
I(i64),
U(u64),
F(f64),
}
fn cell_num(c: &Cell<'_>) -> Option<Num> {
match c {
Cell::Int(v) => Some(Num::I(*v)),
Cell::UInt(v) => Some(Num::U(*v)),
Cell::Float(v) => Some(Num::F(*v)),
_ => None,
}
}
fn scalar_num(s: &Scalar) -> Option<Num> {
match s {
Scalar::Int(v) => Some(Num::I(*v)),
Scalar::UInt(v) => Some(Num::U(*v)),
Scalar::Float(v) => Some(Num::F(*v)),
_ => None,
}
}
fn order(cell: &Cell<'_>, value: &Scalar) -> Result<Option<Ordering>, EvalError> {
match (cell, value) {
(Cell::Bool(a), Scalar::Bool(b)) => Ok(Some(a.cmp(b))),
(Cell::Str(a), Scalar::Str(b)) => Ok(Some(a.as_ref().cmp(b.as_str()))),
_ => match (cell_num(cell), scalar_num(value)) {
(Some(a), Some(b)) => Ok(cmp_num(a, b)),
_ => Err(EvalError::Unsupported(format!(
"cannot compare cell {cell:?} with literal {value:?}"
))),
},
}
}
fn cmp_num(a: Num, b: Num) -> Option<Ordering> {
match (a, b) {
(Num::I(x), Num::I(y)) => Some(x.cmp(&y)),
(Num::U(x), Num::U(y)) => Some(x.cmp(&y)),
(Num::I(x), Num::U(y)) => Some(if x < 0 {
Ordering::Less
} else {
(x as u64).cmp(&y)
}),
(Num::U(x), Num::I(y)) => Some(if y < 0 {
Ordering::Greater
} else {
x.cmp(&(y as u64))
}),
(Num::F(x), Num::F(y)) => x.partial_cmp(&y),
(Num::I(x), Num::F(y)) => cmp_i64_f64(x, y),
(Num::F(x), Num::I(y)) => cmp_i64_f64(y, x).map(Ordering::reverse),
(Num::U(x), Num::F(y)) => cmp_u64_f64(x, y),
(Num::F(x), Num::U(y)) => cmp_u64_f64(y, x).map(Ordering::reverse),
}
}
const TWO_POW_63: f64 = 9_223_372_036_854_775_808.0;
const TWO_POW_64: f64 = 18_446_744_073_709_551_616.0;
fn cmp_i64_f64(i: i64, f: f64) -> Option<Ordering> {
if f.is_nan() {
return None;
}
if f >= TWO_POW_63 {
return Some(Ordering::Less);
}
if f < -TWO_POW_63 {
return Some(Ordering::Greater);
}
let floor = f.floor();
Some(match i.cmp(&(floor as i64)) {
Ordering::Equal if f > floor => Ordering::Less,
o => o,
})
}
fn cmp_u64_f64(u: u64, f: f64) -> Option<Ordering> {
if f.is_nan() {
return None;
}
if f < 0.0 {
return Some(Ordering::Greater);
}
if f >= TWO_POW_64 {
return Some(Ordering::Less);
}
let floor = f.floor();
Some(match u.cmp(&(floor as u64)) {
Ordering::Equal if f > floor => Ordering::Less,
o => o,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backend::types::{CmpOp, FilterExpr, Scalar, StringMatchKind, ToSqlError};
use arrow_array::builder::{ListBuilder, StringBuilder};
use arrow_array::{
BooleanArray, Float64Array, Int64Array, RecordBatchIterator, StringArray, UInt64Array,
};
use arrow_schema::{Field, Schema};
use futures::TryStreamExt;
use proptest::prelude::*;
use std::sync::Arc;
const INTS: [Option<i64>; 6] = [Some(-2), Some(-1), Some(0), Some(1), Some(2), None];
const UINTS: [Option<u64>; 5] = [Some(0), Some(1), Some(2), Some(3), None];
const FLOATS: [Option<f64>; 5] = [Some(-1.5), Some(0.0), Some(1.0), Some(2.5), None];
const STRS: [Option<&str>; 5] = [Some("a"), Some("b"), Some("it's"), Some("a%b_c"), None];
const BOOLS: [Option<bool>; 3] = [Some(true), Some(false), None];
const LABEL_SETS: [Option<&[&str]>; 5] = [
Some(&["Person"]),
Some(&["Person", "Admin"]),
Some(&[]),
Some(&["it's"]),
None,
];
const ROWS: usize = 120;
fn schema() -> Arc<Schema> {
Arc::new(Schema::new(vec![
Field::new("rid", arrow_schema::DataType::UInt64, false),
Field::new("i", arrow_schema::DataType::Int64, true),
Field::new("u", arrow_schema::DataType::UInt64, true),
Field::new("f", arrow_schema::DataType::Float64, true),
Field::new("s", arrow_schema::DataType::Utf8, true),
Field::new("b", arrow_schema::DataType::Boolean, true),
Field::new(
"labels",
arrow_schema::DataType::List(Arc::new(Field::new(
"item",
arrow_schema::DataType::Utf8,
true,
))),
true,
),
]))
}
fn batch() -> RecordBatch {
let mut labels = ListBuilder::new(StringBuilder::new());
for r in 0..ROWS {
match LABEL_SETS[r % LABEL_SETS.len()] {
Some(items) => {
for it in items {
labels.values().append_value(it);
}
labels.append(true);
}
None => labels.append(false),
}
}
let labels = labels.finish();
let schema = Arc::new(Schema::new(vec![
schema().field(0).clone(),
schema().field(1).clone(),
schema().field(2).clone(),
schema().field(3).clone(),
schema().field(4).clone(),
schema().field(5).clone(),
Field::new("labels", labels.data_type().clone(), true),
]));
RecordBatch::try_new(
schema,
vec![
Arc::new(UInt64Array::from_iter_values(0..ROWS as u64)),
Arc::new(Int64Array::from_iter(
(0..ROWS).map(|r| INTS[r % INTS.len()]),
)),
Arc::new(UInt64Array::from_iter(
(0..ROWS).map(|r| UINTS[(r / 2) % UINTS.len()]),
)),
Arc::new(Float64Array::from_iter(
(0..ROWS).map(|r| FLOATS[(r / 3) % FLOATS.len()]),
)),
Arc::new(StringArray::from_iter(
(0..ROWS).map(|r| STRS[(r / 5) % STRS.len()]),
)),
Arc::new(BooleanArray::from_iter(
(0..ROWS).map(|r| BOOLS[(r / 7) % BOOLS.len()]),
)),
Arc::new(labels),
],
)
.expect("batch")
}
fn eval_rids(expr: &FilterExpr, batch: &RecordBatch) -> Result<Vec<u64>, EvalError> {
let rid = batch.column(0).as_primitive::<UInt64Type>();
let mut out = Vec::new();
for r in 0..batch.num_rows() {
if expr.eval(&ArrowRow::new(batch, r))? == Some(true) {
out.push(rid.value(r));
}
}
Ok(out)
}
async fn sql_rids(ds: &lance::Dataset, sql: &str) -> anyhow::Result<Vec<u64>> {
let mut scanner = ds.scan();
scanner.project(&["rid"])?;
scanner.filter(sql)?;
let batches: Vec<RecordBatch> = scanner.try_into_stream().await?.try_collect().await?;
let mut out: Vec<u64> = batches
.iter()
.flat_map(|b| {
let a = b.column(0).as_primitive::<UInt64Type>();
(0..b.num_rows()).map(|i| a.value(i)).collect::<Vec<_>>()
})
.collect();
out.sort_unstable();
Ok(out)
}
fn cmp_op() -> impl Strategy<Value = CmpOp> {
prop_oneof![
Just(CmpOp::Eq),
Just(CmpOp::NotEq),
Just(CmpOp::Lt),
Just(CmpOp::LtEq),
Just(CmpOp::Gt),
Just(CmpOp::GtEq),
]
}
fn typed_operand() -> impl Strategy<Value = (String, Scalar)> {
prop_oneof![
(-3i64..4).prop_map(|v| ("i".to_string(), Scalar::Int(v))),
(0u64..5).prop_map(|v| ("u".to_string(), Scalar::UInt(v))),
prop_oneof![
Just(-1.5f64),
Just(0.0),
Just(1.0),
Just(2.5),
Just(3.0),
Just(0.5)
]
.prop_map(|v| ("f".to_string(), Scalar::Float(v))),
prop_oneof![Just("a"), Just("b"), Just("it's"), Just("z")]
.prop_map(|v| ("s".to_string(), Scalar::Str(v.to_string()))),
any::<bool>().prop_map(|v| ("b".to_string(), Scalar::Bool(v))),
Just(("i".to_string(), Scalar::Null)),
Just(("s".to_string(), Scalar::Null)),
]
}
fn expr_strategy() -> impl Strategy<Value = FilterExpr> {
let leaf = prop_oneof![
8 => (typed_operand(), cmp_op())
.prop_map(|((column, value), op)| FilterExpr::Compare { column, op, value }),
3 => proptest::collection::vec(typed_operand(), 0..4).prop_map(|ops| {
let column = ops
.first()
.map(|(c, _)| c.clone())
.unwrap_or_else(|| "i".to_string());
let values = ops
.into_iter()
.filter(|(c, _)| *c == column)
.map(|(_, v)| v)
.collect();
FilterExpr::In { column, values }
}),
3 => prop_oneof![Just("Person"), Just("Admin"), Just("it's"), Just("Ghost")]
.prop_map(|v| FilterExpr::ArrayContains {
column: "labels".to_string(),
value: Scalar::Str(v.to_string()),
}),
3 => (
prop_oneof![
Just(StringMatchKind::Contains),
Just(StringMatchKind::StartsWith),
Just(StringMatchKind::EndsWith),
],
prop_oneof![Just("a"), Just("b"), Just("it"), Just("'"), Just("z")],
)
.prop_map(|(kind, pattern)| FilterExpr::StringMatch {
column: "s".to_string(),
kind,
pattern: pattern.to_string(),
}),
2 => prop_oneof![Just("i"), Just("s"), Just("b"), Just("labels")].prop_map(|c| {
FilterExpr::IsNull(c.to_string())
}),
2 => prop_oneof![Just("u"), Just("f"), Just("s")].prop_map(|c| {
FilterExpr::IsNotNull(c.to_string())
}),
1 => any::<bool>().prop_map(FilterExpr::Literal),
];
leaf.prop_recursive(3, 16, 4, |inner| {
prop_oneof![
3 => proptest::collection::vec(inner.clone(), 0..4).prop_map(FilterExpr::And),
3 => proptest::collection::vec(inner.clone(), 0..4).prop_map(FilterExpr::Or),
1 => inner.prop_map(FilterExpr::negate),
]
})
}
#[test]
fn eval_agrees_with_lance_sql() {
let rt = tokio::runtime::Runtime::new().expect("runtime");
let tmp = tempfile::tempdir().expect("tempdir");
let uri = tmp.path().join("t.lance").to_string_lossy().to_string();
let batch = batch();
let schema = batch.schema();
rt.block_on(async {
lance::Dataset::write(
RecordBatchIterator::new(vec![Ok(batch.clone())], schema),
&uri,
None,
)
.await
.expect("write");
});
let ds = rt.block_on(lance::Dataset::open(&uri)).expect("open");
proptest!(ProptestConfig::with_cases(200), |(expr in expr_strategy())| {
let sql = expr.to_sql().expect("renderable");
let native = eval_rids(&expr, &batch).expect("evaluable");
let lance = rt.block_on(sql_rids(&ds, &sql)).expect("scan");
prop_assert_eq!(
native, lance,
"\nexpr: {:?}\nsql: {}\n", expr, sql
);
});
}
#[test]
fn in_list_with_null_renders_as_explicit_unknown() {
let all_null = FilterExpr::In {
column: "i".into(),
values: vec![Scalar::Null],
};
assert_eq!(all_null.to_sql().unwrap(), "CAST(NULL AS BOOLEAN)");
let mixed = FilterExpr::In {
column: "i".into(),
values: vec![Scalar::Int(0), Scalar::Null],
};
assert_eq!(
mixed.to_sql().unwrap(),
"(i IN (0) OR CAST(NULL AS BOOLEAN))"
);
let plain = FilterExpr::In {
column: "i".into(),
values: vec![Scalar::Int(0), Scalar::Int(1)],
};
assert_eq!(plain.to_sql().unwrap(), "i IN (0, 1)");
}
#[test]
fn not_and_of_two_in_lists_with_null_agrees_with_lance() {
let rt = tokio::runtime::Runtime::new().expect("runtime");
let tmp = tempfile::tempdir().expect("tempdir");
let uri = tmp.path().join("t.lance").to_string_lossy().to_string();
let b = batch();
let schema = b.schema();
rt.block_on(async {
lance::Dataset::write(
RecordBatchIterator::new(vec![Ok(b.clone())], schema),
&uri,
None,
)
.await
.expect("write");
});
let ds = rt.block_on(lance::Dataset::open(&uri)).expect("open");
let expr = FilterExpr::Not(Box::new(FilterExpr::And(vec![
FilterExpr::In {
column: "i".into(),
values: vec![Scalar::Null],
},
FilterExpr::In {
column: "i".into(),
values: vec![Scalar::Int(0)],
},
])));
let native = eval_rids(&expr, &b).expect("evaluable");
let lance = rt
.block_on(sql_rids(&ds, &expr.to_sql().expect("renderable")))
.expect("scan");
assert_eq!(native, lance);
assert_eq!(native.len(), 80, "expected the non-null, non-zero rows");
}
#[test]
fn null_yields_unknown_not_false() {
let b = batch();
let expr = FilterExpr::Compare {
column: "s".to_string(),
op: CmpOp::NotEq,
value: Scalar::Str("a".to_string()),
};
let kept = eval_rids(&expr, &b).unwrap();
let s = b.column(4).as_string::<i32>();
for r in 0..b.num_rows() {
if s.is_null(r) {
assert!(!kept.contains(&(r as u64)), "NULL row {r} must not survive");
}
}
assert!(!kept.is_empty(), "non-NULL non-'a' rows must survive");
let with_null = FilterExpr::all([
FilterExpr::Compare {
column: "i".to_string(),
op: CmpOp::Eq,
value: Scalar::Null,
},
FilterExpr::Literal(true),
]);
assert!(eval_rids(&with_null, &b).unwrap().is_empty());
}
#[test]
fn empty_in_is_false_and_renders() {
let expr = FilterExpr::In {
column: "i".to_string(),
values: vec![],
};
assert_eq!(expr.to_sql().unwrap(), "false");
assert!(eval_rids(&expr, &batch()).unwrap().is_empty());
}
#[test]
fn in_with_null_is_unknown_only_on_miss() {
let b = batch();
let hit = FilterExpr::In {
column: "i".to_string(),
values: vec![Scalar::Int(1), Scalar::Null],
};
assert!(!eval_rids(&hit, &b).unwrap().is_empty());
let miss = FilterExpr::In {
column: "i".to_string(),
values: vec![Scalar::Int(99), Scalar::Null],
};
assert!(eval_rids(&miss, &b).unwrap().is_empty());
}
#[test]
fn mixed_numeric_compares_by_value_not_representation() {
assert_ne!(Scalar::Int(1), Scalar::Float(1.0));
assert_eq!(cmp_num(Num::I(1), Num::F(1.0)), Some(Ordering::Equal));
assert_eq!(cmp_num(Num::U(1), Num::F(1.5)), Some(Ordering::Less));
assert_eq!(cmp_num(Num::I(-1), Num::U(0)), Some(Ordering::Less));
}
#[test]
fn u64_sentinels_survive_the_int_domain() {
assert_eq!(
cmp_num(Num::U(u64::MAX), Num::I(-1)),
Some(Ordering::Greater)
);
assert_eq!(
FilterExpr::version_at_most(u64::MAX).to_sql().unwrap(),
"_version <= 18446744073709551615"
);
}
#[test]
fn large_integers_compare_exactly_against_floats() {
let big = (1u64 << 53) + 1;
assert_eq!(
cmp_num(Num::U(big), Num::F((1u64 << 53) as f64)),
Some(Ordering::Greater)
);
assert_eq!(
cmp_num(Num::I(i64::MAX), Num::F(f64::MAX)),
Some(Ordering::Less)
);
assert_eq!(
cmp_num(Num::I(i64::MIN), Num::F(f64::NEG_INFINITY)),
Some(Ordering::Greater)
);
}
#[test]
fn nan_is_unordered_so_every_comparison_is_unknown() {
assert_eq!(cmp_num(Num::F(f64::NAN), Num::F(1.0)), None);
assert_eq!(cmp_num(Num::I(1), Num::F(f64::NAN)), None);
let unknown = compare(&Cell::Float(f64::NAN), CmpOp::Eq, &Scalar::Float(f64::NAN));
assert_eq!(unknown.unwrap(), None);
}
#[test]
fn raw_and_cross_domain_refuse_loudly() {
let b = batch();
let raw = FilterExpr::Raw("i > 0".to_string());
assert!(matches!(
raw.eval(&ArrowRow::new(&b, 0)),
Err(EvalError::Unsupported(_))
));
let crossed = FilterExpr::Compare {
column: "i".to_string(),
op: CmpOp::Eq,
value: Scalar::Str("a".to_string()),
};
assert!(matches!(
crossed.eval(&ArrowRow::new(&b, 0)),
Err(EvalError::Unsupported(_))
));
let missing = FilterExpr::Compare {
column: "nope".to_string(),
op: CmpOp::Eq,
value: Scalar::Int(1),
};
assert!(matches!(
missing.eval(&ArrowRow::new(&b, 0)),
Err(EvalError::UnknownColumn(_))
));
}
#[test]
fn apostrophes_round_trip_through_the_single_escaper() {
let expr = FilterExpr::Compare {
column: "s".to_string(),
op: CmpOp::Eq,
value: Scalar::Str("it's".to_string()),
};
assert_eq!(expr.to_sql().unwrap(), "s = 'it''s'");
assert!(!eval_rids(&expr, &batch()).unwrap().is_empty());
}
#[test]
fn wildcard_pattern_refuses_in_sql_but_evaluates_exactly() {
let b = batch();
let expr = FilterExpr::StringMatch {
column: "s".to_string(),
kind: StringMatchKind::Contains,
pattern: "%b_".to_string(),
};
assert!(matches!(expr.to_sql(), Err(ToSqlError::Unsupported(_))));
let kept = eval_rids(&expr, &b).unwrap();
assert!(
!kept.is_empty(),
"the literal substring `%b_` occurs in `a%b_c` rows"
);
let s = b.column(4).as_string::<i32>();
for r in 0..b.num_rows() {
let expect = !s.is_null(r) && s.value(r).contains("%b_");
assert_eq!(kept.contains(&(r as u64)), expect, "row {r}");
}
}
#[test]
fn sql_pushable_drops_an_or_whole_never_one_branch() {
let bad = FilterExpr::StringMatch {
column: "s".to_string(),
kind: StringMatchKind::Contains,
pattern: "%".to_string(),
};
let good = FilterExpr::equals("i", Scalar::Int(1));
let anded = FilterExpr::And(vec![good.clone(), bad.clone()]);
assert_eq!(anded.sql_pushable(), good);
let ored = FilterExpr::Or(vec![good.clone(), bad.clone()]);
assert_eq!(ored.sql_pushable(), FilterExpr::Literal(true));
let nested = FilterExpr::And(vec![good.clone(), ored]);
assert_eq!(nested.sql_pushable(), good);
let fine = FilterExpr::And(vec![good.clone(), FilterExpr::IsNull("s".to_string())]);
assert_eq!(fine.sql_pushable(), fine);
}
#[test]
fn scalar_from_value_excludes_null_so_producers_keep_bailing() {
use uni_common::Value;
assert_eq!(
Scalar::from_value(&Value::String("x".into())),
Some(Scalar::Str("x".into()))
);
assert_eq!(Scalar::from_value(&Value::Int(1)), Some(Scalar::Int(1)));
assert_eq!(
Scalar::from_value(&Value::Bool(true)),
Some(Scalar::Bool(true))
);
assert_eq!(Scalar::from_value(&Value::Null), None);
assert_eq!(Scalar::from_value(&Value::List(vec![])), None);
}
#[test]
fn rendered_compare_narrows_a_real_scan() {
let rt = tokio::runtime::Runtime::new().unwrap();
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().join("c.lance").to_string_lossy().to_string();
let schema = Arc::new(Schema::new(vec![
Field::new("rid", arrow_schema::DataType::UInt64, false),
Field::new("createdAt", arrow_schema::DataType::Int64, true),
]));
let b = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(UInt64Array::from_iter_values(0..5u64)),
Arc::new(Int64Array::from_iter_values(1..6i64)),
],
)
.unwrap();
rt.block_on(async {
lance::Dataset::write(RecordBatchIterator::new(vec![Ok(b)], schema), &uri, None)
.await
.unwrap();
});
let ds = rt.block_on(lance::Dataset::open(&uri)).unwrap();
let range = FilterExpr::all([
FilterExpr::compare("createdAt", CmpOp::GtEq, Scalar::Int(2)),
FilterExpr::compare("createdAt", CmpOp::LtEq, Scalar::Int(4)),
]);
let sql = range.to_sql().unwrap();
assert!(!sql.contains('"'), "column must be bare: {sql}");
assert_eq!(rt.block_on(sql_rids(&ds, &sql)).unwrap(), vec![1, 2, 3]);
let quoted = "\"createdAt\" >= 2 AND \"createdAt\" <= 4";
assert_eq!(
rt.block_on(sql_rids(&ds, quoted)).unwrap(),
Vec::<u64>::new(),
"a double-quoted column is a string literal, not an identifier"
);
}
}