1use std::{
28 collections::HashSet,
29 convert::Infallible,
30 fmt::Debug,
31 io::{Read, Write},
32 marker::PhantomData,
33};
34
35use serde::{Deserialize, Serialize};
36
37#[cfg(feature = "extensions-draft")]
39mod app_data_dict_extension;
40mod application_id_extension;
41mod codec;
42mod extension_in;
43mod external_pub_extension;
44mod external_sender_extension;
45mod last_resort;
46mod ratchet_tree_extension;
47mod required_capabilities;
48use errors::*;
49
50pub mod errors;
52
53#[cfg(feature = "extensions-draft")]
55pub use app_data_dict_extension::{AppDataDictionary, AppDataDictionaryExtension};
56pub use application_id_extension::ApplicationIdExtension;
57pub use external_pub_extension::ExternalPubExtension;
58pub use external_sender_extension::{
59 ExternalSender, ExternalSendersExtension, SenderExtensionIndex,
60};
61pub use last_resort::LastResortExtension;
62pub use ratchet_tree_extension::RatchetTreeExtension;
63pub use required_capabilities::RequiredCapabilitiesExtension;
64
65use tls_codec::{
66 Deserialize as TlsDeserializeTrait, DeserializeBytes, Error, Serialize as TlsSerializeTrait,
67 Size, TlsDeserialize, TlsSerialize, TlsSize,
68};
69
70use crate::{
71 extensions::extension_in::ExtensionIn, group::GroupContext, key_packages::KeyPackage,
72 messages::group_info::GroupInfo, treesync::LeafNode,
73};
74
75#[cfg(test)]
76mod tests;
77
78#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash, Ord, PartialOrd)]
94#[cfg_attr(
95 feature = "0-8-1-storage-format",
96 derive(serde::Serialize, serde::Deserialize)
97)]
98#[cfg_attr(
99 not(feature = "0-8-1-storage-format"),
100 derive(
101 openmls_serialization_helpers::Serialize,
102 openmls_serialization_helpers::Deserialize,
103 )
104)]
105pub enum ExtensionType {
106 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 0)]
107 ApplicationId,
110
111 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 1)]
112 RatchetTree,
115
116 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 2)]
117 RequiredCapabilities,
120
121 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 3)]
122 ExternalPub,
125
126 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 4)]
127 ExternalSenders,
130
131 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 5)]
132 LastResort,
135
136 #[cfg(feature = "extensions-draft")]
137 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 8)]
138 AppDataDictionary,
140
141 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 7)]
142 Grease(u16),
144
145 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 6)]
146 Unknown(u16),
148}
149
150impl ExtensionType {
151 pub(crate) fn is_default(self) -> bool {
153 match self {
154 ExtensionType::ApplicationId
155 | ExtensionType::RatchetTree
156 | ExtensionType::RequiredCapabilities
157 | ExtensionType::ExternalPub
158 | ExtensionType::ExternalSenders => true,
159 ExtensionType::LastResort | ExtensionType::Grease(_) | ExtensionType::Unknown(_) => {
160 false
161 }
162 #[cfg(feature = "extensions-draft")]
163 ExtensionType::AppDataDictionary => false,
164 }
165 }
166
167 pub(crate) fn is_valid_in_leaf_node(self) -> bool {
171 match self {
172 ExtensionType::LastResort
173 | ExtensionType::RatchetTree
174 | ExtensionType::RequiredCapabilities
175 | ExtensionType::ExternalPub
176 | ExtensionType::ExternalSenders => false,
177 ExtensionType::Grease(_) | ExtensionType::Unknown(_) | ExtensionType::ApplicationId => {
182 true
183 }
184 #[cfg(feature = "extensions-draft")]
185 ExtensionType::AppDataDictionary => true,
186 }
187 }
188 pub(crate) fn is_valid_in_group_info(self) -> Option<bool> {
189 match self {
190 ExtensionType::LastResort
191 | ExtensionType::RequiredCapabilities
192 | ExtensionType::ExternalSenders
193 | ExtensionType::ApplicationId => Some(false),
194 ExtensionType::RatchetTree | ExtensionType::ExternalPub => Some(true),
195 ExtensionType::Grease(_) | ExtensionType::Unknown(_) => None,
199 #[cfg(feature = "extensions-draft")]
200 ExtensionType::AppDataDictionary => Some(true),
201 }
202 }
203
204 pub(crate) fn is_valid_in_key_package(self) -> bool {
205 match self {
206 ExtensionType::RatchetTree
207 | ExtensionType::RequiredCapabilities
208 | ExtensionType::ExternalPub
209 | ExtensionType::ExternalSenders
210 | ExtensionType::ApplicationId => false,
211 ExtensionType::Grease(_) | ExtensionType::Unknown(_) | ExtensionType::LastResort => {
216 true
217 }
218 #[cfg(feature = "extensions-draft")]
219 ExtensionType::AppDataDictionary => true,
220 }
221 }
222
223 pub(crate) fn is_valid_in_group_context(self) -> bool {
224 match self {
225 ExtensionType::RequiredCapabilities
226 | ExtensionType::ExternalSenders
227 | ExtensionType::Unknown(_) => true,
228 ExtensionType::Grease(_) => true,
234 #[cfg(feature = "extensions-draft")]
235 ExtensionType::AppDataDictionary => true,
236 _ => false,
237 }
238 }
239
240 pub fn is_grease(&self) -> bool {
245 matches!(self, ExtensionType::Grease(_))
246 }
247}
248
249impl Size for ExtensionType {
250 fn tls_serialized_len(&self) -> usize {
251 2
252 }
253}
254
255impl TlsDeserializeTrait for ExtensionType {
256 fn tls_deserialize<R: Read>(bytes: &mut R) -> Result<Self, Error>
257 where
258 Self: Sized,
259 {
260 let mut extension_type = [0u8; 2];
261 bytes.read_exact(&mut extension_type)?;
262
263 Ok(ExtensionType::from(u16::from_be_bytes(extension_type)))
264 }
265}
266
267impl DeserializeBytes for ExtensionType {
268 fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), Error>
269 where
270 Self: Sized,
271 {
272 let mut bytes_ref = bytes;
273 let extension_type = ExtensionType::tls_deserialize(&mut bytes_ref)?;
274 Ok((extension_type, bytes_ref))
275 }
276}
277
278impl TlsSerializeTrait for ExtensionType {
279 fn tls_serialize<W: Write>(&self, writer: &mut W) -> Result<usize, Error> {
280 writer.write_all(&u16::from(*self).to_be_bytes())?;
281
282 Ok(2)
283 }
284}
285
286impl From<u16> for ExtensionType {
287 fn from(a: u16) -> Self {
288 match a {
289 1 => ExtensionType::ApplicationId,
290 2 => ExtensionType::RatchetTree,
291 3 => ExtensionType::RequiredCapabilities,
292 4 => ExtensionType::ExternalPub,
293 5 => ExtensionType::ExternalSenders,
294 #[cfg(feature = "extensions-draft")]
295 6 => ExtensionType::AppDataDictionary,
296 10 => ExtensionType::LastResort,
297 unknown if crate::grease::is_grease_value(unknown) => ExtensionType::Grease(unknown),
298 unknown => ExtensionType::Unknown(unknown),
299 }
300 }
301}
302
303impl From<ExtensionType> for u16 {
304 fn from(value: ExtensionType) -> Self {
305 match value {
306 ExtensionType::ApplicationId => 1,
307 ExtensionType::RatchetTree => 2,
308 ExtensionType::RequiredCapabilities => 3,
309 ExtensionType::ExternalPub => 4,
310 ExtensionType::ExternalSenders => 5,
311 #[cfg(feature = "extensions-draft")]
312 ExtensionType::AppDataDictionary => 6,
313 ExtensionType::LastResort => 10,
314 ExtensionType::Grease(value) => value,
315 ExtensionType::Unknown(unknown) => unknown,
316 }
317 }
318}
319
320#[derive(Debug, Clone, PartialEq, Eq)]
335#[cfg_attr(
336 feature = "0-8-1-storage-format",
337 derive(serde::Serialize, serde::Deserialize)
338)]
339#[cfg_attr(
340 not(feature = "0-8-1-storage-format"),
341 derive(
342 openmls_serialization_helpers::Serialize,
343 openmls_serialization_helpers::Deserialize,
344 )
345)]
346pub enum Extension {
347 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 0)]
348 ApplicationId(ApplicationIdExtension),
350
351 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 1)]
352 RatchetTree(RatchetTreeExtension),
354
355 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 2)]
356 RequiredCapabilities(RequiredCapabilitiesExtension),
358
359 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 3)]
360 ExternalPub(ExternalPubExtension),
362
363 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 4)]
364 ExternalSenders(ExternalSendersExtension),
366
367 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 7)]
368 #[cfg(feature = "extensions-draft")]
370 AppDataDictionary(AppDataDictionaryExtension),
371
372 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 5)]
373 LastResort(LastResortExtension),
375
376 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 6)]
377 Unknown(u16, UnknownExtension),
379}
380
381#[derive(
383 PartialEq, Eq, Clone, Debug, Serialize, Deserialize, TlsSize, TlsSerialize, TlsDeserialize,
384)]
385pub struct UnknownExtension(pub Vec<u8>);
386
387#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
389pub struct Extensions<T> {
390 unique: Vec<Extension>,
391 #[serde(skip)]
392 _object: core::marker::PhantomData<T>,
393}
394
395#[derive(Clone, Copy, PartialEq, Eq, Debug, Default, TlsSize, TlsSerialize, TlsDeserialize)]
396pub struct AnyObject;
398
399impl<T> Default for Extensions<T> {
400 fn default() -> Self {
401 Self {
402 unique: vec![],
403 _object: PhantomData,
404 }
405 }
406}
407
408impl<T> Size for Extensions<T> {
409 fn tls_serialized_len(&self) -> usize {
410 Vec::tls_serialized_len(&self.unique)
411 }
412}
413
414impl<T> TlsSerializeTrait for Extensions<T> {
415 fn tls_serialize<W: Write>(&self, writer: &mut W) -> Result<usize, Error> {
416 self.unique.tls_serialize(writer)
417 }
418}
419
420impl<T: ExtensionValidator<Error: ToString>> TlsDeserializeTrait for Extensions<T>
421where
422 InvalidExtensionError: From<T::Error>,
423{
424 fn tls_deserialize<R: Read>(bytes: &mut R) -> Result<Self, Error>
425 where
426 Self: Sized,
427 {
428 let candidate: Vec<ExtensionIn<T>> = Vec::tls_deserialize(bytes)?;
429 Extensions::<T>::try_from(candidate)
430 .map_err(|_| Error::DecodingError("Found duplicate extensions".into()))
431 }
432}
433
434impl<T: ExtensionValidator<Error: ToString>> DeserializeBytes for Extensions<T>
435where
436 InvalidExtensionError: From<T::Error>,
437{
438 fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), Error>
439 where
440 Self: Sized,
441 {
442 let mut bytes_ref = bytes;
443 let extensions = Extensions::<T>::tls_deserialize(&mut bytes_ref)?;
444 Ok((extensions, bytes_ref))
445 }
446}
447
448impl<T: ExtensionValidator> Extensions<T> {
449 pub fn empty() -> Self {
451 Self {
452 unique: vec![],
453 _object: PhantomData,
454 }
455 }
456
457 pub fn iter(&self) -> impl Iterator<Item = &Extension> {
459 self.unique.iter()
460 }
461
462 pub fn remove(&mut self, extension_type: ExtensionType) -> Option<Extension> {
467 if let Some(pos) = self
468 .unique
469 .iter()
470 .position(|ext| ext.extension_type() == extension_type)
471 {
472 Some(self.unique.remove(pos))
473 } else {
474 None
475 }
476 }
477
478 pub fn contains(&self, extension_type: ExtensionType) -> bool {
481 self.unique
482 .iter()
483 .any(|ext| ext.extension_type() == extension_type)
484 }
485}
486
487impl<T> Extensions<T>
488where
489 T: ExtensionValidator,
490 InvalidExtensionError: From<T::Error>,
491{
492 pub fn single(extension: Extension) -> Result<Self, InvalidExtensionError> {
494 T::validate_extension_type(extension.extension_type())?;
495 Ok(Self {
496 unique: vec![extension],
497 _object: PhantomData,
498 })
499 }
500
501 pub fn from_vec(extensions: Vec<Extension>) -> Result<Self, InvalidExtensionError> {
506 extensions.try_into()
507 }
508
509 pub fn validate<'a>(
511 extensions: impl Iterator<Item = &'a Extension>,
512 ) -> Result<(), InvalidExtensionError> {
513 for ext in extensions {
514 T::validate_extension_type(ext.extension_type())?;
515 }
516 Ok(())
517 }
518
519 pub fn add(&mut self, extension: Extension) -> Result<(), InvalidExtensionError> {
524 T::validate_extension_type(extension.extension_type())?;
525 if self.contains(extension.extension_type()) {
526 return Err(InvalidExtensionError::Duplicate);
527 }
528
529 self.unique.push(extension);
530
531 Ok(())
532 }
533
534 pub fn add_or_replace(
538 &mut self,
539 extension: Extension,
540 ) -> Result<Option<Extension>, InvalidExtensionError> {
541 T::validate_extension_type(extension.extension_type())?;
542 let replaced = self.remove(extension.extension_type());
543 self.unique.push(extension);
544 Ok(replaced)
545 }
546}
547
548impl Extensions<AnyObject> {
549 #[cfg(feature = "unchecked-conversions")]
555 pub fn into_unchecked<T>(self) -> Extensions<T> {
556 Extensions {
557 unique: self.unique,
558 _object: PhantomData,
559 }
560 }
561}
562
563mod private {
564 pub trait Sealed {}
566}
567
568pub trait ExtensionValidator: private::Sealed {
570 type Error;
572
573 fn validate_extension_type(ext: ExtensionType) -> Result<(), Self::Error>;
575}
576
577impl private::Sealed for AnyObject {}
578
579impl ExtensionValidator for AnyObject {
580 type Error = Infallible;
581
582 fn validate_extension_type(_ext: ExtensionType) -> Result<(), Infallible> {
583 Ok(())
584 }
585}
586
587impl<T: ExtensionValidator> TryFrom<Vec<Extension>> for Extensions<T>
588where
589 InvalidExtensionError: From<T::Error>,
590{
591 type Error = InvalidExtensionError;
592
593 fn try_from(candidate: Vec<Extension>) -> Result<Self, Self::Error> {
594 let mut seen = HashSet::with_capacity(candidate.len());
595 for extension in candidate.iter() {
596 T::validate_extension_type(extension.extension_type())?;
597
598 if !seen.insert(extension.extension_type()) {
599 return Err(InvalidExtensionError::Duplicate);
600 }
601 }
602
603 Ok(Self {
604 unique: candidate,
605 _object: PhantomData,
606 })
607 }
608}
609
610impl private::Sealed for GroupInfo {}
611
612impl ExtensionValidator for GroupInfo {
614 type Error = ExtensionTypeNotValidInGroupInfoError;
615
616 fn validate_extension_type(
617 extension_type: ExtensionType,
618 ) -> Result<(), ExtensionTypeNotValidInGroupInfoError> {
619 if extension_type.is_valid_in_group_info() == Some(true)
620 || extension_type.is_valid_in_group_info().is_none()
621 {
622 Ok(())
623 } else {
624 Err(ExtensionTypeNotValidInGroupInfoError(extension_type))
625 }
626 }
627}
628
629impl private::Sealed for GroupContext {}
630
631impl ExtensionValidator for GroupContext {
633 type Error = ExtensionTypeNotValidInGroupContextError;
634
635 fn validate_extension_type(
636 extension_type: ExtensionType,
637 ) -> Result<(), ExtensionTypeNotValidInGroupContextError> {
638 if extension_type.is_valid_in_group_context() {
639 Ok(())
640 } else {
641 Err(ExtensionTypeNotValidInGroupContextError(extension_type))
642 }
643 }
644}
645
646impl private::Sealed for KeyPackage {}
647
648impl ExtensionValidator for KeyPackage {
650 type Error = ExtensionTypeNotValidInKeyPackageError;
651
652 fn validate_extension_type(
653 extension_type: ExtensionType,
654 ) -> Result<(), ExtensionTypeNotValidInKeyPackageError> {
655 if extension_type.is_valid_in_key_package() {
656 Ok(())
657 } else {
658 Err(ExtensionTypeNotValidInKeyPackageError(extension_type))
659 }
660 }
661}
662
663impl private::Sealed for LeafNode {}
664
665impl ExtensionValidator for LeafNode {
667 type Error = ExtensionTypeNotValidInLeafNodeError;
668
669 fn validate_extension_type(
670 extension_type: ExtensionType,
671 ) -> Result<(), ExtensionTypeNotValidInLeafNodeError> {
672 if extension_type.is_valid_in_leaf_node() {
673 Ok(())
674 } else {
675 Err(ExtensionTypeNotValidInLeafNodeError(extension_type))
676 }
677 }
678}
679
680impl<T> Extensions<T> {
681 fn find_by_type(&self, extension_type: ExtensionType) -> Option<&Extension> {
682 self.unique
683 .iter()
684 .find(|ext| ext.extension_type() == extension_type)
685 }
686
687 pub fn application_id(&self) -> Option<&ApplicationIdExtension> {
689 self.find_by_type(ExtensionType::ApplicationId)
690 .and_then(|e| match e {
691 Extension::ApplicationId(e) => Some(e),
692 _ => None,
693 })
694 }
695
696 pub fn ratchet_tree(&self) -> Option<&RatchetTreeExtension> {
698 self.find_by_type(ExtensionType::RatchetTree)
699 .and_then(|e| match e {
700 Extension::RatchetTree(e) => Some(e),
701 _ => None,
702 })
703 }
704
705 pub fn required_capabilities(&self) -> Option<&RequiredCapabilitiesExtension> {
708 self.find_by_type(ExtensionType::RequiredCapabilities)
709 .and_then(|e| match e {
710 Extension::RequiredCapabilities(e) => Some(e),
711 _ => None,
712 })
713 }
714
715 pub fn external_pub(&self) -> Option<&ExternalPubExtension> {
717 self.find_by_type(ExtensionType::ExternalPub)
718 .and_then(|e| match e {
719 Extension::ExternalPub(e) => Some(e),
720 _ => None,
721 })
722 }
723
724 pub fn external_senders(&self) -> Option<&ExternalSendersExtension> {
726 self.find_by_type(ExtensionType::ExternalSenders)
727 .and_then(|e| match e {
728 Extension::ExternalSenders(e) => Some(e),
729 _ => None,
730 })
731 }
732
733 #[cfg(feature = "extensions-draft")]
734 pub fn app_data_dictionary(&self) -> Option<&AppDataDictionaryExtension> {
736 self.find_by_type(ExtensionType::AppDataDictionary)
737 .and_then(|e| match e {
738 Extension::AppDataDictionary(e) => Some(e),
739 _ => None,
740 })
741 }
742
743 pub fn unknown(&self, extension_type_id: u16) -> Option<&UnknownExtension> {
745 let extension_type: ExtensionType = extension_type_id.into();
746
747 match extension_type {
748 ExtensionType::Grease(_) | ExtensionType::Unknown(_) => {
749 self.find_by_type(extension_type).and_then(|e| match e {
750 Extension::Unknown(_, e) => Some(e),
751 _ => None,
752 })
753 }
754 _ => None,
755 }
756 }
757}
758
759impl Extension {
760 pub fn as_application_id_extension(&self) -> Result<&ApplicationIdExtension, ExtensionError> {
764 match self {
765 Self::ApplicationId(e) => Ok(e),
766 _ => Err(ExtensionError::InvalidExtensionType(
767 "This is not an ApplicationIdExtension".into(),
768 )),
769 }
770 }
771 #[cfg(feature = "extensions-draft")]
772 pub fn as_app_data_dictionary_extension(
776 &self,
777 ) -> Result<&AppDataDictionaryExtension, ExtensionError> {
778 match self {
779 Self::AppDataDictionary(e) => Ok(e),
780 _ => Err(ExtensionError::InvalidExtensionType(
781 "This is not an AppDataDictionaryExtension".into(),
782 )),
783 }
784 }
785
786 pub fn as_ratchet_tree_extension(&self) -> Result<&RatchetTreeExtension, ExtensionError> {
790 match self {
791 Self::RatchetTree(rte) => Ok(rte),
792 _ => Err(ExtensionError::InvalidExtensionType(
793 "This is not a RatchetTreeExtension".into(),
794 )),
795 }
796 }
797
798 pub fn as_required_capabilities_extension(
802 &self,
803 ) -> Result<&RequiredCapabilitiesExtension, ExtensionError> {
804 match self {
805 Self::RequiredCapabilities(e) => Ok(e),
806 _ => Err(ExtensionError::InvalidExtensionType(
807 "This is not a RequiredCapabilitiesExtension".into(),
808 )),
809 }
810 }
811
812 pub fn as_external_pub_extension(&self) -> Result<&ExternalPubExtension, ExtensionError> {
816 match self {
817 Self::ExternalPub(e) => Ok(e),
818 _ => Err(ExtensionError::InvalidExtensionType(
819 "This is not an ExternalPubExtension".into(),
820 )),
821 }
822 }
823
824 pub fn as_external_senders_extension(
828 &self,
829 ) -> Result<&ExternalSendersExtension, ExtensionError> {
830 match self {
831 Self::ExternalSenders(e) => Ok(e),
832 _ => Err(ExtensionError::InvalidExtensionType(
833 "This is not an ExternalSendersExtension".into(),
834 )),
835 }
836 }
837
838 #[inline]
840 pub const fn extension_type(&self) -> ExtensionType {
841 match self {
842 Extension::ApplicationId(_) => ExtensionType::ApplicationId,
843 Extension::RatchetTree(_) => ExtensionType::RatchetTree,
844 Extension::RequiredCapabilities(_) => ExtensionType::RequiredCapabilities,
845 Extension::ExternalPub(_) => ExtensionType::ExternalPub,
846 Extension::ExternalSenders(_) => ExtensionType::ExternalSenders,
847 #[cfg(feature = "extensions-draft")]
848 Extension::AppDataDictionary(_) => ExtensionType::AppDataDictionary,
849 Extension::LastResort(_) => ExtensionType::LastResort,
850 Extension::Unknown(kind, _) if crate::grease::is_grease_value(*kind) => {
855 ExtensionType::Grease(*kind)
856 }
857 Extension::Unknown(kind, _) => ExtensionType::Unknown(*kind),
858 }
859 }
860}
861
862macro_rules! impl_from_extensions_validator {
863 ($validator:ty, $error:ty) => {
864 impl From<Extensions<$validator>> for Extensions<AnyObject> {
865 fn from(value: Extensions<$validator>) -> Self {
866 Extensions {
867 unique: value.unique,
868 _object: PhantomData,
869 }
870 }
871 }
872
873 impl TryFrom<Extensions<AnyObject>> for Extensions<$validator> {
874 type Error = $error;
875
876 fn try_from(value: Extensions<AnyObject>) -> Result<Self, $error> {
877 value
878 .unique
879 .iter()
880 .map(Extension::extension_type)
881 .try_for_each(<$validator as ExtensionValidator>::validate_extension_type)?;
882
883 Ok(Extensions {
884 unique: value.unique,
885 _object: PhantomData,
886 })
887 }
888 }
889 };
890}
891
892impl_from_extensions_validator!(GroupContext, ExtensionTypeNotValidInGroupContextError);
893impl_from_extensions_validator!(LeafNode, ExtensionTypeNotValidInLeafNodeError);
894impl_from_extensions_validator!(KeyPackage, ExtensionTypeNotValidInKeyPackageError);
895
896#[cfg(any(feature = "test-utils", test))]
897impl Extensions<AnyObject> {
898 pub(crate) fn coerce<T: ExtensionValidator>(self) -> Extensions<T> {
900 Extensions {
901 unique: self.unique,
902 _object: PhantomData,
903 }
904 }
905}
906#[cfg(test)]
907mod test {
908 use itertools::Itertools;
909 use tls_codec::{Deserialize, Serialize, VLBytes};
910
911 use crate::{ciphersuite::HpkePublicKey, extensions::*};
912
913 #[test]
914 fn add() {
915 let mut extensions: Extensions<AnyObject> = Extensions::default();
916 extensions
917 .add(Extension::RequiredCapabilities(
918 RequiredCapabilitiesExtension::default(),
919 ))
920 .unwrap();
921 assert!(extensions
922 .add(Extension::RequiredCapabilities(
923 RequiredCapabilitiesExtension::default()
924 ))
925 .is_err());
926 }
927
928 #[test]
929 fn grease_extension_type_mapping() {
930 let grease = Extension::Unknown(0x5A5A, UnknownExtension(vec![1, 2, 3]));
935 assert_eq!(grease.extension_type(), ExtensionType::Grease(0x5A5A));
936 assert!(grease.extension_type().is_grease());
937
938 let unknown = Extension::Unknown(0xABCD, UnknownExtension(vec![]));
940 assert_eq!(unknown.extension_type(), ExtensionType::Unknown(0xABCD));
941 }
942
943 #[test]
944 fn grease_extension_must_be_declared_in_capabilities() {
945 let advertised = crate::treesync::node::leaf_node::Capabilities::new(
952 None,
953 None,
954 Some(&[ExtensionType::Grease(0x5A5A)]),
955 None,
956 None,
957 );
958 assert!(advertised.contains_extension_type(&ExtensionType::Grease(0x5A5A)));
959 assert!(!advertised.contains_extension_type(&ExtensionType::Grease(0xAAAA)));
961
962 let empty =
964 crate::treesync::node::leaf_node::Capabilities::new(None, None, None, None, None);
965 assert!(!empty.contains_extension_type(&ExtensionType::Grease(0x5A5A)));
966 assert!(!empty.contains_extension_type(&ExtensionType::Unknown(0xABCD)));
967 }
968
969 #[test]
970 fn add_try_from() {
971 let ext_x = Extension::ApplicationId(ApplicationIdExtension::new(b"Test"));
974 let ext_y = Extension::RequiredCapabilities(RequiredCapabilitiesExtension::default());
975
976 let tests = [
977 (vec![], true),
978 (vec![ext_x.clone()], true),
979 (vec![ext_x.clone(), ext_x.clone()], false),
980 (vec![ext_x.clone(), ext_x.clone(), ext_x.clone()], false),
981 (vec![ext_y.clone()], true),
982 (vec![ext_y.clone(), ext_y.clone()], false),
983 (vec![ext_y.clone(), ext_y.clone(), ext_y.clone()], false),
984 (vec![ext_x.clone(), ext_y.clone()], true),
985 (vec![ext_y.clone(), ext_x.clone()], true),
986 (vec![ext_x.clone(), ext_x.clone(), ext_y.clone()], false),
987 (vec![ext_y.clone(), ext_y.clone(), ext_x.clone()], false),
988 (vec![ext_x.clone(), ext_y.clone(), ext_y.clone()], false),
989 (vec![ext_y.clone(), ext_x.clone(), ext_x.clone()], false),
990 (vec![ext_x.clone(), ext_y.clone(), ext_x.clone()], false),
991 (vec![ext_y.clone(), ext_x, ext_y], false),
992 ];
993
994 for (test, should_work) in tests.into_iter() {
995 {
997 let mut extensions: Extensions<AnyObject> = Extensions::default();
998
999 let mut works = true;
1000 for ext in test.iter() {
1001 match extensions.add(ext.clone()) {
1002 Ok(_) => {}
1003 Err(InvalidExtensionError::Duplicate) => {
1004 works = false;
1005 }
1006 _ => panic!("This should have never happened."),
1007 }
1008 }
1009
1010 println!("{:?}, {:?}", test.clone(), should_work);
1011 assert_eq!(works, should_work);
1012 }
1013
1014 if should_work {
1016 assert!(Extensions::<AnyObject>::try_from(test).is_ok());
1017 } else {
1018 assert!(Extensions::<AnyObject>::try_from(test).is_err());
1019 }
1020 }
1021 }
1022
1023 #[test]
1024 fn ensure_ordering() {
1025 let ext_x = Extension::ApplicationId(ApplicationIdExtension::new(b"Test"));
1029 let ext_y = Extension::ExternalPub(ExternalPubExtension::new(HpkePublicKey::new(vec![])));
1030 let ext_z = Extension::RequiredCapabilities(RequiredCapabilitiesExtension::default());
1031
1032 for candidate in [ext_x, ext_y, ext_z]
1033 .into_iter()
1034 .permutations(3)
1035 .collect::<Vec<_>>()
1036 {
1037 let candidate: Extensions<AnyObject> = Extensions::try_from(candidate).unwrap();
1038 let bytes = candidate.tls_serialize_detached().unwrap();
1039 let got = Extensions::tls_deserialize(&mut bytes.as_slice()).unwrap();
1040 assert_eq!(candidate, got);
1041 }
1042 }
1043
1044 #[test]
1045 fn that_unknown_extensions_are_de_serialized_correctly() {
1046 let extension_types = [0x0000u16, 0x0A0A, 0x7A7A, 0xF100, 0xFFFF];
1047 let extension_datas = [vec![], vec![0], vec![1, 2, 3]];
1048
1049 for extension_type in extension_types.into_iter() {
1050 for extension_data in extension_datas.iter() {
1051 let test = {
1053 let mut buf = extension_type.to_be_bytes().to_vec();
1054 buf.append(
1055 &mut VLBytes::new(extension_data.clone())
1056 .tls_serialize_detached()
1057 .unwrap(),
1058 );
1059 buf
1060 };
1061
1062 let got = Extension::tls_deserialize_exact(&test).unwrap();
1064
1065 match got {
1066 Extension::Unknown(got_extension_type, ref got_extension_data) => {
1067 assert_eq!(extension_type, got_extension_type);
1068 assert_eq!(extension_data, &got_extension_data.0);
1069 }
1070 other => panic!("Expected `Extension::Unknown`, got {other:?}"),
1071 }
1072
1073 let got_serialized = got.tls_serialize_detached().unwrap();
1075 assert_eq!(test, got_serialized);
1076 }
1077 }
1078 }
1079}