1use crate::error::{Error, ErrorContext};
2
3pub(crate) const SIGMA: [u32; 4] = [0x61707865, 0x3320646e, 0x79622d32, 0x6b206574];
7
8#[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
53pub(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
70pub(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
86pub(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
102pub(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
113pub(crate) fn zeroize_bytes(bytes: &mut [u8]) {
118 let (head, words, tail) = unsafe { bytes.align_to_mut::<u128>() };
123 zeroize_wide(head, words, tail);
124 zeroize::optimization_barrier(bytes);
125}
126
127pub(crate) fn zeroize_u64s(words: &mut [u64]) {
131 let (head, wide, tail) = unsafe { words.align_to_mut::<u128>() };
134 zeroize_wide(head, wide, tail);
135 zeroize::optimization_barrier(words);
136}
137
138pub(crate) fn zeroize_u32s(words: &mut [u32]) {
142 let (head, wide, tail) = unsafe { words.align_to_mut::<u128>() };
145 zeroize_wide(head, wide, tail);
146 zeroize::optimization_barrier(words);
147}
148
149pub(crate) fn zeroize_i16s(words: &mut [i16]) {
153 let (head, wide, tail) = unsafe { words.align_to_mut::<u128>() };
156 zeroize_wide(head, wide, tail);
157 zeroize::optimization_barrier(words);
158}
159
160pub(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
190pub(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#[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 let (groups, rest) = words.as_chunks_mut::<4>();
231 for group in groups {
232 for word in group {
233 unsafe { core::ptr::write_volatile(word, 0) };
235 }
236 }
237 for word in rest {
238 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 #[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 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 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 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#[cfg(test)]
467pub(crate) mod test_util {
468 use crate::error::{Error, ErrorContext, LengthConstraint};
469 use crate::test_prelude::*;
470
471 #[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 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 pub(crate) fn hex(s: &str) -> Vec<u8> {
507 hex::decode(s.replace(' ', "")).expect("hex")
508 }
509
510 pub(crate) fn hex_array<const N: usize>(s: &str) -> [u8; N] {
512 hex(s).try_into().expect("hex array length")
513 }
514
515 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 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}