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 std::sync::atomic::{AtomicBool, Ordering};
#[test]
fn test_new_array() {
let x = Protected::new([0u8; 32]);
assert_eq!(x.0, [0u8; 32]);
}
#[test]
fn drop_zeroizes_inner() {
struct Tracked<'a>(&'a AtomicBool);
impl Zeroize for Tracked<'_> {
fn zeroize(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
let zeroized = AtomicBool::new(false);
{
let _p = Protected::new(Tracked(&zeroized));
assert!(!zeroized.load(Ordering::SeqCst));
}
assert!(
zeroized.load(Ordering::SeqCst),
"Protected::drop should zeroize the inner value"
);
}
#[test]
fn risky_unwrap_does_not_zeroize() {
struct Tracked<'a>(&'a AtomicBool);
impl Zeroize for Tracked<'_> {
fn zeroize(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
let zeroized = AtomicBool::new(false);
let _recovered = Protected::new(Tracked(&zeroized)).risky_unwrap();
assert!(
!zeroized.load(Ordering::SeqCst),
"risky_unwrap must not zeroize the value it hands back"
);
}
#[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]);
}
}