use std::ops::{Deref, DerefMut};
#[cfg(unix)]
use crate::ffi::mem_noaccess;
#[cfg(target_os = "linux")]
use crate::ffi::{mem_no_dump, mem_wipe_on_fork};
use crate::{
MemoryError,
ffi::{mem_alloc, mem_dealloc, mem_lock, mem_readonly, mem_readwrite, mem_unlock},
ptr_ops::{ptr_deref, ptr_deref_mut, ptr_drop_in_place, ptr_fill_zero, secure_zero},
};
pub struct Cell<T> {
ptr: *mut T,
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum PartialState {
Allocated,
Locked,
Written,
}
struct PartialCell<T> {
ptr: *mut T,
len: usize,
state: PartialState,
}
impl<T> PartialCell<T> {
fn new(ptr: *mut T, len: usize) -> Self {
Self {
ptr,
len,
state: PartialState::Allocated,
}
}
fn mark_locked(&mut self) {
debug_assert!(self.state == PartialState::Allocated);
self.state = PartialState::Locked;
}
fn mark_written(&mut self) {
debug_assert!(self.state == PartialState::Locked);
self.state = PartialState::Written;
}
fn disarm(self) -> *mut T {
let ptr = self.ptr;
std::mem::forget(self);
ptr
}
}
impl<T> Drop for PartialCell<T> {
fn drop(&mut self) {
if self.state == PartialState::Written {
let _ = mem_readwrite(self.ptr, self.len);
ptr_drop_in_place(self.ptr);
ptr_fill_zero(self.ptr);
}
if matches!(self.state, PartialState::Locked | PartialState::Written) {
let _ = mem_unlock(self.ptr, self.len);
}
let _ = mem_dealloc(self.ptr, self.len);
}
}
impl<T> Cell<T> {
pub fn new(mut value: T) -> Result<Cell<T>, MemoryError> {
let len = std::mem::size_of::<T>();
if len == 0 {
return Err(MemoryError::from(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"zero-sized values cannot be placed in protected memory",
)));
}
let ptr = mem_alloc(len)?;
let mut guard = PartialCell::new(ptr, len);
mem_lock(ptr, len)?;
guard.mark_locked();
#[cfg(target_os = "linux")]
mem_no_dump(ptr, len)?;
#[cfg(target_os = "linux")]
mem_wipe_on_fork(ptr, len)?;
let val_ptr = &mut value as *mut T;
unsafe {
std::ptr::copy_nonoverlapping(val_ptr as *const u8, ptr as *mut u8, len);
}
guard.mark_written();
ptr_fill_zero(val_ptr);
std::mem::forget(value);
#[cfg(windows)]
mem_readonly(ptr, len)?;
#[cfg(unix)]
mem_noaccess(ptr, len)?;
Ok(Cell {
ptr: guard.disarm(),
})
}
pub fn low_priv(&mut self) -> Result<(), MemoryError> {
#[cfg(windows)]
let ret = self.read_only();
#[cfg(unix)]
let ret = self.no_access();
ret
}
#[cfg(unix)]
pub fn no_access(&mut self) -> Result<(), MemoryError> {
mem_noaccess(self.ptr, std::mem::size_of::<T>())
}
pub fn read_only(&mut self) -> Result<(), MemoryError> {
mem_readonly(self.ptr, std::mem::size_of::<T>())
}
pub fn read_write(&mut self) -> Result<(), MemoryError> {
mem_readwrite(self.ptr, std::mem::size_of::<T>())
}
}
impl<const N: usize> Cell<[u8; N]> {
pub fn new_with<F>(init: F) -> Result<Self, MemoryError>
where
F: FnOnce(&mut [u8; N]),
{
let len = std::mem::size_of::<[u8; N]>();
if len == 0 {
return Err(MemoryError::from(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"zero-sized values cannot be placed in protected memory",
)));
}
let ptr: *mut [u8; N] = mem_alloc(len)?;
let mut guard = PartialCell::new(ptr, len);
mem_lock(ptr, len)?;
guard.mark_locked();
#[cfg(target_os = "linux")]
mem_no_dump(ptr, len)?;
#[cfg(target_os = "linux")]
mem_wipe_on_fork(ptr, len)?;
guard.mark_written();
init(unsafe { &mut *ptr });
#[cfg(windows)]
mem_readonly(ptr, len)?;
#[cfg(unix)]
mem_noaccess(ptr, len)?;
Ok(Cell {
ptr: guard.disarm(),
})
}
pub fn from_bytes<T: AsMut<[u8]>>(mut bytes: T) -> Result<Self, (T, MemoryError)> {
let len = bytes.as_mut().len();
if len > N {
return Err((
bytes,
MemoryError::from(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"byte slice exceeds buffer size",
)),
));
}
Self::new_with(|page| {
let slice = bytes.as_mut();
unsafe {
std::ptr::copy_nonoverlapping(slice.as_ptr(), page.as_mut_ptr(), len);
}
secure_zero(slice);
})
.map_err(|e| (bytes, e))
}
}
impl<T> Deref for Cell<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
ptr_deref(self.ptr)
}
}
impl<T> DerefMut for Cell<T> {
fn deref_mut(&mut self) -> &mut Self::Target {
ptr_deref_mut(self.ptr)
}
}
impl<T> Drop for Cell<T> {
fn drop(&mut self) {
let len = std::mem::size_of::<T>();
if mem_readwrite(self.ptr, len).is_err() {
return;
}
ptr_drop_in_place(self.ptr);
ptr_fill_zero(self.ptr);
let _ = mem_unlock(self.ptr, len);
let _ = mem_dealloc(self.ptr, len);
}
}