use crate::traits::snark::SpartanDigest;
use bincode::Options;
use serde::Serialize;
use sha2::{Digest, Sha256};
use std::io;
pub trait Digestible {
fn write_bytes<W: Sized + io::Write>(&self, byte_sink: &mut W) -> Result<(), io::Error>;
}
pub trait SimpleDigestible: Serialize {}
impl<T: SimpleDigestible> Digestible for T {
fn write_bytes<W: Sized + io::Write>(&self, byte_sink: &mut W) -> Result<(), io::Error> {
let config = bincode::DefaultOptions::new()
.with_little_endian()
.with_fixint_encoding();
config
.serialize_into(byte_sink, self)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))
}
}
pub struct DigestComputer<'a, T> {
inner: &'a T,
}
impl<'a, T: Digestible> DigestComputer<'a, T> {
fn hasher() -> Sha256 {
Sha256::new()
}
pub fn new(inner: &'a T) -> Self {
DigestComputer { inner }
}
pub fn digest(&self) -> Result<SpartanDigest, io::Error> {
let mut writer = io::BufWriter::with_capacity(64 * 1024, HashWriter(Self::hasher()));
self.inner.write_bytes(&mut writer)?;
let hasher = writer
.into_inner()
.map_err(|e| io::Error::other(e.to_string()))?
.0;
Ok(hasher.finalize().into())
}
}
struct HashWriter(Sha256);
impl io::Write for HashWriter {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.0.update(buf);
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::{DigestComputer, SimpleDigestible};
use once_cell::sync::OnceCell;
use serde::{Deserialize, Serialize};
#[derive(Serialize, Deserialize)]
struct S {
i: usize,
#[serde(skip, default = "OnceCell::new")]
digest: OnceCell<[u8; 32]>,
}
impl SimpleDigestible for S {}
impl S {
fn new(i: usize) -> Self {
S {
i,
digest: OnceCell::new(),
}
}
fn digest(&self) -> [u8; 32] {
self
.digest
.get_or_try_init(|| DigestComputer::new(self).digest())
.cloned()
.unwrap()
}
}
#[test]
fn test_digest_field_not_ingested_in_computation() {
let s1 = S::new(42);
let oc = OnceCell::new();
oc.set([1u8; 32]).unwrap();
let s2 = S { i: 42, digest: oc };
assert_eq!(
DigestComputer::<_>::new(&s1).digest().unwrap(),
DigestComputer::<_>::new(&s2).digest().unwrap()
);
assert_ne!(s2.digest(), DigestComputer::<_>::new(&s2).digest().unwrap());
}
#[test]
fn test_digest_impervious_to_serialization() {
let good_s = S::new(42);
let oc = OnceCell::new();
oc.set([2u8; 32]).unwrap();
let bad_s: S = S { i: 42, digest: oc };
assert_ne!(good_s.digest(), bad_s.digest());
let naughty_bytes = bincode::serialize(&bad_s).unwrap();
let retrieved_s: S = bincode::deserialize(&naughty_bytes).unwrap();
assert_eq!(good_s.digest(), retrieved_s.digest())
}
}