Skip to main content

openmls/extensions/
mod.rs

1//! # Extensions
2//!
3//! In MLS, extensions appear in the following places:
4//!
5//! - In [`KeyPackages`](`crate::key_packages`) and [`LeafNode`](`crate::treesync::node::leaf_node::LeafNode`),
6//!   to describe client capabilities
7//!   and aspects of their participation in the group.
8//!
9//! - In `GroupInfo`, to inform new members of the group's parameters and to
10//!   provide any additional information required to join the group.
11//!
12//! - In the `GroupContext` object, to ensure that all members of the group have
13//!   a consistent view of the parameters in use.
14//!
15//! Note that `GroupInfo` and `GroupContext` are not exposed via OpenMLS' public
16//! API.
17//!
18//! OpenMLS supports the following extensions:
19//!
20//! - [`ApplicationIdExtension`] (KeyPackage extension)
21//! - [`RatchetTreeExtension`] (GroupInfo extension)
22//! - [`RequiredCapabilitiesExtension`] (GroupContext extension)
23//! - [`ExternalPubExtension`] (GroupInfo extension)
24//! - [`ExternalSendersExtension`] (GroupContext extension)
25//! - [`LastResortExtension`] (KeyPackage extension)
26
27use 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// Private
38#[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
50// Public
51pub mod errors;
52
53// Public re-exports
54#[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/// MLS Extension Types
79///
80/// Copied from draft-ietf-mls-protocol-16:
81///
82/// | Value            | Name                     | Message(s) | Recommended | Reference |
83/// |:-----------------|:-------------------------|:-----------|:------------|:----------|
84/// | 0x0000           | RESERVED                 | N/A        | N/A         | RFC XXXX  |
85/// | 0x0001           | application_id           | LN         | Y           | RFC XXXX  |
86/// | 0x0002           | ratchet_tree             | GI         | Y           | RFC XXXX  |
87/// | 0x0003           | required_capabilities    | GC         | Y           | RFC XXXX  |
88/// | 0x0004           | external_pub             | GI         | Y           | RFC XXXX  |
89/// | 0x0005           | external_senders         | GC         | Y           | RFC XXXX  |
90/// | 0xff00  - 0xffff | Reserved for Private Use | N/A        | N/A         | RFC XXXX  |
91///
92/// Note: OpenMLS does not provide a `Reserved` variant in [ExtensionType].
93#[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    /// The application id extension allows applications to add an explicit,
108    /// application-defined identifier to a KeyPackage.
109    ApplicationId,
110
111    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 1)]
112    /// The ratchet tree extensions provides the whole public state of the
113    /// ratchet tree.
114    RatchetTree,
115
116    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 2)]
117    /// The required capabilities extension defines the configuration of a group
118    /// that imposes certain requirements on clients in the group.
119    RequiredCapabilities,
120
121    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 3)]
122    /// To join a group via an External Commit, a new member needs a GroupInfo
123    /// with an ExternalPub extension present in its extensions field.
124    ExternalPub,
125
126    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 4)]
127    /// Group context extension that contains the credentials and signature keys
128    /// of senders that are permitted to send external proposals to the group.
129    ExternalSenders,
130
131    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 5)]
132    /// KeyPackage extension that marks a KeyPackage for use in a last resort
133    /// scenario.
134    LastResort,
135
136    #[cfg(feature = "extensions-draft")]
137    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 8)]
138    /// AppDataDictionary extension
139    AppDataDictionary,
140
141    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 7)]
142    /// A GREASE extension type for ensuring extensibility.
143    Grease(u16),
144
145    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 6)]
146    /// A currently unknown extension type.
147    Unknown(u16),
148}
149
150impl ExtensionType {
151    /// Returns true for all extension types that are considered "default" by the spec.
152    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    /// Returns whether an extension type is valid when used in leaf nodes.
168    /// Returns [`true`] for unknown extensions.
169    //  https://validation.openmls.tech/#valn1601
170    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            // GREASE may appear as an extension type in `leaf_node.extensions`
178            // and must be tolerated there. It is still subject to the normal rule
179            // that it must be declared in `capabilities` (checked separately),
180            // so this only permits the type, it does not exempt it from that check.
181            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            // GREASE is treated like an unknown extension type (tolerated): a
196            // GREASE-valued extension used to be reported as `Unknown` here, and
197            // must not become stricter now that it maps to `Grease`.
198            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            // GREASE may appear as an extension type in `key_package.extensions`
212            // and must be tolerated there. It is still subject to the normal rule
213            // that it be declared in `capabilities` (checked separately), so this
214            // only permits the type, it does not exempt it from that check.
215            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            // GREASE is treated like an unknown extension type: structurally
229            // allowed to appear here, but (like any unknown extension in the
230            // GroupContext) still subject to the per-member support check, which
231            // enforces the consensus rule that every member support it. GREASE
232            // does not bypass that check.
233            ExtensionType::Grease(_) => true,
234            #[cfg(feature = "extensions-draft")]
235            ExtensionType::AppDataDictionary => true,
236            _ => false,
237        }
238    }
239
240    /// Returns true if this is a GREASE extension type.
241    ///
242    /// GREASE values are used to ensure implementations properly handle unknown
243    /// extension types. See [RFC 9420 Section 13.5](https://www.rfc-editor.org/rfc/rfc9420.html#section-13.5).
244    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/// # Extension
321///
322/// An extension is one of the [`Extension`] enum values.
323/// The enum provides a set of common functionality for all extensions.
324///
325/// See the individual extensions for more details on each extension.
326///
327/// ```c
328/// // draft-ietf-mls-protocol-16
329/// struct {
330///     ExtensionType extension_type;
331///     opaque extension_data<V>;
332/// } Extension;
333/// ```
334#[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    /// An [`ApplicationIdExtension`]
349    ApplicationId(ApplicationIdExtension),
350
351    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 1)]
352    /// A [`RatchetTreeExtension`]
353    RatchetTree(RatchetTreeExtension),
354
355    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 2)]
356    /// A [`RequiredCapabilitiesExtension`]
357    RequiredCapabilities(RequiredCapabilitiesExtension),
358
359    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 3)]
360    /// An [`ExternalPubExtension`]
361    ExternalPub(ExternalPubExtension),
362
363    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 4)]
364    /// An [`ExternalSendersExtension`]
365    ExternalSenders(ExternalSendersExtension),
366
367    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 7)]
368    /// An [`AppDataDictionaryExtension`]
369    #[cfg(feature = "extensions-draft")]
370    AppDataDictionary(AppDataDictionaryExtension),
371
372    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 5)]
373    /// A [`LastResortExtension`]
374    LastResort(LastResortExtension),
375
376    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 6)]
377    /// A currently unknown extension.
378    Unknown(u16, UnknownExtension),
379}
380
381/// A unknown/unparsed extension represented by raw bytes.
382#[derive(
383    PartialEq, Eq, Clone, Debug, Serialize, Deserialize, TlsSize, TlsSerialize, TlsDeserialize,
384)]
385pub struct UnknownExtension(pub Vec<u8>);
386
387/// A Extension for Object of type T
388#[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)]
396/// Any object
397pub 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    /// Create an empty extension list.
450    pub fn empty() -> Self {
451        Self {
452            unique: vec![],
453            _object: PhantomData,
454        }
455    }
456
457    /// Returns an iterator over the extension list.
458    pub fn iter(&self) -> impl Iterator<Item = &Extension> {
459        self.unique.iter()
460    }
461
462    /// Remove an extension from the extension list.
463    ///
464    /// Returns the removed extension or `None` when there is no extension with
465    /// the given extension type.
466    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    /// Returns `true` iff the extension list contains an extension with the
479    /// given extension type.
480    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    /// Create an extension list with a single extension.
493    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    /// Create an extension list with multiple extensions.
502    ///
503    /// This function will fail when the list of extensions contains duplicate
504    /// extension types.
505    pub fn from_vec(extensions: Vec<Extension>) -> Result<Self, InvalidExtensionError> {
506        extensions.try_into()
507    }
508
509    /// Validate if the extensions are valid for this context
510    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    /// Add an extension to the extension list.
520    ///
521    /// Returns an error when there already is an extension with the same
522    /// extension type.
523    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    /// Add an extension to the extension list (or replace an existing one.)
535    ///
536    /// Returns the replaced extension (if any).
537    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    /// Assume that the extensions contain the given extension type.
550    ///
551    /// # Safety
552    ///
553    /// The caller must guarantee that the extensions are of the correct type.
554    #[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    /// Used to seal other traits
565    pub trait Sealed {}
566}
567
568/// Can be implemented by a type to validate extensions.
569pub trait ExtensionValidator: private::Sealed {
570    /// The error returned by the validator
571    type Error;
572
573    /// Check if the extension is valid.
574    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
612// https://validation.openmls.tech/#valn1602
613impl 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
631// https://validation.openmls.tech/#valn1603
632impl 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
648// https://validation.openmls.tech/#valn1604
649impl 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
665// https://validation.openmls.tech/#valn1601
666impl 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    /// Get a reference to the [`ApplicationIdExtension`] if there is any.
688    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    /// Get a reference to the [`RatchetTreeExtension`] if there is any.
697    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    /// Get a reference to the [`RequiredCapabilitiesExtension`] if there is
706    /// any.
707    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    /// Get a reference to the [`ExternalPubExtension`] if there is any.
716    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    /// Get a reference to the [`ExternalSendersExtension`] if there is any.
725    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    /// Get a reference to the [`AppDataDictionaryExtension`] if there is any.
735    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    /// Get a reference to the [`UnknownExtension`] with the given type id, if there is any.
744    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    /// Get a reference to this extension as [`ApplicationIdExtension`].
761    /// Returns an [`ExtensionError::InvalidExtensionType`] if called on an
762    /// [`Extension`] that's not an [`ApplicationIdExtension`].
763    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    /// Get a reference to this extension as [`AppDataDictionaryExtension`].
773    /// Returns an [`ExtensionError::InvalidExtensionType`] if called on an
774    /// [`Extension`] that's not an [`AppDataDictionaryExtension`].
775    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    /// Get a reference to this extension as [`RatchetTreeExtension`].
787    /// Returns an [`ExtensionError::InvalidExtensionType`] if called on
788    /// an [`Extension`] that's not a [`RatchetTreeExtension`].
789    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    /// Get a reference to this extension as [`RequiredCapabilitiesExtension`].
799    /// Returns an [`ExtensionError::InvalidExtensionType`] error if called on
800    /// an [`Extension`] that's not a [`RequiredCapabilitiesExtension`].
801    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    /// Get a reference to this extension as [`ExternalPubExtension`].
813    /// Returns an [`ExtensionError::InvalidExtensionType`] error if called on
814    /// an [`Extension`] that's not a [`ExternalPubExtension`].
815    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    /// Get a reference to this extension as [`ExternalSendersExtension`].
825    /// Returns an [`ExtensionError::InvalidExtensionType`] error if called on
826    /// an [`Extension`] that's not a [`ExternalSendersExtension`].
827    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    /// Returns the [`ExtensionType`]
839    #[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            // Map GREASE-valued extension types to `Grease`, consistent with
851            // `ExtensionType::from(u16)`. Without this an extension carrying a
852            // GREASE value would be reported as `Unknown`, so GREASE-aware
853            // validation (which ignores `Grease(_)`) would not recognize it.
854            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    /// Coerces the extensions to an Extensions with the given validator. Unsafe.
899    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        // A GREASE-valued extension must report a `Grease` extension type
931        // (consistent with `ExtensionType::from(u16)`), not `Unknown`. Otherwise
932        // GREASE-aware validation would fail to recognize it and reject peers
933        // (e.g. MLS++) that decorate leaf/key-package extensions with GREASE.
934        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        // A non-GREASE unknown value stays `Unknown`.
939        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        // GREASE does NOT bypass the capability check: a GREASE extension type is
946        // "contained" only if it is advertised in the capabilities (RFC 9420:
947        // extensions in leaf_node.extensions/key_package.extensions MUST be in
948        // capabilities). The fix that makes this work is that a GREASE-valued
949        // extension now reports a `Grease(_)` type that matches the `Grease(_)`
950        // parsed into the capabilities list.
951        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        // A GREASE value that is not advertised is not contained.
960        assert!(!advertised.contains_extension_type(&ExtensionType::Grease(0xAAAA)));
961
962        // Empty capabilities contain no (non-default) extension, GREASE included.
963        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        // Create some extensions with different extension types and test that
972        // duplicates are rejected. The extension content does not matter in this test.
973        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            // Test `add`.
996            {
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            // Test `try_from`.
1015            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        // Create some extensions with different extension types and test
1026        // that all permutations keep their order after being (de)serialized.
1027        // The extension content does not matter in this test.
1028        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                // Construct an unknown extension manually.
1052                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                // Test deserialization.
1063                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                // Test serialization.
1074                let got_serialized = got.tls_serialize_detached().unwrap();
1075                assert_eq!(test, got_serialized);
1076            }
1077        }
1078    }
1079}