1use 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 #[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 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 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 let name = algo.name();
131 #[cfg(feature = "std")]
132 assert_eq!(name, algo.to_string());
133 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 let version = kem.version();
151 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 pub fn is_enabled(self) -> bool {
204 unsafe { ffi::OQS_KEM_alg_is_enabled(algorithm_to_id(self)) == 1 }
205 }
206
207 pub fn to_id(self) -> *const libc::c_char {
211 algorithm_to_id(self)
212 }
213
214 pub fn name(&self) -> &'static str {
218 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
231pub 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 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 pub fn algorithm(&self) -> Algorithm {
277 self.algorithm
278 }
279
280 pub fn version(&self) -> &'static str {
282 let kem = unsafe { self.kem.as_ref() };
283 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 pub fn claimed_nist_level(&self) -> u8 {
291 let kem = unsafe { self.kem.as_ref() };
292 kem.claimed_nist_level
293 }
294
295 pub fn is_ind_cca(&self) -> bool {
297 let kem = unsafe { self.kem.as_ref() };
298 kem.ind_cca
299 }
300
301 pub fn length_public_key(&self) -> usize {
303 let kem = unsafe { self.kem.as_ref() };
304 kem.length_public_key
305 }
306
307 pub fn length_secret_key(&self) -> usize {
309 let kem = unsafe { self.kem.as_ref() };
310 kem.length_secret_key
311 }
312
313 pub fn length_ciphertext(&self) -> usize {
315 let kem = unsafe { self.kem.as_ref() };
316 kem.length_ciphertext
317 }
318
319 pub fn length_shared_secret(&self) -> usize {
321 let kem = unsafe { self.kem.as_ref() };
322 kem.length_shared_secret
323 }
324
325 pub fn length_keypair_seed(&self) -> usize {
327 let kem = unsafe { self.kem.as_ref() };
328 kem.length_keypair_seed
329 }
330
331 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 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 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 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 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 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 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 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 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 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 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 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 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 let status = unsafe { func(ss.bytes.as_mut_ptr(), ct.bytes.as_ptr(), sk.bytes.as_ptr()) };
495 status_to_result(status)?;
496 unsafe { ss.bytes.set_len(kem.length_shared_secret) };
499 Ok(ss)
500 }
501}