Skip to main content

lance_datafusion/
planner.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4//! Exec plan planner
5
6use 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
56/// Encode a JSON string into a JSONB `LargeBinary` literal expression.
57fn 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
63// The escape in `LIKE/ILIKE ... ESCAPE '<char>'` must be exactly one character.
64// Reject empty or multi-character escape strings rather than silently treating
65// them as "no escape" or truncating to the first character.
66fn 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
183// Adapter that instructs datafusion how lance expects expressions to be interpreted
184struct 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            // TODO: cast should go thru CAST syntax instead of UDF
233            // Going thru UDF makes it hard for the optimizer to find no-ops
234            "_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        // Variables (things like @@LANGUAGE) not supported
241        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    /// If passed with `true`, then the first identifier in column reference
289    /// is parsed as the relation. For example, `table.field.inner` will be
290    /// read as the nested field `field.inner` (`inner` on struct field `field`)
291    /// on the `table` relation. If `false` (the default), then no relations
292    /// are used and all identifiers are assumed to be a nested column path.
293    pub fn with_enable_relations(mut self, enable_relations: bool) -> Self {
294        self.enable_relations = enable_relations;
295        self
296    }
297
298    /// Resolve a column name using case-insensitive matching against the schema.
299    /// Returns the actual field name if found, otherwise returns the original name.
300    fn resolve_column_name(&self, name: &str) -> String {
301        // Try exact match first
302        if self.schema.field_with_name(name).is_ok() {
303            return name.to_string();
304        }
305        // Fall back to case-insensitive match
306        for field in self.schema.fields() {
307            if field.name().eq_ignore_ascii_case(name) {
308                return field.name().clone();
309            }
310        }
311        // Not found in schema - return original (might be computed column, system column, etc.)
312        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            // Create qualified column reference (relation.column)
330            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            // Default behavior - treat as struct field access
338            // Use resolved column name to handle case-insensitive matching
339            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    // See datafusion `sqlToRel::parse_sql_number()`
485    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    // We now use datafusion to parse functions.  This allows us to use datafusion's
528    // entire collection of functions (previously we had just hard-coded support for two functions).
529    //
530    // Unfortunately, one of those two functions was is_valid and the reason we needed it was because
531    // this is a function that comes from duckdb.  Datafusion does not consider is_valid to be a function
532    // but rather an AST node (Expr::IsNotNull) and so we need to handle this case specially.
533    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                    // Default to microsecond to match PyArrow
624                    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                // Users can pass string literals wrapped in `"`.
687                // (Normally SQL only allows single quotes.)
688                if id.quote_style == Some('"') {
689                    Ok(Expr::Literal(
690                        ScalarValue::Utf8(Some(id.value.clone())),
691                        None,
692                    ))
693                // Users can wrap identifiers with ` to reference non-standard
694                // names, such as uppercase or spaces.
695                } 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            // JSONB literal: jsonb '{"key": "value"}'
777            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            // For example, DATE '2020-01-01'
788            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            // JSONB cast: CAST('...' AS JSONB) or '...'::jsonb
844            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                        // x.y or x['y']
882                        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                            // Handle other cases like JSON access
902                            // Note: JSON access is not supported in lance
903                            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                // Parse the main expression and bounds
922                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    /// Create Logical [Expr] from a SQL filter clause.
941    ///
942    /// Note: the returned expression must be passed through `optimize_expr()`
943    /// before being passed to `create_physical_expr()`.
944    pub fn parse_filter(&self, filter: &str) -> Result<Expr> {
945        // Allow sqlparser to parse filter as part of ONE SQL statement.
946        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    /// Create Logical [Expr] from a SQL expression.
957    ///
958    /// Note: the returned expression must be passed through `optimize_filter()`
959    /// before being passed to `create_physical_expr()`.
960    pub fn parse_expr(&self, expr: &str) -> Result<Expr> {
961        // First check if it's a simple column reference (no operators, functions, etc.)
962        // resolve_column_name tries exact match first, then falls back to case-insensitive
963        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        // Parse as SQL expression
969        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    /// Try to decode bytes from hex literal string.
977    ///
978    /// Copied from datafusion because this is not public.
979    ///
980    /// TODO: use SqlToRel from Datafusion directly?
981    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            // The first byte is formed of only one char.
988            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    /// Try to decode a byte from a hex char.
1001    ///
1002    /// None will be returned if the input char is hex-invalid.
1003    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    /// Optimize the filter expression and coerce data types.
1013    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        // DataFusion rewrites arrow_cast to Expr::Cast, whose Arrow kernel does not support
1017        // integer-to-Time32 casts. Convert literal values with Lance's scalar coercion first.
1018        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        // DataFusion needs the coerce and simplify passes to be applied before
1054        // expressions can be handled by the physical planner.
1055        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        // Coerce before simplify to match DataFusion's analyzer-before-optimizer pipeline.
1063        let expr = simplifier.coerce(expr, &df_schema)?;
1064
1065        // Fold each comparison's own operands and rewrite it before anything above
1066        // it folds. `simplify` folds an operand and everything above it in one
1067        // pass, so a fully constant predicate whose zero appears only as a result
1068        // of folding never presents a zero literal to the rewrite:
1069        // `-1.0 * 0.0 < (1.0 - 1.0)` answered `true` where IEEE says false, and a
1070        // wrapper such as `IS TRUE` or a `CAST` did the same to the comparison's
1071        // own result.
1072        let expr = normalize_zero_comparisons(expr, &|operand| simplifier.simplify(operand))?;
1073
1074        // Again after simplify, which is what expands `BETWEEN` into two
1075        // comparisons and folds the casts `coerce` inserts, so those forms only
1076        // become visible on this pass.
1077        //
1078        // Running the rewrite more than once is safe because its output is a fixed
1079        // point of `optimize_expr`; `optimizing_twice_changes_nothing` pins that.
1080        let expr = simplifier.simplify(expr)?;
1081        let expr = rewrite_signed_zero_comparisons(expr)?;
1082
1083        Ok(expr)
1084    }
1085
1086    /// Create the [`PhysicalExpr`] from a logical [`Expr`]
1087    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    /// Create a [`PhysicalExpr`] using the caller's DataFusion session.
1097    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    /// Collect the columns in the expression.
1107    ///
1108    /// The columns are returned in sorted order.
1109    ///
1110    /// If the expr refers to nested columns these will be returned
1111    /// as dotted paths (x.y.z)
1112    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 column path. If this is empty, we are not in a column expression.
1124    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                // Build the field path from the column name and any nested fields
1135                // The nested field names from get_field already come as literal strings,
1136                // so we just need to concatenate them properly
1137                let mut path = name.clone();
1138                for part in self.current_path.drain(..) {
1139                    path.push('.');
1140                    // Check if the part needs quoting (contains dots)
1141                    if part.contains('.') || part.contains('`') {
1142                        // Quote the field name with backticks and escape any existing backticks
1143                        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        // double quotes
1218        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        // single quote
1224        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        // FIXME: This should work, but sqlparser doesn't recognize anything
1447        // after the period for some reason.
1448        // let expr = planner.parse_expr("l[0].f1").unwrap();
1449        // assert_eq!(expr, expected);
1450    }
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        // single quote
1546        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        // single quote
1574        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        // A valid single-character escape is captured for both LIKE and ILIKE.
1600        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        // Empty and multi-character escapes are rejected rather than silently
1608        // dropped or truncated to the first character.
1609        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        // single quote
1632        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            // Get the thing after 'cast(` but before ' as'.
1768            let expected_value_str = sql
1769                .split("cast(")
1770                .nth(1)
1771                .unwrap()
1772                .split(" as")
1773                .next()
1774                .unwrap();
1775            // Remove any quote marks
1776            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        // Test integer BETWEEN
1926        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        // Create timestamp array with values representing:
1932        // 2024-01-01 00:00:00 to 2024-01-01 00:00:09 (in microseconds)
1933        let base_ts = 1704067200000000_i64; // 2024-01-01 00:00:00
1934        let ts_array = TimestampMicrosecondArray::from_iter_values(
1935            (0..10).map(|i| base_ts + i * 1_000_000), // Each value is 1 second apart
1936        );
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        // Test NOT BETWEEN
1957        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        // Test floating point BETWEEN
1971        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        // Test timestamp BETWEEN
1983        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        // Create a batch with all data types
2002        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        // Each expression is meant to select the final 5 rows
2025        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            // convert to physical expression
2037            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            // Evaluate and assert they have correct results
2042            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        // Repro for a bug where regexp_match inside an AND chain wasn't coerced to boolean,
2090        // causing planning/evaluation failures. This should evaluate successfully.
2091        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        // With the fix applied, using parse_filter should coerce regexp_match to boolean
2150        // even when nested in a larger AND expression, so this should plan successfully.
2151        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        // Should not panic
2168        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        // Invalid JSON content
2211        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        // CAST with non-literal expression
2219        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}