Skip to main content

mls_rs/tree_kem/
leaf_node.rs

1// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
2// Copyright by contributors to this project.
3// SPDX-License-Identifier: (Apache-2.0 OR MIT)
4
5use super::{parent_hash::ParentHash, Capabilities, Lifetime};
6use crate::client::MlsError;
7use crate::crypto::{CipherSuiteProvider, HpkePublicKey, HpkeSecretKey, SignatureSecretKey};
8use crate::{identity::SigningIdentity, signer::Signable, ExtensionList};
9use alloc::vec::Vec;
10use core::fmt::{self, Debug};
11use mls_rs_codec::{MlsDecode, MlsEncode, MlsSize};
12use mls_rs_core::error::IntoAnyError;
13
14#[derive(Debug, Clone, MlsSize, MlsEncode, MlsDecode, PartialEq, Eq)]
15#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
16#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
17#[repr(u8)]
18pub enum LeafNodeSource {
19    KeyPackage(Lifetime) = 1u8,
20    Update = 2u8,
21    Commit(ParentHash) = 3u8,
22}
23
24#[derive(Clone, MlsSize, MlsEncode, MlsDecode, PartialEq)]
25#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
26#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
27#[non_exhaustive]
28pub struct LeafNode {
29    pub public_key: HpkePublicKey,
30    pub signing_identity: SigningIdentity,
31    pub capabilities: Capabilities,
32    pub leaf_node_source: LeafNodeSource,
33    pub extensions: ExtensionList,
34    #[mls_codec(with = "mls_rs_codec::byte_vec")]
35    #[cfg_attr(feature = "serde", serde(with = "mls_rs_core::vec_serde"))]
36    pub signature: Vec<u8>,
37}
38
39impl Debug for LeafNode {
40    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
41        f.debug_struct("LeafNode")
42            .field("public_key", &self.public_key)
43            .field("signing_identity", &self.signing_identity)
44            .field("capabilities", &self.capabilities)
45            .field("leaf_node_source", &self.leaf_node_source)
46            .field("extensions", &self.extensions)
47            .field(
48                "signature",
49                &mls_rs_core::debug::pretty_bytes(&self.signature),
50            )
51            .finish()
52    }
53}
54
55#[derive(Clone, Debug)]
56pub struct ConfigProperties {
57    pub capabilities: Capabilities,
58    pub extensions: ExtensionList,
59}
60
61impl LeafNode {
62    #[cfg_attr(not(mls_build_async), maybe_async::must_be_sync)]
63    pub async fn generate<CSP>(
64        cipher_suite_provider: &CSP,
65        properties: ConfigProperties,
66        signing_identity: SigningIdentity,
67        signer: &SignatureSecretKey,
68        lifetime: Lifetime,
69    ) -> Result<(Self, HpkeSecretKey), MlsError>
70    where
71        CSP: CipherSuiteProvider,
72    {
73        let (secret_key, public_key) = cipher_suite_provider
74            .kem_generate()
75            .await
76            .map_err(|e| MlsError::CryptoProviderError(e.into_any_error()))?;
77
78        let mut leaf_node = LeafNode {
79            public_key,
80            signing_identity,
81            capabilities: properties.capabilities,
82            leaf_node_source: LeafNodeSource::KeyPackage(lifetime),
83            extensions: properties.extensions,
84            signature: Default::default(),
85        };
86
87        leaf_node.grease(cipher_suite_provider)?;
88
89        leaf_node
90            .sign(
91                cipher_suite_provider,
92                signer,
93                &LeafNodeSigningContext::default(),
94            )
95            .await?;
96
97        Ok((leaf_node, secret_key))
98    }
99
100    #[cfg_attr(not(mls_build_async), maybe_async::must_be_sync)]
101    pub async fn update<P: CipherSuiteProvider>(
102        &mut self,
103        cipher_suite_provider: &P,
104        group_id: &[u8],
105        leaf_index: u32,
106        new_properties: Option<ConfigProperties>,
107        signing_identity: Option<SigningIdentity>,
108        signer: &SignatureSecretKey,
109    ) -> Result<HpkeSecretKey, MlsError> {
110        let (secret, public) = cipher_suite_provider
111            .kem_generate()
112            .await
113            .map_err(|e| MlsError::CryptoProviderError(e.into_any_error()))?;
114
115        self.public_key = public;
116
117        if let Some(new_properties) = new_properties {
118            self.capabilities = new_properties.capabilities;
119            self.extensions = new_properties.extensions;
120        }
121
122        self.leaf_node_source = LeafNodeSource::Update;
123
124        self.grease(cipher_suite_provider)?;
125
126        if let Some(signing_identity) = signing_identity {
127            self.signing_identity = signing_identity;
128        }
129
130        self.sign(
131            cipher_suite_provider,
132            signer,
133            &(group_id, leaf_index).into(),
134        )
135        .await?;
136
137        Ok(secret)
138    }
139
140    #[allow(clippy::too_many_arguments)]
141    #[cfg_attr(not(mls_build_async), maybe_async::must_be_sync)]
142    pub async fn commit<P: CipherSuiteProvider>(
143        &mut self,
144        cipher_suite_provider: &P,
145        group_id: &[u8],
146        leaf_index: u32,
147        new_properties: Option<ConfigProperties>,
148        new_signing_identity: Option<SigningIdentity>,
149        signer: &SignatureSecretKey,
150    ) -> Result<HpkeSecretKey, MlsError> {
151        let (secret, public) = cipher_suite_provider
152            .kem_generate()
153            .await
154            .map_err(|e| MlsError::CryptoProviderError(e.into_any_error()))?;
155
156        self.public_key = public;
157
158        if let Some(new_properties) = new_properties {
159            self.capabilities = new_properties.capabilities;
160            self.extensions = new_properties.extensions;
161        }
162
163        if let Some(new_signing_identity) = new_signing_identity {
164            self.signing_identity = new_signing_identity;
165        }
166
167        self.sign(
168            cipher_suite_provider,
169            signer,
170            &(group_id, leaf_index).into(),
171        )
172        .await?;
173
174        Ok(secret)
175    }
176}
177
178#[derive(Debug)]
179struct LeafNodeTBS<'a> {
180    public_key: &'a HpkePublicKey,
181    signing_identity: &'a SigningIdentity,
182    capabilities: &'a Capabilities,
183    leaf_node_source: &'a LeafNodeSource,
184    extensions: &'a ExtensionList,
185    group_id: Option<&'a [u8]>,
186    leaf_index: Option<u32>,
187}
188
189impl MlsSize for LeafNodeTBS<'_> {
190    fn mls_encoded_len(&self) -> usize {
191        self.public_key.mls_encoded_len()
192            + self.signing_identity.mls_encoded_len()
193            + self.capabilities.mls_encoded_len()
194            + self.leaf_node_source.mls_encoded_len()
195            + self.extensions.mls_encoded_len()
196            + self
197                .group_id
198                .as_ref()
199                .map_or(0, mls_rs_codec::byte_vec::mls_encoded_len)
200            + self.leaf_index.map_or(0, |i| i.mls_encoded_len())
201    }
202}
203
204impl MlsEncode for LeafNodeTBS<'_> {
205    fn mls_encode(&self, writer: &mut Vec<u8>) -> Result<(), mls_rs_codec::Error> {
206        self.public_key.mls_encode(writer)?;
207        self.signing_identity.mls_encode(writer)?;
208        self.capabilities.mls_encode(writer)?;
209        self.leaf_node_source.mls_encode(writer)?;
210        self.extensions.mls_encode(writer)?;
211
212        if let Some(ref group_id) = self.group_id {
213            mls_rs_codec::byte_vec::mls_encode(group_id, writer)?;
214        }
215
216        if let Some(leaf_index) = self.leaf_index {
217            leaf_index.mls_encode(writer)?;
218        }
219
220        Ok(())
221    }
222}
223
224#[derive(Clone, Debug, Default)]
225pub(crate) struct LeafNodeSigningContext<'a> {
226    pub group_id: Option<&'a [u8]>,
227    pub leaf_index: Option<u32>,
228}
229
230impl<'a> From<(&'a [u8], u32)> for LeafNodeSigningContext<'a> {
231    fn from((group_id, leaf_index): (&'a [u8], u32)) -> Self {
232        Self {
233            group_id: Some(group_id),
234            leaf_index: Some(leaf_index),
235        }
236    }
237}
238
239impl<'a> Signable<'a> for LeafNode {
240    const SIGN_LABEL: &'static str = "LeafNodeTBS";
241
242    type SigningContext = LeafNodeSigningContext<'a>;
243
244    fn signature(&self) -> &[u8] {
245        &self.signature
246    }
247
248    fn signable_content(
249        &self,
250        context: &Self::SigningContext,
251    ) -> Result<Vec<u8>, mls_rs_codec::Error> {
252        LeafNodeTBS {
253            public_key: &self.public_key,
254            signing_identity: &self.signing_identity,
255            capabilities: &self.capabilities,
256            leaf_node_source: &self.leaf_node_source,
257            extensions: &self.extensions,
258            group_id: context.group_id,
259            leaf_index: context.leaf_index,
260        }
261        .mls_encode_to_vec()
262    }
263
264    fn write_signature(&mut self, signature: Vec<u8>) {
265        self.signature = signature
266    }
267}
268
269#[cfg(test)]
270pub(crate) mod test_utils {
271    use alloc::vec;
272    use mls_rs_core::identity::{BasicCredential, CredentialType};
273
274    use crate::{
275        cipher_suite::CipherSuite,
276        crypto::test_utils::{test_cipher_suite_provider, TestCryptoProvider},
277        identity::test_utils::{get_test_signing_identity, BasicWithCustomProvider},
278    };
279
280    use crate::extension::ApplicationIdExt;
281
282    use super::*;
283
284    #[allow(unused)]
285    #[cfg_attr(not(mls_build_async), maybe_async::must_be_sync)]
286    pub async fn get_test_node(
287        cipher_suite: CipherSuite,
288        signing_identity: SigningIdentity,
289        secret: &SignatureSecretKey,
290        capabilities: Option<Capabilities>,
291        extensions: Option<ExtensionList>,
292    ) -> (LeafNode, HpkeSecretKey) {
293        get_test_node_with_lifetime(
294            cipher_suite,
295            signing_identity,
296            secret,
297            capabilities.unwrap_or_else(get_test_capabilities),
298            extensions.unwrap_or_default(),
299            Lifetime::years(1, None).unwrap(),
300        )
301        .await
302    }
303
304    #[cfg_attr(not(mls_build_async), maybe_async::must_be_sync)]
305    pub async fn get_test_node_with_lifetime(
306        cipher_suite: CipherSuite,
307        signing_identity: SigningIdentity,
308        secret: &SignatureSecretKey,
309        capabilities: Capabilities,
310        extensions: ExtensionList,
311        lifetime: Lifetime,
312    ) -> (LeafNode, HpkeSecretKey) {
313        let properties = ConfigProperties {
314            capabilities,
315            extensions,
316        };
317
318        LeafNode::generate(
319            &test_cipher_suite_provider(cipher_suite),
320            properties,
321            signing_identity,
322            secret,
323            lifetime,
324        )
325        .await
326        .unwrap()
327    }
328
329    #[allow(unused)]
330    #[cfg_attr(not(mls_build_async), maybe_async::must_be_sync)]
331    pub async fn get_basic_test_node(cipher_suite: CipherSuite, id: &str) -> LeafNode {
332        get_basic_test_node_sig_key(cipher_suite, id).await.0
333    }
334
335    #[allow(unused)]
336    pub fn default_properties() -> ConfigProperties {
337        ConfigProperties {
338            capabilities: get_test_capabilities(),
339            extensions: Default::default(),
340        }
341    }
342
343    #[cfg_attr(not(mls_build_async), maybe_async::must_be_sync)]
344    pub async fn get_basic_test_node_capabilities(
345        cipher_suite: CipherSuite,
346        id: &str,
347        capabilities: Capabilities,
348    ) -> (LeafNode, HpkeSecretKey, SignatureSecretKey) {
349        let (signing_identity, signature_key) =
350            get_test_signing_identity(cipher_suite, id.as_bytes()).await;
351
352        LeafNode::generate(
353            &test_cipher_suite_provider(cipher_suite),
354            ConfigProperties {
355                capabilities,
356                extensions: Default::default(),
357            },
358            signing_identity,
359            &signature_key,
360            Lifetime::years(1, None).unwrap(),
361        )
362        .await
363        .map(|(leaf, hpke_secret_key)| (leaf, hpke_secret_key, signature_key))
364        .unwrap()
365    }
366
367    #[cfg_attr(not(mls_build_async), maybe_async::must_be_sync)]
368    pub async fn get_basic_test_node_sig_key(
369        cipher_suite: CipherSuite,
370        id: &str,
371    ) -> (LeafNode, HpkeSecretKey, SignatureSecretKey) {
372        get_basic_test_node_capabilities(cipher_suite, id, get_test_capabilities()).await
373    }
374
375    #[allow(unused)]
376    pub fn get_test_extensions() -> ExtensionList {
377        let mut extension_list = ExtensionList::new();
378
379        extension_list
380            .set_from(ApplicationIdExt {
381                identifier: b"identifier".to_vec(),
382            })
383            .unwrap();
384
385        extension_list
386    }
387
388    pub fn get_test_capabilities() -> Capabilities {
389        Capabilities {
390            credentials: vec![
391                BasicCredential::credential_type(),
392                CredentialType::from(BasicWithCustomProvider::CUSTOM_CREDENTIAL_TYPE),
393            ],
394            cipher_suites: TestCryptoProvider::all_supported_cipher_suites(),
395            ..Default::default()
396        }
397    }
398
399    #[allow(unused)]
400    pub fn get_test_client_identity(leaf: &LeafNode) -> Vec<u8> {
401        leaf.signing_identity
402            .credential
403            .mls_encode_to_vec()
404            .unwrap()
405    }
406}
407
408#[cfg(test)]
409mod tests {
410    use super::test_utils::*;
411    use super::*;
412
413    use crate::client::test_utils::TEST_CIPHER_SUITE;
414    use crate::crypto::test_utils::test_cipher_suite_provider;
415    use crate::crypto::test_utils::TestCryptoProvider;
416    use crate::group::test_utils::random_bytes;
417    use crate::identity::test_utils::get_test_signing_identity;
418    use assert_matches::assert_matches;
419
420    #[maybe_async::test(not(mls_build_async), async(mls_build_async, crate::futures_test))]
421    async fn test_node_generation() {
422        let capabilities = get_test_capabilities();
423        let extensions = get_test_extensions();
424        let lifetime = Lifetime::years(1, None).unwrap();
425
426        for cipher_suite in TestCryptoProvider::all_supported_cipher_suites() {
427            let (signing_identity, secret) = get_test_signing_identity(cipher_suite, b"foo").await;
428
429            let (leaf_node, secret_key) = get_test_node_with_lifetime(
430                cipher_suite,
431                signing_identity.clone(),
432                &secret,
433                capabilities.clone(),
434                extensions.clone(),
435                lifetime.clone(),
436            )
437            .await;
438
439            assert_eq!(leaf_node.ungreased_capabilities(), capabilities);
440            assert_eq!(leaf_node.ungreased_extensions(), extensions);
441            assert_eq!(leaf_node.signing_identity, signing_identity);
442
443            assert_matches!(
444                &leaf_node.leaf_node_source,
445                LeafNodeSource::KeyPackage(lt) if lt == &lifetime,
446                "Expected {:?}, got {:?}", LeafNodeSource::KeyPackage(lifetime),
447                leaf_node.leaf_node_source
448            );
449
450            let provider = test_cipher_suite_provider(cipher_suite);
451
452            // Verify that the hpke key pair generated will work
453            let test_data = random_bytes(32);
454
455            let sealed = provider
456                .hpke_seal(&leaf_node.public_key, &[], None, &test_data)
457                .await
458                .unwrap();
459
460            let opened = provider
461                .hpke_open(&sealed, &secret_key, &leaf_node.public_key, &[], None)
462                .await
463                .unwrap();
464
465            assert_eq!(*opened, test_data);
466
467            leaf_node
468                .verify(
469                    &test_cipher_suite_provider(cipher_suite),
470                    &signing_identity.signature_key,
471                    &LeafNodeSigningContext::default(),
472                )
473                .await
474                .unwrap();
475        }
476    }
477
478    #[maybe_async::test(not(mls_build_async), async(mls_build_async, crate::futures_test))]
479    async fn test_node_generation_randomness() {
480        let cipher_suite = TEST_CIPHER_SUITE;
481
482        let (signing_identity, secret) = get_test_signing_identity(cipher_suite, b"foo").await;
483
484        let (first_leaf, first_secret) =
485            get_test_node(cipher_suite, signing_identity.clone(), &secret, None, None).await;
486
487        for _ in 0..100 {
488            let (next_leaf, next_secret) =
489                get_test_node(cipher_suite, signing_identity.clone(), &secret, None, None).await;
490
491            assert_ne!(first_secret, next_secret);
492            assert_ne!(first_leaf.public_key, next_leaf.public_key);
493        }
494    }
495
496    #[maybe_async::test(not(mls_build_async), async(mls_build_async, crate::futures_test))]
497    async fn test_node_update_no_meta_changes() {
498        for cipher_suite in TestCryptoProvider::all_supported_cipher_suites() {
499            let cipher_suite_provider = test_cipher_suite_provider(cipher_suite);
500
501            let (signing_identity, secret) = get_test_signing_identity(cipher_suite, b"foo").await;
502
503            let (mut leaf, leaf_secret) =
504                get_test_node(cipher_suite, signing_identity.clone(), &secret, None, None).await;
505
506            let original_leaf = leaf.clone();
507
508            let new_secret = leaf
509                .update(
510                    &cipher_suite_provider,
511                    b"group",
512                    0,
513                    Some(default_properties()),
514                    None,
515                    &secret,
516                )
517                .await
518                .unwrap();
519
520            assert_ne!(new_secret, leaf_secret);
521            assert_ne!(original_leaf.public_key, leaf.public_key);
522
523            assert_eq!(
524                leaf.ungreased_capabilities(),
525                original_leaf.ungreased_capabilities()
526            );
527
528            assert_eq!(
529                leaf.ungreased_extensions(),
530                original_leaf.ungreased_extensions()
531            );
532
533            assert_eq!(leaf.signing_identity, original_leaf.signing_identity);
534            assert_matches!(&leaf.leaf_node_source, LeafNodeSource::Update);
535
536            leaf.verify(
537                &cipher_suite_provider,
538                &signing_identity.signature_key,
539                &(b"group".as_slice(), 0).into(),
540            )
541            .await
542            .unwrap();
543        }
544    }
545
546    #[maybe_async::test(not(mls_build_async), async(mls_build_async, crate::futures_test))]
547    async fn test_node_update_meta_changes() {
548        let cipher_suite = TEST_CIPHER_SUITE;
549
550        let (signing_identity, secret) = get_test_signing_identity(cipher_suite, b"foo").await;
551
552        let new_properties = ConfigProperties {
553            capabilities: get_test_capabilities(),
554            extensions: get_test_extensions(),
555        };
556
557        let (mut leaf, _) =
558            get_test_node(cipher_suite, signing_identity, &secret, None, None).await;
559
560        leaf.update(
561            &test_cipher_suite_provider(cipher_suite),
562            b"group",
563            0,
564            Some(new_properties.clone()),
565            None,
566            &secret,
567        )
568        .await
569        .unwrap();
570
571        assert_eq!(leaf.ungreased_capabilities(), new_properties.capabilities);
572        assert_eq!(leaf.ungreased_extensions(), new_properties.extensions);
573    }
574
575    #[maybe_async::test(not(mls_build_async), async(mls_build_async, crate::futures_test))]
576    async fn test_node_commit_no_meta_changes() {
577        for cipher_suite in TestCryptoProvider::all_supported_cipher_suites() {
578            let cipher_suite_provider = test_cipher_suite_provider(cipher_suite);
579
580            let (signing_identity, secret) = get_test_signing_identity(cipher_suite, b"foo").await;
581
582            let (mut leaf, leaf_secret) =
583                get_test_node(cipher_suite, signing_identity.clone(), &secret, None, None).await;
584
585            let original_leaf = leaf.clone();
586
587            let new_secret = leaf
588                .commit(
589                    &cipher_suite_provider,
590                    b"group",
591                    0,
592                    Some(default_properties()),
593                    None,
594                    &secret,
595                )
596                .await
597                .unwrap();
598
599            assert_ne!(new_secret, leaf_secret);
600            assert_ne!(original_leaf.public_key, leaf.public_key);
601
602            assert_eq!(
603                leaf.ungreased_capabilities(),
604                original_leaf.ungreased_capabilities()
605            );
606
607            assert_eq!(
608                leaf.ungreased_extensions(),
609                original_leaf.ungreased_extensions()
610            );
611
612            assert_eq!(leaf.signing_identity, original_leaf.signing_identity);
613
614            leaf.verify(
615                &cipher_suite_provider,
616                &signing_identity.signature_key,
617                &(b"group".as_slice(), 0).into(),
618            )
619            .await
620            .unwrap();
621        }
622    }
623
624    #[maybe_async::test(not(mls_build_async), async(mls_build_async, crate::futures_test))]
625    async fn test_node_commit_meta_changes() {
626        let cipher_suite = TEST_CIPHER_SUITE;
627
628        let (signing_identity, secret) = get_test_signing_identity(cipher_suite, b"foo").await;
629        let (mut leaf, _) =
630            get_test_node(cipher_suite, signing_identity, &secret, None, None).await;
631
632        let new_properties = ConfigProperties {
633            capabilities: get_test_capabilities(),
634            extensions: get_test_extensions(),
635        };
636
637        // The new identity has a fresh public key
638        let new_signing_identity = get_test_signing_identity(cipher_suite, b"foo").await.0;
639
640        leaf.commit(
641            &test_cipher_suite_provider(cipher_suite),
642            b"group",
643            0,
644            Some(new_properties.clone()),
645            Some(new_signing_identity.clone()),
646            &secret,
647        )
648        .await
649        .unwrap();
650
651        assert_eq!(leaf.capabilities, new_properties.capabilities);
652        assert_eq!(leaf.extensions, new_properties.extensions);
653        assert_eq!(leaf.signing_identity, new_signing_identity);
654    }
655
656    #[maybe_async::test(not(mls_build_async), async(mls_build_async, crate::futures_test))]
657    async fn context_is_signed() {
658        let provider = test_cipher_suite_provider(TEST_CIPHER_SUITE);
659
660        let (signing_identity, secret) = get_test_signing_identity(TEST_CIPHER_SUITE, b"foo").await;
661
662        let (mut leaf, _) = get_test_node(
663            TEST_CIPHER_SUITE,
664            signing_identity.clone(),
665            &secret,
666            None,
667            None,
668        )
669        .await;
670
671        leaf.sign(&provider, &secret, &(b"foo".as_slice(), 0).into())
672            .await
673            .unwrap();
674
675        let res = leaf
676            .verify(
677                &provider,
678                &signing_identity.signature_key,
679                &(b"foo".as_slice(), 1).into(),
680            )
681            .await;
682
683        assert_matches!(res, Err(MlsError::InvalidSignature));
684
685        let res = leaf
686            .verify(
687                &provider,
688                &signing_identity.signature_key,
689                &(b"bar".as_slice(), 0).into(),
690            )
691            .await;
692
693        assert_matches!(res, Err(MlsError::InvalidSignature));
694    }
695}