diff --git a/digest/src/block_api/ct_variable.rs b/digest/src/block_api/ct_variable.rs index 8c8dba038..da7fec423 100644 --- a/digest/src/block_api/ct_variable.rs +++ b/digest/src/block_api/ct_variable.rs @@ -13,6 +13,30 @@ use common::{ }; use core::{fmt, marker::PhantomData}; +#[cfg(feature = "zeroize")] +struct ScopedFullResult(Array); + +#[cfg(feature = "zeroize")] +impl Default for ScopedFullResult { + fn default() -> Self { + Self(Default::default()) + } +} + +#[cfg(feature = "zeroize")] +impl Drop for ScopedFullResult { + fn drop(&mut self) { + use zeroize::Zeroize; + self.0.as_mut_slice().zeroize(); + #[cfg(test)] + SCOPED_FULL_RESULT_DROPS.fetch_add(1, core::sync::atomic::Ordering::SeqCst); + } +} + +#[cfg(all(test, feature = "zeroize"))] +static SCOPED_FULL_RESULT_DROPS: core::sync::atomic::AtomicUsize = + core::sync::atomic::AtomicUsize::new(0); + /// Wrapper around [`VariableOutputCore`] which selects output size at compile time. pub struct CtOutWrapper where @@ -105,8 +129,15 @@ where buffer: &mut Buffer, out: &mut Array, ) { + #[cfg(feature = "zeroize")] + let mut scoped_full_res = ScopedFullResult::::default(); + #[cfg(feature = "zeroize")] + let full_res = &mut scoped_full_res.0; + #[cfg(not(feature = "zeroize"))] let mut full_res = Default::default(); - self.inner.finalize_variable_core(buffer, &mut full_res); + #[cfg(not(feature = "zeroize"))] + let full_res = &mut full_res; + self.inner.finalize_variable_core(buffer, full_res); let n = out.len(); let m = full_res.len() - n; match T::TRUNC_SIDE { @@ -116,6 +147,30 @@ where } } +#[cfg(all(test, feature = "zeroize"))] +mod tests { + use super::{SCOPED_FULL_RESULT_DROPS, ScopedFullResult}; + use common::typenum::U32; + use core::sync::atomic::Ordering; + + extern crate std; + + #[test] + fn scoped_full_result_zeroizes_on_return_and_unwind() { + SCOPED_FULL_RESULT_DROPS.store(0, Ordering::SeqCst); + drop(ScopedFullResult::::default()); + assert_eq!(SCOPED_FULL_RESULT_DROPS.load(Ordering::SeqCst), 1); + + SCOPED_FULL_RESULT_DROPS.store(0, Ordering::SeqCst); + let result = std::panic::catch_unwind(|| { + let _result = ScopedFullResult::::default(); + panic!("test-only finalization unwind"); + }); + assert!(result.is_err()); + assert_eq!(SCOPED_FULL_RESULT_DROPS.load(Ordering::SeqCst), 1); + } +} + impl Default for CtOutWrapper where T: VariableOutputCore,