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)
}
}
#[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] {
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) => {
#[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 }] {
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> {
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> {
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> {
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 {
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 {
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 {
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)
}
}