use rand::rngs::OsRng;
use rand::RngCore;
use zeroize::Zeroize;
use crate::vault::VaultError;
const KEY_LEN: usize = 32;
const BUF_LEN: usize = KEY_LEN * 2;
const _: fn() = || {
fn assert_send<T: Send>() {}
assert_send::<MaskedDek>();
};
pub struct MaskedDek {
buf: Box<[u8; BUF_LEN]>,
warnings: Vec<String>,
}
impl Drop for MaskedDek {
fn drop(&mut self) {
self.buf.zeroize();
}
}
fn unmask_into(buf: &[u8; BUF_LEN], out: &mut [u8]) {
let (masked, mask) = buf.split_at(KEY_LEN);
for (o, (m, k)) in out.iter_mut().zip(masked.iter().zip(mask.iter())) {
*o = m ^ k;
}
}
impl MaskedDek {
pub fn new(dek: zeroize::Zeroizing<Vec<u8>>) -> Result<Self, VaultError> {
if dek.len() != KEY_LEN {
return Err(VaultError::Crypto("DEK length must be 32 bytes".into()));
}
let mut buf = Box::new([0u8; BUF_LEN]);
{
let (masked, mask) = buf.split_at_mut(KEY_LEN);
OsRng
.try_fill_bytes(mask)
.map_err(|e| VaultError::Crypto(format!("OsRng failed: {e}")))?;
for (out, (d, k)) in masked.iter_mut().zip(dek.iter().zip(mask.iter())) {
*out = d ^ k;
}
}
let mut warnings = Vec::new();
match region::lock(buf.as_ptr(), BUF_LEN) {
Ok(guard) => core::mem::forget(guard),
Err(_) => warnings.push(
"could not lock the key page in this environment; \
pages holding the key could reach swap"
.to_string(),
),
}
Ok(Self { buf, warnings })
}
pub fn warnings(&self) -> &[String] {
&self.warnings
}
pub fn with_dek<R>(&mut self, f: impl FnOnce(&[u8]) -> R) -> R {
let mut plain = zeroize::Zeroizing::new(vec![0u8; KEY_LEN]);
unmask_into(&self.buf, &mut plain);
struct Remask<'a> {
m: &'a mut MaskedDek,
plain: &'a [u8],
}
impl Drop for Remask<'_> {
fn drop(&mut self) {
self.m.remask_from(self.plain);
}
}
let guard = Remask {
m: self,
plain: &plain,
};
let out = f(&plain);
drop(guard);
out
}
pub fn duplicate(&mut self) -> Result<Self, VaultError> {
self.with_dek(|dek| Self::new(zeroize::Zeroizing::new(dek.to_vec())))
}
fn remask_from(&mut self, plain: &[u8]) {
let mut new_mask = [0u8; KEY_LEN];
let ok = OsRng.try_fill_bytes(&mut new_mask).is_ok();
let (masked, mask) = self.buf.split_at_mut(KEY_LEN);
if ok {
for (out, (p, nk)) in masked.iter_mut().zip(plain.iter().zip(new_mask.iter())) {
*out = p ^ nk;
}
for (k, nk) in mask.iter_mut().zip(new_mask.iter()) {
*k = *nk;
}
} else {
for (out, (p, k)) in masked.iter_mut().zip(plain.iter().zip(mask.iter())) {
*out = p ^ k;
}
self.warnings
.push("OsRng failed during mask rotation; mask not rotated".to_string());
}
new_mask.zeroize();
}
#[cfg(test)]
pub(crate) fn debug_masked_snapshot(&self) -> Vec<u8> {
let (masked, _) = self.buf.split_at(KEY_LEN);
masked.to_vec()
}
}
pub fn harden_process() -> Vec<String> {
let warnings = Vec::new();
#[cfg(unix)]
let mut warnings = warnings;
#[cfg(unix)]
{
use rlimit::Resource;
if let Err(e) = rlimit::setrlimit(Resource::CORE, 0, 0) {
warnings.push(format!("could not set RLIMIT_CORE=0: {e}"));
}
}
#[cfg(target_os = "linux")]
{
if prctl::set_dumpable(false).is_err() {
warnings.push("could not set PR_SET_DUMPABLE=0".to_string());
}
}
warnings
}
#[cfg(test)]
mod tests {
use super::{harden_process, MaskedDek, KEY_LEN};
use zeroize::Zeroizing;
fn dek() -> Zeroizing<Vec<u8>> {
Zeroizing::new((0u8..32).collect())
}
#[test]
fn test_with_dek_exposes_the_original_value_every_time() {
let mut m = MaskedDek::new(dek()).expect("32B");
let a = m.with_dek(|d| d.to_vec());
let b = m.with_dek(|d| d.to_vec());
assert_eq!(a, (0u8..32).collect::<Vec<u8>>());
assert_eq!(a, b);
}
#[test]
fn test_masked_representation_rotates_between_accesses() {
let mut m = MaskedDek::new(dek()).expect("32B");
let snap1 = m.debug_masked_snapshot();
m.with_dek(|_| ());
let snap2 = m.debug_masked_snapshot();
assert_ne!(snap1, snap2);
}
#[test]
fn test_masked_representation_never_equals_plain_dek() {
let mut m = MaskedDek::new(dek()).expect("32B");
let plain = (0u8..32).collect::<Vec<u8>>();
assert_ne!(m.debug_masked_snapshot(), plain);
m.with_dek(|_| ());
assert_ne!(m.debug_masked_snapshot(), plain);
}
#[test]
fn test_duplicate_preserves_value_with_independent_mask() {
let mut m = MaskedDek::new(dek()).expect("32B");
let mut d = m.duplicate().expect("dup");
assert_eq!(m.with_dek(|x| x.to_vec()), d.with_dek(|x| x.to_vec()));
assert_ne!(m.debug_masked_snapshot(), d.debug_masked_snapshot());
}
#[test]
fn test_new_and_harden_never_fail_in_restricted_environments() {
let m = MaskedDek::new(dek()).expect("32B");
let _ = m.warnings();
let _warns: Vec<String> = harden_process();
}
#[test]
fn test_new_rejects_wrong_length_with_typed_error() {
assert!(MaskedDek::new(Zeroizing::new(vec![0u8; 16])).is_err());
assert!(MaskedDek::new(Zeroizing::new(vec![0u8; 33])).is_err());
assert!(MaskedDek::new(Zeroizing::new(vec![0u8; KEY_LEN])).is_ok());
}
#[test]
fn test_with_dek_remasks_even_if_the_operation_panics() {
let mut m = MaskedDek::new(dek()).expect("32B");
let before = m.debug_masked_snapshot();
let plain = (0u8..32).collect::<Vec<u8>>();
let r = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
m.with_dek(|_k| panic!("boom inside the op"));
}));
assert!(r.is_err());
let after = m.debug_masked_snapshot();
assert_ne!(after, plain);
assert_ne!(after, before);
}
}