Skip to main content

query/optimizer/
constant_term.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::fmt;
16use std::fmt::Formatter;
17use std::hash::{Hash, Hasher};
18use std::sync::Arc;
19
20use arrow::array::BooleanArray;
21use common_function::scalars::matches_term::MatchesTermFinder;
22use datafusion::config::ConfigOptions;
23use datafusion::error::Result as DfResult;
24use datafusion::physical_optimizer::PhysicalOptimizerRule;
25use datafusion::physical_plan::ExecutionPlan;
26use datafusion::physical_plan::filter::{FilterExec, FilterExecBuilder};
27use datafusion_common::ScalarValue;
28use datafusion_common::tree_node::{Transformed, TreeNode};
29use datafusion_expr::ColumnarValue;
30use datafusion_physical_expr::expressions::Literal;
31use datafusion_physical_expr::{PhysicalExpr, ScalarFunctionExpr};
32use datatypes::arrow_array::string_array_value_at_index;
33
34/// A physical expression that uses a pre-compiled term finder for the `matches_term` function.
35///
36/// This expression optimizes the `matches_term` function by pre-compiling the term
37/// when the term is a constant value. This avoids recompiling the term for each row
38/// during execution.
39#[derive(Debug)]
40pub struct PreCompiledMatchesTermExpr {
41    /// The text column expression to search in
42    text: Arc<dyn PhysicalExpr>,
43    /// The constant term to search for
44    term: String,
45    /// The pre-compiled term finder
46    finder: MatchesTermFinder,
47
48    /// No used but show how index tokenizes the term basically.
49    /// Not precise due to column options is unknown but for debugging purpose in most cases it's enough.
50    probes: Vec<String>,
51}
52
53impl fmt::Display for PreCompiledMatchesTermExpr {
54    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
55        write!(
56            f,
57            "MatchesConstTerm({}, term: \"{}\", probes: {:?})",
58            self.text, self.term, self.probes
59        )
60    }
61}
62
63impl Hash for PreCompiledMatchesTermExpr {
64    fn hash<H: Hasher>(&self, state: &mut H) {
65        self.text.hash(state);
66        self.term.hash(state);
67    }
68}
69
70impl PartialEq for PreCompiledMatchesTermExpr {
71    fn eq(&self, other: &Self) -> bool {
72        self.text.eq(&other.text) && self.term.eq(&other.term)
73    }
74}
75
76impl Eq for PreCompiledMatchesTermExpr {}
77
78impl PhysicalExpr for PreCompiledMatchesTermExpr {
79    fn data_type(
80        &self,
81        _input_schema: &arrow_schema::Schema,
82    ) -> datafusion::error::Result<arrow_schema::DataType> {
83        Ok(arrow_schema::DataType::Boolean)
84    }
85
86    fn nullable(&self, input_schema: &arrow_schema::Schema) -> datafusion::error::Result<bool> {
87        self.text.nullable(input_schema)
88    }
89
90    fn evaluate(
91        &self,
92        batch: &common_recordbatch::DfRecordBatch,
93    ) -> datafusion::error::Result<ColumnarValue> {
94        let num_rows = batch.num_rows();
95
96        let text_value = self.text.evaluate(batch)?;
97        let array = text_value.into_array(num_rows)?;
98
99        let mut result = BooleanArray::builder(num_rows);
100        for index in 0..array.len() {
101            match string_array_value_at_index(&array, index) {
102                Some(text) => {
103                    result.append_value(self.finder.find(text));
104                }
105                None => {
106                    result.append_null();
107                }
108            }
109        }
110
111        Ok(ColumnarValue::Array(Arc::new(result.finish())))
112    }
113
114    fn children(&self) -> Vec<&Arc<dyn PhysicalExpr>> {
115        vec![&self.text]
116    }
117
118    fn with_new_children(
119        self: Arc<Self>,
120        children: Vec<Arc<dyn PhysicalExpr>>,
121    ) -> datafusion::error::Result<Arc<dyn PhysicalExpr>> {
122        Ok(Arc::new(PreCompiledMatchesTermExpr {
123            text: children[0].clone(),
124            term: self.term.clone(),
125            finder: self.finder.clone(),
126            probes: self.probes.clone(),
127        }))
128    }
129
130    fn fmt_sql(&self, f: &mut Formatter<'_>) -> fmt::Result {
131        write!(f, "{}", self)
132    }
133}
134
135/// Optimizer rule that pre-compiles constant term in `matches_term` function.
136///
137/// This optimizer looks for `matches_term` function calls where the second argument
138/// (the term to match) is a constant value. When found, it replaces the function
139/// call with a specialized `PreCompiledMatchesTermExpr` that uses a pre-compiled
140/// term finder.
141///
142/// Example:
143/// ```sql
144/// -- Before optimization:
145/// matches_term(text_column, 'constant_term')
146///
147/// -- After optimization:
148/// PreCompiledMatchesTermExpr(text_column, 'constant_term')
149/// ```
150///
151/// This optimization improves performance by:
152/// 1. Pre-compiling the term once instead of for each row
153/// 2. Using a specialized expression that avoids function call overhead
154#[derive(Debug)]
155pub struct MatchesConstantTermOptimizer;
156
157impl PhysicalOptimizerRule for MatchesConstantTermOptimizer {
158    fn optimize(
159        &self,
160        plan: Arc<dyn ExecutionPlan>,
161        _config: &ConfigOptions,
162    ) -> DfResult<Arc<dyn ExecutionPlan>> {
163        let res = plan
164            .transform_down(&|plan: Arc<dyn ExecutionPlan>| {
165                if let Some(filter) = plan.downcast_ref::<FilterExec>() {
166                    let pred = filter.predicate().clone();
167                    let new_pred = pred.transform_down(&|expr: Arc<dyn PhysicalExpr>| {
168                        if let Some(func) = expr.downcast_ref::<ScalarFunctionExpr>() {
169                            if !func.name().eq_ignore_ascii_case("matches_term") {
170                                return Ok(Transformed::no(expr));
171                            }
172                            let args = func.args();
173                            if args.len() != 2 {
174                                return Ok(Transformed::no(expr));
175                            }
176
177                            if let Some(lit) = args[1].downcast_ref::<Literal>()
178                                && let ScalarValue::Utf8(Some(term)) = lit.value()
179                            {
180                                let finder = MatchesTermFinder::new(term);
181
182                                // For debugging purpose. Not really precise but enough for most cases.
183                                let probes = term
184                                    .split(|c: char| !c.is_alphanumeric() && c != '_')
185                                    .filter(|s| !s.is_empty())
186                                    .map(|s| s.to_string())
187                                    .collect();
188
189                                let expr = PreCompiledMatchesTermExpr {
190                                    text: args[0].clone(),
191                                    term: term.clone(),
192                                    finder,
193                                    probes,
194                                };
195
196                                return Ok(Transformed::yes(Arc::new(expr)));
197                            }
198                        }
199
200                        Ok(Transformed::no(expr))
201                    })?;
202
203                    if new_pred.transformed {
204                        let exec = FilterExecBuilder::new(new_pred.data, filter.input().clone())
205                            .with_default_selectivity(filter.default_selectivity())
206                            .apply_projection_by_ref(filter.projection().as_ref())
207                            .and_then(|x| x.build())?;
208                        return Ok(Transformed::yes(Arc::new(exec) as _));
209                    }
210                }
211
212                Ok(Transformed::no(plan))
213            })?
214            .data;
215
216        Ok(res)
217    }
218
219    fn name(&self) -> &str {
220        "MatchesConstantTerm"
221    }
222
223    fn schema_check(&self) -> bool {
224        false
225    }
226}
227
228#[cfg(test)]
229mod tests {
230    use std::sync::Arc;
231
232    use arrow::array::{ArrayRef, StringArray, StringDictionaryBuilder};
233    use arrow::datatypes::{DataType, Field, Schema, UInt32Type};
234    use arrow::record_batch::RecordBatch;
235    use catalog::RegisterTableRequest;
236    use catalog::memory::MemoryCatalogManager;
237    use common_catalog::consts::{DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME};
238    use common_function::scalars::matches_term::MatchesTermFunction;
239    use common_function::scalars::udf::create_udf;
240    use datafusion::datasource::memory::MemorySourceConfig;
241    use datafusion::datasource::source::DataSourceExec;
242    use datafusion::physical_optimizer::PhysicalOptimizerRule;
243    use datafusion::physical_plan::filter::FilterExec;
244    use datafusion::physical_plan::get_plan_string;
245    use datafusion_common::{Column, DFSchema};
246    use datafusion_expr::expr::ScalarFunction;
247    use datafusion_expr::physical_planning_context::PhysicalPlanningContext;
248    use datafusion_expr::{Expr, Literal, ScalarUDF};
249    use datafusion_physical_expr::{ScalarFunctionExpr, create_physical_expr};
250    use datatypes::prelude::ConcreteDataType;
251    use datatypes::schema::ColumnSchema;
252    use session::context::QueryContext;
253    use table::metadata::{TableInfoBuilder, TableMetaBuilder};
254    use table::test_util::EmptyTable;
255
256    use super::*;
257    use crate::parser::QueryLanguageParser;
258    use crate::{QueryEngineFactory, QueryEngineRef};
259
260    fn create_test_batch() -> RecordBatch {
261        let schema = Schema::new(vec![Field::new("text", DataType::Utf8, true)]);
262
263        let text_array = StringArray::from(vec![
264            Some("hello world"),
265            Some("greeting"),
266            Some("hello there"),
267            None,
268        ]);
269
270        RecordBatch::try_new(Arc::new(schema), vec![Arc::new(text_array) as ArrayRef]).unwrap()
271    }
272
273    fn create_test_engine() -> QueryEngineRef {
274        let table_name = "test".to_string();
275        let columns = vec![
276            ColumnSchema::new(
277                "text".to_string(),
278                ConcreteDataType::string_datatype(),
279                false,
280            ),
281            ColumnSchema::new(
282                "timestamp".to_string(),
283                ConcreteDataType::timestamp_millisecond_datatype(),
284                false,
285            )
286            .with_time_index(true),
287        ];
288
289        let schema = Arc::new(datatypes::schema::Schema::new(columns));
290        let table_meta = TableMetaBuilder::empty()
291            .schema(schema)
292            .primary_key_indices(vec![])
293            .value_indices(vec![0])
294            .next_column_id(2)
295            .build()
296            .unwrap();
297        let table_info = TableInfoBuilder::default()
298            .name(&table_name)
299            .meta(table_meta)
300            .build()
301            .unwrap();
302        let table = EmptyTable::from_table_info(&table_info);
303        let catalog_list = MemoryCatalogManager::with_default_setup();
304        assert!(
305            catalog_list
306                .register_table_sync(RegisterTableRequest {
307                    catalog: DEFAULT_CATALOG_NAME.to_string(),
308                    schema: DEFAULT_SCHEMA_NAME.to_string(),
309                    table_name,
310                    table_id: 1024,
311                    table,
312                })
313                .is_ok()
314        );
315        QueryEngineFactory::new(
316            catalog_list,
317            None,
318            None,
319            None,
320            None,
321            false,
322            Default::default(),
323        )
324        .query_engine()
325    }
326
327    fn matches_term_udf() -> Arc<ScalarUDF> {
328        Arc::new(create_udf(Arc::new(MatchesTermFunction::default())))
329    }
330
331    #[test]
332    fn test_matches_term_optimization() {
333        let batch = create_test_batch();
334
335        // Create a predicate with a constant pattern
336        let predicate = create_physical_expr(
337            &Expr::ScalarFunction(ScalarFunction::new_udf(
338                matches_term_udf(),
339                vec![Expr::Column(Column::from_name("text")), "hello".lit()],
340            )),
341            &DFSchema::try_from(batch.schema().clone()).unwrap(),
342            &Default::default(),
343            &PhysicalPlanningContext::default(),
344        )
345        .unwrap();
346
347        let input = DataSourceExec::from_data_source(
348            MemorySourceConfig::try_new(&[vec![batch.clone()]], batch.schema(), None).unwrap(),
349        );
350        let filter = FilterExec::try_new(predicate, input).unwrap();
351
352        // Apply the optimizer
353        let optimizer = MatchesConstantTermOptimizer;
354        let optimized_plan = optimizer
355            .optimize(Arc::new(filter), &Default::default())
356            .unwrap();
357
358        let optimized_filter = optimized_plan.downcast_ref::<FilterExec>().unwrap();
359        let predicate = optimized_filter.predicate();
360
361        // The predicate should be a PreCompiledMatchesTermExpr
362        assert!(predicate.is::<PreCompiledMatchesTermExpr>());
363    }
364
365    #[test]
366    fn test_precompiled_matches_term_with_dictionary() {
367        let mut text = StringDictionaryBuilder::<UInt32Type>::new();
368        text.append_value("hello world");
369        text.append_value("greeting");
370        text.append_value("hello there");
371        text.append_null();
372        let schema = Arc::new(Schema::new(vec![Field::new(
373            "text",
374            DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
375            true,
376        )]));
377        let batch = RecordBatch::try_new(schema.clone(), vec![Arc::new(text.finish())]).unwrap();
378        let expr = PreCompiledMatchesTermExpr {
379            text: Arc::new(datafusion_physical_expr::expressions::Column::new(
380                "text", 0,
381            )),
382            term: "hello".to_string(),
383            finder: MatchesTermFinder::new("hello"),
384            probes: vec!["hello".to_string()],
385        };
386
387        let result = expr.evaluate(&batch).unwrap().into_array(4).unwrap();
388        assert_eq!(
389            result.as_any().downcast_ref::<BooleanArray>().unwrap(),
390            &BooleanArray::from(vec![Some(true), Some(false), Some(true), None])
391        );
392    }
393
394    #[test]
395    fn test_matches_term_no_optimization() {
396        let batch = create_test_batch();
397
398        // Create a predicate with a non-constant pattern
399        let predicate = create_physical_expr(
400            &Expr::ScalarFunction(ScalarFunction::new_udf(
401                matches_term_udf(),
402                vec![
403                    Expr::Column(Column::from_name("text")),
404                    Expr::Column(Column::from_name("text")),
405                ],
406            )),
407            &DFSchema::try_from(batch.schema().clone()).unwrap(),
408            &Default::default(),
409            &PhysicalPlanningContext::default(),
410        )
411        .unwrap();
412
413        let input = DataSourceExec::from_data_source(
414            MemorySourceConfig::try_new(&[vec![batch.clone()]], batch.schema(), None).unwrap(),
415        );
416        let filter = FilterExec::try_new(predicate, input).unwrap();
417
418        let optimizer = MatchesConstantTermOptimizer;
419        let optimized_plan = optimizer
420            .optimize(Arc::new(filter), &Default::default())
421            .unwrap();
422
423        let optimized_filter = optimized_plan.downcast_ref::<FilterExec>().unwrap();
424        let predicate = optimized_filter.predicate();
425
426        // The predicate should still be a ScalarFunctionExpr
427        assert!(predicate.is::<ScalarFunctionExpr>());
428    }
429
430    #[tokio::test]
431    async fn test_matches_term_optimization_from_sql() {
432        let sql = "WITH base AS (
433        SELECT text, timestamp FROM test 
434        WHERE MATCHES_TERM(text, 'hello wo_rld') 
435        AND timestamp > '2025-01-01 00:00:00'
436    ),
437    subquery1 AS (
438        SELECT * FROM base 
439        WHERE MATCHES_TERM(text, 'world')
440    ),
441    subquery2 AS (
442        SELECT * FROM test 
443        WHERE MATCHES_TERM(text, 'greeting') 
444        AND timestamp < '2025-01-02 00:00:00'
445    ),
446    union_result AS (
447        SELECT * FROM subquery1 
448        UNION ALL 
449        SELECT * FROM subquery2
450    ),
451    joined_data AS (
452        SELECT a.text, a.timestamp, b.text as other_text 
453        FROM union_result a 
454        JOIN test b ON a.timestamp = b.timestamp 
455        WHERE MATCHES_TERM(a.text, 'there')
456    )
457    SELECT text, other_text 
458    FROM joined_data 
459    WHERE MATCHES_TERM(text, '42') 
460    AND MATCHES_TERM(other_text, 'foo')";
461
462        let query_ctx = QueryContext::arc();
463
464        let stmt = QueryLanguageParser::parse_sql(sql, &query_ctx).unwrap();
465        let engine = create_test_engine();
466        let logical_plan = engine
467            .planner()
468            .plan(&stmt, query_ctx.clone())
469            .await
470            .unwrap();
471
472        let engine_ctx = engine.engine_context(query_ctx);
473        let state = engine_ctx.state();
474
475        let analyzed_plan = state
476            .analyzer()
477            .execute_and_check(logical_plan.clone(), state.config_options(), |_, _| {})
478            .unwrap();
479
480        let optimized_plan = state
481            .optimizer()
482            .optimize(analyzed_plan, state, |_, _| {})
483            .unwrap();
484
485        let physical_plan = state
486            .query_planner()
487            .create_physical_plan(&optimized_plan, state)
488            .await
489            .unwrap();
490
491        let plan_str = get_plan_string(&physical_plan).join("\n");
492        assert!(plan_str.contains("MatchesConstTerm(text@0, term: \"foo\", probes: [\"foo\"]"));
493        assert!(plan_str.contains(
494            "MatchesConstTerm(text@0, term: \"hello wo_rld\", probes: [\"hello\", \"wo_rld\"]"
495        ));
496        assert!(plan_str.contains("MatchesConstTerm(text@0, term: \"world\", probes: [\"world\"]"));
497        assert!(
498            plan_str
499                .contains("MatchesConstTerm(text@0, term: \"greeting\", probes: [\"greeting\"]")
500        );
501        assert!(plan_str.contains("MatchesConstTerm(text@0, term: \"there\", probes: [\"there\"]"));
502        assert!(plan_str.contains("MatchesConstTerm(text@0, term: \"42\", probes: [\"42\"]"));
503        assert!(!plan_str.contains("matches_term"))
504    }
505}