#![allow(dead_code)]
use std::num::NonZeroU32;
use p256::elliptic_curve::sec1::ToEncodedPoint;
use p256::{ProjectivePoint, Scalar};
use ring::hkdf;
use ring::pbkdf2;
use crate::error::{Error, Result};
const PBKDF_OUTPUT_LEN: usize = 80;
const W_HALF_LEN: usize = 40;
const PBKDF_MIN_ITERATIONS: u32 = 1_000;
const PBKDF_MAX_ITERATIONS: u32 = 100_000;
const PBKDF_SALT_MIN: usize = 16;
const PBKDF_SALT_MAX: usize = 32;
pub(crate) fn validate_params(iterations: u32, salt: &[u8]) -> Result<()> {
if iterations < PBKDF_MIN_ITERATIONS {
return Err(Error::PbkdfIterationsTooLow(iterations));
}
if iterations > PBKDF_MAX_ITERATIONS {
return Err(Error::PbkdfIterationsTooHigh {
iterations,
max: PBKDF_MAX_ITERATIONS,
});
}
if salt.len() < PBKDF_SALT_MIN || salt.len() > PBKDF_SALT_MAX {
return Err(Error::PbkdfSaltLengthInvalid(salt.len()));
}
Ok(())
}
pub(crate) fn derive_w0_w1(pin: u32, salt: &[u8], iterations: u32) -> Result<(Scalar, Scalar)> {
validate_params(iterations, salt)?;
let pin_bytes = pin.to_le_bytes();
let mut out = [0u8; PBKDF_OUTPUT_LEN];
let iter_nz = NonZeroU32::new(iterations).ok_or(Error::PinDerivationFailed)?;
pbkdf2::derive(
pbkdf2::PBKDF2_HMAC_SHA256,
iter_nz,
salt,
&pin_bytes,
&mut out,
);
let w0 = reduce_40_bytes_mod_q(&out[..W_HALF_LEN])?;
let w1 = reduce_40_bytes_mod_q(&out[W_HALF_LEN..])?;
Ok((w0, w1))
}
fn reduce_40_bytes_mod_q(input: &[u8]) -> Result<Scalar> {
use p256::elliptic_curve::bigint::{Encoding, NonZero, U256};
use p256::elliptic_curve::ops::Reduce;
use p256::elliptic_curve::{bigint::ArrayEncoding, Curve};
type U320 = p256::elliptic_curve::bigint::Uint<5>;
if input.len() != W_HALF_LEN {
return Err(Error::PinDerivationFailed);
}
let n320 = U320::from_be_slice(input);
let order_u256: U256 = p256::NistP256::ORDER;
let order_be = order_u256.to_be_byte_array(); let mut order_buf = [0u8; 40];
order_buf[8..].copy_from_slice(&order_be); let order_u320 = U320::from_be_slice(&order_buf);
let order_nz: NonZero<U320> = NonZero::new(order_u320)
.into_option()
.ok_or(Error::PinDerivationFailed)?;
let rem_u320 = n320.rem(&order_nz);
let rem_be: [u8; 40] = rem_u320.to_be_bytes();
let mut scalar_bytes = [0u8; 32];
scalar_bytes.copy_from_slice(&rem_be[8..]);
let scalar = <Scalar as Reduce<U256>>::reduce(U256::from_be_slice(&scalar_bytes));
Ok(scalar)
}
pub(crate) fn derive_l(w1: &Scalar) -> [u8; 65] {
let l_point = ProjectivePoint::GENERATOR * w1;
let encoded = l_point.to_affine().to_encoded_point(false);
let mut out = [0u8; 65];
out.copy_from_slice(encoded.as_bytes());
out
}
pub fn pake_passcode_verifier(passcode: u32, salt: &[u8], iterations: u32) -> Result<[u8; 97]> {
validate_params(iterations, salt)?;
let (w0, w1) = derive_w0_w1(passcode, salt, iterations)?;
let l = derive_l(&w1);
let w0_be: p256::FieldBytes = w0.to_bytes();
let mut out = [0u8; 97];
out[..32].copy_from_slice(&w0_be);
out[32..].copy_from_slice(&l);
Ok(out)
}
pub(crate) fn hkdf_expand(prk: &[u8], info: &[u8], out: &mut [u8]) -> Result<()> {
let salt = hkdf::Salt::new(hkdf::HKDF_SHA256, &[]);
let prk_obj = salt.extract(prk);
let info_arr = [info];
let okm = prk_obj
.expand(&info_arr, OutLen(out.len()))
.map_err(|_| Error::PinDerivationFailed)?;
okm.fill(out).map_err(|_| Error::PinDerivationFailed)?;
Ok(())
}
struct OutLen(usize);
impl hkdf::KeyType for OutLen {
fn len(&self) -> usize {
self.0
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)] mod tests {
use super::*;
#[test]
fn validate_rejects_iter_below_min() {
assert!(matches!(
validate_params(999, &[0u8; 16]),
Err(Error::PbkdfIterationsTooLow(999))
));
}
#[test]
fn validate_rejects_iter_zero() {
assert!(matches!(
validate_params(0, &[0u8; 16]),
Err(Error::PbkdfIterationsTooLow(0))
));
}
#[test]
fn validate_rejects_salt_too_short() {
assert!(matches!(
validate_params(1_000, &[0u8; 15]),
Err(Error::PbkdfSaltLengthInvalid(15))
));
}
#[test]
fn validate_rejects_salt_too_long() {
assert!(matches!(
validate_params(1_000, &[0u8; 33]),
Err(Error::PbkdfSaltLengthInvalid(33))
));
}
#[test]
fn validate_accepts_boundary_values() {
validate_params(1_000, &[0u8; 16]).unwrap();
validate_params(10_000, &[0u8; 32]).unwrap();
}
#[test]
fn validate_rejects_iter_above_max() {
assert!(matches!(
validate_params(u32::MAX, &[0u8; 16]),
Err(Error::PbkdfIterationsTooHigh {
iterations: u32::MAX,
max: PBKDF_MAX_ITERATIONS,
})
));
let over = PBKDF_MAX_ITERATIONS + 1;
assert!(matches!(
validate_params(over, &[0u8; 16]),
Err(Error::PbkdfIterationsTooHigh {
iterations,
max,
}) if iterations == over && max == PBKDF_MAX_ITERATIONS
));
}
#[test]
fn validate_accepts_max_boundary() {
validate_params(PBKDF_MAX_ITERATIONS, &[0u8; 16]).unwrap();
}
#[test]
#[allow(clippy::similar_names)]
fn derive_w0_w1_is_deterministic() {
let salt = [0x42u8; 16];
let (w0a, w1a) = derive_w0_w1(20_202_021, &salt, 1_000).unwrap();
let (w0b, w1b) = derive_w0_w1(20_202_021, &salt, 1_000).unwrap();
assert_eq!(w0a.to_bytes(), w0b.to_bytes());
assert_eq!(w1a.to_bytes(), w1b.to_bytes());
}
#[test]
#[allow(clippy::similar_names)]
fn derive_w0_w1_changes_with_pin() {
let salt = [0x42u8; 16];
let (w0a, _) = derive_w0_w1(20_202_021, &salt, 1_000).unwrap();
let (w0b, _) = derive_w0_w1(20_202_022, &salt, 1_000).unwrap();
assert_ne!(w0a.to_bytes(), w0b.to_bytes());
}
#[test]
#[allow(clippy::similar_names)]
fn derive_w0_w1_changes_with_salt() {
let (w0a, _) = derive_w0_w1(20_202_021, &[0x42u8; 16], 1_000).unwrap();
let (w0b, _) = derive_w0_w1(20_202_021, &[0x43u8; 16], 1_000).unwrap();
assert_ne!(w0a.to_bytes(), w0b.to_bytes());
}
#[test]
#[allow(clippy::similar_names)]
fn derive_w0_w1_changes_with_iterations() {
let salt = [0x42u8; 16];
let (w0a, _) = derive_w0_w1(20_202_021, &salt, 1_000).unwrap();
let (w0b, _) = derive_w0_w1(20_202_021, &salt, 2_000).unwrap();
assert_ne!(w0a.to_bytes(), w0b.to_bytes());
}
#[test]
fn derive_w0_w1_rejects_bad_params() {
assert!(derive_w0_w1(20_202_021, &[0u8; 16], 999).is_err());
assert!(derive_w0_w1(20_202_021, &[0u8; 15], 1_000).is_err());
}
#[test]
fn derive_l_produces_uncompressed_p256_point() {
let salt = [0x42u8; 16];
let (_, w1) = derive_w0_w1(20_202_021, &salt, 1_000).unwrap();
let l = derive_l(&w1);
assert_eq!(l[0], 0x04, "SEC1 uncompressed prefix");
assert_eq!(l.len(), 65);
}
#[test]
fn derive_l_is_deterministic() {
let salt = [0x42u8; 16];
let (_, w1) = derive_w0_w1(20_202_021, &salt, 1_000).unwrap();
let l1 = derive_l(&w1);
let l2 = derive_l(&w1);
assert_eq!(l1, l2);
}
#[test]
#[allow(clippy::similar_names)]
fn derive_l_changes_with_w1() {
let (_, w1a) = derive_w0_w1(20_202_021, &[0x42u8; 16], 1_000).unwrap();
let (_, w1b) = derive_w0_w1(20_202_022, &[0x42u8; 16], 1_000).unwrap();
let la = derive_l(&w1a);
let lb = derive_l(&w1b);
assert_ne!(la, lb);
}
#[test]
fn hkdf_expand_is_deterministic() {
let prk = [0x11u8; 32];
let info = b"ConfirmationKeys";
let mut out_a = [0u8; 32];
let mut out_b = [0u8; 32];
hkdf_expand(&prk, info, &mut out_a).unwrap();
hkdf_expand(&prk, info, &mut out_b).unwrap();
assert_eq!(out_a, out_b);
}
#[test]
fn hkdf_expand_differs_with_info() {
let prk = [0x11u8; 32];
let mut out_a = [0u8; 32];
let mut out_b = [0u8; 32];
hkdf_expand(&prk, b"ConfirmationKeys", &mut out_a).unwrap();
hkdf_expand(&prk, b"SessionKeys", &mut out_b).unwrap();
assert_ne!(out_a, out_b);
}
#[test]
fn hkdf_expand_differs_with_prk() {
let mut out_a = [0u8; 32];
let mut out_b = [0u8; 32];
hkdf_expand(&[0x11u8; 32], b"ConfirmationKeys", &mut out_a).unwrap();
hkdf_expand(&[0x22u8; 32], b"ConfirmationKeys", &mut out_b).unwrap();
assert_ne!(out_a, out_b);
}
#[test]
fn hkdf_expand_variable_output_length() {
let prk = [0x11u8; 32];
let mut out_16 = [0u8; 16];
let mut out_48 = [0u8; 48];
hkdf_expand(&prk, b"SessionKeys", &mut out_16).unwrap();
hkdf_expand(&prk, b"SessionKeys", &mut out_48).unwrap();
assert_eq!(&out_48[..16], &out_16[..]);
}
#[test]
#[allow(clippy::unreadable_literal)] fn pake_passcode_verifier_matches_known_params_vector() {
let salt = hex_to_vec("abcdabcdabcdabcdabcdabcdabcdabcdabcdabcdabcdabcd");
let w0_hex = "d8e14a916650ff651dcd0c34d15fc1ed9b8232d550827be4816cc8e0fbd31bfa";
let l_hex = "04cd104598fca43b1a17bcc78d51ad2ec542bdd3f8541ecfbe5b66c7ea714f505d7e4626053087b2980c37876053c431600662fe09af442d6ca49525334dbf59e5";
let mut expected = hex_to_vec(w0_hex);
expected.extend_from_slice(&hex_to_vec(l_hex));
let got = super::pake_passcode_verifier(123456, &salt, 2000).unwrap();
assert_eq!(got.len(), 97);
assert_eq!(&got[..], &expected[..]);
}
fn hex_to_vec(s: &str) -> Vec<u8> {
(0..s.len())
.step_by(2)
.map(|i| u8::from_str_radix(&s[i..i + 2], 16).unwrap())
.collect()
}
}