diff --git a/Cargo.lock b/Cargo.lock index a8f16cf..54ea827 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,6 +2,43 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "aead" +version = "0.6.0-rc.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6b657e772794c6b04730ea897b66a058ccd866c16d1967da05eeeecec39043fe" +dependencies = [ + "crypto-common", + "inout", +] + +[[package]] +name = "anyhow" +version = "1.0.102" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" + +[[package]] +name = "autocfg" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c08606f8c3cbf4ce6ec8e28fb0014a2c086708fe954eaa885384a6165172e7e8" + +[[package]] +name = "bitflags" +version = "2.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4512299f36f043ab09a583e57bceb5a5aab7a73db1805848e8fef3c9e8c78b3" + +[[package]] +name = "block-buffer" +version = "0.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cdd35008169921d80bc60d3d0ab416eecb028c4cd653352907921d95084790be" +dependencies = [ + "hybrid-array", +] + [[package]] name = "cfg-if" version = "1.0.4" @@ -23,6 +60,16 @@ dependencies = [ "libc", ] +[[package]] +name = "crypto-common" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77727bb15fa921304124b128af125e7e3b968275d1b108b379190264f4423710" +dependencies = [ + "hybrid-array", + "rand_core", +] + [[package]] name = "ctutils" version = "0.4.2" @@ -32,6 +79,63 @@ dependencies = [ "cmov", ] +[[package]] +name = "digest" +version = "0.11.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4850db49bf08e663084f7fb5c87d202ef91a3907271aff24a94eb97ff039153c" +dependencies = [ + "block-buffer", + "crypto-common", +] + +[[package]] +name = "equivalent" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" + +[[package]] +name = "foldhash" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" + +[[package]] +name = "getrandom" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0de51e6874e94e7bf76d726fc5d13ba782deca734ff60d5bb2fb2607c7406555" +dependencies = [ + "cfg-if", + "libc", + "r-efi", + "rand_core", + "wasip2", + "wasip3", +] + +[[package]] +name = "hashbrown" +version = "0.15.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9229cfe53dfd69f0609a49f65461bd93001ea1ef889cd5529dd176593f5338a1" +dependencies = [ + "foldhash", +] + +[[package]] +name = "hashbrown" +version = "0.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4f467dd6dccf739c208452f8014c75c18bb8301b050ad1cfb27153803edb0f51" + +[[package]] +name = "heck" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" + [[package]] name = "hex" version = "0.4.3" @@ -44,7 +148,36 @@ version = "0.4.10" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3944cf8cf766b40e2a1a333ee5e9b563f854d5fa49d6a8ca2764e97c6eddb214" dependencies = [ + "ctutils", "typenum", + "zeroize", +] + +[[package]] +name = "id-arena" +version = "2.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3d3067d79b975e8844ca9eb072e16b31c3c1c36928edf9c6789548c524d0d954" + +[[package]] +name = "indexmap" +version = "2.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" +dependencies = [ + "equivalent", + "hashbrown 0.17.0", + "serde", + "serde_core", +] + +[[package]] +name = "inout" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4250ce6452e92010fdf7268ccc5d14faa80bb12fc741938534c58f16804e03c7" +dependencies = [ + "hybrid-array", ] [[package]] @@ -64,18 +197,85 @@ dependencies = [ "hybrid-array", ] +[[package]] +name = "kem" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "01737161ba802849cfd486b5bd209d38ba4943494c249a8126005170c7621edd" +dependencies = [ + "crypto-common", + "rand_core", +] + +[[package]] +name = "leb128fmt" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09edd9e8b54e49e587e4f6295a7d29c3ea94d469cb40ab8ca70b288248a81db2" + [[package]] name = "libc" version = "0.2.184" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "48f5d2a454e16a5ea0f4ced81bd44e4cfc7bd3a507b61887c99fd3538b28e4af" +[[package]] +name = "log" +version = "0.4.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" + [[package]] name = "memchr" version = "2.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79" +[[package]] +name = "ml-kem" +version = "0.3.0-rc.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "04437cb1a66c0b78740927b76cc61f218344b9f6ef3dd430e283274a718ef0e9" +dependencies = [ + "hybrid-array", + "kem", + "module-lattice", + "rand_core", + "sha3", + "zeroize", +] + +[[package]] +name = "module-lattice" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "164eb3faeaecbd14b0b2a917c1b4d0c035097a9c559b0bed85c2cdd032bc8faa" +dependencies = [ + "ctutils", + "hybrid-array", + "num-traits", + "zeroize", +] + +[[package]] +name = "num-traits" +version = "0.2.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841" +dependencies = [ + "autocfg", +] + +[[package]] +name = "prettyplease" +version = "0.2.37" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b" +dependencies = [ + "proc-macro2", + "syn", +] + [[package]] name = "proc-macro2" version = "1.0.106" @@ -94,6 +294,24 @@ dependencies = [ "proc-macro2", ] +[[package]] +name = "r-efi" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" + +[[package]] +name = "rand_core" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" + +[[package]] +name = "semver" +version = "1.0.28" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8a7852d02fc848982e0c167ef163aaff9cd91dc640ba85e263cb1ce46fae51cd" + [[package]] name = "serde" version = "1.0.228" @@ -146,6 +364,16 @@ dependencies = [ "zmij", ] +[[package]] +name = "sha3" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "be176f1a57ce4e3d31c1a166222d9768de5954f811601fb7ca06fc8203905ce1" +dependencies = [ + "digest", + "keccak", +] + [[package]] name = "syn" version = "2.0.117" @@ -169,19 +397,176 @@ version = "1.0.24" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" +[[package]] +name = "unicode-xid" +version = "0.2.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853" + +[[package]] +name = "wasip2" +version = "1.0.3+wasi-0.2.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "20064672db26d7cdc89c7798c48a0fdfac8213434a1186e5ef29fd560ae223d6" +dependencies = [ + "wit-bindgen 0.57.1", +] + +[[package]] +name = "wasip3" +version = "0.4.0+wasi-0.3.0-rc-2026-01-06" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5428f8bf88ea5ddc08faddef2ac4a67e390b88186c703ce6dbd955e1c145aca5" +dependencies = [ + "wit-bindgen 0.51.0", +] + +[[package]] +name = "wasm-encoder" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "990065f2fe63003fe337b932cfb5e3b80e0b4d0f5ff650e6985b1048f62c8319" +dependencies = [ + "leb128fmt", + "wasmparser", +] + +[[package]] +name = "wasm-metadata" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bb0e353e6a2fbdc176932bbaab493762eb1255a7900fe0fea1a2f96c296cc909" +dependencies = [ + "anyhow", + "indexmap", + "wasm-encoder", + "wasmparser", +] + +[[package]] +name = "wasmparser" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47b807c72e1bac69382b3a6fb3dbe8ea4c0ed87ff5629b8685ae6b9a611028fe" +dependencies = [ + "bitflags", + "hashbrown 0.15.5", + "indexmap", + "semver", +] + [[package]] name = "wharrgarbl" version = "0.1.0" dependencies = [ + "aead", "ctutils", + "getrandom", "hex", + "hybrid-array", "keccak", + "ml-kem", + "rand_core", "serde", "serde-big-array", "serde_json", "zeroize", ] +[[package]] +name = "wit-bindgen" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d7249219f66ced02969388cf2bb044a09756a083d0fab1e566056b04d9fbcaa5" +dependencies = [ + "wit-bindgen-rust-macro", +] + +[[package]] +name = "wit-bindgen" +version = "0.57.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" + +[[package]] +name = "wit-bindgen-core" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ea61de684c3ea68cb082b7a88508a8b27fcc8b797d738bfc99a82facf1d752dc" +dependencies = [ + "anyhow", + "heck", + "wit-parser", +] + +[[package]] +name = "wit-bindgen-rust" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7c566e0f4b284dd6561c786d9cb0142da491f46a9fbed79ea69cdad5db17f21" +dependencies = [ + "anyhow", + "heck", + "indexmap", + "prettyplease", + "syn", + "wasm-metadata", + "wit-bindgen-core", + "wit-component", +] + +[[package]] +name = "wit-bindgen-rust-macro" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c0f9bfd77e6a48eccf51359e3ae77140a7f50b1e2ebfe62422d8afdaffab17a" +dependencies = [ + "anyhow", + "prettyplease", + "proc-macro2", + "quote", + "syn", + "wit-bindgen-core", + "wit-bindgen-rust", +] + +[[package]] +name = "wit-component" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d66ea20e9553b30172b5e831994e35fbde2d165325bec84fc43dbf6f4eb9cb2" +dependencies = [ + "anyhow", + "bitflags", + "indexmap", + "log", + "serde", + "serde_derive", + "serde_json", + "wasm-encoder", + "wasm-metadata", + "wasmparser", + "wit-parser", +] + +[[package]] +name = "wit-parser" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ecc8ac4bc1dc3381b7f59c34f00b67e18f910c2c0f50015669dde7def656a736" +dependencies = [ + "anyhow", + "id-arena", + "indexmap", + "log", + "semver", + "serde", + "serde_derive", + "serde_json", + "unicode-xid", + "wasmparser", +] + [[package]] name = "zeroize" version = "1.8.2" diff --git a/Cargo.toml b/Cargo.toml index cf14d38..426ac60 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -13,9 +13,15 @@ parallel = ["keccak/parallel"] keccak = "0.2" ctutils = { version = "0.4.2", default-features = false } zeroize = { version = "1.8.2", default-features = false } +ml-kem = { version = "0.3.0-rc.2", features = ["zeroize"] } +hybrid-array = { version = "0.4.10", features = ["alloc"] } +rand_core = { version = "0.10", default-features = false } +aead = { version = "0.6.0-rc.10" } [dev-dependencies] serde_json = "1" hex = "0.4" serde = { version = "1.0.210", default-features = false, features = ["derive"] } serde-big-array = { version = "0.5" } +getrandom = { version = "0.4.2", features = ["sys_rng"] } +aead = { version = "0.6.0-rc.10", features = ["alloc"] } diff --git a/src/basic_kats.rs b/src/basic_kats.rs index df1b3da..100e392 100644 --- a/src/basic_kats.rs +++ b/src/basic_kats.rs @@ -4,16 +4,19 @@ //! Some tests in the original repo are omitted because we can no longer panic from this //! implementation's public API surface. +use aead::consts::{U16, U65}; +use hybrid_array::Array; + use crate::{ keccakf::KECCAK_BUFFER_SIZE, - strobe::{Role, SecurityParameter, StrobeState}, + strobe::{Role, StrobeSecurity, StrobeState}, }; extern crate std; #[test] fn test_init_128() { - let s = StrobeState::new(b"", SecurityParameter::B128, Role::Sender); + let s = StrobeState::new(b"", StrobeSecurity::B128, Role::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 +40,7 @@ fn test_init_128() { #[test] fn test_init_256() { - let s = StrobeState::new(b"", SecurityParameter::B256, Role::Sender); + let s = StrobeState::new(b"", StrobeSecurity::B256, Role::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 +65,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", SecurityParameter::B256, Role::Sender); + let mut s = StrobeState::new(b"metadatatest", StrobeSecurity::B256, Role::Sender); let mut output = std::vec::Vec::new(); let buf = b"meta1"; @@ -116,7 +119,7 @@ fn test_metadata() { #[test] fn test_seq() { - let mut s = StrobeState::new(b"seqtest", SecurityParameter::B256, Role::Sender); + let mut s = StrobeState::new(b"seqtest", StrobeSecurity::B256, Role::Sender); let mut buf = [0u8; 10]; s.prf(&mut buf[..]); @@ -172,12 +175,8 @@ fn test_seq() { #[test] fn test_enc_correctness() { let orig_msg = b"Hello there"; - let mut tx = StrobeState::new(b"enccorrectnesstest", SecurityParameter::B256, Role::Sender); - let mut rx = StrobeState::new( - b"enccorrectnesstest", - SecurityParameter::B256, - Role::Receiver, - ); + let mut tx = StrobeState::new(b"enccorrectnesstest", StrobeSecurity::B256, Role::Sender); + let mut rx = StrobeState::new(b"enccorrectnesstest", StrobeSecurity::B256, Role::Receiver); tx.key(b"the-combination-on-my-luggage"); rx.key(b"the-combination-on-my-luggage"); @@ -192,8 +191,8 @@ fn test_enc_correctness() { #[test] fn test_mac_correctness_and_soundness() { - let mut tx = StrobeState::new(b"mactest", SecurityParameter::B256, Role::Sender); - let mut rx = StrobeState::new(b"mactest", SecurityParameter::B256, Role::Receiver); + let mut tx = StrobeState::new(b"mactest", StrobeSecurity::B256, Role::Sender); + let mut rx = StrobeState::new(b"mactest", StrobeSecurity::B256, Role::Receiver); // Just do some stuff with the state @@ -201,7 +200,7 @@ fn test_mac_correctness_and_soundness() { let mut msg = b"attack at dawn".to_vec(); tx.send_enc(msg.as_mut_slice()); - let mut mac = [0u8; 16]; + let mut mac: Array = Array([0u8; 16]); tx.send_mac(&mut mac[..]); rx.key(b"secretsauce"); @@ -221,7 +220,7 @@ fn test_mac_correctness_and_soundness() { #[test] fn test_long_inputs() { - let mut s = StrobeState::new(b"bigtest", SecurityParameter::B256, Role::Sender); + let mut s = StrobeState::new(b"bigtest", StrobeSecurity::B256, Role::Sender); const BIG_N: usize = 9823; const SMALL_N: usize = 65; let big_data = [0x34u8; BIG_N]; @@ -241,8 +240,8 @@ fn test_long_inputs() { s.send_enc(big_data.to_vec().as_mut_slice()); s.meta_recv_enc(big_data.to_vec().as_mut_slice()); s.recv_enc(big_data.to_vec().as_mut_slice()); - let _ = s.meta_recv_mac(&small_data); - let _ = s.recv_mac(&small_data); + let _ = s.meta_recv_mac::(&small_data.into()); + let _ = s.recv_mac::(&small_data.into()); let mut big_buf = [0u8; BIG_N]; let mut small_buf = [0u8; SMALL_N]; @@ -278,7 +277,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", SecurityParameter::B256, Role::Receiver); + let mut s = StrobeState::new(b"streamingtest", StrobeSecurity::B256, Role::Receiver); s.ad(b"mynonce"); @@ -294,7 +293,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", SecurityParameter::B256, Role::Receiver); + let mut s = StrobeState::new(b"streamingtest", StrobeSecurity::B256, Role::Receiver); s.ad(b"my"); s.ad(b"nonce"); diff --git a/src/buffer_slice.rs b/src/buffer_slice.rs new file mode 100644 index 0000000..c83c94a --- /dev/null +++ b/src/buffer_slice.rs @@ -0,0 +1,112 @@ +#[derive(Debug)] +pub struct BufferSlice<'slice> { + buffer: &'slice mut [u8], + end: usize, +} + +impl<'slice> BufferSlice<'slice> { + pub const fn new(buffer: &'slice mut [u8]) -> Self { + Self { + end: buffer.len(), + buffer, + } + } + + pub const fn reset(&mut self) { + self.end = self.buffer.len(); + } +} + +impl AsRef<[u8]> for BufferSlice<'_> { + fn as_ref(&self) -> &[u8] { + &self.buffer[..self.end] + } +} + +impl AsMut<[u8]> for BufferSlice<'_> { + fn as_mut(&mut self) -> &mut [u8] { + &mut self.buffer[..self.end] + } +} + +impl aead::Buffer for BufferSlice<'_> { + fn extend_from_slice(&mut self, other: &[u8]) -> aead::Result<()> { + let index = self.end + other.len(); + + if index > self.buffer.len() { + return Err(aead::Error); + } + + self.buffer[self.end..index].copy_from_slice(other); + + self.end = index; + + Ok(()) + } + + fn truncate(&mut self, len: usize) { + self.end = len.min(self.buffer.len()); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use aead::Buffer; + + #[test] + fn defaults_to_full_buffer_size() { + let mut buf = alloc::vec![0u8; 128]; + + let buf_slice = BufferSlice::new(&mut buf); + + assert_eq!(buf_slice.len(), 128); + } + + #[test] + fn truncate_doesnt_extend_past_max_buffer_length() { + let mut buf = alloc::vec![0u8; 128]; + + let mut buf_slice = BufferSlice::new(&mut buf); + + buf_slice.truncate(64); + + assert_eq!(buf_slice.len(), 64); + + buf_slice.truncate(256); + + assert_eq!(buf_slice.len(), 128); + } + + #[test] + fn extend_from_slice_only_works_within_slice_length() { + let mut buf = alloc::vec![0u8; 128]; + + 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.extend_from_slice(&[0, 0, 0, 0, 0, 0]), Ok(())); + assert_eq!(buf_slice.len(), 70); + } + + #[test] + fn reset_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); + + buf_slice.reset(); + + assert_eq!(buf_slice.len(), 128); + } +} diff --git a/src/handshake.rs b/src/handshake.rs new file mode 100644 index 0000000..c2ba0dd --- /dev/null +++ b/src/handshake.rs @@ -0,0 +1,278 @@ +use aead::Buffer; +use hybrid_array::typenum::Unsigned; +use rand_core::CryptoRng; + +use crate::{ + WHARRGHARBL_PROTO, + kem::{KemEncap, KemSecurity}, + strobe::{Role, StrobeSecurity, StrobeState}, + transport::{self, AeadState, AeadStrobe}, +}; + +pub struct ClientHandshake { + kem_sec: KemSecurity, + sec_param: StrobeSecurity, + 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, Role::Sender); + + if let Some(psk) = psk { + strobe.key(psk); + } + + strobe.meta_ad(&kem_sec.to_bytes()); + strobe.meta_ad(&sec_param.to_bytes()); + + Self { + kem_sec, + sec_param, + 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 written = encap.serialize(buf.as_mut())?; + + buf.truncate(written); + + self.strobe.send_clr(buf.as_ref()); + self.strobe.send_mac(&mut tag); + self.strobe.ratchet(self.sec_param.rachet_bytes()); + + buf.extend_from_slice(&tag)?; + + self.decap = Some(decap); + + Ok(()) + } + + pub fn receive(&mut self, ciphertext: &[u8]) -> aead::Result { + let decap = self.decap.as_mut().ok_or(aead::Error)?; + + let tag = ciphertext + .len() + .checked_sub(::TagSize::to_usize()) + .ok_or(aead::Error)?; + + let (ciphertext, tag) = ciphertext.split_at(tag); + + let tag: aead::Tag = tag.try_into().unwrap(); + + self.strobe.recv_clr(ciphertext); + self.strobe.recv_mac(&tag)?; + self.strobe.ratchet(self.sec_param.rachet_bytes()); + + let shared = decap.decapsulate(ciphertext)?; + + Ok(self.finish(shared)) + } + + fn finish(&mut self, shared: ml_kem::SharedKey) -> 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(); + + self.strobe.prf(&mut key); + self.strobe.prf(&mut inbound); + self.strobe.prf(&mut outbound); + + AeadState::new(key, self.sec_param, outbound, inbound, Role::Sender) + } +} + +pub struct ServerHandshake { + kem_sec: KemSecurity, + sec_param: StrobeSecurity, + 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, Role::Receiver); + + if let Some(psk) = psk { + strobe.key(psk); + } + + strobe.meta_ad(&kem_sec.to_bytes()); + strobe.meta_ad(&sec_param.to_bytes()); + + Self { + kem_sec, + sec_param, + strobe, + } + } + + pub fn respond( + &mut self, + rng: &mut impl CryptoRng, + buf: &mut dyn Buffer, + ) -> aead::Result { + let slice = buf.as_ref(); + + let tag = slice + .len() + .checked_sub(::TagSize::to_usize()) + .ok_or(aead::Error)?; + + let (encap, tag) = buf.as_ref().split_at(tag); + + let tag: aead::Tag = tag.try_into().unwrap(); + + self.strobe.recv_clr(encap); + self.strobe.recv_mac(&tag)?; + self.strobe.ratchet(self.sec_param.rachet_bytes()); + + let (cipher, shared) = KemEncap::encapsulate_from_slice(self.kem_sec, encap, rng)?; + + 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()); + + buf.extend_from_slice(cipher.as_ref())?; + buf.extend_from_slice(&tag)?; + + Ok(self.finish(shared)) + } + + fn finish(&mut self, shared: ml_kem::SharedKey) -> 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(); + + self.strobe.prf(&mut key); + self.strobe.prf(&mut inbound); + self.strobe.prf(&mut outbound); + + AeadState::new(key, self.sec_param, outbound, inbound, Role::Receiver) + } +} + +#[cfg(test)] +mod tests { + use crate::buffer_slice::BufferSlice; + + use super::*; + + #[test] + fn handshake_protocol_happy_path() -> aead::Result<()> { + let psk: [u8; 32] = [ + 31, 48, 29, 177, 88, 236, 186, 84, 65, 51, 214, 243, 174, 24, 45, 101, 229, 129, 62, + 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)); + + // BufferSlice acts as our transport across the webz + let mut buf = alloc::vec![0u8; 2048]; + let mut buf = BufferSlice::new(&mut buf); + + let mut rng = rand_core::UnwrapErr(getrandom::SysRng); + + alice.send(&mut rng, &mut buf)?; + + // Pretend to send ek across the webz: client -> server + let bob = bob.respond(&mut rng, &mut buf)?; + + // Pretend to send ciphertext across the webz: server -> client + let alice = alice.receive(buf.as_ref())?; + + assert_eq!(alice.aead.key, bob.aead.key); + + // Both AeadStates have derived base nonces for each context. + // Inbound context nonces will not match Outbound context nonces. + assert_eq!(alice.trump, bob.trump); + assert_eq!(alice.epstein, bob.epstein); + assert_ne!(alice.trump, alice.epstein); + assert_ne!(bob.trump, bob.epstein); + + Ok(()) + } + + #[test] + #[should_panic] + fn handshake_fails_with_strobe_security_mismatch() { + let psk: [u8; 32] = [ + 31, 48, 29, 177, 88, 236, 186, 84, 65, 51, 214, 243, 174, 24, 45, 101, 229, 129, 62, + 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)); + + // BufferSlice acts as our transport across the webz + let mut buf = alloc::vec![0u8; 2048]; + let mut buf = BufferSlice::new(&mut buf); + + let mut rng = rand_core::UnwrapErr(getrandom::SysRng); + + alice.send(&mut rng, &mut buf).unwrap(); + + // Pretend to send ek across the webz: client -> server + let _bob = bob.respond(&mut rng, &mut buf).unwrap(); + } + + #[test] + #[should_panic] + fn handshake_fails_with_kem_security_mismatch() { + let psk: [u8; 32] = [ + 31, 48, 29, 177, 88, 236, 186, 84, 65, 51, 214, 243, 174, 24, 45, 101, 229, 129, 62, + 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)); + + // BufferSlice acts as our transport across the webz + let mut buf = alloc::vec![0u8; 2048]; + let mut buf = BufferSlice::new(&mut buf); + + let mut rng = rand_core::UnwrapErr(getrandom::SysRng); + + alice.send(&mut rng, &mut buf).unwrap(); + + // Pretend to send ek across the webz: client -> server + let _bob = bob.respond(&mut rng, &mut buf).unwrap(); + } + + #[test] + #[should_panic] + fn handshake_fails_with_psk_mismatch() { + let psk: [u8; 32] = [ + 31, 48, 29, 177, 88, 236, 186, 84, 65, 51, 214, 243, 174, 24, 45, 101, 229, 129, 62, + 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)); + + // BufferSlice acts as our transport across the webz + let mut buf = alloc::vec![0u8; 2048]; + let mut buf = BufferSlice::new(&mut buf); + + let mut rng = rand_core::UnwrapErr(getrandom::SysRng); + + alice.send(&mut rng, &mut buf).unwrap(); + + // Pretend to send ek across the webz: client -> server + let _bob = bob.respond(&mut rng, &mut buf).unwrap(); + } +} diff --git a/src/herding_kats/harness.rs b/src/herding_kats/harness.rs index 5379431..f2393c6 100644 --- a/src/herding_kats/harness.rs +++ b/src/herding_kats/harness.rs @@ -2,9 +2,10 @@ extern crate std; use std::{string::String, vec::Vec}; +use aead::consts::U14; use serde::{Deserialize, Deserializer, de}; -use crate::strobe::{Role, SecurityParameter, StrobeState}; +use crate::strobe::{Role, StrobeSecurity, StrobeState}; /// 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) @@ -12,7 +13,7 @@ use crate::strobe::{Role, SecurityParameter, StrobeState}; struct KatHarness { proto_string: String, #[serde(deserialize_with = "security_param_from_bits")] - security: SecurityParameter, + security: StrobeSecurity, operations: Vec, } @@ -79,7 +80,7 @@ fn run_kat_operation( "recv_ENC" => s.recv_enc(data), "send_MAC" => s.send_mac(data), "recv_MAC" => s - .recv_mac::<14>(data.as_ref().try_into().unwrap()) + .recv_mac::(data.as_ref().try_into().unwrap()) .unwrap_or(()), "RATCHET" => panic!("Got RATCHET op without length input"), _ => panic!("Unexpected op name: {}", op_name), @@ -95,7 +96,7 @@ fn run_kat_operation( "recv_ENC" => s.meta_recv_enc(data), "send_MAC" => s.meta_send_mac(data), "recv_MAC" => s - .meta_recv_mac::<14>(data.as_ref().try_into().unwrap()) + .meta_recv_mac::(data.as_ref().try_into().unwrap()) .unwrap_or(()), "RATCHET" => panic!("Got RATCHET op without length input"), _ => panic!("Unexpected op name: {}", op_name), @@ -153,10 +154,10 @@ pub fn test_against_kat>(filename: P) { fn security_param_from_bits<'de, D: Deserializer<'de>>( deserializer: D, -) -> Result { +) -> Result { match u64::deserialize(deserializer)? { - 128 => Ok(SecurityParameter::B128), - 256 => Ok(SecurityParameter::B256), + 128 => Ok(StrobeSecurity::B128), + 256 => Ok(StrobeSecurity::B256), n => Err(de::Error::custom(std::format!( "Invalid security parameter: {}", n diff --git a/src/kem.rs b/src/kem.rs new file mode 100644 index 0000000..c0097c8 --- /dev/null +++ b/src/kem.rs @@ -0,0 +1,137 @@ +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 aac9a9a..e575f24 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -3,23 +3,19 @@ #[cfg(test)] mod basic_kats; +pub mod buffer_slice; +pub mod handshake; #[cfg(test)] mod herding_kats; mod keccakf; +pub mod kem; mod opflags; mod ops; pub mod strobe; +pub mod transport; + +extern crate alloc; /// Version of Strobe that this crate implements. pub static STROBE_VERSION: &str = "1.0.2"; - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -pub struct GarbledError; - -impl core::fmt::Display for GarbledError { - fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result { - f.write_str("Protocol Failure") - } -} - -impl core::error::Error for GarbledError {} +pub static WHARRGHARBL_PROTO: &str = "WGBL-v0.0-STv1.0.2"; diff --git a/src/opflags.rs b/src/opflags.rs index a252724..6a7cc16 100644 --- a/src/opflags.rs +++ b/src/opflags.rs @@ -2,8 +2,6 @@ use core::ops::{BitAnd, BitAndAssign, BitOr, BitOrAssign, BitXor, BitXorAssign, use ctutils::{Choice, CtAssign, CtEq}; -use crate::strobe::Role; - #[derive(Clone, Copy, PartialEq, Eq, Hash)] #[repr(transparent)] pub struct OpFlags(u8); @@ -38,7 +36,7 @@ impl OpFlags { } pub fn set(&mut self, flags: OpFlags, cond: Choice) { - const OPS: [fn(&mut OpFlags, OpFlags); 2] = [OpFlags::remove, OpFlags::insert]; + static OPS: [fn(&mut OpFlags, OpFlags); 2] = [OpFlags::remove, OpFlags::insert]; let index = (cond.to_u8() & 1) as usize; @@ -123,19 +121,6 @@ impl BitXorAssign for OpFlags { } } -impl BitXorAssign for OpFlags { - fn bitxor_assign(&mut self, rhs: Role) { - self.0 ^= rhs as u8; - } -} - -// This is for a specific case to toggle INBOUND -impl BitXorAssign for OpFlags { - fn bitxor_assign(&mut self, rhs: Choice) { - self.0 ^= rhs.to_u8(); - } -} - impl Not for OpFlags { type Output = OpFlags; diff --git a/src/strobe.rs b/src/strobe.rs index bce7f54..419bcb8 100644 --- a/src/strobe.rs +++ b/src/strobe.rs @@ -1,8 +1,11 @@ +use core::ops::BitXor; + use ctutils::{Choice, CtAssign, CtEq, CtLt, CtSelect}; +use hybrid_array::{Array, ArraySize}; use zeroize::Zeroize; use crate::{ - GarbledError, STROBE_VERSION, + STROBE_VERSION, keccakf::{KECCAK_BUFFER_SIZE, KeccakF1600}, opflags::OpFlags, ops, @@ -18,23 +21,41 @@ mod role { #[derive(Debug, Clone, Copy, PartialEq, Eq)] #[repr(u8)] pub enum Role { - Sender, - Receiver, + Sender = 0, + Receiver = 1, +} + +impl BitXor for Role { + fn bitxor(self, rhs: Self) -> Self::Output { + (self as u8) ^ (rhs as u8) + } + + type Output = u8; } #[derive(Debug, Clone, Copy)] -#[repr(usize)] -pub enum SecurityParameter { +#[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() + } +} + #[derive(Clone)] pub struct StrobeState { /// Internal Keccak state pub(crate) state: KeccakF1600, /// Security parameter (128 or 256 bits) - sec: SecurityParameter, + sec: StrobeSecurity, /// This is the `R` parameter in the Strobe spec rate: usize, /// Index into `state` @@ -103,8 +124,8 @@ 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 { - SecurityParameter::B128 => f.write_str("128")?, - SecurityParameter::B256 => f.write_str("256")?, + StrobeSecurity::B128 => f.write_str("128")?, + StrobeSecurity::B256 => f.write_str("256")?, } f.write_str("/1600-v")?; f.write_str(STROBE_VERSION) @@ -122,8 +143,8 @@ impl core::fmt::Debug for StrobeState { } impl StrobeState { - /// Makes a new `StrobeTransport` object with a given protocol byte string and security parameter. - pub fn new(protocol: &[u8], sec: SecurityParameter, role: Role) -> Self { + /// Makes a new `StrobeState` object with a given protocol byte string and security parameter. + pub fn new(protocol: &[u8], sec: StrobeSecurity, role: Role) -> Self { let rate = KECCAK_BUFFER_SIZE - (sec as usize) / 4 - 2; assert!((1..254).contains(&rate)); @@ -324,22 +345,24 @@ impl StrobeState { // RATCHET is special-cased to never call operate directly debug_assert!(flags != ops::KEY && flags.contains(OpFlags::CIPHER).to_bool()); - const SPECIAL_CASES: [OpFlags; 3] = [ops::PRF, ops::SEND_MAC, ops::SEND_ENC]; - const OPS: [fn(&mut StrobeState, data: &mut [u8]); 4] = [ + static SPECIAL_CASES: [OpFlags; 3] = [ops::PRF, ops::SEND_MAC, ops::SEND_ENC]; + static OPS: [fn(&mut StrobeState, data: &mut [u8]); 4] = [ StrobeState::squeeze, StrobeState::copy_state, StrobeState::absorb_and_set, StrobeState::exchange, ]; - // Constant time resolution of op index + // Constant time resolution of op index, applied with a mask to ensure no value greater than + // 3 will ever be produced let index = SPECIAL_CASES .iter() .enumerate() .fold(3usize, |mut res, (index, op)| { res.ct_assign(&index, flags.ct_eq(op)); res - }); + }) + & 0b11; OPS[index](self, data); } @@ -362,7 +385,7 @@ 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()); - const OPS: [fn(&mut StrobeState, data: &[u8]); 2] = + static OPS: [fn(&mut StrobeState, data: &[u8]); 2] = [StrobeState::absorb, StrobeState::overwrite]; let index = (flags.ct_eq(&ops::KEY).to_u8() & 1) as usize; @@ -370,7 +393,7 @@ impl StrobeState { OPS[index](self, data); } - fn recv_mac_inner(&mut self, flags: OpFlags, mac_copy: &mut [u8]) -> Result<(), GarbledError> { + fn recv_mac_inner(&mut self, flags: OpFlags, mac_copy: &mut [u8]) -> Result<(), aead::Error> { // recv_mac can never be streamed self.operate(flags, mac_copy, Choice::FALSE); @@ -382,18 +405,18 @@ impl StrobeState { if all_zero.to_bool() { Ok(()) } else { - Err(GarbledError) + Err(aead::Error) } } - pub fn recv_mac(&mut self, mac: &[u8; N]) -> Result<(), GarbledError> { - let mut mac_copy = *mac; + pub fn recv_mac(&mut self, mac: &Array) -> Result<(), aead::Error> { + let mut mac_copy = mac.clone(); self.recv_mac_inner(ops::RECV_MAC, &mut mac_copy) } - pub fn meta_recv_mac(&mut self, mac: &[u8; N]) -> Result<(), GarbledError> { - let mut mac_copy = *mac; + pub fn meta_recv_mac(&mut self, mac: &Array) -> Result<(), aead::Error> { + let mut mac_copy = mac.clone(); self.recv_mac_inner(ops::META_RECV_MAC, &mut mac_copy) } @@ -469,7 +492,7 @@ mod tests { #[test] fn version_formatting() { - let s = StrobeState::new(b"", SecurityParameter::B128, Role::Sender); + let s = StrobeState::new(b"", StrobeSecurity::B128, Role::Sender); let display = std::format!("{s}"); let debug = std::format!("{s:?}"); diff --git a/src/transport.rs b/src/transport.rs new file mode 100644 index 0000000..2d3fbe0 --- /dev/null +++ b/src/transport.rs @@ -0,0 +1,308 @@ +use aead::{ + AeadInOut, Buffer, Key, TagPosition, + consts::{U16, U32}, +}; +use ctutils::{CtEq, CtSelect}; + +use crate::{ + WHARRGHARBL_PROTO, + strobe::{Role, StrobeSecurity, StrobeState}, +}; + +pub struct AeadStrobe { + pub(crate) key: Key, + param: StrobeSecurity, +} + +impl aead::AeadCore for AeadStrobe { + type NonceSize = U16; + type TagSize = U16; + const TAG_POSITION: TagPosition = TagPosition::Postfix; +} + +impl aead::KeySizeUser for AeadStrobe { + type KeySize = U32; +} + +impl aead::AeadInOut for AeadStrobe { + fn encrypt_inout_detached( + &self, + nonce: &aead::Nonce, + associated_data: &[u8], + 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, + crate::strobe::Role::Sender, + ); + + strobe.key(&self.key); + strobe.meta_ad(&self.param.to_bytes()); + strobe.meta_ad(nonce); + strobe.ad(associated_data); + strobe.send_enc(buffer.get_out()); + + strobe.send_mac(&mut tag); + + Ok(tag) + } + + fn decrypt_inout_detached( + &self, + nonce: &aead::Nonce, + associated_data: &[u8], + mut buffer: aead::inout::InOutBuf<'_, '_, u8>, + tag: &aead::Tag, + ) -> aead::Result<()> { + let mut strobe = StrobeState::new( + WHARRGHARBL_PROTO.as_bytes(), + self.param, + crate::strobe::Role::Receiver, + ); + + strobe.key(&self.key); + strobe.meta_ad(&self.param.to_bytes()); + strobe.meta_ad(nonce); + strobe.ad(associated_data); + strobe.recv_enc(buffer.get_out()); + + strobe.recv_mac(tag) + } +} + +pub struct AeadState { + pub(crate) aead: AeadStrobe, + pub(crate) epstein: aead::Nonce, + pub(crate) trump: aead::Nonce, + role: Role, +} + +impl AeadState { + pub fn new( + key: Key, + sec: StrobeSecurity, + outbound: aead::Nonce, + inbound: aead::Nonce, + role: Role, + ) -> Self { + assert_ne!( + &inbound, &outbound, + "The Base Nonces MUST NOT equal to each other" + ); + + Self { + aead: AeadStrobe { key, param: sec }, + epstein: outbound, + trump: inbound, + role, + } + } + + fn select_nonce(&self, role: Role) -> aead::Nonce { + let role_context = self.role ^ role; + + self.epstein.ct_select(&self.trump, role_context.ct_eq(&1)) + } + + fn mix_nonce(&self, position: [u8; 8], role: Role) -> aead::Nonce { + let mut nonce = self.select_nonce(role); + + let mid = nonce.len() - position.len(); + + nonce[mid..] + .iter_mut() + .zip(position) + .for_each(|(n, p)| *n ^= p); + + nonce + } + + pub fn split(&self) -> (SendState<'_>, RecvState<'_>) { + ( + SendState { + transport: self, + counter: 0, + }, + RecvState { + transport: self, + counter: 0, + }, + ) + } +} + +pub struct SendState<'a> { + transport: &'a AeadState, + counter: u64, +} + +impl SendState<'_> { + 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); + } + + let encryption_result = self.transport.aead.encrypt_in_place( + &self + .transport + .mix_nonce(self.counter.to_be_bytes(), Role::Sender), + ad, + buffer, + ); + + self.counter = self.counter.wrapping_add(1); + + encryption_result + } +} + +pub struct RecvState<'a> { + transport: &'a AeadState, + counter: u64, +} + +impl RecvState<'_> { + 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); + } + + // If the message's MAC doesn't evaluate successfully, this op will fail. + let decryption_result = self.transport.aead.decrypt_in_place( + &self + .transport + .mix_nonce(self.counter.to_be_bytes(), Role::Receiver), + ad, + buffer, + ); + + self.counter = self.counter.wrapping_add(1); + + decryption_result + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn two_way_transport_sync_works() -> aead::Result<()> { + let shared_secret = [ + 0x80, 0x81, 0x82, 0x83, 0x84, 0x85, 0x86, 0x87, 0x88, 0x89, 0x8a, 0x8b, 0x8c, 0x8d, + 0x8e, 0x8f, 0x90, 0x91, 0x92, 0x93, 0x94, 0x95, 0x96, 0x97, 0x98, 0x99, 0x9a, 0x9b, + 0x9c, 0x9d, 0x9e, 0x9f, + ]; + + let outbound = 123u128.to_ne_bytes(); + let inbound = 234u128.to_ne_bytes(); + + let alice = AeadState::new( + shared_secret.into(), + StrobeSecurity::B128, + outbound.into(), + inbound.into(), + Role::Sender, + ); + let bob = AeadState::new( + shared_secret.into(), + StrobeSecurity::B128, + outbound.into(), + inbound.into(), + Role::Receiver, + ); + + let (mut alice_send, mut alice_recv) = alice.split(); + let (mut bob_send, mut bob_recv) = bob.split(); + + let orig = b"Test Message, Please ignore."; + + let ad = b"random"; + + let mut msg = orig.to_vec(); + + // a -> b + alice_send.encrypt(&mut msg, ad)?; + + assert_ne!(orig.as_slice(), msg.as_slice()); + let ct1 = msg.clone(); + + bob_recv.decrypt(&mut msg, ad)?; + + // a -> b + alice_send.encrypt(&mut msg, b"")?; + + assert_ne!(msg.as_slice(), ct1.as_slice()); + let ct2 = msg.clone(); + + bob_recv.decrypt(&mut msg, b"")?; + + // b -> a + bob_send.encrypt(&mut msg, ad)?; + + // None of the ciphertexts should match each other + assert_ne!(msg.as_slice(), ct1.as_slice()); + assert_ne!(msg.as_slice(), ct2.as_slice()); + assert_ne!(ct1.as_slice(), ct2.as_slice()); + + alice_recv.decrypt(&mut msg, ad)?; + + assert_eq!(orig.as_slice(), msg.as_slice()); + + // Counters are tracked from sender to receiver + assert_eq!(alice_send.counter, bob_recv.counter); + assert_eq!(bob_send.counter, alice_recv.counter); + + // Counters are not linked on the same side + assert_ne!(alice_send.counter, alice_recv.counter); + assert_ne!(bob_send.counter, bob_recv.counter); + + Ok(()) + } + + #[test] + #[should_panic] + fn two_way_transport_fails_with_security_level_mismatch() { + let shared_secret = [ + 0x80, 0x81, 0x82, 0x83, 0x84, 0x85, 0x86, 0x87, 0x88, 0x89, 0x8a, 0x8b, 0x8c, 0x8d, + 0x8e, 0x8f, 0x90, 0x91, 0x92, 0x93, 0x94, 0x95, 0x96, 0x97, 0x98, 0x99, 0x9a, 0x9b, + 0x9c, 0x9d, 0x9e, 0x9f, + ]; + + let outbound = 123u128.to_ne_bytes(); + let inbound = 234u128.to_ne_bytes(); + + let alice = AeadState::new( + shared_secret.into(), + StrobeSecurity::B128, + outbound.into(), + inbound.into(), + Role::Sender, + ); + let bob = AeadState::new( + shared_secret.into(), + StrobeSecurity::B256, + outbound.into(), + inbound.into(), + Role::Receiver, + ); + + let (mut alice_send, mut _alice_recv) = alice.split(); + let (mut _bob_send, mut bob_recv) = bob.split(); + + let orig = b"Test Message, Please ignore."; + + let ad = b"random"; + + let mut msg = orig.to_vec(); + + // a -> b + alice_send.encrypt(&mut msg, ad).unwrap(); + + assert_ne!(orig.as_slice(), msg.as_slice()); + + bob_recv.decrypt(&mut msg, ad).unwrap(); + } +}