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