arcis-compiler 0.14.1

A framework for writing secure multi-party computation (MPC) circuits to be executed on the Arcium network.
Documentation
use crate::{
    core::{
        actually_used_field::ActuallyUsedField,
        bounds::FieldBounds,
        circuits::{
            boolean::{
                boolean_value::{Boolean, BooleanValue},
                byte::Byte,
                sha2::Sha256,
                sha3::{Sha3_256, Sha3_512},
            },
            traits::arithmetic_circuit::ArithmeticCircuit,
        },
        expressions::expr::EvalFailure,
        global_value::value::FieldValue,
    },
    utils::{
        crypto::{
            hmac::{Hmac, Hmac_RescuePrime, Hmac_Sha256, Hmac_Sha3_256, Hmac_Sha3_512},
            key::RESCUE_KEY_COUNT,
            rescue_desc::RescueArg,
        },
        field::BaseField,
    },
};

pub trait Hkdf<const L: usize, T> {
    fn extract(&self, salt: Vec<T>, ikm: Vec<T>) -> Vec<T>;

    fn expand(&self, prk: Vec<T>, info: Vec<T>) -> [T; L];

    fn okm(&self, salt: Vec<T>, ikm: Vec<T>, info: Vec<T>) -> [T; L] {
        let prk = self.extract(salt, ikm);
        self.expand(prk, info)
    }
}

/// The Arcis HKDF based on the Rescue-Prime hash function.
/// We follow <https://datatracker.ietf.org/doc/html/rfc5869>.
/// We only support L = HashLen.
#[allow(non_camel_case_types)]
pub struct Hkdf_RescuePrime<F: ActuallyUsedField, T: RescueArg<F>> {
    hmac: Hmac_RescuePrime<F, T>,
}

impl<F: ActuallyUsedField, T: RescueArg<F>> Hkdf_RescuePrime<F, T> {
    pub fn new() -> Self {
        Self {
            hmac: Hmac_RescuePrime::new(),
        }
    }
}

impl<F: ActuallyUsedField, T: RescueArg<F>> Hkdf<RESCUE_KEY_COUNT, T> for Hkdf_RescuePrime<F, T> {
    fn extract(&self, salt: Vec<T>, ikm: Vec<T>) -> Vec<T> {
        let salt = if salt.is_empty() {
            vec![T::from(F::ZERO); self.hmac.hasher.rate]
        } else {
            salt
        };
        self.hmac.digest(salt, ikm)
    }

    fn expand(&self, prk: Vec<T>, info: Vec<T>) -> [T; RESCUE_KEY_COUNT] {
        // we only support L = HashLen = RESCUE_KEY_COUNT = 5, i.e. N = 1
        // message = empty string | info | 0x01
        let mut info = info;
        info.push(T::from(F::ONE));
        self.hmac
            .digest(prk, info)
            .try_into()
            .unwrap_or_else(|v: Vec<T>| {
                panic!(
                    "Expected a Vec of length {} (found {})",
                    RESCUE_KEY_COUNT,
                    v.len()
                )
            })
    }
}

impl<F: ActuallyUsedField, T: RescueArg<F>> Default for Hkdf_RescuePrime<F, T> {
    fn default() -> Self {
        Self::new()
    }
}

macro_rules! impl_hkdf {
    ($t: ident, $hmac: ident, $hasher: ident) => {
        /// The Arcis HKDF. We follow <https://datatracker.ietf.org/doc/html/rfc5869>.
        /// We only support L = HashLen.
        #[derive(Clone, Debug)]
        #[allow(non_camel_case_types)]
        pub struct $t {
            pub hmac: $hmac,
        }

        impl $t {
            pub fn new() -> Self {
                Self { hmac: $hmac::new() }
            }
        }

        impl<B: Boolean> Hkdf<{ $hasher::DIGEST_BYTES }, Byte<B>> for $t {
            fn extract(&self, salt: Vec<Byte<B>>, ikm: Vec<Byte<B>>) -> Vec<Byte<B>> {
                let salt = if salt.is_empty() {
                    vec![Byte::from(0); $hasher::DIGEST_BYTES]
                } else {
                    salt
                };
                self.hmac.digest(salt, ikm)
            }

            fn expand(
                &self,
                prk: Vec<Byte<B>>,
                info: Vec<Byte<B>>,
            ) -> [Byte<B>; { $hasher::DIGEST_BYTES }] {
                // we only support L = HashLen, i.e. N = 1
                // message = empty string | info | 0x01
                let mut info = info;
                info.push(Byte::from(1u8));
                self.hmac
                    .digest(prk, info)
                    .try_into()
                    .unwrap_or_else(|v: Vec<Byte<B>>| {
                        panic!(
                            "Expected a Vec of length {} (found {})",
                            $hasher::DIGEST_BYTES,
                            v.len()
                        )
                    })
            }
        }

        impl Default for $t {
            fn default() -> Self {
                Self::new()
            }
        }
    };
}

