use crate::error::Result;
use crate::zeroize::Zeroize;
use core::mem::ManuallyDrop;
pub struct Secret<T: Zeroize> {
inner: ManuallyDrop<T>,
locked: bool,
}
impl<T: Zeroize> Secret<T> {
#[inline]
pub fn new(val: T) -> Self {
Self {
inner: ManuallyDrop::new(val),
locked: false,
}
}
#[inline]
pub fn lock(mut self) -> Result<Self> {
let size = core::mem::size_of::<T>();
if size == 0 {
return Ok(self);
}
let ptr = &*self.inner as *const T as *const u8;
unsafe {
crate::mlock::lock(ptr, size)?;
}
self.locked = true;
Ok(self)
}
#[inline]
pub fn expose<R, F: FnOnce(&T) -> R>(&self, f: F) -> R {
f(&self.inner)
}
#[inline]
pub fn expose_mut<R, F: FnOnce(&mut T) -> R>(&mut self, f: F) -> R {
f(&mut self.inner)
}
#[inline]
pub fn is_locked(&self) -> bool {
self.locked
}
}
impl<T: Zeroize> Drop for Secret<T> {
fn drop(&mut self) {
if self.locked {
let size = core::mem::size_of::<T>();
if size > 0 {
let ptr = &*self.inner as *const T as *const u8;
let _ = unsafe { crate::mlock::unlock(ptr, size) };
}
}
(*self.inner).zeroize();
unsafe {
ManuallyDrop::drop(&mut self.inner);
}
}
}
#[cfg(feature = "alloc")]
mod boxed {
use super::*;
use alloc::boxed::Box;
pub struct SecretBox<T: Zeroize> {
inner: ManuallyDrop<Box<T>>,
locked: bool,
}
impl<T: Zeroize> SecretBox<T> {
#[inline]
pub fn new(val: T) -> Result<Self> {
let bx = Box::new(val);
let size = core::mem::size_of::<T>();
let mut secret = Self {
inner: ManuallyDrop::new(bx),
locked: false,
};
if size > 0 {
let ptr = &**secret.inner as *const T as *const u8;
match unsafe { crate::mlock::lock(ptr, size) } {
Ok(()) => secret.locked = true,
Err(e) => {
unsafe { ManuallyDrop::drop(&mut secret.inner) };
return Err(e);
}
}
}
Ok(secret)
}
#[inline]
pub fn expose<R, F: FnOnce(&T) -> R>(&self, f: F) -> R {
f(&self.inner)
}
#[inline]
pub fn expose_mut<R, F: FnOnce(&mut T) -> R>(&mut self, f: F) -> R {
f(&mut self.inner)
}
#[inline]
pub fn is_locked(&self) -> bool {
self.locked
}
}
impl<T: Zeroize> Drop for SecretBox<T> {
fn drop(&mut self) {
if self.locked {
let size = core::mem::size_of::<T>();
if size > 0 {
let ptr = &**self.inner as *const T as *const u8;
let _ = unsafe { crate::mlock::unlock(ptr, size) };
}
}
(*self.inner).zeroize();
unsafe {
ManuallyDrop::drop(&mut self.inner);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[cfg_attr(miri, ignore)]
fn secretbox_basic() {
let key = SecretBox::new([0xABu8; 32]).unwrap();
key.expose(|k| {
assert!(k.iter().all(|&b| b == 0xAB));
});
}
}
}
#[cfg(feature = "alloc")]
pub use boxed::SecretBox;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn secret_basic() {
let mut key = Secret::new([0xFFu8; 16]);
key.expose(|k| {
assert!(k.iter().all(|&b| b == 0xFF));
});
key.expose_mut(|k| {
k[0] = 0x00;
});
key.expose(|k| {
assert_eq!(k[0], 0x00);
assert_eq!(k[1], 0xFF);
});
}
#[test]
fn secret_drops_without_panic() {
{
let _key = Secret::new([0x42u8; 32]);
}
}
#[test]
#[cfg_attr(miri, ignore)]
fn secret_lock_unlock() {
let key = Secret::new([0u8; 4096]);
if let Ok(locked) = key.lock() {
assert!(locked.is_locked());
}
}
}