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