impl_hkdf!(Hkdf_Sha256, Hmac_Sha256, Sha256);
impl_hkdf!(Hkdf_Sha3_256, Hmac_Sha3_256, Sha3_256);
impl_hkdf!(Hkdf_Sha3_512, Hmac_Sha3_512, Sha3_512);

impl ArithmeticCircuit<BaseField> for Hkdf_Sha256 {
    fn eval(&self, x: Vec<BaseField>) -> Result<Vec<BaseField>, EvalFailure> {
        // all inputs are expected to be bytes
        x.iter()
            .for_each(|byte| assert!(*byte <= BaseField::from(255)));
        let mut salt = x
            .into_iter()
            .map(|val| val.to_le_bytes()[0])
            .collect::<Vec<u8>>();
        let mut ikm = salt.split_off(Sha256::DIGEST_BYTES);
        let info = ikm.split_off(Sha256::DIGEST_BYTES);

        let hkdf = hkdf::Hkdf::<sha2::Sha256>::new(Some(&salt), &ikm);
        let mut okm = [0u8; Sha256::DIGEST_BYTES];
        hkdf.expand(&info, &mut okm).unwrap_or_else(|_| {
            panic!(
                "{} is a valid length for Sha256 to output",
                Sha256::DIGEST_BYTES
            )
        });

        Ok(okm
            .iter()
            .map(|byte| BaseField::from(*byte as u64))
            .collect::<Vec<BaseField>>())
    }

    fn bounds(&self, _bounds: Vec<FieldBounds<BaseField>>) -> Vec<FieldBounds<BaseField>> {
        vec![FieldBounds::new(BaseField::from(0), BaseField::from(255)); Sha256::DIGEST_BYTES]
    }

    fn run(&self, vals: Vec<FieldValue<BaseField>>) -> Vec<FieldValue<BaseField>> {
        let mut salt = vals
            .into_iter()
            .map(Byte::from)
            .collect::<Vec<Byte<BooleanValue>>>();
        let mut ikm = salt.split_off(Sha256::DIGEST_BYTES);
        let info = ikm.split_off(Sha256::DIGEST_BYTES);

        let hkdf = Hkdf_Sha256::new();
        let okm = hkdf.okm(salt, ikm, info);

        okm.into_iter()
            .map(FieldValue::<BaseField>::from)
            .collect::<Vec<FieldValue<BaseField>>>()
    }
}

impl ArithmeticCircuit<BaseField> for Hkdf_Sha3_256 {
    fn eval(&self, x: Vec<BaseField>) -> Result<Vec<BaseField>, EvalFailure> {
        // all inputs are expected to be bytes
        x.iter()
            .for_each(|byte| assert!(*byte <= BaseField::from(255)));
        let mut salt = x
            .into_iter()
            .map(|val| val.to_le_bytes()[0])
            .collect::<Vec<u8>>();
        let mut ikm = salt.split_off(Sha3_256::DIGEST_BYTES);
        let info = ikm.split_off(Sha3_256::DIGEST_BYTES);

        let hkdf = hkdf::Hkdf::<sha3::Sha3_256>::new(Some(&salt), &ikm);
        let mut okm = [0u8; Sha3_256::DIGEST_BYTES];
        hkdf.expand(&info, &mut okm).unwrap_or_else(|_| {
            panic!(
                "{} is a valid length for Sha3_256 to output",
                Sha3_256::DIGEST_BYTES
            )
        });

        Ok(okm
            .iter()
            .map(|byte| BaseField::from(*byte as u64))
            .collect::<Vec<BaseField>>())
    }

    fn bounds(&self, _bounds: Vec<FieldBounds<BaseField>>) -> Vec<FieldBounds<BaseField>> {
        vec![FieldBounds::new(BaseField::from(0), BaseField::from(255)); Sha3_256::DIGEST_BYTES]
    }

    fn run(&self, vals: Vec<FieldValue<BaseField>>) -> Vec<FieldValue<BaseField>> {
        let mut salt = vals
            .into_iter()
            .map(Byte::from)
            .collect::<Vec<Byte<BooleanValue>>>();
        let mut ikm = salt.split_off(Sha3_256::DIGEST_BYTES);
        let info = ikm.split_off(Sha3_256::DIGEST_BYTES);

        let hkdf = Hkdf_Sha3_256::new();
        let okm = hkdf.okm(salt, ikm, info);

        okm.into_iter()
            .map(FieldValue::<BaseField>::from)
            .collect::<Vec<FieldValue<BaseField>>>()
    }
}

