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