use crate::Digest;
use crate::hmac::{hmac, hmac_multi};
use alloc::vec;
use alloc::vec::Vec;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct InvalidLength;
pub fn extract<H: Digest>(salt: &[u8], ikm: &[u8]) -> Vec<u8> {
let mut prk = vec![0u8; H::OUTPUT_LEN];
hmac::<H>(salt, ikm, &mut prk);
prk
}
pub fn expand<H: Digest>(prk: &[u8], info: &[u8], okm: &mut [u8]) -> Result<(), InvalidLength> {
let hlen = H::OUTPUT_LEN;
if okm.len() > 255 * hlen {
return Err(InvalidLength);
}
let mut t: Vec<u8> = Vec::new();
let mut offset = 0;
let mut block: usize = 1; while offset < okm.len() {
let mut ti = vec![0u8; hlen];
hmac_multi::<H>(prk, &[&t, info, &[block as u8]], &mut ti);
let take = core::cmp::min(hlen, okm.len() - offset);
okm[offset..offset + take].copy_from_slice(&ti[..take]);
offset += take;
t = ti;
block += 1;
}
Ok(())
}
pub fn derive<H: Digest>(salt: &[u8], ikm: &[u8], info: &[u8], okm: &mut [u8]) -> Result<(), InvalidLength> {
let prk = extract::<H>(salt, ikm);
expand::<H>(&prk, info, okm)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sha2::Sha256;
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 rfc5869_a1_sha256() {
let ikm = hx("0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b");
let salt = hx("000102030405060708090a0b0c");
let info = hx("f0f1f2f3f4f5f6f7f8f9");
let prk = extract::<Sha256>(&salt, &ikm);
assert_eq!(
prk,
hx("077709362c2e32df0ddc3f0dc47bba6390b6c73bb50f9c3122ec844ad7c2b3e5")
);
let mut okm = [0u8; 42];
derive::<Sha256>(&salt, &ikm, &info, &mut okm).unwrap();
assert_eq!(
okm[..],
hx("3cb25f25faacd57a90434f64d0362f2a2d2d0a90cf1a5a4c5db02d56ecc4c5bf34007208d5b887185865")[..]
);
}
#[test]
fn output_too_long_errors() {
let prk = extract::<Sha256>(b"salt", b"ikm");
let mut too_long = vec![0u8; 255 * 32 + 1];
assert_eq!(expand::<Sha256>(&prk, b"", &mut too_long), Err(InvalidLength));
let mut ok = vec![0u8; 255 * 32];
assert!(expand::<Sha256>(&prk, b"", &mut ok).is_ok());
}
}