use crate::CryptoError;
use crypto_bigint::modular::{BoxedMontyForm, BoxedMontyParams};
use crypto_bigint::{BoxedUint, Odd};
use sha2::{Digest, Sha256};
pub const MODULUS_LEN: usize = 256;
const BITS: u32 = 2048;
const HLEN: usize = 32;
const SLEN: usize = 32;
const DB_LEN: usize = 223;
const PS_LEN: usize = 190;
const EXPONENT: [u8; 3] = [0x01, 0x00, 0x01];
fn mgf1(seed: &[u8], out: &mut [u8]) {
let mut counter: u32 = 0;
for block in out.chunks_mut(HLEN) {
let mut hasher = Sha256::new();
hasher.update(seed);
hasher.update(counter.to_be_bytes());
let digest = hasher.finalize();
for (dst, src) in block.iter_mut().zip(digest.iter()) {
*dst = *src;
}
counter = counter.wrapping_add(1);
}
}
pub fn verify_pss_sha256(
modulus: &[u8; MODULUS_LEN],
message: &[u8],
signature: &[u8],
) -> Result<(), CryptoError> {
if signature.len() != MODULUS_LEN {
return Err(CryptoError::BadLength);
}
let n = BoxedUint::from_be_slice(modulus.as_slice(), BITS).map_err(|_| CryptoError::BadLength)?;
let s = BoxedUint::from_be_slice(signature, BITS).map_err(|_| CryptoError::BadLength)?;
if s >= n {
return Err(CryptoError::BadSignature);
}
let odd = Odd::new(n).into_option().ok_or(CryptoError::BadLength)?;
let params = BoxedMontyParams::new(odd);
let exponent =
BoxedUint::from_be_slice(&EXPONENT, BITS).map_err(|_| CryptoError::BadLength)?;
let em = BoxedMontyForm::new(s, ¶ms).pow(&exponent).retrieve().to_be_bytes();
if em.len() != MODULUS_LEN {
return Err(CryptoError::BadLength);
}
if em.last() != Some(&0xbc) {
return Err(CryptoError::BadSignature);
}
let (masked_db, rest) = em.split_at_checked(DB_LEN).ok_or(CryptoError::BadLength)?;
let h = rest.get(..HLEN).ok_or(CryptoError::BadLength)?;
let first = masked_db.first().copied().ok_or(CryptoError::BadLength)?;
if first & 0x80 != 0 {
return Err(CryptoError::BadSignature);
}
let mut db = vec![0u8; DB_LEN];
mgf1(h, &mut db);
for (dst, src) in db.iter_mut().zip(masked_db.iter()) {
*dst ^= *src;
}
if let Some(head) = db.first_mut() {
*head &= 0x7f;
}
let (padding, tail) = db.split_at_checked(PS_LEN).ok_or(CryptoError::BadLength)?;
if padding.iter().any(|byte| *byte != 0) {
return Err(CryptoError::BadSignature);
}
let (separator, salt) = tail.split_first().ok_or(CryptoError::BadLength)?;
if *separator != 0x01 || salt.len() != SLEN {
return Err(CryptoError::BadSignature);
}
let mut hasher = Sha256::new();
hasher.update([0u8; 8]);
hasher.update(Sha256::digest(message));
hasher.update(salt);
let computed = hasher.finalize();
let expected = <[u8; HLEN]>::try_from(computed.as_slice()).map_err(|_| CryptoError::BadLength)?;
let found = <[u8; HLEN]>::try_from(h).map_err(|_| CryptoError::BadLength)?;
if crate::digest_eq(&expected, &found) {
Ok(())
} else {
Err(CryptoError::BadSignature)
}
}
pub const ATTESTATION_MODULUS_LENS: [usize; 3] = [256, 384, 512];
const SHA256_DIGEST_INFO: [u8; 19] = [
0x30, 0x31, 0x30, 0x0d, 0x06, 0x09, 0x60, 0x86, 0x48, 0x01, 0x65, 0x03, 0x04, 0x02, 0x01,
0x05, 0x00, 0x04, 0x20,
];
fn public_op(modulus: &[u8], input: &[u8]) -> Result<zeroize::Zeroizing<Vec<u8>>, CryptoError> {
let k = modulus.len();
if !ATTESTATION_MODULUS_LENS.contains(&k) || input.len() != k {
return Err(CryptoError::BadLength);
}
if modulus.first().is_none_or(|byte| byte & 0x80 == 0) {
return Err(CryptoError::BadLength);
}
let bits = u32::try_from(k.saturating_mul(8)).map_err(|_| CryptoError::BadLength)?;
let n = BoxedUint::from_be_slice(modulus, bits).map_err(|_| CryptoError::BadLength)?;
let x = BoxedUint::from_be_slice(input, bits).map_err(|_| CryptoError::BadLength)?;
if x >= n {
return Err(CryptoError::BadSignature);
}
let odd = Odd::new(n).into_option().ok_or(CryptoError::BadLength)?;
let params = BoxedMontyParams::new(odd);
let exponent = BoxedUint::from_be_slice(&EXPONENT, bits).map_err(|_| CryptoError::BadLength)?;
let out = BoxedMontyForm::new(x, ¶ms).pow(&exponent).retrieve().to_be_bytes();
if out.len() != k {
return Err(CryptoError::BadLength);
}
Ok(zeroize::Zeroizing::new(out.to_vec()))
}
pub fn verify_pkcs1v15_sha256(
modulus: &[u8],
message: &[u8],
signature: &[u8],
) -> Result<(), CryptoError> {
use subtle::ConstantTimeEq as _;
let k = modulus.len();
if signature.len() != k {
return Err(CryptoError::BadLength);
}
let recovered = public_op(modulus, signature)?;
let t_len = SHA256_DIGEST_INFO.len().saturating_add(HLEN);
let pad_end = k.checked_sub(t_len).ok_or(CryptoError::BadLength)?;
let mut expected = vec![0xffu8; k];
let (head, tail) = expected.split_at_mut(pad_end);
if let Some(first) = head.first_mut() {
*first = 0x00;
}
if let Some(second) = head.get_mut(1) {
*second = 0x01;
}
if let Some(separator) = head.last_mut() {
*separator = 0x00;
}
let (info, digest) = tail.split_at_mut(SHA256_DIGEST_INFO.len());
info.copy_from_slice(&SHA256_DIGEST_INFO);
digest.copy_from_slice(&Sha256::digest(message));
if pad_end < 11 {
return Err(CryptoError::BadLength);
}
if bool::from(recovered.as_slice().ct_eq(expected.as_slice())) {
Ok(())
} else {
Err(CryptoError::BadSignature)
}
}
pub fn encrypt_oaep_sha256(
modulus: &[u8],
label: &[u8],
message: &[u8],
seed: &[u8; HLEN],
) -> Result<Vec<u8>, CryptoError> {
let k = modulus.len();
if !ATTESTATION_MODULUS_LENS.contains(&k) {
return Err(CryptoError::BadLength);
}
let room = k.checked_sub(HLEN.saturating_mul(2).saturating_add(2)).ok_or(CryptoError::BadLength)?;
if message.len() > room {
return Err(CryptoError::BadLength);
}
let db_len = k.saturating_sub(HLEN).saturating_sub(1);
let mut db = zeroize::Zeroizing::new(vec![0u8; db_len]);
let (l_hash, rest) = db.split_at_mut(HLEN);
l_hash.copy_from_slice(&Sha256::digest(label));
let marker_at = rest.len().checked_sub(message.len().saturating_add(1)).ok_or(CryptoError::BadLength)?;
let (_, tail) = rest.split_at_mut(marker_at);
let (marker, body) = tail.split_first_mut().ok_or(CryptoError::BadLength)?;
*marker = 0x01;
body.copy_from_slice(message);
let mut db_mask = zeroize::Zeroizing::new(vec![0u8; db_len]);
mgf1(seed, &mut db_mask);
for (byte, mask) in db.iter_mut().zip(db_mask.iter()) {
*byte ^= *mask;
}
let mut seed_mask = zeroize::Zeroizing::new([0u8; HLEN]);
mgf1(&db, seed_mask.as_mut_slice());
let mut masked_seed = zeroize::Zeroizing::new(*seed);
for (byte, mask) in masked_seed.iter_mut().zip(seed_mask.iter()) {
*byte ^= *mask;
}
let mut em = zeroize::Zeroizing::new(vec![0u8; k]);
let (zero_and_seed, masked_db) = em.split_at_mut(HLEN.saturating_add(1));
let (_, seed_slot) = zero_and_seed.split_at_mut(1);
seed_slot.copy_from_slice(masked_seed.as_slice());
masked_db.copy_from_slice(&db);
Ok(public_op(modulus, &em)?.to_vec())
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic, clippy::indexing_slicing)]
mod tests {
use super::*;
#[test]
fn the_derived_lengths_match_the_pss_layout() {
assert_eq!(DB_LEN, MODULUS_LEN - HLEN - 1);
assert_eq!(PS_LEN, MODULUS_LEN - SLEN - HLEN - 2);
assert_eq!(PS_LEN + 1 + SLEN, DB_LEN);
}
#[test]
fn a_signature_of_the_wrong_length_is_refused_by_length_not_by_verdict() {
let modulus = [0xffu8; MODULUS_LEN];
assert_eq!(
verify_pss_sha256(&modulus, b"x", &[0u8; MODULUS_LEN - 1]),
Err(CryptoError::BadLength)
);
assert_eq!(
verify_pss_sha256(&modulus, b"x", &[0u8; MODULUS_LEN + 1]),
Err(CryptoError::BadLength)
);
}
#[test]
fn an_even_modulus_is_refused() {
let mut modulus = [0xffu8; MODULUS_LEN];
modulus[MODULUS_LEN - 1] = 0xfe;
let signature = [0u8; MODULUS_LEN];
assert_eq!(verify_pss_sha256(&modulus, b"x", &signature), Err(CryptoError::BadLength));
}
#[test]
fn attestation_moduli_are_limited_to_full_width_known_lengths() {
let short = [0xffu8; 128];
assert_eq!(verify_pkcs1v15_sha256(&short, b"x", &[0u8; 128]), Err(CryptoError::BadLength));
let mut hollow = [0xffu8; 256];
hollow[0] = 0x7f;
assert_eq!(verify_pkcs1v15_sha256(&hollow, b"x", &[0u8; 256]), Err(CryptoError::BadLength));
assert_eq!(encrypt_oaep_sha256(&short, b"", b"x", &[0u8; 32]), Err(CryptoError::BadLength));
let modulus = [0xffu8; 256];
assert_eq!(
encrypt_oaep_sha256(&modulus, b"", &[0u8; 256 - 64 - 1], &[0u8; 32]),
Err(CryptoError::BadLength)
);
}
#[test]
fn a_pkcs1_signature_not_below_the_modulus_is_refused() {
let modulus = [0xffu8; 256];
assert_eq!(
verify_pkcs1v15_sha256(&modulus, b"x", &[0xffu8; 256]),
Err(CryptoError::BadSignature)
);
}
#[test]
fn a_zero_signature_never_verifies() {
let modulus = [0xffu8; MODULUS_LEN];
assert_eq!(
verify_pss_sha256(&modulus, b"x", &[0u8; MODULUS_LEN]),
Err(CryptoError::BadSignature)
);
}
}