Skip to main content

core_crypto/transaction_context/
key_package.rs

1//! This module contains all transactional behavior related to key packages
2
3use std::time::Duration;
4
5use core_crypto_keystore::{
6    entities::{StoredEncryptionKeyPair, StoredHpkePrivateKey, StoredKeyPackage},
7    traits::EntityDeleteBorrowed as _,
8};
9use openmls::prelude::{CryptoConfig, Lifetime};
10
11use super::{Error, Result, TransactionContext};
12use crate::{
13    ConversationConfiguration, CredentialRef, Keypackage, KeypackageRef, KeystoreError, RecursiveError,
14    mls::key_package::KeypackageExt as _,
15};
16
17/// Default lifetime of all generated Keypackages. Matches the limit defined in openmls
18pub const KEYPACKAGE_DEFAULT_LIFETIME: Duration = Duration::from_secs(60 * 60 * 24 * 28 * 3); // ~3 months
19
20impl TransactionContext {
21    /// Generate a [Keypackage] from the referenced credential.
22    ///
23    /// Makes no attempt to look up or prune existing keypackges.
24    ///
25    /// If `lifetime` is set, the keypackages will expire that span into the future.
26    /// If it is unset, [`KEYPACKAGE_DEFAULT_LIFETIME`]
27    /// is used.
28    ///
29    /// As a side effect, stores the keypackages and some related data in the keystore.
30    pub async fn generate_key_package(
31        &self,
32        credential_ref: &CredentialRef,
33        lifetime: Option<Duration>,
34    ) -> Result<Keypackage> {
35        let inner = self.inner().await?;
36        let lifetime = Lifetime::new(lifetime.unwrap_or(KEYPACKAGE_DEFAULT_LIFETIME).as_secs());
37        let credential = credential_ref
38            .load(&inner.transaction)
39            .await
40            .map_err(RecursiveError::mls_credential_ref("loading credential"))?;
41        let config = CryptoConfig {
42            ciphersuite: credential.cipher_suite.into(),
43            version: openmls::versions::ProtocolVersion::default(),
44        };
45
46        Keypackage::builder()
47            .leaf_node_capabilities(ConversationConfiguration::default_leaf_capabilities())
48            .key_package_lifetime(lifetime)
49            .build(
50                config,
51                &self.crypto_provider().await?,
52                &credential.signature_key_pair,
53                credential.to_mls_credential_with_key(),
54            )
55            .await
56            .map_err(Error::key_package_new())
57    }
58
59    /// Get all [`KeypackageRef`]s known to the keystore.
60    pub async fn get_key_package_refs(&self) -> Result<Vec<KeypackageRef>> {
61        let session = self.session().await?;
62        session
63            .get_keypackage_refs()
64            .await
65            .map_err(RecursiveError::mls_client(
66                "getting all key package refs for transaction",
67            ))
68            .map_err(Into::into)
69    }
70
71    /// Remove one [`Keypackage`] from the database.
72    ///
73    /// Succeeds silently if the keypackage does not exist in the database.
74    ///
75    /// Implementation note: this must first load and deserialize the keypackage,
76    /// then remove items from three distinct tables.
77    pub async fn remove_key_package(&self, kp_ref: &KeypackageRef) -> Result<()> {
78        let Some(kp) = self
79            .session()
80            .await?
81            .load_key_package(kp_ref)
82            .await
83            .map_err(RecursiveError::mls_client("loading key packages on session"))?
84        else {
85            return Ok(());
86        };
87
88        let inner = self.inner().await?;
89        let tx = inner.transaction();
90        StoredKeyPackage::delete_borrowed(tx, kp_ref.hash_ref())
91            .map_err(KeystoreError::wrap("removing key package from keystore"))?;
92        StoredHpkePrivateKey::delete_borrowed(tx, kp.hpke_init_key().as_slice())
93            .map_err(KeystoreError::wrap("removing private key from keystore"))?;
94        StoredEncryptionKeyPair::delete_borrowed(tx, kp.leaf_node().encryption_key().as_slice())
95            .map_err(KeystoreError::wrap("removing encryption keypair from keystore"))?;
96
97        Ok(())
98    }
99
100    /// Remove all keypackages associated with this credential.
101    ///
102    /// This is fairly expensive as it must first load all keypackages, then delete those matching the credential.
103    ///
104    /// Implementation note: once it makes it as far as having a list of keypackages, does _not_ short-circuit
105    /// if removing one returns an error. In that case, only the first produced error is returned.
106    /// This helps ensure that as many keypackages for the given credential ref are removed as possible.
107    pub async fn remove_key_packages_for(&self, credential_ref: &CredentialRef) -> Result<()> {
108        let inner = self.inner().await?;
109        let credential = credential_ref
110            .load(&inner.transaction)
111            .await
112            .map_err(RecursiveError::mls_credential_ref("loading credential"))?;
113        let signature_public_key = credential.signature_key_pair.public();
114
115        let mut first_err = None;
116        macro_rules! try_retain_err {
117            ($e:expr) => {
118                match $e {
119                    Err(err) => {
120                        if first_err.is_none() {
121                            first_err = Some(Error::from(err));
122                        }
123                        continue;
124                    }
125                    Ok(val) => val,
126                }
127            };
128        }
129
130        let session = self.session().await?;
131        for keypackage in session
132            .get_key_packages()
133            .await
134            .map_err(RecursiveError::mls_client("loading key packages"))?
135            .into_iter()
136            .filter(|keypackage| keypackage.leaf_node().signature_key().as_slice() == signature_public_key)
137        {
138            let kp_ref = try_retain_err!(keypackage.make_ref());
139            try_retain_err!(self.remove_key_package(&kp_ref).await);
140        }
141
142        match first_err {
143            None => Ok(()),
144            Some(err) => Err(err),
145        }
146    }
147}