1use 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#[derive(Debug)]
40pub struct PreCompiledMatchesTermExpr {
41 text: Arc<dyn PhysicalExpr>,
43 term: String,
45 finder: MatchesTermFinder,
47
48 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#[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 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 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 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 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 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 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}