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