Skip to main content

xmtp_mls_common/mls_ext/
payload_encryption.rs

1use openmls::ciphersuite::hpke::Error as OpenmlsHpkeError;
2use openmls::prelude::tls_codec::Error as TlsCodecError;
3use openmls_traits::crypto::OpenMlsCrypto;
4use openmls_traits::types::HpkeCiphertext;
5use thiserror::Error;
6use tls_codec::{Deserialize, Serialize};
7use xmtp_common::RetryableError;
8use xmtp_id::key_package::WrapperAlgorithm;
9
10static LIBCRUX_CRYPTO_PROVIDER: std::sync::LazyLock<openmls_libcrux_crypto::CryptoProvider> =
11    std::sync::LazyLock::new(|| {
12        openmls_libcrux_crypto::CryptoProvider::new().expect("Failed to create CryptoProvider")
13    });
14
15#[derive(Debug, Error)]
16pub enum WrapPayloadError {
17    #[error("OpenMLS HPKE error: {0}")]
18    Hpke(#[from] OpenmlsHpkeError),
19    #[error("TLS Codec error: {0}")]
20    TlsError(#[from] TlsCodecError),
21    #[error(transparent)]
22    Crypto(#[from] openmls_traits::types::CryptoError),
23}
24
25#[derive(Debug, Error)]
26pub enum UnwrapPayloadError {
27    #[error("OpenMLS HPKE error: {0}")]
28    Hpke(#[from] OpenmlsHpkeError),
29    #[error("TLS Codec error: {0}")]
30    TlsError(#[from] TlsCodecError),
31    #[error(transparent)]
32    Crypto(#[from] openmls_traits::types::CryptoError),
33}
34
35impl RetryableError for WrapPayloadError {
36    fn is_retryable(&self) -> bool {
37        false
38    }
39}
40
41impl RetryableError for UnwrapPayloadError {
42    fn is_retryable(&self) -> bool {
43        false
44    }
45}
46
47/// Wrap a payload (plus optional secondary payload) in an outer layer of HPKE
48/// encryption using the specified [WrapperAlgorithm]. The algorithm and public
49/// key type MUST match.
50///
51/// `label` is fed to the HPKE `EncryptContext` as the domain-separation label.
52/// Use [`xmtp_configuration::WELCOME_HPKE_LABEL`] for welcome-flow compatibility.
53///
54/// For the `XWingMLKEM768Draft6` algorithm, `payload` and `secondary_payload`
55/// are wrapped using the same HPKE public key. The first returned vec is the
56/// `HpkeCiphertext` with TLS serialization. The second vec is just ciphertext.
57pub fn wrap_payload_hpke(
58    payload: &[u8],
59    secondary_payload: &[u8],
60    hpke_public_key: &[u8],
61    wrapper_algorithm: WrapperAlgorithm,
62    label: &str,
63) -> Result<(Vec<u8>, Vec<u8>), WrapPayloadError> {
64    // The following implementation is the same as calling openmls_libcrux_crypto::CryptoProvider::hpke_seal(...)
65    // but uses the context to encrypt multiple messages at once using the same context
66    // because openmls only supports one shot messages.
67
68    let context = openmls::prelude::hpke::EncryptContext::from((label, [].as_slice()));
69    let info = context.tls_serialize_detached()?;
70    let aad = &[];
71
72    let map_hpke_error = |e| match e {
73        hpke_rs::HpkeError::InvalidConfig => openmls::prelude::CryptoError::SenderSetupError,
74        _ => openmls::prelude::CryptoError::HpkeEncryptionError,
75    };
76
77    let pk_r = hpke_rs::HpkePublicKey::new(hpke_public_key.to_vec());
78    let mut config = wrapper_algorithm.to_hpke_config();
79
80    let (enc, mut ctxt) = config
81        .setup_sender(&pk_r, &info, None, None, None)
82        .map_err(map_hpke_error)?;
83
84    let encrypted_payload = ctxt
85        .seal(aad, payload)
86        .map(|ct| HpkeCiphertext {
87            kem_output: enc.into(),
88            ciphertext: ct.into(),
89        })
90        .map_err(map_hpke_error)?;
91    let encrypted_secondary_payload = ctxt.seal(aad, secondary_payload).map_err(map_hpke_error)?;
92
93    Ok((
94        encrypted_payload.tls_serialize_detached()?,
95        encrypted_secondary_payload,
96    ))
97}
98
99/// Unwrap a payload that was wrapped using the specified [WrapperAlgorithm].
100/// The algorithm and private key type MUST match. `label` MUST match the value
101/// used at wrap time.
102pub fn unwrap_payload_hpke(
103    wrapped_payload: &[u8],
104    wrapped_secondary_payload: &[u8],
105    private_key: &[u8],
106    wrapper_algorithm: WrapperAlgorithm,
107    label: &str,
108) -> Result<(Vec<u8>, Vec<u8>), UnwrapPayloadError> {
109    let ciphertext = HpkeCiphertext::tls_deserialize_exact(wrapped_payload)?;
110
111    // The following implementation is the same as calling openmls_libcrux_crypto::CryptoProvider::hpke_open(...)
112    // but uses the context to decrypt multiple messages at once using the same context
113    // because openmls only supports one shot messages.
114
115    let context = openmls::prelude::hpke::EncryptContext::from((label, [].as_slice()));
116    let info = context.tls_serialize_detached()?;
117    let aad = &[];
118
119    let config = wrapper_algorithm.to_hpke_config();
120
121    let sk_r = hpke_rs::HpkePrivateKey::new(private_key.to_vec());
122
123    let map_hpke_error = |_| openmls::ciphersuite::hpke::Error::DecryptionFailed;
124
125    let mut ctxt = config
126        .setup_receiver(
127            ciphertext.kem_output.as_ref(),
128            &sk_r,
129            &info,
130            None,
131            None,
132            None,
133        )
134        .map_err(map_hpke_error)?;
135
136    let payload = ctxt
137        .open(aad, ciphertext.ciphertext.as_ref())
138        .map_err(map_hpke_error)?;
139    let secondary_payload = if wrapped_secondary_payload.is_empty() {
140        vec![]
141    } else {
142        ctxt.open(aad, wrapped_secondary_payload)
143            .map_err(map_hpke_error)?
144    };
145
146    Ok((payload, secondary_payload))
147}
148
149/// Wrap a payload with symmetric AEAD encryption (caller-supplied key + nonce).
150///
151/// Domain separation is handled by construction: callers MUST scope the
152/// symmetric key to a single use-case.
153pub fn wrap_payload_symmetric(
154    data: &[u8],
155    aead_type: openmls::prelude::AeadType,
156    symmetric_key: &[u8],
157    nonce: &[u8],
158) -> Result<Vec<u8>, WrapPayloadError> {
159    (*LIBCRUX_CRYPTO_PROVIDER)
160        .aead_encrypt(aead_type, symmetric_key, data, nonce, &[])
161        .map_err(Into::into)
162}
163
164/// Unwrap a payload that was wrapped with [`wrap_payload_symmetric`].
165pub fn unwrap_payload_symmetric(
166    data: &[u8],
167    aead_type: openmls::prelude::AeadType,
168    symmetric_key: &[u8],
169    nonce: &[u8],
170) -> Result<Vec<u8>, UnwrapPayloadError> {
171    (*LIBCRUX_CRYPTO_PROVIDER)
172        .aead_decrypt(aead_type, symmetric_key, data, nonce, &[])
173        .map_err(Into::into)
174}
175
176#[cfg(test)]
177mod tests {
178    use super::*;
179    use openmls_traits::{crypto::OpenMlsCrypto, random::OpenMlsRand};
180    use xmtp_configuration::{CIPHERSUITE, POST_QUANTUM_CIPHERSUITE, WELCOME_HPKE_LABEL};
181
182    const TEST_LABEL: &str = "test xmtp payload";
183
184    fn fresh_curve25519_keypair() -> (Vec<u8>, Vec<u8>) {
185        let crypto = openmls_rust_crypto::RustCrypto::default();
186        let ikm = crypto.random_vec(CIPHERSUITE.hash_length()).unwrap();
187        let kp = crypto
188            .derive_hpke_keypair(CIPHERSUITE.hpke_config(), &ikm)
189            .unwrap();
190        (kp.public, kp.private.to_vec())
191    }
192
193    fn fresh_xwing_keypair() -> (Vec<u8>, Vec<u8>) {
194        let crypto = openmls_libcrux_crypto::CryptoProvider::new().unwrap();
195        let ikm = crypto
196            .random_vec(POST_QUANTUM_CIPHERSUITE.hash_length())
197            .unwrap();
198        let kp = crypto
199            .derive_hpke_keypair(POST_QUANTUM_CIPHERSUITE.hpke_config(), &ikm)
200            .unwrap();
201        (kp.public, kp.private.to_vec())
202    }
203
204    #[xmtp_common::test]
205    fn round_trip_curve25519_hpke() {
206        let (pk, sk) = fresh_curve25519_keypair();
207
208        let payload = xmtp_common::rand_vec::<1000>();
209        let secondary = xmtp_common::rand_vec::<32>();
210
211        let wrapped = wrap_payload_hpke(
212            &payload,
213            &secondary,
214            &pk,
215            WrapperAlgorithm::Curve25519,
216            TEST_LABEL,
217        )
218        .unwrap();
219
220        assert_ne!(payload, wrapped.0);
221        assert_ne!(secondary, wrapped.1);
222
223        let unwrapped = unwrap_payload_hpke(
224            &wrapped.0,
225            &wrapped.1,
226            &sk,
227            WrapperAlgorithm::Curve25519,
228            TEST_LABEL,
229        )
230        .unwrap();
231
232        assert_eq!(unwrapped, (payload, secondary));
233    }
234
235    #[xmtp_common::test]
236    fn round_trip_xwing_hpke() {
237        let (pk, sk) = fresh_xwing_keypair();
238
239        let payload = xmtp_common::rand_vec::<1000>();
240        let secondary = xmtp_common::rand_vec::<32>();
241
242        let wrapped = wrap_payload_hpke(
243            &payload,
244            &secondary,
245            &pk,
246            WrapperAlgorithm::XWingMLKEM768Draft6,
247            TEST_LABEL,
248        )
249        .unwrap();
250
251        assert_ne!(payload, wrapped.0);
252        assert_ne!(secondary, wrapped.1);
253
254        let unwrapped = unwrap_payload_hpke(
255            &wrapped.0,
256            &wrapped.1,
257            &sk,
258            WrapperAlgorithm::XWingMLKEM768Draft6,
259            TEST_LABEL,
260        )
261        .unwrap();
262
263        assert_eq!(unwrapped, (payload.clone(), secondary));
264
265        // Empty secondary payload short-circuits to vec![].
266        let unwrapped = unwrap_payload_hpke(
267            &wrapped.0,
268            &[],
269            &sk,
270            WrapperAlgorithm::XWingMLKEM768Draft6,
271            TEST_LABEL,
272        )
273        .unwrap();
274
275        assert_eq!(unwrapped, (payload, vec![]));
276    }
277
278    #[xmtp_common::test]
279    fn wrong_key_fails_curve25519() {
280        let (pk, _sk) = fresh_curve25519_keypair();
281        let (_pk2, sk2) = fresh_curve25519_keypair();
282
283        let payload = xmtp_common::rand_vec::<128>();
284        let secondary = xmtp_common::rand_vec::<32>();
285
286        let wrapped = wrap_payload_hpke(
287            &payload,
288            &secondary,
289            &pk,
290            WrapperAlgorithm::Curve25519,
291            TEST_LABEL,
292        )
293        .unwrap();
294
295        unwrap_payload_hpke(
296            &wrapped.0,
297            &wrapped.1,
298            &sk2,
299            WrapperAlgorithm::Curve25519,
300            TEST_LABEL,
301        )
302        .unwrap_err();
303    }
304
305    #[xmtp_common::test]
306    fn wrong_label_fails_curve25519() {
307        let (pk, sk) = fresh_curve25519_keypair();
308
309        let payload = xmtp_common::rand_vec::<128>();
310        let secondary = xmtp_common::rand_vec::<32>();
311
312        let wrapped = wrap_payload_hpke(
313            &payload,
314            &secondary,
315            &pk,
316            WrapperAlgorithm::Curve25519,
317            TEST_LABEL,
318        )
319        .unwrap();
320
321        unwrap_payload_hpke(
322            &wrapped.0,
323            &wrapped.1,
324            &sk,
325            WrapperAlgorithm::Curve25519,
326            "different label",
327        )
328        .unwrap_err();
329    }
330
331    #[xmtp_common::test]
332    fn welcome_label_round_trip_matches_xmtp_configuration() {
333        // Sanity-check that the welcome label still round-trips correctly through
334        // the generalized API. This is the configuration the welcome flow uses.
335        let (pk, sk) = fresh_curve25519_keypair();
336
337        let payload = xmtp_common::rand_vec::<256>();
338        let secondary = xmtp_common::rand_vec::<32>();
339
340        let wrapped = wrap_payload_hpke(
341            &payload,
342            &secondary,
343            &pk,
344            WrapperAlgorithm::Curve25519,
345            WELCOME_HPKE_LABEL,
346        )
347        .unwrap();
348
349        let unwrapped = unwrap_payload_hpke(
350            &wrapped.0,
351            &wrapped.1,
352            &sk,
353            WrapperAlgorithm::Curve25519,
354            WELCOME_HPKE_LABEL,
355        )
356        .unwrap();
357
358        assert_eq!(unwrapped, (payload, secondary));
359    }
360
361    #[xmtp_common::test]
362    fn round_trip_symmetric() {
363        let symmetric_key = xmtp_common::rand_array::<32>();
364        let nonce = xmtp_common::rand_array::<12>();
365        let data = xmtp_common::rand_array::<1000>();
366
367        let wrapped = wrap_payload_symmetric(
368            &data,
369            openmls::prelude::AeadType::ChaCha20Poly1305,
370            &symmetric_key,
371            &nonce,
372        )
373        .unwrap();
374        let unwrapped = unwrap_payload_symmetric(
375            &wrapped,
376            openmls::prelude::AeadType::ChaCha20Poly1305,
377            &symmetric_key,
378            &nonce,
379        )
380        .unwrap();
381        assert_eq!(data.as_slice(), unwrapped.as_slice());
382    }
383
384    #[xmtp_common::test]
385    fn symmetric_wrong_key_fails() {
386        let symmetric_key = xmtp_common::rand_array::<32>();
387        let wrong_key = xmtp_common::rand_array::<32>();
388        let nonce = xmtp_common::rand_array::<12>();
389        let data = xmtp_common::rand_array::<1000>();
390
391        let wrapped = wrap_payload_symmetric(
392            &data,
393            openmls::prelude::AeadType::ChaCha20Poly1305,
394            &symmetric_key,
395            &nonce,
396        )
397        .unwrap();
398        unwrap_payload_symmetric(
399            &wrapped,
400            openmls::prelude::AeadType::ChaCha20Poly1305,
401            &wrong_key,
402            &nonce,
403        )
404        .unwrap_err();
405    }
406}