Skip to main content

common_function/scalars/
matches.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::collections::HashMap;
16use std::fmt;
17use std::sync::Arc;
18
19use common_query::error::{InvalidFuncArgsSnafu, Result};
20use datafusion::arrow::array::{Array, ArrayRef, AsArray, BooleanArray};
21use datafusion::common::tree_node::{Transformed, TreeNode, TreeNodeIterator, TreeNodeRecursion};
22use datafusion::common::{DFSchema, Result as DfResult};
23use datafusion::execution::SessionStateBuilder;
24use datafusion::logical_expr::{self, ColumnarValue, Expr, Volatility};
25use datafusion::physical_planner::{DefaultPhysicalPlanner, PhysicalPlanner};
26use datafusion_common::DataFusionError;
27use datafusion_expr::physical_planning_context::PhysicalPlanningContext;
28use datafusion_expr::{ScalarFunctionArgs, Signature};
29use datatypes::arrow::array::RecordBatch;
30use datatypes::arrow::datatypes::{DataType, Field};
31use snafu::{OptionExt, ensure};
32
33use crate::function::{Function, extract_args};
34use crate::function_registry::FunctionRegistry;
35
36/// `matches` for full text search.
37///
38/// Usage: matches(`<col>`, `<pattern>`) -> boolean
39#[derive(Clone, Debug)]
40pub struct MatchesFunction {
41    signature: Signature,
42}
43
44impl MatchesFunction {
45    pub fn register(registry: &FunctionRegistry) {
46        registry.register_scalar(MatchesFunction::default());
47    }
48}
49
50impl Default for MatchesFunction {
51    fn default() -> Self {
52        Self {
53            signature: Signature::string(2, Volatility::Immutable),
54        }
55    }
56}
57
58impl fmt::Display for MatchesFunction {
59    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
60        write!(f, "MATCHES")
61    }
62}
63
64impl Function for MatchesFunction {
65    fn name(&self) -> &str {
66        "matches"
67    }
68
69    fn return_type(&self, _: &[DataType]) -> datafusion_common::Result<DataType> {
70        Ok(DataType::Boolean)
71    }
72
73    fn signature(&self) -> &Signature {
74        &self.signature
75    }
76
77    // TODO: read case-sensitive config
78    fn invoke_with_args(&self, args: ScalarFunctionArgs) -> DfResult<ColumnarValue> {
79        let [data_column, patterns] = extract_args(self.name(), &args)?;
80
81        if data_column.is_empty() {
82            return Ok(ColumnarValue::Array(Arc::new(BooleanArray::from(
83                Vec::<bool>::with_capacity(0),
84            ))));
85        }
86
87        // Safety: both length and type are checked before
88        let pattern = match patterns.data_type() {
89            DataType::Utf8View => patterns.as_string_view().value(0),
90            DataType::Utf8 => patterns.as_string::<i32>().value(0),
91            DataType::LargeUtf8 => patterns.as_string::<i64>().value(0),
92            t => {
93                return Err(DataFusionError::Execution(format!(
94                    "unsupported datatype {t}"
95                )));
96            }
97        };
98        self.eval(data_column, pattern)
99    }
100}
101
102impl MatchesFunction {
103    fn eval(&self, data_array: ArrayRef, pattern: &str) -> DfResult<ColumnarValue> {
104        let col_name = "data";
105        let parser_context = ParserContext::default();
106        let raw_ast = parser_context.parse_pattern(pattern)?;
107        let ast = raw_ast.transform_ast()?;
108
109        let like_expr = ast.into_like_expr(col_name);
110
111        let input_schema = Self::input_schema();
112        let session_state = SessionStateBuilder::new().with_default_features().build();
113        let planner = DefaultPhysicalPlanner::default();
114        let physical_expr = planner.create_physical_expr(
115            &like_expr,
116            &input_schema,
117            &session_state,
118            &PhysicalPlanningContext::default(),
119        )?;
120
121        let arrow_schema = Arc::new(input_schema.as_arrow().clone());
122        let input_record_batch = RecordBatch::try_new(arrow_schema, vec![data_array]).unwrap();
123
124        let num_rows = input_record_batch.num_rows();
125        let result = physical_expr.evaluate(&input_record_batch)?;
126        let result_array = result.into_array(num_rows)?;
127
128        Ok(ColumnarValue::Array(Arc::new(result_array)))
129    }
130
131    fn input_schema() -> DFSchema {
132        DFSchema::from_unqualified_fields(
133            [Arc::new(Field::new("data", DataType::Utf8, true))].into(),
134            HashMap::new(),
135        )
136        .unwrap()
137    }
138}
139
140#[derive(Debug, Clone, PartialEq, Eq)]
141enum PatternAst {
142    // Distinguish this with `Group` for simplicity
143    /// A leaf node that matches a column with `pattern`
144    Literal { op: UnaryOp, pattern: String },
145    /// Flattened binary chains
146    Binary {
147        op: BinaryOp,
148        children: Vec<PatternAst>,
149    },
150    /// A sub-tree enclosed by parenthesis
151    Group { op: UnaryOp, child: Box<PatternAst> },
152}
153
154#[derive(Debug, Copy, Clone, PartialEq, Eq)]
155enum UnaryOp {
156    Must,
157    Optional,
158    Negative,
159}
160
161#[derive(Debug, Copy, Clone, PartialEq, Eq)]
162enum BinaryOp {
163    And,
164    Or,
165}
166
167impl PatternAst {
168    fn into_like_expr(self, column: &str) -> Expr {
169        match self {
170            PatternAst::Literal { op, pattern } => {
171                let expr = Self::convert_literal(column, &pattern);
172                match op {
173                    UnaryOp::Must => expr,
174                    UnaryOp::Optional => expr,
175                    UnaryOp::Negative => logical_expr::not(expr),
176                }
177            }
178            PatternAst::Binary { op, children } => {
179                if children.is_empty() {
180                    return logical_expr::lit(true);
181                }
182                let exprs = children
183                    .into_iter()
184                    .map(|child| child.into_like_expr(column));
185                // safety: children is not empty
186                match op {
187                    BinaryOp::And => exprs.reduce(Expr::and).unwrap(),
188                    BinaryOp::Or => exprs.reduce(Expr::or).unwrap(),
189                }
190            }
191            PatternAst::Group { op, child } => {
192                let child = child.into_like_expr(column);
193                match op {
194                    UnaryOp::Must => child,
195                    UnaryOp::Optional => child,
196                    UnaryOp::Negative => logical_expr::not(child),
197                }
198            }
199        }
200    }
201
202    fn convert_literal(column: &str, pattern: &str) -> Expr {
203        logical_expr::col(column).like(logical_expr::lit(format!(
204            "%{}%",
205            crate::utils::escape_like_pattern(pattern)
206        )))
207    }
208
209    /// Transform this AST with preset rules to make it correct.
210    fn transform_ast(self) -> Result<Self> {
211        self.transform_up(Self::collapse_binary_branch_fn)
212            .map(|data| data.data)?
213            .transform_up(Self::eliminate_optional_fn)
214            .map(|data| data.data)?
215            .transform_down(Self::eliminate_single_child_fn)
216            .map(|data| data.data)
217            .map_err(Into::into)
218    }
219
220    /// Collapse binary branch with the same operator. I.e., this transformer
221    /// changes the binary-tree AST into a multiple branching AST.
222    ///
223    /// This function is expected to be called in a bottom-up manner as
224    /// it won't recursion.
225    fn collapse_binary_branch_fn(self) -> DfResult<Transformed<Self>> {
226        let PatternAst::Binary {
227            op: parent_op,
228            children,
229        } = self
230        else {
231            return Ok(Transformed::no(self));
232        };
233
234        let mut collapsed = vec![];
235        let mut remains = vec![];
236
237        for child in children {
238            match child {
239                PatternAst::Literal { .. } | PatternAst::Group { .. } => {
240                    collapsed.push(child);
241                }
242                PatternAst::Binary { op, children } => {
243                    // no need to recursion because this function is expected to be called
244                    // in a bottom-up manner
245                    if op == parent_op {
246                        collapsed.extend(children);
247                    } else {
248                        remains.push(PatternAst::Binary { op, children });
249                    }
250                }
251            }
252        }
253
254        if collapsed.is_empty() {
255            Ok(Transformed::no(PatternAst::Binary {
256                op: parent_op,
257                children: remains,
258            }))
259        } else {
260            collapsed.extend(remains);
261            Ok(Transformed::yes(PatternAst::Binary {
262                op: parent_op,
263                children: collapsed,
264            }))
265        }
266    }
267
268    /// Eliminate optional pattern. An optional pattern can always be
269    /// omitted or transformed into a must pattern follows the following rules:
270    /// - If there is only one pattern and it's optional, change it to must
271    /// - If there is any must pattern, remove all other optional patterns
272    fn eliminate_optional_fn(self) -> DfResult<Transformed<Self>> {
273        let PatternAst::Binary {
274            op: parent_op,
275            children,
276        } = self
277        else {
278            return Ok(Transformed::no(self));
279        };
280
281        if parent_op == BinaryOp::Or {
282            let mut must_list = vec![];
283            let mut must_not_list = vec![];
284            let mut optional_list = vec![];
285            let mut compound_list = vec![];
286
287            for child in children {
288                match child {
289                    PatternAst::Literal { op, .. } | PatternAst::Group { op, .. } => match op {
290                        UnaryOp::Must => must_list.push(child),
291                        UnaryOp::Optional => optional_list.push(child),
292                        UnaryOp::Negative => must_not_list.push(child),
293                    },
294                    PatternAst::Binary { .. } => {
295                        compound_list.push(child);
296                    }
297                }
298            }
299
300            // Eliminate optional list if there is MUST.
301            if !must_list.is_empty() {
302                optional_list.clear();
303            }
304
305            let children_this_level = optional_list.into_iter().chain(compound_list).collect();
306            let new_node = if !must_list.is_empty() || !must_not_list.is_empty() {
307                let new_children = must_list
308                    .into_iter()
309                    .chain(must_not_list)
310                    .chain(Some(PatternAst::Binary {
311                        op: BinaryOp::Or,
312                        children: children_this_level,
313                    }))
314                    .collect();
315                PatternAst::Binary {
316                    op: BinaryOp::And,
317                    children: new_children,
318                }
319            } else {
320                PatternAst::Binary {
321                    op: BinaryOp::Or,
322                    children: children_this_level,
323                }
324            };
325
326            return Ok(Transformed::yes(new_node));
327        }
328
329        Ok(Transformed::no(PatternAst::Binary {
330            op: parent_op,
331            children,
332        }))
333    }
334
335    /// Eliminate single child [`PatternAst::Binary`] node. If a binary node has only one child, it can be
336    /// replaced by its only child.
337    ///
338    /// This function prefers to be applied in a top-down manner. But it's not required.
339    fn eliminate_single_child_fn(self) -> DfResult<Transformed<Self>> {
340        let PatternAst::Binary { op, mut children } = self else {
341            return Ok(Transformed::no(self));
342        };
343
344        // remove empty grand children
345        children.retain(|child| match child {
346            PatternAst::Binary {
347                children: grand_children,
348                ..
349            } => !grand_children.is_empty(),
350            PatternAst::Literal { .. } | PatternAst::Group { .. } => true,
351        });
352
353        if children.len() == 1 {
354            Ok(Transformed::yes(children.into_iter().next().unwrap()))
355        } else {
356            Ok(Transformed::no(PatternAst::Binary { op, children }))
357        }
358    }
359}
360
361impl TreeNode for PatternAst {
362    fn apply_children<'n, F: FnMut(&'n Self) -> DfResult<TreeNodeRecursion>>(
363        &'n self,
364        mut f: F,
365    ) -> DfResult<TreeNodeRecursion> {
366        match self {
367            PatternAst::Literal { .. } => Ok(TreeNodeRecursion::Continue),
368            PatternAst::Binary { op: _, children } => {
369                for child in children {
370                    if TreeNodeRecursion::Stop == f(child)? {
371                        return Ok(TreeNodeRecursion::Stop);
372                    }
373                }
374                Ok(TreeNodeRecursion::Continue)
375            }
376            PatternAst::Group { op: _, child } => f(child),
377        }
378    }
379
380    fn map_children<F: FnMut(Self) -> DfResult<Transformed<Self>>>(
381        self,
382        mut f: F,
383    ) -> DfResult<Transformed<Self>> {
384        match self {
385            PatternAst::Literal { .. } => Ok(Transformed::no(self)),
386            PatternAst::Binary { op, children } => children
387                .into_iter()
388                .map_until_stop_and_collect(&mut f)?
389                .map_data(|new_children| {
390                    Ok(PatternAst::Binary {
391                        op,
392                        children: new_children,
393                    })
394                }),
395            PatternAst::Group { op, child } => f(*child)?.map_data(|new_child| {
396                Ok(PatternAst::Group {
397                    op,
398                    child: Box::new(new_child),
399                })
400            }),
401        }
402    }
403}
404
405#[derive(Default)]
406struct ParserContext {
407    stack: Vec<PatternAst>,
408}
409
410impl ParserContext {
411    pub fn parse_pattern(mut self, pattern: &str) -> Result<PatternAst> {
412        let tokenizer = Tokenizer::default();
413        let raw_tokens = tokenizer.tokenize(pattern)?;
414        let raw_tokens = Self::accomplish_optional_unary_op(raw_tokens)?;
415        let mut tokens = Self::to_rpn(raw_tokens)?;
416
417        while !tokens.is_empty() {
418            self.parse_one_impl(&mut tokens)?;
419        }
420
421        ensure!(
422            !self.stack.is_empty(),
423            InvalidFuncArgsSnafu {
424                err_msg: "Empty pattern",
425            }
426        );
427
428        // conjoin them together
429        if self.stack.len() == 1 {
430            Ok(self.stack.pop().unwrap())
431        } else {
432            Ok(PatternAst::Binary {
433                op: BinaryOp::Or,
434                children: self.stack,
435            })
436        }
437    }
438
439    /// Add [`Token::Optional`] for all bare [`Token::Phase`] and [`Token::Or`]
440    /// for all adjacent [`Token::Phase`]s.
441    ///
442    /// This function also does some checks by the way. Like if two unary ops are
443    /// adjacent.
444    fn accomplish_optional_unary_op(raw_tokens: Vec<Token>) -> Result<Vec<Token>> {
445        let mut is_prev_unary_op = false;
446        // The first one doesn't need binary op
447        let mut is_binary_op_before = true;
448        let mut is_unary_op_before = false;
449        let mut new_tokens = Vec::with_capacity(raw_tokens.len());
450        for token in raw_tokens {
451            // fill `Token::Or`
452            if !is_binary_op_before
453                && matches!(
454                    token,
455                    Token::Phase(_)
456                        | Token::OpenParen
457                        | Token::Must
458                        | Token::Optional
459                        | Token::Negative
460                )
461            {
462                is_binary_op_before = true;
463                new_tokens.push(Token::Or);
464            }
465            if matches!(
466                token,
467                Token::OpenParen // treat open paren as begin of new group
468                | Token::And | Token::Or
469            ) {
470                is_binary_op_before = true;
471            } else if matches!(token, Token::Phase(_) | Token::CloseParen) {
472                // need binary op next time
473                is_binary_op_before = false;
474            }
475
476            // fill `Token::Optional`
477            if !is_prev_unary_op && matches!(token, Token::Phase(_) | Token::OpenParen) {
478                new_tokens.push(Token::Optional);
479            } else {
480                is_prev_unary_op = matches!(token, Token::Must | Token::Negative);
481            }
482
483            // check if unary ops are adjacent by the way
484            if matches!(token, Token::Must | Token::Optional | Token::Negative) {
485                if is_unary_op_before {
486                    return InvalidFuncArgsSnafu {
487                        err_msg: "Invalid pattern, unary operators should not be adjacent",
488                    }
489                    .fail();
490                }
491                is_unary_op_before = true;
492            } else {
493                is_unary_op_before = false;
494            }
495
496            new_tokens.push(token);
497        }
498
499        Ok(new_tokens)
500    }
501
502    /// Convert infix token stream to RPN
503    fn to_rpn(mut raw_tokens: Vec<Token>) -> Result<Vec<Token>> {
504        let mut operator_stack = vec![];
505        let mut result = vec![];
506        raw_tokens.reverse();
507
508        while let Some(token) = raw_tokens.pop() {
509            match token {
510                Token::Phase(_) => result.push(token),
511                Token::Must | Token::Negative | Token::Optional => {
512                    operator_stack.push(token);
513                }
514                Token::OpenParen => operator_stack.push(token),
515                Token::And | Token::Or => {
516                    // - Or has lower priority than And
517                    // - Binary op have lower priority than unary op
518                    while let Some(stack_top) = operator_stack.last()
519                        && ((*stack_top == Token::And && token == Token::Or)
520                            || matches!(
521                                *stack_top,
522                                Token::Must | Token::Optional | Token::Negative
523                            ))
524                    {
525                        result.push(operator_stack.pop().unwrap());
526                    }
527                    operator_stack.push(token);
528                }
529                Token::CloseParen => {
530                    let mut is_open_paren_found = false;
531                    while let Some(op) = operator_stack.pop() {
532                        if op == Token::OpenParen {
533                            is_open_paren_found = true;
534                            break;
535                        }
536                        result.push(op);
537                    }
538                    if !is_open_paren_found {
539                        return InvalidFuncArgsSnafu {
540                            err_msg: "Unmatched close parentheses",
541                        }
542                        .fail();
543                    }
544                }
545            }
546        }
547
548        while let Some(operator) = operator_stack.pop() {
549            if operator == Token::OpenParen {
550                return InvalidFuncArgsSnafu {
551                    err_msg: "Unmatched parentheses",
552                }
553                .fail();
554            }
555            result.push(operator);
556        }
557
558        Ok(result)
559    }
560
561    fn parse_one_impl(&mut self, tokens: &mut Vec<Token>) -> Result<()> {
562        if let Some(token) = tokens.pop() {
563            match token {
564                Token::Must => {
565                    if self.stack.is_empty() {
566                        self.parse_one_impl(tokens)?;
567                    }
568                    let phase_or_group = self.stack.pop().context(InvalidFuncArgsSnafu {
569                        err_msg: "Invalid pattern, \"+\" operator should have one operand",
570                    })?;
571                    match phase_or_group {
572                        PatternAst::Literal { op: _, pattern } => {
573                            self.stack.push(PatternAst::Literal {
574                                op: UnaryOp::Must,
575                                pattern,
576                            });
577                        }
578                        PatternAst::Binary { .. } | PatternAst::Group { .. } => {
579                            self.stack.push(PatternAst::Group {
580                                op: UnaryOp::Must,
581                                child: Box::new(phase_or_group),
582                            })
583                        }
584                    }
585                    return Ok(());
586                }
587                Token::Negative => {
588                    if self.stack.is_empty() {
589                        self.parse_one_impl(tokens)?;
590                    }
591                    let phase_or_group = self.stack.pop().context(InvalidFuncArgsSnafu {
592                        err_msg: "Invalid pattern, \"-\" operator should have one operand",
593                    })?;
594                    match phase_or_group {
595                        PatternAst::Literal { op: _, pattern } => {
596                            self.stack.push(PatternAst::Literal {
597                                op: UnaryOp::Negative,
598                                pattern,
599                            });
600                        }
601                        PatternAst::Binary { .. } | PatternAst::Group { .. } => {
602                            self.stack.push(PatternAst::Group {
603                                op: UnaryOp::Negative,
604                                child: Box::new(phase_or_group),
605                            })
606                        }
607                    }
608                    return Ok(());
609                }
610                Token::Optional => {
611                    if self.stack.is_empty() {
612                        self.parse_one_impl(tokens)?;
613                    }
614                    let phase_or_group = self.stack.pop().context(InvalidFuncArgsSnafu {
615                        err_msg:
616                            "Invalid pattern, OPTIONAL(space) operator should have one operand",
617                    })?;
618                    match phase_or_group {
619                        PatternAst::Literal { op: _, pattern } => {
620                            self.stack.push(PatternAst::Literal {
621                                op: UnaryOp::Optional,
622                                pattern,
623                            });
624                        }
625                        PatternAst::Binary { .. } | PatternAst::Group { .. } => {
626                            self.stack.push(PatternAst::Group {
627                                op: UnaryOp::Optional,
628                                child: Box::new(phase_or_group),
629                            })
630                        }
631                    }
632                    return Ok(());
633                }
634                Token::Phase(pattern) => {
635                    self.stack.push(PatternAst::Literal {
636                        // Op here is a placeholder
637                        op: UnaryOp::Optional,
638                        pattern,
639                    })
640                }
641                Token::And => {
642                    if self.stack.is_empty() {
643                        self.parse_one_impl(tokens)?;
644                    };
645                    let rhs = self.stack.pop().context(InvalidFuncArgsSnafu {
646                        err_msg: "Invalid pattern, \"AND\" operator should have two operands",
647                    })?;
648                    if self.stack.is_empty() {
649                        self.parse_one_impl(tokens)?
650                    };
651                    let lhs = self.stack.pop().context(InvalidFuncArgsSnafu {
652                        err_msg: "Invalid pattern, \"AND\" operator should have two operands",
653                    })?;
654                    self.stack.push(PatternAst::Binary {
655                        op: BinaryOp::And,
656                        children: vec![lhs, rhs],
657                    });
658                    return Ok(());
659                }
660                Token::Or => {
661                    if self.stack.is_empty() {
662                        self.parse_one_impl(tokens)?
663                    };
664                    let rhs = self.stack.pop().context(InvalidFuncArgsSnafu {
665                        err_msg: "Invalid pattern, \"OR\" operator should have two operands",
666                    })?;
667                    if self.stack.is_empty() {
668                        self.parse_one_impl(tokens)?
669                    };
670                    let lhs = self.stack.pop().context(InvalidFuncArgsSnafu {
671                        err_msg: "Invalid pattern, \"OR\" operator should have two operands",
672                    })?;
673                    self.stack.push(PatternAst::Binary {
674                        op: BinaryOp::Or,
675                        children: vec![lhs, rhs],
676                    });
677                    return Ok(());
678                }
679                Token::OpenParen | Token::CloseParen => {
680                    return InvalidFuncArgsSnafu {
681                        err_msg: "Unexpected parentheses",
682                    }
683                    .fail();
684                }
685            }
686        }
687
688        Ok(())
689    }
690}
691
692#[derive(Clone, Debug, PartialEq, Eq)]
693enum Token {
694    /// "+"
695    Must,
696    /// "-"
697    Negative,
698    /// "AND"
699    And,
700    /// "OR"
701    Or,
702    /// "("
703    OpenParen,
704    /// ")"
705    CloseParen,
706    /// Any other phases
707    Phase(String),
708
709    /// This is not a token from user input, but a placeholder for internal use.
710    /// It's used to accomplish the unary operator class with Must and Negative.
711    /// In user provided pattern, optional is expressed by a bare phase or group
712    /// (simply nothing or writespace).
713    Optional,
714}
715
716#[derive(Default)]
717struct Tokenizer {
718    cursor: usize,
719}
720
721impl Tokenizer {
722    pub fn tokenize(mut self, pattern: &str) -> Result<Vec<Token>> {
723        let mut tokens = vec![];
724        let char_len = pattern.chars().count();
725        while self.cursor < char_len {
726            // TODO: collect pattern into Vec<char> if this tokenizer is bottleneck in the future
727            let c = pattern.chars().nth(self.cursor).unwrap();
728            match c {
729                '+' => tokens.push(Token::Must),
730                '-' => tokens.push(Token::Negative),
731                '(' => tokens.push(Token::OpenParen),
732                ')' => tokens.push(Token::CloseParen),
733                ' ' => {
734                    if let Some(last_token) = tokens.last() {
735                        match last_token {
736                            Token::Must | Token::Negative => {
737                                return InvalidFuncArgsSnafu {
738                                    err_msg: format!("Unexpected space after {:?}", last_token),
739                                }
740                                .fail();
741                            }
742                            _ => {}
743                        }
744                    }
745                }
746                '\"' => {
747                    self.step_next();
748                    let phase = self.consume_next_phase(true, pattern)?;
749                    tokens.push(Token::Phase(phase));
750                    // consume a writespace (or EOF) after quotes
751                    if let Some(ending_separator) = self.consume_next(pattern)
752                        && ending_separator != ' '
753                    {
754                        return InvalidFuncArgsSnafu {
755                            err_msg: "Expect a space after quotes ('\"')",
756                        }
757                        .fail();
758                    }
759                }
760                _ => {
761                    let phase = self.consume_next_phase(false, pattern)?;
762                    match phase.to_uppercase().as_str() {
763                        "AND" => tokens.push(Token::And),
764                        "OR" => tokens.push(Token::Or),
765                        _ => tokens.push(Token::Phase(phase)),
766                    }
767                }
768            }
769            self.cursor += 1;
770        }
771        Ok(tokens)
772    }
773
774    fn consume_next(&mut self, pattern: &str) -> Option<char> {
775        self.cursor += 1;
776        pattern.chars().nth(self.cursor)
777    }
778
779    fn step_next(&mut self) {
780        self.cursor += 1;
781    }
782
783    fn rewind_one(&mut self) {
784        self.cursor -= 1;
785    }
786
787    /// Current `cursor` points to the first character of the phase.
788    /// If the phase is enclosed by double quotes, consume the start quote before calling this.
789    fn consume_next_phase(&mut self, is_quoted: bool, pattern: &str) -> Result<String> {
790        let mut phase = String::new();
791        let mut is_quote_present = false;
792
793        let char_len = pattern.chars().count();
794        while self.cursor < char_len {
795            let mut c = pattern.chars().nth(self.cursor).unwrap();
796
797            match c {
798                '\"' => {
799                    is_quote_present = true;
800                    break;
801                }
802                ' ' if !is_quoted => {
803                    break;
804                }
805                '(' | ')' | '+' | '-' if !is_quoted => {
806                    self.rewind_one();
807                    break;
808                }
809                '\\' => {
810                    let Some(next) = self.consume_next(pattern) else {
811                        return InvalidFuncArgsSnafu {
812                            err_msg: "Unexpected end of pattern, expected a character after escape ('\\')",
813                        }.fail();
814                    };
815                    // it doesn't check whether the escaped character is valid or not
816                    c = next;
817                }
818                _ => {}
819            }
820
821            phase.push(c);
822            self.cursor += 1;
823        }
824
825        if is_quoted ^ is_quote_present {
826            return InvalidFuncArgsSnafu {
827                err_msg: "Unclosed quotes ('\"')",
828            }
829            .fail();
830        }
831
832        Ok(phase)
833    }
834}
835
836#[cfg(test)]
837mod test {
838    use datafusion::arrow::array::StringArray;
839    use datafusion_common::ScalarValue;
840    use datafusion_common::config::ConfigOptions;
841
842    use super::*;
843
844    #[test]
845    fn valid_matches_tokenizer() {
846        use Token::*;
847        let cases = [
848            (
849                "a +b -c",
850                vec![
851                    Phase("a".to_string()),
852                    Must,
853                    Phase("b".to_string()),
854                    Negative,
855                    Phase("c".to_string()),
856                ],
857            ),
858            (
859                "+a(b-c)",
860                vec![
861                    Must,
862                    Phase("a".to_string()),
863                    OpenParen,
864                    Phase("b".to_string()),
865                    Negative,
866                    Phase("c".to_string()),
867                    CloseParen,
868                ],
869            ),
870            (
871                r#"Barack Obama"#,
872                vec![Phase("Barack".to_string()), Phase("Obama".to_string())],
873            ),
874            (
875                r#"+apple +fruit"#,
876                vec![
877                    Must,
878                    Phase("apple".to_string()),
879                    Must,
880                    Phase("fruit".to_string()),
881                ],
882            ),
883            (
884                r#""He said \"hello\"""#,
885                vec![Phase("He said \"hello\"".to_string())],
886            ),
887            (
888                r#"a AND b OR c"#,
889                vec![
890                    Phase("a".to_string()),
891                    And,
892                    Phase("b".to_string()),
893                    Or,
894                    Phase("c".to_string()),
895                ],
896            ),
897            (
898                r#"中文 测试"#,
899                vec![Phase("中文".to_string()), Phase("测试".to_string())],
900            ),
901            (
902                r#"中文 AND 测试"#,
903                vec![Phase("中文".to_string()), And, Phase("测试".to_string())],
904            ),
905            (
906                r#"中文 +测试"#,
907                vec![Phase("中文".to_string()), Must, Phase("测试".to_string())],
908            ),
909            (
910                r#"中文 -测试"#,
911                vec![
912                    Phase("中文".to_string()),
913                    Negative,
914                    Phase("测试".to_string()),
915                ],
916            ),
917        ];
918
919        for (query, expected) in cases {
920            let tokenizer = Tokenizer::default();
921            let tokens = tokenizer.tokenize(query).unwrap();
922            assert_eq!(expected, tokens, "{query}");
923        }
924    }
925
926    #[test]
927    fn invalid_matches_tokenizer() {
928        let cases = [
929            (r#""He said "hello""#, "Expect a space after quotes"),
930            (r#""He said hello"#, "Unclosed quotes"),
931            (r#"a + b - c"#, "Unexpected space after"),
932            (r#"ab "c"def"#, "Expect a space after quotes"),
933        ];
934
935        for (query, expected) in cases {
936            let tokenizer = Tokenizer::default();
937            let result = tokenizer.tokenize(query);
938            assert!(result.is_err(), "{query}");
939            let actual_error = result.unwrap_err().to_string();
940            assert!(actual_error.contains(expected), "{query}, {actual_error}");
941        }
942    }
943
944    #[test]
945    fn valid_ast_transformer() {
946        let cases = [
947            (
948                "a AND b OR c",
949                PatternAst::Binary {
950                    op: BinaryOp::Or,
951                    children: vec![
952                        PatternAst::Literal {
953                            op: UnaryOp::Optional,
954                            pattern: "c".to_string(),
955                        },
956                        PatternAst::Binary {
957                            op: BinaryOp::And,
958                            children: vec![
959                                PatternAst::Literal {
960                                    op: UnaryOp::Optional,
961                                    pattern: "a".to_string(),
962                                },
963                                PatternAst::Literal {
964                                    op: UnaryOp::Optional,
965                                    pattern: "b".to_string(),
966                                },
967                            ],
968                        },
969                    ],
970                },
971            ),
972            (
973                "a -b",
974                PatternAst::Binary {
975                    op: BinaryOp::And,
976                    children: vec![
977                        PatternAst::Literal {
978                            op: UnaryOp::Negative,
979                            pattern: "b".to_string(),
980                        },
981                        PatternAst::Literal {
982                            op: UnaryOp::Optional,
983                            pattern: "a".to_string(),
984                        },
985                    ],
986                },
987            ),
988            (
989                "a +b",
990                PatternAst::Literal {
991                    op: UnaryOp::Must,
992                    pattern: "b".to_string(),
993                },
994            ),
995            (
996                "a b c d",
997                PatternAst::Binary {
998                    op: BinaryOp::Or,
999                    children: vec![
1000                        PatternAst::Literal {
1001                            op: UnaryOp::Optional,
1002                            pattern: "a".to_string(),
1003                        },
1004                        PatternAst::Literal {
1005                            op: UnaryOp::Optional,
1006                            pattern: "b".to_string(),
1007                        },
1008                        PatternAst::Literal {
1009                            op: UnaryOp::Optional,
1010                            pattern: "c".to_string(),
1011                        },
1012                        PatternAst::Literal {
1013                            op: UnaryOp::Optional,
1014                            pattern: "d".to_string(),
1015                        },
1016                    ],
1017                },
1018            ),
1019            (
1020                "a b c AND d",
1021                PatternAst::Binary {
1022                    op: BinaryOp::Or,
1023                    children: vec![
1024                        PatternAst::Literal {
1025                            op: UnaryOp::Optional,
1026                            pattern: "a".to_string(),
1027                        },
1028                        PatternAst::Literal {
1029                            op: UnaryOp::Optional,
1030                            pattern: "b".to_string(),
1031                        },
1032                        PatternAst::Binary {
1033                            op: BinaryOp::And,
1034                            children: vec![
1035                                PatternAst::Literal {
1036                                    op: UnaryOp::Optional,
1037                                    pattern: "c".to_string(),
1038                                },
1039                                PatternAst::Literal {
1040                                    op: UnaryOp::Optional,
1041                                    pattern: "d".to_string(),
1042                                },
1043                            ],
1044                        },
1045                    ],
1046                },
1047            ),
1048            (
1049                r#"中文 测试"#,
1050                PatternAst::Binary {
1051                    op: BinaryOp::Or,
1052                    children: vec![
1053                        PatternAst::Literal {
1054                            op: UnaryOp::Optional,
1055                            pattern: "中文".to_string(),
1056                        },
1057                        PatternAst::Literal {
1058                            op: UnaryOp::Optional,
1059                            pattern: "测试".to_string(),
1060                        },
1061                    ],
1062                },
1063            ),
1064            (
1065                r#"中文 AND 测试"#,
1066                PatternAst::Binary {
1067                    op: BinaryOp::And,
1068                    children: vec![
1069                        PatternAst::Literal {
1070                            op: UnaryOp::Optional,
1071                            pattern: "中文".to_string(),
1072                        },
1073                        PatternAst::Literal {
1074                            op: UnaryOp::Optional,
1075                            pattern: "测试".to_string(),
1076                        },
1077                    ],
1078                },
1079            ),
1080            (
1081                r#"中文 +测试"#,
1082                PatternAst::Literal {
1083                    op: UnaryOp::Must,
1084                    pattern: "测试".to_string(),
1085                },
1086            ),
1087            (
1088                r#"中文 -测试"#,
1089                PatternAst::Binary {
1090                    op: BinaryOp::And,
1091                    children: vec![
1092                        PatternAst::Literal {
1093                            op: UnaryOp::Negative,
1094                            pattern: "测试".to_string(),
1095                        },
1096                        PatternAst::Literal {
1097                            op: UnaryOp::Optional,
1098                            pattern: "中文".to_string(),
1099                        },
1100                    ],
1101                },
1102            ),
1103        ];
1104
1105        for (query, expected) in cases {
1106            let parser = ParserContext { stack: vec![] };
1107            let ast = parser.parse_pattern(query).unwrap();
1108            let ast = ast.transform_ast().unwrap();
1109            assert_eq!(expected, ast, "{query}");
1110        }
1111    }
1112
1113    #[test]
1114    fn invalid_ast() {
1115        let cases = [
1116            (r#"a b (c"#, "Unmatched parentheses"),
1117            (r#"a b) c"#, "Unmatched close parentheses"),
1118            (r#"a +-b"#, "unary operators should not be adjacent"),
1119        ];
1120
1121        for (query, expected) in cases {
1122            let result: Result<()> = (|| {
1123                let parser = ParserContext { stack: vec![] };
1124                let ast = parser.parse_pattern(query)?;
1125                let _ast = ast.transform_ast()?;
1126                Ok(())
1127            })();
1128
1129            assert!(result.is_err(), "{query}");
1130            let actual_error = result.unwrap_err().to_string();
1131            assert!(actual_error.contains(expected), "{query}, {actual_error}");
1132        }
1133    }
1134
1135    #[test]
1136    fn valid_matches_parser() {
1137        let cases = [
1138            (
1139                "a AND b OR c",
1140                PatternAst::Binary {
1141                    op: BinaryOp::Or,
1142                    children: vec![
1143                        PatternAst::Binary {
1144                            op: BinaryOp::And,
1145                            children: vec![
1146                                PatternAst::Literal {
1147                                    op: UnaryOp::Optional,
1148                                    pattern: "a".to_string(),
1149                                },
1150                                PatternAst::Literal {
1151                                    op: UnaryOp::Optional,
1152                                    pattern: "b".to_string(),
1153                                },
1154                            ],
1155                        },
1156                        PatternAst::Literal {
1157                            op: UnaryOp::Optional,
1158                            pattern: "c".to_string(),
1159                        },
1160                    ],
1161                },
1162            ),
1163            (
1164                "(a AND b) OR c",
1165                PatternAst::Binary {
1166                    op: BinaryOp::Or,
1167                    children: vec![
1168                        PatternAst::Group {
1169                            op: UnaryOp::Optional,
1170                            child: Box::new(PatternAst::Binary {
1171                                op: BinaryOp::And,
1172                                children: vec![
1173                                    PatternAst::Literal {
1174                                        op: UnaryOp::Optional,
1175                                        pattern: "a".to_string(),
1176                                    },
1177                                    PatternAst::Literal {
1178                                        op: UnaryOp::Optional,
1179                                        pattern: "b".to_string(),
1180                                    },
1181                                ],
1182                            }),
1183                        },
1184                        PatternAst::Literal {
1185                            op: UnaryOp::Optional,
1186                            pattern: "c".to_string(),
1187                        },
1188                    ],
1189                },
1190            ),
1191            (
1192                "a AND (b OR c)",
1193                PatternAst::Binary {
1194                    op: BinaryOp::And,
1195                    children: vec![
1196                        PatternAst::Literal {
1197                            op: UnaryOp::Optional,
1198                            pattern: "a".to_string(),
1199                        },
1200                        PatternAst::Group {
1201                            op: UnaryOp::Optional,
1202                            child: Box::new(PatternAst::Binary {
1203                                op: BinaryOp::Or,
1204                                children: vec![
1205                                    PatternAst::Literal {
1206                                        op: UnaryOp::Optional,
1207                                        pattern: "b".to_string(),
1208                                    },
1209                                    PatternAst::Literal {
1210                                        op: UnaryOp::Optional,
1211                                        pattern: "c".to_string(),
1212                                    },
1213                                ],
1214                            }),
1215                        },
1216                    ],
1217                },
1218            ),
1219            (
1220                "a +b -c",
1221                PatternAst::Binary {
1222                    op: BinaryOp::Or,
1223                    children: vec![
1224                        PatternAst::Literal {
1225                            op: UnaryOp::Optional,
1226                            pattern: "a".to_string(),
1227                        },
1228                        PatternAst::Binary {
1229                            op: BinaryOp::Or,
1230                            children: vec![
1231                                PatternAst::Literal {
1232                                    op: UnaryOp::Must,
1233                                    pattern: "b".to_string(),
1234                                },
1235                                PatternAst::Literal {
1236                                    op: UnaryOp::Negative,
1237                                    pattern: "c".to_string(),
1238                                },
1239                            ],
1240                        },
1241                    ],
1242                },
1243            ),
1244            (
1245                "(+a +b) c",
1246                PatternAst::Binary {
1247                    op: BinaryOp::Or,
1248                    children: vec![
1249                        PatternAst::Group {
1250                            op: UnaryOp::Optional,
1251                            child: Box::new(PatternAst::Binary {
1252                                op: BinaryOp::Or,
1253                                children: vec![
1254                                    PatternAst::Literal {
1255                                        op: UnaryOp::Must,
1256                                        pattern: "a".to_string(),
1257                                    },
1258                                    PatternAst::Literal {
1259                                        op: UnaryOp::Must,
1260                                        pattern: "b".to_string(),
1261                                    },
1262                                ],
1263                            }),
1264                        },
1265                        PatternAst::Literal {
1266                            op: UnaryOp::Optional,
1267                            pattern: "c".to_string(),
1268                        },
1269                    ],
1270                },
1271            ),
1272            (
1273                "\"AND\" AnD \"OR\"",
1274                PatternAst::Binary {
1275                    op: BinaryOp::And,
1276                    children: vec![
1277                        PatternAst::Literal {
1278                            op: UnaryOp::Optional,
1279                            pattern: "AND".to_string(),
1280                        },
1281                        PatternAst::Literal {
1282                            op: UnaryOp::Optional,
1283                            pattern: "OR".to_string(),
1284                        },
1285                    ],
1286                },
1287            ),
1288        ];
1289
1290        for (query, expected) in cases {
1291            let parser = ParserContext { stack: vec![] };
1292            let ast = parser.parse_pattern(query).unwrap();
1293            assert_eq!(expected, ast, "{query}");
1294        }
1295    }
1296
1297    #[test]
1298    fn evaluate_matches() {
1299        let input_data = vec![
1300            "The quick brown fox jumps over the lazy dog",
1301            "The             fox jumps over the lazy dog",
1302            "The quick brown     jumps over the lazy dog",
1303            "The quick brown fox       over the lazy dog",
1304            "The quick brown fox jumps      the lazy dog",
1305            "The quick brown fox jumps over          dog",
1306            "The quick brown fox jumps over the      dog",
1307        ];
1308        let col: ArrayRef = Arc::new(StringArray::from(input_data));
1309        let cases = [
1310            // basic cases
1311            ("quick", vec![true, false, true, true, true, true, true]),
1312            (
1313                "\"quick brown\"",
1314                vec![true, false, true, true, true, true, true],
1315            ),
1316            (
1317                "\"fox jumps\"",
1318                vec![true, true, false, false, true, true, true],
1319            ),
1320            (
1321                "fox OR lazy",
1322                vec![true, true, true, true, true, true, true],
1323            ),
1324            (
1325                "fox AND lazy",
1326                vec![true, true, false, true, true, false, false],
1327            ),
1328            (
1329                "-over -lazy",
1330                vec![false, false, false, false, false, false, false],
1331            ),
1332            (
1333                "-over AND -lazy",
1334                vec![false, false, false, false, false, false, false],
1335            ),
1336            // priority between AND & OR
1337            (
1338                "fox AND jumps OR over",
1339                vec![true, true, true, true, true, true, true],
1340            ),
1341            (
1342                "fox OR brown AND quick",
1343                vec![true, true, true, true, true, true, true],
1344            ),
1345            (
1346                "(fox OR brown) AND quick",
1347                vec![true, false, true, true, true, true, true],
1348            ),
1349            (
1350                "brown AND quick OR fox",
1351                vec![true, true, true, true, true, true, true],
1352            ),
1353            (
1354                "brown AND (quick OR fox)",
1355                vec![true, false, true, true, true, true, true],
1356            ),
1357            (
1358                "brown AND quick AND fox  OR  jumps AND over AND lazy",
1359                vec![true, true, true, true, true, true, true],
1360            ),
1361            // optional & must conversion
1362            (
1363                "quick brown fox +jumps",
1364                vec![true, true, true, false, true, true, true],
1365            ),
1366            (
1367                "fox +jumps -over",
1368                vec![false, false, false, false, true, false, false],
1369            ),
1370            (
1371                "fox AND +jumps AND -over",
1372                vec![false, false, false, false, true, false, false],
1373            ),
1374            // weird parentheses cases
1375            (
1376                "(+fox +jumps) over",
1377                vec![true, true, true, true, true, true, true],
1378            ),
1379            (
1380                "+(fox jumps) AND over",
1381                vec![true, true, true, true, false, true, true],
1382            ),
1383            (
1384                "over -(fox jumps)",
1385                vec![false, false, false, false, false, false, false],
1386            ),
1387            (
1388                "over -(fox AND jumps)",
1389                vec![false, false, true, true, false, false, false],
1390            ),
1391            (
1392                "over AND -(-(fox OR jumps))",
1393                vec![true, true, true, true, false, true, true],
1394            ),
1395        ];
1396
1397        let f = MatchesFunction::default();
1398        for (pattern, expected) in cases {
1399            let args = ScalarFunctionArgs {
1400                args: vec![
1401                    ColumnarValue::Array(col.clone()),
1402                    ColumnarValue::Scalar(ScalarValue::Utf8View(Some(pattern.to_string()))),
1403                ],
1404                arg_fields: vec![],
1405                number_rows: col.len(),
1406                return_field: Arc::new(Field::new("x", col.data_type().clone(), true)),
1407                config_options: Arc::new(ConfigOptions::new()),
1408            };
1409            let actual = f
1410                .invoke_with_args(args)
1411                .and_then(|x| x.to_array(col.len()))
1412                .unwrap();
1413            let expected: ArrayRef = Arc::new(BooleanArray::from(expected));
1414            assert_eq!(expected.as_ref(), actual.as_ref(), "{pattern}");
1415        }
1416    }
1417}