Skip to main content

zenith_foundation/
ct_compare.rs

1//! 恒定时间比较工具
2//!
3//! 防止时序侧信道攻击(Timing Side-Channel Attack),
4//! 确保安全关键路径的比较操作不受数据内容影响执行时间。
5//!
6//! # 攻击原理
7//! 普通比较操作在遇到首个不匹配字节时会提前返回,
8//! 攻击者通过测量比较耗时可逐字节推断密钥内容。
9//!
10//! # 防御策略
11//! 恒定时间比较遍历所有字节,使用位操作累积差异,
12//! 不进行提前返回,使执行时间仅取决于输入长度。
13
14/// 恒定时间字节切片比较
15///
16/// 遍历所有字节,使用 XOR + OR 累积差异位,
17/// 不进行提前返回。
18///
19/// # Arguments
20/// * `a` - 第一个字节切片
21/// * `b` - 第二个字节切片
22///
23/// # Returns
24/// * `bool` - 是否相等
25#[inline(never)]
26pub fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
27    let mut diff: u8 = 0;
28    // 长度差异并入 diff,不提前返回;长度不等时结果恒为 false
29    diff |= u8::from(a.len() != b.len());
30    // 遍历较长一方的长度,越界字节按 0 参与比较,
31    // 使耗时仅取决于较长的长度,与内容、长度差无关。
32    let max_len = a.len().max(b.len());
33    for i in 0..max_len {
34        let av = a.get(i).copied().unwrap_or(0);
35        let bv = b.get(i).copied().unwrap_or(0);
36        diff |= av ^ bv;
37        // 阻止编译器优化破坏恒定时间性质(防止 diff 被消除)
38        core::hint::black_box(&mut diff);
39    }
40    diff == 0
41}
42
43/// 恒定时间 u32 比较
44///
45/// # Arguments
46/// * `a` - 第一个值
47/// * `b` - 第二个值
48///
49/// # Returns
50/// * `bool` - 是否相等
51#[inline(never)]
52pub fn constant_time_eq_u32(a: u32, b: u32) -> bool {
53    // 使用 XOR 累积差异,避免提前返回
54    let mut diff = a ^ b;
55    core::hint::black_box(&mut diff);
56    diff == 0
57}
58
59/// 恒定时间 u64 比较
60///
61/// # Arguments
62/// * `a` - 第一个值
63/// * `b` - 第二个值
64///
65/// # Returns
66/// * `bool` - 是否相等
67#[inline(never)]
68pub fn constant_time_eq_u64(a: u64, b: u64) -> bool {
69    let mut diff = a ^ b;
70    core::hint::black_box(&mut diff);
71    diff == 0
72}
73
74/// 恒定时间 u128 比较
75///
76/// # Arguments
77/// * `a` - 第一个值
78/// * `b` - 第二个值
79///
80/// # Returns
81/// * `bool` - 是否相等
82#[inline(never)]
83pub fn constant_time_eq_u128(a: u128, b: u128) -> bool {
84    let mut diff = a ^ b;
85    core::hint::black_box(&mut diff);
86    diff == 0
87}
88
89/// 恒定时间布尔累积 OR(用于多项验证)
90///
91/// 即使某项失败也继续检查后续项,使所有验证
92/// 总耗时保持一致。
93///
94/// # Arguments
95/// * `checks` - 布尔检查结果数组
96///
97/// # Returns
98/// * `bool` - 是否全部通过
99#[inline(never)]
100pub fn constant_time_all_pass(checks: &[bool]) -> bool {
101    let mut result: u8 = 0;
102    for &c in checks {
103        // u8::from(!c) 由布尔直接转整数,无数据相关分支,语义同 if c {0} else {1}
104        result |= u8::from(!c);
105        core::hint::black_box(&mut result);
106    }
107    result == 0
108}
109
110/// 恒定时间 ASCII 大小写不敏感比较
111///
112/// 对两个字节切片做恒定时间、不提前退出的大小写不敏感比较。
113/// 用于检测公开命名的敏感模式(如 `Password=` / `Authorization: Bearer `)时,
114/// 避免因提前退出导致的模式存在性侧信道泄露。
115///
116/// # Arguments
117/// * `a` - 待比较的第一个字节切片(任意大小写)
118/// * `lower_b` - 待比较的第二个字节切片(必须为 ASCII 小写)
119///
120/// # Safety
121/// 调用者需保证 `lower_b` 是已经转换为小写的 ASCII;若 `lower_b` 含大写,
122/// 结果恒为 false,但执行时间依然恒定。
123#[inline(never)]
124pub fn constant_time_eq_ascii_lower(a: &[u8], lower_b: &[u8]) -> bool {
125    let mut diff: u8 = 0;
126    // 长度差异并入 diff,不提前返回;长度不等时结果恒为 false
127    diff |= u8::from(a.len() != lower_b.len());
128    let max_len = a.len().max(lower_b.len());
129    for i in 0..max_len {
130        diff |= a
131            .get(i)
132            .copied()
133            .unwrap_or(0)
134            .to_ascii_lowercase()
135            ^ lower_b.get(i).copied().unwrap_or(0);
136        core::hint::black_box(&mut diff);
137    }
138    diff == 0
139}
140
141/// 恒定时间 ASCII 大小写不敏感比较(双侧任意大小写)
142///
143/// 与 [`constant_time_eq_ascii_lower`] 语义相同,但两侧输入均允许任意大小写,
144/// 比较前在寄存器内逐字节归一化为小写,无堆分配、不提前退出。
145///
146/// # Arguments
147/// * `a` - 第一个字节切片(任意大小写)
148/// * `b` - 第二个字节切片(任意大小写)
149///
150/// # Returns
151/// * `bool` - 忽略 ASCII 大小写后是否相等
152#[inline(never)]
153pub fn constant_time_eq_case_insensitive(a: &[u8], b: &[u8]) -> bool {
154    let mut diff: u8 = 0;
155    // 长度差异并入 diff,不提前返回;长度不等时结果恒为 false
156    diff |= u8::from(a.len() != b.len());
157    let max_len = a.len().max(b.len());
158    for i in 0..max_len {
159        diff |= a
160            .get(i)
161            .copied()
162            .unwrap_or(0)
163            .to_ascii_lowercase()
164            ^ b.get(i).copied().unwrap_or(0).to_ascii_lowercase();
165        core::hint::black_box(&mut diff);
166    }
167    diff == 0
168}
169
170/// 恒定时间子串匹配(AGENT.md §4.9 时序安全)
171///
172/// 遍历所有可能的对齐位置,逐字节 XOR 累积差异并用位运算聚合结果,
173/// 执行时间仅取决于输入长度,与是否命中、命中位置无关:
174/// - 不因首个不匹配字节提前退出内层循环
175/// - 不因已命中提前退出外层循环
176/// - 无数据相关分支(匹配结果经位运算聚合,而非条件跳转)
177///
178/// 长度检查仅依赖输入长度(长度非秘密),不泄露内容信息。
179///
180/// # Arguments
181/// * `haystack` - 被搜索的字节串
182/// * `needle` - 待匹配的模式串(空模式恒为命中)
183///
184/// # Returns
185/// * `bool` - `haystack` 中是否包含 `needle`
186#[inline(never)]
187pub fn constant_time_contains(haystack: &[u8], needle: &[u8]) -> bool {
188    // 空 needle 恒命中,但走完整外层循环(内层 0 次),不提前返回;
189    // 空 needle 长度恒为 0(非秘密),完整循环耗时与内容无关。
190    // haystack 短于 needle 时必然不包含:将 last 置 0 仍做一次对齐比较,
191    // 并累加长度差异,保证耗时与内容无关、结果恒为 false。
192    let last = haystack.len().checked_sub(needle.len()).unwrap_or(0);
193    let too_short = haystack.len() < needle.len();
194    let mut found: u8 = 0;
195    for i in 0..=last {
196        // 逐字节比较,不因首个不匹配字节而退出
197        let mut diff: u8 = 0;
198        for j in 0..needle.len() {
199            // 越界字节按 0 参与比较(结果已被 too_short 否决,不影响正确性)
200            diff |= haystack.get(i + j).copied().unwrap_or(0)
201                ^ needle.get(j).copied().unwrap_or(0);
202            core::hint::black_box(&mut diff);
203        }
204        found |= u8::from(diff == 0);
205        core::hint::black_box(&mut found);
206    }
207    // 长度不足时强制为 false;用位运算聚合,避免短路求值泄露时序信息
208    let result = u8::from(!too_short) & found;
209    result != 0
210}
211
212/// 恒定时间前缀比较(AGENT.md §4.9 时序安全)
213///
214/// 逐字节 XOR 累积差异,不因首个不匹配字节提前退出;
215/// `haystack` 长度不足时仍完成全长比较后返回 `false`
216/// (长度检查仅依赖输入长度,长度本身非秘密)。
217///
218/// # Arguments
219/// * `haystack` - 被检查的字节串
220/// * `prefix` - 待匹配的前缀(空前缀恒为命中,与 `str::starts_with` 语义一致)
221///
222/// # Returns
223/// * `bool` - `haystack` 是否以 `prefix` 开头
224#[inline(never)]
225pub fn constant_time_starts_with(haystack: &[u8], prefix: &[u8]) -> bool {
226    // 长度不足必然不匹配(仅依赖长度,长度非秘密),但仍完成全长比较
227    let len_ok = haystack.len() >= prefix.len();
228    let mut diff: u8 = 0;
229    for (i, &p) in prefix.iter().enumerate() {
230        // 越界字节以 0 参与比较(结果已被 len_ok 否决,不影响正确性)
231        let b = haystack.get(i).copied().unwrap_or(0);
232        diff |= b ^ p;
233        core::hint::black_box(&mut diff);
234    }
235    // 使用位运算聚合结果,避免短路求值泄露时序信息
236    let result = u8::from(len_ok) & u8::from(diff == 0);
237    result == 1
238}
239
240/// 恒定时间 ASCII 大小写不敏感子串匹配(AGENT.md §4.9 时序安全)
241///
242/// 与 [`constant_time_contains`] 同一安全语义:
243/// 遍历全部对齐位置、逐字节 XOR 累积差异、无数据相关分支、不提前退出,
244/// 执行时间仅取决于输入长度,与是否命中、命中位置无关。
245/// 比较前在寄存器内逐字节 ASCII 折叠为小写,零堆分配
246/// (替代热路径 `to_lowercase()/to_uppercase()` 的每请求堆分配)。
247///
248/// # Arguments
249/// * `haystack` - 被搜索的字节串(任意大小写)
250/// * `needle` - 待匹配的模式串(任意大小写;空模式恒为命中)
251///
252/// # Returns
253/// * `bool` - 忽略 ASCII 大小写后,`haystack` 中是否包含 `needle`
254#[inline(never)]
255pub fn constant_time_contains_case_insensitive(haystack: &[u8], needle: &[u8]) -> bool {
256    // 空 needle 恒命中,但走完整外层循环(内层 0 次),不提前返回;
257    // haystack 短于 needle 时必然不包含:将 last 置 0 仍做一次对齐比较,
258    // 并累加长度差异,保证耗时与内容无关、结果恒为 false。
259    let last = haystack.len().checked_sub(needle.len()).unwrap_or(0);
260    let too_short = haystack.len() < needle.len();
261    let mut found: u8 = 0;
262    for i in 0..=last {
263        // 逐字节 ASCII 折叠后比较,不因首个不匹配字节而退出
264        let mut diff: u8 = 0;
265        for j in 0..needle.len() {
266            // 越界字节按 0 参与比较(结果已被 too_short 否决,不影响正确性)
267            diff |= haystack
268                .get(i + j)
269                .copied()
270                .unwrap_or(0)
271                .to_ascii_lowercase()
272                ^ needle.get(j).copied().unwrap_or(0).to_ascii_lowercase();
273            core::hint::black_box(&mut diff);
274        }
275        found |= u8::from(diff == 0);
276        core::hint::black_box(&mut found);
277    }
278    // 长度不足时强制为 false;用位运算聚合,避免短路求值泄露时序信息
279    let result = u8::from(!too_short) & found;
280    result != 0
281}
282
283#[cfg(test)]
284mod tests {
285    use super::*;
286
287    #[test]
288    fn test_constant_time_starts_with_basic() {
289        assert!(constant_time_starts_with(b"hello world", b"hello"));
290        assert!(constant_time_starts_with(b"hello", b"hello"));
291        assert!(constant_time_starts_with(b"hello", b""));
292        assert!(!constant_time_starts_with(b"hello world", b"world"));
293        assert!(!constant_time_starts_with(b"hi", b"hello"));
294        assert!(!constant_time_starts_with(b"hellx", b"hello"));
295    }
296
297    #[test]
298    fn test_constant_time_contains_case_insensitive_basic() {
299        assert!(constant_time_contains_case_insensitive(b"Hello World", b"world"));
300        assert!(constant_time_contains_case_insensitive(b"HELLO WORLD", b"hello"));
301        assert!(constant_time_contains_case_insensitive(b"abcdef", b""));
302        assert!(!constant_time_contains_case_insensitive(b"Hello", b"word"));
303        assert!(!constant_time_contains_case_insensitive(b"ab", b"ABC"));
304        // 非 ASCII 字节按原值折叠(to_ascii_lowercase 不影响非字母)
305        assert!(constant_time_contains_case_insensitive(b"a\xFFb", b"A\xFF"));
306    }
307
308    #[test]
309    fn test_constant_time_eq_equal() {
310        let a = b"secret_key_123";
311        let b = b"secret_key_123";
312        assert!(constant_time_eq(a, b));
313    }
314
315    #[test]
316    fn test_constant_time_eq_different() {
317        let a = b"secret_key_123";
318        let b = b"secret_key_124";
319        assert!(!constant_time_eq(a, b));
320    }
321
322    #[test]
323    fn test_constant_time_eq_different_length() {
324        let a = b"short";
325        let b = b"longer";
326        assert!(!constant_time_eq(a, b));
327    }
328
329    #[test]
330    fn test_constant_time_eq_empty() {
331        let a: &[u8] = b"";
332        let b: &[u8] = b"";
333        assert!(constant_time_eq(a, b));
334    }
335
336    #[test]
337    fn test_constant_time_eq_u32() {
338        assert!(constant_time_eq_u32(42, 42));
339        assert!(!constant_time_eq_u32(42, 43));
340        assert!(constant_time_eq_u32(0, 0));
341        assert!(constant_time_eq_u32(u32::MAX, u32::MAX));
342    }
343
344    #[test]
345    fn test_constant_time_eq_u64() {
346        assert!(constant_time_eq_u64(123456789, 123456789));
347        assert!(!constant_time_eq_u64(123456789, 123456790));
348        assert!(constant_time_eq_u64(0, 0));
349        assert!(constant_time_eq_u64(u64::MAX, u64::MAX));
350    }
351
352    #[test]
353    fn test_constant_time_eq_u128() {
354        assert!(constant_time_eq_u128(1, 1));
355        assert!(!constant_time_eq_u128(1, 2));
356        assert!(constant_time_eq_u128(u128::MAX, u128::MAX));
357    }
358
359    #[test]
360    fn test_constant_time_all_pass() {
361        assert!(constant_time_all_pass(&[true, true, true]));
362        assert!(!constant_time_all_pass(&[true, false, true]));
363        assert!(!constant_time_all_pass(&[false, false, false]));
364        assert!(constant_time_all_pass(&[]));
365    }
366
367    #[test]
368    fn test_constant_time_eq_all_byte_values() {
369        for byte in 0u16..=255 {
370            let val = byte as u8;
371            let a = [val; 32];
372            let b = [val; 32];
373            assert!(constant_time_eq(&a, &b));
374
375            let mut c = [val; 32];
376            if val < 255 {
377                c[16] = val + 1;
378                assert!(!constant_time_eq(&a, &c));
379            }
380        }
381    }
382
383    // ===== constant_time_eq_case_insensitive 测试 =====
384
385    #[test]
386    fn test_constant_time_eq_case_insensitive_equal() {
387        assert!(constant_time_eq_case_insensitive(b"Content-Type", b"content-type"));
388        assert!(constant_time_eq_case_insensitive(b"CONTENT-TYPE", b"content-type"));
389        assert!(constant_time_eq_case_insensitive(b"content-type", b"content-type"));
390        assert!(constant_time_eq_case_insensitive(b"AbCdEf", b"aBcDeF"));
391    }
392
393    #[test]
394    fn test_constant_time_eq_case_insensitive_not_equal() {
395        assert!(!constant_time_eq_case_insensitive(b"content-type", b"content-length"));
396        assert!(!constant_time_eq_case_insensitive(b"abc", b"abd"));
397        // 长度不同恒为 false
398        assert!(!constant_time_eq_case_insensitive(b"abc", b"abcd"));
399        assert!(!constant_time_eq_case_insensitive(b"", b"a"));
400    }
401
402    #[test]
403    fn test_constant_time_eq_case_insensitive_non_ascii() {
404        // 非 ASCII 字节按原值比较(to_ascii_lowercase 不影响非字母)
405        assert!(constant_time_eq_case_insensitive(b"a\xFFb", b"A\xFFB"));
406        assert!(!constant_time_eq_case_insensitive(b"a\xFFb", b"A\xFEb"));
407        // 空切片相等
408        assert!(constant_time_eq_case_insensitive(b"", b""));
409    }
410
411    // ===== constant_time_contains 测试 =====
412
413    #[test]
414    fn test_constant_time_contains_basic() {
415        assert!(constant_time_contains(b"hello world", b"world"));
416        assert!(constant_time_contains(b"hello world", b"hello"));
417        assert!(constant_time_contains(b"hello world", b"o w"));
418        assert!(!constant_time_contains(b"hello world", b"word"));
419    }
420
421    #[test]
422    fn test_constant_time_contains_boundaries() {
423        // 完全相等
424        assert!(constant_time_contains(b"abc", b"abc"));
425        assert!(!constant_time_contains(b"abc", b"abd"));
426        // needle 长于 haystack
427        assert!(!constant_time_contains(b"ab", b"abc"));
428        // 空 needle 恒命中;空 haystack 仅容纳空 needle
429        assert!(constant_time_contains(b"ab", b""));
430        assert!(constant_time_contains(b"", b""));
431        assert!(!constant_time_contains(b"", b"a"));
432    }
433
434    #[test]
435    fn test_constant_time_contains_all_positions() {
436        // 命中在末尾位置(最后一个对齐点)
437        assert!(constant_time_contains(b"aaab", b"ab"));
438        // 单字节 needle 遍历所有位置
439        assert!(constant_time_contains(b"xyz", b"z"));
440        assert!(!constant_time_contains(b"xyz", b"w"));
441    }
442
443    // ===== 长度不同但内容相关时结果正确(false) =====
444
445    #[test]
446    fn test_constant_time_eq_different_length_content_related() {
447        // a 是 b 的前缀,仅长度不同,结果必须为 false
448        assert!(!constant_time_eq(b"secret", b"secret_key"));
449        assert!(!constant_time_eq(b"secret_key", b"secret"));
450        // 较长者是较短者的重复/扩展,仍因长度不同而为 false
451        assert!(!constant_time_eq(b"ab", b"abab"));
452        assert!(!constant_time_eq(b"abab", b"ab"));
453        // 空与空的特殊情况
454        assert!(constant_time_eq(b"", b""));
455        assert!(!constant_time_eq(b"", b"a"));
456        assert!(!constant_time_eq(b"a", b""));
457    }
458
459    #[test]
460    fn test_constant_time_eq_ascii_lower_different_length_content_related() {
461        assert!(!constant_time_eq_ascii_lower(b"ABC", b"abcd"));
462        assert!(!constant_time_eq_ascii_lower(b"abcd", b"ABC"));
463        assert!(!constant_time_eq_ascii_lower(b"abc", b"abcD"));
464        assert!(!constant_time_eq_ascii_lower(b"", b"a"));
465        assert!(!constant_time_eq_ascii_lower(b"a", b""));
466    }
467
468    #[test]
469    fn test_constant_time_eq_case_insensitive_different_length_content_related() {
470        assert!(!constant_time_eq_case_insensitive(b"AbC", b"aBcD"));
471        assert!(!constant_time_eq_case_insensitive(b"aBcD", b"AbC"));
472        assert!(!constant_time_eq_case_insensitive(b"AbC", b"aBcCd"));
473        assert!(!constant_time_eq_case_insensitive(b"", b"a"));
474        assert!(!constant_time_eq_case_insensitive(b"a", b""));
475    }
476
477    #[test]
478    fn test_constant_time_contains_needle_longer_than_haystack() {
479        // needle 长于 haystack 时恒为 false,即使内容高度相关
480        assert!(!constant_time_contains(b"abc", b"abcd"));
481        assert!(!constant_time_contains(b"abc", b"abcabc"));
482        assert!(!constant_time_contains(b"", b"a"));
483        assert!(!constant_time_contains(b"a", b"ab"));
484        // 空 needle 恒命中
485        assert!(constant_time_contains(b"", b""));
486        assert!(constant_time_contains(b"abc", b""));
487    }
488
489    #[test]
490    fn test_constant_time_contains_case_insensitive_needle_longer() {
491        assert!(!constant_time_contains_case_insensitive(b"ABC", b"AbCd"));
492        assert!(!constant_time_contains_case_insensitive(b"abc", b"abcabc"));
493        assert!(!constant_time_contains_case_insensitive(b"", b"A"));
494        assert!(!constant_time_contains_case_insensitive(b"a", b"AB"));
495        assert!(constant_time_contains_case_insensitive(b"", b""));
496        assert!(constant_time_contains_case_insensitive(b"ABC", b""));
497    }
498
499    // ===== 形式化验证恒定性(说明性,避免 flaky) =====
500
501    #[test]
502    fn test_timing_constant_eq_different_length() {
503        // 说明性测试:测量不同长度输入的比较耗时,断言耗时差异在容忍范围内。
504        // constant_time_eq 现在遍历 max(a.len(), b.len()) 次,耗时与长度差无关。
505        let a_short = vec![0x5Au8; 64];
506        let b_short = vec![0x5Bu8; 64];
507        let a_long_fixed = vec![0x5Au8; 4096];
508        let b_long_fixed = vec![0x5Bu8; 4096];
509        // 长度差很大但内容高度相关(长串是短串的前缀)
510        let long_prefix = {
511            let mut v = a_short.clone();
512            v.resize(4096, 0x5Au8);
513            v
514        };
515
516        let iters = 200_000u32;
517
518        // 预热(触发分支预测/缓存)
519        for _ in 0..10_000 {
520            core::hint::black_box(constant_time_eq(&a_short, &b_short));
521        }
522
523        let t0 = std::time::Instant::now();
524        for _ in 0..iters {
525            core::hint::black_box(constant_time_eq(&a_short, &b_short));
526        }
527        let short_elapsed = t0.elapsed();
528
529        let t1 = std::time::Instant::now();
530        for _ in 0..iters {
531            core::hint::black_box(constant_time_eq(&a_long_fixed, &b_long_fixed));
532        }
533        // 较长输入的耗时必然更大(字节数更多),此项仅作 sanity 检查
534        let long_fixed_elapsed = t1.elapsed();
535        assert!(
536            long_fixed_elapsed >= short_elapsed,
537            "较长输入应花费不少于较短输入的耗时"
538        );
539
540        // 关键断言:同为 4096 字节、仅内容不同(相关 vs 不相关)时耗时应接近
541        let t2 = std::time::Instant::now();
542        for _ in 0..iters {
543            core::hint::black_box(constant_time_eq(&long_prefix, &a_long_fixed));
544        }
545        let related_elapsed = t2.elapsed();
546
547        // 相同长度下,内容相关与否耗时应当几乎一致(恒定时间核心性质)。
548        // 使用宽松阈值(30%)避免 flaky。
549        let lower = long_fixed_elapsed.as_secs_f64() * 0.7;
550        let upper = long_fixed_elapsed.as_secs_f64() * 1.3;
551        let got = related_elapsed.as_secs_f64();
552        assert!(
553            got >= lower && got <= upper,
554            "相同长度下比较耗时应基本恒定:related={got:.6}s, fixed={:.6}s",
555            long_fixed_elapsed.as_secs_f64()
556        );
557        // 避免 unused 警告
558        let _ = short_elapsed;
559    }
560}