#[inline(never)]
pub fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
let mut diff: u8 = 0;
diff |= u8::from(a.len() != b.len());
let max_len = a.len().max(b.len());
for i in 0..max_len {
let av = a.get(i).copied().unwrap_or(0);
let bv = b.get(i).copied().unwrap_or(0);
diff |= av ^ bv;
core::hint::black_box(&mut diff);
}
diff == 0
}
#[inline(never)]
pub fn constant_time_eq_u32(a: u32, b: u32) -> bool {
let mut diff = a ^ b;
core::hint::black_box(&mut diff);
diff == 0
}
#[inline(never)]
pub fn constant_time_eq_u64(a: u64, b: u64) -> bool {
let mut diff = a ^ b;
core::hint::black_box(&mut diff);
diff == 0
}
#[inline(never)]
pub fn constant_time_eq_u128(a: u128, b: u128) -> bool {
let mut diff = a ^ b;
core::hint::black_box(&mut diff);
diff == 0
}
#[inline(never)]
pub fn constant_time_all_pass(checks: &[bool]) -> bool {
let mut result: u8 = 0;
for &c in checks {
result |= u8::from(!c);
core::hint::black_box(&mut result);
}
result == 0
}
#[inline(never)]
pub fn constant_time_eq_ascii_lower(a: &[u8], lower_b: &[u8]) -> bool {
let mut diff: u8 = 0;
diff |= u8::from(a.len() != lower_b.len());
let max_len = a.len().max(lower_b.len());
for i in 0..max_len {
diff |= a
.get(i)
.copied()
.unwrap_or(0)
.to_ascii_lowercase()
^ lower_b.get(i).copied().unwrap_or(0);
core::hint::black_box(&mut diff);
}
diff == 0
}
#[inline(never)]
pub fn constant_time_eq_case_insensitive(a: &[u8], b: &[u8]) -> bool {
let mut diff: u8 = 0;
diff |= u8::from(a.len() != b.len());
let max_len = a.len().max(b.len());
for i in 0..max_len {
diff |= a
.get(i)
.copied()
.unwrap_or(0)
.to_ascii_lowercase()
^ b.get(i).copied().unwrap_or(0).to_ascii_lowercase();
core::hint::black_box(&mut diff);
}
diff == 0
}
#[inline(never)]
pub fn constant_time_contains(haystack: &[u8], needle: &[u8]) -> bool {
let last = haystack.len().checked_sub(needle.len()).unwrap_or(0);
let too_short = haystack.len() < needle.len();
let mut found: u8 = 0;
for i in 0..=last {
let mut diff: u8 = 0;
for j in 0..needle.len() {
diff |= haystack.get(i + j).copied().unwrap_or(0)
^ needle.get(j).copied().unwrap_or(0);
core::hint::black_box(&mut diff);
}
found |= u8::from(diff == 0);
core::hint::black_box(&mut found);
}
let result = u8::from(!too_short) & found;
result != 0
}
#[inline(never)]
pub fn constant_time_starts_with(haystack: &[u8], prefix: &[u8]) -> bool {
let len_ok = haystack.len() >= prefix.len();
let mut diff: u8 = 0;
for (i, &p) in prefix.iter().enumerate() {
let b = haystack.get(i).copied().unwrap_or(0);
diff |= b ^ p;
core::hint::black_box(&mut diff);
}
let result = u8::from(len_ok) & u8::from(diff == 0);
result == 1
}
#[inline(never)]
pub fn constant_time_contains_case_insensitive(haystack: &[u8], needle: &[u8]) -> bool {
let last = haystack.len().checked_sub(needle.len()).unwrap_or(0);
let too_short = haystack.len() < needle.len();
let mut found: u8 = 0;
for i in 0..=last {
let mut diff: u8 = 0;
for j in 0..needle.len() {
diff |= haystack
.get(i + j)
.copied()
.unwrap_or(0)
.to_ascii_lowercase()
^ needle.get(j).copied().unwrap_or(0).to_ascii_lowercase();
core::hint::black_box(&mut diff);
}
found |= u8::from(diff == 0);
core::hint::black_box(&mut found);
}
let result = u8::from(!too_short) & found;
result != 0
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_constant_time_starts_with_basic() {
assert!(constant_time_starts_with(b"hello world", b"hello"));
assert!(constant_time_starts_with(b"hello", b"hello"));
assert!(constant_time_starts_with(b"hello", b""));
assert!(!constant_time_starts_with(b"hello world", b"world"));
assert!(!constant_time_starts_with(b"hi", b"hello"));
assert!(!constant_time_starts_with(b"hellx", b"hello"));
}
#[test]
fn test_constant_time_contains_case_insensitive_basic() {
assert!(constant_time_contains_case_insensitive(b"Hello World", b"world"));
assert!(constant_time_contains_case_insensitive(b"HELLO WORLD", b"hello"));
assert!(constant_time_contains_case_insensitive(b"abcdef", b""));
assert!(!constant_time_contains_case_insensitive(b"Hello", b"word"));
assert!(!constant_time_contains_case_insensitive(b"ab", b"ABC"));
assert!(constant_time_contains_case_insensitive(b"a\xFFb", b"A\xFF"));
}
#[test]
fn test_constant_time_eq_equal() {
let a = b"secret_key_123";
let b = b"secret_key_123";
assert!(constant_time_eq(a, b));
}
#[test]
fn test_constant_time_eq_different() {
let a = b"secret_key_123";
let b = b"secret_key_124";
assert!(!constant_time_eq(a, b));
}
#[test]
fn test_constant_time_eq_different_length() {
let a = b"short";
let b = b"longer";
assert!(!constant_time_eq(a, b));
}
#[test]
fn test_constant_time_eq_empty() {
let a: &[u8] = b"";
let b: &[u8] = b"";
assert!(constant_time_eq(a, b));
}
#[test]
fn test_constant_time_eq_u32() {
assert!(constant_time_eq_u32(42, 42));
assert!(!constant_time_eq_u32(42, 43));
assert!(constant_time_eq_u32(0, 0));
assert!(constant_time_eq_u32(u32::MAX, u32::MAX));
}
#[test]
fn test_constant_time_eq_u64() {
assert!(constant_time_eq_u64(123456789, 123456789));
assert!(!constant_time_eq_u64(123456789, 123456790));
assert!(constant_time_eq_u64(0, 0));
assert!(constant_time_eq_u64(u64::MAX, u64::MAX));
}
#[test]
fn test_constant_time_eq_u128() {
assert!(constant_time_eq_u128(1, 1));
assert!(!constant_time_eq_u128(1, 2));
assert!(constant_time_eq_u128(u128::MAX, u128::MAX));
}
#[test]
fn test_constant_time_all_pass() {
assert!(constant_time_all_pass(&[true, true, true]));
assert!(!constant_time_all_pass(&[true, false, true]));
assert!(!constant_time_all_pass(&[false, false, false]));
assert!(constant_time_all_pass(&[]));
}
#[test]
fn test_constant_time_eq_all_byte_values() {
for byte in 0u16..=255 {
let val = byte as u8;
let a = [val; 32];
let b = [val; 32];
assert!(constant_time_eq(&a, &b));
let mut c = [val; 32];
if val < 255 {
c[16] = val + 1;
assert!(!constant_time_eq(&a, &c));
}
}
}
#[test]
fn test_constant_time_eq_case_insensitive_equal() {
assert!(constant_time_eq_case_insensitive(b"Content-Type", b"content-type"));
assert!(constant_time_eq_case_insensitive(b"CONTENT-TYPE", b"content-type"));
assert!(constant_time_eq_case_insensitive(b"content-type", b"content-type"));
assert!(constant_time_eq_case_insensitive(b"AbCdEf", b"aBcDeF"));
}
#[test]
fn test_constant_time_eq_case_insensitive_not_equal() {
assert!(!constant_time_eq_case_insensitive(b"content-type", b"content-length"));
assert!(!constant_time_eq_case_insensitive(b"abc", b"abd"));
assert!(!constant_time_eq_case_insensitive(b"abc", b"abcd"));
assert!(!constant_time_eq_case_insensitive(b"", b"a"));
}
#[test]
fn test_constant_time_eq_case_insensitive_non_ascii() {
assert!(constant_time_eq_case_insensitive(b"a\xFFb", b"A\xFFB"));
assert!(!constant_time_eq_case_insensitive(b"a\xFFb", b"A\xFEb"));
assert!(constant_time_eq_case_insensitive(b"", b""));
}
#[test]
fn test_constant_time_contains_basic() {
assert!(constant_time_contains(b"hello world", b"world"));
assert!(constant_time_contains(b"hello world", b"hello"));
assert!(constant_time_contains(b"hello world", b"o w"));
assert!(!constant_time_contains(b"hello world", b"word"));
}
#[test]
fn test_constant_time_contains_boundaries() {
assert!(constant_time_contains(b"abc", b"abc"));
assert!(!constant_time_contains(b"abc", b"abd"));
assert!(!constant_time_contains(b"ab", b"abc"));
assert!(constant_time_contains(b"ab", b""));
assert!(constant_time_contains(b"", b""));
assert!(!constant_time_contains(b"", b"a"));
}
#[test]
fn test_constant_time_contains_all_positions() {
assert!(constant_time_contains(b"aaab", b"ab"));
assert!(constant_time_contains(b"xyz", b"z"));
assert!(!constant_time_contains(b"xyz", b"w"));
}
#[test]
fn test_constant_time_eq_different_length_content_related() {
assert!(!constant_time_eq(b"secret", b"secret_key"));
assert!(!constant_time_eq(b"secret_key", b"secret"));
assert!(!constant_time_eq(b"ab", b"abab"));
assert!(!constant_time_eq(b"abab", b"ab"));
assert!(constant_time_eq(b"", b""));
assert!(!constant_time_eq(b"", b"a"));
assert!(!constant_time_eq(b"a", b""));
}
#[test]
fn test_constant_time_eq_ascii_lower_different_length_content_related() {
assert!(!constant_time_eq_ascii_lower(b"ABC", b"abcd"));
assert!(!constant_time_eq_ascii_lower(b"abcd", b"ABC"));
assert!(!constant_time_eq_ascii_lower(b"abc", b"abcD"));
assert!(!constant_time_eq_ascii_lower(b"", b"a"));
assert!(!constant_time_eq_ascii_lower(b"a", b""));
}
#[test]
fn test_constant_time_eq_case_insensitive_different_length_content_related() {
assert!(!constant_time_eq_case_insensitive(b"AbC", b"aBcD"));
assert!(!constant_time_eq_case_insensitive(b"aBcD", b"AbC"));
assert!(!constant_time_eq_case_insensitive(b"AbC", b"aBcCd"));
assert!(!constant_time_eq_case_insensitive(b"", b"a"));
assert!(!constant_time_eq_case_insensitive(b"a", b""));
}
#[test]
fn test_constant_time_contains_needle_longer_than_haystack() {
assert!(!constant_time_contains(b"abc", b"abcd"));
assert!(!constant_time_contains(b"abc", b"abcabc"));
assert!(!constant_time_contains(b"", b"a"));
assert!(!constant_time_contains(b"a", b"ab"));
assert!(constant_time_contains(b"", b""));
assert!(constant_time_contains(b"abc", b""));
}
#[test]
fn test_constant_time_contains_case_insensitive_needle_longer() {
assert!(!constant_time_contains_case_insensitive(b"ABC", b"AbCd"));
assert!(!constant_time_contains_case_insensitive(b"abc", b"abcabc"));
assert!(!constant_time_contains_case_insensitive(b"", b"A"));
assert!(!constant_time_contains_case_insensitive(b"a", b"AB"));
assert!(constant_time_contains_case_insensitive(b"", b""));
assert!(constant_time_contains_case_insensitive(b"ABC", b""));
}
#[test]
fn test_timing_constant_eq_different_length() {
let a_short = vec![0x5Au8; 64];
let b_short = vec![0x5Bu8; 64];
let a_long_fixed = vec![0x5Au8; 4096];
let b_long_fixed = vec![0x5Bu8; 4096];
let long_prefix = {
let mut v = a_short.clone();
v.resize(4096, 0x5Au8);
v
};
let iters = 200_000u32;
for _ in 0..10_000 {
core::hint::black_box(constant_time_eq(&a_short, &b_short));
}
let t0 = std::time::Instant::now();
for _ in 0..iters {
core::hint::black_box(constant_time_eq(&a_short, &b_short));
}
let short_elapsed = t0.elapsed();
let t1 = std::time::Instant::now();
for _ in 0..iters {
core::hint::black_box(constant_time_eq(&a_long_fixed, &b_long_fixed));
}
let long_fixed_elapsed = t1.elapsed();
assert!(
long_fixed_elapsed >= short_elapsed,
"较长输入应花费不少于较短输入的耗时"
);
let t2 = std::time::Instant::now();
for _ in 0..iters {
core::hint::black_box(constant_time_eq(&long_prefix, &a_long_fixed));
}
let related_elapsed = t2.elapsed();
let lower = long_fixed_elapsed.as_secs_f64() * 0.7;
let upper = long_fixed_elapsed.as_secs_f64() * 1.3;
let got = related_elapsed.as_secs_f64();
assert!(
got >= lower && got <= upper,
"相同长度下比较耗时应基本恒定:related={got:.6}s, fixed={:.6}s",
long_fixed_elapsed.as_secs_f64()
);
let _ = short_elapsed;
}
}