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