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
47pub 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 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
99pub 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 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
149pub 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
164pub 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 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 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}