1use 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;
31#[cfg(feature = "ai_functions")]
32use crate::scalars::ai;
33use crate::scalars::anomaly::AnomalyFunction;
34use crate::scalars::avg_calc::AvgCalcFunction;
35use crate::scalars::date::DateFunction;
36use crate::scalars::expression::ExpressionFunction;
37use crate::scalars::hll_count::HllCalcFunction;
38use crate::scalars::ip::IpFunctions;
39use crate::scalars::json::JsonFunction;
40use crate::scalars::matches::MatchesFunction;
41use crate::scalars::matches_term::MatchesTermFunction;
42use crate::scalars::math::MathFunction;
43use crate::scalars::primary_key::DecodePrimaryKeyFunction;
44use crate::scalars::string::register_string_functions;
45use crate::scalars::timestamp::TimestampFunction;
46use crate::scalars::uddsketch_calc::UddSketchCalcFunction;
47use crate::scalars::uddsketch_rank::UddSketchRankFunction;
48use crate::scalars::vector::VectorFunction as VectorScalarFunction;
49use crate::scalars::welford_stddev::WelfordStddevFunction;
50use crate::system::SystemFunction;
51
52#[derive(Default)]
53pub struct FunctionRegistry {
54 functions: RwLock<HashMap<String, ScalarFunctionFactory>>,
55 aggregate_functions: RwLock<HashMap<String, AggregateUDF>>,
56 table_functions: RwLock<HashMap<String, Arc<TableFunction>>>,
57 function_rewrites: RwLock<Vec<Arc<dyn FunctionRewrite + Send + Sync>>>,
58 window_functions: RwLock<HashMap<String, WindowUDF>>,
59}
60
61#[derive(Debug, Clone, Copy, PartialEq, Eq)]
63pub enum FunctionRegistrationResult {
64 Registered,
66 AlreadyExists,
68}
69
70impl FunctionRegistry {
71 pub fn register(&self, func: impl Into<ScalarFunctionFactory>) {
80 let func = func.into();
81 let _ = self
82 .functions
83 .write()
84 .unwrap()
85 .insert(func.name().to_string(), func);
86 }
87
88 pub fn register_if_absent(
97 &self,
98 func: impl Into<ScalarFunctionFactory>,
99 ) -> FunctionRegistrationResult {
100 let func = func.into();
101 let mut functions = self.functions.write().unwrap();
102 match functions.entry(func.name().to_string()) {
103 Entry::Occupied(_) => FunctionRegistrationResult::AlreadyExists,
104 Entry::Vacant(entry) => {
105 entry.insert(func);
106 FunctionRegistrationResult::Registered
107 }
108 }
109 }
110
111 pub fn register_scalar(&self, func: impl Function + 'static) {
113 let func = Arc::new(func) as FunctionRef;
114
115 for alias in func.aliases() {
116 let func: ScalarFunctionFactory = func.clone().into();
117 let alias = ScalarFunctionFactory {
118 name: alias.clone(),
119 ..func
120 };
121 self.register(alias);
122 }
123
124 self.register(func)
125 }
126
127 pub fn register_aggr(&self, func: AggregateUDF) {
129 let _ = self
130 .aggregate_functions
131 .write()
132 .unwrap()
133 .insert(func.name().to_string(), func);
134 }
135
136 pub fn register_table_function(&self, func: TableFunction) {
138 let _ = self
139 .table_functions
140 .write()
141 .unwrap()
142 .insert(func.name().to_string(), Arc::new(func));
143 }
144
145 pub fn register_function_rewrite(&self, func: impl FunctionRewrite + Send + Sync + 'static) {
147 self.function_rewrites.write().unwrap().push(Arc::new(func));
148 }
149
150 pub fn register_window(&self, func: WindowUDF) {
152 let _ = self
153 .window_functions
154 .write()
155 .unwrap()
156 .insert(func.name().to_string(), func);
157 }
158
159 pub fn get_function(&self, name: &str) -> Option<ScalarFunctionFactory> {
160 self.functions.read().unwrap().get(name).cloned()
161 }
162
163 pub fn scalar_functions(&self) -> Vec<ScalarFunctionFactory> {
165 self.functions.read().unwrap().values().cloned().collect()
166 }
167
168 pub fn aggregate_functions(&self) -> Vec<AggregateUDF> {
170 self.aggregate_functions
171 .read()
172 .unwrap()
173 .values()
174 .cloned()
175 .collect()
176 }
177
178 pub fn table_functions(&self) -> Vec<Arc<TableFunction>> {
179 self.table_functions
180 .read()
181 .unwrap()
182 .values()
183 .cloned()
184 .collect()
185 }
186
187 pub fn window_functions(&self) -> Vec<WindowUDF> {
189 self.window_functions
190 .read()
191 .unwrap()
192 .values()
193 .cloned()
194 .collect()
195 }
196
197 pub fn get_aggr_func(&self, name: &str) -> Option<AggregateUDF> {
199 self.aggregate_functions.read().unwrap().get(name).cloned()
200 }
201
202 pub fn is_aggr_func_exist(&self, name: &str) -> bool {
204 self.aggregate_functions.read().unwrap().contains_key(name)
205 }
206
207 pub fn function_rewrites(&self) -> Vec<Arc<dyn FunctionRewrite + Send + Sync>> {
209 self.function_rewrites.read().unwrap().clone()
210 }
211}
212
213pub static FUNCTION_REGISTRY: LazyLock<Arc<FunctionRegistry>> = LazyLock::new(|| {
214 let function_registry = FunctionRegistry::default();
215
216 MathFunction::register(&function_registry);
218 TimestampFunction::register(&function_registry);
219 DateFunction::register(&function_registry);
220 ExpressionFunction::register(&function_registry);
221 AvgCalcFunction::register(&function_registry);
222 UddSketchCalcFunction::register(&function_registry);
223 UddSketchRankFunction::register(&function_registry);
224 HllCalcFunction::register(&function_registry);
225 WelfordStddevFunction::register(&function_registry);
226 DecodePrimaryKeyFunction::register(&function_registry);
227
228 MatchesFunction::register(&function_registry);
230 MatchesTermFunction::register(&function_registry);
231 #[cfg(feature = "ai_functions")]
232 ai::register(&function_registry);
233
234 SystemFunction::register(&function_registry);
236 AdminFunction::register(&function_registry);
237
238 JsonFunction::register(&function_registry);
240
241 register_string_functions(&function_registry);
243
244 VectorScalarFunction::register(&function_registry);
246 VectorAggrFunction::register(&function_registry);
247
248 #[cfg(feature = "geo")]
250 crate::scalars::geo::GeoFunctions::register(&function_registry);
251 #[cfg(feature = "geo")]
252 crate::aggrs::geo::GeoFunction::register(&function_registry);
253
254 IpFunctions::register(&function_registry);
256
257 ApproximateFunction::register(&function_registry);
259
260 CountHash::register(&function_registry);
262
263 StateMergeHelper::register(&function_registry);
265
266 AnomalyFunction::register(&function_registry);
268
269 Arc::new(function_registry)
270});
271
272static ADMIN_FUNCTION_REGISTRY: LazyLock<FunctionRegistry> = LazyLock::new(|| {
273 let registry = FunctionRegistry::default();
274 AdminFunction::register_admin_only(®istry);
275 registry
276});
277
278pub fn get_admin_function(name: &str) -> Option<ScalarFunctionFactory> {
280 ADMIN_FUNCTION_REGISTRY.get_function(name)
281}
282
283pub fn register_admin_function(
300 func: impl Into<ScalarFunctionFactory>,
301) -> FunctionRegistrationResult {
302 register_admin_function_in(&ADMIN_FUNCTION_REGISTRY, &FUNCTION_REGISTRY, func)
303}
304
305fn register_admin_function_in(
322 admin_registry: &FunctionRegistry,
323 normal_registry: &FunctionRegistry,
324 func: impl Into<ScalarFunctionFactory>,
325) -> FunctionRegistrationResult {
326 let func = func.into();
327 let mut admin_functions = admin_registry.functions.write().unwrap();
328 let normal_functions = normal_registry.functions.read().unwrap();
334 if normal_functions.contains_key(func.name()) {
335 drop(normal_functions);
336 return FunctionRegistrationResult::AlreadyExists;
337 }
338 let result = match admin_functions.entry(func.name().to_string()) {
339 Entry::Occupied(_) => FunctionRegistrationResult::AlreadyExists,
340 Entry::Vacant(entry) => {
341 entry.insert(func);
342 FunctionRegistrationResult::Registered
343 }
344 };
345 drop(normal_functions);
348 result
349}
350
351#[cfg(test)]
352mod tests {
353 use std::sync::{Arc, Barrier};
354 use std::thread;
355
356 use super::*;
357 use crate::scalars::test::TestAndFunction;
358 use crate::scalars::udf::create_udf;
359
360 fn named_factory(name: &str) -> ScalarFunctionFactory {
364 ScalarFunctionFactory {
365 name: name.to_string(),
366 factory: Arc::new(|_ctx| create_udf(Arc::new(TestAndFunction::default()))),
367 }
368 }
369
370 #[test]
371 fn test_function_registry() {
372 let registry = FunctionRegistry::default();
373
374 assert!(registry.get_function("test_and").is_none());
375 assert!(registry.scalar_functions().is_empty());
376 registry.register_scalar(TestAndFunction::default());
377 let _ = registry.get_function("test_and").unwrap();
378 assert_eq!(1, registry.scalar_functions().len());
379 }
380
381 #[test]
382 fn test_uddsketch_rank_registered() {
383 assert!(FUNCTION_REGISTRY.get_function("uddsketch_rank").is_some());
384 }
385
386 #[test]
387 fn test_ai_registration_matches_feature() {
388 for name in ["ai_match", "ai_choose", "ai_score"] {
389 assert_eq!(
390 FUNCTION_REGISTRY.get_function(name).is_some(),
391 cfg!(feature = "ai_functions"),
392 "{name}"
393 );
394 }
395 for name in ["jev", "jev_choice", "jev_score"] {
396 assert!(FUNCTION_REGISTRY.get_function(name).is_none(), "{name}");
397 }
398 }
399
400 #[test]
401 fn test_register_if_absent_registers_new_function() {
402 let registry = FunctionRegistry::default();
403 let name = "pr3_register_if_absent_new";
404 let factory = named_factory(name);
405
406 assert_eq!(
407 registry.register_if_absent(factory.clone()),
408 FunctionRegistrationResult::Registered
409 );
410 let registered = registry
411 .get_function(name)
412 .expect("function should be registered");
413 assert!(Arc::ptr_eq(®istered.factory, &factory.factory));
414 }
415
416 #[test]
417 fn test_register_if_absent_first_registration_wins() {
418 let registry = FunctionRegistry::default();
419 let name = "pr3_register_if_absent_duplicate";
420 let first = named_factory(name);
421 let second = named_factory(name);
422
423 assert_eq!(
424 registry.register_if_absent(first.clone()),
425 FunctionRegistrationResult::Registered
426 );
427 assert_eq!(
428 registry.register_if_absent(second.clone()),
429 FunctionRegistrationResult::AlreadyExists
430 );
431
432 let stored = registry
433 .get_function(name)
434 .expect("function should be registered");
435 assert!(Arc::ptr_eq(&stored.factory, &first.factory));
436 assert!(!Arc::ptr_eq(&stored.factory, &second.factory));
437 }
438
439 #[test]
440 fn test_register_replaces_existing_function() {
441 let registry = FunctionRegistry::default();
443 let name = "pr3_register_replaces";
444 let first = named_factory(name);
445 let second = named_factory(name);
446
447 registry.register(first.clone());
448 registry.register(second.clone());
449
450 let stored = registry
451 .get_function(name)
452 .expect("function should be registered");
453 assert!(Arc::ptr_eq(&stored.factory, &second.factory));
454 assert!(!Arc::ptr_eq(&stored.factory, &first.factory));
455 }
456
457 #[test]
458 fn test_concurrent_register_if_absent_same_name() {
459 const THREADS: usize = 8;
460 let registry = Arc::new(FunctionRegistry::default());
461 let name = "pr3_concurrent_same_name";
462 let barrier = Arc::new(Barrier::new(THREADS));
463
464 let handles: Vec<_> = (0..THREADS)
465 .map(|_| {
466 let registry = Arc::clone(®istry);
467 let barrier = Arc::clone(&barrier);
468 thread::spawn(move || {
469 let factory = named_factory(name);
470 barrier.wait();
473 let result = registry.register_if_absent(factory.clone());
474 (result, factory)
475 })
476 })
477 .collect();
478
479 let mut results: Vec<(FunctionRegistrationResult, ScalarFunctionFactory)> =
480 Vec::with_capacity(THREADS);
481 for handle in handles {
482 results.push(handle.join().expect("thread should not panic"));
483 }
484
485 let registered = results
486 .iter()
487 .filter(|(result, _)| *result == FunctionRegistrationResult::Registered)
488 .count();
489 let already_exists = results
490 .iter()
491 .filter(|(result, _)| *result == FunctionRegistrationResult::AlreadyExists)
492 .count();
493 assert_eq!(registered, 1);
494 assert_eq!(already_exists, THREADS - 1);
495
496 let winner = results
497 .iter()
498 .find(|(result, _)| *result == FunctionRegistrationResult::Registered)
499 .map(|(_, factory)| factory)
500 .expect("exactly one registration must win");
501
502 let stored = registry
503 .get_function(name)
504 .expect("function should be registered");
505 assert!(
506 Arc::ptr_eq(&stored.factory, &winner.factory),
507 "the stored factory must be the factory of the winning registration"
508 );
509 }
510
511 #[test]
512 fn test_register_admin_function_first_wins() {
513 let name = "pr3_admin_runtime_register";
515 let first = named_factory(name);
516 let second = named_factory(name);
517
518 assert_eq!(
519 register_admin_function(first.clone()),
520 FunctionRegistrationResult::Registered
521 );
522 assert_eq!(
523 register_admin_function(second.clone()),
524 FunctionRegistrationResult::AlreadyExists
525 );
526
527 let stored = get_admin_function(name).expect("admin function should be queryable");
528 assert!(Arc::ptr_eq(&stored.factory, &first.factory));
529 assert!(!Arc::ptr_eq(&stored.factory, &second.factory));
530 }
531
532 #[test]
533 fn test_register_admin_function_duplicate_same_factory() {
534 let name = "pr3_admin_runtime_register_same_factory";
539 let factory = named_factory(name);
540
541 assert_eq!(
542 register_admin_function(factory.clone()),
543 FunctionRegistrationResult::Registered
544 );
545 assert_eq!(
546 register_admin_function(factory.clone()),
547 FunctionRegistrationResult::AlreadyExists
548 );
549
550 let stored = get_admin_function(name).expect("admin function should be queryable");
551 assert!(Arc::ptr_eq(&stored.factory, &factory.factory));
552 }
553
554 #[test]
555 fn test_register_admin_function_in_is_one_way_later_normal_register_allowed() {
556 let admin = FunctionRegistry::default();
561 let normal = FunctionRegistry::default();
562 let name = "pr3_admin_in_one_way";
563 let admin_factory = named_factory(name);
564 let normal_factory = named_factory(name);
565
566 assert_eq!(
567 register_admin_function_in(&admin, &normal, admin_factory.clone()),
568 FunctionRegistrationResult::Registered
569 );
570
571 normal.register(normal_factory.clone());
572
573 let stored_admin = admin
574 .get_function(name)
575 .expect("the ADMIN registration must be kept");
576 assert!(Arc::ptr_eq(&stored_admin.factory, &admin_factory.factory));
577 let stored_normal = normal
578 .get_function(name)
579 .expect("the later normal registration must succeed");
580 assert!(Arc::ptr_eq(&stored_normal.factory, &normal_factory.factory));
581 }
582
583 #[test]
584 fn test_register_admin_function_rejects_normal_registry_builtin_name() {
585 let factory = named_factory("flush_table");
593
594 assert_eq!(
595 register_admin_function(factory.clone()),
596 FunctionRegistrationResult::AlreadyExists
597 );
598 assert!(
599 get_admin_function("flush_table").is_none(),
600 "a normal-registry built-in must not be shadowed into the ADMIN registry"
601 );
602 assert!(
603 FUNCTION_REGISTRY.get_function("flush_table").is_some(),
604 "the normal-registry built-in must remain registered"
605 );
606 }
607
608 #[test]
609 fn test_builtin_admin_functions_remain_queryable() {
610 #[cfg(feature = "enterprise")]
613 {
614 assert!(get_admin_function("purge_table").is_some());
615 }
616 #[cfg(not(feature = "enterprise"))]
617 {
618 assert!(get_admin_function("purge_table").is_none());
619 }
620 }
621}