From d169a41836ffd2cf5382a5a0be54b7f18e6c2c94 Mon Sep 17 00:00:00 2001 From: Sachymetsu Date: Sun, 12 Apr 2026 12:11:47 +0200 Subject: [PATCH] Refactor to constant-time opflags --- Cargo.lock | 21 ---- Cargo.toml | 1 - src/basic_kats.rs | 28 ++--- src/herding_kats/harness.rs | 4 +- src/lib.rs | 2 + src/opflags.rs | 188 +++++++++++++++++++++++++++++++++ src/ops.rs | 42 ++++++++ src/strobe.rs | 203 +++++++++++++++--------------------- 8 files changed, 334 insertions(+), 155 deletions(-) create mode 100644 src/opflags.rs create mode 100644 src/ops.rs diff --git a/Cargo.lock b/Cargo.lock index 54e3d23..60ec358 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -17,26 +17,6 @@ dependencies = [ "libc", ] -[[package]] -name = "enumflags2" -version = "0.7.12" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1027f7680c853e056ebcec683615fb6fbbc07dbaa13b4d5d9442b146ded4ecef" -dependencies = [ - "enumflags2_derive", -] - -[[package]] -name = "enumflags2_derive" -version = "0.7.12" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "67c78a4d8fdf9953a5c9d458f9efe940fd97a0cab0941c075a813ac594733827" -dependencies = [ - "proc-macro2", - "quote", - "syn", -] - [[package]] name = "hex" version = "0.4.3" @@ -168,7 +148,6 @@ checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" name = "wharrgarbl" version = "0.1.0" dependencies = [ - "enumflags2", "hex", "keccak", "serde", diff --git a/Cargo.toml b/Cargo.toml index 0e3edb3..127e0ba 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -8,7 +8,6 @@ edition = "2024" [dependencies] keccak = "0.2" -enumflags2 = "0.7.12" subtle = { version = "2.6", default-features = false } [dev-dependencies] diff --git a/src/basic_kats.rs b/src/basic_kats.rs index 2b87244..df1b3da 100644 --- a/src/basic_kats.rs +++ b/src/basic_kats.rs @@ -6,14 +6,14 @@ use crate::{ keccakf::KECCAK_BUFFER_SIZE, - strobe::{SecurityParameter, StrobeState}, + strobe::{Role, SecurityParameter, StrobeState}, }; extern crate std; #[test] fn test_init_128() { - let s = StrobeState::new(b"", SecurityParameter::B128); + let s = StrobeState::new(b"", SecurityParameter::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 +37,7 @@ fn test_init_128() { #[test] fn test_init_256() { - let s = StrobeState::new(b"", SecurityParameter::B256); + let s = StrobeState::new(b"", SecurityParameter::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 +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", SecurityParameter::B256); + let mut s = StrobeState::new(b"metadatatest", SecurityParameter::B256, Role::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", SecurityParameter::B256); + let mut s = StrobeState::new(b"seqtest", SecurityParameter::B256, Role::Sender); let mut buf = [0u8; 10]; s.prf(&mut buf[..]); @@ -172,8 +172,12 @@ fn test_seq() { #[test] fn test_enc_correctness() { let orig_msg = b"Hello there"; - let mut tx = StrobeState::new(b"enccorrectnesstest", SecurityParameter::B256); - let mut rx = StrobeState::new(b"enccorrectnesstest", SecurityParameter::B256); + let mut tx = StrobeState::new(b"enccorrectnesstest", SecurityParameter::B256, Role::Sender); + let mut rx = StrobeState::new( + b"enccorrectnesstest", + SecurityParameter::B256, + Role::Receiver, + ); tx.key(b"the-combination-on-my-luggage"); rx.key(b"the-combination-on-my-luggage"); @@ -188,8 +192,8 @@ fn test_enc_correctness() { #[test] fn test_mac_correctness_and_soundness() { - let mut tx = StrobeState::new(b"mactest", SecurityParameter::B256); - let mut rx = StrobeState::new(b"mactest", SecurityParameter::B256); + let mut tx = StrobeState::new(b"mactest", SecurityParameter::B256, Role::Sender); + let mut rx = StrobeState::new(b"mactest", SecurityParameter::B256, Role::Receiver); // Just do some stuff with the state @@ -217,7 +221,7 @@ fn test_mac_correctness_and_soundness() { #[test] fn test_long_inputs() { - let mut s = StrobeState::new(b"bigtest", SecurityParameter::B256); + let mut s = StrobeState::new(b"bigtest", SecurityParameter::B256, Role::Sender); const BIG_N: usize = 9823; const SMALL_N: usize = 65; let big_data = [0x34u8; BIG_N]; @@ -274,7 +278,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); + let mut s = StrobeState::new(b"streamingtest", SecurityParameter::B256, Role::Receiver); s.ad(b"mynonce"); @@ -290,7 +294,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); + let mut s = StrobeState::new(b"streamingtest", SecurityParameter::B256, Role::Receiver); s.ad(b"my"); s.ad(b"nonce"); diff --git a/src/herding_kats/harness.rs b/src/herding_kats/harness.rs index a856c54..5379431 100644 --- a/src/herding_kats/harness.rs +++ b/src/herding_kats/harness.rs @@ -4,7 +4,7 @@ use std::{string::String, vec::Vec}; use serde::{Deserialize, Deserializer, de}; -use crate::strobe::{SecurityParameter, StrobeState}; +use crate::strobe::{Role, SecurityParameter, 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) @@ -113,7 +113,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); + let mut strobe = StrobeState::new(proto_string.as_bytes(), security, Role::Sender); operations.into_iter().for_each( |KatOperation { diff --git a/src/lib.rs b/src/lib.rs index 16b97bb..aac9a9a 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -6,6 +6,8 @@ mod basic_kats; #[cfg(test)] mod herding_kats; mod keccakf; +mod opflags; +mod ops; pub mod strobe; /// Version of Strobe that this crate implements. diff --git a/src/opflags.rs b/src/opflags.rs new file mode 100644 index 0000000..02ff57c --- /dev/null +++ b/src/opflags.rs @@ -0,0 +1,188 @@ +use core::ops::{BitAnd, BitAndAssign, BitOr, BitOrAssign, BitXor, BitXorAssign, Not}; + +use subtle::{Choice, ConstantTimeEq}; + +use crate::strobe::Role; + +#[derive(Clone, Copy, PartialEq, Eq, Hash)] +pub struct OpFlags(u8); + +impl OpFlags { + pub const EMPTY: OpFlags = OpFlags(0); + pub const INBOUND: OpFlags = OpFlags(1 << 0); + pub const APP: OpFlags = OpFlags(1 << 1); + pub const CIPHER: OpFlags = OpFlags(1 << 2); + pub const TRANSPORT: OpFlags = OpFlags(1 << 3); + pub const META: OpFlags = OpFlags(1 << 4); + pub const KEYTREE: OpFlags = OpFlags(1 << 5); + + pub(crate) const fn new(val: u8) -> Self { + Self(val) + } + + pub fn contains(&self, other: OpFlags) -> Choice { + (*self & other).ct_eq(&other) + } + + pub fn intersects(&self, other: OpFlags) -> Choice { + (*self & other).ct_ne(&OpFlags::EMPTY) + } + + pub fn set(&mut self, flags: OpFlags, cond: Choice) { + if cond.into() { + *self |= flags; + } else { + *self &= !flags; + } + } + + #[inline(always)] + pub const fn bits(&self) -> u8 { + self.0 + } +} + +impl core::fmt::Debug for OpFlags { + fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result { + write!(f, "OpFlags({:#08b})", self.0) + } +} + +impl ConstantTimeEq for OpFlags { + fn ct_eq(&self, other: &Self) -> subtle::Choice { + self.0.ct_eq(&other.0) + } +} + +impl BitAnd for OpFlags { + type Output = OpFlags; + + fn bitand(self, rhs: Self) -> Self::Output { + Self(self.0 & rhs.0) + } +} + +impl BitAndAssign for OpFlags { + fn bitand_assign(&mut self, rhs: Self) { + self.0 &= rhs.0; + } +} + +impl BitOr for OpFlags { + type Output = OpFlags; + + fn bitor(self, rhs: Self) -> Self::Output { + Self(self.0 | rhs.0) + } +} + +impl BitOrAssign for OpFlags { + fn bitor_assign(&mut self, rhs: Self) { + self.0 |= rhs.0; + } +} + +impl BitXor for OpFlags { + type Output = Self; + + fn bitxor(self, rhs: Self) -> Self::Output { + Self(self.0 ^ rhs.0) + } +} + +impl BitXorAssign for OpFlags { + fn bitxor_assign(&mut self, rhs: Self) { + self.0 ^= rhs.0; + } +} + +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.unwrap_u8(); + } +} + +impl Not for OpFlags { + type Output = OpFlags; + + fn not(self) -> Self::Output { + Self(!self.0) + } +} + +#[cfg(test)] +mod flag_tests { + use super::*; + + extern crate std; + + #[test] + fn debug_shows_toggled_bits() { + let test_case = OpFlags::INBOUND | OpFlags::APP | OpFlags::META; + + let debug = std::format!("{test_case:?}"); + + assert_eq!(&debug, "OpFlags(0b010011)"); + } + + #[test] + fn contains_works() { + assert!(!bool::from(OpFlags::EMPTY.contains(OpFlags::INBOUND))); + assert!(bool::from(OpFlags(0b110).contains(OpFlags::CIPHER))); + assert!(bool::from(OpFlags(0b110).contains(OpFlags::APP))); + assert!(bool::from( + OpFlags(0b110).contains(OpFlags::CIPHER | OpFlags::APP) + )); + assert!(!bool::from( + OpFlags(0b100).contains(OpFlags::CIPHER | OpFlags::APP) + )); + assert!(!bool::from(OpFlags(0b110).contains(OpFlags::INBOUND))); + } + + #[test] + fn intersects_works() { + let intersection = OpFlags::CIPHER | OpFlags::KEYTREE; + + let test_1 = OpFlags::CIPHER | OpFlags::TRANSPORT; + let test_2 = OpFlags::KEYTREE | OpFlags::META; + let test_3 = OpFlags::APP | OpFlags::TRANSPORT; + + assert!(bool::from(test_1.intersects(intersection))); + assert_eq!( + bool::from(test_1.intersects(intersection)), + bool::from(test_1.contains(OpFlags::KEYTREE)) + || bool::from(test_1.contains(OpFlags::CIPHER)) + ); + + assert!(bool::from(test_2.intersects(intersection))); + + assert!(!bool::from(test_3.intersects(intersection))); + assert_eq!( + bool::from(test_3.intersects(intersection)), + bool::from(test_3.contains(OpFlags::KEYTREE)) + || bool::from(test_3.contains(OpFlags::CIPHER)) + ); + } + + #[test] + fn set_works() { + let mut test_case = OpFlags::INBOUND | OpFlags::APP | OpFlags::META; + + test_case.set(OpFlags::META, Choice::from(0)); + + assert!(!bool::from(test_case.contains(OpFlags::META))); + + assert_eq!(test_case, OpFlags::INBOUND | OpFlags::APP); + + test_case.set(OpFlags::META, Choice::from(1)); + + assert!(bool::from(test_case.contains(OpFlags::META))); + } +} diff --git a/src/ops.rs b/src/ops.rs new file mode 100644 index 0000000..e216a8c --- /dev/null +++ b/src/ops.rs @@ -0,0 +1,42 @@ +use crate::opflags::OpFlags; + +macro_rules! define_ops { + { + $( + $name:ident($flags:expr); + )+ + } => { + $( + pub(crate) const $name: OpFlags = OpFlags::new($flags); + )+ + }; +} + +define_ops! { + SEND_ENC(OpFlags::APP.bits() | OpFlags::CIPHER.bits() | OpFlags::TRANSPORT.bits()); + META_SEND_ENC(OpFlags::APP.bits() | OpFlags::CIPHER.bits() | OpFlags::TRANSPORT.bits() | OpFlags::META.bits()); + RECV_ENC(OpFlags::INBOUND.bits() | OpFlags::APP.bits() | OpFlags::CIPHER.bits() | OpFlags::TRANSPORT.bits()); + META_RECV_ENC(OpFlags::INBOUND.bits() | OpFlags::APP.bits() | OpFlags::CIPHER.bits() | OpFlags::TRANSPORT.bits() | OpFlags::META.bits()); + + SEND_MAC(OpFlags::CIPHER.bits() | OpFlags::TRANSPORT.bits()); + META_SEND_MAC(OpFlags::CIPHER.bits() | OpFlags::TRANSPORT.bits() | OpFlags::META.bits()); + RECV_MAC(OpFlags::INBOUND.bits() | OpFlags::CIPHER.bits() | OpFlags::TRANSPORT.bits()); + META_RECV_MAC(OpFlags::INBOUND.bits() | OpFlags::CIPHER.bits() | OpFlags::TRANSPORT.bits() | OpFlags::META.bits()); + + PRF(OpFlags::INBOUND.bits() | OpFlags::APP.bits() | OpFlags::CIPHER.bits()); + META_PRF(OpFlags::INBOUND.bits() | OpFlags::APP.bits() | OpFlags::CIPHER.bits() | OpFlags::META.bits()); + + RATCHET(OpFlags::CIPHER.bits()); + META_RATCHET(OpFlags::CIPHER.bits() | OpFlags::META.bits()); + + AD(OpFlags::APP.bits()); + META_AD(OpFlags::APP.bits() | OpFlags::META.bits()); + + KEY(OpFlags::APP.bits() | OpFlags::CIPHER.bits()); + META_KEY(OpFlags::APP.bits() | OpFlags::CIPHER.bits() | OpFlags::META.bits()); + + SEND_CLR(OpFlags::APP.bits() | OpFlags::TRANSPORT.bits()); + META_SEND_CLR(OpFlags::APP.bits() | OpFlags::TRANSPORT.bits() | OpFlags::META.bits()); + RECV_CLR(OpFlags::INBOUND.bits() | OpFlags::APP.bits() | OpFlags::TRANSPORT.bits()); + META_RECV_CLR(OpFlags::INBOUND.bits() | OpFlags::APP.bits() | OpFlags::TRANSPORT.bits() | OpFlags::META.bits()); +} diff --git a/src/strobe.rs b/src/strobe.rs index e2d2fdd..d3f35ba 100644 --- a/src/strobe.rs +++ b/src/strobe.rs @@ -1,35 +1,25 @@ -use enumflags2::{BitFlag, BitFlags}; -use subtle::ConstantTimeEq; +use subtle::{Choice, ConstantTimeEq}; use crate::{ GarbledError, STROBE_VERSION, keccakf::{KECCAK_BUFFER_SIZE, KeccakF1600}, + opflags::OpFlags, + ops, }; -#[enumflags2::bitflags] -#[repr(u8)] -#[derive(Copy, Clone, Debug, PartialEq)] -enum OpFlags { - /// Is data being moved inbound - Inbound = 0b000001, // 1<<0 - /// Is data being sent to the application - App = 0b000010, // 1<<1 - /// Does this operation use cipher output - Cipher = 0b000100, // 1<<2 - /// Is data being sent for transport - Transport = 0b001000, // 1<<3 - /// Use exclusively for metadata operations - Meta = 0b010000, // 1<<4 - /// Reserved and currently unimplemented. Using this will cause a panic. - KeyTree = 0b100000, // 1<<5 -} - #[derive(Debug, Clone, Copy, PartialEq, Eq)] -enum Role { +#[repr(u8)] +pub enum Role { Sender, Receiver, } +impl ConstantTimeEq for Role { + fn ct_eq(&self, other: &Self) -> Choice { + (*self as u8).ct_eq(&(*other as u8)) + } +} + #[derive(Debug, Clone, Copy)] #[repr(usize)] pub enum SecurityParameter { @@ -49,11 +39,11 @@ pub struct StrobeState { position: usize, /// Index into `state` start: usize, - /// Represents whether we're a sender or a receiver or uninitialized - role: Option, + /// Represents whether we're a sender or a receiver + role: Role, /// The last operation performed. This is to verify that the `more` flag is only used across /// identical operations. - prev_flags: BitFlags, + prev_flags: OpFlags, } macro_rules! define_mut_operations { @@ -64,9 +54,9 @@ macro_rules! define_mut_operations { #[$doc] pub fn $name(&mut self, data: &mut [u8]) { let flags = $flags; - let prev_flags = self.prev_flags.bits(); - let more = prev_flags.ct_eq(&flags.bits()); - self.operate(flags, data, bool::from(more)); + let prev_flags = self.prev_flags; + let more = prev_flags.ct_eq(&flags); + self.operate(flags, data, more); } )* }; @@ -80,9 +70,9 @@ macro_rules! define_non_mut_operations { #[$doc] pub fn $name(&mut self, data: &[u8]) { let flags = $flags; - let prev_flags = self.prev_flags.bits(); - let more = prev_flags.ct_eq(&flags.bits()); - self.operate_no_mutate(flags, data, bool::from(more)); + let prev_flags = self.prev_flags; + let more = prev_flags.ct_eq(&flags); + self.operate_no_mutate(flags, data, more); } )* }; @@ -112,7 +102,7 @@ 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) -> Self { + pub fn new(protocol: &[u8], sec: SecurityParameter, role: Role) -> Self { let rate = KECCAK_BUFFER_SIZE - (sec as usize) / 4 - 2; assert!((1..254).contains(&rate)); @@ -132,8 +122,8 @@ impl StrobeState { rate, position: 0, start: 0, - role: None, - prev_flags: OpFlags::empty(), + role, + prev_flags: OpFlags::EMPTY, }; // Mix the protocol into the state @@ -154,7 +144,7 @@ impl StrobeState { pub fn reset_ops(&mut self) { // This prevents streaming so to always make the prev_flags == flags // comparison always fail - self.prev_flags = OpFlags::empty(); + self.prev_flags = OpFlags::EMPTY; } // Runs the permutation function on the internal state @@ -263,22 +253,16 @@ impl StrobeState { /// Mixes the current state index and flags into the state, accounting for whether we are /// sending or receiving - fn begin_op(&mut self, mut flags: BitFlags) { - if flags.contains(OpFlags::Transport) { - let op_role = if flags.contains(OpFlags::Inbound) { + fn begin_op(&mut self, mut flags: OpFlags) { + if flags.contains(OpFlags::TRANSPORT).into() { + let op_role = if flags.contains(OpFlags::INBOUND).into() { Role::Receiver } else { Role::Sender }; - // If uninitialized, take on the direction of the first directional operation we get - if self.role.is_none() { - self.role = Some(op_role); - } - // So that the sender and receiver agree, toggle the I flag as necessary - // This is equivalent to flags ^= is_receiver - flags.set(OpFlags::Inbound, self.role.unwrap() != op_role); + flags.set(OpFlags::INBOUND, self.role.ct_ne(&op_role)); } let old_start = self.start; @@ -287,8 +271,10 @@ impl StrobeState { // Mix in the position and flags self.absorb(&[old_start as u8, flags.bits()]); - let force_permutation = flags.contains(OpFlags::Cipher) || flags.contains(OpFlags::KeyTree); - if force_permutation && self.position != 0 { + let mut force_permutation = flags.intersects(OpFlags::CIPHER | OpFlags::KEYTREE); + force_permutation &= self.position.ct_ne(&0); + + if force_permutation.into() { self.permutation_f(); } } @@ -296,58 +282,53 @@ impl StrobeState { /// Performs the state / data transformation that corresponds to the given flags. If `more` is /// given, this will treat `data` as a continuation of the data given in the previous /// call to `operate`. - fn operate(&mut self, flags: BitFlags, data: &mut [u8], more: bool) { + fn operate(&mut self, flags: OpFlags, data: &mut [u8], more: Choice) { self.prev_flags = flags; // If `more` isn't set, this is a new operation. Do the begin_op sequence - if !more { + if !bool::from(more) { self.begin_op(flags); } // Meta-ness is only relevant for `begin_op`. Remove it to simplify the below logic. - let flags = flags & !OpFlags::Meta; - - // TODO?: Assert that input is empty under some flag conditions - if flags.contains(OpFlags::Cipher | OpFlags::Transport) && !flags.contains(OpFlags::Inbound) - { - // This is equivalent to the `duplex` operation in the Python implementation, with - // `cafter = True` - if flags == OpFlags::Cipher | OpFlags::Transport { - // This is `send_mac`. Pretend the input is all zeros - self.copy_state(data); - } else { - self.absorb_and_set(data); - } - } else if flags == OpFlags::Inbound | OpFlags::App | OpFlags::Cipher { - // Special case of case below. This is PRF. Use `squeeze` instead of `exchange`. - self.squeeze(data); - } else if flags.contains(OpFlags::Cipher) { - // This is equivalent to the `duplex` operation in the Python implementation, with - // `cbefore = True` - self.exchange(data); - } else { - // This should normally call `absorb`, but `absorb` does not mutate, so the implementor - // should have used operate_no_mutate instead - unreachable!("operate should not be called for operations that do not require mutation") + let flags = flags & !OpFlags::META; + + // Flags that don't pass this assertion should normally call `absorb`, but `absorb` does not mutate, + // so the implementor should have used operate_no_mutate instead + // RATCHET is special-cased to never call operate directly + debug_assert!(flags != ops::KEY && bool::from(flags.contains(OpFlags::CIPHER))); + + match flags { + ops::PRF => self.squeeze(data), + ops::SEND_MAC => self.copy_state(data), + ops::SEND_ENC => self.absorb_and_set(data), + _ => self.exchange(data), } } /// Performs the state transformation that corresponds to the given flags. If `more` is given, /// this will treat `data` as a continuation of the data given in the previous call to /// `operate`. This uses non-mutating variants of the specializations of the `duplex` function. - fn operate_no_mutate(&mut self, flags: BitFlags, data: &[u8], more: bool) { + fn operate_no_mutate(&mut self, flags: OpFlags, data: &[u8], more: Choice) { self.prev_flags = flags; // If `more` isn't set, this is a new operation. Do the begin_op sequence - if !more { + if !bool::from(more) { self.begin_op(flags); } + // Meta-ness is only relevant for `begin_op`. Remove it to simplify the below logic. + let flags = flags & !OpFlags::META; + + // Flags that trigger the assertion to fail are mutating operations. + // RATCHET is special cased to never call operate/operate_no_mutate directly + debug_assert!( + flags != ops::PRF && !bool::from(flags.contains(OpFlags::CIPHER | OpFlags::TRANSPORT)) + || bool::from(flags.contains(OpFlags::INBOUND)) + ); + // There are no non-mutating variants of things with flags & (C | T | I) == C | T - if flags.contains(OpFlags::Cipher | OpFlags::Transport) && !flags.contains(OpFlags::Inbound) - { - unreachable!("operate_no_mutate called on something that requires mutation") - } else if flags.contains(OpFlags::Cipher) { + if flags.contains(OpFlags::CIPHER).into() { // This is equivalent to a non-mutating form of the `duplex` operation in the Python // implementation, with `cbefore = True` self.overwrite(data); @@ -358,20 +339,14 @@ impl StrobeState { } } - fn recv_mac_inner( - &mut self, - mac_copy: &mut [u8], - flags: BitFlags, - ) -> Result<(), GarbledError> { + fn recv_mac_inner(&mut self, mac_copy: &mut [u8], flags: OpFlags) -> Result<(), GarbledError> { // recv_mac can never be streamed - self.operate(flags, mac_copy, false); + self.operate(flags, mac_copy, Choice::from(0u8)); // Constant-time MAC check. This accumulates the truth values of byte == 0 let all_zero: bool = mac_copy .iter() - .fold(subtle::Choice::from(1u8), |all_zero, b| { - all_zero & 0u8.ct_eq(b) - }) + .fold(Choice::from(1u8), |all_zero, b| all_zero & 0u8.ct_eq(b)) .into(); if all_zero { Ok(()) } else { Err(GarbledError) } @@ -380,22 +355,16 @@ impl StrobeState { pub fn recv_mac(&mut self, mac: &[u8; N]) -> Result<(), GarbledError> { let mut mac_copy = *mac; - self.recv_mac_inner( - &mut mac_copy, - OpFlags::Inbound | OpFlags::Cipher | OpFlags::Transport, - ) + self.recv_mac_inner(&mut mac_copy, ops::RECV_MAC) } pub fn meta_recv_mac(&mut self, mac: &[u8; N]) -> Result<(), GarbledError> { let mut mac_copy = *mac; - self.recv_mac_inner( - &mut mac_copy, - OpFlags::Inbound | OpFlags::Cipher | OpFlags::Transport | OpFlags::Meta, - ) + self.recv_mac_inner(&mut mac_copy, ops::META_RECV_MAC) } - fn ratchet_inner(&mut self, num_bytes_to_zero: usize, flags: BitFlags) { + fn ratchet_inner(&mut self, num_bytes_to_zero: usize, flags: OpFlags) { let more = self.prev_flags.bits().ct_eq(&flags.bits()); // We don't make an `operate` call, since this is a super special case. That means we have @@ -410,53 +379,49 @@ impl StrobeState { } pub fn ratchet(&mut self, num_bytes_to_zero: usize) { - let flags = BitFlags::from(OpFlags::Cipher); - - self.ratchet_inner(num_bytes_to_zero, flags); + self.ratchet_inner(num_bytes_to_zero, ops::RATCHET); } pub fn meta_ratchet(&mut self, num_bytes_to_zero: usize) { - let flags = OpFlags::Cipher | OpFlags::Meta; - - self.ratchet_inner(num_bytes_to_zero, flags); + self.ratchet_inner(num_bytes_to_zero, ops::META_RATCHET); } define_mut_operations! { /// SEND ENC - pub fn send_enc(OpFlags::App | OpFlags::Cipher | OpFlags::Transport); + pub fn send_enc(ops::SEND_ENC); /// META SEND ENC - pub fn meta_send_enc(OpFlags::App | OpFlags::Cipher | OpFlags::Transport | OpFlags::Meta); + pub fn meta_send_enc(ops::META_SEND_ENC); /// RECV ENV - pub fn recv_enc(OpFlags::Inbound | OpFlags::App | OpFlags::Cipher | OpFlags::Transport); + pub fn recv_enc(ops::RECV_ENC); /// META RECV ENC - pub fn meta_recv_enc(OpFlags::Inbound | OpFlags::App | OpFlags::Cipher | OpFlags::Transport | OpFlags::Meta); + pub fn meta_recv_enc(ops::META_RECV_ENC); /// SEND MAC - pub fn send_mac(OpFlags::Cipher | OpFlags::Transport); + pub fn send_mac(ops::SEND_MAC); /// META SEND MAC - pub fn meta_send_mac(OpFlags::Cipher | OpFlags::Transport | OpFlags::Meta); + pub fn meta_send_mac(ops::META_SEND_MAC); /// PRF - pub fn prf(OpFlags::Inbound | OpFlags::App | OpFlags::Cipher); + pub fn prf(ops::PRF); /// META PRF - pub fn meta_prf(OpFlags::Inbound | OpFlags::App | OpFlags::Cipher | OpFlags::Meta); + pub fn meta_prf(ops::META_PRF); } define_non_mut_operations! { /// AD - pub fn ad(BitFlags::from(OpFlags::App)); + pub fn ad(ops::AD); /// META AD - pub fn meta_ad(OpFlags::App | OpFlags::Meta); + pub fn meta_ad(ops::META_AD); /// KEY - pub fn key(OpFlags::App | OpFlags::Cipher); + pub fn key(ops::KEY); /// META KEY - pub fn meta_key(OpFlags::App | OpFlags::Cipher | OpFlags::Meta); + pub fn meta_key(ops::META_KEY); /// SEND CLR - pub fn send_clr(OpFlags::App | OpFlags::Transport); + pub fn send_clr(ops::SEND_CLR); /// META SEND CLR - pub fn meta_send_clr(OpFlags::App | OpFlags::Transport | OpFlags::Meta); + pub fn meta_send_clr(ops::META_SEND_CLR); /// RECV CLR - pub fn recv_clr(OpFlags::Inbound | OpFlags::App | OpFlags::Transport); + pub fn recv_clr(ops::RECV_CLR); /// META RECV CLR - pub fn meta_recv_clr(OpFlags::Inbound | OpFlags::App | OpFlags::Transport | OpFlags::Meta); + pub fn meta_recv_clr(ops::META_RECV_CLR); } } @@ -468,7 +433,7 @@ mod tests { #[test] fn version_formatting() { - let s = StrobeState::new(b"", SecurityParameter::B128); + let s = StrobeState::new(b"", SecurityParameter::B128, Role::Sender); let display = std::format!("{s}"); let debug = std::format!("{s:?}"); -- 2.51.2