diff --git a/.tangled/workflows/test.yml b/.tangled/workflows/test.yml index 7913fec..d966f88 100644 --- a/.tangled/workflows/test.yml +++ b/.tangled/workflows/test.yml @@ -15,6 +15,6 @@ steps: - name: Format check command: cargo fmt --all --check - name: Clippy - command: cargo clippy --locked --no-deps + command: cargo clippy --workspace --locked --no-deps - name: Tests - command: cargo test --locked --no-fail-fast + command: cargo test --workspace --locked --no-fail-fast diff --git a/src/handshake.rs b/src/handshake.rs index a7e9b8f..58f9242 100644 --- a/src/handshake.rs +++ b/src/handshake.rs @@ -1,52 +1,54 @@ +use core::marker::PhantomData; + use aead::Buffer; -use hybrid_array::typenum::Unsigned; +use hybrid_array::{Array, typenum::Unsigned}; +use ml_kem::{Encapsulate, Kem, KeyExport, ParameterSet, TryKeyInit, kem::Decapsulate}; use rand_core::CryptoRng; -use wharrgarbl_strobe::{StrobeRole, StrobeSecurity, StrobeState}; +use wharrgarbl_strobe::{StrobeRole, StrobeState, traits::StrobeSecurity}; use crate::{ WHARRGHARBL_PROTO, - kem::{KemEncap, KemSecurity}, transport::{AeadState, AeadStrobe}, }; -pub struct ClientHandshake { - kem_sec: KemSecurity, - sec_param: StrobeSecurity, - strobe: StrobeState, - decap: Option, +pub struct ClientHandshake { + kem_sec: PhantomData, + strobe: StrobeState, + decap: Option, } -impl ClientHandshake { - pub fn new(kem_sec: KemSecurity, sec_param: StrobeSecurity, psk: Option<&[u8; 32]>) -> Self { - let mut strobe = - StrobeState::new(WHARRGHARBL_PROTO.as_bytes(), sec_param, StrobeRole::Sender); +impl ClientHandshake +where + K::DecapsulationKey: Decapsulate, +{ + pub fn new(psk: Option<&[u8; 32]>) -> Self { + let mut strobe = StrobeState::::new(WHARRGHARBL_PROTO.as_bytes(), StrobeRole::Sender); if let Some(psk) = psk { strobe.key(psk); } - strobe.meta_ad(&kem_sec.to_bytes()); - strobe.meta_ad(&sec_param.to_bytes()); + strobe.meta_ad(&K::K::to_u16().to_le_bytes()); + strobe.meta_ad(&S::to_bytes()); Self { - kem_sec, - sec_param, + kem_sec: PhantomData, strobe, decap: None, } } pub fn send(&mut self, rng: &mut impl CryptoRng, buf: &mut dyn Buffer) -> aead::Result<()> { - let mut tag: aead::Tag = Default::default(); - let (decap, encap) = self.kem_sec.generate_with_rng(rng); + let mut tag: aead::Tag> = Default::default(); + let (decap, encap) = K::generate_keypair_from_rng(rng); - let written = encap.serialize(buf.as_mut())?; + buf.extend_from_slice(&encap.to_bytes())?; - buf.truncate(written); + let rachet_bytes = S::to_usize() >> 3; self.strobe.send_clr(buf.as_ref()); self.strobe.send_mac(&mut tag); - self.strobe.ratchet(self.sec_param.rachet_bytes()); + self.strobe.ratchet(rachet_bytes); buf.extend_from_slice(&tag)?; @@ -55,66 +57,63 @@ impl ClientHandshake { Ok(()) } - pub fn receive(&mut self, ciphertext: &[u8]) -> aead::Result { - let decap = self.decap.as_mut().ok_or(aead::Error)?; + pub fn receive(&mut self, ciphertext: &[u8]) -> aead::Result> { + let decap = self.decap.as_ref().ok_or(aead::Error)?; let tag = ciphertext .len() - .checked_sub(::TagSize::to_usize()) + .checked_sub( as aead::AeadCore>::TagSize::to_usize()) .ok_or(aead::Error)?; let (ciphertext, tag) = ciphertext.split_at(tag); - let tag: aead::Tag = tag.try_into().unwrap(); + let tag: aead::Tag> = tag.try_into().unwrap(); + let rachet_bytes = S::to_usize() >> 3; self.strobe.recv_clr(ciphertext); self.strobe.recv_mac(&tag)?; - self.strobe.ratchet(self.sec_param.rachet_bytes()); + self.strobe.ratchet(rachet_bytes); - let shared = decap.decapsulate(ciphertext)?; + let shared = decap + .decapsulate_slice(ciphertext) + .map_err(|_| aead::Error)?; Ok(self.finish(shared)) } - fn finish(&mut self, shared: ml_kem::SharedKey) -> AeadState { + fn finish(&mut self, shared: Array) -> AeadState { self.strobe.ad(&shared); - let mut key: aead::Key = Default::default(); - let mut inbound: aead::Nonce = Default::default(); - let mut outbound: aead::Nonce = Default::default(); + let mut key: aead::Key> = Default::default(); + let mut inbound: aead::Nonce> = Default::default(); + let mut outbound: aead::Nonce> = Default::default(); self.strobe.prf(&mut key); self.strobe.prf(&mut inbound); self.strobe.prf(&mut outbound); - AeadState::new(key, self.sec_param, outbound, inbound, StrobeRole::Sender) + AeadState::new(key, outbound, inbound, StrobeRole::Sender) } } -pub struct ServerHandshake { - kem_sec: KemSecurity, - sec_param: StrobeSecurity, - strobe: StrobeState, +pub struct ServerHandshake { + kem_sec: PhantomData, + strobe: StrobeState, } -impl ServerHandshake { - pub fn new(kem_sec: KemSecurity, sec_param: StrobeSecurity, psk: Option<&[u8; 32]>) -> Self { - let mut strobe = StrobeState::new( - WHARRGHARBL_PROTO.as_bytes(), - sec_param, - StrobeRole::Receiver, - ); +impl ServerHandshake { + pub fn new(psk: Option<&[u8; 32]>) -> Self { + let mut strobe = StrobeState::::new(WHARRGHARBL_PROTO.as_bytes(), StrobeRole::Receiver); if let Some(psk) = psk { strobe.key(psk); } - strobe.meta_ad(&kem_sec.to_bytes()); - strobe.meta_ad(&sec_param.to_bytes()); + strobe.meta_ad(&K::K::to_u16().to_le_bytes()); + strobe.meta_ad(&S::to_bytes()); Self { - kem_sec, - sec_param, + kem_sec: PhantomData, strobe, } } @@ -123,31 +122,34 @@ impl ServerHandshake { &mut self, rng: &mut impl CryptoRng, buf: &mut dyn Buffer, - ) -> aead::Result { + ) -> aead::Result> { let slice = buf.as_ref(); let tag = slice .len() - .checked_sub(::TagSize::to_usize()) + .checked_sub( as aead::AeadCore>::TagSize::to_usize()) .ok_or(aead::Error)?; let (encap, tag) = buf.as_ref().split_at(tag); - let tag: aead::Tag = tag.try_into().unwrap(); + let tag: aead::Tag> = tag.try_into().unwrap(); + let rachet_bytes = S::to_usize() >> 3; self.strobe.recv_clr(encap); self.strobe.recv_mac(&tag)?; - self.strobe.ratchet(self.sec_param.rachet_bytes()); + self.strobe.ratchet(rachet_bytes); + + let encap = K::EncapsulationKey::new_from_slice(encap).map_err(|_| aead::Error)?; - let (cipher, shared) = KemEncap::encapsulate_from_slice(self.kem_sec, encap, rng)?; + let (cipher, shared) = encap.encapsulate_with_rng(rng); - let mut tag: aead::Tag = Default::default(); + let mut tag: aead::Tag> = Default::default(); buf.truncate(0); self.strobe.send_clr(cipher.as_ref()); self.strobe.send_mac(&mut tag); - self.strobe.ratchet(self.sec_param.rachet_bytes()); + self.strobe.ratchet(rachet_bytes); buf.extend_from_slice(cipher.as_ref())?; buf.extend_from_slice(&tag)?; @@ -155,23 +157,26 @@ impl ServerHandshake { Ok(self.finish(shared)) } - fn finish(&mut self, shared: ml_kem::SharedKey) -> AeadState { + fn finish(&mut self, shared: Array) -> AeadState { self.strobe.ad(&shared); - let mut key: aead::Key = Default::default(); - let mut inbound: aead::Nonce = Default::default(); - let mut outbound: aead::Nonce = Default::default(); + let mut key: aead::Key> = Default::default(); + let mut inbound: aead::Nonce> = Default::default(); + let mut outbound: aead::Nonce> = Default::default(); self.strobe.prf(&mut key); self.strobe.prf(&mut inbound); self.strobe.prf(&mut outbound); - AeadState::new(key, self.sec_param, outbound, inbound, StrobeRole::Receiver) + AeadState::new(key, outbound, inbound, StrobeRole::Receiver) } } #[cfg(test)] mod tests { + use ml_kem::{MlKem512, MlKem768}; + use wharrgarbl_strobe::{Sec128, Sec256}; + use crate::utils::BufferSlice; use super::*; @@ -183,8 +188,8 @@ mod tests { 132, 45, 174, 183, 65, 89, 73, 107, 177, 77, 90, 164, 251, ]; - let mut alice = ClientHandshake::new(KemSecurity::Level1, StrobeSecurity::B128, Some(&psk)); - let mut bob = ServerHandshake::new(KemSecurity::Level1, StrobeSecurity::B128, Some(&psk)); + let mut alice = ClientHandshake::::new(Some(&psk)); + let mut bob = ServerHandshake::::new(Some(&psk)); // BufferSlice acts as our transport across the webz let mut buf = alloc::vec![0u8; 2048]; @@ -192,13 +197,13 @@ mod tests { let mut rng = rand_core::UnwrapErr(getrandom::SysRng); - alice.send(&mut rng, &mut buf)?; + alice.send(&mut rng, &mut buf).unwrap(); // Pretend to send ek across the webz: client -> server - let bob = bob.respond(&mut rng, &mut buf)?; + let bob = bob.respond(&mut rng, &mut buf).unwrap(); // Pretend to send ciphertext across the webz: server -> client - let alice = alice.receive(buf.as_ref())?; + let alice = alice.receive(buf.as_ref()).unwrap(); assert_eq!(alice.aead.key, bob.aead.key); @@ -220,8 +225,8 @@ mod tests { 132, 45, 174, 183, 65, 89, 73, 107, 177, 77, 90, 164, 251, ]; - let mut alice = ClientHandshake::new(KemSecurity::Level1, StrobeSecurity::B128, Some(&psk)); - let mut bob = ServerHandshake::new(KemSecurity::Level1, StrobeSecurity::B256, Some(&psk)); + let mut alice = ClientHandshake::::new(Some(&psk)); + let mut bob = ServerHandshake::::new(Some(&psk)); // BufferSlice acts as our transport across the webz let mut buf = alloc::vec![0u8; 2048]; @@ -243,8 +248,8 @@ mod tests { 132, 45, 174, 183, 65, 89, 73, 107, 177, 77, 90, 164, 251, ]; - let mut alice = ClientHandshake::new(KemSecurity::Level1, StrobeSecurity::B128, Some(&psk)); - let mut bob = ServerHandshake::new(KemSecurity::Level3, StrobeSecurity::B128, Some(&psk)); + let mut alice = ClientHandshake::::new(Some(&psk)); + let mut bob = ServerHandshake::::new(Some(&psk)); // BufferSlice acts as our transport across the webz let mut buf = alloc::vec![0u8; 2048]; @@ -266,8 +271,8 @@ mod tests { 132, 45, 174, 183, 65, 89, 73, 107, 177, 77, 90, 164, 251, ]; - let mut alice = ClientHandshake::new(KemSecurity::Level1, StrobeSecurity::B128, None); - let mut bob = ServerHandshake::new(KemSecurity::Level1, StrobeSecurity::B128, Some(&psk)); + let mut alice = ClientHandshake::::new(None); + let mut bob = ServerHandshake::::new(Some(&psk)); // BufferSlice acts as our transport across the webz let mut buf = alloc::vec![0u8; 2048]; diff --git a/src/kem.rs b/src/kem.rs deleted file mode 100644 index c0097c8..0000000 --- a/src/kem.rs +++ /dev/null @@ -1,137 +0,0 @@ -use alloc::boxed::Box; -use ml_kem::{Decapsulate, Encapsulate, Kem, KeyExport, SharedKey, TryKeyInit}; -use rand_core::CryptoRng; -use zeroize::Zeroize; - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -#[repr(u8)] -pub enum KemSecurity { - Level1 = 0, - Level3 = 1, -} - -pub enum KemEncap { - Level1(Box), - Level3(Box), -} - -pub struct KemCipher(Box<[u8]>); - -pub enum KemDecap { - Level1(Box), - Level3(Box), -} - -impl KemSecurity { - pub fn generate_with_rng(&self, rng: &mut impl CryptoRng) -> (KemDecap, KemEncap) { - match self { - Self::Level1 => { - let (decap, encap) = ml_kem::MlKem512::generate_keypair_from_rng(rng); - - ( - KemDecap::Level1(Box::new(decap)), - KemEncap::Level1(Box::new(encap)), - ) - } - Self::Level3 => { - let (decap, encap) = ml_kem::MlKem768::generate_keypair_from_rng(rng); - - ( - KemDecap::Level3(Box::new(decap)), - KemEncap::Level3(Box::new(encap)), - ) - } - } - } - - pub const fn to_bytes(self) -> [u8; 1] { - (self as u8).to_le_bytes() - } -} - -impl KemEncap { - pub fn encapsulate_from_slice( - sec: KemSecurity, - buf: &[u8], - rng: &mut impl CryptoRng, - ) -> Result<(KemCipher, ml_kem::SharedKey), aead::Error> { - match sec { - KemSecurity::Level1 => { - let encap = - ml_kem::EncapsulationKey512::new_from_slice(buf).map_err(|_| aead::Error)?; - - let (ct, shared) = encap.encapsulate_with_rng(rng); - - Ok((KemCipher(ct.into()), shared)) - } - KemSecurity::Level3 => { - let encap = - ml_kem::EncapsulationKey768::new_from_slice(buf).map_err(|_| aead::Error)?; - - let (ct, shared) = encap.encapsulate_with_rng(rng); - - Ok((KemCipher(ct.into()), shared)) - } - } - } - - pub fn serialize(&self, buf: &mut [u8]) -> Result { - match self { - Self::Level1(encap) => { - let data = encap.to_bytes(); - let (slice, _) = buf.split_at_mut_checked(data.len()).ok_or(aead::Error)?; - - slice.copy_from_slice(&data); - - Ok(data.len()) - } - Self::Level3(encap) => { - let data = encap.to_bytes(); - let (slice, _) = buf.split_at_mut_checked(data.len()).ok_or(aead::Error)?; - - slice.copy_from_slice(&data); - - Ok(data.len()) - } - } - } -} - -impl AsRef<[u8]> for KemCipher { - fn as_ref(&self) -> &[u8] { - self.0.as_ref() - } -} - -impl KemCipher { - pub const fn len(&self) -> usize { - self.0.len() - } - - pub const fn is_empty(&self) -> bool { - self.0.is_empty() - } -} - -impl Zeroize for KemCipher { - fn zeroize(&mut self) { - self.0.zeroize(); - } -} - -impl Drop for KemCipher { - fn drop(&mut self) { - self.zeroize(); - } -} - -impl KemDecap { - pub fn decapsulate(&mut self, ciphertext: &[u8]) -> Result { - let key = match self { - Self::Level1(decap) => decap.decapsulate_slice(ciphertext), - Self::Level3(decap) => decap.decapsulate_slice(ciphertext), - }; - - key.map_err(|_| aead::Error) - } -} diff --git a/src/lib.rs b/src/lib.rs index 7b9f08f..e898173 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,8 +1,10 @@ #![no_std] #![forbid(unsafe_code)] +use ml_kem::{MlKem512, MlKem768}; +use wharrgarbl_strobe::{Sec128, Sec256}; + pub mod handshake; -mod kem; pub mod transport; extern crate alloc; @@ -13,3 +15,8 @@ pub static WHARRGHARBL_PROTO: &str = "WGBL-v0.0-STv1.0.2"; pub mod utils { pub use wharrgarbl_utils::BufferSlice; } + +pub type ClientHandshake128L1 = handshake::ClientHandshake; +pub type ServerHandshake128L1 = handshake::ServerHandshake; +pub type ClientHandshake256L3 = handshake::ClientHandshake; +pub type ServerHandshake256L3 = handshake::ServerHandshake; diff --git a/src/transport.rs b/src/transport.rs index a8ef2c2..5c43cb0 100644 --- a/src/transport.rs +++ b/src/transport.rs @@ -1,28 +1,30 @@ +use core::marker::PhantomData; + use aead::{ AeadInOut, Buffer, Key, TagPosition, consts::{U16, U32}, }; use ctutils::{CtEq, CtSelect}; -use wharrgarbl_strobe::{StrobeRole, StrobeSecurity, StrobeState}; +use wharrgarbl_strobe::{StrobeRole, StrobeState, traits::StrobeSecurity}; use crate::WHARRGHARBL_PROTO; -pub struct AeadStrobe { +pub struct AeadStrobe { pub(crate) key: Key, - param: StrobeSecurity, + param: PhantomData, } -impl aead::AeadCore for AeadStrobe { +impl aead::AeadCore for AeadStrobe { type NonceSize = U16; type TagSize = U16; const TAG_POSITION: TagPosition = TagPosition::Postfix; } -impl aead::KeySizeUser for AeadStrobe { +impl aead::KeySizeUser for AeadStrobe { type KeySize = U32; } -impl aead::AeadInOut for AeadStrobe { +impl aead::AeadInOut for AeadStrobe { fn encrypt_inout_detached( &self, nonce: &aead::Nonce, @@ -30,11 +32,10 @@ impl aead::AeadInOut for AeadStrobe { mut buffer: aead::inout::InOutBuf<'_, '_, u8>, ) -> aead::Result> { let mut tag: aead::Tag = Default::default(); - let mut strobe = - StrobeState::new(WHARRGHARBL_PROTO.as_bytes(), self.param, StrobeRole::Sender); + let mut strobe = StrobeState::::new(WHARRGHARBL_PROTO.as_bytes(), StrobeRole::Sender); strobe.key(&self.key); - strobe.meta_ad(&self.param.to_bytes()); + strobe.meta_ad(&S::to_bytes()); strobe.meta_ad(nonce); strobe.ad(associated_data); strobe.send_enc(buffer.get_out()); @@ -51,14 +52,10 @@ impl aead::AeadInOut for AeadStrobe { mut buffer: aead::inout::InOutBuf<'_, '_, u8>, tag: &aead::Tag, ) -> aead::Result<()> { - let mut strobe = StrobeState::new( - WHARRGHARBL_PROTO.as_bytes(), - self.param, - StrobeRole::Receiver, - ); + let mut strobe = StrobeState::::new(WHARRGHARBL_PROTO.as_bytes(), StrobeRole::Receiver); strobe.key(&self.key); - strobe.meta_ad(&self.param.to_bytes()); + strobe.meta_ad(&S::to_bytes()); strobe.meta_ad(nonce); strobe.ad(associated_data); strobe.recv_enc(buffer.get_out()); @@ -67,20 +64,20 @@ impl aead::AeadInOut for AeadStrobe { } } -impl zeroize::Zeroize for AeadStrobe { +impl zeroize::Zeroize for AeadStrobe { fn zeroize(&mut self) { self.key.zeroize(); } } -pub struct AeadState { - pub(crate) aead: AeadStrobe, - pub(crate) epstein: aead::Nonce, - pub(crate) trump: aead::Nonce, +pub struct AeadState { + pub(crate) aead: AeadStrobe, + pub(crate) epstein: aead::Nonce>, + pub(crate) trump: aead::Nonce>, handshake_role: StrobeRole, } -impl zeroize::Zeroize for AeadState { +impl zeroize::Zeroize for AeadState { fn zeroize(&mut self) { self.aead.zeroize(); self.epstein.zeroize(); @@ -88,20 +85,19 @@ impl zeroize::Zeroize for AeadState { } } -impl zeroize::ZeroizeOnDrop for AeadState {} +impl zeroize::ZeroizeOnDrop for AeadState {} -impl Drop for AeadState { +impl Drop for AeadState { fn drop(&mut self) { zeroize::Zeroize::zeroize(self); } } -impl AeadState { +impl AeadState { pub(crate) fn new( - key: Key, - sec: StrobeSecurity, - outbound: aead::Nonce, - inbound: aead::Nonce, + key: Key>, + outbound: aead::Nonce>, + inbound: aead::Nonce>, role: StrobeRole, ) -> Self { assert_ne!( @@ -110,20 +106,23 @@ impl AeadState { ); Self { - aead: AeadStrobe { key, param: sec }, + aead: AeadStrobe { + key, + param: PhantomData, + }, epstein: outbound, trump: inbound, handshake_role: role, } } - fn select_nonce(&self, sending_role: StrobeRole) -> aead::Nonce { + fn select_nonce(&self, sending_role: StrobeRole) -> aead::Nonce> { let role_context = self.handshake_role ^ sending_role; self.epstein.ct_select(&self.trump, role_context.ct_eq(&1)) } - fn mix_nonce(&self, position: [u8; 8], sending_role: StrobeRole) -> aead::Nonce { + fn mix_nonce(&self, position: [u8; 8], sending_role: StrobeRole) -> aead::Nonce> { let mut nonce = self.select_nonce(sending_role); let mid = nonce.len() - position.len(); @@ -136,7 +135,7 @@ impl AeadState { nonce } - pub fn split(&self) -> (SendState<'_>, RecvState<'_>) { + pub fn split(&self) -> (SendState<'_, S>, RecvState<'_, S>) { ( SendState { transport: self, @@ -150,12 +149,12 @@ impl AeadState { } } -pub struct SendState<'a> { - transport: &'a AeadState, +pub struct SendState<'a, S: StrobeSecurity> { + transport: &'a AeadState, counter: u64, } -impl SendState<'_> { +impl SendState<'_, S> { pub fn encrypt(&mut self, buffer: &mut dyn Buffer, ad: &[u8]) -> aead::Result<()> { if self.counter.ct_eq(&u64::MAX).into() { return Err(aead::Error); @@ -175,12 +174,12 @@ impl SendState<'_> { } } -pub struct RecvState<'a> { - transport: &'a AeadState, +pub struct RecvState<'a, S: StrobeSecurity> { + transport: &'a AeadState, counter: u64, } -impl RecvState<'_> { +impl RecvState<'_, S> { pub fn decrypt(&mut self, buffer: &mut dyn Buffer, ad: &[u8]) -> aead::Result<()> { if self.counter.ct_eq(&u64::MAX).into() { return Err(aead::Error); @@ -203,6 +202,8 @@ impl RecvState<'_> { #[cfg(test)] mod tests { + use wharrgarbl_strobe::{Sec128, Sec256}; + use super::*; #[test] @@ -216,16 +217,14 @@ mod tests { let outbound = 123u128.to_ne_bytes(); let inbound = 234u128.to_ne_bytes(); - let alice = AeadState::new( + let alice = AeadState::::new( shared_secret.into(), - StrobeSecurity::B128, outbound.into(), inbound.into(), StrobeRole::Sender, ); - let bob = AeadState::new( + let bob = AeadState::::new( shared_secret.into(), - StrobeSecurity::B128, outbound.into(), inbound.into(), StrobeRole::Receiver, @@ -291,16 +290,14 @@ mod tests { let outbound = 123u128.to_ne_bytes(); let inbound = 234u128.to_ne_bytes(); - let alice = AeadState::new( + let alice = AeadState::::new( shared_secret.into(), - StrobeSecurity::B128, outbound.into(), inbound.into(), StrobeRole::Sender, ); - let bob = AeadState::new( + let bob = AeadState::::new( shared_secret.into(), - StrobeSecurity::B256, outbound.into(), inbound.into(), StrobeRole::Receiver, diff --git a/wharrgarbl-strobe/src/basic_kats.rs b/wharrgarbl-strobe/src/basic_kats.rs index 0e5409c..2579842 100644 --- a/wharrgarbl-strobe/src/basic_kats.rs +++ b/wharrgarbl-strobe/src/basic_kats.rs @@ -7,13 +7,13 @@ use aead::consts::{U16, U65}; use hybrid_array::Array; -use crate::{StrobeRole, StrobeSecurity, keccakf::KECCAK_BUFFER_SIZE, strobe::StrobeState}; +use crate::{Sec128, Sec256, StrobeRole, keccakf::KECCAK_BUFFER_SIZE, strobe::StrobeState}; extern crate std; #[test] fn test_init_128() { - let s = StrobeState::new(b"", StrobeSecurity::B128, StrobeRole::Sender); + let s = StrobeState::::new(b"", StrobeRole::Sender); let expected_st: [u8; KECCAK_BUFFER_SIZE] = [ 0x9c, 0x7f, 0x16, 0x8f, 0xf8, 0xfd, 0x55, 0xda, 0x2a, 0xa7, 0x3c, 0x23, 0x55, 0x65, 0x35, @@ -37,7 +37,7 @@ fn test_init_128() { #[test] fn test_init_256() { - let s = StrobeState::new(b"", StrobeSecurity::B256, StrobeRole::Sender); + let s = StrobeState::::new(b"", StrobeRole::Sender); let expected_st: [u8; KECCAK_BUFFER_SIZE] = [ 0x37, 0xc1, 0x15, 0x06, 0xed, 0x61, 0xe7, 0xda, 0x7c, 0x1a, 0x2f, 0x2c, 0x1f, 0x49, 0x74, @@ -62,7 +62,7 @@ fn test_init_256() { #[test] fn test_metadata() { // We will accumulate output over 3 operations and 3 meta-operations - let mut s = StrobeState::new(b"metadatatest", StrobeSecurity::B256, StrobeRole::Sender); + let mut s = StrobeState::::new(b"metadatatest", StrobeRole::Sender); let mut output = std::vec::Vec::new(); let buf = b"meta1"; @@ -116,7 +116,7 @@ fn test_metadata() { #[test] fn test_seq() { - let mut s = StrobeState::new(b"seqtest", StrobeSecurity::B256, StrobeRole::Sender); + let mut s = StrobeState::::new(b"seqtest", StrobeRole::Sender); let mut buf = [0u8; 10]; s.prf(&mut buf[..]); @@ -172,16 +172,8 @@ fn test_seq() { #[test] fn test_enc_correctness() { let orig_msg = b"Hello there"; - let mut tx = StrobeState::new( - b"enccorrectnesstest", - StrobeSecurity::B256, - StrobeRole::Sender, - ); - let mut rx = StrobeState::new( - b"enccorrectnesstest", - StrobeSecurity::B256, - StrobeRole::Receiver, - ); + let mut tx = StrobeState::::new(b"enccorrectnesstest", StrobeRole::Sender); + let mut rx = StrobeState::::new(b"enccorrectnesstest", StrobeRole::Receiver); tx.key(b"the-combination-on-my-luggage"); rx.key(b"the-combination-on-my-luggage"); @@ -196,8 +188,8 @@ fn test_enc_correctness() { #[test] fn test_mac_correctness_and_soundness() { - let mut tx = StrobeState::new(b"mactest", StrobeSecurity::B256, StrobeRole::Sender); - let mut rx = StrobeState::new(b"mactest", StrobeSecurity::B256, StrobeRole::Receiver); + let mut tx = StrobeState::::new(b"mactest", StrobeRole::Sender); + let mut rx = StrobeState::::new(b"mactest", StrobeRole::Receiver); // Just do some stuff with the state @@ -225,7 +217,7 @@ fn test_mac_correctness_and_soundness() { #[test] fn test_long_inputs() { - let mut s = StrobeState::new(b"bigtest", StrobeSecurity::B256, StrobeRole::Sender); + let mut s = StrobeState::::new(b"bigtest", StrobeRole::Sender); const BIG_N: usize = 9823; const SMALL_N: usize = 65; let big_data = [0x34u8; BIG_N]; @@ -282,7 +274,7 @@ fn test_long_inputs() { fn test_streaming_correctness() { // Compute a few things without breaking up their inputs let one_shot_st: std::vec::Vec = { - let mut s = StrobeState::new(b"streamingtest", StrobeSecurity::B256, StrobeRole::Receiver); + let mut s = StrobeState::::new(b"streamingtest", StrobeRole::Receiver); s.ad(b"mynonce"); @@ -298,7 +290,7 @@ fn test_streaming_correctness() { }; // Now do the same thing but stream the inputs let streamed_st: std::vec::Vec = { - let mut s = StrobeState::new(b"streamingtest", StrobeSecurity::B256, StrobeRole::Receiver); + let mut s = StrobeState::::new(b"streamingtest", StrobeRole::Receiver); s.ad(b"my"); s.ad(b"nonce"); diff --git a/wharrgarbl-strobe/src/herding_kats/harness.rs b/wharrgarbl-strobe/src/herding_kats/harness.rs index 99c6b2e..80c2492 100644 --- a/wharrgarbl-strobe/src/herding_kats/harness.rs +++ b/wharrgarbl-strobe/src/herding_kats/harness.rs @@ -2,18 +2,17 @@ extern crate std; use std::{string::String, vec::Vec}; -use aead::consts::U14; +use aead::consts::{U14, U128, U256}; use serde::{Deserialize, Deserializer, de}; -use crate::{StrobeRole, StrobeSecurity, strobe::StrobeState}; +use crate::{StrobeRole, strobe::StrobeState, traits::StrobeSecurity}; /// The harness we will put on our KATs so we can herd them and make them do tests. /// (This is the top-level structure of the JSON we find in the test vectors) #[derive(Deserialize)] struct KatHarness { proto_string: String, - #[serde(deserialize_with = "security_param_from_bits")] - security: StrobeSecurity, + security: u64, operations: Vec, } @@ -38,11 +37,43 @@ pub enum DataOrLength<'a> { Length(usize), } -// Given the name of the operation and its required parameters, run the STROBE operation. -fn run_kat_operation( +enum StrobeKinds { + U128(StrobeState), + U256(StrobeState), +} + +impl StrobeKinds { + fn new(protocol: &[u8], security: u64) -> Self { + match security { + 128 => Self::U128(StrobeState::new(protocol, StrobeRole::Sender)), + 256 => Self::U256(StrobeState::new(protocol, StrobeRole::Sender)), + _ => panic!("Invalid Security parameter"), + } + } + + fn get_state(&self) -> &[u8] { + match self { + Self::U128(s) => &s.state.0, + Self::U256(s) => &s.state.0, + } + } + + fn run_kat_operation(&mut self, op_name: &str, meta: bool, dol: DataOrLength, more: bool) { + match self { + Self::U128(s) => { + exec_kat(s, op_name, meta, dol, more); + } + Self::U256(s) => { + exec_kat(s, op_name, meta, dol, more); + } + } + } +} + +fn exec_kat( + s: &mut StrobeState, op_name: &str, meta: bool, - s: &mut StrobeState, dol: DataOrLength, more: bool, ) { @@ -114,7 +145,7 @@ pub fn test_against_kat>(filename: P) { operations, } = serde_json::from_reader(file).unwrap(); - let mut strobe = StrobeState::new(proto_string.as_bytes(), security, StrobeRole::Sender); + let mut strobe = StrobeKinds::new(proto_string.as_bytes(), security); operations.into_iter().for_each( |KatOperation { @@ -127,7 +158,7 @@ pub fn test_against_kat>(filename: P) { }| match name.as_str() { "init" => { // Check that the initial state matches what is expected in the KAT operation. - assert_eq!(&strobe.state.0[..], expected_state_after.as_slice()); + assert_eq!(strobe.get_state(), expected_state_after.as_slice()); } name => { // RATCHET inputs are given as strings of zeros instead of lengths. So just take the @@ -138,9 +169,9 @@ pub fn test_against_kat>(filename: P) { DataOrLength::Data(input_data.as_mut_slice()) }; - run_kat_operation(name, meta, &mut strobe, input, stream); + strobe.run_kat_operation(name, meta, input, stream); - assert_eq!(&strobe.state.0[..], expected_state_after.as_slice()); + assert_eq!(strobe.get_state(), expected_state_after.as_slice()); // Only test expected output if the test vector has output to test against. if let Some(expected) = expected_output { @@ -152,18 +183,18 @@ pub fn test_against_kat>(filename: P) { ); } -fn security_param_from_bits<'de, D: Deserializer<'de>>( - deserializer: D, -) -> Result { - match u64::deserialize(deserializer)? { - 128 => Ok(StrobeSecurity::B128), - 256 => Ok(StrobeSecurity::B256), - n => Err(de::Error::custom(std::format!( - "Invalid security parameter: {}", - n - ))), - } -} +// fn security_param_from_bits<'de, D: Deserializer<'de>>( +// deserializer: D, +// ) -> Result { +// match u64::deserialize(deserializer)? { +// 128 => Ok(StrobeSecurity::B128), +// 256 => Ok(StrobeSecurity::B256), +// n => Err(de::Error::custom(std::format!( +// "Invalid security parameter: {}", +// n +// ))), +// } +// } fn bytes_from_hex<'de, D>(deserializer: D) -> Result, D::Error> where diff --git a/wharrgarbl-strobe/src/lib.rs b/wharrgarbl-strobe/src/lib.rs index 2df7215..6108fca 100644 --- a/wharrgarbl-strobe/src/lib.rs +++ b/wharrgarbl-strobe/src/lib.rs @@ -7,30 +7,18 @@ mod strobe; mod basic_kats; #[cfg(test)] mod herding_kats; +pub mod traits; use core::ops::BitXor; +use aead::consts::{U128, U256}; pub use strobe::StrobeState; /// Version of Strobe that this crate implements. pub static STROBE_VERSION: &str = "1.0.2"; -#[derive(Debug, Clone, Copy)] -#[repr(u16)] -pub enum StrobeSecurity { - B128 = 128, - B256 = 256, -} - -impl StrobeSecurity { - pub const fn rachet_bytes(self) -> usize { - (self as usize) >> 3 - } - - pub const fn to_bytes(self) -> [u8; 2] { - (self as u16).to_le_bytes() - } -} +pub type Sec128 = U128; +pub type Sec256 = U256; #[derive(Debug, Clone, Copy, PartialEq, Eq)] #[repr(u8)] diff --git a/wharrgarbl-strobe/src/strobe.rs b/wharrgarbl-strobe/src/strobe.rs index ac64a4e..a0d315d 100644 --- a/wharrgarbl-strobe/src/strobe.rs +++ b/wharrgarbl-strobe/src/strobe.rs @@ -1,12 +1,15 @@ +use core::marker::PhantomData; + use ctutils::{Choice, CtAssign, CtEq, CtLt, CtSelect}; use hybrid_array::{Array, ArraySize}; use zeroize::Zeroize; use crate::{ - STROBE_VERSION, StrobeRole, StrobeSecurity, + STROBE_VERSION, StrobeRole, keccakf::{KECCAK_BUFFER_SIZE, KeccakF1600}, opflags::OpFlags, ops, + traits::StrobeSecurity, }; /// Private integer representations for Role, to allow for better constant time compat. @@ -16,11 +19,11 @@ mod role { } #[derive(Clone)] -pub struct StrobeState { +pub struct StrobeState { /// Internal Keccak state pub(crate) state: KeccakF1600, /// Security parameter (128 or 256 bits) - sec: StrobeSecurity, + sec: PhantomData, /// This is the `R` parameter in the Strobe spec rate: usize, /// Index into `state` @@ -66,7 +69,7 @@ macro_rules! define_non_mut_operations { }; } -impl Zeroize for StrobeState { +impl Zeroize for StrobeState { fn zeroize(&mut self) { self.state.zeroize(); self.rate.zeroize(); @@ -77,41 +80,37 @@ impl Zeroize for StrobeState { } } -impl zeroize::ZeroizeOnDrop for StrobeState {} +impl zeroize::ZeroizeOnDrop for StrobeState {} -impl Drop for StrobeState { +impl Drop for StrobeState { fn drop(&mut self) { self.zeroize(); } } -impl core::fmt::Display for StrobeState { +impl core::fmt::Display for StrobeState { fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result { f.write_str("Strobe-Keccak-")?; - match self.sec { - StrobeSecurity::B128 => f.write_str("128")?, - StrobeSecurity::B256 => f.write_str("256")?, - } + write!(f, "{}", S::to_usize())?; f.write_str("/1600-v")?; f.write_str(STROBE_VERSION) } } -impl core::fmt::Debug for StrobeState { +impl core::fmt::Debug for StrobeState { fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result { // Do not reveal internal state of StrobeState, other than its security level f.debug_struct("StrobeState") - .field("sec", &self.sec) + .field("sec", &S::to_usize()) .field("version", &STROBE_VERSION) .finish_non_exhaustive() } } -impl StrobeState { +impl StrobeState { /// Makes a new `StrobeState` object with a given protocol byte string and security parameter. - pub fn new(protocol: &[u8], sec: StrobeSecurity, role: StrobeRole) -> Self { - let rate = KECCAK_BUFFER_SIZE - (sec as usize) / 4 - 2; - assert!((1..254).contains(&rate)); + pub fn new(protocol: &[u8], role: StrobeRole) -> Self { + let rate = S::rate(); // Initialize state: st = F([0x01, R+2, 0x01, 0x00, 0x01, 0x60] + b"STROBEvX.Y.Z") let mut state_buffer = [0u8; KECCAK_BUFFER_SIZE]; @@ -125,7 +124,7 @@ impl StrobeState { let mut strobe = Self { state, - sec, + sec: PhantomData, rate, position: 0, start: 0, @@ -311,7 +310,7 @@ impl StrobeState { debug_assert!(flags != ops::KEY && flags.contains(OpFlags::CIPHER).to_bool()); static SPECIAL_CASES: [OpFlags; 3] = [ops::PRF, ops::SEND_MAC, ops::SEND_ENC]; - static OPS: [fn(&mut StrobeState, data: &mut [u8]); 4] = [ + let ops: [fn(&mut Self, data: &mut [u8]); 4] = [ StrobeState::squeeze, StrobeState::copy_state, StrobeState::absorb_and_set, @@ -329,7 +328,7 @@ impl StrobeState { }) & 0b11; - OPS[index](self, data); + ops[index](self, data); } /// Performs the state transformation that corresponds to the given flags. If `more` is given, @@ -350,12 +349,11 @@ impl StrobeState { // RATCHET is special cased to never call operate/operate_no_mutate directly debug_assert!(flags == ops::KEY || !flags.contains(OpFlags::CIPHER).to_bool()); - static OPS: [fn(&mut StrobeState, data: &[u8]); 2] = - [StrobeState::absorb, StrobeState::overwrite]; + let ops: [fn(&mut Self, data: &[u8]); 2] = [StrobeState::absorb, StrobeState::overwrite]; let index = (flags.ct_eq(&ops::KEY).to_u8() & 1) as usize; - OPS[index](self, data); + ops[index](self, data); } fn recv_mac_inner(&mut self, flags: OpFlags, mac_copy: &mut [u8]) -> Result<(), aead::Error> { @@ -451,18 +449,20 @@ impl StrobeState { #[cfg(test)] mod tests { + use aead::consts::U128; + use super::*; extern crate std; #[test] fn version_formatting() { - let s = StrobeState::new(b"", StrobeSecurity::B128, StrobeRole::Sender); + let s = StrobeState::::new(b"", StrobeRole::Sender); let display = std::format!("{s}"); let debug = std::format!("{s:?}"); assert_eq!(&display, "Strobe-Keccak-128/1600-v1.0.2"); - assert_eq!(&debug, "StrobeState { sec: B128, version: \"1.0.2\", .. }"); + assert_eq!(&debug, "StrobeState { sec: 128, version: \"1.0.2\", .. }"); } } diff --git a/wharrgarbl-strobe/src/traits.rs b/wharrgarbl-strobe/src/traits.rs new file mode 100644 index 0000000..716fe63 --- /dev/null +++ b/wharrgarbl-strobe/src/traits.rs @@ -0,0 +1,29 @@ +use aead::consts::{U128, U256}; +use hybrid_array::typenum::Unsigned; + +use crate::keccakf::KECCAK_BUFFER_SIZE; + +pub trait StrobeSecurity: Unsigned { + fn to_bytes() -> [u8; 2]; + fn rate() -> usize; +} + +impl StrobeSecurity for U128 { + fn to_bytes() -> [u8; 2] { + Self::to_u16().to_le_bytes() + } + + fn rate() -> usize { + KECCAK_BUFFER_SIZE - (Self::to_usize()) / 4 - 2 + } +} + +impl StrobeSecurity for U256 { + fn to_bytes() -> [u8; 2] { + Self::to_u16().to_le_bytes() + } + + fn rate() -> usize { + KECCAK_BUFFER_SIZE - (Self::to_usize()) / 4 - 2 + } +} diff --git a/wharrgarbl-utils/src/lib.rs b/wharrgarbl-utils/src/lib.rs index 36f1e01..962f969 100644 --- a/wharrgarbl-utils/src/lib.rs +++ b/wharrgarbl-utils/src/lib.rs @@ -8,13 +8,10 @@ pub struct BufferSlice<'slice> { impl<'slice> BufferSlice<'slice> { pub const fn new(buffer: &'slice mut [u8]) -> Self { - Self { - end: buffer.len(), - buffer, - } + Self { end: 0, buffer } } - pub const fn reset(&mut self) { + pub const fn fill(&mut self) { self.end = self.buffer.len(); } } @@ -47,7 +44,7 @@ impl aead::Buffer for BufferSlice<'_> { } fn truncate(&mut self, len: usize) { - self.end = len.min(self.buffer.len()); + self.end = len.min(self.end); } } @@ -59,12 +56,12 @@ mod tests { extern crate alloc; #[test] - fn defaults_to_full_buffer_size() { + fn defaults_to_empty_buffer_size() { let mut buf = alloc::vec![0u8; 128]; let buf_slice = BufferSlice::new(&mut buf); - assert_eq!(buf_slice.len(), 128); + assert_eq!(buf_slice.len(), 0); } #[test] @@ -73,13 +70,15 @@ mod tests { let mut buf_slice = BufferSlice::new(&mut buf); + buf_slice.extend_from_slice(&[0; 70]).unwrap(); + buf_slice.truncate(64); assert_eq!(buf_slice.len(), 64); buf_slice.truncate(256); - assert_eq!(buf_slice.len(), 128); + assert_eq!(buf_slice.len(), 64); } #[test] @@ -88,28 +87,22 @@ mod tests { let mut buf_slice = BufferSlice::new(&mut buf); - assert_eq!(buf_slice.len(), 128); - assert_eq!(buf_slice.extend_from_slice(&[0, 0, 0]), Err(aead::Error)); - - buf_slice.truncate(64); + assert_eq!(buf_slice.len(), 0); + assert_eq!(buf_slice.extend_from_slice(&[0; 129]), Err(aead::Error)); assert_eq!(buf_slice.extend_from_slice(&[0, 0, 0, 0, 0, 0]), Ok(())); - assert_eq!(buf_slice.len(), 70); + assert_eq!(buf_slice.len(), 6); } #[test] - fn reset_sets_length_to_buffer_max_length() { + fn fill_sets_length_to_buffer_max_length() { let mut buf = alloc::vec![0u8; 128]; let mut buf_slice = BufferSlice::new(&mut buf); - assert_eq!(buf_slice.len(), 128); - - buf_slice.truncate(64); - - assert_eq!(buf_slice.len(), 64); + assert_eq!(buf_slice.len(), 0); - buf_slice.reset(); + buf_slice.fill(); assert_eq!(buf_slice.len(), 128); }