Skip to main content

substrait/parse/text/simple_extensions/
registry.rs

1// SPDX-License-Identifier: Apache-2.0
2
3//! Substrait Extension Registry
4//!
5//! This module provides registries for Substrait extensions:
6//! - **Global Registry**: Immutable, reusable across plans, URI+name based lookup
7//! - **Local Registry**: Per-plan, anchor-based, references Global Registry (TODO)
8//!
9//! Currently only type definitions are supported. Function parsing will be added in a future update.
10//!
11//! This module is only available when the `parse` feature is enabled.
12
13use std::collections::{HashMap, hash_map::Entry};
14
15use super::{ExtensionFile, SimpleExtensions, SimpleExtensionsError, types::CustomType};
16use crate::urn::Urn;
17
18/// Extension Registry that manages Substrait extensions
19///
20/// This registry is immutable and reusable across multiple plans.
21/// It provides URN + name based lookup for extension types. Function parsing will be added in a future update.
22#[derive(Debug)]
23pub struct Registry {
24    /// Pre-validated extension files
25    extensions: HashMap<Urn, SimpleExtensions>,
26}
27
28impl Registry {
29    /// Create a new Global Registry from validated extension files.
30    ///
31    /// Any duplicate URNs will raise an error.
32    pub fn new<I: IntoIterator<Item = ExtensionFile>>(
33        extensions: I,
34    ) -> Result<Self, SimpleExtensionsError> {
35        let mut map = HashMap::new();
36        for ExtensionFile { urn, extension } in extensions {
37            match map.entry(urn.clone()) {
38                Entry::Occupied(_) => return Err(SimpleExtensionsError::DuplicateUrn(urn)),
39                Entry::Vacant(entry) => {
40                    entry.insert(extension);
41                }
42            }
43        }
44        Ok(Self { extensions: map })
45    }
46
47    /// Get an iterator over all extension files in this registry
48    pub fn extensions(&self) -> impl Iterator<Item = (&Urn, &SimpleExtensions)> {
49        self.extensions.iter()
50    }
51
52    /// Create a Global Registry from the built-in core extensions.
53    ///
54    /// Most core extensions are included. Some are skipped due to bugs in the upstream
55    /// YAML files (see <https://github.com/substrait-io/substrait/issues/935>).
56    #[cfg(feature = "extensions")]
57    pub fn from_core_extensions() -> Self {
58        use crate::extensions::EXTENSIONS;
59
60        // Parse the core extensions from the raw extensions format to the parsed format
61        let extensions: HashMap<Urn, SimpleExtensions> = EXTENSIONS
62            .iter()
63            .filter_map(|(orig_urn, simple_extensions)| {
64                // Skip specific core extensions that have bugs (missing u! prefix on type references).
65                // Most core extensions are included; only these problematic ones are filtered out.
66                // See: https://github.com/substrait-io/substrait/issues/935
67                let urn_str = orig_urn.to_string();
68                if urn_str == "extension:io.substrait:extension_types" ||
69                   urn_str == "extension:io.substrait:unknown" {
70                    return None;
71                }
72
73                let ExtensionFile { urn, extension } = ExtensionFile::create(simple_extensions.clone())
74                    .unwrap_or_else(|err| panic!("Core extensions should be valid, but failed to create extension file for {orig_urn}: {err}"));
75                debug_assert_eq!(orig_urn, &urn);
76                Some((urn, extension))
77            })
78            .collect();
79
80        Self { extensions }
81    }
82
83    fn get_extension(&self, urn: &Urn) -> Option<&SimpleExtensions> {
84        self.extensions.get(urn)
85    }
86
87    /// Get a type by URN and name
88    pub fn get_type(&self, urn: &Urn, name: &str) -> Option<&CustomType> {
89        self.get_extension(urn)?.get_type(name)
90    }
91
92    /// Get a scalar function by URN and name.
93    ///
94    /// TODO: Add support for retrieving functions by their full signature shorthand
95    /// (e.g., "add:i32_i32").
96    pub fn get_scalar_function(&self, urn: &Urn, name: &str) -> Option<&super::ScalarFunction> {
97        self.get_extension(urn)?.get_scalar_function(name)
98    }
99}
100
101#[cfg(test)]
102mod tests {
103    use super::{ExtensionFile, Registry};
104    use crate::parse::text::simple_extensions::{
105        SimpleExtensionsError, scalar_functions::ScalarFunctionError, types::ExtensionTypeError,
106    };
107    use crate::text::simple_extensions::{SimpleExtensions, SimpleExtensionsTypesItem};
108    use crate::urn::Urn;
109    use std::str::FromStr;
110
111    fn extension_file(urn: &str, type_names: &[&str]) -> ExtensionFile {
112        let types = type_names
113            .iter()
114            .map(|name| SimpleExtensionsTypesItem {
115                name: (*name).to_string(),
116                deprecated: None,
117                description: None,
118                metadata: Default::default(),
119                parameters: None,
120                structure: None,
121                variadic: None,
122            })
123            .collect();
124
125        let raw = SimpleExtensions {
126            scalar_functions: vec![],
127            aggregate_functions: vec![],
128            window_functions: vec![],
129            dependencies: Default::default(),
130            metadata: Default::default(),
131            type_variations: vec![],
132            types,
133            urn: urn.to_string(),
134        };
135
136        ExtensionFile::create(raw).expect("valid extension file")
137    }
138
139    #[test]
140    fn test_registry_iteration() {
141        let urns = vec![
142            "extension:example.com:first",
143            "extension:example.com:second",
144        ];
145        let registry =
146            Registry::new(urns.iter().map(|&urn| extension_file(urn, &["type"]))).unwrap();
147
148        let collected: Vec<&Urn> = registry.extensions().map(|(urn, _)| urn).collect();
149        assert_eq!(collected.len(), 2);
150        for urn in urns {
151            assert!(
152                collected
153                    .iter()
154                    .any(|candidate| candidate.to_string() == urn)
155            );
156        }
157    }
158
159    #[test]
160    fn test_type_lookup() {
161        let urn = Urn::from_str("extension:example.com:test").unwrap();
162        let registry =
163            Registry::new(vec![extension_file(&urn.to_string(), &["test_type"])]).unwrap();
164        let other_urn = Urn::from_str("extension:example.com:other").unwrap();
165
166        let cases = vec![
167            (&urn, "test_type", true),
168            (&urn, "missing", false),
169            (&other_urn, "test_type", false),
170        ];
171
172        for (query_urn, type_name, expected) in cases {
173            assert_eq!(
174                registry.get_type(query_urn, type_name).is_some(),
175                expected,
176                "unexpected lookup result for {query_urn}:{type_name}"
177            );
178        }
179    }
180
181    #[cfg(feature = "extensions")]
182    #[test]
183    fn test_from_core_extensions() {
184        let registry = Registry::from_core_extensions();
185        assert!(registry.extensions().count() > 0);
186
187        // Test that functions_geometry.yaml loaded correctly with its geometry type
188        let urn = Urn::from_str("extension:io.substrait:functions_geometry").unwrap();
189        let core_extension = registry
190            .get_extension(&urn)
191            .expect("Should find functions_geometry extension");
192
193        let geometry_type = core_extension.get_type("geometry");
194        assert!(
195            geometry_type.is_some(),
196            "Should find 'geometry' type in functions_geometry extension"
197        );
198
199        // Also test the registry's get_type method with the actual URN
200        let type_via_registry = registry.get_type(&urn, "geometry");
201        assert!(type_via_registry.is_some());
202
203        // Verify extension_types is skipped due to u! prefix bug (substrait#935)
204        let extension_types_urn = Urn::from_str("extension:io.substrait:extension_types").unwrap();
205        assert!(
206            registry.get_extension(&extension_types_urn).is_none(),
207            "extension_types should be skipped due to missing u! prefix bug"
208        );
209    }
210
211    #[test]
212    fn test_unknown_type_without_prefix_fails() {
213        use crate::text::simple_extensions;
214
215        // Function that references a type without u! prefix - should fail with UnknownTypeName
216        let invalid_extension = SimpleExtensions {
217            scalar_functions: vec![simple_extensions::ScalarFunction {
218                name: "bad_function".to_string(),
219                description: None,
220                metadata: Default::default(),
221                deprecated: None,
222                impls: vec![simple_extensions::ScalarFunctionImplsItem {
223                    args: None,
224                    deprecated: None,
225                    description: None,
226                    options: None,
227                    variadic: None,
228                    session_dependent: None,
229                    deterministic: None,
230                    nullability: None,
231                    return_: simple_extensions::ReturnValue(simple_extensions::Type::String(
232                        "point".to_string(), // Missing u! prefix - this is an error, not NYI
233                    )),
234                    implementation: None,
235                }],
236            }],
237            aggregate_functions: vec![],
238            window_functions: vec![],
239            dependencies: Default::default(),
240            metadata: Default::default(),
241            type_variations: vec![],
242            types: vec![],
243            urn: "extension:example.com:invalid".to_string(),
244        };
245
246        let result = ExtensionFile::create(invalid_extension);
247        assert!(
248            result.is_err(),
249            "Should fail when type is missing u! prefix"
250        );
251
252        match result {
253            Err(SimpleExtensionsError::ScalarFunctionError(ScalarFunctionError::TypeError(
254                ExtensionTypeError::UnknownTypeName { name },
255            ))) => {
256                assert_eq!(name, "point");
257            }
258            other => panic!("Expected UnknownTypeName error, got {:?}", other),
259        }
260    }
261
262    /// Helper to create a minimal extension with a scalar function returning a custom type
263    fn extension_with_custom_type_reference(
264        urn: &str,
265        function_name: &str,
266        return_type: &str,
267        defined_types: Vec<&str>,
268    ) -> SimpleExtensions {
269        use crate::text::simple_extensions;
270
271        SimpleExtensions {
272            scalar_functions: vec![simple_extensions::ScalarFunction {
273                name: function_name.to_string(),
274                description: None,
275                metadata: Default::default(),
276                deprecated: None,
277                impls: vec![simple_extensions::ScalarFunctionImplsItem {
278                    args: None,
279                    deprecated: None,
280                    description: None,
281                    options: None,
282                    variadic: None,
283                    session_dependent: None,
284                    deterministic: None,
285                    nullability: None,
286                    return_: simple_extensions::ReturnValue(simple_extensions::Type::String(
287                        return_type.to_string(),
288                    )),
289                    implementation: None,
290                }],
291            }],
292            aggregate_functions: vec![],
293            window_functions: vec![],
294            dependencies: Default::default(),
295            metadata: Default::default(),
296            type_variations: vec![],
297            types: defined_types
298                .into_iter()
299                .map(|name| SimpleExtensionsTypesItem {
300                    name: name.to_string(),
301                    deprecated: None,
302                    description: None,
303                    metadata: Default::default(),
304                    parameters: None,
305                    structure: None,
306                    variadic: None,
307                })
308                .collect(),
309            urn: urn.to_string(),
310        }
311    }
312
313    #[test]
314    fn test_custom_type_reference_valid() {
315        let extension = extension_with_custom_type_reference(
316            "extension:example.com:valid",
317            "get_point",
318            "u!point",
319            vec!["point"],
320        );
321
322        let result = ExtensionFile::create(extension);
323        assert!(
324            result.is_ok(),
325            "Should succeed when referenced type exists with u! prefix"
326        );
327    }
328
329    #[test]
330    fn test_custom_type_reference_missing() {
331        let extension = extension_with_custom_type_reference(
332            "extension:example.com:invalid",
333            "get_rectangle",
334            "u!rectangle",
335            vec![], // rectangle type not defined
336        );
337
338        let result = ExtensionFile::create(extension);
339        assert!(
340            result.is_err(),
341            "Should fail when referenced type doesn't exist"
342        );
343
344        match result {
345            Err(SimpleExtensionsError::UnresolvedTypeReference { type_name }) => {
346                assert_eq!(type_name, "rectangle");
347            }
348            other => panic!("Expected UnresolvedTypeReference error, got {:?}", other),
349        }
350    }
351
352    #[cfg(feature = "extensions")]
353    #[test]
354    fn test_scalar_function_parses_completely() {
355        use super::super::{
356            argument::ArgumentsItem,
357            scalar_functions::{Impl, NullabilityHandling, Options},
358            types::*,
359        };
360        use crate::parse::Parse;
361        use crate::text::simple_extensions;
362        use std::collections::HashMap;
363
364        let registry = Registry::from_core_extensions();
365        let functions_arithmetic_urn =
366            Urn::from_str("extension:io.substrait:functions_arithmetic").unwrap();
367
368        let add = registry
369            .get_scalar_function(&functions_arithmetic_urn, "add")
370            .expect("add function should exist");
371
372        // Verify function-level metadata
373        assert_eq!(add.name, "add");
374        assert_eq!(add.description, Some("Add two values.".to_string()));
375        assert!(
376            !add.impls.is_empty(),
377            "add should have at least one implementation"
378        );
379
380        // Create the expected first implementation (i8 + i8 -> i8)
381        let mut ctx = super::super::extensions::TypeContext::default();
382        let expected_impl = Impl {
383            args: vec![
384                ArgumentsItem::ValueArgument(
385                    simple_extensions::ValueArg {
386                        name: Some("x".to_string()),
387                        description: None,
388                        value: simple_extensions::Type::String("i8".to_string()),
389                        constant: None,
390                    }
391                    .parse(&mut ctx)
392                    .unwrap(),
393                ),
394                ArgumentsItem::ValueArgument(
395                    simple_extensions::ValueArg {
396                        name: Some("y".to_string()),
397                        description: None,
398                        value: simple_extensions::Type::String("i8".to_string()),
399                        constant: None,
400                    }
401                    .parse(&mut ctx)
402                    .unwrap(),
403                ),
404            ],
405            options: Options({
406                let mut map = HashMap::new();
407                map.insert(
408                    "overflow".to_string(),
409                    vec![
410                        "SILENT".to_string(),
411                        "SATURATE".to_string(),
412                        "ERROR".to_string(),
413                    ],
414                );
415                map
416            }),
417            variadic: None,
418            session_dependent: false,
419            deterministic: true,
420            nullability: NullabilityHandling::Mirror,
421            return_type: ConcreteType {
422                kind: ConcreteTypeKind::Builtin(BasicBuiltinType::I8),
423                nullable: false,
424            },
425            implementation: HashMap::new(),
426        };
427
428        assert_eq!(&add.impls[0], &expected_impl);
429    }
430}