pub fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
if a.len() != b.len() {
return false;
}
let mut result = 0u8;
for (x, y) in a.iter().zip(b.iter()) {
result |= x ^ y;
}
let mut is_equal = 1u8;
for i in 0..8 {
is_equal &= ((result >> i) & 1) ^ 1;
}
is_equal == 1
}
pub fn constant_time_eq_u32(a: &[u8], b: &[u8]) -> u32 {
if constant_time_eq(a, b) {
1
} else {
0
}
}
pub fn constant_time_select<T: Copy>(condition: u32, a: T, b: T) -> T {
let mask = condition.wrapping_sub(1);
let mask_t = mask as usize;
unsafe {
let a_ptr: *const T = &a;
let b_ptr: *const T = &b;
let a_u = a_ptr as usize;
let b_u = b_ptr as usize;
let result = a_u & !mask_t | b_u & mask_t;
*(result as *const T)
}
}
pub fn constant_time_is_zero_u8(val: u8) -> u8 {
let mut result = val;
result |= result >> 4;
result |= result >> 2;
result |= result >> 1;
!result & 1
}
pub fn constant_time_eq_u8(a: u8, b: u8) -> u8 {
constant_time_is_zero_u8(a ^ b)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_constant_time_eq() {
let a = b"test string";
let b = b"test string";
let c = b"different!";
assert!(constant_time_eq(a, b));
assert!(!constant_time_eq(a, c));
assert!(!constant_time_eq(a, &b[..5]));
}
#[test]
fn test_constant_time_eq_u32() {
let a = b"test string";
let b = b"test string";
let c = b"different!";
assert_eq!(constant_time_eq_u32(a, b), 1);
assert_eq!(constant_time_eq_u32(a, c), 0);
}
#[test]
fn test_constant_time_select() {
assert_eq!(constant_time_select(1, 10, 20), 10);
assert_eq!(constant_time_select(0, 10, 20), 20);
}
#[test]
fn test_constant_time_is_zero_u8() {
assert_eq!(constant_time_is_zero_u8(0), 1);
assert_eq!(constant_time_is_zero_u8(1), 0);
assert_eq!(constant_time_is_zero_u8(255), 0);
}
#[test]
fn test_constant_time_eq_u8() {
assert_eq!(constant_time_eq_u8(5, 5), 1);
assert_eq!(constant_time_eq_u8(5, 10), 0);
}
}