use alloc::borrow::Cow;
use alloc::vec;
use alloc::vec::Vec;
use super::DecoyStrategy;
use crate::Result;
use crate::error::Error;
use crate::fetcher::RawKey;
#[derive(Debug, Default, Clone, Copy)]
pub struct KeyDerivedDecoy;
impl DecoyStrategy for KeyDerivedDecoy {
fn generate(&self, key: &RawKey, output_len: usize) -> Result<Vec<u8>> {
if output_len == 0 {
return Ok(Vec::new());
}
let mut nonce = [0u8; 32];
getrandom::getrandom(&mut nonce).map_err(|_| Error::Internal("OS RNG failed"))?;
let mut hasher = blake3::Hasher::new();
let _ = hasher.update(key.as_bytes());
let _ = hasher.update(&nonce);
let mut out = vec![0u8; output_len];
let mut reader = hasher.finalize_xof();
reader.fill(&mut out);
unsafe {
let ptr = nonce.as_mut_ptr();
for i in 0..nonce.len() {
core::ptr::write_volatile(ptr.add(i), 0u8);
}
}
core::sync::atomic::compiler_fence(core::sync::atomic::Ordering::SeqCst);
Ok(out)
}
fn describe(&self) -> Cow<'_, str> {
Cow::Borrowed("key-derived")
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
mod tests {
use super::*;
fn raw(bytes: &[u8]) -> RawKey {
RawKey::new(bytes.to_vec())
}
#[test]
fn produces_requested_length() {
let key = raw(b"k");
for n in [0usize, 1, 7, 32, 256, 4096] {
let out = KeyDerivedDecoy.generate(&key, n).unwrap();
assert_eq!(out.len(), n, "wrong length for n = {n}");
}
}
#[test]
fn two_calls_with_same_key_produce_different_outputs() {
let key = raw(b"deterministic seed");
let a = KeyDerivedDecoy.generate(&key, 64).unwrap();
let b = KeyDerivedDecoy.generate(&key, 64).unwrap();
assert_ne!(a, b);
}
#[test]
fn different_keys_produce_different_outputs() {
let a = KeyDerivedDecoy.generate(&raw(b"key one"), 32).unwrap();
let b = KeyDerivedDecoy.generate(&raw(b"key two"), 32).unwrap();
assert_ne!(a, b);
}
#[test]
fn empty_key_is_accepted() {
let out = KeyDerivedDecoy.generate(&raw(&[]), 32).unwrap();
assert_eq!(out.len(), 32);
}
#[test]
fn describe_returns_key_derived() {
assert_eq!(KeyDerivedDecoy.describe(), "key-derived");
}
}