use alloc::vec::Vec;
use core::fmt;
use core::sync::atomic;
#[cfg(unix)]
mod unix;
#[cfg(windows)]
mod windows;
pub(crate) struct LockedBytes {
data: Vec<u8>,
locked: bool,
}
impl LockedBytes {
pub(crate) fn from_slice(bytes: &[u8]) -> Self {
let mut data: Vec<u8> = Vec::with_capacity(bytes.len());
data.extend_from_slice(bytes);
debug_assert_eq!(data.capacity(), data.len());
let locked = if data.is_empty() {
false
} else {
unsafe { lock_pages(data.as_ptr(), data.len()) }
};
Self { data, locked }
}
pub(crate) fn as_bytes(&self) -> &[u8] {
&self.data
}
#[allow(dead_code)] pub(crate) fn len(&self) -> usize {
self.data.len()
}
#[allow(dead_code)] pub(crate) fn is_locked(&self) -> bool {
self.locked
}
}
impl Drop for LockedBytes {
fn drop(&mut self) {
if !self.data.is_empty() {
unsafe {
let ptr = self.data.as_mut_ptr();
for i in 0..self.data.len() {
core::ptr::write_volatile(ptr.add(i), 0u8);
}
}
atomic::compiler_fence(atomic::Ordering::SeqCst);
if self.locked {
unsafe {
unlock_pages(self.data.as_ptr(), self.data.len());
}
}
}
}
}
impl fmt::Debug for LockedBytes {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("LockedBytes")
.field("len", &self.data.len())
.field("locked", &self.locked)
.field("bytes", &"<redacted>")
.finish()
}
}
#[cfg(unix)]
use self::unix::{lock_pages, unlock_pages};
#[cfg(windows)]
use self::windows::{lock_pages, unlock_pages};
#[cfg(not(any(unix, windows)))]
unsafe fn lock_pages(_ptr: *const u8, _len: usize) -> bool {
false
}
#[cfg(not(any(unix, windows)))]
unsafe fn unlock_pages(_ptr: *const u8, _len: usize) {}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
mod tests {
use super::*;
use alloc::format;
#[test]
fn round_trips_bytes() {
let input = [0xa1, 0xb2, 0xc3, 0xd4, 0xe5];
let buf = LockedBytes::from_slice(&input);
assert_eq!(buf.as_bytes(), &input);
assert_eq!(buf.len(), 5);
}
#[test]
fn empty_buffer_is_unlocked() {
let buf = LockedBytes::from_slice(&[]);
assert_eq!(buf.len(), 0);
assert!(!buf.is_locked());
assert!(buf.as_bytes().is_empty());
}
#[test]
fn debug_is_redacted() {
let buf = LockedBytes::from_slice(&[0xde, 0xad, 0xbe, 0xef]);
let rendered = format!("{buf:?}");
assert!(rendered.contains("<redacted>"));
assert!(!rendered.contains("de"));
assert!(!rendered.contains("ad"));
assert!(rendered.contains("len"));
assert!(rendered.contains("locked"));
}
#[test]
fn many_small_buffers_do_not_leak_within_run() {
for size in [1, 7, 32, 64, 256, 4096] {
let bytes: Vec<u8> = (0..size).map(|i| (i & 0xff) as u8).collect();
let buf = LockedBytes::from_slice(&bytes);
assert_eq!(buf.as_bytes(), &bytes[..]);
drop(buf);
}
}
}