use crate::{AsProtectedRef, Choice, OpaqueDebug, ProtectedRef, TimingSafeEq, Zeroed};
use core::alloc::Layout;
use core::any::type_name;
use core::fmt;
use core::marker::PhantomData;
use core::mem::{self, ManuallyDrop, MaybeUninit};
use core::ptr;
use core::sync::atomic::{AtomicU8, Ordering};
use std::io;
use zeroize::{Zeroize, ZeroizeOnDrop};
mod layout;
#[cfg(kani)]
mod proofs;
#[cfg(all(test, target_os = "linux", not(miri)))]
#[path = "../../tests/common/mod.rs"]
mod common;
#[cfg(all(unix, not(kani)))]
mod unix;
#[cfg(all(unix, not(kani)))]
use unix::Region;
#[cfg(any(test, not(unix), kani))]
mod fallback;
#[cfg(any(not(unix), kani))]
use fallback::Region;
#[derive(Debug)]
#[non_exhaustive]
pub enum LockError {
Map {
bytes: usize,
source: io::Error,
},
Guard {
source: io::Error,
},
Refused {
bytes: usize,
limit: Option<u64>,
source: io::Error,
},
Dump {
source: io::Error,
},
Unavailable,
Forked,
Untracked {
source: io::Error,
},
Alignment {
align: usize,
page: usize,
},
}
impl fmt::Display for LockError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Map { bytes, source } => {
write!(f, "could not map {bytes} bytes for locked storage: {source}")
}
Self::Guard { source } => write!(f, "could not protect a guard page: {source}"),
Self::Refused {
bytes,
limit: Some(limit),
source,
} => write!(
f,
"locking {bytes} bytes was refused ({source}); the RLIMIT_MEMLOCK soft limit is {limit} bytes"
),
Self::Refused {
bytes,
limit: None,
source,
} => write!(f, "locking {bytes} bytes was refused ({source})"),
Self::Dump { source } => {
write!(f, "excluding the region from core dumps was refused ({source})")
}
Self::Unavailable => f.write_str("memory locking is not available on this platform"),
Self::Forked => f.write_str(
"the lock belongs to the process that created the value; a forked child inherits the memory but not the lock",
),
Self::Untracked { source } => write!(
f,
"forks cannot be tracked in this process ({source}), so a forked child would misreport the lock"
),
Self::Alignment { align, page } => write!(
f,
"alignment {align} exceeds the page size {page}; locked storage cannot hold this type"
),
}
}
}
impl std::error::Error for LockError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Map { source, .. }
| Self::Guard { source }
| Self::Refused { source, .. }
| Self::Dump { source }
| Self::Untracked { source } => Some(source),
Self::Unavailable | Self::Forked | Self::Alignment { .. } => None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum LockPolicy {
#[default]
BestEffort,
Strict,
}
const POLICY_UNSET: u8 = 0;
const POLICY_BEST_EFFORT: u8 = 1;
const POLICY_STRICT: u8 = 2;
static POLICY: AtomicU8 = AtomicU8::new(POLICY_UNSET);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct LockPolicyError {
pub current: LockPolicy,
}
impl fmt::Display for LockPolicyError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "the lock policy is already set to {:?}", self.current)
}
}
impl std::error::Error for LockPolicyError {}
impl LockPolicy {
fn encode(self) -> u8 {
match self {
Self::BestEffort => POLICY_BEST_EFFORT,
Self::Strict => POLICY_STRICT,
}
}
fn decode(raw: u8) -> Self {
match raw {
POLICY_STRICT => Self::Strict,
_ => Self::BestEffort,
}
}
pub fn set(self) -> Result<(), LockPolicyError> {
let wanted = self.encode();
match POLICY.compare_exchange(POLICY_UNSET, wanted, Ordering::AcqRel, Ordering::Acquire) {
Ok(_) => Ok(()),
Err(current) if current == wanted => Ok(()),
Err(current) => Err(LockPolicyError {
current: Self::decode(current),
}),
}
}
pub fn current() -> Self {
Self::decode(POLICY.load(Ordering::Acquire))
}
}
pub struct Locked<T: Zeroize> {
storage: Storage<T>,
}
unsafe impl<T: Zeroize + Send> Send for Locked<T> {}
unsafe impl<T: Zeroize + Sync> Sync for Locked<T> {}
struct Storage<T> {
region: Region,
lock: Option<Box<LockError>>,
_value: PhantomData<T>,
}
impl<T> Storage<T> {
fn allocate() -> Result<Self, LockError> {
let (region, lock) = Region::allocate(Layout::new::<T>())?;
if LockPolicy::current() == LockPolicy::Strict {
if let Some(err) = lock {
return Err(err);
}
}
Ok(Self {
region,
lock: lock.map(Box::new),
_value: PhantomData,
})
}
fn as_ptr(&self) -> *mut T {
self.region
.value_ptr(Layout::new::<T>())
.as_ptr()
.cast::<T>()
}
unsafe fn filled(self) -> Locked<T>
where
T: Zeroize,
{
Locked { storage: self }
}
unsafe fn move_in(self, slot: &mut ManuallyDrop<T>) -> Locked<T>
where
T: Zeroize,
{
ptr::copy_nonoverlapping(&**slot as *const T, self.as_ptr(), 1);
wipe_bytes(&mut **slot as *mut T);
self.filled()
}
}
impl<T: Zeroize> Locked<T> {
fn as_ptr(&self) -> *mut T {
self.storage.as_ptr()
}
pub fn new(mut value: T) -> Result<Self, LockError> {
let storage = match Storage::allocate() {
Ok(storage) => storage,
Err(err) => {
value.zeroize();
return Err(err);
}
};
let mut slot = ManuallyDrop::new(value);
Ok(unsafe { storage.move_in(&mut slot) })
}
pub fn generate<F>(f: F) -> Result<Self, LockError>
where
T: Zeroed,
F: FnOnce(&mut T),
{
let mut this = Self::zeroed()?;
f(this.risky_mut());
Ok(this)
}
pub fn try_generate<F, E>(f: F) -> Result<Self, E>
where
T: Zeroed,
E: From<LockError>,
F: FnOnce(&mut T) -> Result<(), E>,
{
let mut this = Self::zeroed()?;
f(this.risky_mut())?;
Ok(this)
}
pub fn zeroed() -> Result<Self, LockError>
where
T: Zeroed,
{
let storage = Storage::allocate()?;
let zeroed = T::zeroed();
unsafe {
ptr::write(storage.as_ptr(), zeroed);
Ok(storage.filled())
}
}
pub fn locked(&self) -> bool {
self.storage.lock.is_none() && self.storage.region.same_process()
}
pub fn lock_error(&self) -> Option<&LockError> {
static FORKED: LockError = LockError::Forked;
if !self.storage.region.same_process() {
return Some(&FORKED);
}
self.storage.lock.as_deref()
}
pub fn relock(&mut self) -> Result<(), &LockError> {
self.storage.lock = self.storage.region.relock().map(Box::new);
match self.storage.lock.as_deref() {
None => Ok(()),
Some(err) => Err(err),
}
}
pub fn require_locked(mut self) -> Result<Self, LockError> {
if !self.storage.region.same_process() {
return Err(LockError::Forked);
}
match self.storage.lock.take() {
None => Ok(self),
Some(err) => Err(*err),
}
}
pub fn risky_ref(&self) -> &T {
unsafe { &*self.as_ptr() }
}
pub fn risky_mut(&mut self) -> &mut T {
unsafe { &mut *self.as_ptr() }
}
pub fn with<R>(&self, f: impl FnOnce(&T) -> R) -> R {
f(self.risky_ref())
}
pub fn update(&mut self, f: impl FnOnce(&mut T)) {
f(self.risky_mut())
}
pub fn try_clone(&self) -> Result<Self, LockError>
where
T: Clone,
{
Self::new(self.risky_ref().clone())
}
#[cfg(test)]
fn into_wiped_region(self) -> Region {
let mut this = ManuallyDrop::new(self);
unsafe {
(*this.as_ptr()).zeroize();
ptr::drop_in_place(this.as_ptr());
}
this.storage.region.wipe();
drop(this.storage.lock.take());
unsafe { ptr::read(&this.storage.region) }
}
}
pub(super) unsafe fn wipe_raw(ptr: *mut u8, len: usize) {
let bytes = core::slice::from_raw_parts_mut(ptr.cast::<MaybeUninit<u8>>(), len);
for b in bytes {
ptr::write_volatile(b, MaybeUninit::new(0));
}
core::sync::atomic::compiler_fence(Ordering::SeqCst);
}
unsafe fn wipe_bytes<T>(slot: *mut T) {
wipe_raw(slot.cast::<u8>(), mem::size_of::<T>());
}
impl<T: Zeroize> Drop for Locked<T> {
fn drop(&mut self) {
struct DropValue<T>(*mut T);
impl<T> Drop for DropValue<T> {
fn drop(&mut self) {
unsafe { ptr::drop_in_place(self.0) };
}
}
let value = DropValue(self.as_ptr());
unsafe { (*value.0).zeroize() };
drop(value);
}
}
impl<T: Zeroize> Zeroize for Locked<T> {
fn zeroize(&mut self) {
self.risky_mut().zeroize();
}
}
impl<T: Zeroize> ZeroizeOnDrop for Locked<T> {}
impl<T: Zeroize> OpaqueDebug for Locked<T> {}
impl<T: Zeroize> fmt::Debug for Locked<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"Locked<{}>({})",
type_name::<T>(),
if self.locked() { "locked" } else { "unlocked" }
)
}
}
impl<T: Zeroize + TimingSafeEq> TimingSafeEq for Locked<T> {
fn ts_eq(&self, other: &Self) -> Choice {
self.risky_ref().ts_eq(other.risky_ref())
}
}
impl<'a, T: Zeroize> AsProtectedRef<'a, T> for Locked<T> {
fn as_protected_ref(&'a self) -> ProtectedRef<'a, T> {
ProtectedRef(self.risky_ref())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_value_moved_in_reads_back() {
let key = Locked::new([0xA5u8; 32]).unwrap();
assert_eq!(key.risky_ref(), &[0xA5u8; 32]);
}
#[test]
fn a_generated_value_is_built_in_place() {
let key = Locked::<[u8; 32]>::generate(|k| k.fill(9)).unwrap();
assert_eq!(key.risky_ref(), &[9u8; 32]);
assert!(std::ptr::eq(
key.risky_ref().as_ptr(),
key.storage
.region
.value_ptr(Layout::new::<[u8; 32]>())
.as_ptr()
));
}
#[test]
fn a_failed_generation_returns_the_callers_error() {
#[derive(Debug, PartialEq)]
enum E {
Lock,
Rng,
}
impl From<LockError> for E {
fn from(_: LockError) -> Self {
E::Lock
}
}
let r: Result<Locked<[u8; 16]>, E> = Locked::try_generate(|_| Err(E::Rng));
assert_eq!(r.unwrap_err(), E::Rng);
}
#[test]
fn update_changes_the_value_in_place() {
let mut key = Locked::new([0u8; 8]).unwrap();
let before = key.risky_ref().as_ptr();
key.update(|k| k[3] = 42);
assert_eq!(key.risky_ref()[3], 42);
assert_eq!(before, key.risky_ref().as_ptr());
}
#[test]
fn the_region_is_all_zero_after_the_wipe() {
let key = Locked::new([0xFFu8; 64]).unwrap();
let region = key.into_wiped_region();
let bytes = unsafe {
core::slice::from_raw_parts(region.value_ptr(Layout::new::<[u8; 64]>()).as_ptr(), 64)
};
assert!(bytes.iter().all(|&b| b == 0));
}
#[test]
fn the_value_owns_its_own_drop() {
use std::sync::atomic::{AtomicUsize, Ordering};
static DROPS: AtomicUsize = AtomicUsize::new(0);
struct Counted(u8);
impl Zeroize for Counted {
fn zeroize(&mut self) {
self.0 = 0;
}
}
impl Drop for Counted {
fn drop(&mut self) {
let _ = DROPS.fetch_add(1, Ordering::SeqCst);
}
}
drop(Locked::new(Counted(1)).unwrap());
assert_eq!(DROPS.load(Ordering::SeqCst), 1);
}
#[test]
fn a_panicking_zeroed_releases_the_region_without_dropping_a_value() {
struct Boom(#[allow(dead_code)] Box<u8>);
impl Zeroize for Boom {
fn zeroize(&mut self) {
*self.0 = 0;
}
}
impl Zeroed for Boom {
fn zeroed() -> Self {
panic!("no zero value");
}
}
let r = std::panic::catch_unwind(Locked::<Boom>::zeroed);
assert!(r.is_err());
let r = std::panic::catch_unwind(|| Locked::<Boom>::generate(|_| {}));
assert!(r.is_err());
}
#[test]
fn what_the_value_owns_elsewhere_is_zeroized_before_it_is_dropped() {
use std::sync::atomic::{AtomicBool, Ordering};
static ZEROIZED_FIRST: AtomicBool = AtomicBool::new(false);
struct Owning(Vec<u8>);
impl Zeroize for Owning {
fn zeroize(&mut self) {
self.0.zeroize();
}
}
impl Drop for Owning {
fn drop(&mut self) {
ZEROIZED_FIRST.store(self.0.is_empty(), Ordering::SeqCst);
}
}
drop(Locked::new(Owning(vec![0xAB; 64])).unwrap());
assert!(ZEROIZED_FIRST.load(Ordering::SeqCst));
}
#[test]
fn a_panicking_zeroize_still_drops_the_value_and_releases_the_region() {
use std::sync::atomic::{AtomicUsize, Ordering};
static DROPS: AtomicUsize = AtomicUsize::new(0);
struct Stubborn(#[allow(dead_code)] [u8; 8]);
impl Zeroize for Stubborn {
fn zeroize(&mut self) {
panic!("will not be wiped");
}
}
impl Drop for Stubborn {
fn drop(&mut self) {
let _ = DROPS.fetch_add(1, Ordering::SeqCst);
}
}
let key = Locked::new(Stubborn([1; 8])).unwrap();
let r = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| drop(key)));
assert!(r.is_err());
assert_eq!(DROPS.load(Ordering::SeqCst), 1);
}
#[test]
fn a_panicking_destructor_still_releases_the_region_exactly_once() {
use std::sync::atomic::{AtomicUsize, Ordering};
static DROPS: AtomicUsize = AtomicUsize::new(0);
struct Angry([u8; 16]);
impl Zeroize for Angry {
fn zeroize(&mut self) {
self.0.zeroize();
}
}
impl Drop for Angry {
fn drop(&mut self) {
let _ = DROPS.fetch_add(1, Ordering::SeqCst);
panic!("angry");
}
}
let key = Locked::new(Angry([0xEE; 16])).unwrap();
let r = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| drop(key)));
assert!(r.is_err());
assert_eq!(DROPS.load(Ordering::SeqCst), 1);
}
#[cfg(all(unix, not(miri)))]
#[test]
fn a_failed_allocation_zeroizes_the_input() {
use std::sync::atomic::{AtomicBool, Ordering};
static ZEROIZED: AtomicBool = AtomicBool::new(false);
#[repr(align(131072))]
struct Wide(u8);
impl Zeroize for Wide {
fn zeroize(&mut self) {
self.0 = 0;
ZEROIZED.store(true, Ordering::SeqCst);
}
}
impl Drop for Wide {
fn drop(&mut self) {
assert_eq!(self.0, 0, "dropped before being zeroized");
}
}
std::thread::Builder::new()
.stack_size(16 << 20)
.spawn(|| {
let r = Locked::new(Wide(0x5A));
assert!(matches!(r, Err(LockError::Alignment { .. })), "{r:?}");
})
.unwrap()
.join()
.unwrap();
assert!(ZEROIZED.load(Ordering::SeqCst));
}
#[test]
fn every_error_displays_and_only_os_errors_have_a_source() {
use std::error::Error;
let os = || io::Error::from_raw_os_error(1);
let errors = [
LockError::Map {
bytes: 3,
source: os(),
},
LockError::Guard { source: os() },
LockError::Refused {
bytes: 3,
limit: Some(0),
source: os(),
},
LockError::Refused {
bytes: 3,
limit: None,
source: os(),
},
LockError::Dump { source: os() },
LockError::Untracked { source: os() },
LockError::Unavailable,
LockError::Forked,
LockError::Alignment { align: 8, page: 4 },
];
for err in &errors {
assert!(!err.to_string().is_empty());
let has_source = !matches!(
err,
LockError::Unavailable | LockError::Forked | LockError::Alignment { .. }
);
assert_eq!(err.source().is_some(), has_source, "{err}");
}
assert!(errors[4].to_string().contains("core dumps"));
}
#[test]
fn zeroize_clears_the_value_in_place() {
let mut key = Locked::new([0x77u8; 24]).unwrap();
let before = key.risky_ref().as_ptr();
key.zeroize();
assert_eq!(key.risky_ref(), &[0u8; 24]);
assert_eq!(before, key.risky_ref().as_ptr());
}
#[test]
fn the_policy_error_names_the_policy_in_force() {
let err = LockPolicyError {
current: LockPolicy::Strict,
};
assert_eq!(err.to_string(), "the lock policy is already set to Strict");
}
#[derive(Zeroize, Clone, PartialEq, Eq, Debug)]
#[repr(C)]
struct Padded {
a: u8,
b: u32,
c: u16,
}
impl Zeroed for Padded {
fn zeroed() -> Self {
Padded { a: 0, b: 0, c: 0 }
}
}
const PADDED_SIZE: usize = mem::size_of::<Padded>();
const _: () = assert!(
PADDED_SIZE > 1 + 4 + 2,
"the type must actually have padding"
);
#[test]
fn a_padded_value_moves_in_clones_generates_zeroizes_and_wipes() {
const SIZE: usize = PADDED_SIZE;
let key = Locked::new(Padded { a: 1, b: 2, c: 3 }).unwrap();
assert_eq!(*key.risky_ref(), Padded { a: 1, b: 2, c: 3 });
let copy = key.try_clone().unwrap();
assert_eq!(copy.risky_ref(), key.risky_ref());
let mut made = Locked::<Padded>::generate(|p| p.b = 9).unwrap();
assert_eq!(made.risky_ref().b, 9);
made.zeroize();
assert_eq!(*made.risky_ref(), Padded::zeroed());
let region = key.into_wiped_region();
let bytes = unsafe {
core::slice::from_raw_parts(region.value_ptr(Layout::new::<Padded>()).as_ptr(), SIZE)
};
assert!(bytes.iter().all(|&b| b == 0));
}
#[test]
fn move_in_wipes_the_source_slot() {
let mut slot = ManuallyDrop::new([0xA5u8; 16]);
let storage = Storage::<[u8; 16]>::allocate().unwrap();
let key = unsafe { storage.move_in(&mut slot) };
assert_eq!(key.risky_ref(), &[0xA5u8; 16]);
let left_behind: [u8; 16] = unsafe { ptr::read((&*slot as *const [u8; 16]).cast()) };
assert_eq!(left_behind, [0u8; 16]);
}
#[test]
fn new_moves_without_running_drop_on_the_source() {
use std::sync::atomic::{AtomicUsize, Ordering};
static DROPS: AtomicUsize = AtomicUsize::new(0);
struct Counted([u8; 4]);
impl Zeroize for Counted {
fn zeroize(&mut self) {
self.0.zeroize();
}
}
impl Drop for Counted {
fn drop(&mut self) {
assert_eq!(self.0, [0; 4], "drop runs after zeroize");
let _ = DROPS.fetch_add(1, Ordering::SeqCst);
}
}
let l = Locked::new(Counted([1, 2, 3, 4])).unwrap();
assert_eq!(l.risky_ref().0, [1, 2, 3, 4]);
assert_eq!(
DROPS.load(Ordering::SeqCst),
0,
"the source slot is not dropped"
);
drop(l);
assert_eq!(DROPS.load(Ordering::SeqCst), 1);
}
#[test]
fn try_clone_is_a_second_region() {
let a = Locked::new([3u8; 16]).unwrap();
let b = a.try_clone().unwrap();
assert_eq!(a.risky_ref(), b.risky_ref());
assert_ne!(a.storage.region.ptr(), b.storage.region.ptr());
}
#[test]
fn relock_in_the_same_process_reports_the_known_state() {
let mut key = Locked::new([1u8; 8]).unwrap();
let before = key.locked();
assert_eq!(key.relock().is_ok(), before);
assert_eq!(key.locked(), before);
}
#[test]
fn require_locked_passes_a_locked_value_through() {
let key = Locked::new([1u8; 8]).unwrap();
if key.locked() {
assert!(key.require_locked().is_ok());
} else {
assert!(key.require_locked().is_err());
}
}
#[test]
fn debug_reveals_nothing_but_the_type_and_lock_state() {
let key = Locked::new([0x42u8; 8]).unwrap();
let s = format!("{key:?}");
assert!(s.starts_with("Locked<[u8; 8]>("), "{s}");
assert!(!s.contains("42"));
}
#[test]
fn timing_safe_eq_compares_the_values() {
let a = Locked::new([1u8; 8]).unwrap();
let b = Locked::new([1u8; 8]).unwrap();
let c = Locked::new([2u8; 8]).unwrap();
assert!(bool::from(a.ts_eq(&b)));
assert!(!bool::from(a.ts_eq(&c)));
}
#[test]
fn zero_sized_and_odd_sized_types_are_fine() {
#[derive(Zeroize)]
struct Nothing;
let _ = Locked::new(Nothing).unwrap();
let odd = Locked::new([7u8; 33]).unwrap();
assert_eq!(odd.risky_ref().len(), 33);
}
#[test]
fn the_handle_is_send_and_sync() {
fn assert_send_sync<S: Send + Sync>() {}
assert_send_sync::<Locked<[u8; 32]>>();
}
#[test]
fn the_handle_is_a_few_words_whatever_the_value_size() {
let words = mem::size_of::<Locked<[u8; 1024]>>() / mem::size_of::<usize>();
assert!(words <= 4, "{words} words");
assert_eq!(
mem::size_of::<Locked<[u8; 1024]>>(),
mem::size_of::<Locked<u8>>()
);
}
#[cfg(all(unix, not(miri)))]
fn lock_required() -> bool {
std::env::var_os("LOCKED_TESTS_REQUIRE_LOCK").is_some()
}
#[cfg(all(unix, not(miri)))]
#[test]
fn a_small_value_is_locked_when_the_limit_allows() {
let key = Locked::new([1u8; 32]).unwrap();
let constrained = matches!(unix::memlock_limit(), Some(limit) if limit < 1 << 20);
if lock_required() || !constrained {
assert!(key.locked(), "{:?}", key.lock_error());
}
}
#[cfg(all(target_os = "linux", not(miri)))]
#[test]
fn the_kernel_reports_the_region_locked_and_not_dumpable() {
let key = Locked::new([1u8; 32]).unwrap();
assert!(
!matches!(key.lock_error(), Some(LockError::Dump { .. })),
"{:?}",
key.lock_error()
);
let expect_locked = lock_required() || (key.locked() && !super::common::mlock_is_a_no_op());
let addr = key.storage.region.ptr().as_ptr() as usize;
if expect_locked {
let locked_kb = super::common::smaps_field(addr, "Locked:");
assert!(locked_kb.unwrap_or(0) > 0, "Locked: {locked_kb:?}");
}
let flags = super::common::vm_flags(addr);
assert!(flags.split(' ').any(|f| f == "dd"), "VmFlags: {flags}");
}
}