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`), to describe client capabilities
6//!   and aspects of their participation in the group.
7//!
8//! - In `GroupInfo`, to inform new members of the group's parameters and to
9//!   provide any additional information required to join the group.
10//!
11//! - In the `GroupContext` object, to ensure that all members of the group have
12//!   a consistent view of the parameters in use.
13//!
14//! Note that `GroupInfo` and `GroupContext` are not exposed via OpenMLS' public
15//! API.
16//!
17//! OpenMLS supports the following extensions:
18//!
19//! - [`ApplicationIdExtension`] (KeyPackage extension)
20//! - [`RatchetTreeExtension`] (GroupInfo extension)
21//! - [`RequiredCapabilitiesExtension`] (GroupContext extension)
22//! - [`ExternalPubExtension`] (GroupInfo extension)
23
24use 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// Private
35#[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
46// Public
47pub mod errors;
48
49// Public re-exports
50#[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/// MLS Extension Types
75///
76/// Copied from draft-ietf-mls-protocol-16:
77///
78/// | Value            | Name                     | Message(s) | Recommended | Reference |
79/// |:-----------------|:-------------------------|:-----------|:------------|:----------|
80/// | 0x0000           | RESERVED                 | N/A        | N/A         | RFC XXXX  |
81/// | 0x0001           | application_id           | LN         | Y           | RFC XXXX  |
82/// | 0x0002           | ratchet_tree             | GI         | Y           | RFC XXXX  |
83/// | 0x0003           | required_capabilities    | GC         | Y           | RFC XXXX  |
84/// | 0x0004           | external_pub             | GI         | Y           | RFC XXXX  |
85/// | 0x0005           | external_senders         | GC         | Y           | RFC XXXX  |
86/// | 0xff00  - 0xffff | Reserved for Private Use | N/A        | N/A         | RFC XXXX  |
87///
88/// Note: OpenMLS does not provide a `Reserved` variant in [ExtensionType].
89#[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    /// The application id extension allows applications to add an explicit,
104    /// application-defined identifier to a KeyPackage.
105    ApplicationId,
106
107    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 1)]
108    /// The ratchet tree extensions provides the whole public state of the
109    /// ratchet tree.
110    RatchetTree,
111
112    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 2)]
113    /// The required capabilities extension defines the configuration of a group
114    /// that imposes certain requirements on clients in the group.
115    RequiredCapabilities,
116
117    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 3)]
118    /// To join a group via an External Commit, a new member needs a GroupInfo
119    /// with an ExternalPub extension present in its extensions field.
120    ExternalPub,
121
122    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 4)]
123    /// Group context extension that contains the credentials and signature keys
124    /// of senders that are permitted to send external proposals to the group.
125    ExternalSenders,
126
127    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 5)]
128    /// KeyPackage extension that marks a KeyPackage for use in a last resort
129    /// scenario.
130    LastResort,
131
132    #[cfg(feature = "extensions-draft")]
133    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 8)]
134    /// AppDataDictionary extension
135    AppDataDictionary,
136
137    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 7)]
138    /// A GREASE extension type for ensuring extensibility.
139    Grease(u16),
140
141    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 6)]
142    /// A currently unknown extension type.
143    Unknown(u16),
144}
145
146impl ExtensionType {
147    /// Returns true for all extension types that are considered "default" by the spec.
148    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    /// Returns whether an extension type is valid when used in leaf nodes.
164    /// Returns None if validity can not be determined.
165    /// This is the case for unknown extensions.
166    //  https://validation.openmls.tech/#valn1601
167    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            // GREASE may appear as an extension type in `leaf_node.extensions`
175            // and must be tolerated there. It is still subject to the normal rule
176            // that it must be declared in `capabilities` (checked separately),
177            // so this only permits the type, it does not exempt it from that check.
178            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            // GREASE is treated like an unknown extension type (tolerated): a
193            // GREASE-valued extension used to be reported as `Unknown` here, and
194            // must not become stricter now that it maps to `Grease`.
195            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            // GREASE may appear as an extension type in `key_package.extensions`
209            // and must be tolerated there. It is still subject to the normal rule
210            // that it be declared in `capabilities` (checked separately), so this
211            // only permits the type, it does not exempt it from that check.
212            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            // GREASE is treated like an unknown extension type: structurally
226            // allowed to appear here, but (like any unknown extension in the
227            // GroupContext) still subject to the per-member support check, which
228            // enforces the consensus rule that every member support it. GREASE
229            // does not bypass that check.
230            ExtensionType::Grease(_) => true,
231            #[cfg(feature = "extensions-draft")]
232            ExtensionType::AppDataDictionary => true,
233            _ => false,
234        }
235    }
236
237    /// Returns true if this is a GREASE extension type.
238    ///
239    /// GREASE values are used to ensure implementations properly handle unknown
240    /// extension types. See [RFC 9420 Section 13.5](https://www.rfc-editor.org/rfc/rfc9420.html#section-13.5).
241    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/// # Extension
318///
319/// An extension is one of the [`Extension`] enum values.
320/// The enum provides a set of common functionality for all extensions.
321///
322/// See the individual extensions for more details on each extension.
323///
324/// ```c
325/// // draft-ietf-mls-protocol-16
326/// struct {
327///     ExtensionType extension_type;
328///     opaque extension_data<V>;
329/// } Extension;
330/// ```
331#[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    /// An [`ApplicationIdExtension`]
346    ApplicationId(ApplicationIdExtension),
347
348    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 1)]
349    /// A [`RatchetTreeExtension`]
350    RatchetTree(RatchetTreeExtension),
351
352    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 2)]
353    /// A [`RequiredCapabilitiesExtension`]
354    RequiredCapabilities(RequiredCapabilitiesExtension),
355
356    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 3)]
357    /// An [`ExternalPubExtension`]
358    ExternalPub(ExternalPubExtension),
359
360    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 4)]
361    /// An [`ExternalSendersExtension`]
362    ExternalSenders(ExternalSendersExtension),
363
364    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 7)]
365    /// An [`AppDataDictionaryExtension`]
366    #[cfg(feature = "extensions-draft")]
367    AppDataDictionary(AppDataDictionaryExtension),
368
369    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 5)]
370    /// A [`LastResortExtension`]
371    LastResort(LastResortExtension),
372
373    #[cfg_attr(not(feature = "0-8-1-storage-format"), storage_tag = 6)]
374    /// A currently unknown extension.
375    Unknown(u16, UnknownExtension),
376}
377
378/// A unknown/unparsed extension represented by raw bytes.
379#[derive(
380    PartialEq, Eq, Clone, Debug, Serialize, Deserialize, TlsSize, TlsSerialize, TlsDeserialize,
381)]
382pub struct UnknownExtension(pub Vec<u8>);
383
384/// A Extension for Object of type T
385#[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)]
393/// Any object
394pub 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    /// Create an empty extension list.
447    pub fn empty() -> Self {
448        Self {
449            unique: vec![],
450            _object: PhantomData,
451        }
452    }
453
454    /// Returns an iterator over the extension list.
455    pub fn iter(&self) -> impl Iterator<Item = &Extension> {
456        self.unique.iter()
457    }
458
459    /// Remove an extension from the extension list.
460    ///
461    /// Returns the removed extension or `None` when there is no extension with
462    /// the given extension type.
463    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    /// Returns `true` iff the extension list contains an extension with the
476    /// given extension type.
477    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    /// Create an extension list with a single extension.
490    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    /// Create an extension list with multiple extensions.
499    ///
500    /// This function will fail when the list of extensions contains duplicate
501    /// extension types.
502    pub fn from_vec(extensions: Vec<Extension>) -> Result<Self, InvalidExtensionError> {
503        extensions.try_into()
504    }
505
506    /// Validate if the extensions are valid for this context
507    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    /// Add an extension to the extension list.
517    ///
518    /// Returns an error when there already is an extension with the same
519    /// extension type.
520    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    /// Add an extension to the extension list (or replace an existing one.)
532    ///
533    /// Returns the replaced extension (if any).
534    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    /// Assume that the extensions contain the given extension type.
547    ///
548    /// # Safety
549    ///
550    /// The caller must guarantee that the extensions are of the correct type.
551    #[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
560/// Can be implemented by a type to validate extensions.
561pub trait ExtensionValidator {
562    /// The error returned by the validator
563    type Error;
564
565    /// Check if the extension is valid.
566    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
600// https://validation.openmls.tech/#valn1602
601impl 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
617// https://validation.openmls.tech/#valn1603
618impl 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
634// https://validation.openmls.tech/#valn1604
635impl 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
649// https://validation.openmls.tech/#valn1601
650impl 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    /// Get a reference to the [`ApplicationIdExtension`] if there is any.
672    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    /// Get a reference to the [`RatchetTreeExtension`] if there is any.
681    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    /// Get a reference to the [`RequiredCapabilitiesExtension`] if there is
690    /// any.
691    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    /// Get a reference to the [`ExternalPubExtension`] if there is any.
700    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    /// Get a reference to the [`ExternalSendersExtension`] if there is any.
709    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    /// Get a reference to the [`AppDataDictionaryExtension`] if there is any.
719    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    /// Get a reference to the [`UnknownExtension`] with the given type id, if there is any.
728    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    /// Get a reference to this extension as [`ApplicationIdExtension`].
745    /// Returns an [`ExtensionError::InvalidExtensionType`] if called on an
746    /// [`Extension`] that's not an [`ApplicationIdExtension`].
747    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    /// Get a reference to this extension as [`AppDataDictionaryExtension`].
757    /// Returns an [`ExtensionError::InvalidExtensionType`] if called on an
758    /// [`Extension`] that's not an [`AppDataDictionaryExtension`].
759    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    /// Get a reference to this extension as [`RatchetTreeExtension`].
771    /// Returns an [`ExtensionError::InvalidExtensionType`] if called on
772    /// an [`Extension`] that's not a [`RatchetTreeExtension`].
773    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    /// Get a reference to this extension as [`RequiredCapabilitiesExtension`].
783    /// Returns an [`ExtensionError::InvalidExtensionType`] error if called on
784    /// an [`Extension`] that's not a [`RequiredCapabilitiesExtension`].
785    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    /// Get a reference to this extension as [`ExternalPubExtension`].
797    /// Returns an [`ExtensionError::InvalidExtensionType`] error if called on
798    /// an [`Extension`] that's not a [`ExternalPubExtension`].
799    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    /// Get a reference to this extension as [`ExternalSendersExtension`].
809    /// Returns an [`ExtensionError::InvalidExtensionType`] error if called on
810    /// an [`Extension`] that's not a [`ExternalSendersExtension`].
811    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    /// Returns the [`ExtensionType`]
823    #[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            // Map GREASE-valued extension types to `Grease`, consistent with
835            // `ExtensionType::from(u16)`. Without this an extension carrying a
836            // GREASE value would be reported as `Unknown`, so GREASE-aware
837            // validation (which ignores `Grease(_)`) would not recognize it.
838            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    /// Coerces the extensions to an Extensions with the given validator. Unsafe.
882    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        // A GREASE-valued extension must report a `Grease` extension type
914        // (consistent with `ExtensionType::from(u16)`), not `Unknown`. Otherwise
915        // GREASE-aware validation would fail to recognize it and reject peers
916        // (e.g. MLS++) that decorate leaf/key-package extensions with GREASE.
917        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        // A non-GREASE unknown value stays `Unknown`.
922        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        // GREASE does NOT bypass the capability check: a GREASE extension type is
929        // "contained" only if it is advertised in the capabilities (RFC 9420:
930        // extensions in leaf_node.extensions/key_package.extensions MUST be in
931        // capabilities). The fix that makes this work is that a GREASE-valued
932        // extension now reports a `Grease(_)` type that matches the `Grease(_)`
933        // parsed into the capabilities list.
934        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        // A GREASE value that is not advertised is not contained.
943        assert!(!advertised.contains_extension_type(&ExtensionType::Grease(0xAAAA)));
944
945        // Empty capabilities contain no (non-default) extension, GREASE included.
946        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        // Create some extensions with different extension types and test that
955        // duplicates are rejected. The extension content does not matter in this test.
956        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            // Test `add`.
979            {
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            // Test `try_from`.
998            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        // Create some extensions with different extension types and test
1009        // that all permutations keep their order after being (de)serialized.
1010        // The extension content does not matter in this test.
1011        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                // Construct an unknown extension manually.
1035                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                // Test deserialization.
1046                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                // Test serialization.
1057                let got_serialized = got.tls_serialize_detached().unwrap();
1058                assert_eq!(test, got_serialized);
1059            }
1060        }
1061    }
1062}