1use std::collections::HashMap;
88use std::fmt;
89use std::marker::PhantomData;
90use std::sync::Arc;
91
92use substrait::proto::NamedStruct;
93use substrait::proto::r#type::{Nullability, Struct};
94use thiserror::Error;
95
96use crate::extensions::any::{Any, AnyRef};
97use crate::extensions::args::{ExtensionArgs, ExtensionColumn, ExtensionValueKind};
98
99#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
101pub enum ExtensionType {
102 Relation,
104 ExtensionTable,
106 Enhancement,
108 Optimization,
110}
111
112#[derive(Debug, Clone, Copy, PartialEq, Eq)]
114pub struct ExtensionInput {
115 emitted_column_count: usize,
116}
117
118impl ExtensionInput {
119 pub(crate) fn new(emitted_column_count: usize) -> Self {
120 Self {
121 emitted_column_count,
122 }
123 }
124
125 pub fn emitted_column_count(&self) -> usize {
127 self.emitted_column_count
128 }
129}
130
131#[derive(Debug, Clone, Copy, Default)]
137pub struct ExtensionContext<'a> {
138 inputs: &'a [ExtensionInput],
139}
140
141impl<'a> ExtensionContext<'a> {
142 pub(crate) fn new(inputs: &'a [ExtensionInput]) -> Self {
143 Self { inputs }
144 }
145
146 pub fn inputs(&self) -> &'a [ExtensionInput] {
148 self.inputs
149 }
150}
151
152#[derive(Debug, Error, Clone)]
154pub enum RegistrationError {
155 #[error("{ext_type:?} extension '{name}' already registered")]
156 DuplicateName {
157 ext_type: ExtensionType,
158 name: String,
159 },
160
161 #[error("Type URL '{type_url}' already registered to {ext_type:?} extension '{existing_name}'")]
162 ConflictingTypeUrl {
163 type_url: String,
164 ext_type: ExtensionType,
165 existing_name: String,
166 },
167}
168
169#[derive(Debug, Error, Clone)]
171pub enum ExtensionError {
172 #[error("Extension '{name}' not found in registry")]
174 NotFound { name: String },
175
176 #[error("Missing required argument: {name}")]
178 MissingArgument { name: String },
179
180 #[error("Invalid argument: expected {expected}, got {actual}")]
182 InvalidArgumentType {
183 expected: ExtensionValueKind,
184 actual: ExtensionValueKind,
185 },
186
187 #[error("Invalid argument: {0}")]
193 InvalidArgument(String),
194
195 #[error("Type URL mismatch: expected {expected}, got {actual}")]
197 TypeUrlMismatch { expected: String, actual: String },
198
199 #[error("Failed to decode protobuf message")]
201 DecodeFailed(#[source] prost::DecodeError),
202
203 #[error("Failed to encode protobuf message")]
205 EncodeFailed(#[source] prost::EncodeError),
206
207 #[error("Extension detail is missing")]
209 MissingDetail,
210
211 #[error("{0}")]
213 Custom(String),
214}
215
216pub trait AnyConvertible: Sized {
220 fn to_any(&self) -> Result<Any, ExtensionError>;
222
223 fn from_any<'a>(any: AnyRef<'a>) -> Result<Self, ExtensionError>;
225
226 fn type_url() -> String;
230}
231
232impl<T> AnyConvertible for T
234where
235 T: prost::Message + prost::Name + Default,
236{
237 fn to_any(&self) -> Result<Any, ExtensionError> {
238 Any::encode(self)
239 }
240
241 fn from_any<'a>(any: AnyRef<'a>) -> Result<Self, ExtensionError> {
242 any.decode()
243 }
244
245 fn type_url() -> String {
246 T::type_url()
247 }
248}
249
250pub trait ExtensionProtoConvert<T> {
261 fn convert(&self) -> Result<T, ExtensionError>;
263}
264
265impl ExtensionProtoConvert<NamedStruct> for [ExtensionColumn] {
266 fn convert(&self) -> Result<NamedStruct, ExtensionError> {
267 let mut names = Vec::with_capacity(self.len());
268 let mut types = Vec::with_capacity(self.len());
269 for col in self {
270 match col {
271 ExtensionColumn::Named { name, r#type: ty } => {
272 names.push(name.clone());
273 types.push(ty.clone());
274 }
275 other => {
276 return Err(ExtensionError::InvalidArgument(format!(
277 "Expected named column, got {other:?}"
278 )));
279 }
280 }
281 }
282 Ok(NamedStruct {
283 names,
284 r#struct: Some(Struct {
285 types,
286 type_variation_reference: 0,
287 nullability: Nullability::Required as i32,
291 }),
292 })
293 }
294}
295
296impl ExtensionProtoConvert<Vec<ExtensionColumn>> for NamedStruct {
297 fn convert(&self) -> Result<Vec<ExtensionColumn>, ExtensionError> {
298 let types = self
299 .r#struct
300 .as_ref()
301 .map(|s| s.types.as_slice())
302 .unwrap_or_default();
303 if self.names.len() != types.len() {
304 return Err(ExtensionError::InvalidArgument(format!(
305 "NamedStruct has {} names but {} types",
306 self.names.len(),
307 types.len()
308 )));
309 }
310 Ok(self
311 .names
312 .iter()
313 .zip(types.iter())
314 .map(|(name, ty)| ExtensionColumn::Named {
315 name: name.clone(),
316 r#type: ty.clone(),
317 })
318 .collect())
319 }
320}
321
322pub trait Explainable: Sized {
324 fn name() -> &'static str;
327
328 fn from_args(args: &ExtensionArgs) -> Result<Self, ExtensionError>;
330
331 fn to_args(&self, context: &ExtensionContext<'_>) -> Result<ExtensionArgs, ExtensionError>;
333}
334
335trait ExtensionConverter: Send + Sync {
349 fn parse_detail(&self, args: &ExtensionArgs) -> Result<Any, ExtensionError>;
350
351 fn textify_detail(
352 &self,
353 detail: AnyRef<'_>,
354 context: &ExtensionContext<'_>,
355 ) -> Result<ExtensionArgs, ExtensionError>;
356}
357
358struct ExtensionAdapter<T>(PhantomData<T>);
375
376impl<T: AnyConvertible + Explainable + Send + Sync> ExtensionConverter for ExtensionAdapter<T> {
377 fn parse_detail(&self, args: &ExtensionArgs) -> Result<Any, ExtensionError> {
378 T::from_args(args)?.to_any()
379 }
380
381 fn textify_detail(
382 &self,
383 detail: AnyRef<'_>,
384 context: &ExtensionContext<'_>,
385 ) -> Result<ExtensionArgs, ExtensionError> {
386 let owned_any = Any::new(detail.type_url.to_string(), detail.value.to_vec());
387 T::from_any(owned_any.as_ref())?.to_args(context)
388 }
389}
390
391pub trait Extension: AnyConvertible + Explainable + Send + Sync + 'static {}
392
393impl<T> Extension for T where T: AnyConvertible + Explainable + Send + Sync + 'static {}
394
395#[derive(Default, Clone)]
397pub struct ExtensionRegistry {
398 handlers: HashMap<(ExtensionType, String), Arc<dyn ExtensionConverter>>,
400 type_urls: HashMap<(ExtensionType, String), String>,
402 descriptors: Vec<Vec<u8>>,
408}
409
410impl ExtensionRegistry {
411 pub fn new() -> Self {
413 Self {
414 handlers: HashMap::new(),
415 type_urls: HashMap::new(),
416 descriptors: Vec::new(),
417 }
418 }
419
420 pub fn add_descriptor(&mut self, bytes: Vec<u8>) {
428 self.descriptors.push(bytes);
429 }
430
431 pub fn descriptors(&self) -> Vec<&[u8]> {
433 self.descriptors.iter().map(|b| b.as_slice()).collect()
434 }
435
436 fn register<T>(&mut self, ext_type: ExtensionType) -> Result<(), RegistrationError>
438 where
439 T: Extension,
440 {
441 let canonical_name = T::name();
442 let type_url = T::type_url();
443 let handler: Arc<dyn ExtensionConverter> = Arc::new(ExtensionAdapter::<T>(PhantomData));
444
445 let key = (ext_type, canonical_name.to_string());
446 if self.handlers.contains_key(&key) {
447 return Err(RegistrationError::DuplicateName {
448 ext_type,
449 name: canonical_name.to_string(),
450 });
451 }
452
453 let type_url_key = (ext_type, type_url.clone());
455 if let Some(existing) = self.type_urls.get(&type_url_key)
456 && existing != canonical_name
457 {
458 return Err(RegistrationError::ConflictingTypeUrl {
459 type_url,
460 ext_type,
461 existing_name: existing.clone(),
462 });
463 }
464
465 self.handlers.insert(key, Arc::clone(&handler));
467 self.type_urls
468 .insert(type_url_key, canonical_name.to_string());
469 Ok(())
470 }
471
472 pub fn register_relation<T>(&mut self) -> Result<(), RegistrationError>
476 where
477 T: Extension,
478 {
479 self.register::<T>(ExtensionType::Relation)
480 }
481
482 pub fn register_extension_table<T>(&mut self) -> Result<(), RegistrationError>
490 where
491 T: Extension,
492 {
493 self.register::<T>(ExtensionType::ExtensionTable)
494 }
495
496 pub fn register_enhancement<T>(&mut self) -> Result<(), RegistrationError>
503 where
504 T: Extension,
505 {
506 self.register::<T>(ExtensionType::Enhancement)
507 }
508
509 pub fn register_optimization<T>(&mut self) -> Result<(), RegistrationError>
516 where
517 T: Extension,
518 {
519 self.register::<T>(ExtensionType::Optimization)
520 }
521
522 pub fn parse_extension(
524 &self,
525 extension_name: &str,
526 args: &ExtensionArgs,
527 ) -> Result<Any, ExtensionError> {
528 self.parse_with_type(ExtensionType::Relation, extension_name, args)
529 }
530
531 pub fn parse_extension_table(
536 &self,
537 extension_table_name: &str,
538 args: &ExtensionArgs,
539 ) -> Result<Any, ExtensionError> {
540 self.parse_with_type(ExtensionType::ExtensionTable, extension_table_name, args)
541 }
542
543 pub fn parse_enhancement(
548 &self,
549 enhancement_name: &str,
550 args: &ExtensionArgs,
551 ) -> Result<Any, ExtensionError> {
552 self.parse_with_type(ExtensionType::Enhancement, enhancement_name, args)
553 }
554
555 pub fn parse_optimization(
560 &self,
561 optimization_name: &str,
562 args: &ExtensionArgs,
563 ) -> Result<Any, ExtensionError> {
564 self.parse_with_type(ExtensionType::Optimization, optimization_name, args)
565 }
566
567 fn parse_with_type(
569 &self,
570 ext_type: ExtensionType,
571 name: &str,
572 args: &ExtensionArgs,
573 ) -> Result<Any, ExtensionError> {
574 let key = (ext_type, name.to_string());
575 let handler = self
576 .handlers
577 .get(&key)
578 .ok_or_else(|| ExtensionError::NotFound {
579 name: name.to_string(),
580 })?;
581 handler.parse_detail(args)
582 }
583
584 pub fn decode(&self, detail: AnyRef<'_>) -> Result<(String, ExtensionArgs), ExtensionError> {
588 self.decode_with_type(ExtensionType::Relation, detail)
589 }
590
591 pub(crate) fn decode_with_context(
593 &self,
594 detail: AnyRef<'_>,
595 context: &ExtensionContext<'_>,
596 ) -> Result<(String, ExtensionArgs), ExtensionError> {
597 self.decode_with_type_and_context(ExtensionType::Relation, detail, context)
598 }
599
600 pub fn decode_extension_table(
606 &self,
607 detail: AnyRef<'_>,
608 ) -> Result<(String, ExtensionArgs), ExtensionError> {
609 self.decode_with_type(ExtensionType::ExtensionTable, detail)
610 }
611
612 pub fn decode_enhancement(
620 &self,
621 detail: AnyRef<'_>,
622 ) -> Result<(String, ExtensionArgs), ExtensionError> {
623 self.decode_with_type(ExtensionType::Enhancement, detail)
624 }
625
626 pub fn decode_optimization(
634 &self,
635 detail: AnyRef<'_>,
636 ) -> Result<(String, ExtensionArgs), ExtensionError> {
637 self.decode_with_type(ExtensionType::Optimization, detail)
638 }
639
640 fn decode_with_type(
642 &self,
643 ext_type: ExtensionType,
644 detail: AnyRef<'_>,
645 ) -> Result<(String, ExtensionArgs), ExtensionError> {
646 self.decode_with_type_and_context(ext_type, detail, &ExtensionContext::default())
647 }
648
649 fn decode_with_type_and_context(
650 &self,
651 ext_type: ExtensionType,
652 detail: AnyRef<'_>,
653 context: &ExtensionContext<'_>,
654 ) -> Result<(String, ExtensionArgs), ExtensionError> {
655 let type_url_key = (ext_type, detail.type_url.to_string());
657 let extension_name =
658 self.type_urls
659 .get(&type_url_key)
660 .ok_or_else(|| ExtensionError::NotFound {
661 name: detail.type_url.to_string(),
662 })?;
663
664 let name_key = (ext_type, extension_name.clone());
666 let handler = self
667 .handlers
668 .get(&name_key)
669 .ok_or_else(|| ExtensionError::NotFound {
670 name: extension_name.clone(),
671 })?;
672
673 let args = handler.textify_detail(detail, context)?;
674
675 Ok((extension_name.clone(), args))
676 }
677
678 pub fn extension_names(&self, ext_type: ExtensionType) -> Vec<&str> {
680 let mut names: Vec<&str> = self
681 .type_urls
682 .iter()
683 .filter_map(|((t, _), name)| {
684 if *t == ext_type {
685 Some(name.as_str())
686 } else {
687 None
688 }
689 })
690 .collect();
691 names.sort_unstable();
692 names.dedup();
693 names
694 }
695
696 pub fn has_extension(&self, ext_type: ExtensionType, name: &str) -> bool {
698 self.handlers.contains_key(&(ext_type, name.to_string()))
699 }
700}
701
702impl fmt::Debug for ExtensionRegistry {
703 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
704 let mut keys: Vec<_> = self
705 .handlers
706 .keys()
707 .map(|(t, n)| (format!("{t:?}"), n.as_str()))
708 .collect();
709 keys.sort();
710 f.debug_struct("ExtensionRegistry")
711 .field("handlers", &keys)
712 .finish()
713 }
714}
715
716#[cfg(test)]
717mod tests {
718 use super::*;
719 use crate::extensions::ExtensionColumn;
720 use crate::fixtures::parse_type;
721 use crate::textify::expressions::Reference;
722
723 struct TestExtension {
725 path: String,
726 batch_size: i64,
727 }
728
729 impl AnyConvertible for TestExtension {
731 fn to_any(&self) -> Result<Any, ExtensionError> {
732 let json_str = format!(
734 r#"{{"path":"{}","batch_size":{}}}"#,
735 self.path, self.batch_size
736 );
737 Ok(Any::new(Self::type_url(), json_str.into_bytes()))
738 }
739
740 fn type_url() -> String {
741 "test.TestExtension".to_string()
742 }
743
744 fn from_any<'a>(any: AnyRef<'a>) -> Result<Self, ExtensionError> {
745 let json_str = String::from_utf8(any.value.to_vec())
747 .map_err(|e| ExtensionError::Custom(format!("Invalid UTF-8: {e}")))?;
748
749 if json_str.contains("path") && json_str.contains("batch_size") {
751 Ok(TestExtension {
752 path: "test.parquet".to_string(),
753 batch_size: 1024,
754 })
755 } else {
756 Err(ExtensionError::Custom("Missing fields".to_string()))
757 }
758 }
759 }
760
761 impl Explainable for TestExtension {
762 fn name() -> &'static str {
763 "TestExtension"
764 }
765
766 fn from_args(args: &ExtensionArgs) -> Result<Self, ExtensionError> {
767 let mut extractor = args.extractor();
768 let path: String = extractor.expect_named_arg::<&str>("path")?.to_string();
769 let batch_size: i64 = extractor.expect_named_arg("batch_size")?;
770 extractor.check_exhausted()?;
771
772 Ok(TestExtension {
773 path: path.to_string(),
774 batch_size,
775 })
776 }
777
778 fn to_args(&self, context: &ExtensionContext<'_>) -> Result<ExtensionArgs, ExtensionError> {
779 let mut args = ExtensionArgs::default();
780 args.insert("path", self.path.clone());
781 args.insert("batch_size", self.batch_size);
782 if !context.inputs().is_empty() {
783 args.insert("input_count", context.inputs().len() as i64);
784 }
785 Ok(args)
786 }
787 }
788
789 #[test]
790 fn test_extension_registry_basic() {
791 let mut registry = ExtensionRegistry::new();
792
793 assert_eq!(registry.extension_names(ExtensionType::Relation).len(), 0);
795 assert_eq!(
796 registry
797 .extension_names(ExtensionType::ExtensionTable)
798 .len(),
799 0
800 );
801 assert!(!registry.has_extension(ExtensionType::Relation, "TestExtension"));
802
803 registry.register_relation::<TestExtension>().unwrap();
805
806 assert_eq!(registry.extension_names(ExtensionType::Relation).len(), 1);
808 assert!(registry.has_extension(ExtensionType::Relation, "TestExtension"));
809
810 let mut args = ExtensionArgs::default();
812 args.insert("path", "data.parquet");
813 args.insert("batch_size", 2048_i64);
814
815 let any = registry.parse_extension("TestExtension", &args).unwrap();
816 assert_eq!(any.type_url, "test.TestExtension");
817
818 let any_ref = any.as_ref();
819 let result = registry.decode(any_ref).unwrap();
820 assert_eq!(result.0, "TestExtension");
821 assert_eq!(
822 <&str>::try_from(result.1.named.get("path").unwrap()).unwrap(),
823 "test.parquet"
824 );
825 assert!(!result.1.named.contains_key("input_count"));
826 }
827
828 #[test]
829 fn test_extension_registry_decode_with_context() {
830 let mut registry = ExtensionRegistry::new();
831 registry.register_relation::<TestExtension>().unwrap();
832
833 let mut args = ExtensionArgs::default();
834 args.insert("path", "data.parquet");
835 args.insert("batch_size", 2048_i64);
836 let any = registry.parse_extension("TestExtension", &args).unwrap();
837 let inputs = [ExtensionInput::new(4)];
838 let context = ExtensionContext::new(&inputs);
839
840 let (_, decoded_args) = registry
841 .decode_with_context(any.as_ref(), &context)
842 .unwrap();
843 assert_eq!(
844 i64::try_from(decoded_args.named.get("input_count").unwrap()).unwrap(),
845 1
846 );
847 }
848
849 #[test]
850 fn test_extension_table_registry_basic() {
851 let mut registry = ExtensionRegistry::new();
852
853 registry
854 .register_extension_table::<TestExtension>()
855 .unwrap();
856
857 assert_eq!(
858 registry.extension_names(ExtensionType::ExtensionTable),
859 vec!["TestExtension"]
860 );
861 assert!(registry.has_extension(ExtensionType::ExtensionTable, "TestExtension"));
862
863 let mut args = ExtensionArgs::default();
864 args.insert("path", "data.parquet");
865 args.insert("batch_size", 2048_i64);
866
867 let any = registry
868 .parse_extension_table("TestExtension", &args)
869 .unwrap();
870 assert_eq!(any.type_url, "test.TestExtension");
871
872 let (name, decoded_args) = registry.decode_extension_table(any.as_ref()).unwrap();
873 assert_eq!(name, "TestExtension");
874 assert_eq!(
875 <&str>::try_from(decoded_args.named.get("path").unwrap()).unwrap(),
876 "test.parquet"
877 );
878 assert!(!decoded_args.named.contains_key("input_count"));
879 }
880
881 #[test]
882 fn test_extension_args() {
883 let mut args = ExtensionArgs::default();
884
885 args.insert("path", "data/*.parquet");
887 args.insert("batch_size", 1024_i64);
888
889 args.push(Reference(0));
891
892 args.output_columns.push(ExtensionColumn::Named {
894 name: "col1".to_string(),
895 r#type: parse_type("i32"),
896 });
897
898 let mut extractor = args.extractor();
900
901 let path = extractor.get_named_arg("path").unwrap();
902 assert_eq!(<&str>::try_from(path).unwrap(), "data/*.parquet");
903
904 let batch_size = extractor.get_named_arg("batch_size").unwrap();
905 assert_eq!(i64::try_from(batch_size).unwrap(), 1024);
906
907 assert!(extractor.check_exhausted().is_ok());
909
910 assert_eq!(args.positional.len(), 1);
911 assert_eq!(args.output_columns.len(), 1);
912 }
913
914 #[test]
915 fn test_extension_error_cases() {
916 let registry = ExtensionRegistry::new();
917
918 let args = ExtensionArgs::default();
920 let result = registry.parse_extension("NonExistent", &args);
921 assert!(matches!(result, Err(ExtensionError::NotFound { .. })));
922
923 let args = ExtensionArgs::default();
925 let mut extractor = args.extractor();
926 let result = extractor.get_named_arg("missing");
927 assert!(result.is_none());
928 assert!(extractor.check_exhausted().is_ok());
929
930 let mut args = ExtensionArgs::default();
932 args.insert("test", 42_i64);
933 let mut extractor = args.extractor();
934 let result = extractor.get_named_arg("test");
935 assert_eq!(i64::try_from(result.unwrap()).unwrap(), 42);
936 assert!(extractor.check_exhausted().is_ok());
937 }
938
939 struct TestEnhancement {
941 hint: String,
942 }
943
944 impl AnyConvertible for TestEnhancement {
945 fn to_any(&self) -> Result<Any, ExtensionError> {
946 let json_str = format!(r#"{{"hint":"{}"}}"#, self.hint);
947 Ok(Any::new(Self::type_url(), json_str.into_bytes()))
948 }
949
950 fn type_url() -> String {
951 "test.TestExtension".to_string()
953 }
954
955 fn from_any<'a>(any: AnyRef<'a>) -> Result<Self, ExtensionError> {
956 let json_str = String::from_utf8(any.value.to_vec())
957 .map_err(|e| ExtensionError::Custom(format!("Invalid UTF-8: {e}")))?;
958 if json_str.contains("hint") {
959 Ok(TestEnhancement {
960 hint: "test_hint".to_string(),
961 })
962 } else {
963 Err(ExtensionError::Custom("Missing hint field".to_string()))
964 }
965 }
966 }
967
968 impl Explainable for TestEnhancement {
969 fn name() -> &'static str {
970 "TestEnhancement"
971 }
972
973 fn from_args(args: &ExtensionArgs) -> Result<Self, ExtensionError> {
974 let mut extractor = args.extractor();
975 let hint: String = extractor.expect_named_arg::<&str>("hint")?.to_string();
976 extractor.check_exhausted()?;
977 Ok(TestEnhancement { hint })
978 }
979
980 fn to_args(
981 &self,
982 _context: &ExtensionContext<'_>,
983 ) -> Result<ExtensionArgs, ExtensionError> {
984 let mut args = ExtensionArgs::default();
985 args.insert("hint", self.hint.clone());
986 Ok(args)
987 }
988 }
989
990 #[test]
991 fn test_namespace_separation() {
992 let mut registry = ExtensionRegistry::new();
993
994 registry.register_relation::<TestExtension>().unwrap();
996 registry
997 .register_extension_table::<TestExtension>()
998 .unwrap();
999 registry.register_enhancement::<TestEnhancement>().unwrap();
1000
1001 assert!(registry.has_extension(ExtensionType::Relation, "TestExtension"));
1003 assert!(registry.has_extension(ExtensionType::ExtensionTable, "TestExtension"));
1004 assert!(registry.has_extension(ExtensionType::Enhancement, "TestEnhancement"));
1005 assert_eq!(registry.extension_names(ExtensionType::Relation).len(), 1);
1006 assert_eq!(
1007 registry
1008 .extension_names(ExtensionType::ExtensionTable)
1009 .len(),
1010 1
1011 );
1012 assert_eq!(
1013 registry.extension_names(ExtensionType::Enhancement).len(),
1014 1
1015 );
1016
1017 let mut ext_args = ExtensionArgs::default();
1019 ext_args.insert("path", "data.parquet");
1020 ext_args.insert("batch_size", 2048_i64);
1021
1022 let ext_any = registry
1023 .parse_extension("TestExtension", &ext_args)
1024 .unwrap();
1025 assert_eq!(ext_any.type_url, "test.TestExtension");
1026
1027 let table_any = registry
1029 .parse_extension_table("TestExtension", &ext_args)
1030 .unwrap();
1031 assert_eq!(table_any.type_url, "test.TestExtension");
1032
1033 let mut enh_args = ExtensionArgs::default();
1035 enh_args.insert("hint", "optimize");
1036
1037 let enh_any = registry
1038 .parse_enhancement("TestEnhancement", &enh_args)
1039 .unwrap();
1040 assert_eq!(enh_any.type_url, "test.TestExtension"); let enh_ref = enh_any.as_ref();
1044 let (name, args) = registry.decode_enhancement(enh_ref).unwrap();
1045 assert_eq!(name, "TestEnhancement");
1046 assert_eq!(
1047 <&str>::try_from(args.named.get("hint").unwrap()).unwrap(),
1048 "test_hint"
1049 );
1050 }
1051
1052 #[test]
1053 fn test_enhancement_duplicate_registration_returns_error() {
1054 let mut registry = ExtensionRegistry::new();
1055 registry.register_enhancement::<TestEnhancement>().unwrap();
1056 let result = registry.register_enhancement::<TestEnhancement>();
1057 assert!(matches!(
1058 result,
1059 Err(RegistrationError::DuplicateName { .. })
1060 ));
1061 }
1062
1063 #[test]
1064 fn test_extension_table_duplicate_registration_returns_error() {
1065 let mut registry = ExtensionRegistry::new();
1066 registry
1067 .register_extension_table::<TestExtension>()
1068 .unwrap();
1069 let result = registry.register_extension_table::<TestExtension>();
1070 assert!(matches!(
1071 result,
1072 Err(RegistrationError::DuplicateName { .. })
1073 ));
1074 }
1075
1076 #[test]
1077 fn test_extension_table_not_found_error() {
1078 let registry = ExtensionRegistry::new();
1079 let args = ExtensionArgs::default();
1080 let result = registry.parse_extension_table("NonExistentExtensionTable", &args);
1081 assert!(matches!(result, Err(ExtensionError::NotFound { .. })));
1082 }
1083
1084 #[test]
1085 fn test_enhancement_not_found_error() {
1086 let registry = ExtensionRegistry::new();
1087 let args = ExtensionArgs::default();
1088 let result = registry.parse_enhancement("NonExistentEnhancement", &args);
1089 assert!(matches!(result, Err(ExtensionError::NotFound { .. })));
1090 }
1091
1092 struct ConflictingExtension;
1095
1096 impl AnyConvertible for ConflictingExtension {
1097 fn to_any(&self) -> Result<Any, ExtensionError> {
1098 Ok(Any::new(Self::type_url(), vec![]))
1099 }
1100
1101 fn type_url() -> String {
1102 "test.TestExtension".to_string()
1104 }
1105
1106 fn from_any<'a>(_any: AnyRef<'a>) -> Result<Self, ExtensionError> {
1107 Ok(ConflictingExtension)
1108 }
1109 }
1110
1111 impl Explainable for ConflictingExtension {
1112 fn name() -> &'static str {
1113 "ConflictingExtension"
1114 }
1115
1116 fn from_args(_args: &ExtensionArgs) -> Result<Self, ExtensionError> {
1117 Ok(ConflictingExtension)
1118 }
1119
1120 fn to_args(
1121 &self,
1122 _context: &ExtensionContext<'_>,
1123 ) -> Result<ExtensionArgs, ExtensionError> {
1124 Ok(ExtensionArgs::default())
1125 }
1126 }
1127
1128 #[test]
1129 fn test_conflicting_type_url_leaves_registry_unchanged() {
1130 let mut registry = ExtensionRegistry::new();
1131 registry.register_relation::<TestExtension>().unwrap();
1132
1133 let result = registry.register_relation::<ConflictingExtension>();
1135 assert!(matches!(
1136 result,
1137 Err(RegistrationError::ConflictingTypeUrl { .. })
1138 ));
1139
1140 assert!(registry.has_extension(ExtensionType::Relation, "TestExtension"));
1142 assert!(!registry.has_extension(ExtensionType::Relation, "ConflictingExtension"));
1143 assert_eq!(
1144 registry.extension_names(ExtensionType::Relation),
1145 vec!["TestExtension"]
1146 );
1147 }
1148
1149 #[test]
1150 fn test_extension_table_conflicting_type_url_leaves_registry_unchanged() {
1151 let mut registry = ExtensionRegistry::new();
1152 registry
1153 .register_extension_table::<TestExtension>()
1154 .unwrap();
1155
1156 let result = registry.register_extension_table::<ConflictingExtension>();
1158 assert!(matches!(
1159 result,
1160 Err(RegistrationError::ConflictingTypeUrl { .. })
1161 ));
1162
1163 assert!(registry.has_extension(ExtensionType::ExtensionTable, "TestExtension"));
1165 assert!(!registry.has_extension(ExtensionType::ExtensionTable, "ConflictingExtension"));
1166 assert_eq!(
1167 registry.extension_names(ExtensionType::ExtensionTable),
1168 vec!["TestExtension"]
1169 );
1170 }
1171}