diff --git a/src/pss.rs b/src/pss.rs index 363989a0..4b5af1d9 100644 --- a/src/pss.rs +++ b/src/pss.rs @@ -39,7 +39,10 @@ use { crate::encoding::ID_RSASSA_PSS, const_oid::AssociatedOid, pkcs1::RsaPssParams, - spki::{der::Any, AlgorithmIdentifierOwned}, + spki::{ + der::{Any, ErrorKind}, + AlgorithmIdentifierOwned, + }, }; /// Digital signatures using PSS padding. @@ -272,15 +275,16 @@ pub fn get_default_pss_signature_algo_id() -> spki::Result::output_size() as u8; - get_pss_signature_algo_id::(salt_len) + get_pss_signature_algo_id::(::output_size()) } #[cfg(feature = "encoding")] -fn get_pss_signature_algo_id(salt_len: u8) -> spki::Result +fn get_pss_signature_algo_id(salt_len: usize) -> spki::Result where D: Digest + AssociatedOid, { + // RsaPssParams encodes salt_len in a single byte; reject rather than truncate. + let salt_len = u8::try_from(salt_len).map_err(|_| ErrorKind::Overflow.to_error())?; let pss_params = RsaPssParams::new::(salt_len); Ok(AlgorithmIdentifierOwned { @@ -672,4 +676,24 @@ tAboUGBxTDq3ZroNism3DaMIbKPyYrAqhKov1h5V .expect("verification to succeed"); } } + + // A salt length that does not fit in a byte must be rejected, not truncated (#703). + #[test] + fn signature_algorithm_identifier_rejects_oversized_salt_len() { + use spki::DynSignatureAlgorithmIdentifier; + + let priv_key = get_private_key(); + + // largest value that fits in a byte + let ok_key = SigningKey::::new_with_salt_len(priv_key.clone(), 255); + assert!(ok_key.signature_algorithm_identifier().is_ok()); + + // would wrap to 0 + let bad_key = SigningKey::::new_with_salt_len(priv_key.clone(), 256); + assert!(bad_key.signature_algorithm_identifier().is_err()); + + // blinded key uses the same path + let bad_blinded = BlindedSigningKey::::new_with_salt_len(priv_key, 256); + assert!(bad_blinded.signature_algorithm_identifier().is_err()); + } } diff --git a/src/pss/blinded_signing_key.rs b/src/pss/blinded_signing_key.rs index 3063c8a6..d260122a 100644 --- a/src/pss/blinded_signing_key.rs +++ b/src/pss/blinded_signing_key.rs @@ -184,7 +184,7 @@ where D: Digest + AssociatedOid, { fn signature_algorithm_identifier(&self) -> spki::Result { - get_pss_signature_algo_id::(self.salt_len as u8) + get_pss_signature_algo_id::(self.salt_len) } } diff --git a/src/pss/signing_key.rs b/src/pss/signing_key.rs index c159f0b4..4a44bf8f 100644 --- a/src/pss/signing_key.rs +++ b/src/pss/signing_key.rs @@ -221,7 +221,7 @@ where D: Digest + AssociatedOid, { fn signature_algorithm_identifier(&self) -> spki::Result { - get_pss_signature_algo_id::(self.salt_len as u8) + get_pss_signature_algo_id::(self.salt_len) } }