use core::hint::black_box;
#[must_use = "a Choice carries the outcome of a constant-time comparison"]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Choice(u8);
impl Choice {
pub const FALSE: Choice = Choice(0);
pub const TRUE: Choice = Choice(1);
#[inline]
pub fn from_u8(v: u8) -> Self {
Choice(((v | v.wrapping_neg()) >> 7) & 1)
}
#[inline]
pub fn unwrap_u8(self) -> u8 {
black_box(self.0)
}
#[inline]
pub fn mask(self) -> u8 {
black_box(self.0.wrapping_neg())
}
#[allow(clippy::should_implement_trait)]
#[inline]
pub fn not(self) -> Self {
Choice(self.0 ^ 1)
}
#[inline]
pub fn and(self, other: Self) -> Self {
Choice(self.0 & other.0)
}
#[inline]
pub fn or(self, other: Self) -> Self {
Choice(self.0 | other.0)
}
}
impl From<Choice> for bool {
#[inline]
fn from(c: Choice) -> bool {
c.unwrap_u8() == 1
}
}
#[inline]
pub fn eq(a: &[u8], b: &[u8]) -> Choice {
if a.len() != b.len() {
return Choice::FALSE;
}
let mut acc: u8 = 0;
for i in 0..a.len() {
acc |= a[i] ^ b[i];
}
Choice::from_u8(acc).not()
}
#[inline]
#[must_use = "this is the result of a cryptographic verification; discarding it accepts everything"]
pub fn verify(expected: &[u8], actual: &[u8]) -> bool {
eq(expected, actual).into()
}
#[inline]
pub fn select_u8(c: Choice, a: u8, b: u8) -> u8 {
let m = c.mask();
b ^ (m & (a ^ b))
}
#[inline]
pub fn select_u32(c: Choice, a: u32, b: u32) -> u32 {
let m = (c.unwrap_u8() as u32).wrapping_neg();
b ^ (m & (a ^ b))
}
#[inline]
pub fn select_u64(c: Choice, a: u64, b: u64) -> u64 {
let m = (c.unwrap_u8() as u64).wrapping_neg();
b ^ (m & (a ^ b))
}
#[inline]
pub fn cswap(c: Choice, a: &mut [u8], b: &mut [u8]) {
debug_assert_eq!(a.len(), b.len());
let m = c.mask();
let n = core::cmp::min(a.len(), b.len());
for i in 0..n {
let t = m & (a[i] ^ b[i]);
a[i] ^= t;
b[i] ^= t;
}
}
#[inline]
pub fn cmov(c: Choice, dst: &mut [u8], src: &[u8]) {
debug_assert_eq!(dst.len(), src.len());
let m = c.mask();
let n = core::cmp::min(dst.len(), src.len());
for i in 0..n {
dst[i] ^= m & (dst[i] ^ src[i]);
}
}
pub fn lt_be(a: &[u8], b: &[u8]) -> Choice {
debug_assert_eq!(a.len(), b.len());
let mut borrow: u16 = 0;
for i in (0..a.len()).rev() {
let d = (a[i] as u16).wrapping_sub(b[i] as u16).wrapping_sub(borrow);
borrow = (d >> 8) & 1;
}
Choice::from_u8(borrow as u8)
}
#[inline]
pub fn is_zero(x: &[u8]) -> Choice {
let mut acc = 0u8;
for &b in x {
acc |= b;
}
Choice::from_u8(acc).not()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn eq_matches_semantics() {
assert!(bool::from(eq(b"abc", b"abc")));
assert!(!bool::from(eq(b"abc", b"abd")));
assert!(!bool::from(eq(b"abc", b"ab")));
assert!(bool::from(eq(b"", b"")));
}
#[test]
fn select_picks_correct_branch() {
assert_eq!(select_u8(Choice::TRUE, 0xAA, 0x55), 0xAA);
assert_eq!(select_u8(Choice::FALSE, 0xAA, 0x55), 0x55);
assert_eq!(select_u32(Choice::TRUE, 1, 2), 1);
assert_eq!(select_u64(Choice::FALSE, 1, 2), 2);
}
#[test]
fn cswap_is_conditional() {
let (mut a, mut b) = ([1u8, 2, 3], [4u8, 5, 6]);
cswap(Choice::FALSE, &mut a, &mut b);
assert_eq!((a, b), ([1, 2, 3], [4, 5, 6]));
cswap(Choice::TRUE, &mut a, &mut b);
assert_eq!((a, b), ([4, 5, 6], [1, 2, 3]));
}
#[test]
fn cmov_is_conditional() {
let mut dst = [0u8; 4];
cmov(Choice::FALSE, &mut dst, &[9, 9, 9, 9]);
assert_eq!(dst, [0, 0, 0, 0]);
cmov(Choice::TRUE, &mut dst, &[9, 9, 9, 9]);
assert_eq!(dst, [9, 9, 9, 9]);
}
#[test]
fn lt_be_orders_correctly() {
assert!(bool::from(lt_be(&[0, 1], &[0, 2])));
assert!(!bool::from(lt_be(&[0, 2], &[0, 1])));
assert!(!bool::from(lt_be(&[0, 2], &[0, 2])));
assert!(bool::from(lt_be(&[0x00, 0xFF], &[0x01, 0x00])));
}
#[test]
fn is_zero_detects_all_zero() {
assert!(bool::from(is_zero(&[0, 0, 0])));
assert!(!bool::from(is_zero(&[0, 1, 0])));
}
}