1use std::collections::{HashMap, hash_map::Entry};
14
15use super::{ExtensionFile, SimpleExtensions, SimpleExtensionsError, types::CustomType};
16use crate::urn::Urn;
17
18#[derive(Debug)]
23pub struct Registry {
24 extensions: HashMap<Urn, SimpleExtensions>,
26}
27
28impl Registry {
29 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 pub fn extensions(&self) -> impl Iterator<Item = (&Urn, &SimpleExtensions)> {
49 self.extensions.iter()
50 }
51
52 #[cfg(feature = "extensions")]
57 pub fn from_core_extensions() -> Self {
58 use crate::extensions::EXTENSIONS;
59
60 let extensions: HashMap<Urn, SimpleExtensions> = EXTENSIONS
62 .iter()
63 .filter_map(|(orig_urn, simple_extensions)| {
64 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 pub fn get_type(&self, urn: &Urn, name: &str) -> Option<&CustomType> {
89 self.get_extension(urn)?.get_type(name)
90 }
91
92 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 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 let type_via_registry = registry.get_type(&urn, "geometry");
201 assert!(type_via_registry.is_some());
202
203 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 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(), )),
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 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![], );
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 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 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}