Skip to main content

oqs/
kem.rs

1//! KEM API
2//!
3//! See [`Kem`] for the main functionality.
4//! [`Algorithm`] lists the available algorithms.
5use alloc::vec::Vec;
6use core::ptr::NonNull;
7use core::str::FromStr;
8
9#[cfg(not(feature = "std"))]
10use cstr_core::CStr;
11#[cfg(feature = "std")]
12use std::ffi::CStr;
13
14#[cfg(feature = "serde")]
15use serde::{Deserialize, Serialize};
16
17use crate::ffi::kem as ffi;
18use crate::newtype_buffer;
19use crate::*;
20
21newtype_buffer!(PublicKey, PublicKeyRef);
22newtype_buffer!(SecretKey, SecretKeyRef);
23newtype_buffer!(Ciphertext, CiphertextRef);
24newtype_buffer!(SharedSecret, SharedSecretRef);
25newtype_buffer!(KeypairSeed, KeypairSeedRef);
26
27macro_rules! implement_kems {
28    { $(($feat: literal) $kem: ident: $oqs_id: ident),* $(,)? } => (
29
30        /// Supported algorithms by OQS
31        ///
32        /// Note that this doesn't mean that they'll be available.
33        ///
34        /// Optional support for `serde` if that feature is enabled.
35        #[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
36        #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
37        #[allow(missing_docs)]
38        pub enum Algorithm {
39            $(
40                $kem,
41            )*
42        }
43
44        fn algorithm_to_id(algorithm: Algorithm) -> *const libc::c_char {
45            let id: &[u8] = match algorithm {
46                $(
47                    Algorithm::$kem => &ffi::$oqs_id[..],
48                )*
49            };
50            id as *const _ as *const libc::c_char
51        }
52
53        impl FromStr for Algorithm {
54            type Err = crate::Error;
55
56            fn from_str(s: &str) -> Result<Self> {
57                $(
58                    if s == Algorithm::$kem.name() {
59                        return Ok(Algorithm::$kem);
60                    }
61                )*
62                Err(crate::Error::AlgorithmParsingError)
63            }
64        }
65
66        $(
67            #[cfg(test)]
68            #[allow(non_snake_case)]
69            mod $kem {
70                use super::*;
71
72                #[test]
73                #[cfg(feature = $feat)]
74                fn test_encaps_decaps() -> Result<()> {
75                    crate::init();
76
77                    let alg = Algorithm::$kem;
78                    let kem = Kem::new(alg)?;
79                    let (pk, sk) = kem.keypair()?;
80                    let (ct, ss1) = kem.encapsulate(&pk)?;
81                    let ss2 = kem.decapsulate(&sk, &ct)?;
82                    assert_eq!(ss1, ss2, "shared secret not equal!");
83                    Ok(())
84                }
85
86                #[test]
87                #[cfg(feature = $feat)]
88                fn test_encaps_decaps_derand() -> Result<()> {
89                    use crate::ffi::rand::OQS_randombytes;
90                    crate::init();
91
92                    let alg = Algorithm::$kem;
93                    let kem = Kem::new(alg)?;
94                    let mut seed = KeypairSeed {
95                        bytes: Vec::with_capacity(kem.length_keypair_seed()),
96                    };
97                    unsafe {
98                        // On some systems, getentropy fails if given a zero-length array
99                        if (kem.length_keypair_seed() > 0) {
100                            OQS_randombytes(seed.bytes.as_mut_ptr(), kem.length_keypair_seed());
101                        }
102                        seed.bytes.set_len(kem.length_keypair_seed());
103                    }
104                    let result = kem.keypair_derand(&seed);
105                    // expect Error::Error for KEMs with this API disabled
106                    if (kem.length_keypair_seed() == 0) {
107                        return result.map_or_else(|e| { match e { Error::Error => Ok(()), _ => Err(Error::Error) } }, |_| Err(Error::Error));
108                    }
109                    let (pk, sk) = result?;
110                    let (ct, ss1) = kem.encapsulate(&pk)?;
111                    let ss2 = kem.decapsulate(&sk, &ct)?;
112                    assert_eq!(ss1, ss2, "shared secret not equal!");
113                    Ok(())
114                }
115
116                #[test]
117                fn test_enabled() {
118                    crate::init();
119                    if cfg!(feature = $feat) {
120                        assert!(Algorithm::$kem.is_enabled());
121                    } else {
122                        assert!(!Algorithm::$kem.is_enabled())
123                    }
124                }
125
126                #[test]
127                fn test_name() {
128                    let algo = Algorithm::$kem;
129                    // Just make sure the name impl does not panic or crash.
130                    let name = algo.name();
131                    #[cfg(feature = "std")]
132                    assert_eq!(name, algo.to_string());
133                    // ... And actually contains something.
134                    assert!(!name.is_empty());
135                }
136
137                #[test]
138                fn test_get_algorithm_back() {
139                    let algorithm = Algorithm::$kem;
140                    if algorithm.is_enabled() {
141                        let kem = Kem::new(algorithm).unwrap();
142                        assert_eq!(algorithm, kem.algorithm());
143                    }
144                }
145
146                #[test]
147                fn test_version() {
148                    if let Ok(kem) = Kem::new(Algorithm::$kem) {
149                        // Just make sure the version can be called without panic
150                        let version = kem.version();
151                        // ... And actually contains something.
152                        assert!(!version.is_empty());
153                    }
154                }
155
156                #[test]
157                fn test_from_str() {
158                    let algorithm = Algorithm::$kem;
159                    let name = algorithm.name();
160                    let parsed = Algorithm::from_str(name).unwrap();
161                    assert_eq!(algorithm, parsed);
162                }
163            }
164        )*
165    )
166}
167
168implement_kems! {
169    ("bike") BikeL1: OQS_KEM_alg_bike_l1,
170    ("bike") BikeL3: OQS_KEM_alg_bike_l3,
171    ("bike") BikeL5: OQS_KEM_alg_bike_l5,
172    ("classic_mceliece") ClassicMcEliece348864: OQS_KEM_alg_classic_mceliece_348864,
173    ("classic_mceliece") ClassicMcEliece348864f: OQS_KEM_alg_classic_mceliece_348864f,
174    ("classic_mceliece") ClassicMcEliece460896: OQS_KEM_alg_classic_mceliece_460896,
175    ("classic_mceliece") ClassicMcEliece460896f: OQS_KEM_alg_classic_mceliece_460896f,
176    ("classic_mceliece") ClassicMcEliece6688128: OQS_KEM_alg_classic_mceliece_6688128,
177    ("classic_mceliece") ClassicMcEliece6688128f: OQS_KEM_alg_classic_mceliece_6688128f,
178    ("classic_mceliece") ClassicMcEliece6960119: OQS_KEM_alg_classic_mceliece_6960119,
179    ("classic_mceliece") ClassicMcEliece6960119f: OQS_KEM_alg_classic_mceliece_6960119f,
180    ("classic_mceliece") ClassicMcEliece8192128: OQS_KEM_alg_classic_mceliece_8192128,
181    ("classic_mceliece") ClassicMcEliece8192128f: OQS_KEM_alg_classic_mceliece_8192128f,
182    ("hqc") Hqc1: OQS_KEM_alg_hqc_1,
183    ("hqc") Hqc3: OQS_KEM_alg_hqc_3,
184    ("hqc") Hqc5: OQS_KEM_alg_hqc_5,
185    ("kyber") Kyber512: OQS_KEM_alg_kyber_512,
186    ("kyber") Kyber768: OQS_KEM_alg_kyber_768,
187    ("kyber") Kyber1024: OQS_KEM_alg_kyber_1024,
188    ("ml_kem") MlKem512: OQS_KEM_alg_ml_kem_512,
189    ("ml_kem") MlKem768: OQS_KEM_alg_ml_kem_768,
190    ("ml_kem") MlKem1024: OQS_KEM_alg_ml_kem_1024,
191    ("ntruprime") NtruPrimeSntrup761: OQS_KEM_alg_ntruprime_sntrup761,
192    ("frodokem") FrodoKem640Aes: OQS_KEM_alg_frodokem_640_aes,
193    ("frodokem") FrodoKem640Shake: OQS_KEM_alg_frodokem_640_shake,
194    ("frodokem") FrodoKem976Aes: OQS_KEM_alg_frodokem_976_aes,
195    ("frodokem") FrodoKem976Shake: OQS_KEM_alg_frodokem_976_shake,
196    ("frodokem") FrodoKem1344Aes: OQS_KEM_alg_frodokem_1344_aes,
197    ("frodokem") FrodoKem1344Shake: OQS_KEM_alg_frodokem_1344_shake,
198}
199
200impl Algorithm {
201    /// Returns true if this algorithm is enabled in the linked version
202    /// of liboqs
203    pub fn is_enabled(self) -> bool {
204        unsafe { ffi::OQS_KEM_alg_is_enabled(algorithm_to_id(self)) == 1 }
205    }
206
207    /// Provides a pointer to the id of the algorithm
208    ///
209    /// For use with the FFI api methods
210    pub fn to_id(self) -> *const libc::c_char {
211        algorithm_to_id(self)
212    }
213
214    /// Returns the algorithm's name as a static Rust string.
215    ///
216    /// This is the same as the `to_id`, but as a safe Rust string.
217    pub fn name(&self) -> &'static str {
218        // SAFETY: The id from ffi must be a proper null terminated C string
219        let id = unsafe { CStr::from_ptr(self.to_id()) };
220        id.to_str().expect("OQS algorithm names must be UTF-8")
221    }
222}
223
224#[cfg(feature = "std")]
225impl std::fmt::Display for Algorithm {
226    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
227        self.name().fmt(f)
228    }
229}
230
231/// KEM algorithm
232///
233/// # Example
234/// ```rust
235/// # if !cfg!(feature = "ml_kem") { return; }
236/// use oqs;
237/// oqs::init();
238/// let kem = oqs::kem::Kem::new(oqs::kem::Algorithm::MlKem512).unwrap();
239/// let (pk, sk) = kem.keypair().unwrap();
240/// let (ct, ss) = kem.encapsulate(&pk).unwrap();
241/// let ss2 = kem.decapsulate(&sk, &ct).unwrap();
242/// assert_eq!(ss, ss2);
243/// ```
244pub struct Kem {
245    algorithm: Algorithm,
246    kem: NonNull<ffi::OQS_KEM>,
247}
248
249unsafe impl Sync for Kem {}
250unsafe impl Send for Kem {}
251
252impl Drop for Kem {
253    fn drop(&mut self) {
254        unsafe { ffi::OQS_KEM_free(self.kem.as_ptr()) };
255    }
256}
257
258impl core::convert::TryFrom<Algorithm> for Kem {
259    type Error = crate::Error;
260    fn try_from(alg: Algorithm) -> Result<Kem> {
261        Kem::new(alg)
262    }
263}
264
265impl Kem {
266    /// Construct a new algorithm
267    pub fn new(algorithm: Algorithm) -> Result<Self> {
268        let kem = unsafe { ffi::OQS_KEM_new(algorithm_to_id(algorithm)) };
269        NonNull::new(kem).map_or_else(
270            || Err(Error::AlgorithmDisabled),
271            |kem| Ok(Self { algorithm, kem }),
272        )
273    }
274
275    /// Get the algorithm used by this `Kem`
276    pub fn algorithm(&self) -> Algorithm {
277        self.algorithm
278    }
279
280    /// Get the version of the implementation
281    pub fn version(&self) -> &'static str {
282        let kem = unsafe { self.kem.as_ref() };
283        // SAFETY: The alg_version from ffi must be a proper null terminated C string
284        let cstr = unsafe { CStr::from_ptr(kem.alg_version) };
285        cstr.to_str()
286            .expect("Algorithm version strings must be UTF-8")
287    }
288
289    /// Get the claimed nist level
290    pub fn claimed_nist_level(&self) -> u8 {
291        let kem = unsafe { self.kem.as_ref() };
292        kem.claimed_nist_level
293    }
294
295    /// Is the algorithm ind_cca secure
296    pub fn is_ind_cca(&self) -> bool {
297        let kem = unsafe { self.kem.as_ref() };
298        kem.ind_cca
299    }
300
301    /// Get the length of the public key
302    pub fn length_public_key(&self) -> usize {
303        let kem = unsafe { self.kem.as_ref() };
304        kem.length_public_key
305    }
306
307    /// Get the length of the secret key
308    pub fn length_secret_key(&self) -> usize {
309        let kem = unsafe { self.kem.as_ref() };
310        kem.length_secret_key
311    }
312
313    /// Get the length of the ciphertext
314    pub fn length_ciphertext(&self) -> usize {
315        let kem = unsafe { self.kem.as_ref() };
316        kem.length_ciphertext
317    }
318
319    /// Get the length of a shared secret
320    pub fn length_shared_secret(&self) -> usize {
321        let kem = unsafe { self.kem.as_ref() };
322        kem.length_shared_secret
323    }
324
325    /// Get the length of a keypair seed
326    pub fn length_keypair_seed(&self) -> usize {
327        let kem = unsafe { self.kem.as_ref() };
328        kem.length_keypair_seed
329    }
330
331    /// Obtain a secret key objects from bytes
332    ///
333    /// Returns None if the secret key is not the correct length.
334    pub fn secret_key_from_bytes<'a>(&self, buf: &'a [u8]) -> Option<SecretKeyRef<'a>> {
335        if self.length_secret_key() != buf.len() {
336            None
337        } else {
338            Some(SecretKeyRef::new(buf))
339        }
340    }
341
342    /// Obtain a public key from bytes
343    ///
344    /// Returns None if the public key is not the correct length.
345    pub fn public_key_from_bytes<'a>(&self, buf: &'a [u8]) -> Option<PublicKeyRef<'a>> {
346        if self.length_public_key() != buf.len() {
347            None
348        } else {
349            Some(PublicKeyRef::new(buf))
350        }
351    }
352
353    /// Obtain a ciphertext from bytes
354    ///
355    /// Returns None if the ciphertext is not the correct length.
356    pub fn ciphertext_from_bytes<'a>(&self, buf: &'a [u8]) -> Option<CiphertextRef<'a>> {
357        if self.length_ciphertext() != buf.len() {
358            None
359        } else {
360            Some(CiphertextRef::new(buf))
361        }
362    }
363
364    /// Obtain a secret key from bytes
365    ///
366    /// Returns None if the shared secret is not the correct length.
367    pub fn shared_secret_from_bytes<'a>(&self, buf: &'a [u8]) -> Option<SharedSecretRef<'a>> {
368        if self.length_shared_secret() != buf.len() {
369            None
370        } else {
371            Some(SharedSecretRef::new(buf))
372        }
373    }
374
375    /// Obtain a keypair seed from bytes
376    ///
377    /// Returns None if the shared secret is not the correct length.
378    pub fn keypair_seed_from_bytes<'a>(&self, buf: &'a [u8]) -> Option<KeypairSeedRef<'a>> {
379        if self.length_keypair_seed() != buf.len() {
380            None
381        } else {
382            Some(KeypairSeedRef::new(buf))
383        }
384    }
385
386    /// Generate a new keypair
387    pub fn keypair(&self) -> Result<(PublicKey, SecretKey)> {
388        let kem = unsafe { self.kem.as_ref() };
389        let func = kem.keypair.unwrap();
390        let mut pk = PublicKey {
391            bytes: Vec::with_capacity(kem.length_public_key),
392        };
393        let mut sk = SecretKey {
394            bytes: Vec::with_capacity(kem.length_secret_key),
395        };
396        let status = unsafe { func(pk.bytes.as_mut_ptr(), sk.bytes.as_mut_ptr()) };
397        status_to_result(status)?;
398        // update the lengths of the vecs
399        // this is safe to do, as we have initialised them now.
400        unsafe {
401            pk.bytes.set_len(kem.length_public_key);
402            sk.bytes.set_len(kem.length_secret_key);
403        }
404        Ok((pk, sk))
405    }
406
407    /// Generate a new keypair from a seed
408    pub fn keypair_derand<'a, S: Into<KeypairSeedRef<'a>>>(
409        &self,
410        seed: S,
411    ) -> Result<(PublicKey, SecretKey)> {
412        let seed = seed.into();
413        if seed.bytes.len() != self.length_keypair_seed() {
414            return Err(Error::InvalidLength);
415        }
416        let kem = unsafe { self.kem.as_ref() };
417        let func = kem.keypair_derand.unwrap();
418        let mut pk = PublicKey {
419            bytes: Vec::with_capacity(kem.length_public_key),
420        };
421        let mut sk = SecretKey {
422            bytes: Vec::with_capacity(kem.length_secret_key),
423        };
424        let status = unsafe {
425            func(
426                pk.bytes.as_mut_ptr(),
427                sk.bytes.as_mut_ptr(),
428                seed.bytes.as_ptr(),
429            )
430        };
431        status_to_result(status)?;
432        // update the lengths of the vecs
433        // this is safe to do, as we have initialised them now.
434        unsafe {
435            pk.bytes.set_len(kem.length_public_key);
436            sk.bytes.set_len(kem.length_secret_key);
437        }
438        Ok((pk, sk))
439    }
440
441    /// Encapsulate to the provided public key
442    pub fn encapsulate<'a, P: Into<PublicKeyRef<'a>>>(
443        &self,
444        pk: P,
445    ) -> Result<(Ciphertext, SharedSecret)> {
446        let pk = pk.into();
447        if pk.bytes.len() != self.length_public_key() {
448            return Err(Error::InvalidLength);
449        }
450        let kem = unsafe { self.kem.as_ref() };
451        let func = kem.encaps.unwrap();
452        let mut ct = Ciphertext {
453            bytes: Vec::with_capacity(kem.length_ciphertext),
454        };
455        let mut ss = SharedSecret {
456            bytes: Vec::with_capacity(kem.length_shared_secret),
457        };
458        // call encapsulate
459        let status = unsafe {
460            func(
461                ct.bytes.as_mut_ptr(),
462                ss.bytes.as_mut_ptr(),
463                pk.bytes.as_ptr(),
464            )
465        };
466        status_to_result(status)?;
467        // update the lengths of the vecs
468        // this is safe to do, as we have initialised them now.
469        unsafe {
470            ct.bytes.set_len(kem.length_ciphertext);
471            ss.bytes.set_len(kem.length_shared_secret);
472        }
473        Ok((ct, ss))
474    }
475
476    /// Decapsulate the provided ciphertext
477    pub fn decapsulate<'a, 'b, S: Into<SecretKeyRef<'a>>, C: Into<CiphertextRef<'b>>>(
478        &self,
479        sk: S,
480        ct: C,
481    ) -> Result<SharedSecret> {
482        let kem = unsafe { self.kem.as_ref() };
483        let sk = sk.into();
484        let ct = ct.into();
485        if sk.bytes.len() != self.length_secret_key() || ct.bytes.len() != self.length_ciphertext()
486        {
487            return Err(Error::InvalidLength);
488        }
489        let mut ss = SharedSecret {
490            bytes: Vec::with_capacity(kem.length_shared_secret),
491        };
492        let func = kem.decaps.unwrap();
493        // Call decapsulate
494        let status = unsafe { func(ss.bytes.as_mut_ptr(), ct.bytes.as_ptr(), sk.bytes.as_ptr()) };
495        status_to_result(status)?;
496        // update the lengths of the vecs
497        // this is safe to do, as we have initialised them now.
498        unsafe { ss.bytes.set_len(kem.length_shared_secret) };
499        Ok(ss)
500    }
501}