use super::Controlled;
use crate::private::ControlledPrivate;
use crate::OpaqueDebug;
use zeroize::{Zeroize, ZeroizeOnDrop};
#[derive(Zeroize, ZeroizeOnDrop, OpaqueDebug)]
pub struct Protected<T: Zeroize>(pub(crate) T);
impl<T: Zeroize> Protected<T> {
pub const fn new(x: T) -> Self {
Self(x)
}
fn into_inner_unchecked(self) -> T {
crate::move_inner_out(self)
}
}
unsafe impl<T: Zeroize> crate::MoveInner for Protected<T> {
type Inner = T;
fn inner_ptr(&self) -> *const T {
&self.0
}
}
impl<T: Zeroize> Protected<Protected<T>> {
#[inline]
pub fn flatten(self) -> Protected<T> {
self.into_inner_unchecked()
}
}
impl<T: Zeroize> Protected<Option<T>> {
#[inline]
pub fn transpose(self) -> Option<Protected<T>> {
self.into_inner_unchecked().map(Protected::new)
}
}
impl<T: Zeroize> ControlledPrivate for Protected<T> {}
impl<T> Controlled for Protected<T>
where
T: Zeroize,
{
fn risky_unwrap(self) -> Self::Inner {
self.into_inner_unchecked()
}
type Inner = T;
fn init_from_inner(x: Self::Inner) -> Self {
Self(x)
}
fn risky_ref(&self) -> &T {
&self.0
}
fn inner_mut(&mut self) -> &mut Self::Inner {
&mut self.0
}
}
impl<T> Clone for Protected<T>
where
T: Clone + Zeroize,
{
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
impl<T, A> Extend<A> for Protected<T>
where
T: Extend<A> + Zeroize,
{
fn extend<I>(&mut self, iter: I)
where
I: IntoIterator<Item = A>,
{
self.0.extend(iter);
}
}
#[cfg(feature = "arbitrary")]
impl<T> quickcheck::Arbitrary for Protected<T>
where
T: quickcheck::Arbitrary + Zeroize,
{
fn arbitrary(g: &mut quickcheck::Gen) -> Self {
let inner = T::arbitrary(g);
Self::new(inner)
}
}
pub fn flatten_array<const N: usize, T>(array: [Protected<T>; N]) -> Protected<[T; N]>
where
T: Zeroize,
{
Protected::new(array.map(|p| p.risky_unwrap()))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_util::{assert_zeroize_on_drop, Counted, Tracked};
use std::sync::atomic::{AtomicBool, AtomicUsize};
#[test]
fn test_new_array() {
let x = Protected::new([0u8; 32]);
assert_eq!(x.0, [0u8; 32]);
}
#[test]
fn drop_zeroizes_inner() {
let zeroized = AtomicBool::new(false);
let tracked = Tracked(&zeroized);
{
let _p = Protected::new(tracked);
assert!(!tracked.was_zeroized());
}
assert!(
tracked.was_zeroized(),
"Protected::drop should zeroize the inner value"
);
}
#[test]
fn zeroize_wipes_array_payload() {
let mut p = Protected::new([0xAB_u8; 32]);
p.zeroize();
assert_eq!(p.0, [0_u8; 32], "Protected::zeroize must wipe the array");
}
#[test]
fn zeroize_on_drop_is_backed_by_drop_glue() {
assert_zeroize_on_drop::<Protected<[u8; 32]>>();
assert_zeroize_on_drop::<Protected<Vec<u8>>>();
assert!(!std::mem::needs_drop::<[u8; 32]>());
assert!(std::mem::needs_drop::<Protected<[u8; 32]>>());
assert!(!std::mem::needs_drop::<u64>());
assert!(std::mem::needs_drop::<Protected<u64>>());
}
#[test]
fn nested_wrappers_keep_zeroizing_drop_glue() {
use crate::{Equatable, Exportable};
assert!(!std::mem::needs_drop::<[u8; 32]>());
assert_zeroize_on_drop::<Protected<[u8; 32]>>();
assert!(std::mem::needs_drop::<Protected<[u8; 32]>>());
assert_zeroize_on_drop::<Equatable<[u8; 32]>>();
assert!(std::mem::needs_drop::<Equatable<[u8; 32]>>());
assert_zeroize_on_drop::<Exportable<[u8; 32]>>();
assert!(std::mem::needs_drop::<Exportable<[u8; 32]>>());
assert_zeroize_on_drop::<Exportable<Equatable<Protected<[u8; 32]>>>>();
assert_zeroize_on_drop::<Equatable<Exportable<Protected<[u8; 32]>>>>();
assert!(std::mem::needs_drop::<
Exportable<Equatable<Protected<[u8; 32]>>>,
>());
assert!(std::mem::needs_drop::<
Equatable<Exportable<Protected<[u8; 32]>>>,
>());
}
#[test]
fn each_clone_is_wiped_independently() {
let wipes = AtomicUsize::new(0);
let counted = Counted(&wipes);
let original = Protected::new(counted.clone());
let dup = original.clone();
drop(original);
assert_eq!(counted.wipes(), 1, "original wiped on its own drop");
drop(dup);
assert_eq!(counted.wipes(), 2, "clone wiped independently");
}
#[test]
fn risky_unwrap_does_not_zeroize() {
let zeroized = AtomicBool::new(false);
let tracked = Tracked(&zeroized);
let _recovered = Protected::new(tracked).risky_unwrap();
assert!(
!tracked.was_zeroized(),
"risky_unwrap must not zeroize the value it hands back"
);
}
#[test]
fn flatten_and_transpose_move_without_zeroizing() {
let zeroized = AtomicBool::new(false);
let tracked = Tracked(&zeroized);
let flattened = Protected::new(Protected::new(tracked)).flatten();
assert!(
!tracked.was_zeroized(),
"flatten must not zeroize the value it hands on"
);
drop(flattened);
assert!(
tracked.was_zeroized(),
"the flattened wrapper still wipes on drop"
);
let zeroized = AtomicBool::new(false);
let tracked = Tracked(&zeroized);
let transposed = Protected::new(Some(tracked)).transpose();
assert!(
!tracked.was_zeroized(),
"transpose must not zeroize the value it hands on"
);
drop(transposed);
assert!(
tracked.was_zeroized(),
"the transposed wrapper still wipes on drop"
);
}
#[test]
fn test_opaque_debug() {
let x = Protected::new([0u8; 32]);
assert_eq!(
format!("{x:?}"),
"vitaminc_protected::protected::Protected<[u8; 32]>(\"***\")"
);
}
#[test]
fn test_flatten() {
let x = Protected::new(Protected::new([0u8; 32]));
let y = x.flatten();
assert_eq!(y.risky_unwrap(), [0u8; 32]);
}
#[test]
fn test_flatten_array() {
let x = Protected::new(1);
let y = Protected::new(2);
let z = Protected::new(3);
let array: [Protected<u8>; 3] = [x, y, z];
let flattened = flatten_array(array);
assert!(matches!(flattened, Protected(_)));
assert_eq!(flattened.risky_unwrap(), [1, 2, 3]);
}
}