common_function/
function_registry.rs1use std::collections::HashMap;
17use std::collections::hash_map::Entry;
18use std::sync::{Arc, LazyLock, RwLock};
19
20use datafusion::catalog::TableFunction;
21use datafusion_expr::expr_rewriter::FunctionRewrite;
22use datafusion_expr::{AggregateUDF, WindowUDF};
23
24use crate::admin::AdminFunction;
25use crate::aggrs::aggr_wrapper::StateMergeHelper;
26use crate::aggrs::approximate::ApproximateFunction;
27use crate::aggrs::count_hash::CountHash;
28use crate::aggrs::vector::VectorFunction as VectorAggrFunction;
29use crate::function::{Function, FunctionRef};
30use crate::function_factory::ScalarFunctionFactory;
31use crate::scalars::anomaly::AnomalyFunction;
32use crate::scalars::date::DateFunction;
33use crate::scalars::expression::ExpressionFunction;
34use crate::scalars::hll_count::HllCalcFunction;
35use crate::scalars::ip::IpFunctions;
36use crate::scalars::json::JsonFunction;
37use crate::scalars::matches::MatchesFunction;
38use crate::scalars::matches_term::MatchesTermFunction;
39use crate::scalars::math::MathFunction;
40use crate::scalars::primary_key::DecodePrimaryKeyFunction;
41use crate::scalars::string::register_string_functions;
42use crate::scalars::timestamp::TimestampFunction;
43use crate::scalars::uddsketch_calc::UddSketchCalcFunction;
44use crate::scalars::vector::VectorFunction as VectorScalarFunction;
45use crate::system::SystemFunction;
46
47#[derive(Default)]
48pub struct FunctionRegistry {
49 functions: RwLock<HashMap<String, ScalarFunctionFactory>>,
50 aggregate_functions: RwLock<HashMap<String, AggregateUDF>>,
51 table_functions: RwLock<HashMap<String, Arc<TableFunction>>>,
52 function_rewrites: RwLock<Vec<Arc<dyn FunctionRewrite + Send + Sync>>>,
53 window_functions: RwLock<HashMap<String, WindowUDF>>,
54}
55
56#[derive(Debug, Clone, Copy, PartialEq, Eq)]
58pub enum FunctionRegistrationResult {
59 Registered,
61 AlreadyExists,
63}
64
65impl FunctionRegistry {
66 pub fn register(&self, func: impl Into<ScalarFunctionFactory>) {
75 let func = func.into();
76 let _ = self
77 .functions
78 .write()
79 .unwrap()
80 .insert(func.name().to_string(), func);
81 }
82
83 pub fn register_if_absent(
92 &self,
93 func: impl Into<ScalarFunctionFactory>,
94 ) -> FunctionRegistrationResult {
95 let func = func.into();
96 let mut functions = self.functions.write().unwrap();
97 match functions.entry(func.name().to_string()) {
98 Entry::Occupied(_) => FunctionRegistrationResult::AlreadyExists,
99 Entry::Vacant(entry) => {
100 entry.insert(func);
101 FunctionRegistrationResult::Registered
102 }
103 }
104 }
105
106 pub fn register_scalar(&self, func: impl Function + 'static) {
108 let func = Arc::new(func) as FunctionRef;
109
110 for alias in func.aliases() {
111 let func: ScalarFunctionFactory = func.clone().into();
112 let alias = ScalarFunctionFactory {
113 name: alias.clone(),
114 ..func
115 };
116 self.register(alias);
117 }
118
119 self.register(func)
120 }
121
122 pub fn register_aggr(&self, func: AggregateUDF) {
124 let _ = self
125 .aggregate_functions
126 .write()
127 .unwrap()
128 .insert(func.name().to_string(), func);
129 }
130
131 pub fn register_table_function(&self, func: TableFunction) {
133 let _ = self
134 .table_functions
135 .write()
136 .unwrap()
137 .insert(func.name().to_string(), Arc::new(func));
138 }
139
140 pub fn register_function_rewrite(&self, func: impl FunctionRewrite + Send + Sync + 'static) {
142 self.function_rewrites.write().unwrap().push(Arc::new(func));
143 }
144
145 pub fn register_window(&self, func: WindowUDF) {
147 let _ = self
148 .window_functions
149 .write()
150 .unwrap()
151 .insert(func.name().to_string(), func);
152 }
153
154 pub fn get_function(&self, name: &str) -> Option<ScalarFunctionFactory> {
155 self.functions.read().unwrap().get(name).cloned()
156 }
157
158 pub fn scalar_functions(&self) -> Vec<ScalarFunctionFactory> {
160 self.functions.read().unwrap().values().cloned().collect()
161 }
162
163 pub fn aggregate_functions(&self) -> Vec<AggregateUDF> {
165 self.aggregate_functions
166 .read()
167 .unwrap()
168 .values()
169 .cloned()
170 .collect()
171 }
172
173 pub fn table_functions(&self) -> Vec<Arc<TableFunction>> {
174 self.table_functions
175 .read()
176 .unwrap()
177 .values()
178 .cloned()
179 .collect()
180 }
181
182 pub fn window_functions(&self) -> Vec<WindowUDF> {
184 self.window_functions
185 .read()
186 .unwrap()
187 .values()
188 .cloned()
189 .collect()
190 }
191
192 pub fn is_aggr_func_exist(&self, name: &str) -> bool {
194 self.aggregate_functions.read().unwrap().contains_key(name)
195 }
196
197 pub fn function_rewrites(&self) -> Vec<Arc<dyn FunctionRewrite + Send + Sync>> {
199 self.function_rewrites.read().unwrap().clone()
200 }
201}
202
203pub static FUNCTION_REGISTRY: LazyLock<Arc<FunctionRegistry>> = LazyLock::new(|| {
204 let function_registry = FunctionRegistry::default();
205
206 MathFunction::register(&function_registry);
208 TimestampFunction::register(&function_registry);
209 DateFunction::register(&function_registry);
210 ExpressionFunction::register(&function_registry);
211 UddSketchCalcFunction::register(&function_registry);
212 HllCalcFunction::register(&function_registry);
213 DecodePrimaryKeyFunction::register(&function_registry);
214
215 MatchesFunction::register(&function_registry);
217 MatchesTermFunction::register(&function_registry);
218
219 SystemFunction::register(&function_registry);
221 AdminFunction::register(&function_registry);
222
223 JsonFunction::register(&function_registry);
225
226 register_string_functions(&function_registry);
228
229 VectorScalarFunction::register(&function_registry);
231 VectorAggrFunction::register(&function_registry);
232
233 #[cfg(feature = "geo")]
235 crate::scalars::geo::GeoFunctions::register(&function_registry);
236 #[cfg(feature = "geo")]
237 crate::aggrs::geo::GeoFunction::register(&function_registry);
238
239 IpFunctions::register(&function_registry);
241
242 ApproximateFunction::register(&function_registry);
244
245 CountHash::register(&function_registry);
247
248 StateMergeHelper::register(&function_registry);
250
251 AnomalyFunction::register(&function_registry);
253
254 Arc::new(function_registry)
255});
256
257static ADMIN_FUNCTION_REGISTRY: LazyLock<FunctionRegistry> = LazyLock::new(|| {
258 let registry = FunctionRegistry::default();
259 AdminFunction::register_admin_only(®istry);
260 registry
261});
262
263pub fn get_admin_function(name: &str) -> Option<ScalarFunctionFactory> {
265 ADMIN_FUNCTION_REGISTRY.get_function(name)
266}
267
268pub fn register_admin_function(
285 func: impl Into<ScalarFunctionFactory>,
286) -> FunctionRegistrationResult {
287 register_admin_function_in(&ADMIN_FUNCTION_REGISTRY, &FUNCTION_REGISTRY, func)
288}
289
290fn register_admin_function_in(
307 admin_registry: &FunctionRegistry,
308 normal_registry: &FunctionRegistry,
309 func: impl Into<ScalarFunctionFactory>,
310) -> FunctionRegistrationResult {
311 let func = func.into();
312 let mut admin_functions = admin_registry.functions.write().unwrap();
313 let normal_functions = normal_registry.functions.read().unwrap();
319 if normal_functions.contains_key(func.name()) {
320 drop(normal_functions);
321 return FunctionRegistrationResult::AlreadyExists;
322 }
323 let result = match admin_functions.entry(func.name().to_string()) {
324 Entry::Occupied(_) => FunctionRegistrationResult::AlreadyExists,
325 Entry::Vacant(entry) => {
326 entry.insert(func);
327 FunctionRegistrationResult::Registered
328 }
329 };
330 drop(normal_functions);
333 result
334}
335
336#[cfg(test)]
337mod tests {
338 use std::sync::{Arc, Barrier};
339 use std::thread;
340
341 use super::*;
342 use crate::scalars::test::TestAndFunction;
343 use crate::scalars::udf::create_udf;
344
345 fn named_factory(name: &str) -> ScalarFunctionFactory {
349 ScalarFunctionFactory {
350 name: name.to_string(),
351 factory: Arc::new(|_ctx| create_udf(Arc::new(TestAndFunction::default()))),
352 }
353 }
354
355 #[test]
356 fn test_function_registry() {
357 let registry = FunctionRegistry::default();
358
359 assert!(registry.get_function("test_and").is_none());
360 assert!(registry.scalar_functions().is_empty());
361 registry.register_scalar(TestAndFunction::default());
362 let _ = registry.get_function("test_and").unwrap();
363 assert_eq!(1, registry.scalar_functions().len());
364 }
365
366 #[test]
367 fn test_register_if_absent_registers_new_function() {
368 let registry = FunctionRegistry::default();
369 let name = "pr3_register_if_absent_new";
370 let factory = named_factory(name);
371
372 assert_eq!(
373 registry.register_if_absent(factory.clone()),
374 FunctionRegistrationResult::Registered
375 );
376 let registered = registry
377 .get_function(name)
378 .expect("function should be registered");
379 assert!(Arc::ptr_eq(®istered.factory, &factory.factory));
380 }
381
382 #[test]
383 fn test_register_if_absent_first_registration_wins() {
384 let registry = FunctionRegistry::default();
385 let name = "pr3_register_if_absent_duplicate";
386 let first = named_factory(name);
387 let second = named_factory(name);
388
389 assert_eq!(
390 registry.register_if_absent(first.clone()),
391 FunctionRegistrationResult::Registered
392 );
393 assert_eq!(
394 registry.register_if_absent(second.clone()),
395 FunctionRegistrationResult::AlreadyExists
396 );
397
398 let stored = registry
399 .get_function(name)
400 .expect("function should be registered");
401 assert!(Arc::ptr_eq(&stored.factory, &first.factory));
402 assert!(!Arc::ptr_eq(&stored.factory, &second.factory));
403 }
404
405 #[test]
406 fn test_register_replaces_existing_function() {
407 let registry = FunctionRegistry::default();
409 let name = "pr3_register_replaces";
410 let first = named_factory(name);
411 let second = named_factory(name);
412
413 registry.register(first.clone());
414 registry.register(second.clone());
415
416 let stored = registry
417 .get_function(name)
418 .expect("function should be registered");
419 assert!(Arc::ptr_eq(&stored.factory, &second.factory));
420 assert!(!Arc::ptr_eq(&stored.factory, &first.factory));
421 }
422
423 #[test]
424 fn test_concurrent_register_if_absent_same_name() {
425 const THREADS: usize = 8;
426 let registry = Arc::new(FunctionRegistry::default());
427 let name = "pr3_concurrent_same_name";
428 let barrier = Arc::new(Barrier::new(THREADS));
429
430 let handles: Vec<_> = (0..THREADS)
431 .map(|_| {
432 let registry = Arc::clone(®istry);
433 let barrier = Arc::clone(&barrier);
434 thread::spawn(move || {
435 let factory = named_factory(name);
436 barrier.wait();
439 let result = registry.register_if_absent(factory.clone());
440 (result, factory)
441 })
442 })
443 .collect();
444
445 let mut results: Vec<(FunctionRegistrationResult, ScalarFunctionFactory)> =
446 Vec::with_capacity(THREADS);
447 for handle in handles {
448 results.push(handle.join().expect("thread should not panic"));
449 }
450
451 let registered = results
452 .iter()
453 .filter(|(result, _)| *result == FunctionRegistrationResult::Registered)
454 .count();
455 let already_exists = results
456 .iter()
457 .filter(|(result, _)| *result == FunctionRegistrationResult::AlreadyExists)
458 .count();
459 assert_eq!(registered, 1);
460 assert_eq!(already_exists, THREADS - 1);
461
462 let winner = results
463 .iter()
464 .find(|(result, _)| *result == FunctionRegistrationResult::Registered)
465 .map(|(_, factory)| factory)
466 .expect("exactly one registration must win");
467
468 let stored = registry
469 .get_function(name)
470 .expect("function should be registered");
471 assert!(
472 Arc::ptr_eq(&stored.factory, &winner.factory),
473 "the stored factory must be the factory of the winning registration"
474 );
475 }
476
477 #[test]
478 fn test_register_admin_function_first_wins() {
479 let name = "pr3_admin_runtime_register";
481 let first = named_factory(name);
482 let second = named_factory(name);
483
484 assert_eq!(
485 register_admin_function(first.clone()),
486 FunctionRegistrationResult::Registered
487 );
488 assert_eq!(
489 register_admin_function(second.clone()),
490 FunctionRegistrationResult::AlreadyExists
491 );
492
493 let stored = get_admin_function(name).expect("admin function should be queryable");
494 assert!(Arc::ptr_eq(&stored.factory, &first.factory));
495 assert!(!Arc::ptr_eq(&stored.factory, &second.factory));
496 }
497
498 #[test]
499 fn test_register_admin_function_duplicate_same_factory() {
500 let name = "pr3_admin_runtime_register_same_factory";
505 let factory = named_factory(name);
506
507 assert_eq!(
508 register_admin_function(factory.clone()),
509 FunctionRegistrationResult::Registered
510 );
511 assert_eq!(
512 register_admin_function(factory.clone()),
513 FunctionRegistrationResult::AlreadyExists
514 );
515
516 let stored = get_admin_function(name).expect("admin function should be queryable");
517 assert!(Arc::ptr_eq(&stored.factory, &factory.factory));
518 }
519
520 #[test]
521 fn test_register_admin_function_in_is_one_way_later_normal_register_allowed() {
522 let admin = FunctionRegistry::default();
527 let normal = FunctionRegistry::default();
528 let name = "pr3_admin_in_one_way";
529 let admin_factory = named_factory(name);
530 let normal_factory = named_factory(name);
531
532 assert_eq!(
533 register_admin_function_in(&admin, &normal, admin_factory.clone()),
534 FunctionRegistrationResult::Registered
535 );
536
537 normal.register(normal_factory.clone());
538
539 let stored_admin = admin
540 .get_function(name)
541 .expect("the ADMIN registration must be kept");
542 assert!(Arc::ptr_eq(&stored_admin.factory, &admin_factory.factory));
543 let stored_normal = normal
544 .get_function(name)
545 .expect("the later normal registration must succeed");
546 assert!(Arc::ptr_eq(&stored_normal.factory, &normal_factory.factory));
547 }
548
549 #[test]
550 fn test_register_admin_function_rejects_normal_registry_builtin_name() {
551 let factory = named_factory("flush_table");
559
560 assert_eq!(
561 register_admin_function(factory.clone()),
562 FunctionRegistrationResult::AlreadyExists
563 );
564 assert!(
565 get_admin_function("flush_table").is_none(),
566 "a normal-registry built-in must not be shadowed into the ADMIN registry"
567 );
568 assert!(
569 FUNCTION_REGISTRY.get_function("flush_table").is_some(),
570 "the normal-registry built-in must remain registered"
571 );
572 }
573
574 #[test]
575 fn test_builtin_admin_functions_remain_queryable() {
576 #[cfg(feature = "enterprise")]
579 {
580 assert!(get_admin_function("purge_table").is_some());
581 }
582 #[cfg(not(feature = "enterprise"))]
583 {
584 assert!(get_admin_function("purge_table").is_none());
585 }
586 }
587}