#[derive(Debug)]
pub(crate) struct EntropyUnavailable(pub String);
impl std::fmt::Display for EntropyUnavailable {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.0)
}
}
impl std::error::Error for EntropyUnavailable {}
#[derive(Debug)]
pub(crate) struct MaskKeySource {
#[cfg(feature = "ring-crypto")]
rng: ring::rand::SystemRandom,
#[cfg(all(feature = "aws-lc-crypto", not(feature = "ring-crypto")))]
rng: aws_lc_rs::rand::SystemRandom,
}
impl MaskKeySource {
pub(crate) fn new() -> Result<Self, EntropyUnavailable> {
let me = Self::new_uninit();
let mut probe = [0u8; 4];
me.fill(&mut probe)?;
Ok(me)
}
#[cfg(feature = "ring-crypto")]
fn new_uninit() -> Self {
Self {
rng: ring::rand::SystemRandom::new(),
}
}
#[cfg(all(feature = "aws-lc-crypto", not(feature = "ring-crypto")))]
fn new_uninit() -> Self {
Self {
rng: aws_lc_rs::rand::SystemRandom::new(),
}
}
#[inline]
#[allow(dead_code)]
pub(crate) fn next_key(&self) -> Result<[u8; 4], EntropyUnavailable> {
let mut key = [0u8; 4];
self.fill(&mut key)?;
Ok(key)
}
#[cfg(feature = "ring-crypto")]
pub(crate) fn fill(&self, buf: &mut [u8]) -> Result<(), EntropyUnavailable> {
use ring::rand::SecureRandom;
self.rng
.fill(buf)
.map_err(|e| EntropyUnavailable(format!("system entropy source unavailable: {e:?}")))
}
#[cfg(all(feature = "aws-lc-crypto", not(feature = "ring-crypto")))]
pub(crate) fn fill(&self, buf: &mut [u8]) -> Result<(), EntropyUnavailable> {
use aws_lc_rs::rand::SecureRandom;
self.rng
.fill(buf)
.map_err(|e| EntropyUnavailable(format!("system entropy source unavailable: {e:?}")))
}
}
#[inline]
pub(crate) fn apply_mask(buf: &mut [u8], mask_key: [u8; 4], start_offset: usize) {
let phase = start_offset & 3;
let rotated_mask = [
mask_key[phase],
mask_key[(phase + 1) & 3],
mask_key[(phase + 2) & 3],
mask_key[(phase + 3) & 3],
];
apply_mask_rotated(buf, rotated_mask);
}
#[inline]
fn apply_mask_rotated(buf: &mut [u8], mask: [u8; 4]) {
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") {
unsafe { apply_mask_avx2(buf, mask) };
} else {
unsafe { apply_mask_sse2(buf, mask) };
}
}
#[cfg(target_arch = "aarch64")]
{
unsafe { apply_mask_neon(buf, mask) };
}
#[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
apply_mask_scalar(buf, mask);
}
#[inline]
fn apply_mask_scalar(buf: &mut [u8], mask: [u8; 4]) {
for (i, b) in buf.iter_mut().enumerate() {
*b ^= mask[i & 3];
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse2")]
unsafe fn apply_mask_sse2(buf: &mut [u8], mask: [u8; 4]) {
use std::arch::x86_64::{
__m128i, _mm_loadu_si128, _mm_set1_epi32, _mm_storeu_si128, _mm_xor_si128,
};
let mask_vec = _mm_set1_epi32(i32::from_le_bytes(mask));
let len = buf.len();
let mut i = 0;
while i + 16 <= len {
unsafe {
let p = buf.as_mut_ptr().add(i) as *mut __m128i;
let v = _mm_loadu_si128(p);
let x = _mm_xor_si128(v, mask_vec);
_mm_storeu_si128(p, x);
}
i += 16;
}
apply_mask_scalar(&mut buf[i..], mask);
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn apply_mask_avx2(buf: &mut [u8], mask: [u8; 4]) {
use std::arch::x86_64::{
__m128i, __m256i, _mm_loadu_si128, _mm_set1_epi32, _mm_storeu_si128, _mm_xor_si128,
_mm256_loadu_si256, _mm256_set1_epi32, _mm256_storeu_si256, _mm256_xor_si256,
};
let mask_u32 = i32::from_le_bytes(mask);
let mask256 = _mm256_set1_epi32(mask_u32);
let mask128 = _mm_set1_epi32(mask_u32);
let len = buf.len();
let mut i = 0;
while i + 32 <= len {
unsafe {
let p = buf.as_mut_ptr().add(i) as *mut __m256i;
let v = _mm256_loadu_si256(p);
let x = _mm256_xor_si256(v, mask256);
_mm256_storeu_si256(p, x);
}
i += 32;
}
while i + 16 <= len {
unsafe {
let p = buf.as_mut_ptr().add(i) as *mut __m128i;
let v = _mm_loadu_si128(p);
let x = _mm_xor_si128(v, mask128);
_mm_storeu_si128(p, x);
}
i += 16;
}
apply_mask_scalar(&mut buf[i..], mask);
}
#[cfg(target_arch = "aarch64")]
unsafe fn apply_mask_neon(buf: &mut [u8], mask: [u8; 4]) {
use std::arch::aarch64::{
uint8x16_t, vdupq_n_u32, veorq_u8, vld1q_u8, vreinterpretq_u8_u32, vst1q_u8,
};
let mask_vec: uint8x16_t =
unsafe { vreinterpretq_u8_u32(vdupq_n_u32(u32::from_le_bytes(mask))) };
let len = buf.len();
let mut i = 0;
while i + 16 <= len {
unsafe {
let p = buf.as_mut_ptr().add(i);
let v = vld1q_u8(p);
let x = veorq_u8(v, mask_vec);
vst1q_u8(p, x);
}
i += 16;
}
apply_mask_scalar(&mut buf[i..], mask);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mask_key_source_constructs() {
let _ = MaskKeySource::new().expect("system entropy must be available in tests");
}
#[test]
fn mask_keys_are_non_zero_in_aggregate() {
let rng = MaskKeySource::new().expect("system entropy");
let mut all_zero_streak = 0;
for _ in 0..10 {
if rng.next_key().expect("entropy draw") == [0; 4] {
all_zero_streak += 1;
}
}
assert!(all_zero_streak < 10, "OS CSPRNG appears to be broken");
}
#[test]
fn mask_keys_are_independently_sampled() {
let rng = MaskKeySource::new().expect("system entropy");
let mut seen = std::collections::HashSet::new();
let mut collisions = 0;
for _ in 0..10_000 {
if !seen.insert(rng.next_key().expect("entropy draw")) {
collisions += 1;
}
}
assert!(collisions <= 5, "{collisions} duplicates in 10000");
}
#[test]
fn apply_mask_round_trips() {
let key = [0x11, 0x22, 0x33, 0x44];
let plaintext = b"the quick brown fox jumps over the lazy dog";
let mut buf = plaintext.to_vec();
apply_mask(&mut buf, key, 0);
assert_ne!(buf, plaintext); apply_mask(&mut buf, key, 0); assert_eq!(buf, plaintext);
}
#[test]
fn apply_mask_chunks_match_full() {
let key = [0xAA, 0xBB, 0xCC, 0xDD];
let plaintext: Vec<u8> = (0..1000u32).map(|i| i as u8).collect();
let mut full = plaintext.clone();
apply_mask(&mut full, key, 0);
let mut chunked = plaintext.clone();
let mut off = 0;
for c in chunked.chunks_mut(7) {
apply_mask(c, key, off);
off += c.len();
}
assert_eq!(full, chunked);
}
#[test]
fn apply_mask_handles_short_buffers() {
for len in 0..16 {
let key = [0x5A; 4];
let mut buf = vec![0u8; len];
apply_mask(&mut buf, key, 0);
for (i, b) in buf.iter().enumerate() {
assert_eq!(*b, key[i & 3], "len={len} i={i}");
}
}
}
fn apply_mask_reference(buf: &mut [u8], mask_key: [u8; 4], start_offset: usize) {
for (i, b) in buf.iter_mut().enumerate() {
*b ^= mask_key[(i + start_offset) & 3];
}
}
#[test]
fn apply_mask_simd_matches_scalar_across_lengths() {
let key = [0xAA, 0xBB, 0xCC, 0xDD];
for len in 0..=160usize {
for phase in 0..4 {
let plaintext: Vec<u8> = (0..len).map(|i| (i * 7 + 11) as u8).collect();
let mut expected = plaintext.clone();
apply_mask_reference(&mut expected, key, phase);
let mut actual = plaintext.clone();
apply_mask(&mut actual, key, phase);
assert_eq!(actual, expected, "len={len} phase={phase}");
}
}
}
#[test]
fn apply_mask_simd_matches_scalar_at_size_boundaries() {
let key = [0x01, 0x23, 0x45, 0x67];
for len in [
0_usize, 1, 3, 4, 7, 8, 15, 16, 17, 31, 32, 33, 47, 48, 49, 63, 64, 65, 95, 96, 97,
127, 128, 129,
] {
for phase in 0..4 {
let plaintext: Vec<u8> = (0..len).map(|i| i as u8).collect();
let mut expected = plaintext.clone();
apply_mask_reference(&mut expected, key, phase);
let mut actual = plaintext.clone();
apply_mask(&mut actual, key, phase);
assert_eq!(actual, expected, "len={len} phase={phase}");
}
}
}
#[test]
fn apply_mask_simd_matches_scalar_for_large_payload() {
let key = [0xDE, 0xAD, 0xBE, 0xEF];
let len = (1 << 20) + 37; let plaintext: Vec<u8> = (0..len).map(|i| ((i * 13) ^ 0x5A) as u8).collect();
let mut expected = plaintext.clone();
apply_mask_reference(&mut expected, key, 0);
let mut actual = plaintext.clone();
apply_mask(&mut actual, key, 0);
assert_eq!(actual.len(), expected.len());
assert_eq!(&actual[..256], &expected[..256], "head mismatch");
assert_eq!(
&actual[len - 256..],
&expected[len - 256..],
"tail mismatch"
);
assert!(actual == expected, "bulk mismatch in 1 MiB payload");
}
fn apply_mask_dispatched_scalar(buf: &mut [u8], mask_key: [u8; 4], start_offset: usize) {
let phase = start_offset & 3;
let rotated = [
mask_key[phase],
mask_key[(phase + 1) & 3],
mask_key[(phase + 2) & 3],
mask_key[(phase + 3) & 3],
];
apply_mask_scalar(buf, rotated);
}
#[test]
#[ignore]
fn apply_mask_bench() {
use std::time::Instant;
let key = [0x10, 0x20, 0x30, 0x40];
let sizes_kib = [1, 4, 16, 64, 256, 1024];
let iterations = 100;
println!("\napply_mask bench (single-threaded, {iterations} iterations per size)");
println!(
" {:>10} {:>14} {:>14} {:>10}",
"size", "scalar GB/s", "simd GB/s", "speedup"
);
for &kib in &sizes_kib {
let len = kib * 1024;
let plaintext: Vec<u8> = (0..len).map(|i| (i ^ 0x5A) as u8).collect();
let mut buf = plaintext.clone();
let start = Instant::now();
for _ in 0..iterations {
apply_mask_dispatched_scalar(&mut buf, key, 0);
}
let scalar_elapsed = start.elapsed();
let mut buf = plaintext.clone();
let start = Instant::now();
for _ in 0..iterations {
apply_mask(&mut buf, key, 0);
}
let simd_elapsed = start.elapsed();
let total_bytes = (len * iterations) as f64;
let scalar_gbps = total_bytes / scalar_elapsed.as_secs_f64() / 1e9;
let simd_gbps = total_bytes / simd_elapsed.as_secs_f64() / 1e9;
let speedup = scalar_elapsed.as_secs_f64() / simd_elapsed.as_secs_f64();
println!(
" {:>8} K {:>14.2} {:>14.2} {:>9.2}x",
kib, scalar_gbps, simd_gbps, speedup
);
}
}
#[test]
fn apply_mask_simd_matches_scalar_under_chunked_calls() {
let key = [0x10, 0x20, 0x30, 0x40];
let plaintext: Vec<u8> = (0..2048u32).map(|i| (i ^ 0x37) as u8).collect();
let mut full = plaintext.clone();
apply_mask(&mut full, key, 0);
for chunk_size in [1usize, 3, 7, 13, 16, 17, 31, 32, 33] {
let mut chunked = plaintext.clone();
let mut off = 0;
for c in chunked.chunks_mut(chunk_size) {
apply_mask(c, key, off);
off += c.len();
}
assert_eq!(full, chunked, "chunk_size={chunk_size}");
}
}
}