diff --git a/light-poseidon/Cargo.toml b/light-poseidon/Cargo.toml index 6b86875..56105a1 100644 --- a/light-poseidon/Cargo.toml +++ b/light-poseidon/Cargo.toml @@ -12,7 +12,7 @@ edition = "2021" [dependencies] ark-bn254 = "0.5.0" ark-ff = "0.5.0" -num-bigint = "0.4.4" +arrayvec = "0.7.8" thiserror = "1.0" [dev-dependencies] diff --git a/light-poseidon/src/lib.rs b/light-poseidon/src/lib.rs index d3a9d02..e2dbd1f 100644 --- a/light-poseidon/src/lib.rs +++ b/light-poseidon/src/lib.rs @@ -130,12 +130,14 @@ //! read the audit report [here](https://github.com/Lightprotocol/light-poseidon/blob/main/assets/audit.pdf). use ark_bn254::Fr; use ark_ff::{BigInteger, PrimeField, Zero}; +use arrayvec::ArrayVec; use thiserror::Error; pub mod parameters; pub const HASH_LEN: usize = 32; pub const MAX_X5_LEN: usize = 13; +pub const MAX_INPUTS: usize = 12; #[derive(Error, Debug, PartialEq)] pub enum PoseidonError { @@ -312,7 +314,7 @@ pub trait PoseidonBytesHasher { pub struct Poseidon { params: PoseidonParameters, domain_tag: F, - state: Vec, + state: ArrayVec, } impl Poseidon { @@ -325,14 +327,24 @@ impl Poseidon { } fn with_domain_tag(params: PoseidonParameters, domain_tag: F) -> Self { - let width = params.width; Self { domain_tag, params, - state: Vec::with_capacity(width), + state: ArrayVec::new(), } } + fn validate_inputs_length(&self, inputs: &[T]) -> Result<(), PoseidonError> { + if inputs.len() != self.params.width - 1 { + return Err(PoseidonError::InvalidNumberOfInputs { + inputs: inputs.len(), + max_limit: self.params.width - 1, + width: self.params.width, + }); + } + Ok(()) + } + #[inline(always)] fn apply_ark(&mut self, round: usize) { self.state.iter_mut().enumerate().for_each(|(i, a)| { @@ -388,14 +400,7 @@ impl Poseidon { impl PoseidonHasher for Poseidon { fn hash(&mut self, inputs: &[F]) -> Result { - if inputs.len() != self.params.width - 1 { - return Err(PoseidonError::InvalidNumberOfInputs { - inputs: inputs.len(), - max_limit: self.params.width - 1, - width: self.params.width, - }); - } - + self.validate_inputs_length(inputs)?; self.state.push(self.domain_tag); for input in inputs { @@ -430,32 +435,81 @@ impl PoseidonHasher for Poseidon { } } +/// Serializes prime field elements into the fixed-size hash output byte +/// array ([`HASH_LEN`]) directly from their limbs, without the intermediate +/// `Vec` produced by `to_bytes_be`/`to_bytes_le`. +trait IntoHashBytes { + /// Serializes the element in big-endian byte order, returning `None` + /// when the element's byte length does not match [`HASH_LEN`]. + fn into_hash_bytes_be(self) -> Option<[u8; HASH_LEN]>; + /// Serializes the element in little-endian byte order, returning `None` + /// when the element's byte length does not match [`HASH_LEN`]. + fn into_hash_bytes_le(self) -> Option<[u8; HASH_LEN]>; +} + +fn fill_hash_bytes_be(element: F) -> Option<[u8; HASH_LEN]> { + let bigint = element.into_bigint(); + let limbs = bigint.as_ref(); + if limbs.len() * 8 != HASH_LEN { + return None; + } + let mut hash_bytes = [0u8; HASH_LEN]; + for (i, limb) in limbs.iter().rev().enumerate() { + hash_bytes[i * 8..i * 8 + 8].copy_from_slice(&limb.to_be_bytes()); + } + Some(hash_bytes) +} + +fn fill_hash_bytes_le(element: F) -> Option<[u8; HASH_LEN]> { + let bigint = element.into_bigint(); + let limbs = bigint.as_ref(); + if limbs.len() * 8 != HASH_LEN { + return None; + } + let mut hash_bytes = [0u8; HASH_LEN]; + for (i, limb) in limbs.iter().enumerate() { + hash_bytes[i * 8..i * 8 + 8].copy_from_slice(&limb.to_le_bytes()); + } + Some(hash_bytes) +} + +impl IntoHashBytes for F { + fn into_hash_bytes_be(self) -> Option<[u8; HASH_LEN]> { + fill_hash_bytes_be(self) + } + + fn into_hash_bytes_le(self) -> Option<[u8; HASH_LEN]> { + fill_hash_bytes_le(self) + } +} + macro_rules! impl_hash_bytes { - ($fn_name:ident, $bytes_to_prime_field_element_fn:ident, $to_bytes_fn:ident) => { + ($fn_name:ident, $bytes_to_prime_field_element_fn:ident, $into_hash_bytes_fn:ident) => { fn $fn_name(&mut self, inputs: &[&[u8]]) -> Result<[u8; HASH_LEN], PoseidonError> { - let inputs: Result, _> = inputs - .iter() - .map(|input| validate_bytes_length::(input)) - .collect(); - let inputs = inputs?; - let inputs: Result, _> = inputs - .iter() - .map(|input| $bytes_to_prime_field_element_fn(input)) - .collect(); - let inputs = inputs?; - let hash = self.hash(&inputs)?; - - hash.into_bigint() - .$to_bytes_fn() - .try_into() - .map_err(|_| PoseidonError::VecToArray) + let mut deserialized_inputs: ArrayVec = ArrayVec::new(); + for input in inputs { + validate_bytes_length::(input)?; + let deserialized_input = $bytes_to_prime_field_element_fn(input)?; + deserialized_inputs.push(deserialized_input); + } + let hash = self.hash(deserialized_inputs.as_slice())?; + + hash.$into_hash_bytes_fn().ok_or(PoseidonError::VecToArray) } }; } impl PoseidonBytesHasher for Poseidon { - impl_hash_bytes!(hash_bytes_le, bytes_to_prime_field_element_le, to_bytes_le); - impl_hash_bytes!(hash_bytes_be, bytes_to_prime_field_element_be, to_bytes_be); + impl_hash_bytes!( + hash_bytes_le, + bytes_to_prime_field_element_le, + into_hash_bytes_le + ); + impl_hash_bytes!( + hash_bytes_be, + bytes_to_prime_field_element_be, + into_hash_bytes_be + ); } /// Checks whether a slice of bytes is not empty or its length does not exceed @@ -468,7 +522,7 @@ impl PoseidonBytesHasher for Poseidon { /// to collisions. The purpose of this function is to prevent them by returning /// and error. It should be always used before converting byte slices to /// prime field elements. -pub fn validate_bytes_length(input: &[u8]) -> Result<&[u8], PoseidonError> +pub fn validate_bytes_length(input: &[u8]) -> Result<(), PoseidonError> where F: PrimeField, { @@ -482,37 +536,58 @@ where modulus_bytes_len, }); } - Ok(input) + Ok(()) } macro_rules! impl_bytes_to_prime_field_element { - ($name:ident, $from_bytes_method:ident, $endianess:expr) => { + ($name:ident, $is_be:literal, $endian:expr) => { #[doc = "Converts a slice of "] - #[doc = $endianess] + #[doc = $endian] #[doc = "-endian bytes into a prime field element, \ represented by the [`ark_ff::PrimeField`](ark_ff::PrimeField) trait."] pub fn $name(input: &[u8]) -> Result where F: PrimeField, { - let element = num_bigint::BigUint::$from_bytes_method(input); - let element = F::BigInt::try_from(element).map_err(|_| PoseidonError::BytesToBigInt)?; - - // In theory, `F::from_bigint` should also perform a check whether input is - // larger than modulus (and return `None` if it is), but it's not reliable... - // To be sure, we check it ourselves. - if element >= F::MODULUS { - return Err(PoseidonError::InputLargerThanModulus); + let max_len = F::BigInt::NUM_LIMBS * 8; + // Trim the leading (most-significant-end) zero bytes so that inputs + // longer than the modulus size but representing a smaller value are + // still accepted. + let trimmed = if $is_be { + let start = input.iter().position(|b| *b != 0).unwrap_or(input.len()); + &input[start..] + } else { + let end = input + .iter() + .rposition(|b| *b != 0) + .map(|i| i + 1) + .unwrap_or(0); + &input[..end] + }; + if trimmed.len() > max_len { + return Err(PoseidonError::BytesToBigInt); + } + + // Build the value in `F::BigInt` (a fixed-size stack array of limbs) + // directly from the bytes, most significant byte first. + let mut value = F::BigInt::default(); + if $is_be { + for byte in trimmed.iter() { + value = (value << 8) | F::BigInt::from(*byte); + } + } else { + for byte in trimmed.iter().rev() { + value = (value << 8) | F::BigInt::from(*byte); + } } - let element = F::from_bigint(element).ok_or(PoseidonError::InputLargerThanModulus)?; - Ok(element) + F::from_bigint(value).ok_or(PoseidonError::InputLargerThanModulus) } }; } -impl_bytes_to_prime_field_element!(bytes_to_prime_field_element_le, from_bytes_le, "little"); -impl_bytes_to_prime_field_element!(bytes_to_prime_field_element_be, from_bytes_be, "big"); +impl_bytes_to_prime_field_element!(bytes_to_prime_field_element_le, false, "little"); +impl_bytes_to_prime_field_element!(bytes_to_prime_field_element_be, true, "big"); impl Poseidon { pub fn new_circom(nr_inputs: usize) -> Result, PoseidonError> { diff --git a/light-poseidon/tests/bn254_fq_x5.rs b/light-poseidon/tests/bn254_fq_x5.rs index b7e764c..adac60a 100644 --- a/light-poseidon/tests/bn254_fq_x5.rs +++ b/light-poseidon/tests/bn254_fq_x5.rs @@ -177,8 +177,7 @@ fn test_poseidon_bn254_x5_fq_validate_bytes_length() { } let input = vec![1u8; 32]; - let res = validate_bytes_length::(&input).unwrap(); - assert_eq!(res, &input); + validate_bytes_length::(&input).unwrap(); for i in 33..64 { let input = vec![1u8; i];