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