1use std::borrow::Cow;
7use std::collections::{BTreeSet, VecDeque};
8use std::sync::Arc;
9
10use crate::exec::{LanceExecutionOptions, get_session_context};
11use crate::expr::safe_coerce_scalar;
12use crate::logical_expr::{coerce_filter_type_to_boolean, get_as_string_scalar_opt, resolve_expr};
13use crate::sql::{parse_sql_expr, parse_sql_filter};
14use arrow::compute::CastOptions;
15use arrow_array::ListArray;
16use arrow_buffer::OffsetBuffer;
17use arrow_cast::cast_with_options;
18use arrow_schema::{DataType as ArrowDataType, Field, SchemaRef, TimeUnit};
19use arrow_select::concat::concat;
20use datafusion::common::DFSchema;
21use datafusion::common::tree_node::{TreeNode, TreeNodeRecursion, TreeNodeVisitor};
22use datafusion::config::ConfigOptions;
23use datafusion::error::Result as DFResult;
24use datafusion::execution::context::SessionState;
25use datafusion::logical_expr::expr::ScalarFunction;
26use datafusion::logical_expr::planner::{ExprPlanner, PlannerResult, RawFieldAccessExpr};
27use datafusion::logical_expr::{
28 AggregateUDF, ColumnarValue, GetFieldAccess, ScalarFunctionArgs, ScalarUDF, ScalarUDFImpl,
29 Signature, Volatility, WindowUDF,
30};
31use datafusion::optimizer::simplify_expressions::SimplifyContext;
32use datafusion::sql::planner::{
33 ContextProvider, NullOrdering, ParserOptions, PlannerContext, SqlToRel,
34};
35use datafusion::sql::sqlparser::ast::{
36 AccessExpr, Array as SQLArray, BinaryOperator, DataType as SQLDataType, ExactNumberInfo,
37 Expr as SQLExpr, Function, FunctionArg, FunctionArgExpr, FunctionArguments, Ident,
38 ObjectNamePart, Subscript, TimezoneInfo, TypedString, UnaryOperator, Value, ValueWithSpan,
39};
40use datafusion::{
41 common::Column,
42 logical_expr::{Between, BinaryExpr, Like, Operator},
43 physical_plan::PhysicalExpr,
44 prelude::Expr,
45 scalar::ScalarValue,
46};
47use datafusion_functions::core::getfield::GetFieldFunc;
48use lance_core::datatypes::Schema;
49use lance_core::error::LanceOptionExt;
50
51use chrono::Utc;
52use lance_core::{Error, Result};
53
54fn encode_jsonb(json_str: &str) -> Result<Expr> {
56 let bytes = lance_arrow::json::encode_json(json_str)
57 .map_err(|e| Error::invalid_input(format!("Failed to encode JSONB: {e}")))?;
58 Ok(Expr::Literal(ScalarValue::LargeBinary(Some(bytes)), None))
59}
60
61fn parse_like_escape_char(escape_char: &Option<ValueWithSpan>) -> Result<Option<char>> {
65 let Some(value) = escape_char else {
66 return Ok(None);
67 };
68 let ValueWithSpan {
69 value: Value::SingleQuotedString(escape),
70 ..
71 } = value
72 else {
73 return Err(Error::invalid_input(format!(
74 "Invalid escape character in LIKE expression. Expected a single character wrapped with single quotes, got {value}"
75 )));
76 };
77 let mut chars = escape.chars();
78 match (chars.next(), chars.next()) {
79 (Some(c), None) => Ok(Some(c)),
80 _ => Err(Error::invalid_input(format!(
81 "Invalid escape character in LIKE expression. Expected a single character, got '{escape}'"
82 ))),
83 }
84}
85
86#[derive(Debug, Clone, Eq, PartialEq, Hash)]
87struct CastListF16Udf {
88 signature: Signature,
89}
90
91impl CastListF16Udf {
92 pub fn new() -> Self {
93 Self {
94 signature: Signature::any(1, Volatility::Immutable),
95 }
96 }
97}
98
99impl ScalarUDFImpl for CastListF16Udf {
100 fn name(&self) -> &str {
101 "_cast_list_f16"
102 }
103
104 fn signature(&self) -> &Signature {
105 &self.signature
106 }
107
108 fn return_type(&self, arg_types: &[ArrowDataType]) -> DFResult<ArrowDataType> {
109 let input = &arg_types[0];
110 match input {
111 ArrowDataType::FixedSizeList(field, size) => {
112 if field.data_type() != &ArrowDataType::Float32
113 && field.data_type() != &ArrowDataType::Float16
114 {
115 return Err(datafusion::error::DataFusionError::Execution(
116 "cast_list_f16 only supports list of float32 or float16".to_string(),
117 ));
118 }
119 Ok(ArrowDataType::FixedSizeList(
120 Arc::new(Field::new(
121 field.name(),
122 ArrowDataType::Float16,
123 field.is_nullable(),
124 )),
125 *size,
126 ))
127 }
128 ArrowDataType::List(field) => {
129 if field.data_type() != &ArrowDataType::Float32
130 && field.data_type() != &ArrowDataType::Float16
131 {
132 return Err(datafusion::error::DataFusionError::Execution(
133 "cast_list_f16 only supports list of float32 or float16".to_string(),
134 ));
135 }
136 Ok(ArrowDataType::List(Arc::new(Field::new(
137 field.name(),
138 ArrowDataType::Float16,
139 field.is_nullable(),
140 ))))
141 }
142 _ => Err(datafusion::error::DataFusionError::Execution(
143 "cast_list_f16 only supports FixedSizeList/List arguments".to_string(),
144 )),
145 }
146 }
147
148 fn invoke_with_args(&self, func_args: ScalarFunctionArgs) -> DFResult<ColumnarValue> {
149 let ColumnarValue::Array(arr) = &func_args.args[0] else {
150 return Err(datafusion::error::DataFusionError::Execution(
151 "cast_list_f16 only supports array arguments".to_string(),
152 ));
153 };
154
155 let to_type = match arr.data_type() {
156 ArrowDataType::FixedSizeList(field, size) => ArrowDataType::FixedSizeList(
157 Arc::new(Field::new(
158 field.name(),
159 ArrowDataType::Float16,
160 field.is_nullable(),
161 )),
162 *size,
163 ),
164 ArrowDataType::List(field) => ArrowDataType::List(Arc::new(Field::new(
165 field.name(),
166 ArrowDataType::Float16,
167 field.is_nullable(),
168 ))),
169 _ => {
170 return Err(datafusion::error::DataFusionError::Execution(
171 "cast_list_f16 only supports array arguments".to_string(),
172 ));
173 }
174 };
175
176 let res = cast_with_options(arr.as_ref(), &to_type, &CastOptions::default())?;
177 Ok(ColumnarValue::Array(res))
178 }
179}
180
181struct LanceContextProvider {
183 options: datafusion::config::ConfigOptions,
184 state: SessionState,
185 expr_planners: Vec<Arc<dyn ExprPlanner>>,
186}
187
188impl Default for LanceContextProvider {
189 fn default() -> Self {
190 let ctx = get_session_context(&LanceExecutionOptions::default());
191 let state = ctx.state();
192 let expr_planners = state.expr_planners().to_vec();
193
194 Self {
195 options: ConfigOptions::default(),
196 state,
197 expr_planners,
198 }
199 }
200}
201
202impl ContextProvider for LanceContextProvider {
203 fn get_table_source(
204 &self,
205 name: datafusion::sql::TableReference,
206 ) -> DFResult<Arc<dyn datafusion::logical_expr::TableSource>> {
207 Err(datafusion::error::DataFusionError::NotImplemented(format!(
208 "Attempt to reference inner table {} not supported",
209 name
210 )))
211 }
212
213 fn get_aggregate_meta(&self, name: &str) -> Option<Arc<AggregateUDF>> {
214 self.state.aggregate_functions().get(name).cloned()
215 }
216
217 fn get_window_meta(&self, name: &str) -> Option<Arc<WindowUDF>> {
218 self.state.window_functions().get(name).cloned()
219 }
220
221 fn get_higher_order_meta(
222 &self,
223 name: &str,
224 ) -> Option<Arc<datafusion::logical_expr::HigherOrderUDF>> {
225 self.state.higher_order_functions().get(name).cloned()
226 }
227
228 fn get_function_meta(&self, f: &str) -> Option<Arc<ScalarUDF>> {
229 match f {
230 "_cast_list_f16" => Some(Arc::new(ScalarUDF::new_from_impl(CastListF16Udf::new()))),
233 _ => self.state.scalar_functions().get(f).cloned(),
234 }
235 }
236
237 fn get_variable_type(&self, _: &[String]) -> Option<ArrowDataType> {
238 None
240 }
241
242 fn options(&self) -> &datafusion::config::ConfigOptions {
243 &self.options
244 }
245
246 fn udf_names(&self) -> Vec<String> {
247 self.state.scalar_functions().keys().cloned().collect()
248 }
249
250 fn udaf_names(&self) -> Vec<String> {
251 self.state.aggregate_functions().keys().cloned().collect()
252 }
253
254 fn udwf_names(&self) -> Vec<String> {
255 self.state.window_functions().keys().cloned().collect()
256 }
257
258 fn higher_order_function_names(&self) -> Vec<String> {
259 self.state
260 .higher_order_functions()
261 .keys()
262 .cloned()
263 .collect()
264 }
265
266 fn get_expr_planners(&self) -> &[Arc<dyn ExprPlanner>] {
267 &self.expr_planners
268 }
269}
270
271pub struct Planner {
272 schema: SchemaRef,
273 context_provider: LanceContextProvider,
274 enable_relations: bool,
275}
276
277impl Planner {
278 pub fn new(schema: SchemaRef) -> Self {
279 Self {
280 schema,
281 context_provider: LanceContextProvider::default(),
282 enable_relations: false,
283 }
284 }
285
286 pub fn with_enable_relations(mut self, enable_relations: bool) -> Self {
292 self.enable_relations = enable_relations;
293 self
294 }
295
296 fn resolve_column_name(&self, name: &str) -> String {
299 if self.schema.field_with_name(name).is_ok() {
301 return name.to_string();
302 }
303 for field in self.schema.fields() {
305 if field.name().eq_ignore_ascii_case(name) {
306 return field.name().clone();
307 }
308 }
309 name.to_string()
311 }
312
313 fn column(&self, idents: &[Ident]) -> Expr {
314 fn handle_remaining_idents(expr: &mut Expr, idents: &[Ident]) {
315 for ident in idents {
316 *expr = Expr::ScalarFunction(ScalarFunction {
317 args: vec![
318 std::mem::take(expr),
319 Expr::Literal(ScalarValue::Utf8(Some(ident.value.clone())), None),
320 ],
321 func: Arc::new(ScalarUDF::new_from_impl(GetFieldFunc::default())),
322 });
323 }
324 }
325
326 if self.enable_relations && idents.len() > 1 {
327 let relation = &idents[0].value;
329 let column_name = self.resolve_column_name(&idents[1].value);
330 let column = Expr::Column(Column::new(Some(relation.clone()), column_name));
331 let mut result = column;
332 handle_remaining_idents(&mut result, &idents[2..]);
333 result
334 } else {
335 let resolved_name = self.resolve_column_name(&idents[0].value);
338 let mut column = Expr::Column(Column::from_name(resolved_name));
339 handle_remaining_idents(&mut column, &idents[1..]);
340 column
341 }
342 }
343
344 fn binary_op(&self, op: &BinaryOperator) -> Result<Operator> {
345 Ok(match op {
346 BinaryOperator::Plus => Operator::Plus,
347 BinaryOperator::Minus => Operator::Minus,
348 BinaryOperator::Multiply => Operator::Multiply,
349 BinaryOperator::Divide => Operator::Divide,
350 BinaryOperator::Modulo => Operator::Modulo,
351 BinaryOperator::StringConcat => Operator::StringConcat,
352 BinaryOperator::Gt => Operator::Gt,
353 BinaryOperator::Lt => Operator::Lt,
354 BinaryOperator::GtEq => Operator::GtEq,
355 BinaryOperator::LtEq => Operator::LtEq,
356 BinaryOperator::Eq => Operator::Eq,
357 BinaryOperator::NotEq => Operator::NotEq,
358 BinaryOperator::And => Operator::And,
359 BinaryOperator::Or => Operator::Or,
360 _ => {
361 return Err(Error::invalid_input(format!(
362 "Operator {op} is not supported"
363 )));
364 }
365 })
366 }
367
368 fn is_logical_binary_op(op: &BinaryOperator) -> bool {
369 matches!(op, BinaryOperator::And | BinaryOperator::Or)
370 }
371
372 fn is_same_logical_binary_op(left: &BinaryOperator, right: &BinaryOperator) -> bool {
373 matches!(
374 (left, right),
375 (BinaryOperator::And, BinaryOperator::And) | (BinaryOperator::Or, BinaryOperator::Or)
376 )
377 }
378
379 fn flatten_logical_binary_exprs<'a>(
380 left: &'a SQLExpr,
381 op: &BinaryOperator,
382 right: &'a SQLExpr,
383 ) -> Vec<&'a SQLExpr> {
384 let mut leaves = Vec::new();
385 let mut stack = vec![right, left];
386
387 while let Some(expr) = stack.pop() {
388 match expr {
389 SQLExpr::BinaryOp {
390 left,
391 op: child_op,
392 right,
393 } if Self::is_same_logical_binary_op(op, child_op) => {
394 stack.push(right.as_ref());
395 stack.push(left.as_ref());
396 }
397 _ => leaves.push(expr),
398 }
399 }
400
401 leaves
402 }
403
404 fn balanced_binary_expr(mut exprs: VecDeque<Expr>, op: Operator) -> Result<Expr> {
405 if exprs.is_empty() {
406 return Err(Error::invalid_input("Binary expression has no operands"));
407 }
408
409 while exprs.len() > 1 {
410 let mut next = VecDeque::with_capacity(exprs.len().div_ceil(2));
411 while let Some(left) = exprs.pop_front() {
412 if let Some(right) = exprs.pop_front() {
413 next.push_back(Expr::BinaryExpr(BinaryExpr::new(
414 Box::new(left),
415 op,
416 Box::new(right),
417 )));
418 } else {
419 next.push_back(left);
420 }
421 }
422 exprs = next;
423 }
424
425 exprs
426 .pop_front()
427 .ok_or_else(|| Error::invalid_input("Binary expression has no operands"))
428 }
429
430 fn binary_expr(&self, left: &SQLExpr, op: &BinaryOperator, right: &SQLExpr) -> Result<Expr> {
431 let df_op = self.binary_op(op)?;
432 if Self::is_logical_binary_op(op) {
433 let leaves = Self::flatten_logical_binary_exprs(left, op, right);
434 let mut exprs = VecDeque::with_capacity(leaves.len());
435 for leaf in leaves {
436 exprs.push_back(self.parse_sql_expr(leaf)?);
437 }
438 return Self::balanced_binary_expr(exprs, df_op);
439 }
440
441 Ok(Expr::BinaryExpr(BinaryExpr::new(
442 Box::new(self.parse_sql_expr(left)?),
443 df_op,
444 Box::new(self.parse_sql_expr(right)?),
445 )))
446 }
447
448 fn unary_expr(&self, op: &UnaryOperator, expr: &SQLExpr) -> Result<Expr> {
449 Ok(match op {
450 UnaryOperator::Not | UnaryOperator::BitwiseNot => {
451 Expr::Not(Box::new(self.parse_sql_expr(expr)?))
452 }
453
454 UnaryOperator::Minus => {
455 use datafusion::logical_expr::lit;
456 match expr {
457 SQLExpr::Value(ValueWithSpan { value: Value::Number(n, _), ..}) => match n.parse::<i64>() {
458 Ok(n) => lit(-n),
459 Err(_) => lit(-n
460 .parse::<f64>()
461 .map_err(|_e| {
462 Error::invalid_input(format!("negative operator can be only applied to integer and float operands, got: {n}"))
463 })?),
464 },
465 _ => {
466 Expr::Negative(Box::new(self.parse_sql_expr(expr)?))
467 }
468 }
469 }
470
471 _ => {
472 return Err(Error::invalid_input(format!(
473 "Unary operator '{:?}' is not supported",
474 op
475 )));
476 }
477 })
478 }
479
480 fn number(&self, value: &str, negative: bool) -> Result<Expr> {
482 use datafusion::logical_expr::lit;
483 let value: Cow<str> = if negative {
484 Cow::Owned(format!("-{}", value))
485 } else {
486 Cow::Borrowed(value)
487 };
488 if let Ok(n) = value.parse::<i64>() {
489 Ok(lit(n))
490 } else {
491 value.parse::<f64>().map(lit).map_err(|_| {
492 Error::invalid_input(format!("'{value}' is not supported number value."))
493 })
494 }
495 }
496
497 fn value(&self, value: &Value) -> Result<Expr> {
498 Ok(match value {
499 Value::Number(v, _) => self.number(v.as_str(), false)?,
500 Value::SingleQuotedString(s) => Expr::Literal(ScalarValue::Utf8(Some(s.clone())), None),
501 Value::HexStringLiteral(hsl) => {
502 Expr::Literal(ScalarValue::Binary(Self::try_decode_hex_literal(hsl)), None)
503 }
504 Value::DoubleQuotedString(s) => Expr::Literal(ScalarValue::Utf8(Some(s.clone())), None),
505 Value::Boolean(v) => Expr::Literal(ScalarValue::Boolean(Some(*v)), None),
506 Value::Null => Expr::Literal(ScalarValue::Null, None),
507 _ => todo!(),
508 })
509 }
510
511 fn parse_function_args(&self, func_args: &FunctionArg) -> Result<Expr> {
512 match func_args {
513 FunctionArg::Unnamed(FunctionArgExpr::Expr(expr)) => self.parse_sql_expr(expr),
514 _ => Err(Error::invalid_input(format!(
515 "Unsupported function args: {:?}",
516 func_args
517 ))),
518 }
519 }
520
521 fn legacy_parse_function(&self, func: &Function) -> Result<Expr> {
528 match &func.args {
529 FunctionArguments::List(args) => {
530 if func.name.0.len() != 1 {
531 return Err(Error::invalid_input(format!(
532 "Function name must have 1 part, got: {:?}",
533 func.name.0
534 )));
535 }
536 Ok(Expr::IsNotNull(Box::new(
537 self.parse_function_args(&args.args[0])?,
538 )))
539 }
540 _ => Err(Error::invalid_input(format!(
541 "Unsupported function args: {:?}",
542 func.args
543 ))),
544 }
545 }
546
547 fn parse_function(&self, function: SQLExpr) -> Result<Expr> {
548 if let SQLExpr::Function(function) = &function
549 && let Some(ObjectNamePart::Identifier(name)) = &function.name.0.first()
550 && &name.value == "is_valid"
551 {
552 return self.legacy_parse_function(function);
553 }
554 let sql_to_rel = SqlToRel::new_with_options(
555 &self.context_provider,
556 ParserOptions {
557 parse_float_as_decimal: false,
558 enable_ident_normalization: false,
559 support_varchar_with_length: false,
560 enable_options_value_normalization: false,
561 collect_spans: false,
562 map_string_types_to_utf8view: false,
563 default_null_ordering: NullOrdering::NullsMax,
564 },
565 );
566
567 let mut planner_context = PlannerContext::default();
568 let schema = DFSchema::try_from(self.schema.as_ref().clone())?;
569 sql_to_rel
570 .sql_to_expr(function, &schema, &mut planner_context)
571 .map_err(|e| Error::invalid_input(format!("Error parsing function: {e}")))
572 }
573
574 fn parse_type(&self, data_type: &SQLDataType) -> Result<ArrowDataType> {
575 const SUPPORTED_TYPES: [&str; 13] = [
576 "int [unsigned]",
577 "tinyint [unsigned]",
578 "smallint [unsigned]",
579 "bigint [unsigned]",
580 "float",
581 "double",
582 "string",
583 "binary",
584 "date",
585 "timestamp(precision)",
586 "datetime(precision)",
587 "decimal(precision,scale)",
588 "boolean",
589 ];
590 match data_type {
591 SQLDataType::String(_) => Ok(ArrowDataType::Utf8),
592 SQLDataType::Binary(_) => Ok(ArrowDataType::Binary),
593 SQLDataType::Float(_) => Ok(ArrowDataType::Float32),
594 SQLDataType::Double(_) => Ok(ArrowDataType::Float64),
595 SQLDataType::Boolean => Ok(ArrowDataType::Boolean),
596 SQLDataType::TinyInt(_) => Ok(ArrowDataType::Int8),
597 SQLDataType::SmallInt(_) => Ok(ArrowDataType::Int16),
598 SQLDataType::Int(_) | SQLDataType::Integer(_) => Ok(ArrowDataType::Int32),
599 SQLDataType::BigInt(_) => Ok(ArrowDataType::Int64),
600 SQLDataType::TinyIntUnsigned(_) => Ok(ArrowDataType::UInt8),
601 SQLDataType::SmallIntUnsigned(_) => Ok(ArrowDataType::UInt16),
602 SQLDataType::IntUnsigned(_) | SQLDataType::IntegerUnsigned(_) => {
603 Ok(ArrowDataType::UInt32)
604 }
605 SQLDataType::BigIntUnsigned(_) => Ok(ArrowDataType::UInt64),
606 SQLDataType::Date => Ok(ArrowDataType::Date32),
607 SQLDataType::Timestamp(resolution, tz) => {
608 match tz {
609 TimezoneInfo::None => {}
610 _ => {
611 return Err(Error::invalid_input(
612 "Timezone not supported in timestamp".to_string(),
613 ));
614 }
615 };
616 let time_unit = match resolution {
617 None => TimeUnit::Microsecond,
619 Some(0) => TimeUnit::Second,
620 Some(3) => TimeUnit::Millisecond,
621 Some(6) => TimeUnit::Microsecond,
622 Some(9) => TimeUnit::Nanosecond,
623 _ => {
624 return Err(Error::invalid_input(format!(
625 "Unsupported datetime resolution: {:?}",
626 resolution
627 )));
628 }
629 };
630 Ok(ArrowDataType::Timestamp(time_unit, None))
631 }
632 SQLDataType::Datetime(resolution) => {
633 let time_unit = match resolution {
634 None => TimeUnit::Microsecond,
635 Some(0) => TimeUnit::Second,
636 Some(3) => TimeUnit::Millisecond,
637 Some(6) => TimeUnit::Microsecond,
638 Some(9) => TimeUnit::Nanosecond,
639 _ => {
640 return Err(Error::invalid_input(format!(
641 "Unsupported datetime resolution: {:?}",
642 resolution
643 )));
644 }
645 };
646 Ok(ArrowDataType::Timestamp(time_unit, None))
647 }
648 SQLDataType::Decimal(number_info) => match number_info {
649 ExactNumberInfo::PrecisionAndScale(precision, scale) => {
650 Ok(ArrowDataType::Decimal128(*precision as u8, *scale as i8))
651 }
652 _ => Err(Error::invalid_input(format!(
653 "Must provide precision and scale for decimal: {:?}",
654 number_info
655 ))),
656 },
657 _ => Err(Error::invalid_input(format!(
658 "Unsupported data type: {:?}. Supported types: {:?}",
659 data_type, SUPPORTED_TYPES
660 ))),
661 }
662 }
663
664 fn plan_field_access(&self, mut field_access_expr: RawFieldAccessExpr) -> Result<Expr> {
665 let df_schema = DFSchema::try_from(self.schema.as_ref().clone())?;
666 for planner in self.context_provider.get_expr_planners() {
667 match planner.plan_field_access(field_access_expr, &df_schema)? {
668 PlannerResult::Planned(expr) => return Ok(expr),
669 PlannerResult::Original(expr) => {
670 field_access_expr = expr;
671 }
672 }
673 }
674 Err(Error::invalid_input("Field access could not be planned"))
675 }
676
677 fn parse_sql_expr(&self, expr: &SQLExpr) -> Result<Expr> {
678 match expr {
679 SQLExpr::Identifier(id) => {
680 if id.quote_style == Some('"') {
683 Ok(Expr::Literal(
684 ScalarValue::Utf8(Some(id.value.clone())),
685 None,
686 ))
687 } else if id.quote_style == Some('`') {
690 Ok(Expr::Column(Column::from_name(id.value.clone())))
691 } else {
692 Ok(self.column(vec![id.clone()].as_slice()))
693 }
694 }
695 SQLExpr::CompoundIdentifier(ids) => Ok(self.column(ids.as_slice())),
696 SQLExpr::BinaryOp { left, op, right } => self.binary_expr(left, op, right),
697 SQLExpr::UnaryOp { op, expr } => self.unary_expr(op, expr),
698 SQLExpr::Value(value) => self.value(&value.value),
699 SQLExpr::Array(SQLArray { elem, .. }) => {
700 let mut values = vec![];
701
702 let array_literal_error = |pos: usize, value: &_| {
703 Err(Error::invalid_input(format!(
704 "Expected a literal value in array, instead got {} at position {}",
705 value, pos
706 )))
707 };
708
709 for (pos, expr) in elem.iter().enumerate() {
710 match expr {
711 SQLExpr::Value(value) => {
712 if let Expr::Literal(value, _) = self.value(&value.value)? {
713 values.push(value);
714 } else {
715 return array_literal_error(pos, expr);
716 }
717 }
718 SQLExpr::UnaryOp {
719 op: UnaryOperator::Minus,
720 expr,
721 } => {
722 if let SQLExpr::Value(ValueWithSpan {
723 value: Value::Number(number, _),
724 ..
725 }) = expr.as_ref()
726 {
727 if let Expr::Literal(value, _) = self.number(number, true)? {
728 values.push(value);
729 } else {
730 return array_literal_error(pos, expr);
731 }
732 } else {
733 return array_literal_error(pos, expr);
734 }
735 }
736 _ => {
737 return array_literal_error(pos, expr);
738 }
739 }
740 }
741
742 let field = if !values.is_empty() {
743 let data_type = values[0].data_type();
744
745 for value in &mut values {
746 if value.data_type() != data_type {
747 *value = safe_coerce_scalar(value, &data_type).ok_or_else(|| Error::invalid_input(format!("Array expressions must have a consistent datatype. Expected: {}, got: {}", data_type, value.data_type())))?;
748 }
749 }
750 Field::new("item", data_type, true)
751 } else {
752 Field::new("item", ArrowDataType::Null, true)
753 };
754
755 let values = values
756 .into_iter()
757 .map(|v| v.to_array().map_err(Error::from))
758 .collect::<Result<Vec<_>>>()?;
759 let array_refs = values.iter().map(|v| v.as_ref()).collect::<Vec<_>>();
760 let values = concat(&array_refs)?;
761 let values = ListArray::try_new(
762 field.into(),
763 OffsetBuffer::from_lengths([values.len()]),
764 values,
765 None,
766 )?;
767
768 Ok(Expr::Literal(ScalarValue::List(Arc::new(values)), None))
769 }
770 SQLExpr::TypedString(TypedString {
772 data_type: SQLDataType::JSONB,
773 value,
774 ..
775 }) => match &value.value {
776 Value::SingleQuotedString(s) | Value::DoubleQuotedString(s) => encode_jsonb(s),
777 _ => Err(Error::invalid_input(
778 "Expected a string value for JSONB literal",
779 )),
780 },
781 SQLExpr::TypedString(TypedString {
783 data_type, value, ..
784 }) => {
785 let value = value.clone().into_string().expect_ok()?;
786 Ok(Expr::Cast(datafusion::logical_expr::Cast::new(
787 Box::new(Expr::Literal(ScalarValue::Utf8(Some(value)), None)),
788 self.parse_type(data_type)?,
789 )))
790 }
791 SQLExpr::IsFalse(expr) => Ok(Expr::IsFalse(Box::new(self.parse_sql_expr(expr)?))),
792 SQLExpr::IsNotFalse(expr) => Ok(Expr::IsNotFalse(Box::new(self.parse_sql_expr(expr)?))),
793 SQLExpr::IsTrue(expr) => Ok(Expr::IsTrue(Box::new(self.parse_sql_expr(expr)?))),
794 SQLExpr::IsNotTrue(expr) => Ok(Expr::IsNotTrue(Box::new(self.parse_sql_expr(expr)?))),
795 SQLExpr::IsNull(expr) => Ok(Expr::IsNull(Box::new(self.parse_sql_expr(expr)?))),
796 SQLExpr::IsNotNull(expr) => Ok(Expr::IsNotNull(Box::new(self.parse_sql_expr(expr)?))),
797 SQLExpr::InList {
798 expr,
799 list,
800 negated,
801 } => {
802 let value_expr = self.parse_sql_expr(expr)?;
803 let list_exprs = list
804 .iter()
805 .map(|e| self.parse_sql_expr(e))
806 .collect::<Result<Vec<_>>>()?;
807 Ok(value_expr.in_list(list_exprs, *negated))
808 }
809 SQLExpr::Nested(inner) => self.parse_sql_expr(inner.as_ref()),
810 SQLExpr::Function(_) => self.parse_function(expr.clone()),
811 SQLExpr::ILike {
812 negated,
813 expr,
814 pattern,
815 escape_char,
816 any: _,
817 } => Ok(Expr::Like(Like::new(
818 *negated,
819 Box::new(self.parse_sql_expr(expr)?),
820 Box::new(self.parse_sql_expr(pattern)?),
821 parse_like_escape_char(escape_char)?,
822 true,
823 ))),
824 SQLExpr::Like {
825 negated,
826 expr,
827 pattern,
828 escape_char,
829 any: _,
830 } => Ok(Expr::Like(Like::new(
831 *negated,
832 Box::new(self.parse_sql_expr(expr)?),
833 Box::new(self.parse_sql_expr(pattern)?),
834 parse_like_escape_char(escape_char)?,
835 false,
836 ))),
837 SQLExpr::Cast {
839 data_type: SQLDataType::JSONB,
840 expr: inner,
841 ..
842 } => match inner.as_ref() {
843 SQLExpr::Value(ValueWithSpan {
844 value: Value::SingleQuotedString(s) | Value::DoubleQuotedString(s),
845 ..
846 }) => encode_jsonb(s),
847 _ => Err(Error::invalid_input(
848 "CAST to JSONB only supports string literals",
849 )),
850 },
851 SQLExpr::Cast {
852 expr,
853 data_type,
854 kind,
855 ..
856 } => match kind {
857 datafusion::sql::sqlparser::ast::CastKind::TryCast
858 | datafusion::sql::sqlparser::ast::CastKind::SafeCast => {
859 Ok(Expr::TryCast(datafusion::logical_expr::TryCast::new(
860 Box::new(self.parse_sql_expr(expr)?),
861 self.parse_type(data_type)?,
862 )))
863 }
864 _ => Ok(Expr::Cast(datafusion::logical_expr::Cast::new(
865 Box::new(self.parse_sql_expr(expr)?),
866 self.parse_type(data_type)?,
867 ))),
868 },
869 SQLExpr::JsonAccess { .. } => Err(Error::invalid_input("JSON access is not supported")),
870 SQLExpr::CompoundFieldAccess { root, access_chain } => {
871 let mut expr = self.parse_sql_expr(root)?;
872
873 for access in access_chain {
874 let field_access = match access {
875 AccessExpr::Dot(SQLExpr::Identifier(Ident { value: s, .. }))
877 | AccessExpr::Subscript(Subscript::Index {
878 index:
879 SQLExpr::Value(ValueWithSpan {
880 value:
881 Value::SingleQuotedString(s) | Value::DoubleQuotedString(s),
882 ..
883 }),
884 }) => GetFieldAccess::NamedStructField {
885 name: ScalarValue::from(s.as_str()),
886 },
887 AccessExpr::Subscript(Subscript::Index { index }) => {
888 let key = Box::new(self.parse_sql_expr(index)?);
889 GetFieldAccess::ListIndex { key }
890 }
891 AccessExpr::Subscript(Subscript::Slice { .. }) => {
892 return Err(Error::invalid_input("Slice subscript is not supported"));
893 }
894 _ => {
895 return Err(Error::invalid_input(
898 "Only dot notation or index access is supported for field access",
899 ));
900 }
901 };
902
903 let field_access_expr = RawFieldAccessExpr { expr, field_access };
904 expr = self.plan_field_access(field_access_expr)?;
905 }
906
907 Ok(expr)
908 }
909 SQLExpr::Between {
910 expr,
911 negated,
912 low,
913 high,
914 } => {
915 let expr = self.parse_sql_expr(expr)?;
917 let low = self.parse_sql_expr(low)?;
918 let high = self.parse_sql_expr(high)?;
919
920 let between = Expr::Between(Between::new(
921 Box::new(expr),
922 *negated,
923 Box::new(low),
924 Box::new(high),
925 ));
926 Ok(between)
927 }
928 _ => Err(Error::invalid_input(format!(
929 "Expression '{expr}' is not supported SQL in lance"
930 ))),
931 }
932 }
933
934 pub fn parse_filter(&self, filter: &str) -> Result<Expr> {
939 let ast_expr = parse_sql_filter(filter)?;
941 let expr = self.parse_sql_expr(&ast_expr)?;
942 let schema = Schema::try_from(self.schema.as_ref())?;
943 let resolved = resolve_expr(&expr, &schema).map_err(|e| {
944 Error::invalid_input(format!("Error resolving filter expression {filter}: {e}"))
945 })?;
946
947 Ok(coerce_filter_type_to_boolean(resolved))
948 }
949
950 pub fn parse_expr(&self, expr: &str) -> Result<Expr> {
955 let resolved_name = self.resolve_column_name(expr);
958 if self.schema.field_with_name(&resolved_name).is_ok() {
959 return Ok(Expr::Column(Column::from_name(resolved_name)));
960 }
961
962 let ast_expr = parse_sql_expr(expr)?;
964 let expr = self.parse_sql_expr(&ast_expr)?;
965 let schema = Schema::try_from(self.schema.as_ref())?;
966 let resolved = resolve_expr(&expr, &schema)?;
967 Ok(resolved)
968 }
969
970 fn try_decode_hex_literal(s: &str) -> Option<Vec<u8>> {
976 let hex_bytes = s.as_bytes();
977 let mut decoded_bytes = Vec::with_capacity(hex_bytes.len().div_ceil(2));
978
979 let start_idx = hex_bytes.len() % 2;
980 if start_idx > 0 {
981 decoded_bytes.push(Self::try_decode_hex_char(hex_bytes[0])?);
983 }
984
985 for i in (start_idx..hex_bytes.len()).step_by(2) {
986 let high = Self::try_decode_hex_char(hex_bytes[i])?;
987 let low = Self::try_decode_hex_char(hex_bytes[i + 1])?;
988 decoded_bytes.push((high << 4) | low);
989 }
990
991 Some(decoded_bytes)
992 }
993
994 const fn try_decode_hex_char(c: u8) -> Option<u8> {
998 match c {
999 b'A'..=b'F' => Some(c - b'A' + 10),
1000 b'a'..=b'f' => Some(c - b'a' + 10),
1001 b'0'..=b'9' => Some(c - b'0'),
1002 _ => None,
1003 }
1004 }
1005
1006 pub fn optimize_expr(&self, expr: Expr) -> Result<Expr> {
1008 let df_schema = Arc::new(DFSchema::try_from(self.schema.as_ref().clone())?);
1009
1010 let simplify_context = SimplifyContext::builder()
1013 .with_schema(df_schema.clone())
1014 .with_query_execution_start_time(Some(Utc::now()))
1015 .build();
1016 let simplifier =
1017 datafusion::optimizer::simplify_expressions::ExprSimplifier::new(simplify_context);
1018
1019 let expr = simplifier.coerce(expr, &df_schema)?;
1021 let expr = simplifier.simplify(expr)?;
1022
1023 Ok(expr)
1024 }
1025
1026 pub fn create_physical_expr(&self, expr: &Expr) -> Result<Arc<dyn PhysicalExpr>> {
1028 let df_schema = Arc::new(DFSchema::try_from(self.schema.as_ref().clone())?);
1029 Ok(datafusion::physical_expr::create_physical_expr(
1030 expr,
1031 df_schema.as_ref(),
1032 &Default::default(),
1033 )?)
1034 }
1035
1036 pub fn column_names_in_expr(expr: &Expr) -> Vec<String> {
1043 let mut visitor = ColumnCapturingVisitor {
1044 current_path: VecDeque::new(),
1045 columns: BTreeSet::new(),
1046 };
1047 expr.visit(&mut visitor).unwrap();
1048 visitor.columns.into_iter().collect()
1049 }
1050}
1051
1052struct ColumnCapturingVisitor {
1053 current_path: VecDeque<String>,
1055 columns: BTreeSet<String>,
1056}
1057
1058impl TreeNodeVisitor<'_> for ColumnCapturingVisitor {
1059 type Node = Expr;
1060
1061 fn f_down(&mut self, node: &Self::Node) -> DFResult<TreeNodeRecursion> {
1062 match node {
1063 Expr::Column(Column { name, .. }) => {
1064 let mut path = name.clone();
1068 for part in self.current_path.drain(..) {
1069 path.push('.');
1070 if part.contains('.') || part.contains('`') {
1072 let escaped = part.replace('`', "``");
1074 path.push('`');
1075 path.push_str(&escaped);
1076 path.push('`');
1077 } else {
1078 path.push_str(&part);
1079 }
1080 }
1081 self.columns.insert(path);
1082 self.current_path.clear();
1083 }
1084 Expr::ScalarFunction(udf) if udf.name() == GetFieldFunc::default().name() => {
1085 if let Some(name) = get_as_string_scalar_opt(&udf.args[1]) {
1086 self.current_path.push_front(name.to_string())
1087 } else {
1088 self.current_path.clear();
1089 }
1090 }
1091 _ => {
1092 self.current_path.clear();
1093 }
1094 }
1095
1096 Ok(TreeNodeRecursion::Continue)
1097 }
1098}
1099
1100#[cfg(test)]
1101mod tests {
1102
1103 use crate::logical_expr::ExprExt;
1104
1105 use super::*;
1106
1107 use arrow::datatypes::Float64Type;
1108 use arrow_array::{
1109 ArrayRef, BooleanArray, Float32Array, Int32Array, Int64Array, RecordBatch, StringArray,
1110 StructArray, TimestampMicrosecondArray, TimestampMillisecondArray,
1111 TimestampNanosecondArray, TimestampSecondArray,
1112 };
1113 use arrow_schema::{DataType, Fields, Schema};
1114 use datafusion::{
1115 logical_expr::{Cast, col, lit},
1116 prelude::{array_element, get_field},
1117 };
1118 use datafusion_functions::core::expr_ext::FieldAccessor;
1119
1120 #[test]
1121 fn test_parse_filter_simple() {
1122 let schema = Arc::new(Schema::new(vec![
1123 Field::new("i", DataType::Int32, false),
1124 Field::new("s", DataType::Utf8, true),
1125 Field::new(
1126 "st",
1127 DataType::Struct(Fields::from(vec![
1128 Field::new("x", DataType::Float32, false),
1129 Field::new("y", DataType::Float32, false),
1130 ])),
1131 true,
1132 ),
1133 ]));
1134
1135 let planner = Planner::new(schema.clone());
1136
1137 let expected = col("i")
1138 .gt(lit(3_i32))
1139 .and(col("st").field_newstyle("x").lt_eq(lit(5.0_f32)))
1140 .and(
1141 col("s")
1142 .eq(lit("str-4"))
1143 .or(col("s").in_list(vec![lit("str-4"), lit("str-5")], false)),
1144 );
1145
1146 let expr = planner
1148 .parse_filter("i > 3 AND st.x <= 5.0 AND (s == 'str-4' OR s in ('str-4', 'str-5'))")
1149 .unwrap();
1150 assert_eq!(expr, expected);
1151
1152 let expr = planner
1154 .parse_filter("i > 3 AND st.x <= 5.0 AND (s = 'str-4' OR s in ('str-4', 'str-5'))")
1155 .unwrap();
1156
1157 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1158
1159 let batch = RecordBatch::try_new(
1160 schema,
1161 vec![
1162 Arc::new(Int32Array::from_iter_values(0..10)) as ArrayRef,
1163 Arc::new(StringArray::from_iter_values(
1164 (0..10).map(|v| format!("str-{}", v)),
1165 )),
1166 Arc::new(StructArray::from(vec![
1167 (
1168 Arc::new(Field::new("x", DataType::Float32, false)),
1169 Arc::new(Float32Array::from_iter_values((0..10).map(|v| v as f32)))
1170 as ArrayRef,
1171 ),
1172 (
1173 Arc::new(Field::new("y", DataType::Float32, false)),
1174 Arc::new(Float32Array::from_iter_values(
1175 (0..10).map(|v| (v * 10) as f32),
1176 )),
1177 ),
1178 ])),
1179 ],
1180 )
1181 .unwrap();
1182 let predicates = physical_expr.evaluate(&batch).unwrap();
1183 assert_eq!(
1184 predicates.into_array(0).unwrap().as_ref(),
1185 &BooleanArray::from(vec![
1186 false, false, false, false, true, true, false, false, false, false
1187 ])
1188 );
1189 }
1190
1191 #[test]
1192 fn test_parse_deep_logical_filter() {
1193 let planner = Planner::new(Arc::new(Schema::empty()));
1194
1195 for op in ["AND", "OR"] {
1196 let filter = std::iter::repeat_n("true", 1000)
1197 .collect::<Vec<_>>()
1198 .join(&format!(" {op} "));
1199
1200 let expr = planner.parse_filter(&filter).unwrap();
1201 let optimized = planner.optimize_expr(expr).unwrap();
1202
1203 assert_eq!(optimized, lit(true));
1204 }
1205 }
1206
1207 #[derive(Debug, Eq, PartialEq, Hash)]
1208 struct StrictFloat64Udf {
1209 signature: Signature,
1210 }
1211
1212 impl StrictFloat64Udf {
1213 fn new() -> Self {
1214 Self {
1215 signature: Signature::exact(vec![DataType::Float64], Volatility::Immutable),
1216 }
1217 }
1218 }
1219
1220 impl ScalarUDFImpl for StrictFloat64Udf {
1221 fn name(&self) -> &str {
1222 "strict_float64"
1223 }
1224
1225 fn signature(&self) -> &Signature {
1226 &self.signature
1227 }
1228
1229 fn return_type(&self, _arg_types: &[DataType]) -> DFResult<DataType> {
1230 Ok(DataType::Float64)
1231 }
1232
1233 fn invoke_with_args(&self, args: ScalarFunctionArgs) -> DFResult<ColumnarValue> {
1234 let data_type = args.args[0].data_type();
1235 assert_eq!(
1236 data_type,
1237 DataType::Float64,
1238 "strict_float64 expected Float64, got {data_type}"
1239 );
1240 Ok(ColumnarValue::Scalar(ScalarValue::Float64(Some(0.0))))
1241 }
1242 }
1243
1244 #[test]
1245 fn test_coerce_before_simplify() {
1246 let planner = Planner::new(Arc::new(Schema::empty()));
1247 let strict_float64 = Arc::new(ScalarUDF::new_from_impl(StrictFloat64Udf::new()));
1248 let expr = Expr::ScalarFunction(ScalarFunction::new_udf(strict_float64, vec![lit(0_i64)]))
1249 .eq(lit(0.0_f64));
1250
1251 let optimized = planner.optimize_expr(expr).unwrap();
1252
1253 planner.create_physical_expr(&optimized).unwrap();
1254 }
1255
1256 #[test]
1257 fn test_nested_col_refs() {
1258 let schema = Arc::new(Schema::new(vec![
1259 Field::new("s0", DataType::Utf8, true),
1260 Field::new(
1261 "st",
1262 DataType::Struct(Fields::from(vec![
1263 Field::new("s1", DataType::Utf8, true),
1264 Field::new(
1265 "st",
1266 DataType::Struct(Fields::from(vec![Field::new(
1267 "s2",
1268 DataType::Utf8,
1269 true,
1270 )])),
1271 true,
1272 ),
1273 ])),
1274 true,
1275 ),
1276 ]));
1277
1278 let planner = Planner::new(schema);
1279
1280 fn assert_column_eq(planner: &Planner, expr: &str, expected: &Expr) {
1281 let expr = planner.parse_filter(&format!("{expr} = 'val'")).unwrap();
1282 assert!(matches!(
1283 expr,
1284 Expr::BinaryExpr(BinaryExpr {
1285 left: _,
1286 op: Operator::Eq,
1287 right: _
1288 })
1289 ));
1290 if let Expr::BinaryExpr(BinaryExpr { left, .. }) = expr {
1291 assert_eq!(left.as_ref(), expected);
1292 }
1293 }
1294
1295 let expected = Expr::Column(Column::new_unqualified("s0"));
1296 assert_column_eq(&planner, "s0", &expected);
1297 assert_column_eq(&planner, "`s0`", &expected);
1298
1299 let expected = Expr::ScalarFunction(ScalarFunction {
1300 func: Arc::new(ScalarUDF::new_from_impl(GetFieldFunc::default())),
1301 args: vec![
1302 Expr::Column(Column::new_unqualified("st")),
1303 Expr::Literal(ScalarValue::Utf8(Some("s1".to_string())), None),
1304 ],
1305 });
1306 assert_column_eq(&planner, "st.s1", &expected);
1307 assert_column_eq(&planner, "`st`.`s1`", &expected);
1308 assert_column_eq(&planner, "st.`s1`", &expected);
1309
1310 let expected = Expr::ScalarFunction(ScalarFunction {
1311 func: Arc::new(ScalarUDF::new_from_impl(GetFieldFunc::default())),
1312 args: vec![
1313 Expr::ScalarFunction(ScalarFunction {
1314 func: Arc::new(ScalarUDF::new_from_impl(GetFieldFunc::default())),
1315 args: vec![
1316 Expr::Column(Column::new_unqualified("st")),
1317 Expr::Literal(ScalarValue::Utf8(Some("st".to_string())), None),
1318 ],
1319 }),
1320 Expr::Literal(ScalarValue::Utf8(Some("s2".to_string())), None),
1321 ],
1322 });
1323
1324 assert_column_eq(&planner, "st.st.s2", &expected);
1325 assert_column_eq(&planner, "`st`.`st`.`s2`", &expected);
1326 assert_column_eq(&planner, "st.st.`s2`", &expected);
1327 assert_column_eq(&planner, "st['st'][\"s2\"]", &expected);
1328 }
1329
1330 #[test]
1331 fn test_nested_list_refs() {
1332 let schema = Arc::new(Schema::new(vec![Field::new(
1333 "l",
1334 DataType::List(Arc::new(Field::new(
1335 "item",
1336 DataType::Struct(Fields::from(vec![Field::new("f1", DataType::Utf8, true)])),
1337 true,
1338 ))),
1339 true,
1340 )]));
1341
1342 let planner = Planner::new(schema);
1343
1344 let expected = array_element(col("l"), lit(0_i64));
1345 let expr = planner.parse_expr("l[0]").unwrap();
1346 assert_eq!(expr, expected);
1347
1348 let expected = get_field(array_element(col("l"), lit(0_i64)), "f1");
1349 let expr = planner.parse_expr("l[0]['f1']").unwrap();
1350 assert_eq!(expr, expected);
1351
1352 }
1357
1358 #[test]
1359 fn test_negative_expressions() {
1360 let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, false)]));
1361
1362 let planner = Planner::new(schema.clone());
1363
1364 let expected = col("x")
1365 .gt(lit(-3_i64))
1366 .and(col("x").lt(-(lit(-5_i64) + lit(3_i64))));
1367
1368 let expr = planner.parse_filter("x > -3 AND x < -(-5 + 3)").unwrap();
1369
1370 assert_eq!(expr, expected);
1371
1372 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1373
1374 let batch = RecordBatch::try_new(
1375 schema,
1376 vec![Arc::new(Int64Array::from_iter_values(-5..5)) as ArrayRef],
1377 )
1378 .unwrap();
1379 let predicates = physical_expr.evaluate(&batch).unwrap();
1380 assert_eq!(
1381 predicates.into_array(0).unwrap().as_ref(),
1382 &BooleanArray::from(vec![
1383 false, false, false, true, true, true, true, false, false, false
1384 ])
1385 );
1386 }
1387
1388 #[test]
1389 fn test_negative_array_expressions() {
1390 let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, false)]));
1391
1392 let planner = Planner::new(schema);
1393
1394 let expected = Expr::Literal(
1395 ScalarValue::List(Arc::new(
1396 ListArray::from_iter_primitive::<Float64Type, _, _>(vec![Some(
1397 [-1_f64, -2.0, -3.0, -4.0, -5.0].map(Some),
1398 )]),
1399 )),
1400 None,
1401 );
1402
1403 let expr = planner
1404 .parse_expr("[-1.0, -2.0, -3.0, -4.0, -5.0]")
1405 .unwrap();
1406
1407 assert_eq!(expr, expected);
1408 }
1409
1410 #[test]
1411 fn test_sql_like() {
1412 let schema = Arc::new(Schema::new(vec![Field::new("s", DataType::Utf8, true)]));
1413
1414 let planner = Planner::new(schema.clone());
1415
1416 let expected = col("s").like(lit("str-4"));
1417 let expr = planner.parse_filter("s LIKE 'str-4'").unwrap();
1419 assert_eq!(expr, expected);
1420 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1421
1422 let batch = RecordBatch::try_new(
1423 schema,
1424 vec![Arc::new(StringArray::from_iter_values(
1425 (0..10).map(|v| format!("str-{}", v)),
1426 ))],
1427 )
1428 .unwrap();
1429 let predicates = physical_expr.evaluate(&batch).unwrap();
1430 assert_eq!(
1431 predicates.into_array(0).unwrap().as_ref(),
1432 &BooleanArray::from(vec![
1433 false, false, false, false, true, false, false, false, false, false
1434 ])
1435 );
1436 }
1437
1438 #[test]
1439 fn test_not_like() {
1440 let schema = Arc::new(Schema::new(vec![Field::new("s", DataType::Utf8, true)]));
1441
1442 let planner = Planner::new(schema.clone());
1443
1444 let expected = col("s").not_like(lit("str-4"));
1445 let expr = planner.parse_filter("s NOT LIKE 'str-4'").unwrap();
1447 assert_eq!(expr, expected);
1448 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1449
1450 let batch = RecordBatch::try_new(
1451 schema,
1452 vec![Arc::new(StringArray::from_iter_values(
1453 (0..10).map(|v| format!("str-{}", v)),
1454 ))],
1455 )
1456 .unwrap();
1457 let predicates = physical_expr.evaluate(&batch).unwrap();
1458 assert_eq!(
1459 predicates.into_array(0).unwrap().as_ref(),
1460 &BooleanArray::from(vec![
1461 true, true, true, true, false, true, true, true, true, true
1462 ])
1463 );
1464 }
1465
1466 #[test]
1467 fn test_like_escape_char() {
1468 let schema = Arc::new(Schema::new(vec![Field::new("s", DataType::Utf8, true)]));
1469 let planner = Planner::new(schema);
1470
1471 for filter in ["s LIKE 'a!%' ESCAPE '!'", "s ILIKE 'a!%' ESCAPE '!'"] {
1473 match planner.parse_filter(filter).unwrap() {
1474 Expr::Like(like) => assert_eq!(like.escape_char, Some('!'), "{filter}"),
1475 other => panic!("expected a LIKE expression for `{filter}`, got {other:?}"),
1476 }
1477 }
1478
1479 for filter in [
1482 "s LIKE 'x' ESCAPE ''",
1483 "s LIKE 'x' ESCAPE 'ab'",
1484 "s ILIKE 'x' ESCAPE ''",
1485 "s ILIKE 'x' ESCAPE 'ab'",
1486 ] {
1487 let err = planner.parse_filter(filter).unwrap_err();
1488 assert!(
1489 err.to_string()
1490 .contains("Invalid escape character in LIKE expression"),
1491 "unexpected error for `{filter}`: {err}"
1492 );
1493 }
1494 }
1495
1496 #[test]
1497 fn test_sql_is_in() {
1498 let schema = Arc::new(Schema::new(vec![Field::new("s", DataType::Utf8, true)]));
1499
1500 let planner = Planner::new(schema.clone());
1501
1502 let expected = col("s").in_list(vec![lit("str-4"), lit("str-5")], false);
1503 let expr = planner.parse_filter("s IN ('str-4', 'str-5')").unwrap();
1505 assert_eq!(expr, expected);
1506 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1507
1508 let batch = RecordBatch::try_new(
1509 schema,
1510 vec![Arc::new(StringArray::from_iter_values(
1511 (0..10).map(|v| format!("str-{}", v)),
1512 ))],
1513 )
1514 .unwrap();
1515 let predicates = physical_expr.evaluate(&batch).unwrap();
1516 assert_eq!(
1517 predicates.into_array(0).unwrap().as_ref(),
1518 &BooleanArray::from(vec![
1519 false, false, false, false, true, true, false, false, false, false
1520 ])
1521 );
1522 }
1523
1524 #[test]
1525 fn test_sql_is_null() {
1526 let schema = Arc::new(Schema::new(vec![Field::new("s", DataType::Utf8, true)]));
1527
1528 let planner = Planner::new(schema.clone());
1529
1530 let expected = col("s").is_null();
1531 let expr = planner.parse_filter("s IS NULL").unwrap();
1532 assert_eq!(expr, expected);
1533 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1534
1535 let batch = RecordBatch::try_new(
1536 schema,
1537 vec![Arc::new(StringArray::from_iter((0..10).map(|v| {
1538 if v % 3 == 0 {
1539 Some(format!("str-{}", v))
1540 } else {
1541 None
1542 }
1543 })))],
1544 )
1545 .unwrap();
1546 let predicates = physical_expr.evaluate(&batch).unwrap();
1547 assert_eq!(
1548 predicates.into_array(0).unwrap().as_ref(),
1549 &BooleanArray::from(vec![
1550 false, true, true, false, true, true, false, true, true, false
1551 ])
1552 );
1553
1554 let expr = planner.parse_filter("s IS NOT NULL").unwrap();
1555 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1556 let predicates = physical_expr.evaluate(&batch).unwrap();
1557 assert_eq!(
1558 predicates.into_array(0).unwrap().as_ref(),
1559 &BooleanArray::from(vec![
1560 true, false, false, true, false, false, true, false, false, true,
1561 ])
1562 );
1563 }
1564
1565 #[test]
1566 fn test_sql_invert() {
1567 let schema = Arc::new(Schema::new(vec![Field::new("s", DataType::Boolean, true)]));
1568
1569 let planner = Planner::new(schema.clone());
1570
1571 let expr = planner.parse_filter("NOT s").unwrap();
1572 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1573
1574 let batch = RecordBatch::try_new(
1575 schema,
1576 vec![Arc::new(BooleanArray::from_iter(
1577 (0..10).map(|v| Some(v % 3 == 0)),
1578 ))],
1579 )
1580 .unwrap();
1581 let predicates = physical_expr.evaluate(&batch).unwrap();
1582 assert_eq!(
1583 predicates.into_array(0).unwrap().as_ref(),
1584 &BooleanArray::from(vec![
1585 false, true, true, false, true, true, false, true, true, false
1586 ])
1587 );
1588 }
1589
1590 #[test]
1591 fn test_sql_cast() {
1592 let cases = &[
1593 (
1594 "x = cast('2021-01-01 00:00:00' as timestamp)",
1595 ArrowDataType::Timestamp(TimeUnit::Microsecond, None),
1596 ),
1597 (
1598 "x = cast('2021-01-01 00:00:00' as timestamp(0))",
1599 ArrowDataType::Timestamp(TimeUnit::Second, None),
1600 ),
1601 (
1602 "x = cast('2021-01-01 00:00:00.123' as timestamp(9))",
1603 ArrowDataType::Timestamp(TimeUnit::Nanosecond, None),
1604 ),
1605 (
1606 "x = cast('2021-01-01 00:00:00.123' as datetime(9))",
1607 ArrowDataType::Timestamp(TimeUnit::Nanosecond, None),
1608 ),
1609 ("x = cast('2021-01-01' as date)", ArrowDataType::Date32),
1610 (
1611 "x = cast('1.238' as decimal(9,3))",
1612 ArrowDataType::Decimal128(9, 3),
1613 ),
1614 ("x = cast(1 as float)", ArrowDataType::Float32),
1615 ("x = cast(1 as double)", ArrowDataType::Float64),
1616 ("x = cast(1 as tinyint)", ArrowDataType::Int8),
1617 ("x = cast(1 as smallint)", ArrowDataType::Int16),
1618 ("x = cast(1 as int)", ArrowDataType::Int32),
1619 ("x = cast(1 as integer)", ArrowDataType::Int32),
1620 ("x = cast(1 as bigint)", ArrowDataType::Int64),
1621 ("x = cast(1 as tinyint unsigned)", ArrowDataType::UInt8),
1622 ("x = cast(1 as smallint unsigned)", ArrowDataType::UInt16),
1623 ("x = cast(1 as int unsigned)", ArrowDataType::UInt32),
1624 ("x = cast(1 as integer unsigned)", ArrowDataType::UInt32),
1625 ("x = cast(1 as bigint unsigned)", ArrowDataType::UInt64),
1626 ("x = cast(1 as boolean)", ArrowDataType::Boolean),
1627 ("x = cast(1 as string)", ArrowDataType::Utf8),
1628 ];
1629
1630 for (sql, expected_data_type) in cases {
1631 let schema = Arc::new(Schema::new(vec![Field::new(
1632 "x",
1633 expected_data_type.clone(),
1634 true,
1635 )]));
1636 let planner = Planner::new(schema.clone());
1637 let expr = planner.parse_filter(sql).unwrap();
1638
1639 let expected_value_str = sql
1641 .split("cast(")
1642 .nth(1)
1643 .unwrap()
1644 .split(" as")
1645 .next()
1646 .unwrap();
1647 let expected_value_str = expected_value_str.trim_matches('\'');
1649
1650 match expr {
1651 Expr::BinaryExpr(BinaryExpr { right, .. }) => match right.as_ref() {
1652 Expr::Cast(Cast { expr, field }) => {
1653 match expr.as_ref() {
1654 Expr::Literal(ScalarValue::Utf8(Some(value_str)), _) => {
1655 assert_eq!(value_str, expected_value_str);
1656 }
1657 Expr::Literal(ScalarValue::Int64(Some(value)), _) => {
1658 assert_eq!(*value, 1);
1659 }
1660 _ => panic!("Expected cast to be applied to literal"),
1661 }
1662 assert_eq!(field.data_type(), expected_data_type);
1663 }
1664 _ => panic!("Expected right to be a cast"),
1665 },
1666 _ => panic!("Expected binary expression"),
1667 }
1668 }
1669 }
1670
1671 #[test]
1672 fn test_sql_literals() {
1673 let cases = &[
1674 (
1675 "x = timestamp '2021-01-01 00:00:00'",
1676 ArrowDataType::Timestamp(TimeUnit::Microsecond, None),
1677 ),
1678 (
1679 "x = timestamp(0) '2021-01-01 00:00:00'",
1680 ArrowDataType::Timestamp(TimeUnit::Second, None),
1681 ),
1682 (
1683 "x = timestamp(9) '2021-01-01 00:00:00.123'",
1684 ArrowDataType::Timestamp(TimeUnit::Nanosecond, None),
1685 ),
1686 ("x = date '2021-01-01'", ArrowDataType::Date32),
1687 ("x = decimal(9,3) '1.238'", ArrowDataType::Decimal128(9, 3)),
1688 ];
1689
1690 for (sql, expected_data_type) in cases {
1691 let schema = Arc::new(Schema::new(vec![Field::new(
1692 "x",
1693 expected_data_type.clone(),
1694 true,
1695 )]));
1696 let planner = Planner::new(schema.clone());
1697 let expr = planner.parse_filter(sql).unwrap();
1698
1699 let expected_value_str = sql.split('\'').nth(1).unwrap();
1700
1701 match expr {
1702 Expr::BinaryExpr(BinaryExpr { right, .. }) => match right.as_ref() {
1703 Expr::Cast(Cast { expr, field }) => {
1704 match expr.as_ref() {
1705 Expr::Literal(ScalarValue::Utf8(Some(value_str)), _) => {
1706 assert_eq!(value_str, expected_value_str);
1707 }
1708 _ => panic!("Expected cast to be applied to literal"),
1709 }
1710 assert_eq!(field.data_type(), expected_data_type);
1711 }
1712 _ => panic!("Expected right to be a cast"),
1713 },
1714 _ => panic!("Expected binary expression"),
1715 }
1716 }
1717 }
1718
1719 #[test]
1720 fn test_sql_array_literals() {
1721 let cases = [
1722 (
1723 "x = [1, 2, 3]",
1724 ArrowDataType::List(Arc::new(Field::new("item", ArrowDataType::Int64, true))),
1725 ),
1726 (
1727 "x = [1, 2, 3]",
1728 ArrowDataType::FixedSizeList(
1729 Arc::new(Field::new("item", ArrowDataType::Int64, true)),
1730 3,
1731 ),
1732 ),
1733 ];
1734
1735 for (sql, expected_data_type) in cases {
1736 let schema = Arc::new(Schema::new(vec![Field::new(
1737 "x",
1738 expected_data_type.clone(),
1739 true,
1740 )]));
1741 let planner = Planner::new(schema.clone());
1742 let expr = planner.parse_filter(sql).unwrap();
1743 let expr = planner.optimize_expr(expr).unwrap();
1744
1745 match expr {
1746 Expr::BinaryExpr(BinaryExpr { right, .. }) => match right.as_ref() {
1747 Expr::Literal(value, _) => {
1748 assert_eq!(&value.data_type(), &expected_data_type);
1749 }
1750 _ => panic!("Expected right to be a literal"),
1751 },
1752 _ => panic!("Expected binary expression"),
1753 }
1754 }
1755 }
1756
1757 #[test]
1758 fn test_sql_between() {
1759 use arrow_array::{Float64Array, Int32Array, TimestampMicrosecondArray};
1760 use arrow_schema::{DataType, Field, Schema, TimeUnit};
1761 use std::sync::Arc;
1762
1763 let schema = Arc::new(Schema::new(vec![
1764 Field::new("x", DataType::Int32, false),
1765 Field::new("y", DataType::Float64, false),
1766 Field::new(
1767 "ts",
1768 DataType::Timestamp(TimeUnit::Microsecond, None),
1769 false,
1770 ),
1771 ]));
1772
1773 let planner = Planner::new(schema.clone());
1774
1775 let expr = planner
1777 .parse_filter("x BETWEEN CAST(3 AS INT) AND CAST(7 AS INT)")
1778 .unwrap();
1779 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1780
1781 let base_ts = 1704067200000000_i64; let ts_array = TimestampMicrosecondArray::from_iter_values(
1785 (0..10).map(|i| base_ts + i * 1_000_000), );
1787
1788 let batch = RecordBatch::try_new(
1789 schema,
1790 vec![
1791 Arc::new(Int32Array::from_iter_values(0..10)) as ArrayRef,
1792 Arc::new(Float64Array::from_iter_values((0..10).map(|v| v as f64))),
1793 Arc::new(ts_array),
1794 ],
1795 )
1796 .unwrap();
1797
1798 let predicates = physical_expr.evaluate(&batch).unwrap();
1799 assert_eq!(
1800 predicates.into_array(0).unwrap().as_ref(),
1801 &BooleanArray::from(vec![
1802 false, false, false, true, true, true, true, true, false, false
1803 ])
1804 );
1805
1806 let expr = planner
1808 .parse_filter("x NOT BETWEEN CAST(3 AS INT) AND CAST(7 AS INT)")
1809 .unwrap();
1810 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1811
1812 let predicates = physical_expr.evaluate(&batch).unwrap();
1813 assert_eq!(
1814 predicates.into_array(0).unwrap().as_ref(),
1815 &BooleanArray::from(vec![
1816 true, true, true, false, false, false, false, false, true, true
1817 ])
1818 );
1819
1820 let expr = planner.parse_filter("y BETWEEN 2.5 AND 6.5").unwrap();
1822 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1823
1824 let predicates = physical_expr.evaluate(&batch).unwrap();
1825 assert_eq!(
1826 predicates.into_array(0).unwrap().as_ref(),
1827 &BooleanArray::from(vec![
1828 false, false, false, true, true, true, true, false, false, false
1829 ])
1830 );
1831
1832 let expr = planner
1834 .parse_filter(
1835 "ts BETWEEN timestamp '2024-01-01 00:00:03' AND timestamp '2024-01-01 00:00:07'",
1836 )
1837 .unwrap();
1838 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1839
1840 let predicates = physical_expr.evaluate(&batch).unwrap();
1841 assert_eq!(
1842 predicates.into_array(0).unwrap().as_ref(),
1843 &BooleanArray::from(vec![
1844 false, false, false, true, true, true, true, true, false, false
1845 ])
1846 );
1847 }
1848
1849 #[test]
1850 fn test_sql_comparison() {
1851 let batch: Vec<(&str, ArrayRef)> = vec![
1853 (
1854 "timestamp_s",
1855 Arc::new(TimestampSecondArray::from_iter_values(0..10)),
1856 ),
1857 (
1858 "timestamp_ms",
1859 Arc::new(TimestampMillisecondArray::from_iter_values(0..10)),
1860 ),
1861 (
1862 "timestamp_us",
1863 Arc::new(TimestampMicrosecondArray::from_iter_values(0..10)),
1864 ),
1865 (
1866 "timestamp_ns",
1867 Arc::new(TimestampNanosecondArray::from_iter_values(4995..5005)),
1868 ),
1869 ];
1870 let batch = RecordBatch::try_from_iter(batch).unwrap();
1871
1872 let planner = Planner::new(batch.schema());
1873
1874 let expressions = &[
1876 "timestamp_s >= TIMESTAMP '1970-01-01 00:00:05'",
1877 "timestamp_ms >= TIMESTAMP '1970-01-01 00:00:00.005'",
1878 "timestamp_us >= TIMESTAMP '1970-01-01 00:00:00.000005'",
1879 "timestamp_ns >= TIMESTAMP '1970-01-01 00:00:00.000005'",
1880 ];
1881
1882 let expected: ArrayRef = Arc::new(BooleanArray::from_iter(
1883 std::iter::repeat_n(Some(false), 5).chain(std::iter::repeat_n(Some(true), 5)),
1884 ));
1885 for expression in expressions {
1886 let logical_expr = planner.parse_filter(expression).unwrap();
1888 let logical_expr = planner.optimize_expr(logical_expr).unwrap();
1889 let physical_expr = planner.create_physical_expr(&logical_expr).unwrap();
1890
1891 let result = physical_expr.evaluate(&batch).unwrap();
1893 let result = result.into_array(batch.num_rows()).unwrap();
1894 assert_eq!(&expected, &result, "unexpected result for {}", expression);
1895 }
1896 }
1897
1898 #[test]
1899 fn test_columns_in_expr() {
1900 let expr = col("s0").gt(lit("value")).and(
1901 col("st")
1902 .field("st")
1903 .field("s2")
1904 .eq(lit("value"))
1905 .or(col("st")
1906 .field("s1")
1907 .in_list(vec![lit("value 1"), lit("value 2")], false)),
1908 );
1909
1910 let columns = Planner::column_names_in_expr(&expr);
1911 assert_eq!(columns, vec!["s0", "st.s1", "st.st.s2"]);
1912 }
1913
1914 #[test]
1915 fn test_parse_binary_expr() {
1916 let bin_str = "x'616263'";
1917
1918 let schema = Arc::new(Schema::new(vec![Field::new(
1919 "binary",
1920 DataType::Binary,
1921 true,
1922 )]));
1923 let planner = Planner::new(schema);
1924 let expr = planner.parse_expr(bin_str).unwrap();
1925 assert_eq!(
1926 expr,
1927 Expr::Literal(ScalarValue::Binary(Some(vec![b'a', b'b', b'c'])), None)
1928 );
1929 }
1930
1931 #[test]
1932 fn test_lance_context_provider_expr_planners() {
1933 let ctx_provider = LanceContextProvider::default();
1934 assert!(!ctx_provider.get_expr_planners().is_empty());
1935 }
1936
1937 #[test]
1938 fn test_regexp_match_and_non_empty_captions() {
1939 let schema = Arc::new(Schema::new(vec![
1942 Field::new("keywords", DataType::Utf8, true),
1943 Field::new("natural_caption", DataType::Utf8, true),
1944 Field::new("poetic_caption", DataType::Utf8, true),
1945 ]));
1946
1947 let planner = Planner::new(schema.clone());
1948
1949 let expr = planner
1950 .parse_filter(
1951 "regexp_match(keywords, 'Liberty|revolution') AND \
1952 (natural_caption IS NOT NULL AND natural_caption <> '' AND \
1953 poetic_caption IS NOT NULL AND poetic_caption <> '')",
1954 )
1955 .unwrap();
1956
1957 let physical_expr = planner.create_physical_expr(&expr).unwrap();
1958
1959 let batch = RecordBatch::try_new(
1960 schema,
1961 vec![
1962 Arc::new(StringArray::from(vec![
1963 Some("Liberty for all"),
1964 Some("peace"),
1965 Some("revolution now"),
1966 Some("Liberty"),
1967 Some("revolutionary"),
1968 Some("none"),
1969 ])) as ArrayRef,
1970 Arc::new(StringArray::from(vec![
1971 Some("a"),
1972 Some("b"),
1973 None,
1974 Some(""),
1975 Some("c"),
1976 Some("d"),
1977 ])) as ArrayRef,
1978 Arc::new(StringArray::from(vec![
1979 Some("x"),
1980 Some(""),
1981 Some("y"),
1982 Some("z"),
1983 None,
1984 Some("w"),
1985 ])) as ArrayRef,
1986 ],
1987 )
1988 .unwrap();
1989
1990 let result = physical_expr.evaluate(&batch).unwrap();
1991 assert_eq!(
1992 result.into_array(0).unwrap().as_ref(),
1993 &BooleanArray::from(vec![true, false, false, false, false, false])
1994 );
1995 }
1996
1997 #[test]
1998 fn test_regexp_match_infer_error_without_boolean_coercion() {
1999 let schema = Arc::new(Schema::new(vec![
2002 Field::new("keywords", DataType::Utf8, true),
2003 Field::new("natural_caption", DataType::Utf8, true),
2004 Field::new("poetic_caption", DataType::Utf8, true),
2005 ]));
2006
2007 let planner = Planner::new(schema);
2008
2009 let expr = planner
2010 .parse_filter(
2011 "regexp_match(keywords, 'Liberty|revolution') AND \
2012 (natural_caption IS NOT NULL AND natural_caption <> '' AND \
2013 poetic_caption IS NOT NULL AND poetic_caption <> '')",
2014 )
2015 .unwrap();
2016
2017 let _physical = planner.create_physical_expr(&expr).unwrap();
2019 }
2020
2021 #[test]
2022 fn test_jsonb_literals() {
2023 let schema = Arc::new(Schema::new(vec![Field::new(
2024 "j",
2025 DataType::LargeBinary,
2026 true,
2027 )]));
2028 let planner = Planner::new(schema);
2029
2030 let cases = [
2031 ("jsonb '{\"key\": \"value\"}'", r#"{"key":"value"}"#),
2032 ("cast('{\"a\": 1}' as jsonb)", r#"{"a":1}"#),
2033 ("'{\"a\": 1}'::jsonb", r#"{"a":1}"#),
2034 ];
2035 for (sql, expected) in cases {
2036 let ast = parse_sql_expr(sql).unwrap();
2037 let expr = planner.parse_sql_expr(&ast).unwrap();
2038 match expr {
2039 Expr::Literal(ScalarValue::LargeBinary(Some(bytes)), _) => {
2040 assert_eq!(
2041 lance_arrow::json::decode_json(&bytes),
2042 expected,
2043 "failed for: {sql}"
2044 );
2045 }
2046 other => panic!("Expected LargeBinary literal for '{sql}', got: {other:?}"),
2047 }
2048 }
2049 }
2050
2051 #[test]
2052 fn test_jsonb_literal_errors() {
2053 let schema = Arc::new(Schema::new(vec![Field::new(
2054 "j",
2055 DataType::LargeBinary,
2056 true,
2057 )]));
2058 let planner = Planner::new(schema);
2059
2060 let ast = parse_sql_expr("jsonb 'not valid json'").unwrap();
2062 let err = planner.parse_sql_expr(&ast).unwrap_err();
2063 assert!(
2064 err.to_string().contains("Failed to encode JSONB"),
2065 "expected JSONB encoding error, got: {err}"
2066 );
2067
2068 let ast = parse_sql_expr("cast(j as jsonb)").unwrap();
2070 let err = planner.parse_sql_expr(&ast).unwrap_err();
2071 assert!(
2072 err.to_string()
2073 .contains("CAST to JSONB only supports string literals"),
2074 "got: {err}"
2075 );
2076 }
2077}