1#[cfg(feature = "nightly")]
156use std::alloc::{AllocError, Allocator};
157use std::fmt;
158use std::marker::PhantomData;
159use std::ptr::{self, NonNull};
160use std::sync::LazyLock;
161
162use subtle::ConstantTimeEq;
163use zeroize::{Zeroize, ZeroizeOnDrop};
164
165use crate::error;
166use crate::rng::copy_randombytes;
167pub use crate::types::*;
168
169mod int {
170 #[derive(Clone, Debug, PartialEq, Eq)]
171 pub(super) enum LockMode {
172 Locked,
173 Unlocked,
174 }
175
176 #[derive(Clone, Debug, PartialEq, Eq)]
177 pub(super) enum ProtectMode {
178 ReadOnly,
179 ReadWrite,
180 NoAccess,
181 }
182
183 #[derive(Clone, Copy)]
190 pub(super) struct Region {
191 addr: usize,
192 pub(super) len: usize,
193 }
194
195 impl Region {
196 pub(super) fn of(bytes: &[u8]) -> Self {
197 Self {
198 addr: bytes.as_ptr().addr(),
199 len: bytes.len(),
200 }
201 }
202
203 pub(super) fn ptr(self) -> *mut u8 {
204 core::ptr::without_provenance_mut(self.addr)
205 }
206 }
207
208 pub(super) struct InternalData<A> {
209 pub(super) a: A,
210 pub(super) lm: LockMode,
211 pub(super) pm: ProtectMode,
212 pub(super) noaccess_region: Region,
216 }
217
218 impl<A: crate::types::Bytes> InternalData<A> {
219 pub(super) fn region(&self) -> Region {
224 match self.pm {
225 ProtectMode::NoAccess => self.noaccess_region,
226 ProtectMode::ReadOnly | ProtectMode::ReadWrite => Region::of(self.a.as_slice()),
227 }
228 }
229 }
230}
231
232mod sealed {
233 pub trait Sealed {}
237}
238
239pub mod traits {
250 use super::sealed::Sealed;
251
252 pub trait ProtectMode: Sealed {}
257 pub struct ReadOnly;
259 pub struct ReadWrite;
261 pub struct NoAccess;
263
264 impl Sealed for ReadOnly {}
265 impl Sealed for ReadWrite {}
266 impl Sealed for NoAccess {}
267 impl ProtectMode for ReadOnly {}
268 impl ProtectMode for ReadWrite {}
269 impl ProtectMode for NoAccess {}
270
271 pub trait LockMode: Sealed {}
276 pub struct Locked;
279 pub struct Unlocked;
281
282 impl Sealed for Locked {}
283 impl Sealed for Unlocked {}
284 impl LockMode for Locked {}
285 impl LockMode for Unlocked {}
286}
287
288pub trait Lockable<A: Zeroize + Bytes> {
291 fn mlock(self) -> Result<Protected<A, traits::ReadWrite, traits::Locked>, error::Error>;
303}
304
305pub trait Lock<A: Zeroize + Bytes, PM: traits::ProtectMode> {
307 fn mlock(self) -> Result<Protected<A, PM, traits::Locked>, error::Error>;
317}
318
319pub trait Unlock<A: Zeroize + Bytes, PM: traits::ProtectMode> {
321 fn munlock(self) -> Result<Protected<A, PM, traits::Unlocked>, error::Error>;
328}
329
330pub trait ProtectReadOnly<A: Zeroize + Bytes, PM: traits::ProtectMode, LM: traits::LockMode> {
332 fn mprotect_readonly(self) -> Result<Protected<A, traits::ReadOnly, LM>, error::Error>;
339}
340
341pub trait ProtectReadWrite<A: Zeroize + Bytes, PM: traits::ProtectMode, LM: traits::LockMode> {
343 fn mprotect_readwrite(self) -> Result<Protected<A, traits::ReadWrite, LM>, error::Error>;
350}
351
352pub trait ProtectNoAccess<A: Zeroize + Bytes, PM: traits::ProtectMode> {
354 fn mprotect_noaccess(
361 self,
362 ) -> Result<Protected<A, traits::NoAccess, traits::Unlocked>, error::Error>;
363}
364
365pub trait NewLocked<A: Zeroize + NewBytes + Lockable<A>> {
367 fn new_locked() -> Result<Protected<A, traits::ReadWrite, traits::Locked>, error::Error>;
374 fn new_readonly_locked() -> Result<Protected<A, traits::ReadOnly, traits::Locked>, error::Error>;
381 fn generate_locked() -> Result<Protected<A, traits::ReadWrite, traits::Locked>, error::Error>;
387 fn generate_readonly_locked()
394 -> Result<Protected<A, traits::ReadOnly, traits::Locked>, error::Error>;
395}
396
397pub trait NewLockedFromSlice<A: Zeroize + NewBytes + Lockable<A>> {
399 fn from_slice_into_locked(
411 src: &[u8],
412 ) -> Result<Protected<A, traits::ReadWrite, traits::Locked>, crate::error::Error>;
413 fn from_slice_into_readonly_locked(
426 src: &[u8],
427 ) -> Result<Protected<A, traits::ReadOnly, traits::Locked>, crate::error::Error>;
428}
429
430pub struct Protected<A: Zeroize + Bytes, PM: traits::ProtectMode, LM: traits::LockMode> {
434 i: Option<int::InternalData<A>>,
435 p: PhantomData<PM>,
436 l: PhantomData<LM>,
437}
438
439mod ptypes {
441 pub type Locked<T> = super::Protected<T, super::traits::ReadWrite, super::traits::Locked>;
443 pub type LockedRO<T> = super::Protected<T, super::traits::ReadOnly, super::traits::Locked>;
445 pub type NoAccess<T> = super::Protected<T, super::traits::NoAccess, super::traits::Unlocked>;
447 pub type Unlocked<T> = super::Protected<T, super::traits::ReadWrite, super::traits::Unlocked>;
449 pub type UnlockedRO<T> = super::Protected<T, super::traits::ReadOnly, super::traits::Unlocked>;
451 pub type LockedBytes = Locked<super::HeapBytes>;
453}
454
455fn clone_into_locked<
459 S: Bytes,
460 T: Zeroize + NewBytes + ResizableBytes + Lockable<T> + NewLocked<T>,
461>(
462 src: &S,
463) -> Protected<T, traits::ReadWrite, traits::Locked> {
464 let mut cloned = T::new_locked().expect("unable to create new locked instance");
465 cloned.resize(src.len(), 0);
466 cloned.as_mut_slice().copy_from_slice(src.as_slice());
467 cloned
468}
469
470impl<T: Zeroize + NewBytes + ResizableBytes + Lockable<T> + NewLocked<T>> Clone for Locked<T> {
471 fn clone(&self) -> Self {
472 clone_into_locked(self)
473 }
474}
475
476impl<T: Zeroize + NewBytes + ResizableBytes + Lockable<T> + NewLocked<T>> Clone for LockedRO<T> {
477 fn clone(&self) -> Self {
478 clone_into_locked(self)
479 .mprotect_readonly()
480 .expect("unable to protect readonly")
481 }
482}
483
484impl<T: Zeroize + Bytes + Clone> Clone for Unlocked<T> {
485 fn clone(&self) -> Self {
486 Self::new_with(self.i.as_ref().unwrap().a.clone())
487 }
488}
489
490impl<T: Zeroize + NewBytes + Clone> Clone for UnlockedRO<T> {
491 fn clone(&self) -> Self {
492 Unlocked::<T>::new_with(self.i.as_ref().unwrap().a.clone())
493 .mprotect_readonly()
494 .expect("unable to create new readonly instance")
495 }
496}
497
498pub use ptypes::*;
499
500fn dryoc_mlock(region: int::Region) -> Result<(), std::io::Error> {
501 if region.len == 0 {
502 return Ok(());
504 }
505 #[cfg(unix)]
506 {
507 #[cfg(target_os = "linux")]
508 {
509 use libc::{MADV_DONTDUMP, madvise};
511 unsafe {
516 madvise(region.ptr() as *mut c_void, region.len, MADV_DONTDUMP);
517 }
518 }
519
520 use libc::{c_void, mlock as c_mlock};
521 let ret = unsafe { c_mlock(region.ptr() as *const c_void, region.len) };
525 match ret {
526 0 => Ok(()),
527 _ => Err(std::io::Error::last_os_error()),
528 }
529 }
530 #[cfg(windows)]
531 {
532 use winapi::shared::minwindef::LPVOID;
533 use winapi::um::memoryapi::VirtualLock;
534
535 let res = unsafe { VirtualLock(region.ptr() as LPVOID, region.len) };
539 if res != 0 {
540 Ok(())
541 } else {
542 Err(std::io::Error::last_os_error())
543 }
544 }
545}
546
547fn dryoc_munlock(region: int::Region) -> Result<(), std::io::Error> {
548 if region.len == 0 {
549 return Ok(());
551 }
552 #[cfg(unix)]
553 {
554 #[cfg(target_os = "linux")]
555 {
556 use libc::{MADV_DODUMP, madvise};
558 unsafe {
562 madvise(region.ptr() as *mut c_void, region.len, MADV_DODUMP);
563 }
564 }
565
566 use libc::{c_void, munlock as c_munlock};
567 let ret = unsafe { c_munlock(region.ptr() as *const c_void, region.len) };
571 match ret {
572 0 => Ok(()),
573 _ => Err(std::io::Error::last_os_error()),
574 }
575 }
576 #[cfg(windows)]
577 {
578 use winapi::shared::minwindef::LPVOID;
579 use winapi::um::memoryapi::VirtualUnlock;
580
581 let res = unsafe { VirtualUnlock(region.ptr() as LPVOID, region.len) };
585 if res != 0 {
586 Ok(())
587 } else {
588 Err(std::io::Error::last_os_error())
589 }
590 }
591}
592
593fn dryoc_mprotect(region: int::Region, mode: int::ProtectMode) -> Result<(), std::io::Error> {
594 dryoc_mprotect_ptr(region.ptr(), region.len, mode)
595}
596
597fn dryoc_mprotect_ptr(
598 data: *mut u8,
599 len: usize,
600 mode: int::ProtectMode,
601) -> Result<(), std::io::Error> {
602 if len == 0 {
603 return Ok(());
605 }
606 #[cfg(unix)]
607 {
608 use libc::{PROT_NONE, PROT_READ, PROT_WRITE, c_void, mprotect as c_mprotect};
609 let prot = match mode {
610 int::ProtectMode::ReadOnly => PROT_READ,
611 int::ProtectMode::ReadWrite => PROT_READ | PROT_WRITE,
612 int::ProtectMode::NoAccess => PROT_NONE,
613 };
614 let ret = unsafe { c_mprotect(data as *mut c_void, len, prot) };
617 match ret {
618 0 => Ok(()),
619 _ => Err(std::io::Error::last_os_error()),
620 }
621 }
622 #[cfg(windows)]
623 {
624 use winapi::shared::minwindef::{DWORD, LPVOID};
625 use winapi::um::memoryapi::VirtualProtect;
626 use winapi::um::winnt::{PAGE_NOACCESS, PAGE_READONLY, PAGE_READWRITE};
627
628 let protect = match mode {
629 int::ProtectMode::ReadOnly => PAGE_READONLY,
630 int::ProtectMode::ReadWrite => PAGE_READWRITE,
631 int::ProtectMode::NoAccess => PAGE_NOACCESS,
632 };
633 let mut old: DWORD = 0;
634
635 let res = unsafe { VirtualProtect(data as LPVOID, len, protect, &mut old) };
639 if res != 0 {
640 Ok(())
641 } else {
642 Err(std::io::Error::last_os_error())
643 }
644 }
645}
646
647impl<A: Zeroize + Bytes, PM: traits::ProtectMode, LM: traits::LockMode> Protected<A, PM, LM> {
648 fn new() -> Self {
649 Self {
650 i: None,
651 p: PhantomData,
652 l: PhantomData,
653 }
654 }
655
656 fn new_with(a: A) -> Self {
657 let noaccess_region = int::Region::of(a.as_slice());
658 Self {
659 i: Some(int::InternalData {
660 a,
661 lm: int::LockMode::Unlocked,
662 pm: int::ProtectMode::ReadWrite,
663 noaccess_region,
664 }),
665 p: PhantomData,
666 l: PhantomData,
667 }
668 }
669
670 fn swap_some_or_err<F, OPM: traits::ProtectMode, OLM: traits::LockMode>(
671 &mut self,
672 f: F,
673 ) -> Result<Protected<A, OPM, OLM>, error::Error>
674 where
675 F: Fn(&mut int::InternalData<A>) -> Result<Protected<A, OPM, OLM>, error::Error>,
676 {
677 match &mut self.i {
678 Some(d) => {
679 let mut new = f(d)?;
680 std::mem::swap(&mut new.i, &mut self.i);
682 Ok(new)
683 }
684 _ => Err(error::Error::invalid_state(
685 crate::ErrorContext::ProtectedMemory,
686 )),
687 }
688 }
689
690 fn inner(&self) -> &A {
693 match &self.i {
694 Some(d) => &d.a,
695 None => panic!("invalid array"),
696 }
697 }
698
699 fn inner_mut(&mut self) -> &mut A {
701 match &mut self.i {
702 Some(d) => &mut d.a,
703 None => panic!("invalid array"),
704 }
705 }
706}
707
708impl<A: Zeroize + Bytes, PM: traits::ProtectMode> Unlock<A, PM>
709 for Protected<A, PM, traits::Locked>
710{
711 fn munlock(mut self) -> Result<Protected<A, PM, traits::Unlocked>, error::Error> {
712 self.swap_some_or_err(|old| {
713 dryoc_munlock(old.region())?;
714 old.lm = int::LockMode::Unlocked;
716 Ok(Protected::<A, PM, traits::Unlocked>::new())
717 })
718 }
719}
720
721impl<A: Zeroize + Bytes + Default, PM: traits::ProtectMode> Lock<A, PM>
722 for Protected<A, PM, traits::Unlocked>
723{
724 fn mlock(mut self) -> Result<Protected<A, PM, traits::Locked>, error::Error> {
725 self.swap_some_or_err(|old| {
726 dryoc_mlock(old.region())?;
727 old.lm = int::LockMode::Locked;
729 Ok(Protected::<A, PM, traits::Locked>::new())
730 })
731 }
732}
733
734impl<A: Zeroize + Bytes, PM: traits::ProtectMode, LM: traits::LockMode> ProtectReadOnly<A, PM, LM>
735 for Protected<A, PM, LM>
736{
737 fn mprotect_readonly(mut self) -> Result<Protected<A, traits::ReadOnly, LM>, error::Error> {
738 self.swap_some_or_err(|old| {
739 dryoc_mprotect(old.region(), int::ProtectMode::ReadOnly)?;
740 old.pm = int::ProtectMode::ReadOnly;
742 Ok(Protected::<A, traits::ReadOnly, LM>::new())
743 })
744 }
745}
746
747impl<A: Zeroize + Bytes, PM: traits::ProtectMode, LM: traits::LockMode> ProtectReadWrite<A, PM, LM>
748 for Protected<A, PM, LM>
749{
750 fn mprotect_readwrite(mut self) -> Result<Protected<A, traits::ReadWrite, LM>, error::Error> {
751 self.swap_some_or_err(|old| {
752 dryoc_mprotect(old.region(), int::ProtectMode::ReadWrite)?;
753 old.pm = int::ProtectMode::ReadWrite;
755 Ok(Protected::<A, traits::ReadWrite, LM>::new())
756 })
757 }
758}
759
760impl<A: Zeroize + Bytes, PM: traits::ProtectMode> ProtectNoAccess<A, PM>
761 for Protected<A, PM, traits::Unlocked>
762{
763 fn mprotect_noaccess(
764 mut self,
765 ) -> Result<Protected<A, traits::NoAccess, traits::Unlocked>, error::Error> {
766 self.swap_some_or_err(|old| {
767 let region = old.region();
768 dryoc_mprotect(region, int::ProtectMode::NoAccess)?;
769 old.noaccess_region = region;
771 old.pm = int::ProtectMode::NoAccess;
772 Ok(Protected::<A, traits::NoAccess, traits::Unlocked>::new())
773 })
774 }
775}
776
777macro_rules! impl_protected_read_views {
781 ($($pm:ident),*) => {$(
782 impl<A: Zeroize + Bytes + AsRef<[u8]>, LM: traits::LockMode> AsRef<[u8]>
783 for Protected<A, traits::$pm, LM>
784 {
785 fn as_ref(&self) -> &[u8] {
786 self.inner().as_ref()
787 }
788 }
789
790 impl<A: Zeroize + Bytes, LM: traits::LockMode> Bytes for Protected<A, traits::$pm, LM> {
791 #[inline]
792 fn as_slice(&self) -> &[u8] {
793 self.inner().as_slice()
794 }
795
796 #[inline]
797 fn len(&self) -> usize {
798 self.inner().len()
799 }
800
801 #[inline]
802 fn is_empty(&self) -> bool {
803 self.inner().is_empty()
804 }
805 }
806
807 impl<A: Bytes + Zeroize, LM: traits::LockMode> std::ops::Deref
808 for Protected<A, traits::$pm, LM>
809 {
810 type Target = [u8];
811
812 fn deref(&self) -> &Self::Target {
813 self.inner().as_slice()
814 }
815 }
816 )*};
817}
818
819impl_protected_read_views!(ReadOnly, ReadWrite);
820
821impl<A: Zeroize + MutBytes + AsMut<[u8]>, LM: traits::LockMode> AsMut<[u8]>
822 for Protected<A, traits::ReadWrite, LM>
823{
824 fn as_mut(&mut self) -> &mut [u8] {
825 self.inner_mut().as_mut()
826 }
827}
828
829impl<const LENGTH: usize> From<StackByteArray<LENGTH>> for HeapByteArray<LENGTH> {
830 fn from(other: StackByteArray<LENGTH>) -> Self {
831 let mut r = HeapByteArray::<LENGTH>::new_byte_array();
832 let mut s = other;
833 r.copy_from_slice(s.as_slice());
834 s.zeroize();
835 r
836 }
837}
838
839impl<const LENGTH: usize> StackByteArray<LENGTH> {
840 pub fn mlock(
852 self,
853 ) -> Result<Protected<HeapByteArray<LENGTH>, traits::ReadWrite, traits::Locked>, error::Error>
854 {
855 Protected::<HeapByteArray<LENGTH>, traits::ReadWrite, traits::Unlocked>::new_with(
856 self.into(),
857 )
858 .mlock()
859 }
860}
861
862impl<const LENGTH: usize> StackByteArray<LENGTH> {
863 pub fn mprotect_readonly(
875 self,
876 ) -> Result<Protected<HeapByteArray<LENGTH>, traits::ReadOnly, traits::Unlocked>, error::Error>
877 {
878 Protected::<HeapByteArray<LENGTH>, traits::ReadWrite, traits::Unlocked>::new_with(
879 self.into(),
880 )
881 .mprotect_readonly()
882 }
883}
884
885impl<const LENGTH: usize> Lockable<HeapByteArray<LENGTH>> for HeapByteArray<LENGTH> {
886 fn mlock(
888 self,
889 ) -> Result<Protected<HeapByteArray<LENGTH>, traits::ReadWrite, traits::Locked>, error::Error>
890 {
891 Protected::<HeapByteArray<LENGTH>, traits::ReadWrite, traits::Unlocked>::new_with(self)
892 .mlock()
893 }
894}
895
896impl Lockable<HeapBytes> for HeapBytes {
897 fn mlock(
899 self,
900 ) -> Result<Protected<HeapBytes, traits::ReadWrite, traits::Locked>, error::Error> {
901 Protected::<HeapBytes, traits::ReadWrite, traits::Unlocked>::new_with(self).mlock()
902 }
903}
904
905#[derive(Clone)]
906pub struct PageAlignedAllocator;
911
912#[cfg(unix)]
913const DEFAULT_PAGESIZE: usize = 4096;
914
915#[cfg(unix)]
916fn page_size_from_sysconf(page_size: libc::c_long) -> usize {
917 if page_size > 0 {
918 page_size as usize
919 } else {
920 DEFAULT_PAGESIZE
921 }
922}
923
924static PAGESIZE: LazyLock<usize> = LazyLock::new(|| {
925 #[cfg(unix)]
926 {
927 use libc::{_SC_PAGE_SIZE, sysconf};
928 let page_size = unsafe { sysconf(_SC_PAGE_SIZE) };
931 page_size_from_sysconf(page_size)
932 }
933 #[cfg(windows)]
934 {
935 use winapi::um::sysinfoapi::{GetSystemInfo, SYSTEM_INFO};
936 let mut si = SYSTEM_INFO::default();
937 unsafe { GetSystemInfo(&mut si) };
940 si.dwPageSize as usize
941 }
942});
943
944fn _page_round(size: usize, pagesize: usize) -> Option<usize> {
945 let rem = size % pagesize;
946 if rem == 0 {
947 Some(size)
948 } else {
949 size.checked_add(pagesize - rem)
950 }
951}
952
953fn protected_alloc_error() -> std::io::Error {
954 std::io::Error::other("protected memory allocation failed")
955}
956
957#[derive(Clone, Copy)]
958struct RawRegionLayout {
959 rounded_size: usize,
960 total_size: usize,
961}
962
963fn checked_raw_region_layout(
964 user_size: usize,
965 pagesize: usize,
966) -> Result<RawRegionLayout, std::io::Error> {
967 let rounded_size = _page_round(user_size, pagesize).ok_or_else(protected_alloc_error)?;
968 let guard_size = pagesize.checked_mul(2).ok_or_else(protected_alloc_error)?;
969 let total_size = rounded_size
970 .checked_add(guard_size)
971 .ok_or_else(protected_alloc_error)?;
972 Ok(RawRegionLayout {
973 rounded_size,
974 total_size,
975 })
976}
977
978#[derive(Clone, Copy)]
979struct RawProtectedAllocation {
980 base: NonNull<u8>,
981 data: NonNull<u8>,
982 rounded_size: usize,
983 total_size: usize,
984}
985
986fn platform_alloc(total_size: usize, pagesize: usize) -> Result<NonNull<u8>, std::io::Error> {
987 #[cfg(unix)]
988 {
989 use libc::posix_memalign;
990 let mut out = ptr::null_mut();
991
992 let ret = unsafe { posix_memalign(&mut out, pagesize, total_size) };
996 if ret != 0 {
997 return Err(std::io::Error::from_raw_os_error(ret));
998 }
999
1000 NonNull::new(out as *mut u8).ok_or_else(protected_alloc_error)
1001 }
1002 #[cfg(windows)]
1003 {
1004 let _ = pagesize;
1005 use winapi::um::memoryapi::VirtualAlloc;
1006 use winapi::um::winnt::{MEM_COMMIT, MEM_RESERVE, PAGE_READWRITE};
1007
1008 let out = unsafe {
1012 VirtualAlloc(
1013 ptr::null_mut(),
1014 total_size,
1015 MEM_COMMIT | MEM_RESERVE,
1016 PAGE_READWRITE,
1017 )
1018 };
1019
1020 NonNull::new(out as *mut u8).ok_or_else(std::io::Error::last_os_error)
1021 }
1022}
1023
1024fn platform_free(base: NonNull<u8>, total_size: usize) {
1025 #[cfg(unix)]
1026 {
1027 let _ = total_size;
1028 unsafe { libc::free(base.as_ptr() as *mut libc::c_void) };
1031 }
1032 #[cfg(windows)]
1033 {
1034 let _ = total_size;
1035 use winapi::shared::minwindef::LPVOID;
1036 use winapi::um::memoryapi::VirtualFree;
1037 use winapi::um::winnt::MEM_RELEASE;
1038 unsafe { VirtualFree(base.as_ptr() as LPVOID, 0, MEM_RELEASE) };
1041 }
1042}
1043
1044fn allocate_raw_region(user_size: usize) -> Result<RawProtectedAllocation, std::io::Error> {
1045 let pagesize = *PAGESIZE;
1046 let layout = checked_raw_region_layout(user_size, pagesize)?;
1047 let base = platform_alloc(layout.total_size, pagesize)?;
1048 let base_ptr = base.as_ptr();
1049
1050 if let Err(err) = dryoc_mprotect_ptr(base_ptr, pagesize, int::ProtectMode::NoAccess) {
1051 platform_free(base, layout.total_size);
1052 return Err(err);
1053 }
1054
1055 let aft_guard_offset = pagesize
1056 .checked_add(layout.rounded_size)
1057 .ok_or_else(protected_alloc_error)?;
1058 let aft_guard = unsafe { base_ptr.add(aft_guard_offset) };
1061 if let Err(err) = dryoc_mprotect_ptr(aft_guard, pagesize, int::ProtectMode::NoAccess) {
1062 let _ = dryoc_mprotect_ptr(base_ptr, pagesize, int::ProtectMode::ReadWrite);
1063 platform_free(base, layout.total_size);
1064 return Err(err);
1065 }
1066
1067 let data_ptr = unsafe { base_ptr.add(pagesize) };
1070 let data = NonNull::new(data_ptr).ok_or_else(protected_alloc_error)?;
1071
1072 Ok(RawProtectedAllocation {
1073 base,
1074 data,
1075 rounded_size: layout.rounded_size,
1076 total_size: layout.total_size,
1077 })
1078}
1079
1080fn deallocate_raw_region(raw: RawProtectedAllocation) {
1081 let pagesize = *PAGESIZE;
1082 let base_ptr = raw.base.as_ptr();
1083 let _ = dryoc_mprotect_ptr(base_ptr, pagesize, int::ProtectMode::ReadWrite);
1084
1085 if let Some(aft_guard_offset) = pagesize.checked_add(raw.rounded_size) {
1086 let aft_guard = unsafe { base_ptr.add(aft_guard_offset) };
1089 let _ = dryoc_mprotect_ptr(aft_guard, pagesize, int::ProtectMode::ReadWrite);
1090 }
1091
1092 platform_free(raw.base, raw.total_size);
1093}
1094
1095struct ProtectedBuffer {
1096 base: Option<NonNull<u8>>,
1097 data: NonNull<u8>,
1098 len: usize,
1099 capacity: usize,
1100 rounded_size: usize,
1101 total_size: usize,
1102}
1103
1104unsafe impl Send for ProtectedBuffer {}
1108
1109unsafe impl Sync for ProtectedBuffer {}
1112
1113impl ProtectedBuffer {
1114 fn new_filled(len: usize, value: u8) -> Result<Self, std::io::Error> {
1115 if len == 0 {
1116 return Ok(Self::default());
1117 }
1118
1119 let raw = allocate_raw_region(len)?;
1120 unsafe { raw.data.as_ptr().write_bytes(value, len) };
1125 Ok(Self {
1126 base: Some(raw.base),
1127 data: raw.data,
1128 len,
1129 capacity: len,
1130 rounded_size: raw.rounded_size,
1131 total_size: raw.total_size,
1132 })
1133 }
1134
1135 fn from_slice(src: &[u8]) -> Result<Self, std::io::Error> {
1136 let mut buffer = Self::new_filled(src.len(), 0)?;
1137 buffer.as_mut_slice().copy_from_slice(src);
1138 Ok(buffer)
1139 }
1140
1141 fn as_ptr(&self) -> *const u8 {
1142 self.data.as_ptr()
1143 }
1144
1145 fn as_mut_ptr(&mut self) -> *mut u8 {
1146 self.data.as_ptr()
1147 }
1148
1149 fn as_slice(&self) -> &[u8] {
1150 debug_assert!(self.len <= self.capacity);
1151 unsafe { std::slice::from_raw_parts(self.data.as_ptr(), self.len) }
1154 }
1155
1156 fn as_mut_slice(&mut self) -> &mut [u8] {
1157 debug_assert!(self.len <= self.capacity);
1158 unsafe { std::slice::from_raw_parts_mut(self.data.as_ptr(), self.len) }
1161 }
1162
1163 fn len(&self) -> usize {
1164 self.len
1165 }
1166
1167 fn is_empty(&self) -> bool {
1168 self.len == 0
1169 }
1170
1171 fn resize(&mut self, new_len: usize, value: u8) {
1172 self.try_resize(new_len, value)
1173 .expect("protected resize failed");
1174 }
1175
1176 fn try_resize(&mut self, new_len: usize, value: u8) -> Result<(), std::io::Error> {
1179 if new_len == self.len {
1180 return Ok(());
1181 }
1182
1183 let mut resized = Self::new_filled(new_len, value)?;
1184 let len_to_copy = std::cmp::min(self.len, new_len);
1185 resized.as_mut_slice()[..len_to_copy].copy_from_slice(&self.as_slice()[..len_to_copy]);
1186 std::mem::swap(self, &mut resized);
1187 Ok(())
1188 }
1189
1190 fn copy_from_slice(&mut self, other: &[u8]) {
1191 self.as_mut_slice().copy_from_slice(other);
1192 }
1193}
1194
1195impl Default for ProtectedBuffer {
1196 fn default() -> Self {
1197 Self {
1198 base: None,
1199 data: NonNull::dangling(),
1200 len: 0,
1201 capacity: 0,
1202 rounded_size: 0,
1203 total_size: 0,
1204 }
1205 }
1206}
1207
1208impl Clone for ProtectedBuffer {
1209 fn clone(&self) -> Self {
1210 Self::from_slice(self.as_slice()).expect("protected clone failed")
1211 }
1212}
1213
1214impl fmt::Debug for ProtectedBuffer {
1215 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1216 f.debug_struct("ProtectedBuffer")
1217 .field("len", &self.len())
1218 .field("contents", &"[REDACTED]")
1219 .finish()
1220 }
1221}
1222
1223impl PartialEq for ProtectedBuffer {
1224 fn eq(&self, other: &Self) -> bool {
1225 self.as_slice().ct_eq(other.as_slice()).into()
1226 }
1227}
1228
1229impl Eq for ProtectedBuffer {}
1230
1231impl Zeroize for ProtectedBuffer {
1232 fn zeroize(&mut self) {
1233 self.as_mut_slice().zeroize();
1234 }
1235}
1236
1237impl Drop for ProtectedBuffer {
1238 fn drop(&mut self) {
1239 if let Some(base) = self.base.take() {
1240 if self.rounded_size != 0 {
1241 let _ = dryoc_mprotect_ptr(
1242 self.data.as_ptr(),
1243 self.rounded_size,
1244 int::ProtectMode::ReadWrite,
1245 );
1246 }
1247 self.as_mut_slice().zeroize();
1248 deallocate_raw_region(RawProtectedAllocation {
1249 base,
1250 data: self.data,
1251 rounded_size: self.rounded_size,
1252 total_size: self.total_size,
1253 });
1254 }
1255 }
1256}
1257
1258impl AsRef<[u8]> for ProtectedBuffer {
1259 fn as_ref(&self) -> &[u8] {
1260 self.as_slice()
1261 }
1262}
1263
1264impl AsMut<[u8]> for ProtectedBuffer {
1265 fn as_mut(&mut self) -> &mut [u8] {
1266 self.as_mut_slice()
1267 }
1268}
1269
1270impl std::ops::Deref for ProtectedBuffer {
1271 type Target = [u8];
1272
1273 fn deref(&self) -> &Self::Target {
1274 self.as_slice()
1275 }
1276}
1277
1278impl std::ops::DerefMut for ProtectedBuffer {
1279 fn deref_mut(&mut self) -> &mut Self::Target {
1280 self.as_mut_slice()
1281 }
1282}
1283
1284impl_slice_index!(impl[] ProtectedBuffer, |s| s.as_slice(), |s| s.as_mut_slice());
1285
1286#[cfg(feature = "nightly")]
1287unsafe impl Allocator for PageAlignedAllocator {
1292 #[inline]
1293 fn allocate(&self, layout: std::alloc::Layout) -> Result<NonNull<[u8]>, AllocError> {
1294 let pagesize = *PAGESIZE;
1295 if !pagesize.is_multiple_of(layout.align()) {
1296 return Err(AllocError);
1297 }
1298
1299 let raw = allocate_raw_region(layout.size()).map_err(|_| AllocError)?;
1300 unsafe {
1303 Ok(NonNull::new_unchecked(ptr::slice_from_raw_parts_mut(
1304 raw.data.as_ptr(),
1305 layout.size(),
1306 )))
1307 }
1308 }
1309
1310 #[inline]
1315 unsafe fn deallocate(&self, ptr: NonNull<u8>, layout: std::alloc::Layout) {
1318 let pagesize = *PAGESIZE;
1319
1320 let base_ptr = unsafe { ptr.as_ptr().sub(pagesize) };
1323 let Some(base) = NonNull::new(base_ptr) else {
1324 return;
1325 };
1326 let Ok(raw_layout) = checked_raw_region_layout(layout.size(), pagesize) else {
1327 return;
1328 };
1329 deallocate_raw_region(RawProtectedAllocation {
1330 base,
1331 data: ptr,
1332 rounded_size: raw_layout.rounded_size,
1333 total_size: raw_layout.total_size,
1334 });
1335 }
1336}
1337
1338#[derive(Zeroize, ZeroizeOnDrop, Debug, PartialEq, Eq, Clone)]
1343pub struct HeapByteArray<const LENGTH: usize>(ProtectedBuffer);
1344
1345#[derive(Zeroize, ZeroizeOnDrop, Debug, PartialEq, Eq, Clone, Default)]
1350pub struct HeapBytes(ProtectedBuffer);
1351
1352fn expect_locked<T>(result: Result<T, error::Error>) -> T {
1355 match result {
1356 Ok(r) => r,
1357 Err(err) => panic!("Error creating locked bytes: {:?}", err),
1358 }
1359}
1360
1361fn into_readonly_locked<A: Zeroize + Bytes>(
1364 result: Result<Protected<A, traits::ReadWrite, traits::Locked>, error::Error>,
1365) -> Result<Protected<A, traits::ReadOnly, traits::Locked>, error::Error> {
1366 result.and_then(|p| p.mprotect_readonly())
1367}
1368
1369impl<A: Zeroize + NewBytes + Lockable<A>> NewLocked<A> for A {
1370 fn new_locked() -> Result<Protected<Self, traits::ReadWrite, traits::Locked>, error::Error> {
1371 Self::new_bytes().mlock()
1372 }
1373
1374 fn new_readonly_locked()
1375 -> Result<Protected<Self, traits::ReadOnly, traits::Locked>, error::Error> {
1376 into_readonly_locked(Self::new_bytes().mlock())
1377 }
1378
1379 fn generate_locked() -> Result<Protected<Self, traits::ReadWrite, traits::Locked>, error::Error>
1380 {
1381 let mut res = Self::new_bytes().mlock()?;
1382 copy_randombytes(res.as_mut_slice());
1383 Ok(res)
1384 }
1385
1386 fn generate_readonly_locked()
1387 -> Result<Protected<Self, traits::ReadOnly, traits::Locked>, error::Error> {
1388 into_readonly_locked(Self::generate_locked())
1389 }
1390}
1391
1392impl<A: Zeroize + NewBytes + ResizableBytes + Lockable<A>> NewLockedFromSlice<A> for A {
1393 fn from_slice_into_locked(
1395 src: &[u8],
1396 ) -> Result<Protected<Self, traits::ReadWrite, traits::Locked>, crate::error::Error> {
1397 let mut res = Self::new_bytes().mlock()?;
1398 res.resize(src.len(), 0);
1399 res.as_mut_slice().copy_from_slice(src);
1400 Ok(res)
1401 }
1402
1403 fn from_slice_into_readonly_locked(
1405 src: &[u8],
1406 ) -> Result<Protected<Self, traits::ReadOnly, traits::Locked>, crate::error::Error> {
1407 into_readonly_locked(Self::from_slice_into_locked(src))
1408 }
1409}
1410
1411impl<const LENGTH: usize> NewLockedFromSlice<HeapByteArray<LENGTH>> for HeapByteArray<LENGTH> {
1412 fn from_slice_into_locked(
1414 other: &[u8],
1415 ) -> Result<Protected<Self, traits::ReadWrite, traits::Locked>, crate::error::Error> {
1416 validate_length!(exact LENGTH, other.len(), crate::ErrorContext::Slice);
1417 let mut res = Self::new_bytes().mlock()?;
1418 res.as_mut_slice().copy_from_slice(other);
1419 Ok(res)
1420 }
1421
1422 fn from_slice_into_readonly_locked(
1423 other: &[u8],
1424 ) -> Result<Protected<Self, traits::ReadOnly, traits::Locked>, crate::error::Error> {
1425 into_readonly_locked(Self::from_slice_into_locked(other))
1426 }
1427}
1428
1429macro_rules! impl_heap_buffer_views {
1433 ($($t:ident $(<$length:ident: usize>)?;)*) => {$(
1434 impl$(<const $length: usize>)? Bytes for $t$(<$length>)? {
1435 #[inline]
1436 fn as_slice(&self) -> &[u8] {
1437 &self.0
1438 }
1439
1440 #[inline]
1441 fn len(&self) -> usize {
1442 self.0.len()
1443 }
1444
1445 #[inline]
1446 fn is_empty(&self) -> bool {
1447 self.0.is_empty()
1448 }
1449 }
1450
1451 impl$(<const $length: usize>)? MutBytes for $t$(<$length>)? {
1452 #[inline]
1453 fn as_mut_slice(&mut self) -> &mut [u8] {
1454 self.0.as_mut_slice()
1455 }
1456
1457 fn copy_from_slice(&mut self, other: &[u8]) {
1458 self.0.copy_from_slice(other)
1459 }
1460 }
1461
1462 impl$(<const $length: usize>)? std::convert::AsRef<[u8]> for $t$(<$length>)? {
1463 fn as_ref(&self) -> &[u8] {
1464 self.0.as_ref()
1465 }
1466 }
1467
1468 impl$(<const $length: usize>)? std::convert::AsMut<[u8]> for $t$(<$length>)? {
1469 fn as_mut(&mut self) -> &mut [u8] {
1470 self.0.as_mut()
1471 }
1472 }
1473
1474 impl$(<const $length: usize>)? std::ops::Deref for $t$(<$length>)? {
1475 type Target = [u8];
1476
1477 fn deref(&self) -> &Self::Target {
1478 &self.0
1479 }
1480 }
1481
1482 impl$(<const $length: usize>)? std::ops::DerefMut for $t$(<$length>)? {
1483 fn deref_mut(&mut self) -> &mut Self::Target {
1484 &mut self.0
1485 }
1486 }
1487 )*};
1488}
1489
1490impl_heap_buffer_views!(HeapByteArray<LENGTH: usize>; HeapBytes;);
1491
1492impl NewBytes for HeapBytes {
1493 fn new_bytes() -> Self {
1494 Self::default()
1495 }
1496}
1497
1498impl ResizableBytes for HeapBytes {
1499 fn resize(&mut self, new_len: usize, value: u8) {
1500 self.0.resize(new_len, value);
1501 }
1502}
1503
1504#[cfg(feature = "serde")]
1507impl HeapBytes {
1508 pub(crate) fn try_resize(&mut self, new_len: usize, value: u8) -> Result<(), error::Error> {
1509 Ok(self.0.try_resize(new_len, value)?)
1510 }
1511}
1512
1513#[cfg(feature = "serde")]
1514impl Protected<HeapBytes, traits::ReadWrite, traits::Locked> {
1515 pub(crate) fn try_resize(&mut self, new_len: usize, value: u8) -> Result<(), error::Error> {
1516 if new_len == self.len() {
1517 return Ok(());
1518 }
1519 let mut new = HeapBytes::default();
1520 new.try_resize(new_len, value)?;
1521 self.replace_locked(new)
1522 }
1523}
1524
1525impl<A: Zeroize + NewBytes + Lockable<A>> Protected<A, traits::ReadWrite, traits::Locked> {
1526 fn replace_locked(&mut self, new: A) -> Result<(), error::Error> {
1531 let mut locked = new.mlock()?;
1532 let len_to_copy = std::cmp::min(locked.len(), self.len());
1533 locked.as_mut_slice()[..len_to_copy].copy_from_slice(&self.as_slice()[..len_to_copy]);
1534 std::mem::swap(&mut locked.i, &mut self.i);
1535 Ok(())
1536 }
1537}
1538
1539impl<A: Zeroize + NewBytes + ResizableBytes + Lockable<A>> ResizableBytes
1540 for Protected<A, traits::ReadWrite, traits::Locked>
1541{
1542 fn resize(&mut self, new_len: usize, value: u8) {
1543 if new_len == self.len() {
1544 return;
1545 }
1546 let mut new = A::new_bytes();
1547 new.resize(new_len, value);
1548 self.replace_locked(new).expect("unable to lock on resize");
1549 }
1550}
1551
1552impl<A: Zeroize + NewBytes + ResizableBytes + Lockable<A>> ResizableBytes
1553 for Protected<A, traits::ReadWrite, traits::Unlocked>
1554{
1555 fn resize(&mut self, new_len: usize, value: u8) {
1556 self.inner_mut().resize(new_len, value)
1557 }
1558}
1559
1560impl<A: Zeroize + MutBytes, LM: traits::LockMode> MutBytes for Protected<A, traits::ReadWrite, LM> {
1561 #[inline]
1562 fn as_mut_slice(&mut self) -> &mut [u8] {
1563 self.inner_mut().as_mut_slice()
1564 }
1565
1566 fn copy_from_slice(&mut self, other: &[u8]) {
1567 self.inner_mut().copy_from_slice(other)
1568 }
1569}
1570
1571impl<const LENGTH: usize> std::convert::AsRef<[u8; LENGTH]> for HeapByteArray<LENGTH> {
1572 fn as_ref(&self) -> &[u8; LENGTH] {
1573 let arr = self.0.as_ptr() as *const [u8; LENGTH];
1574 unsafe { &*arr }
1577 }
1578}
1579
1580impl<const LENGTH: usize> std::convert::AsMut<[u8; LENGTH]> for HeapByteArray<LENGTH> {
1581 fn as_mut(&mut self) -> &mut [u8; LENGTH] {
1582 let arr = self.0.as_mut_ptr() as *mut [u8; LENGTH];
1583 unsafe { &mut *arr }
1586 }
1587}
1588
1589impl<A: MutBytes + Zeroize, LM: traits::LockMode> std::ops::DerefMut
1590 for Protected<A, traits::ReadWrite, LM>
1591{
1592 fn deref_mut(&mut self) -> &mut Self::Target {
1593 self.inner_mut().as_mut_slice()
1594 }
1595}
1596
1597impl_slice_index!(impl[const LENGTH: usize] HeapByteArray<LENGTH>, |s| s.0, |s| s.0);
1598
1599impl<const LENGTH: usize> Default for HeapByteArray<LENGTH> {
1600 fn default() -> Self {
1601 Self(ProtectedBuffer::new_filled(LENGTH, 0).expect("protected allocation failed"))
1602 }
1603}
1604
1605impl<A: Zeroize + NewBytes + Lockable<A> + NewLocked<A>> Default
1606 for Protected<A, traits::ReadWrite, traits::Locked>
1607{
1608 fn default() -> Self {
1609 A::new_locked().expect("mlock failed")
1610 }
1611}
1612
1613impl_slice_index!(impl[] HeapBytes, |s| s.0, |s| s.0);
1614
1615impl<const LENGTH: usize> From<&[u8; LENGTH]> for HeapByteArray<LENGTH> {
1616 fn from(src: &[u8; LENGTH]) -> Self {
1617 let mut arr = Self::default();
1618 arr.0.copy_from_slice(src);
1619 arr
1620 }
1621}
1622
1623impl<const LENGTH: usize> From<[u8; LENGTH]> for HeapByteArray<LENGTH> {
1624 fn from(mut src: [u8; LENGTH]) -> Self {
1625 let ret = Self::from(&src);
1626 src.zeroize();
1628 ret
1629 }
1630}
1631
1632impl<const LENGTH: usize> TryFrom<&[u8]> for HeapByteArray<LENGTH> {
1633 type Error = error::Error;
1634
1635 fn try_from(src: &[u8]) -> Result<Self, Self::Error> {
1636 validate_length!(exact LENGTH, src.len(), crate::ErrorContext::Slice);
1637 let mut arr = Self::default();
1638 arr.0.copy_from_slice(src);
1639 Ok(arr)
1640 }
1641}
1642
1643impl From<&[u8]> for HeapBytes {
1644 fn from(src: &[u8]) -> Self {
1645 Self(ProtectedBuffer::from_slice(src).expect("protected allocation failed"))
1646 }
1647}
1648
1649impl<const LENGTH: usize> ByteArray<LENGTH> for HeapByteArray<LENGTH> {
1650 #[inline]
1651 fn as_array(&self) -> &[u8; LENGTH] {
1652 let ptr = self.0.as_ptr() as *const [u8; LENGTH];
1653 unsafe { &*ptr }
1656 }
1657}
1658
1659impl<const LENGTH: usize> NewBytes for HeapByteArray<LENGTH> {
1660 fn new_bytes() -> Self {
1661 Self::default()
1662 }
1663}
1664
1665impl NewBytes for Protected<HeapBytes, traits::ReadWrite, traits::Locked> {
1666 fn new_bytes() -> Self {
1667 expect_locked(HeapBytes::new_locked())
1668 }
1669}
1670
1671impl<const LENGTH: usize> NewBytes
1672 for Protected<HeapByteArray<LENGTH>, traits::ReadWrite, traits::Locked>
1673{
1674 fn new_bytes() -> Self {
1675 expect_locked(HeapByteArray::<LENGTH>::new_locked())
1676 }
1677}
1678
1679impl<const LENGTH: usize> NewByteArray<LENGTH>
1680 for Protected<HeapByteArray<LENGTH>, traits::ReadWrite, traits::Locked>
1681{
1682 fn new_byte_array() -> Self {
1683 expect_locked(HeapByteArray::<LENGTH>::new_locked())
1684 }
1685
1686 fn generate() -> Self {
1687 let mut res = expect_locked(HeapByteArray::<LENGTH>::new_locked());
1688 copy_randombytes(res.as_mut_slice());
1689 res
1690 }
1691}
1692
1693impl<const LENGTH: usize> NewByteArray<LENGTH> for HeapByteArray<LENGTH> {
1694 fn new_byte_array() -> Self {
1695 Self::default()
1696 }
1697
1698 fn generate() -> Self {
1700 gen_bytes()
1701 }
1702}
1703
1704impl<const LENGTH: usize> MutByteArray<LENGTH> for HeapByteArray<LENGTH> {
1705 fn as_mut_array(&mut self) -> &mut [u8; LENGTH] {
1706 let ptr = self.0.as_mut_ptr() as *mut [u8; LENGTH];
1707 unsafe { &mut *ptr }
1710 }
1711}
1712
1713macro_rules! impl_protected_array_views {
1717 (readable: $($pm:ident, $lm:ident;)*) => {$(
1718 impl<const LENGTH: usize> ByteArray<LENGTH>
1719 for Protected<HeapByteArray<LENGTH>, traits::$pm, traits::$lm>
1720 {
1721 #[inline]
1722 fn as_array(&self) -> &[u8; LENGTH] {
1723 self.inner().as_array()
1724 }
1725 }
1726 )*};
1727 (writable: $($lm:ident;)*) => {$(
1728 impl<const LENGTH: usize> MutByteArray<LENGTH>
1729 for Protected<HeapByteArray<LENGTH>, traits::ReadWrite, traits::$lm>
1730 {
1731 #[inline]
1732 fn as_mut_array(&mut self) -> &mut [u8; LENGTH] {
1733 self.inner_mut().as_mut_array()
1734 }
1735 }
1736
1737 impl<const LENGTH: usize> AsMut<[u8; LENGTH]>
1738 for Protected<HeapByteArray<LENGTH>, traits::ReadWrite, traits::$lm>
1739 {
1740 fn as_mut(&mut self) -> &mut [u8; LENGTH] {
1741 self.inner_mut().as_mut()
1742 }
1743 }
1744 )*};
1745}
1746
1747impl_protected_array_views!(readable:
1748 ReadOnly, Unlocked;
1749 ReadOnly, Locked;
1750 ReadWrite, Unlocked;
1751 ReadWrite, Locked;
1752);
1753impl_protected_array_views!(writable:
1754 Locked;
1755 Unlocked;
1756);
1757
1758impl<A: Zeroize + Bytes, PM: traits::ProtectMode, LM: traits::LockMode> Drop
1759 for Protected<A, PM, LM>
1760{
1761 fn drop(&mut self) {
1762 let Some(mut data) = self.i.take() else {
1763 return;
1764 };
1765
1766 let region = data.region();
1769 let writable = region.len == 0
1770 || data.pm == int::ProtectMode::ReadWrite
1771 || match dryoc_mprotect(region, int::ProtectMode::ReadWrite) {
1772 Ok(()) => true,
1773 Err(err) => abort_protected_memory_failure("making memory writable for drop", err),
1774 };
1775
1776 if writable {
1777 data.a.zeroize();
1778 }
1779
1780 if data.lm == int::LockMode::Locked {
1781 match dryoc_munlock(region) {
1782 Ok(()) => data.lm = int::LockMode::Unlocked,
1783 Err(err) => abort_protected_memory_failure("unlocking memory for drop", err),
1784 }
1785 }
1786 }
1787}
1788
1789impl<A: Zeroize + Bytes, PM: traits::ProtectMode, LM: traits::LockMode> ZeroizeOnDrop
1790 for Protected<A, PM, LM>
1791{
1792}
1793
1794impl<A: Zeroize + Bytes, PM: traits::ProtectMode, LM: traits::LockMode> Zeroize
1795 for Protected<A, PM, LM>
1796{
1797 fn zeroize(&mut self) {
1798 let Some(data) = &mut self.i else {
1799 return;
1800 };
1801 let region = data.region();
1802 if region.len == 0 {
1803 return;
1804 }
1805
1806 let previous_mode = data.pm.clone();
1807 if previous_mode != int::ProtectMode::ReadWrite
1808 && let Err(error) = dryoc_mprotect(region, int::ProtectMode::ReadWrite)
1809 {
1810 abort_protected_memory_failure("making memory writable for zeroization", error);
1811 }
1812
1813 data.a.zeroize();
1814
1815 if previous_mode != int::ProtectMode::ReadWrite
1816 && let Err(error) = dryoc_mprotect(region, previous_mode)
1817 {
1818 abort_protected_memory_failure("restoring memory protection after zeroization", error);
1819 }
1820 }
1821}
1822
1823fn abort_protected_memory_failure(_operation: &str, _error: std::io::Error) -> ! {
1824 std::process::abort()
1825}
1826
1827#[cfg(test)]
1829pub(crate) mod test_util {
1830 use super::*;
1831 use crate::test_prelude::*;
1832
1833 pub(crate) fn can_lock_pages(pages: usize) -> bool {
1843 let probes: Result<Vec<_>, _> = (0..pages)
1844 .map(|_| HeapBytes::from(&[0u8][..]).mlock())
1845 .collect();
1846 match probes {
1847 Ok(_) => true,
1848 Err(error::Error::Io(err)) if is_lock_quota_error(&err) => {
1849 std::eprintln!("skipping: this process cannot lock {pages} page(s): {err}");
1850 false
1851 }
1852 Err(err) => panic!("locking a fresh page failed: {err}"),
1853 }
1854 }
1855
1856 #[cfg(unix)]
1860 fn is_lock_quota_error(err: &std::io::Error) -> bool {
1861 matches!(
1862 err.raw_os_error(),
1863 Some(libc::ENOMEM | libc::EPERM | libc::EAGAIN)
1864 )
1865 }
1866
1867 #[cfg(windows)]
1870 fn is_lock_quota_error(err: &std::io::Error) -> bool {
1871 err.raw_os_error() == Some(1453)
1872 }
1873}
1874
1875#[cfg(test)]
1876mod tests {
1877 use proptest::prelude::*;
1878
1879 use super::test_util::can_lock_pages;
1880 use super::*;
1881 use crate::test_prelude::*;
1882
1883 #[test]
1884 fn protected_byte_array_debug_redacts_contents() {
1885 let bytes = HeapByteArray::from(StackByteArray::from([0xabu8; 4]));
1886 let debug = format!("{bytes:?}");
1887
1888 assert!(debug.contains("[REDACTED]"));
1889 assert!(!debug.contains("171"));
1890 }
1891
1892 fn interesting_lengths() -> impl Strategy<Value = usize> {
1893 let pagesize = *PAGESIZE;
1894 let max = pagesize.saturating_mul(2).saturating_add(8);
1895
1896 prop_oneof![
1897 Just(0usize),
1898 Just(1),
1899 0usize..=128,
1900 pagesize.saturating_sub(8)..=pagesize.saturating_add(8),
1901 pagesize.saturating_mul(2).saturating_sub(8)..=max,
1902 ]
1903 .boxed()
1904 }
1905
1906 fn small_lengths() -> impl Strategy<Value = usize> {
1907 prop_oneof![Just(0usize), Just(1), 0usize..=256].boxed()
1908 }
1909
1910 fn interesting_bytes() -> impl Strategy<Value = Vec<u8>> {
1911 interesting_lengths()
1912 .prop_flat_map(|len| prop::collection::vec(any::<u8>(), len))
1913 .boxed()
1914 }
1915
1916 fn small_bytes() -> impl Strategy<Value = Vec<u8>> {
1917 small_lengths()
1918 .prop_flat_map(|len| prop::collection::vec(any::<u8>(), len))
1919 .boxed()
1920 }
1921
1922 #[cfg_attr(
1923 tarpaulin,
1924 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
1925 )]
1926 #[test]
1927 fn test_lock_unlock() {
1928 use crate::dryocstream::Key;
1929
1930 let key = Key::generate();
1931 let key_clone = key.clone();
1932
1933 let locked_key = key.mlock().expect("lock failed");
1934
1935 let unlocked_key = locked_key.munlock().expect("unlock failed");
1936
1937 assert_eq!(unlocked_key.as_slice(), key_clone.as_slice());
1938 }
1939
1940 #[cfg_attr(
1941 tarpaulin,
1942 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
1943 )]
1944 #[test]
1945 fn explicit_zeroize_preserves_locked_readwrite_state() {
1946 let mut locked =
1947 HeapBytes::from_slice_into_locked(b"sensitive").expect("locked allocation failed");
1948
1949 locked.zeroize();
1950
1951 assert_eq!(locked.as_slice(), &[0; 9]);
1952 let state = locked.i.as_ref().expect("protected state missing");
1953 assert_eq!(state.lm, int::LockMode::Locked);
1954 assert_eq!(state.pm, int::ProtectMode::ReadWrite);
1955
1956 let unlocked = locked.munlock().expect("unlock after zeroize failed");
1957 assert_eq!(unlocked.as_slice(), &[0; 9]);
1958 }
1959
1960 #[cfg(unix)]
1961 #[cfg_attr(
1962 tarpaulin,
1963 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
1964 )]
1965 #[test]
1966 fn explicit_zeroize_restores_readonly_protection() {
1967 let mut readonly = HeapBytes::from_slice_into_readonly_locked(b"sensitive")
1968 .expect("read-only locked allocation failed");
1969
1970 readonly.zeroize();
1971
1972 assert_eq!(readonly.as_slice(), &[0; 9]);
1973 let state = readonly.i.as_ref().expect("protected state missing");
1974 assert_eq!(state.lm, int::LockMode::Locked);
1975 assert_eq!(state.pm, int::ProtectMode::ReadOnly);
1976
1977 let child = unsafe { libc::fork() };
1979 assert!(child >= 0, "fork failed");
1980 if child == 0 {
1981 let data = readonly.as_slice().as_ptr() as *mut u8;
1982 unsafe {
1985 std::ptr::write_volatile(data, 1);
1986 libc::_exit(0);
1987 }
1988 }
1989
1990 let mut status = 0;
1991 let wait_ret = unsafe { libc::waitpid(child, &mut status, 0) };
1994 assert_eq!(wait_ret, child);
1995 assert!(
1996 libc::WIFSIGNALED(status),
1997 "child unexpectedly wrote to explicitly zeroized read-only memory"
1998 );
1999
2000 let readwrite = readonly
2001 .mprotect_readwrite()
2002 .expect("read-write transition failed");
2003 let unlocked = readwrite.munlock().expect("unlock failed");
2004 assert_eq!(unlocked.as_slice(), &[0; 9]);
2005 }
2006
2007 #[cfg_attr(
2008 tarpaulin,
2009 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2010 )]
2011 #[test]
2012 fn test_protect_unprotect() {
2013 use crate::dryocstream::Key;
2014
2015 let key = Key::generate();
2016 let key_clone = key.clone();
2017
2018 let readonly_key = key.mprotect_readonly().expect("mprotect failed");
2019 assert_eq!(readonly_key.as_slice(), key_clone.as_slice());
2020
2021 let mut readwrite_key = readonly_key.mprotect_readwrite().expect("mprotect failed");
2022 assert_eq!(readwrite_key.as_slice(), key_clone.as_slice());
2023
2024 readwrite_key.as_mut_slice()[0] = 0;
2026 }
2027
2028 #[cfg(feature = "nightly")]
2029 #[test]
2030 fn test_allocator() {
2031 let mut vec: Vec<i32, _> = Vec::new_in(PageAlignedAllocator);
2032
2033 vec.push(1);
2034 vec.push(2);
2035 vec.push(3);
2036
2037 for i in 0..5000 {
2038 vec.push(i);
2039 }
2040
2041 vec.resize(5, 0);
2042
2043 assert_eq!([1, 2, 3, 0, 1], vec.as_slice());
2044 }
2045
2046 #[cfg(feature = "nightly")]
2047 #[test]
2048 fn test_allocator_honors_supported_alignment() {
2049 let allocator = PageAlignedAllocator;
2050 let layout = std::alloc::Layout::from_size_align(1, *PAGESIZE).unwrap();
2051 let allocation = allocator.allocate(layout).unwrap();
2052 let data = allocation.as_ptr() as *mut u8;
2053
2054 assert_eq!(data.addr() % layout.align(), 0);
2055
2056 unsafe { allocator.deallocate(NonNull::new_unchecked(data), layout) };
2058 }
2059
2060 #[cfg(feature = "nightly")]
2061 #[test]
2062 fn test_allocator_rejects_unsupported_alignment() {
2063 let unsupported_alignment = PAGESIZE.checked_mul(2).unwrap();
2064 let layout = std::alloc::Layout::from_size_align(1, unsupported_alignment).unwrap();
2065
2066 assert!(PageAlignedAllocator.allocate(layout).is_err());
2067 }
2068
2069 #[cfg(feature = "nightly")]
2070 #[test]
2071 fn test_allocator_handles_zero_sized_layout() {
2072 let allocator = PageAlignedAllocator;
2073 let layout = std::alloc::Layout::from_size_align(0, 1).unwrap();
2074 let allocation = allocator.allocate(layout).unwrap();
2075 let data = allocation.as_ptr() as *mut u8;
2076
2077 assert_eq!(allocation.len(), 0);
2078 assert_eq!(data.addr() % layout.align(), 0);
2079
2080 unsafe { allocator.deallocate(NonNull::new_unchecked(data), layout) };
2082 }
2083
2084 #[test]
2085 fn test_page_rounding() {
2086 let pagesize = *PAGESIZE;
2087
2088 assert_eq!(_page_round(0, pagesize), Some(0));
2089 assert_eq!(_page_round(1, pagesize), Some(pagesize));
2090 assert_eq!(_page_round(pagesize, pagesize), Some(pagesize));
2091 assert_eq!(_page_round(pagesize + 1, pagesize), Some(pagesize * 2));
2092 assert_eq!(_page_round(usize::MAX, pagesize), None);
2093 }
2094
2095 #[cfg(unix)]
2096 #[test]
2097 fn test_page_size_from_sysconf_handles_error_sentinel() {
2098 assert_eq!(page_size_from_sysconf(-1), DEFAULT_PAGESIZE);
2099 assert_eq!(page_size_from_sysconf(0), DEFAULT_PAGESIZE);
2100 assert_eq!(page_size_from_sysconf(8192), 8192);
2101 }
2102
2103 #[test]
2104 fn test_empty_heapbytes_and_locking() {
2105 let empty = HeapBytes::default();
2106 assert!(empty.is_empty());
2107 assert_eq!(empty.as_slice().len(), 0);
2108
2109 let locked: LockedBytes = HeapBytes::new_locked().expect("empty mlock failed");
2110 assert!(locked.is_empty());
2111
2112 let unlocked = locked.munlock().expect("empty munlock failed");
2113 assert!(unlocked.is_empty());
2114 }
2115
2116 #[test]
2117 fn test_heapbytes_resize_grow_shrink_and_fill() {
2118 let mut bytes = HeapBytes::default();
2119 bytes.resize(3, 0x7a);
2120 assert_eq!(bytes.as_slice(), &[0x7a, 0x7a, 0x7a]);
2121
2122 bytes.as_mut_slice()[1] = 0x11;
2123 bytes.resize(5, 0x5a);
2124 assert_eq!(bytes.as_slice(), &[0x7a, 0x11, 0x7a, 0x5a, 0x5a]);
2125
2126 bytes.resize(2, 0);
2127 assert_eq!(bytes.as_slice(), &[0x7a, 0x11]);
2128
2129 bytes.resize(0, 0);
2130 assert!(bytes.is_empty());
2131 }
2132
2133 proptest! {
2134 #![proptest_config(ProptestConfig::with_cases(64))]
2135
2136 #[test]
2137 fn proptest_heapbytes_roundtrip_clone_and_mutation(data in interesting_bytes()) {
2138 let bytes = HeapBytes::from(data.as_slice());
2139 prop_assert_eq!(bytes.len(), data.len());
2140 prop_assert_eq!(bytes.as_slice(), data.as_slice());
2141 prop_assert_eq!(bytes.as_ref(), data.as_slice());
2142
2143 let mut cloned = bytes.clone();
2144 prop_assert_eq!(&cloned, &bytes);
2145 prop_assert_eq!(cloned.as_slice(), data.as_slice());
2146
2147 if !data.is_empty() {
2148 prop_assert_eq!(cloned[0], data[0]);
2149
2150 let last = data.len() - 1;
2151 prop_assert_eq!(cloned[last], data[last]);
2152
2153 cloned[0] = cloned[0].wrapping_add(1);
2154 prop_assert_ne!(cloned[0], data[0]);
2155 prop_assert_eq!(&cloned[1..], &data[1..]);
2156 }
2157 }
2158
2159 #[test]
2160 fn proptest_heapbytes_resize_matches_vec_model(
2161 initial in interesting_bytes(),
2162 ops in prop::collection::vec((interesting_lengths(), any::<u8>()), 0..12),
2163 ) {
2164 let mut bytes = HeapBytes::from(initial.as_slice());
2165 let mut model = initial;
2166
2167 for (new_len, value) in ops {
2168 bytes.resize(new_len, value);
2169 model.resize(new_len, value);
2170 prop_assert_eq!(bytes.as_slice(), model.as_slice());
2171 }
2172 }
2173
2174 #[test]
2175 fn proptest_protection_transitions_preserve_bytes(data in interesting_bytes()) {
2176 let protected =
2177 Protected::<HeapBytes, traits::ReadWrite, traits::Unlocked>::new_with(
2178 HeapBytes::from(data.as_slice()),
2179 );
2180
2181 let readonly = protected
2182 .mprotect_readonly()
2183 .expect("readonly mprotect failed");
2184 prop_assert_eq!(readonly.as_slice(), data.as_slice());
2185
2186 let readwrite = readonly
2187 .mprotect_readwrite()
2188 .expect("readwrite mprotect failed");
2189 prop_assert_eq!(readwrite.as_slice(), data.as_slice());
2190
2191 let noaccess = readwrite
2192 .mprotect_noaccess()
2193 .expect("noaccess mprotect failed");
2194 let readwrite = noaccess
2195 .mprotect_readwrite()
2196 .expect("readwrite mprotect failed");
2197 prop_assert_eq!(readwrite.as_slice(), data.as_slice());
2198 }
2199 }
2200
2201 proptest! {
2202 #![proptest_config(ProptestConfig::with_cases(32))]
2203
2204 #[test]
2205 fn proptest_locked_heapbytes_resize_matches_vec_model(
2206 initial in small_bytes(),
2207 ops in prop::collection::vec((small_lengths(), any::<u8>()), 0..8),
2208 ) {
2209 let mut locked = HeapBytes::from_slice_into_locked(initial.as_slice())
2210 .expect("locked allocation failed");
2211 let mut model = initial;
2212
2213 for (new_len, value) in ops {
2214 locked.resize(new_len, value);
2215 model.resize(new_len, value);
2216 prop_assert_eq!(locked.as_slice(), model.as_slice());
2217 }
2218
2219 let unlocked = locked.munlock().expect("munlock failed");
2220 prop_assert_eq!(unlocked.as_slice(), model.as_slice());
2221 }
2222
2223 #[test]
2224 fn proptest_heapbytearray_exact_size_views(data in any::<[u8; 32]>()) {
2225 let mut bytes = HeapByteArray::<32>::from(&data);
2226
2227 prop_assert_eq!(bytes.as_array(), &data);
2228 prop_assert_eq!(AsRef::<[u8; 32]>::as_ref(&bytes), &data);
2229 prop_assert_eq!(bytes.as_slice(), &data);
2230
2231 let mut expected = data;
2232 bytes.as_mut_array()[7] ^= 0xa5;
2233 expected[7] ^= 0xa5;
2234 prop_assert_eq!(bytes.as_array(), &expected);
2235
2236 AsMut::<[u8; 32]>::as_mut(&mut bytes)[24] = 0x5a;
2237 expected[24] = 0x5a;
2238 prop_assert_eq!(bytes.as_slice(), &expected);
2239 }
2240 }
2241
2242 #[test]
2243 fn test_heapbytearray_exact_size_views() {
2244 let mut bytes = HeapByteArray::<4>::default();
2245 bytes.as_mut_array().copy_from_slice(&[1, 2, 3, 4]);
2246
2247 assert_eq!(bytes.as_array(), &[1, 2, 3, 4]);
2248 assert_eq!(AsRef::<[u8; 4]>::as_ref(&bytes), &[1, 2, 3, 4]);
2249
2250 AsMut::<[u8; 4]>::as_mut(&mut bytes)[2] = 9;
2251 assert_eq!(bytes.as_slice(), &[1, 2, 9, 4]);
2252 }
2253
2254 #[cfg_attr(
2255 tarpaulin,
2256 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2257 )]
2258 #[test]
2259 fn test_mprotect_handles_single_byte_slice() {
2260 let mut vec = HeapBytes::from(&[1u8][..]);
2261
2262 let region = int::Region::of(vec.as_slice());
2263 dryoc_mprotect(region, int::ProtectMode::ReadOnly).expect("readonly mprotect failed");
2264 dryoc_mprotect(region, int::ProtectMode::ReadWrite).expect("readwrite mprotect failed");
2265 vec[0] = 2;
2266
2267 assert_eq!(vec[0], 2);
2268 }
2269
2270 #[cfg_attr(
2271 tarpaulin,
2272 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2273 )]
2274 #[test]
2275 fn test_mprotect_handles_exact_page_slice() {
2276 let pagesize = *PAGESIZE;
2277 let mut vec = HeapBytes::default();
2278 vec.resize(pagesize, 1);
2279
2280 let region = int::Region::of(vec.as_slice());
2281 dryoc_mprotect(region, int::ProtectMode::ReadOnly).expect("readonly mprotect failed");
2282 dryoc_mprotect(region, int::ProtectMode::ReadWrite).expect("readwrite mprotect failed");
2283 vec[0] = 2;
2284 vec[pagesize - 1] = 3;
2285
2286 assert_eq!(vec[0], 2);
2287 assert_eq!(vec[pagesize - 1], 3);
2288 }
2289
2290 #[cfg(unix)]
2291 #[cfg_attr(
2292 tarpaulin,
2293 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2294 )]
2295 #[test]
2296 fn test_mprotect_noaccess_covers_page_boundary_tail() {
2297 let pagesize = *PAGESIZE;
2298 let mut vec = HeapBytes::default();
2299 vec.resize(pagesize + 1, 0);
2300
2301 let region = int::Region::of(vec.as_slice());
2303 let data = vec.as_mut_slice().as_mut_ptr();
2304 dryoc_mprotect(region, int::ProtectMode::NoAccess).expect("noaccess mprotect failed");
2305
2306 let child = unsafe { libc::fork() };
2307 assert!(child >= 0, "fork failed");
2308
2309 if child == 0 {
2310 let tail = unsafe { data.add(pagesize) };
2311 unsafe {
2312 std::ptr::write_volatile(tail, 1);
2313 libc::_exit(0);
2314 }
2315 }
2316
2317 let mut status = 0;
2318 let wait_ret = unsafe { libc::waitpid(child, &mut status, 0) };
2319 dryoc_mprotect(region, int::ProtectMode::ReadWrite).expect("readwrite mprotect failed");
2320
2321 assert_eq!(wait_ret, child);
2322 assert!(
2323 libc::WIFSIGNALED(status),
2324 "child unexpectedly wrote to protected tail page"
2325 );
2326 }
2327
2328 const SRC: [u8; 6] = [10, 20, 30, 40, 50, 60];
2329
2330 #[cfg(unix)]
2333 fn child_faults(probe: impl FnOnce()) -> bool {
2334 let child = unsafe { libc::fork() };
2338 assert!(child >= 0, "fork failed");
2339 if child == 0 {
2340 probe();
2341 unsafe { libc::_exit(0) };
2343 }
2344
2345 let mut status = 0;
2346 let wait_ret = unsafe { libc::waitpid(child, &mut status, 0) };
2349 assert_eq!(wait_ret, child);
2350 libc::WIFSIGNALED(status)
2351 }
2352
2353 #[test]
2354 fn protected_allocations_are_page_aligned() {
2355 let pagesize = *PAGESIZE;
2356 let lockable = can_lock_pages(2);
2358
2359 for len in [1, pagesize, pagesize + 1] {
2360 let mut bytes = HeapBytes::default();
2361 bytes.resize(len, 0x5a);
2362 assert_eq!(bytes.len(), len);
2363 assert_eq!(bytes.as_slice().as_ptr().addr() % pagesize, 0, "len {len}");
2364
2365 if lockable {
2366 let locked = HeapBytes::from_slice_into_locked(bytes.as_slice()).expect("locked");
2367 assert_eq!(locked.as_slice().as_ptr().addr() % pagesize, 0, "len {len}");
2368 }
2369 }
2370
2371 let array = HeapByteArray::<32>::default();
2372 assert_eq!(array.as_slice().as_ptr().addr() % pagesize, 0);
2373 }
2374
2375 #[test]
2376 fn test_checked_raw_region_layout_boundaries() {
2377 let pagesize = *PAGESIZE;
2378
2379 let ok = |user_size: usize, rounded_size: usize| {
2380 let layout = checked_raw_region_layout(user_size, pagesize).expect("layout should fit");
2381 assert_eq!(layout.rounded_size, rounded_size, "user size {user_size}");
2382 assert_eq!(
2383 layout.total_size,
2384 rounded_size + 2 * pagesize,
2385 "user size {user_size}"
2386 );
2387 };
2388
2389 ok(0, 0);
2390 ok(1, pagesize);
2391 ok(pagesize - 1, pagesize);
2392 ok(pagesize, pagesize);
2393 ok(pagesize + 1, 2 * pagesize);
2394
2395 let largest = usize::MAX - 3 * pagesize + 1;
2398 ok(largest, largest);
2399
2400 assert!(checked_raw_region_layout(largest + 1, pagesize).is_err());
2402 assert!(checked_raw_region_layout(usize::MAX - 2 * pagesize + 1, pagesize).is_err());
2404 assert!(checked_raw_region_layout(usize::MAX - pagesize + 1, pagesize).is_err());
2405 assert!(checked_raw_region_layout(usize::MAX, pagesize).is_err());
2407 assert!(checked_raw_region_layout(0, usize::MAX / 2 + 1).is_err());
2410 }
2411
2412 #[cfg_attr(
2413 tarpaulin,
2414 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2415 )]
2416 #[test]
2417 fn locked_clone_is_a_distinct_independent_locked_copy() {
2418 if !can_lock_pages(2) {
2420 return;
2421 }
2422 let original = HeapBytes::from_slice_into_locked(b"clone me").expect("locked");
2423 let mut cloned = original.clone();
2424
2425 assert_eq!(cloned.as_slice(), original.as_slice());
2426 assert_ne!(cloned.as_slice().as_ptr(), original.as_slice().as_ptr());
2427 let state = cloned.i.as_ref().expect("protected state missing");
2428 assert_eq!(state.lm, int::LockMode::Locked);
2429 assert_eq!(state.pm, int::ProtectMode::ReadWrite);
2430
2431 cloned.as_mut_slice()[0] = b'C';
2432 assert_eq!(original.as_slice(), b"clone me");
2433 assert_eq!(cloned.as_slice(), b"Clone me");
2434
2435 let unlocked = cloned.munlock().expect("unlock failed");
2436 assert_eq!(unlocked.as_slice(), b"Clone me");
2437 assert_eq!(original.as_slice(), b"clone me");
2438 }
2439
2440 #[cfg_attr(
2441 tarpaulin,
2442 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2443 )]
2444 #[test]
2445 fn locked_resize_to_the_same_length_keeps_its_region() {
2446 if !can_lock_pages(1) {
2448 return;
2449 }
2450 let mut locked = HeapBytes::from_slice_into_locked(b"keep").expect("locked");
2451 let data = locked.as_slice().as_ptr();
2452
2453 locked.resize(4, 0);
2454
2455 assert_eq!(locked.as_slice().as_ptr(), data);
2456 assert_eq!(locked.as_slice(), b"keep");
2457 let state = locked.i.as_ref().expect("protected state missing");
2458 assert_eq!(state.lm, int::LockMode::Locked);
2459 }
2460
2461 #[cfg_attr(
2462 tarpaulin,
2463 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2464 )]
2465 #[test]
2466 fn locked_readonly_clone_is_a_distinct_readonly_copy() {
2467 if !can_lock_pages(2) {
2469 return;
2470 }
2471 let original =
2472 HeapBytes::from_slice_into_readonly_locked(b"clone me").expect("read-only locked");
2473 let cloned = original.clone();
2474
2475 assert_eq!(cloned.as_slice(), original.as_slice());
2476 assert_ne!(cloned.as_slice().as_ptr(), original.as_slice().as_ptr());
2477 let state = cloned.i.as_ref().expect("protected state missing");
2478 assert_eq!(state.lm, int::LockMode::Locked);
2479 assert_eq!(state.pm, int::ProtectMode::ReadOnly);
2480
2481 #[cfg(unix)]
2482 {
2483 let data = cloned.as_slice().as_ptr() as *mut u8;
2484 assert!(
2485 child_faults(|| unsafe { ptr::write_volatile(data, 1) }),
2486 "clone's pages are not read-only"
2487 );
2488 }
2489
2490 let mut writable = cloned
2491 .mprotect_readwrite()
2492 .expect("read-write transition failed")
2493 .munlock()
2494 .expect("unlock failed");
2495 writable.as_mut_slice()[0] = b'C';
2496 assert_eq!(writable.as_slice(), b"Clone me");
2497 assert_eq!(original.as_slice(), b"clone me");
2498 }
2499
2500 #[cfg_attr(
2501 tarpaulin,
2502 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2503 )]
2504 #[test]
2505 fn unlocked_clones_are_distinct_copies_preserving_protect_mode() {
2506 let original = Unlocked::<HeapBytes>::new_with(HeapBytes::from(&b"clone me"[..]));
2507 let mut cloned = original.clone();
2508
2509 assert_eq!(cloned.as_slice(), original.as_slice());
2510 assert_ne!(cloned.as_slice().as_ptr(), original.as_slice().as_ptr());
2511 let state = cloned.i.as_ref().expect("protected state missing");
2512 assert_eq!(state.lm, int::LockMode::Unlocked);
2513 assert_eq!(state.pm, int::ProtectMode::ReadWrite);
2514 cloned.as_mut_slice()[0] = b'C';
2515 assert_eq!(original.as_slice(), b"clone me");
2516
2517 let readonly = original
2518 .mprotect_readonly()
2519 .expect("readonly mprotect failed");
2520 let readonly_clone = readonly.clone();
2521 assert_eq!(readonly_clone.as_slice(), b"clone me");
2522 assert_ne!(
2523 readonly_clone.as_slice().as_ptr(),
2524 readonly.as_slice().as_ptr()
2525 );
2526 let state = readonly_clone.i.as_ref().expect("protected state missing");
2527 assert_eq!(state.lm, int::LockMode::Unlocked);
2528 assert_eq!(state.pm, int::ProtectMode::ReadOnly);
2529
2530 let mut writable = readonly_clone
2531 .mprotect_readwrite()
2532 .expect("readwrite mprotect failed");
2533 writable.as_mut_slice()[0] = b'C';
2534 assert_eq!(writable.as_slice(), b"Clone me");
2535 assert_eq!(readonly.as_slice(), b"clone me");
2536 }
2537
2538 #[cfg_attr(
2539 tarpaulin,
2540 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2541 )]
2542 #[test]
2543 fn heap_bytes_move_across_threads_and_share_through_arc() {
2544 use std::sync::Arc;
2545
2546 let bytes = HeapBytes::from(&SRC[..]);
2547 let returned = std::thread::spawn(move || {
2548 let mut bytes = bytes;
2549 assert_eq!(bytes.as_slice(), &SRC);
2550 bytes.as_mut_slice()[0] ^= 0xff;
2551 bytes
2552 })
2553 .join()
2554 .expect("thread panicked");
2555 assert_eq!(returned.as_slice()[0], SRC[0] ^ 0xff);
2556 assert_eq!(&returned.as_slice()[1..], &SRC[1..]);
2557
2558 if !can_lock_pages(1) {
2559 return;
2560 }
2561 let shared = Arc::new(HeapBytes::from_slice_into_locked(&SRC).expect("locked"));
2562 let readers: Vec<_> = (0..4)
2563 .map(|_| {
2564 let shared = Arc::clone(&shared);
2565 std::thread::spawn(move || shared.as_slice().to_vec())
2566 })
2567 .collect();
2568 for reader in readers {
2569 assert_eq!(reader.join().expect("reader panicked"), SRC);
2570 }
2571 assert_eq!(shared.as_slice(), &SRC);
2572 }
2573
2574 #[cfg_attr(
2575 tarpaulin,
2576 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2577 )]
2578 #[test]
2579 fn fixed_size_locked_constructors_require_exact_length() {
2580 use crate::utils::test_util::assert_exact_slice_length_error;
2581
2582 const LENGTH: usize = 8;
2583 let data = [7u8; LENGTH + 1];
2584
2585 let exact_heap = HeapByteArray::<LENGTH>::try_from(&data[..LENGTH]).expect("exact heap");
2586 assert_eq!(exact_heap.as_slice(), &data[..LENGTH]);
2587 if can_lock_pages(1) {
2591 let exact =
2592 HeapByteArray::<LENGTH>::from_slice_into_locked(&data[..LENGTH]).expect("exact");
2593 assert_eq!(exact.as_slice(), &data[..LENGTH]);
2594 drop(exact);
2595 let exact_readonly =
2596 HeapByteArray::<LENGTH>::from_slice_into_readonly_locked(&data[..LENGTH])
2597 .expect("exact read-only");
2598 assert_eq!(exact_readonly.as_slice(), &data[..LENGTH]);
2599 }
2600
2601 for actual in [LENGTH - 1, LENGTH + 1] {
2602 assert_exact_slice_length_error(
2603 HeapByteArray::<LENGTH>::from_slice_into_locked(&data[..actual]),
2604 actual,
2605 LENGTH,
2606 );
2607 assert_exact_slice_length_error(
2608 HeapByteArray::<LENGTH>::from_slice_into_readonly_locked(&data[..actual]),
2609 actual,
2610 LENGTH,
2611 );
2612 assert_exact_slice_length_error(
2613 HeapByteArray::<LENGTH>::try_from(&data[..actual]),
2614 actual,
2615 LENGTH,
2616 );
2617 }
2618 }
2619
2620 #[test]
2621 fn transitions_on_a_taken_protected_value_report_invalid_state() {
2622 fn assert_invalid_state<T>(result: Result<T, error::Error>) {
2623 match result {
2624 Err(error::Error::InvalidState { context }) => {
2625 assert_eq!(context, crate::ErrorContext::ProtectedMemory)
2626 }
2627 Err(other) => panic!("unexpected error {other:?}"),
2628 Ok(_) => panic!("transition succeeded without a backing buffer"),
2629 }
2630 }
2631
2632 assert_invalid_state(Unlocked::<HeapBytes>::new().mprotect_readonly());
2633 assert_invalid_state(Unlocked::<HeapBytes>::new().mprotect_noaccess());
2634 assert_invalid_state(Unlocked::<HeapBytes>::new().mlock());
2635 assert_invalid_state(LockedBytes::new().munlock());
2636 assert_invalid_state(LockedRO::<HeapBytes>::new().mprotect_readwrite());
2637 }
2638
2639 struct SpyBytes {
2642 inner: HeapBytes,
2643 wipes: std::sync::Arc<std::sync::atomic::AtomicUsize>,
2644 }
2645
2646 impl SpyBytes {
2647 fn new(data: &[u8]) -> (Self, std::sync::Arc<std::sync::atomic::AtomicUsize>) {
2648 let wipes = std::sync::Arc::default();
2649 let spy = Self {
2650 inner: HeapBytes::from(data),
2651 wipes: std::sync::Arc::clone(&wipes),
2652 };
2653 (spy, wipes)
2654 }
2655 }
2656
2657 impl Default for SpyBytes {
2658 fn default() -> Self {
2659 Self::new(&[]).0
2660 }
2661 }
2662
2663 impl Zeroize for SpyBytes {
2664 fn zeroize(&mut self) {
2665 self.wipes.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
2666 self.inner.zeroize();
2667 }
2668 }
2669
2670 impl Bytes for SpyBytes {
2671 fn as_slice(&self) -> &[u8] {
2672 self.inner.as_slice()
2673 }
2674
2675 fn len(&self) -> usize {
2676 self.inner.len()
2677 }
2678
2679 fn is_empty(&self) -> bool {
2680 self.inner.is_empty()
2681 }
2682 }
2683
2684 #[cfg_attr(
2685 tarpaulin,
2686 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2687 )]
2688 #[test]
2689 fn drop_wipes_the_backing_store_exactly_once_in_every_state() {
2690 use std::sync::atomic::Ordering::SeqCst;
2691
2692 let (spy, wipes) = SpyBytes::new(b"secret");
2693 drop(Unlocked::<SpyBytes>::new_with(spy));
2694 assert_eq!(wipes.load(SeqCst), 1, "unlocked read-write");
2695
2696 let (spy, wipes) = SpyBytes::new(b"secret");
2699 let readonly = Unlocked::<SpyBytes>::new_with(spy)
2700 .mprotect_readonly()
2701 .expect("readonly mprotect failed");
2702 assert_eq!(wipes.load(SeqCst), 0);
2703 drop(readonly);
2704 assert_eq!(wipes.load(SeqCst), 1, "unlocked read-only");
2705
2706 let (spy, wipes) = SpyBytes::new(b"secret");
2707 let noaccess = Unlocked::<SpyBytes>::new_with(spy)
2708 .mprotect_noaccess()
2709 .expect("noaccess mprotect failed");
2710 assert_eq!(wipes.load(SeqCst), 0);
2711 drop(noaccess);
2712 assert_eq!(wipes.load(SeqCst), 1, "no-access");
2713
2714 if can_lock_pages(1) {
2715 let (spy, wipes) = SpyBytes::new(b"secret");
2716 let locked = Unlocked::<SpyBytes>::new_with(spy)
2717 .mlock()
2718 .expect("mlock failed");
2719 assert_eq!(wipes.load(SeqCst), 0);
2720 drop(locked);
2721 assert_eq!(wipes.load(SeqCst), 1, "locked read-write");
2722
2723 let (spy, wipes) = SpyBytes::new(b"secret");
2724 let locked_readonly = Unlocked::<SpyBytes>::new_with(spy)
2725 .mlock()
2726 .expect("mlock failed")
2727 .mprotect_readonly()
2728 .expect("readonly mprotect failed");
2729 assert_eq!(wipes.load(SeqCst), 0);
2730 drop(locked_readonly);
2731 assert_eq!(wipes.load(SeqCst), 1, "locked read-only");
2732 }
2733
2734 let (spy, wipes) = SpyBytes::new(b"");
2735 drop(Unlocked::<SpyBytes>::new_with(spy));
2736 assert_eq!(wipes.load(SeqCst), 1, "empty");
2737 }
2738
2739 #[cfg_attr(
2740 tarpaulin,
2741 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2742 )]
2743 #[test]
2744 fn explicit_zeroize_wipes_once_and_drop_wipes_again() {
2745 use std::sync::atomic::Ordering::SeqCst;
2746
2747 let (spy, wipes) = SpyBytes::new(b"secret");
2748 let mut protected = Unlocked::<SpyBytes>::new_with(spy);
2749
2750 protected.zeroize();
2751 assert_eq!(wipes.load(SeqCst), 1);
2752 assert_eq!(protected.as_slice(), &[0; 6]);
2753
2754 drop(protected);
2755 assert_eq!(wipes.load(SeqCst), 2);
2756
2757 let (spy, wipes) = SpyBytes::new(b"");
2759 let mut empty = Unlocked::<SpyBytes>::new_with(spy);
2760 empty.zeroize();
2761 assert_eq!(wipes.load(SeqCst), 0);
2762 }
2763
2764 #[cfg(unix)]
2765 #[cfg_attr(
2766 tarpaulin,
2767 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2768 )]
2769 #[test]
2770 fn explicit_zeroize_restores_noaccess_protection() {
2771 use std::sync::atomic::Ordering::SeqCst;
2772
2773 let (spy, wipes) = SpyBytes::new(b"secret");
2774 let readwrite = Unlocked::<SpyBytes>::new_with(spy);
2775 let data = readwrite.as_slice().as_ptr();
2777 let mut noaccess = readwrite
2778 .mprotect_noaccess()
2779 .expect("noaccess mprotect failed");
2780
2781 noaccess.zeroize();
2782
2783 assert_eq!(wipes.load(SeqCst), 1);
2784 let state = noaccess.i.as_ref().expect("protected state missing");
2785 assert_eq!(state.lm, int::LockMode::Unlocked);
2786 assert_eq!(state.pm, int::ProtectMode::NoAccess);
2787
2788 assert!(
2790 child_faults(|| {
2791 std::hint::black_box(unsafe { ptr::read_volatile(data) });
2792 }),
2793 "child unexpectedly read explicitly zeroized no-access memory"
2794 );
2795
2796 let readwrite = noaccess
2797 .mprotect_readwrite()
2798 .expect("readwrite mprotect failed");
2799 assert_eq!(readwrite.as_slice(), &[0; 6]);
2800 drop(readwrite);
2801 assert_eq!(wipes.load(SeqCst), 2);
2802 }
2803
2804 #[cfg(unix)]
2805 #[cfg_attr(
2806 tarpaulin,
2807 ignore = "tarpaulin can segfault while tracing mlock/mprotect tests"
2808 )]
2809 #[test]
2810 fn guard_pages_fault_on_both_sides_of_the_user_region() {
2811 let pagesize = *PAGESIZE;
2812
2813 for len in [1usize, pagesize, pagesize + 1] {
2814 let mut bytes = HeapBytes::default();
2815 bytes.resize(len, 0x5a);
2816 let data = bytes.as_slice().as_ptr() as *mut u8;
2817 let rounded = bytes.0.rounded_size;
2818 assert_eq!(rounded, _page_round(len, pagesize).unwrap(), "len {len}");
2819
2820 assert!(
2822 !child_faults(|| unsafe { ptr::write_volatile(data.add(rounded - 1), 1) }),
2823 "len {len}: last byte of the user region faulted"
2824 );
2825 assert!(
2827 child_faults(|| unsafe { ptr::write_volatile(data.add(rounded), 1) }),
2828 "len {len}: rear guard page did not fault"
2829 );
2830 assert!(
2832 child_faults(|| unsafe { ptr::write_volatile(data.sub(1), 1) }),
2833 "len {len}: front guard page did not fault"
2834 );
2835 assert!(
2836 child_faults(|| {
2837 std::hint::black_box(unsafe { ptr::read_volatile(data.sub(1)) });
2838 }),
2839 "len {len}: front guard page allowed a read"
2840 );
2841
2842 assert_eq!(bytes.as_slice(), vec![0x5a; len].as_slice());
2843 }
2844 }
2845}