use crate::Sha256;
use crate::pbkdf2::pbkdf2;
use alloc::vec;
use alloc::vec::Vec;
use core::num::NonZeroU32;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Error {
InvalidN,
InvalidBlockParams,
}
fn salsa20_8(block: &mut [u8; 64]) {
let mut x = [0u32; 16];
for (i, w) in x.iter_mut().enumerate() {
*w = u32::from_le_bytes([block[4 * i], block[4 * i + 1], block[4 * i + 2], block[4 * i + 3]]);
}
let orig = x;
for _ in 0..4 {
x[4] ^= x[0].wrapping_add(x[12]).rotate_left(7);
x[8] ^= x[4].wrapping_add(x[0]).rotate_left(9);
x[12] ^= x[8].wrapping_add(x[4]).rotate_left(13);
x[0] ^= x[12].wrapping_add(x[8]).rotate_left(18);
x[9] ^= x[5].wrapping_add(x[1]).rotate_left(7);
x[13] ^= x[9].wrapping_add(x[5]).rotate_left(9);
x[1] ^= x[13].wrapping_add(x[9]).rotate_left(13);
x[5] ^= x[1].wrapping_add(x[13]).rotate_left(18);
x[14] ^= x[10].wrapping_add(x[6]).rotate_left(7);
x[2] ^= x[14].wrapping_add(x[10]).rotate_left(9);
x[6] ^= x[2].wrapping_add(x[14]).rotate_left(13);
x[10] ^= x[6].wrapping_add(x[2]).rotate_left(18);
x[3] ^= x[15].wrapping_add(x[11]).rotate_left(7);
x[7] ^= x[3].wrapping_add(x[15]).rotate_left(9);
x[11] ^= x[7].wrapping_add(x[3]).rotate_left(13);
x[15] ^= x[11].wrapping_add(x[7]).rotate_left(18);
x[1] ^= x[0].wrapping_add(x[3]).rotate_left(7);
x[2] ^= x[1].wrapping_add(x[0]).rotate_left(9);
x[3] ^= x[2].wrapping_add(x[1]).rotate_left(13);
x[0] ^= x[3].wrapping_add(x[2]).rotate_left(18);
x[6] ^= x[5].wrapping_add(x[4]).rotate_left(7);
x[7] ^= x[6].wrapping_add(x[5]).rotate_left(9);
x[4] ^= x[7].wrapping_add(x[6]).rotate_left(13);
x[5] ^= x[4].wrapping_add(x[7]).rotate_left(18);
x[11] ^= x[10].wrapping_add(x[9]).rotate_left(7);
x[8] ^= x[11].wrapping_add(x[10]).rotate_left(9);
x[9] ^= x[8].wrapping_add(x[11]).rotate_left(13);
x[10] ^= x[9].wrapping_add(x[8]).rotate_left(18);
x[12] ^= x[15].wrapping_add(x[14]).rotate_left(7);
x[13] ^= x[12].wrapping_add(x[15]).rotate_left(9);
x[14] ^= x[13].wrapping_add(x[12]).rotate_left(13);
x[15] ^= x[14].wrapping_add(x[13]).rotate_left(18);
}
for i in 0..16 {
let v = x[i].wrapping_add(orig[i]);
block[4 * i..4 * i + 4].copy_from_slice(&v.to_le_bytes());
}
}
fn block_mix(b: &[u8], r: usize) -> Vec<u8> {
let two_r = 2 * r;
let mut x = [0u8; 64];
x.copy_from_slice(&b[(two_r - 1) * 64..two_r * 64]);
let mut out = vec![0u8; 128 * r];
for i in 0..two_r {
for k in 0..64 {
x[k] ^= b[i * 64 + k];
}
salsa20_8(&mut x);
let dst = if i % 2 == 0 { i / 2 } else { r + i / 2 };
out[dst * 64..dst * 64 + 64].copy_from_slice(&x);
}
out
}
fn integerify_mod(x: &[u8], r: usize, n: usize) -> usize {
let last = (2 * r - 1) * 64;
let j = u32::from_le_bytes([x[last], x[last + 1], x[last + 2], x[last + 3]]) as usize;
j & (n - 1) }
fn romix(b: &mut [u8], n: usize, r: usize) {
let blen = 128 * r;
let mut v = vec![0u8; n * blen]; let mut x = b.to_vec();
for i in 0..n {
v[i * blen..(i + 1) * blen].copy_from_slice(&x);
x = block_mix(&x, r);
}
for _ in 0..n {
let j = integerify_mod(&x, r, n);
for k in 0..blen {
x[k] ^= v[j * blen + k];
}
x = block_mix(&x, r);
}
b.copy_from_slice(&x);
}
pub fn scrypt(password: &[u8], salt: &[u8], n: usize, r: usize, p: usize, dk: &mut [u8]) -> Result<(), Error> {
if !n.is_power_of_two() || n <= 1 {
return Err(Error::InvalidN);
}
if r == 0 || p == 0 {
return Err(Error::InvalidBlockParams);
}
let blen = (128usize).checked_mul(r).ok_or(Error::InvalidBlockParams)?;
let total = p.checked_mul(blen).ok_or(Error::InvalidBlockParams)?;
if (p as u64) * (blen as u64) > ((u32::MAX as u64) * 32) {
return Err(Error::InvalidBlockParams);
}
let mut b = vec![0u8; total];
pbkdf2::<Sha256>(password, salt, NonZeroU32::MIN, &mut b);
for i in 0..p {
romix(&mut b[i * blen..(i + 1) * blen], n, r);
}
pbkdf2::<Sha256>(password, &b, NonZeroU32::MIN, dk);
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn hx(h: &str) -> Vec<u8> {
(0..h.len())
.step_by(2)
.map(|i| u8::from_str_radix(&h[i..i + 2], 16).unwrap())
.collect()
}
#[test]
fn rfc7914_vector1() {
let mut dk = [0u8; 64];
scrypt(b"", b"", 16, 1, 1, &mut dk).unwrap();
assert_eq!(
dk[..],
hx("77d6576238657b203b19ca42c18a0497f16b4844e3074ae8dfdffa3fede21442\
fcd0069ded0948f8326a753a0fc81f17e8d3e0fb2e0d3628cf35e20c38d18906")[..]
);
}
#[test]
fn invalid_params() {
let mut dk = [0u8; 16];
assert_eq!(scrypt(b"p", b"s", 15, 1, 1, &mut dk), Err(Error::InvalidN)); assert_eq!(scrypt(b"p", b"s", 1, 1, 1, &mut dk), Err(Error::InvalidN)); assert_eq!(scrypt(b"p", b"s", 16, 0, 1, &mut dk), Err(Error::InvalidBlockParams));
assert_eq!(scrypt(b"p", b"s", 16, 1, 0, &mut dk), Err(Error::InvalidBlockParams));
}
}