Skip to main content

common_function/
function_registry.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
15//! functions registry
16use 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/// The result of registering a function.
62#[derive(Debug, Clone, Copy, PartialEq, Eq)]
63pub enum FunctionRegistrationResult {
64    /// The function was newly registered.
65    Registered,
66    /// A function with the same name was already registered and was kept.
67    AlreadyExists,
68}
69
70impl FunctionRegistry {
71    /// Register a function in the registry by converting it into a `ScalarFunctionFactory`.
72    ///
73    /// # Arguments
74    ///
75    /// * `func` - An object that can be converted into a `ScalarFunctionFactory`.
76    ///
77    /// The function is inserted into the internal function map, keyed by its name.
78    /// If a function with the same name already exists, it will be replaced.
79    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    /// Register a function only if no function with the same name exists.
89    ///
90    /// The duplicate check and the insert happen atomically under the same
91    /// write lock of the functions map. If a function with the same name
92    /// already exists, it is kept unchanged and
93    /// [`FunctionRegistrationResult::AlreadyExists`] is returned; otherwise the
94    /// function is registered and [`FunctionRegistrationResult::Registered`] is
95    /// returned.
96    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    /// Register a scalar function in the registry.
112    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    /// Register an aggregate function in the registry.
128    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    /// Register a table function
137    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    /// Register a function rewrite rule.
146    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    /// Register a window function (UDWF).
151    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    /// Returns a list of all scalar functions registered in the registry.
164    pub fn scalar_functions(&self) -> Vec<ScalarFunctionFactory> {
165        self.functions.read().unwrap().values().cloned().collect()
166    }
167
168    /// Returns a list of all aggregate functions registered in the registry.
169    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    /// Returns a list of all window functions registered in the registry.
188    pub fn window_functions(&self) -> Vec<WindowUDF> {
189        self.window_functions
190            .read()
191            .unwrap()
192            .values()
193            .cloned()
194            .collect()
195    }
196
197    /// Returns a registered aggregate function by name.
198    pub fn get_aggr_func(&self, name: &str) -> Option<AggregateUDF> {
199        self.aggregate_functions.read().unwrap().get(name).cloned()
200    }
201
202    /// Returns true if an aggregate function with the given name exists in the registry.
203    pub fn is_aggr_func_exist(&self, name: &str) -> bool {
204        self.aggregate_functions.read().unwrap().contains_key(name)
205    }
206
207    /// Returns a list of all function rewrite rules registered in the registry.
208    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    // Utility functions
217    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    // Full text search function
229    MatchesFunction::register(&function_registry);
230    MatchesTermFunction::register(&function_registry);
231    #[cfg(feature = "ai_functions")]
232    ai::register(&function_registry);
233
234    // System and administration functions
235    SystemFunction::register(&function_registry);
236    AdminFunction::register(&function_registry);
237
238    // Json related functions
239    JsonFunction::register(&function_registry);
240
241    // String related functions
242    register_string_functions(&function_registry);
243
244    // Vector related functions
245    VectorScalarFunction::register(&function_registry);
246    VectorAggrFunction::register(&function_registry);
247
248    // Geo functions
249    #[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    // Ip functions
255    IpFunctions::register(&function_registry);
256
257    // Approximate functions
258    ApproximateFunction::register(&function_registry);
259
260    // CountHash function
261    CountHash::register(&function_registry);
262
263    // state function of supported aggregate functions
264    StateMergeHelper::register(&function_registry);
265
266    // Anomaly detection window functions
267    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(&registry);
275    registry
276});
277
278/// Returns a function that is only available to the ADMIN statement executor.
279pub fn get_admin_function(name: &str) -> Option<ScalarFunctionFactory> {
280    ADMIN_FUNCTION_REGISTRY.get_function(name)
281}
282
283/// Register a function that is only available to the ADMIN statement executor.
284///
285/// If a function with the same name is already registered in the ADMIN
286/// registry, the existing one is kept and
287/// [`FunctionRegistrationResult::AlreadyExists`] is returned. A name that
288/// already exists in the normal [`FUNCTION_REGISTRY`] when this call
289/// linearizes is also rejected: the ADMIN executor resolves admin-only
290/// functions before falling back to the normal registry, so inserting such a
291/// name here would shadow the built-in. Otherwise the function is registered
292/// and [`FunctionRegistrationResult::Registered`] is returned.
293///
294/// The enforced contract is one-way: it only guards the ADMIN registration
295/// against names already present in the normal registry. A later ordinary
296/// [`FunctionRegistry::register`] may still install the same name in the
297/// normal registry because the normal registry keeps its legacy replace
298/// semantics.
299pub fn register_admin_function(
300    func: impl Into<ScalarFunctionFactory>,
301) -> FunctionRegistrationResult {
302    register_admin_function_in(&ADMIN_FUNCTION_REGISTRY, &FUNCTION_REGISTRY, func)
303}
304
305/// Core implementation of [`register_admin_function`] against a pair of
306/// registries, parameterized so tests can exercise it with local registries.
307///
308/// Locking: the ADMIN-registry write lock is acquired first, then a read lock
309/// on the normal registry, and the normal-registry guard (bound to
310/// `normal_functions`) is kept alive through both the normal-name check and
311/// the ADMIN insertion below. This is the only code path that holds both
312/// registries' locks, so the ADMIN -> FUNCTION acquisition order is
313/// consistent and a concurrent normal-registry registration cannot slip in
314/// between the check and the ADMIN insert and be shadowed.
315///
316/// The enforced contract is one-way: it only guards the ADMIN registration
317/// against names already present in the normal registry. A later ordinary
318/// [`FunctionRegistry::register`] may still install the same name in the
319/// normal registry because the normal registry keeps its legacy replace
320/// semantics.
321fn 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    // The normal-registry guard is a read lock: it is held across the
329    // normal-name check and the ADMIN insertion below, and while it is alive
330    // no writer can acquire the normal-registry write lock, so a concurrent
331    // normal-registry registration cannot slip in between the check and the
332    // ADMIN insert and be shadowed.
333    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 the read guard only after the ADMIN insertion, so writers to the
346    // normal registry stay blocked until the check-and-insert is complete.
347    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    /// Creates a [`ScalarFunctionFactory`] with the given name. Each call
361    /// allocates a distinct factory closure, so factories can be told apart by
362    /// [`Arc::ptr_eq`] on their `factory` field even when names are identical.
363    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(&registered.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        // Regression test: `register` keeps its replace semantics.
442        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(&registry);
467                let barrier = Arc::clone(&barrier);
468                thread::spawn(move || {
469                    let factory = named_factory(name);
470                    // Synchronize so every thread attempts registration at the
471                    // same time; only one may win the write lock.
472                    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        // Tests touching the global registry must use unique names.
514        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        // The exact same factory clone registered twice: the first registration
535        // wins, the duplicate is rejected, and the stored factory is
536        // pointer-equal to the original. Tests touching the global registry
537        // must use unique names.
538        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        // The guard is one-way: after the ADMIN registration completes, an
557        // ordinary normal-registry registration of the same name still
558        // succeeds because the normal registry keeps its legacy replace
559        // semantics.
560        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        // Regression test: the ADMIN executor resolves `get_admin_function`
586        // before falling back to `FUNCTION_REGISTRY`, so registering a function
587        // whose name already exists in the normal registry would shadow the
588        // ADMIN-invocable built-in (e.g. `flush_table`). Such registrations
589        // must be rejected with
590        // [`FunctionRegistrationResult::AlreadyExists`] and must not be
591        // inserted into the ADMIN registry.
592        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        // Built-in admin-only functions registered at startup stay queryable
611        // through the same global registry used for runtime registrations.
612        #[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}