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