1use std::{
25 collections::HashSet,
26 convert::Infallible,
27 fmt::Debug,
28 io::{Read, Write},
29 marker::PhantomData,
30};
31
32use serde::{Deserialize, Serialize};
33
34#[cfg(feature = "extensions-draft")]
36mod app_data_dict_extension;
37mod application_id_extension;
38mod codec;
39mod external_pub_extension;
40mod external_sender_extension;
41mod last_resort;
42mod ratchet_tree_extension;
43mod required_capabilities;
44use errors::*;
45
46pub mod errors;
48
49#[cfg(feature = "extensions-draft")]
51pub use app_data_dict_extension::{AppDataDictionary, AppDataDictionaryExtension};
52pub use application_id_extension::ApplicationIdExtension;
53pub use external_pub_extension::ExternalPubExtension;
54pub use external_sender_extension::{
55 ExternalSender, ExternalSendersExtension, SenderExtensionIndex,
56};
57pub use last_resort::LastResortExtension;
58pub use ratchet_tree_extension::RatchetTreeExtension;
59pub use required_capabilities::RequiredCapabilitiesExtension;
60
61use tls_codec::{
62 Deserialize as TlsDeserializeTrait, DeserializeBytes, Error, Serialize as TlsSerializeTrait,
63 Size, TlsDeserialize, TlsSerialize, TlsSize,
64};
65
66use crate::{
67 group::GroupContext, key_packages::KeyPackage, messages::group_info::GroupInfo,
68 treesync::LeafNode,
69};
70
71#[cfg(test)]
72mod tests;
73
74#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash, Ord, PartialOrd)]
90#[cfg_attr(
91 feature = "0-8-1-storage-format",
92 derive(serde::Serialize, serde::Deserialize)
93)]
94#[cfg_attr(
95 not(feature = "0-8-1-storage-format"),
96 derive(
97 openmls_serialization_helpers::Serialize,
98 openmls_serialization_helpers::Deserialize,
99 )
100)]
101pub enum ExtensionType {
102 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 0)]
103 ApplicationId,
106
107 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 1)]
108 RatchetTree,
111
112 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 2)]
113 RequiredCapabilities,
116
117 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 3)]
118 ExternalPub,
121
122 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 4)]
123 ExternalSenders,
126
127 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 5)]
128 LastResort,
131
132 #[cfg(feature = "extensions-draft")]
133 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 8)]
134 AppDataDictionary,
136
137 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 7)]
138 Grease(u16),
140
141 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 6)]
142 Unknown(u16),
144}
145
146impl ExtensionType {
147 pub(crate) fn is_default(self) -> bool {
149 match self {
150 ExtensionType::ApplicationId
151 | ExtensionType::RatchetTree
152 | ExtensionType::RequiredCapabilities
153 | ExtensionType::ExternalPub
154 | ExtensionType::ExternalSenders => true,
155 ExtensionType::LastResort | ExtensionType::Grease(_) | ExtensionType::Unknown(_) => {
156 false
157 }
158 #[cfg(feature = "extensions-draft")]
159 ExtensionType::AppDataDictionary => false,
160 }
161 }
162
163 pub(crate) fn is_valid_in_leaf_node(self) -> bool {
168 match self {
169 ExtensionType::LastResort
170 | ExtensionType::RatchetTree
171 | ExtensionType::RequiredCapabilities
172 | ExtensionType::ExternalPub
173 | ExtensionType::ExternalSenders => false,
174 ExtensionType::Grease(_) | ExtensionType::Unknown(_) | ExtensionType::ApplicationId => {
179 true
180 }
181 #[cfg(feature = "extensions-draft")]
182 ExtensionType::AppDataDictionary => true,
183 }
184 }
185 pub(crate) fn is_valid_in_group_info(self) -> Option<bool> {
186 match self {
187 ExtensionType::LastResort
188 | ExtensionType::RequiredCapabilities
189 | ExtensionType::ExternalSenders
190 | ExtensionType::ApplicationId => Some(false),
191 ExtensionType::RatchetTree | ExtensionType::ExternalPub => Some(true),
192 ExtensionType::Grease(_) | ExtensionType::Unknown(_) => None,
196 #[cfg(feature = "extensions-draft")]
197 ExtensionType::AppDataDictionary => Some(true),
198 }
199 }
200
201 pub(crate) fn is_valid_in_key_package(self) -> bool {
202 match self {
203 ExtensionType::RatchetTree
204 | ExtensionType::RequiredCapabilities
205 | ExtensionType::ExternalPub
206 | ExtensionType::ExternalSenders
207 | ExtensionType::ApplicationId => false,
208 ExtensionType::Grease(_) | ExtensionType::Unknown(_) | ExtensionType::LastResort => {
213 true
214 }
215 #[cfg(feature = "extensions-draft")]
216 ExtensionType::AppDataDictionary => true,
217 }
218 }
219
220 pub(crate) fn is_valid_in_group_context(self) -> bool {
221 match self {
222 ExtensionType::RequiredCapabilities
223 | ExtensionType::ExternalSenders
224 | ExtensionType::Unknown(_) => true,
225 ExtensionType::Grease(_) => true,
231 #[cfg(feature = "extensions-draft")]
232 ExtensionType::AppDataDictionary => true,
233 _ => false,
234 }
235 }
236
237 pub fn is_grease(&self) -> bool {
242 matches!(self, ExtensionType::Grease(_))
243 }
244}
245
246impl Size for ExtensionType {
247 fn tls_serialized_len(&self) -> usize {
248 2
249 }
250}
251
252impl TlsDeserializeTrait for ExtensionType {
253 fn tls_deserialize<R: Read>(bytes: &mut R) -> Result<Self, Error>
254 where
255 Self: Sized,
256 {
257 let mut extension_type = [0u8; 2];
258 bytes.read_exact(&mut extension_type)?;
259
260 Ok(ExtensionType::from(u16::from_be_bytes(extension_type)))
261 }
262}
263
264impl DeserializeBytes for ExtensionType {
265 fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), Error>
266 where
267 Self: Sized,
268 {
269 let mut bytes_ref = bytes;
270 let extension_type = ExtensionType::tls_deserialize(&mut bytes_ref)?;
271 Ok((extension_type, bytes_ref))
272 }
273}
274
275impl TlsSerializeTrait for ExtensionType {
276 fn tls_serialize<W: Write>(&self, writer: &mut W) -> Result<usize, Error> {
277 writer.write_all(&u16::from(*self).to_be_bytes())?;
278
279 Ok(2)
280 }
281}
282
283impl From<u16> for ExtensionType {
284 fn from(a: u16) -> Self {
285 match a {
286 1 => ExtensionType::ApplicationId,
287 2 => ExtensionType::RatchetTree,
288 3 => ExtensionType::RequiredCapabilities,
289 4 => ExtensionType::ExternalPub,
290 5 => ExtensionType::ExternalSenders,
291 #[cfg(feature = "extensions-draft")]
292 6 => ExtensionType::AppDataDictionary,
293 10 => ExtensionType::LastResort,
294 unknown if crate::grease::is_grease_value(unknown) => ExtensionType::Grease(unknown),
295 unknown => ExtensionType::Unknown(unknown),
296 }
297 }
298}
299
300impl From<ExtensionType> for u16 {
301 fn from(value: ExtensionType) -> Self {
302 match value {
303 ExtensionType::ApplicationId => 1,
304 ExtensionType::RatchetTree => 2,
305 ExtensionType::RequiredCapabilities => 3,
306 ExtensionType::ExternalPub => 4,
307 ExtensionType::ExternalSenders => 5,
308 #[cfg(feature = "extensions-draft")]
309 ExtensionType::AppDataDictionary => 6,
310 ExtensionType::LastResort => 10,
311 ExtensionType::Grease(value) => value,
312 ExtensionType::Unknown(unknown) => unknown,
313 }
314 }
315}
316
317#[derive(Debug, Clone, PartialEq, Eq)]
332#[cfg_attr(
333 feature = "0-8-1-storage-format",
334 derive(serde::Serialize, serde::Deserialize)
335)]
336#[cfg_attr(
337 not(feature = "0-8-1-storage-format"),
338 derive(
339 openmls_serialization_helpers::Serialize,
340 openmls_serialization_helpers::Deserialize,
341 )
342)]
343pub enum Extension {
344 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 0)]
345 ApplicationId(ApplicationIdExtension),
347
348 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 1)]
349 RatchetTree(RatchetTreeExtension),
351
352 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 2)]
353 RequiredCapabilities(RequiredCapabilitiesExtension),
355
356 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 3)]
357 ExternalPub(ExternalPubExtension),
359
360 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 4)]
361 ExternalSenders(ExternalSendersExtension),
363
364 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 7)]
365 #[cfg(feature = "extensions-draft")]
367 AppDataDictionary(AppDataDictionaryExtension),
368
369 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 5)]
370 LastResort(LastResortExtension),
372
373 #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 6)]
374 Unknown(u16, UnknownExtension),
376}
377
378#[derive(
380 PartialEq, Eq, Clone, Debug, Serialize, Deserialize, TlsSize, TlsSerialize, TlsDeserialize,
381)]
382pub struct UnknownExtension(pub Vec<u8>);
383
384#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
386pub struct Extensions<T> {
387 unique: Vec<Extension>,
388 #[serde(skip)]
389 _object: core::marker::PhantomData<T>,
390}
391
392#[derive(Clone, Copy, PartialEq, Eq, Debug, Default, TlsSize, TlsSerialize, TlsDeserialize)]
393pub struct AnyObject;
395
396impl<T> Default for Extensions<T> {
397 fn default() -> Self {
398 Self {
399 unique: vec![],
400 _object: PhantomData,
401 }
402 }
403}
404
405impl<T> Size for Extensions<T> {
406 fn tls_serialized_len(&self) -> usize {
407 Vec::tls_serialized_len(&self.unique)
408 }
409}
410
411impl<T> TlsSerializeTrait for Extensions<T> {
412 fn tls_serialize<W: Write>(&self, writer: &mut W) -> Result<usize, Error> {
413 self.unique.tls_serialize(writer)
414 }
415}
416
417impl<T: ExtensionValidator> TlsDeserializeTrait for Extensions<T>
418where
419 InvalidExtensionError: From<T::Error>,
420{
421 fn tls_deserialize<R: Read>(bytes: &mut R) -> Result<Self, Error>
422 where
423 Self: Sized,
424 {
425 let candidate: Vec<Extension> = Vec::tls_deserialize(bytes)?;
426 Extensions::<T>::try_from(candidate)
427 .map_err(|_| Error::DecodingError("Found duplicate extensions".into()))
428 }
429}
430
431impl<T: ExtensionValidator> DeserializeBytes for Extensions<T>
432where
433 InvalidExtensionError: From<T::Error>,
434{
435 fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), Error>
436 where
437 Self: Sized,
438 {
439 let mut bytes_ref = bytes;
440 let extensions = Extensions::<T>::tls_deserialize(&mut bytes_ref)?;
441 Ok((extensions, bytes_ref))
442 }
443}
444
445impl<T: ExtensionValidator> Extensions<T> {
446 pub fn empty() -> Self {
448 Self {
449 unique: vec![],
450 _object: PhantomData,
451 }
452 }
453
454 pub fn iter(&self) -> impl Iterator<Item = &Extension> {
456 self.unique.iter()
457 }
458
459 pub fn remove(&mut self, extension_type: ExtensionType) -> Option<Extension> {
464 if let Some(pos) = self
465 .unique
466 .iter()
467 .position(|ext| ext.extension_type() == extension_type)
468 {
469 Some(self.unique.remove(pos))
470 } else {
471 None
472 }
473 }
474
475 pub fn contains(&self, extension_type: ExtensionType) -> bool {
478 self.unique
479 .iter()
480 .any(|ext| ext.extension_type() == extension_type)
481 }
482}
483
484impl<T> Extensions<T>
485where
486 T: ExtensionValidator,
487 InvalidExtensionError: From<T::Error>,
488{
489 pub fn single(extension: Extension) -> Result<Self, InvalidExtensionError> {
491 T::validate_extension_type(&extension)?;
492 Ok(Self {
493 unique: vec![extension],
494 _object: PhantomData,
495 })
496 }
497
498 pub fn from_vec(extensions: Vec<Extension>) -> Result<Self, InvalidExtensionError> {
503 extensions.try_into()
504 }
505
506 pub fn validate<'a>(
508 extensions: impl Iterator<Item = &'a Extension>,
509 ) -> Result<(), InvalidExtensionError> {
510 for ext in extensions {
511 T::validate_extension_type(ext)?;
512 }
513 Ok(())
514 }
515
516 pub fn add(&mut self, extension: Extension) -> Result<(), InvalidExtensionError> {
521 T::validate_extension_type(&extension)?;
522 if self.contains(extension.extension_type()) {
523 return Err(InvalidExtensionError::Duplicate);
524 }
525
526 self.unique.push(extension);
527
528 Ok(())
529 }
530
531 pub fn add_or_replace(
535 &mut self,
536 extension: Extension,
537 ) -> Result<Option<Extension>, InvalidExtensionError> {
538 T::validate_extension_type(&extension)?;
539 let replaced = self.remove(extension.extension_type());
540 self.unique.push(extension);
541 Ok(replaced)
542 }
543}
544
545impl Extensions<AnyObject> {
546 #[cfg(feature = "unchecked-conversions")]
552 pub fn into_unchecked<T>(self) -> Extensions<T> {
553 Extensions {
554 unique: self.unique,
555 _object: PhantomData,
556 }
557 }
558}
559
560pub trait ExtensionValidator {
562 type Error;
564
565 fn validate_extension_type(ext: &Extension) -> Result<(), Self::Error>;
567}
568
569impl ExtensionValidator for AnyObject {
570 type Error = Infallible;
571
572 fn validate_extension_type(_ext: &Extension) -> Result<(), Infallible> {
573 Ok(())
574 }
575}
576
577impl<T: ExtensionValidator> TryFrom<Vec<Extension>> for Extensions<T>
578where
579 InvalidExtensionError: From<T::Error>,
580{
581 type Error = InvalidExtensionError;
582
583 fn try_from(candidate: Vec<Extension>) -> Result<Self, Self::Error> {
584 let mut seen = HashSet::with_capacity(candidate.len());
585 for extension in candidate.iter() {
586 T::validate_extension_type(extension)?;
587
588 if !seen.insert(extension.extension_type()) {
589 return Err(InvalidExtensionError::Duplicate);
590 }
591 }
592
593 Ok(Self {
594 unique: candidate,
595 _object: PhantomData,
596 })
597 }
598}
599
600impl ExtensionValidator for GroupInfo {
602 type Error = ExtensionTypeNotValidInGroupInfoError;
603
604 fn validate_extension_type(
605 ext: &Extension,
606 ) -> Result<(), ExtensionTypeNotValidInGroupInfoError> {
607 if ext.extension_type().is_valid_in_group_info() == Some(true)
608 || ext.extension_type().is_valid_in_group_info().is_none()
609 {
610 Ok(())
611 } else {
612 Err(ExtensionTypeNotValidInGroupInfoError(ext.extension_type()))
613 }
614 }
615}
616
617impl ExtensionValidator for GroupContext {
619 type Error = ExtensionTypeNotValidInGroupContextError;
620
621 fn validate_extension_type(
622 ext: &Extension,
623 ) -> Result<(), ExtensionTypeNotValidInGroupContextError> {
624 if ext.extension_type().is_valid_in_group_context() {
625 Ok(())
626 } else {
627 Err(ExtensionTypeNotValidInGroupContextError(
628 ext.extension_type(),
629 ))
630 }
631 }
632}
633
634impl ExtensionValidator for KeyPackage {
636 type Error = ExtensionTypeNotValidInKeyPackageError;
637
638 fn validate_extension_type(
639 ext: &Extension,
640 ) -> Result<(), ExtensionTypeNotValidInKeyPackageError> {
641 if ext.extension_type().is_valid_in_key_package() {
642 Ok(())
643 } else {
644 Err(ExtensionTypeNotValidInKeyPackageError(ext.extension_type()))
645 }
646 }
647}
648
649impl ExtensionValidator for LeafNode {
651 type Error = ExtensionTypeNotValidInLeafNodeError;
652
653 fn validate_extension_type(
654 ext: &Extension,
655 ) -> Result<(), ExtensionTypeNotValidInLeafNodeError> {
656 if ext.extension_type().is_valid_in_leaf_node() {
657 Ok(())
658 } else {
659 Err(ExtensionTypeNotValidInLeafNodeError(ext.extension_type()))
660 }
661 }
662}
663
664impl<T> Extensions<T> {
665 fn find_by_type(&self, extension_type: ExtensionType) -> Option<&Extension> {
666 self.unique
667 .iter()
668 .find(|ext| ext.extension_type() == extension_type)
669 }
670
671 pub fn application_id(&self) -> Option<&ApplicationIdExtension> {
673 self.find_by_type(ExtensionType::ApplicationId)
674 .and_then(|e| match e {
675 Extension::ApplicationId(e) => Some(e),
676 _ => None,
677 })
678 }
679
680 pub fn ratchet_tree(&self) -> Option<&RatchetTreeExtension> {
682 self.find_by_type(ExtensionType::RatchetTree)
683 .and_then(|e| match e {
684 Extension::RatchetTree(e) => Some(e),
685 _ => None,
686 })
687 }
688
689 pub fn required_capabilities(&self) -> Option<&RequiredCapabilitiesExtension> {
692 self.find_by_type(ExtensionType::RequiredCapabilities)
693 .and_then(|e| match e {
694 Extension::RequiredCapabilities(e) => Some(e),
695 _ => None,
696 })
697 }
698
699 pub fn external_pub(&self) -> Option<&ExternalPubExtension> {
701 self.find_by_type(ExtensionType::ExternalPub)
702 .and_then(|e| match e {
703 Extension::ExternalPub(e) => Some(e),
704 _ => None,
705 })
706 }
707
708 pub fn external_senders(&self) -> Option<&ExternalSendersExtension> {
710 self.find_by_type(ExtensionType::ExternalSenders)
711 .and_then(|e| match e {
712 Extension::ExternalSenders(e) => Some(e),
713 _ => None,
714 })
715 }
716
717 #[cfg(feature = "extensions-draft")]
718 pub fn app_data_dictionary(&self) -> Option<&AppDataDictionaryExtension> {
720 self.find_by_type(ExtensionType::AppDataDictionary)
721 .and_then(|e| match e {
722 Extension::AppDataDictionary(e) => Some(e),
723 _ => None,
724 })
725 }
726
727 pub fn unknown(&self, extension_type_id: u16) -> Option<&UnknownExtension> {
729 let extension_type: ExtensionType = extension_type_id.into();
730
731 match extension_type {
732 ExtensionType::Grease(_) | ExtensionType::Unknown(_) => {
733 self.find_by_type(extension_type).and_then(|e| match e {
734 Extension::Unknown(_, e) => Some(e),
735 _ => None,
736 })
737 }
738 _ => None,
739 }
740 }
741}
742
743impl Extension {
744 pub fn as_application_id_extension(&self) -> Result<&ApplicationIdExtension, ExtensionError> {
748 match self {
749 Self::ApplicationId(e) => Ok(e),
750 _ => Err(ExtensionError::InvalidExtensionType(
751 "This is not an ApplicationIdExtension".into(),
752 )),
753 }
754 }
755 #[cfg(feature = "extensions-draft")]
756 pub fn as_app_data_dictionary_extension(
760 &self,
761 ) -> Result<&AppDataDictionaryExtension, ExtensionError> {
762 match self {
763 Self::AppDataDictionary(e) => Ok(e),
764 _ => Err(ExtensionError::InvalidExtensionType(
765 "This is not an AppDataDictionaryExtension".into(),
766 )),
767 }
768 }
769
770 pub fn as_ratchet_tree_extension(&self) -> Result<&RatchetTreeExtension, ExtensionError> {
774 match self {
775 Self::RatchetTree(rte) => Ok(rte),
776 _ => Err(ExtensionError::InvalidExtensionType(
777 "This is not a RatchetTreeExtension".into(),
778 )),
779 }
780 }
781
782 pub fn as_required_capabilities_extension(
786 &self,
787 ) -> Result<&RequiredCapabilitiesExtension, ExtensionError> {
788 match self {
789 Self::RequiredCapabilities(e) => Ok(e),
790 _ => Err(ExtensionError::InvalidExtensionType(
791 "This is not a RequiredCapabilitiesExtension".into(),
792 )),
793 }
794 }
795
796 pub fn as_external_pub_extension(&self) -> Result<&ExternalPubExtension, ExtensionError> {
800 match self {
801 Self::ExternalPub(e) => Ok(e),
802 _ => Err(ExtensionError::InvalidExtensionType(
803 "This is not an ExternalPubExtension".into(),
804 )),
805 }
806 }
807
808 pub fn as_external_senders_extension(
812 &self,
813 ) -> Result<&ExternalSendersExtension, ExtensionError> {
814 match self {
815 Self::ExternalSenders(e) => Ok(e),
816 _ => Err(ExtensionError::InvalidExtensionType(
817 "This is not an ExternalSendersExtension".into(),
818 )),
819 }
820 }
821
822 #[inline]
824 pub const fn extension_type(&self) -> ExtensionType {
825 match self {
826 Extension::ApplicationId(_) => ExtensionType::ApplicationId,
827 Extension::RatchetTree(_) => ExtensionType::RatchetTree,
828 Extension::RequiredCapabilities(_) => ExtensionType::RequiredCapabilities,
829 Extension::ExternalPub(_) => ExtensionType::ExternalPub,
830 Extension::ExternalSenders(_) => ExtensionType::ExternalSenders,
831 #[cfg(feature = "extensions-draft")]
832 Extension::AppDataDictionary(_) => ExtensionType::AppDataDictionary,
833 Extension::LastResort(_) => ExtensionType::LastResort,
834 Extension::Unknown(kind, _) if crate::grease::is_grease_value(*kind) => {
839 ExtensionType::Grease(*kind)
840 }
841 Extension::Unknown(kind, _) => ExtensionType::Unknown(*kind),
842 }
843 }
844}
845
846macro_rules! impl_from_extensions_validator {
847 ($validator:ty, $error:ty) => {
848 impl From<Extensions<$validator>> for Extensions<AnyObject> {
849 fn from(value: Extensions<$validator>) -> Self {
850 Extensions {
851 unique: value.unique,
852 _object: PhantomData,
853 }
854 }
855 }
856
857 impl TryFrom<Extensions<AnyObject>> for Extensions<$validator> {
858 type Error = $error;
859
860 fn try_from(value: Extensions<AnyObject>) -> Result<Self, $error> {
861 value
862 .unique
863 .iter()
864 .try_for_each(<$validator as ExtensionValidator>::validate_extension_type)?;
865
866 Ok(Extensions {
867 unique: value.unique,
868 _object: PhantomData,
869 })
870 }
871 }
872 };
873}
874
875impl_from_extensions_validator!(GroupContext, ExtensionTypeNotValidInGroupContextError);
876impl_from_extensions_validator!(LeafNode, ExtensionTypeNotValidInLeafNodeError);
877impl_from_extensions_validator!(KeyPackage, ExtensionTypeNotValidInKeyPackageError);
878
879#[cfg(any(feature = "test-utils", test))]
880impl Extensions<AnyObject> {
881 pub(crate) fn coerce<T: ExtensionValidator>(self) -> Extensions<T> {
883 Extensions {
884 unique: self.unique,
885 _object: PhantomData,
886 }
887 }
888}
889#[cfg(test)]
890mod test {
891 use itertools::Itertools;
892 use tls_codec::{Deserialize, Serialize, VLBytes};
893
894 use crate::{ciphersuite::HpkePublicKey, extensions::*};
895
896 #[test]
897 fn add() {
898 let mut extensions: Extensions<AnyObject> = Extensions::default();
899 extensions
900 .add(Extension::RequiredCapabilities(
901 RequiredCapabilitiesExtension::default(),
902 ))
903 .unwrap();
904 assert!(extensions
905 .add(Extension::RequiredCapabilities(
906 RequiredCapabilitiesExtension::default()
907 ))
908 .is_err());
909 }
910
911 #[test]
912 fn grease_extension_type_mapping() {
913 let grease = Extension::Unknown(0x5A5A, UnknownExtension(vec![1, 2, 3]));
918 assert_eq!(grease.extension_type(), ExtensionType::Grease(0x5A5A));
919 assert!(grease.extension_type().is_grease());
920
921 let unknown = Extension::Unknown(0xABCD, UnknownExtension(vec![]));
923 assert_eq!(unknown.extension_type(), ExtensionType::Unknown(0xABCD));
924 }
925
926 #[test]
927 fn grease_extension_must_be_declared_in_capabilities() {
928 let advertised = crate::treesync::node::leaf_node::Capabilities::new(
935 None,
936 None,
937 Some(&[ExtensionType::Grease(0x5A5A)]),
938 None,
939 None,
940 );
941 assert!(advertised.contains_extension_type(&ExtensionType::Grease(0x5A5A)));
942 assert!(!advertised.contains_extension_type(&ExtensionType::Grease(0xAAAA)));
944
945 let empty =
947 crate::treesync::node::leaf_node::Capabilities::new(None, None, None, None, None);
948 assert!(!empty.contains_extension_type(&ExtensionType::Grease(0x5A5A)));
949 assert!(!empty.contains_extension_type(&ExtensionType::Unknown(0xABCD)));
950 }
951
952 #[test]
953 fn add_try_from() {
954 let ext_x = Extension::ApplicationId(ApplicationIdExtension::new(b"Test"));
957 let ext_y = Extension::RequiredCapabilities(RequiredCapabilitiesExtension::default());
958
959 let tests = [
960 (vec![], true),
961 (vec![ext_x.clone()], true),
962 (vec![ext_x.clone(), ext_x.clone()], false),
963 (vec![ext_x.clone(), ext_x.clone(), ext_x.clone()], false),
964 (vec![ext_y.clone()], true),
965 (vec![ext_y.clone(), ext_y.clone()], false),
966 (vec![ext_y.clone(), ext_y.clone(), ext_y.clone()], false),
967 (vec![ext_x.clone(), ext_y.clone()], true),
968 (vec![ext_y.clone(), ext_x.clone()], true),
969 (vec![ext_x.clone(), ext_x.clone(), ext_y.clone()], false),
970 (vec![ext_y.clone(), ext_y.clone(), ext_x.clone()], false),
971 (vec![ext_x.clone(), ext_y.clone(), ext_y.clone()], false),
972 (vec![ext_y.clone(), ext_x.clone(), ext_x.clone()], false),
973 (vec![ext_x.clone(), ext_y.clone(), ext_x.clone()], false),
974 (vec![ext_y.clone(), ext_x, ext_y], false),
975 ];
976
977 for (test, should_work) in tests.into_iter() {
978 {
980 let mut extensions: Extensions<AnyObject> = Extensions::default();
981
982 let mut works = true;
983 for ext in test.iter() {
984 match extensions.add(ext.clone()) {
985 Ok(_) => {}
986 Err(InvalidExtensionError::Duplicate) => {
987 works = false;
988 }
989 _ => panic!("This should have never happened."),
990 }
991 }
992
993 println!("{:?}, {:?}", test.clone(), should_work);
994 assert_eq!(works, should_work);
995 }
996
997 if should_work {
999 assert!(Extensions::<AnyObject>::try_from(test).is_ok());
1000 } else {
1001 assert!(Extensions::<AnyObject>::try_from(test).is_err());
1002 }
1003 }
1004 }
1005
1006 #[test]
1007 fn ensure_ordering() {
1008 let ext_x = Extension::ApplicationId(ApplicationIdExtension::new(b"Test"));
1012 let ext_y = Extension::ExternalPub(ExternalPubExtension::new(HpkePublicKey::new(vec![])));
1013 let ext_z = Extension::RequiredCapabilities(RequiredCapabilitiesExtension::default());
1014
1015 for candidate in [ext_x, ext_y, ext_z]
1016 .into_iter()
1017 .permutations(3)
1018 .collect::<Vec<_>>()
1019 {
1020 let candidate: Extensions<AnyObject> = Extensions::try_from(candidate).unwrap();
1021 let bytes = candidate.tls_serialize_detached().unwrap();
1022 let got = Extensions::tls_deserialize(&mut bytes.as_slice()).unwrap();
1023 assert_eq!(candidate, got);
1024 }
1025 }
1026
1027 #[test]
1028 fn that_unknown_extensions_are_de_serialized_correctly() {
1029 let extension_types = [0x0000u16, 0x0A0A, 0x7A7A, 0xF100, 0xFFFF];
1030 let extension_datas = [vec![], vec![0], vec![1, 2, 3]];
1031
1032 for extension_type in extension_types.into_iter() {
1033 for extension_data in extension_datas.iter() {
1034 let test = {
1036 let mut buf = extension_type.to_be_bytes().to_vec();
1037 buf.append(
1038 &mut VLBytes::new(extension_data.clone())
1039 .tls_serialize_detached()
1040 .unwrap(),
1041 );
1042 buf
1043 };
1044
1045 let got = Extension::tls_deserialize_exact(&test).unwrap();
1047
1048 match got {
1049 Extension::Unknown(got_extension_type, ref got_extension_data) => {
1050 assert_eq!(extension_type, got_extension_type);
1051 assert_eq!(extension_data, &got_extension_data.0);
1052 }
1053 other => panic!("Expected `Extension::Unknown`, got {other:?}"),
1054 }
1055
1056 let got_serialized = got.tls_serialize_detached().unwrap();
1058 assert_eq!(test, got_serialized);
1059 }
1060 }
1061 }
1062}