impl ArithmeticCircuit<BaseField> for Hkdf_Sha3_512 {
    fn eval(&self, x: Vec<BaseField>) -> Result<Vec<BaseField>, EvalFailure> {
        // all inputs are expected to be bytes
        x.iter()
            .for_each(|byte| assert!(*byte <= BaseField::from(255)));
        let mut salt = x
            .into_iter()
            .map(|val| val.to_le_bytes()[0])
            .collect::<Vec<u8>>();
        let mut ikm = salt.split_off(Sha3_512::DIGEST_BYTES);
        let info = ikm.split_off(Sha3_512::DIGEST_BYTES);

        let hkdf = hkdf::Hkdf::<sha3::Sha3_512>::new(Some(&salt), &ikm);
        let mut okm = [0u8; Sha3_512::DIGEST_BYTES];
        hkdf.expand(&info, &mut okm).unwrap_or_else(|_| {
            panic!(
                "{} is a valid length for Sha3_512 to output",
                Sha3_512::DIGEST_BYTES
            )
        });

        Ok(okm
            .iter()
            .map(|byte| BaseField::from(*byte as u64))
            .collect::<Vec<BaseField>>())
    }

    fn bounds(&self, _bounds: Vec<FieldBounds<BaseField>>) -> Vec<FieldBounds<BaseField>> {
        vec![FieldBounds::new(BaseField::from(0), BaseField::from(255)); Sha3_512::DIGEST_BYTES]
    }

    fn run(&self, vals: Vec<FieldValue<BaseField>>) -> Vec<FieldValue<BaseField>> {
        let mut salt = vals
            .into_iter()
            .map(Byte::from)
            .collect::<Vec<Byte<BooleanValue>>>();
        let mut ikm = salt.split_off(Sha3_512::DIGEST_BYTES);
        let info = ikm.split_off(Sha3_512::DIGEST_BYTES);

        let hkdf = Hkdf_Sha3_512::new();
        let okm = hkdf.okm(salt, ikm, info);

        okm.into_iter()
            .map(FieldValue::<BaseField>::from)
            .collect::<Vec<FieldValue<BaseField>>>()
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::core::circuits::traits::arithmetic_circuit::tests::TestedArithmeticCircuit;
    use rand::Rng;

    impl TestedArithmeticCircuit<BaseField> for Hkdf_Sha256 {
        fn gen_desc<R: Rng + ?Sized>(_rng: &mut R) -> Self {
            Self::new()
        }

        fn gen_n_inputs<R: Rng + ?Sized>(&self, _rng: &mut R) -> usize {
            // DIGEST_BYTES for the salt and the ikm respectively, and 2 bytes for the info
            2 * Sha256::DIGEST_BYTES + 2
        }

        fn gen_input_bounds<R: Rng + ?Sized>(_rng: &mut R) -> FieldBounds<BaseField> {
            FieldBounds::new(BaseField::from(0), BaseField::from(255))
        }
    }

    impl TestedArithmeticCircuit<BaseField> for Hkdf_Sha3_256 {
        fn gen_desc<R: Rng + ?Sized>(_rng: &mut R) -> Self {
            Self::new()
        }

        fn gen_n_inputs<R: Rng + ?Sized>(&self, _rng: &mut R) -> usize {
            // DIGEST_BYTES for the salt and the ikm respectively, and 2 bytes for the info
            2 * Sha3_256::DIGEST_BYTES + 2
        }

        fn gen_input_bounds<R: Rng + ?Sized>(_rng: &mut R) -> FieldBounds<BaseField> {
            FieldBounds::new(BaseField::from(0), BaseField::from(255))
        }
    }

    impl TestedArithmeticCircuit<BaseField> for Hkdf_Sha3_512 {
        fn gen_desc<R: Rng + ?Sized>(_rng: &mut R) -> Self {
            Self::new()
        }

        fn gen_n_inputs<R: Rng + ?Sized>(&self, _rng: &mut R) -> usize {
            // DIGEST_BYTES for the salt and the ikm respectively, and 2 bytes for the info
            2 * Sha3_512::DIGEST_BYTES + 2
        }

        fn gen_input_bounds<R: Rng + ?Sized>(_rng: &mut R) -> FieldBounds<BaseField> {
            FieldBounds::new(BaseField::from(0), BaseField::from(255))
        }
    }

    #[test]
    fn tested_hkdf_sha256() {
        Hkdf_Sha256::test(1, 1)
    }

    #[test]
    fn tested_hkdf_sha3_256() {
        Hkdf_Sha3_256::test(1, 1)
    }

    #[test]
    fn tested_hkdf_sha3_512() {
        Hkdf_Sha3_512::test(1, 1)
    }
}