Skip to main content

dryoc/
utils.rs

1use crate::error::{Error, ErrorContext};
2
3/// The ChaCha/Salsa20 "expand 32-byte k" constant, as four little-endian
4/// words. Shared by the ChaCha20 and XSalsa20 stream ciphers and the
5/// HChaCha20/HSalsa20 defaults in [`crate::classic::crypto_core`].
6pub(crate) const SIGMA: [u32; 4] = [0x61707865, 0x3320646e, 0x79622d32, 0x6b206574];
7
8/// Increments `bytes` in constant time, representing a large little-endian
9/// integer; equivalent to `sodium_increment`.
10#[inline]
11pub fn increment_bytes(bytes: &mut [u8]) {
12    let mut carry: u16 = 1;
13    for b in bytes {
14        carry += *b as u16;
15        *b = (carry & 0xff) as u8;
16        carry >>= 8;
17    }
18}
19
20#[inline]
21pub(crate) fn xor_buf(out: &mut [u8], in_: &[u8]) {
22    let len = core::cmp::min(out.len(), in_.len());
23    for i in 0..len {
24        out[i] ^= in_[i];
25    }
26}
27
28#[inline]
29pub(crate) fn load_u64_le(bytes: &[u8]) -> u64 {
30    (bytes[0] as u64)
31        | ((bytes[1] as u64) << 8)
32        | ((bytes[2] as u64) << 16)
33        | ((bytes[3] as u64) << 24)
34        | ((bytes[4] as u64) << 32)
35        | ((bytes[5] as u64) << 40)
36        | ((bytes[6] as u64) << 48)
37        | ((bytes[7] as u64) << 56)
38}
39
40#[inline]
41pub(crate) fn load_u32_le(bytes: &[u8]) -> u32 {
42    (bytes[0] as u32)
43        | ((bytes[1] as u32) << 8)
44        | ((bytes[2] as u32) << 16)
45        | ((bytes[3] as u32) << 24)
46}
47
48#[inline]
49pub(crate) fn pad16(n: usize) -> usize {
50    (0x10 - (n % 16)) & 0xf
51}
52
53/// Splits `bytes` into a fixed-length prefix and the remainder.
54///
55/// Returns [`Error::InvalidLength`] with `context` when `bytes` is shorter
56/// than `len`. Shared by the Rustaceous `from_bytes` constructors, which all
57/// parse an authentication tag, signature, or nonce off one end of a slice.
58pub(crate) fn split_prefix(
59    bytes: &[u8],
60    len: usize,
61    context: ErrorContext,
62) -> Result<(&[u8], &[u8]), Error> {
63    if bytes.len() < len {
64        Err(length_error!(context, bytes.len(), min len))
65    } else {
66        Ok(bytes.split_at(len))
67    }
68}
69
70/// Splits `bytes` into the leading remainder and a fixed-length suffix.
71///
72/// Returns [`Error::InvalidLength`] with `context` when `bytes` is shorter
73/// than `len`. See [`split_prefix`].
74pub(crate) fn split_suffix(
75    bytes: &[u8],
76    len: usize,
77    context: ErrorContext,
78) -> Result<(&[u8], &[u8]), Error> {
79    if bytes.len() < len {
80        Err(length_error!(context, bytes.len(), min len))
81    } else {
82        Ok(bytes.split_at(bytes.len() - len))
83    }
84}
85
86/// Compares `expected` against `computed` in constant time, returning
87/// [`Error::AuthenticationFailed`] on mismatch.
88///
89/// Shared by the Classic verify paths, which all performed this exact
90/// [`subtle::ConstantTimeEq`] comparison inline. Both slices must have the
91/// same length; every caller compares fixed-size tags or hashes.
92pub(crate) fn verify_ct(expected: &[u8], computed: &[u8]) -> Result<(), Error> {
93    use subtle::ConstantTimeEq;
94
95    if expected.ct_eq(computed).unwrap_u8() == 1 {
96        Ok(())
97    } else {
98        Err(Error::AuthenticationFailed)
99    }
100}
101
102/// Compares `a` and `b` in constant time, returning `true` when equal.
103///
104/// Shared by the constant-time [`PartialEq`] impls of the byte-container
105/// types. Both slices must have the same length; every caller compares
106/// fixed-size values.
107pub(crate) fn ct_eq_bytes(a: &[u8], b: &[u8]) -> bool {
108    use subtle::ConstantTimeEq;
109
110    a.ct_eq(b).unwrap_u8() == 1
111}
112
113/// Zeroizes `bytes` with volatile stores like [`zeroize::Zeroize`], but
114/// sixteen bytes at a time. The `zeroize` crate issues one volatile store per
115/// byte, which dominates the cost of small-message operations whose state
116/// buffers are wiped on every call.
117pub(crate) fn zeroize_bytes(bytes: &mut [u8]) {
118    // SAFETY: `u128` has no bit-validity or padding requirements, so viewing
119    // the 16-byte-aligned middle of a `u8` slice as `u128`s is sound;
120    // `align_to_mut` returns disjoint views of `bytes`, with the unaligned
121    // ends left as bytes.
122    let (head, words, tail) = unsafe { bytes.align_to_mut::<u128>() };
123    zeroize_wide(head, words, tail);
124    zeroize::optimization_barrier(bytes);
125}
126
127/// Zeroizes `words` with volatile stores, sixteen bytes at a time where the
128/// alignment allows. Used for large secret working buffers (Argon2 memory),
129/// where the `zeroize` crate's one store per word is measurable.
130pub(crate) fn zeroize_u64s(words: &mut [u64]) {
131    // SAFETY: as in `zeroize_bytes`; `u128` has no validity requirements and
132    // the three views are disjoint.
133    let (head, wide, tail) = unsafe { words.align_to_mut::<u128>() };
134    zeroize_wide(head, wide, tail);
135    zeroize::optimization_barrier(words);
136}
137
138/// Zeroizes `words` like [`zeroize_u64s`]. Used for the ChaCha20 state
139/// wiped on every drop, where one store per word is a measurable part of a
140/// short message.
141pub(crate) fn zeroize_u32s(words: &mut [u32]) {
142    // SAFETY: as in `zeroize_bytes`; neither `u128` nor `u32` has validity
143    // requirements and the three views are disjoint.
144    let (head, wide, tail) = unsafe { words.align_to_mut::<u128>() };
145    zeroize_wide(head, wide, tail);
146    zeroize::optimization_barrier(words);
147}
148
149/// Zeroizes `words` like [`zeroize_u64s`]. Used for ML-KEM's secret
150/// coefficient buffers, where one store per coefficient cost more than a
151/// tenth of an encapsulation.
152pub(crate) fn zeroize_i16s(words: &mut [i16]) {
153    // SAFETY: as in `zeroize_bytes`; neither `u128` nor `i16` has validity
154    // requirements and the three views are disjoint.
155    let (head, wide, tail) = unsafe { words.align_to_mut::<u128>() };
156    zeroize_wide(head, wide, tail);
157    zeroize::optimization_barrier(words);
158}
159
160/// Fixed-size secret buffers that [`WideZeroizing`] wipes with
161/// [`zeroize_bytes`] or [`zeroize_i16s`].
162pub(crate) trait WideZeroize {
163    fn wide_zeroize(&mut self);
164}
165
166impl<const L: usize> WideZeroize for [u8; L] {
167    fn wide_zeroize(&mut self) {
168        zeroize_bytes(self);
169    }
170}
171
172impl<const L: usize, const M: usize> WideZeroize for [[u8; L]; M] {
173    fn wide_zeroize(&mut self) {
174        zeroize_bytes(self.as_flattened_mut());
175    }
176}
177
178impl<const L: usize> WideZeroize for [i16; L] {
179    fn wide_zeroize(&mut self) {
180        zeroize_i16s(self);
181    }
182}
183
184impl<const L: usize, const M: usize> WideZeroize for [[i16; L]; M] {
185    fn wide_zeroize(&mut self) {
186        zeroize_i16s(self.as_flattened_mut());
187    }
188}
189
190/// [`zeroize::Zeroizing`] for [`WideZeroize`] buffers: wiped on drop sixteen
191/// bytes per volatile store instead of one element per store.
192pub(crate) struct WideZeroizing<T: WideZeroize>(T);
193
194impl<T: WideZeroize> WideZeroizing<T> {
195    pub(crate) fn new(value: T) -> Self {
196        Self(value)
197    }
198}
199
200impl<T: WideZeroize> core::ops::Deref for WideZeroizing<T> {
201    type Target = T;
202
203    fn deref(&self) -> &T {
204        &self.0
205    }
206}
207
208impl<T: WideZeroize> core::ops::DerefMut for WideZeroizing<T> {
209    fn deref_mut(&mut self) -> &mut T {
210        &mut self.0
211    }
212}
213
214impl<T: WideZeroize> Drop for WideZeroizing<T> {
215    fn drop(&mut self) {
216        self.0.wide_zeroize();
217    }
218}
219
220/// Clears the three views of one buffer with volatile stores. Callers follow
221/// this with [`zeroize::optimization_barrier`] over the whole buffer, matching
222/// what the `zeroize` crate does after its own volatile writes.
223#[inline]
224fn zeroize_wide<T: zeroize::DefaultIsZeroes>(head: &mut [T], words: &mut [u128], tail: &mut [T]) {
225    use zeroize::Zeroize;
226
227    head.zeroize();
228    // Four stores per iteration: one per iteration spent most of a large
229    // wipe (Argon2's memory) on the loop's own two instructions.
230    let (groups, rest) = words.as_chunks_mut::<4>();
231    for group in groups {
232        for word in group {
233            // SAFETY: `word` is a valid, aligned, exclusively borrowed `u128`.
234            unsafe { core::ptr::write_volatile(word, 0) };
235        }
236    }
237    for word in rest {
238        // SAFETY: as above.
239        unsafe { core::ptr::write_volatile(word, 0) };
240    }
241    tail.zeroize();
242}
243
244#[cfg(test)]
245mod tests {
246    use super::*;
247
248    #[test]
249    fn test_zeroize_u64s_covers_unaligned_ends_and_odd_lengths() {
250        let mut buffer = [0xa5a5_a5a5_a5a5_a5a5u64; 19];
251        for start in 0..3 {
252            for len in [0, 1, 2, 3, 4, 7, 8, 9, 16] {
253                buffer.fill(0xa5a5_a5a5_a5a5_a5a5);
254                zeroize_u64s(&mut buffer[start..start + len]);
255                assert!(
256                    buffer[..start].iter().all(|&w| w == 0xa5a5_a5a5_a5a5_a5a5),
257                    "{start} {len}"
258                );
259                assert!(
260                    buffer[start..start + len].iter().all(|&w| w == 0),
261                    "{start} {len}"
262                );
263                assert!(
264                    buffer[start + len..]
265                        .iter()
266                        .all(|&w| w == 0xa5a5_a5a5_a5a5_a5a5),
267                    "{start} {len}"
268                );
269            }
270        }
271    }
272
273    /// Every start offset within a 16-byte-aligned buffer (so the `u128`
274    /// middle starts at each of the four `u32` positions) and lengths around
275    /// the 16- and 64-byte boundaries: exactly the range is cleared.
276    #[test]
277    fn test_zeroize_u32s_covers_unaligned_ends_and_odd_lengths() {
278        #[repr(align(16))]
279        struct Aligned([u32; 40]);
280        const FILL: u32 = 0xa5a5_a5a5;
281        let mut buffer = Aligned([FILL; 40]);
282        for start in 0..4 {
283            for len in [0, 1, 2, 3, 4, 5, 7, 8, 12, 15, 16, 17, 20, 32] {
284                buffer.0.fill(FILL);
285                zeroize_u32s(&mut buffer.0[start..start + len]);
286                assert!(
287                    buffer.0[..start].iter().all(|&w| w == FILL),
288                    "{start} {len}"
289                );
290                assert!(
291                    buffer.0[start..start + len].iter().all(|&w| w == 0),
292                    "{start} {len}"
293                );
294                assert!(
295                    buffer.0[start + len..].iter().all(|&w| w == FILL),
296                    "{start} {len}"
297                );
298            }
299        }
300    }
301
302    #[test]
303    fn test_zeroize_bytes_covers_unaligned_ends_and_odd_lengths() {
304        let mut buffer = [0xa5u8; 71];
305        for start in 0..9 {
306            for len in [0, 1, 7, 8, 9, 15, 16, 17, 31, 40, 62] {
307                buffer.fill(0xa5);
308                zeroize_bytes(&mut buffer[start..start + len]);
309                assert!(buffer[..start].iter().all(|&b| b == 0xa5), "{start} {len}");
310                assert!(
311                    buffer[start..start + len].iter().all(|&b| b == 0),
312                    "{start} {len}"
313                );
314                assert!(
315                    buffer[start + len..].iter().all(|&b| b == 0xa5),
316                    "{start} {len}"
317                );
318            }
319        }
320    }
321
322    #[test]
323    fn test_zeroize_i16s_covers_unaligned_ends_and_odd_lengths() {
324        let mut buffer = [0x5a5au16 as i16; 43];
325        for start in 0..9 {
326            for len in [0, 1, 7, 8, 9, 15, 16, 17, 34] {
327                buffer.fill(0x5a5a);
328                zeroize_i16s(&mut buffer[start..start + len]);
329                assert!(
330                    buffer[..start].iter().all(|&w| w == 0x5a5a),
331                    "{start} {len}"
332                );
333                assert!(
334                    buffer[start..start + len].iter().all(|&w| w == 0),
335                    "{start} {len}"
336                );
337                assert!(
338                    buffer[start + len..].iter().all(|&w| w == 0x5a5a),
339                    "{start} {len}"
340                );
341            }
342        }
343    }
344
345    /// `sodium_increment` vectors: little-endian, the carry propagates
346    /// through every `0xff` byte, and the all-ones value wraps to zero.
347    const INCREMENT_VECTORS: &[(&[u8], &[u8])] = &[
348        (&[], &[]),
349        (&[0], &[1]),
350        (&[1], &[2]),
351        (&[0xff], &[0]),
352        (&[0xff, 0], &[0, 1]),
353        (&[0x00, 0xff], &[0x01, 0xff]),
354        (&[0xff, 0xff, 0x00], &[0, 0, 1]),
355        (
356            &[0xfe, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff],
357            &[0xff; 8],
358        ),
359        (&[0xff; 8], &[0; 8]),
360        (&[0xff; 24], &[0; 24]),
361    ];
362
363    #[test]
364    fn test_increment_bytes() {
365        for (input, expected) in INCREMENT_VECTORS {
366            let mut bytes = input.to_vec();
367            increment_bytes(&mut bytes);
368            assert_eq!(bytes.as_slice(), *expected, "increment of {input:02x?}");
369        }
370
371        let mut b = [0xff, 0];
372        increment_bytes(&mut b);
373        assert_eq!(b, [0, 1]);
374        increment_bytes(&mut b);
375        assert_eq!(b, [1, 1]);
376        increment_bytes(&mut b);
377        assert_eq!(b, [2, 1]);
378    }
379
380    #[test]
381    fn test_xor_buf() {
382        let mut a = [0];
383        let b = [0];
384
385        xor_buf(&mut a, &b);
386        assert_eq!([0], a);
387
388        let mut a = [1];
389        let b = [0];
390
391        xor_buf(&mut a, &b);
392        assert_eq!([1], a);
393
394        let mut a = [1, 1, 1];
395        let b = [0];
396
397        xor_buf(&mut a, &b);
398        assert_eq!([1, 1, 1], a);
399
400        let mut a = [1, 1, 1];
401        let b = [0, 1, 1];
402
403        xor_buf(&mut a, &b);
404        assert_eq!([1, 0, 0], a);
405    }
406
407    #[test]
408    fn test_pad16() {
409        assert_eq!(pad16(0), 0);
410        assert_eq!(pad16(1), 15);
411        assert_eq!(pad16(2), 14);
412        assert_eq!(pad16(15), 1);
413        assert_eq!(pad16(16), 0);
414        assert_eq!(pad16(17), 15);
415        assert_eq!(pad16(32), 0);
416        assert_eq!(pad16(33), 15);
417    }
418
419    #[cfg(dryoc_native_tests)]
420    mod native_tests {
421        use super::*;
422
423        #[test]
424        fn test_increment_bytes_matches_libsodium() {
425            use libsodium_sys::sodium_increment as so_sodium_increment;
426
427            use crate::utils::test_util::XorShift64;
428
429            crate::native_test_util::init();
430
431            fn assert_matches_libsodium(input: &[u8]) {
432                let mut ours = input.to_vec();
433                let mut theirs = input.to_vec();
434                increment_bytes(&mut ours);
435                // SAFETY: `theirs` is a valid, writable buffer of exactly
436                // `theirs.len()` bytes for the duration of the call.
437                unsafe { so_sodium_increment(theirs.as_mut_ptr(), theirs.len()) };
438                assert_eq!(ours, theirs, "input {input:02x?}");
439            }
440
441            for (input, _) in INCREMENT_VECTORS {
442                assert_matches_libsodium(input);
443            }
444
445            let mut rng = XorShift64::new(0x9e37_79b9_7f4a_7c15);
446            for len in 0..=64 {
447                let mut data = vec![0u8; len];
448                for b in &mut data {
449                    *b = rng.next_u64() as u8;
450                }
451                assert_matches_libsodium(&data);
452                assert_matches_libsodium(&vec![0xff; len]);
453
454                // Carry chain that stops at a non-`0xff` final byte.
455                if let Some((last, head)) = data.split_last_mut() {
456                    head.fill(0xff);
457                    *last &= 0x7f;
458                    assert_matches_libsodium(&data);
459                }
460            }
461        }
462    }
463}
464
465/// Helpers shared by the unit tests.
466#[cfg(test)]
467pub(crate) mod test_util {
468    use crate::error::{Error, ErrorContext, LengthConstraint};
469    use crate::test_prelude::*;
470
471    /// Bounds Miri runs and disables filesystem-backed failure persistence.
472    #[cfg(not(all(target_arch = "wasm32", target_os = "unknown")))]
473    pub(crate) fn proptest_config(cases: u32) -> proptest::test_runner::Config {
474        let mut config = proptest::test_runner::Config::with_cases(cases);
475        if cfg!(miri) {
476            config.cases = 8;
477            config.failure_persistence = None;
478        }
479        config
480    }
481
482    /// Asserts that `result` is `Error::InvalidLength` for a slice of
483    /// `actual` bytes where exactly `expected` were required, matching on the
484    /// variant rather than its message.
485    pub(crate) fn assert_exact_slice_length_error<T>(
486        result: Result<T, Error>,
487        actual: usize,
488        expected: usize,
489    ) {
490        match result {
491            Err(Error::InvalidLength {
492                context,
493                actual: got,
494                constraint,
495            }) => {
496                assert_eq!(context, ErrorContext::Slice);
497                assert_eq!(got, actual);
498                assert_eq!(constraint, LengthConstraint::Exact(expected));
499            }
500            Err(other) => panic!("unexpected error {other:?}"),
501            Ok(_) => panic!("length {actual} accepted where exactly {expected} is required"),
502        }
503    }
504
505    /// Decodes a hexadecimal string into bytes, ignoring embedded ASCII spaces.
506    pub(crate) fn hex(s: &str) -> Vec<u8> {
507        hex::decode(s.replace(' ', "")).expect("hex")
508    }
509
510    /// Decodes a hexadecimal string into an exact-length byte array.
511    pub(crate) fn hex_array<const N: usize>(s: &str) -> [u8; N] {
512        hex(s).try_into().expect("hex array length")
513    }
514
515    /// Deterministic xorshift64 generator for reproducible random test inputs.
516    pub(crate) struct XorShift64(u64);
517
518    impl XorShift64 {
519        pub(crate) fn new(seed: u64) -> Self {
520            Self(seed)
521        }
522
523        pub(crate) fn next_u64(&mut self) -> u64 {
524            self.0 ^= self.0 << 13;
525            self.0 ^= self.0 >> 7;
526            self.0 ^= self.0 << 17;
527            self.0
528        }
529
530        /// Four successive outputs, little-endian, as 32 bytes.
531        pub(crate) fn next_bytes32(&mut self) -> [u8; 32] {
532            let mut bytes = [0u8; 32];
533            for chunk in bytes.chunks_mut(8) {
534                chunk.copy_from_slice(&self.next_u64().to_le_bytes());
535            }
536            bytes
537        }
538    }
539}