use aes_siv::aead::{AeadInPlace, KeyInit};
use aes_siv::{Aes256SivAead, Nonce};
use crate::error::{Error, Result};
pub fn sha3_256_hex(data: &[u8]) -> String {
use sha3::Digest;
let mut h = sha3::Sha3_256::new();
h.update(data);
let out = h.finalize();
out.iter().fold(String::with_capacity(64), |mut s, b| {
use std::fmt::Write;
let _ = write!(s, "{b:02x}");
s
})
}
pub fn aes_256_siv_encrypt(key: &[u8], nonce: &[u8], aad: &[u8], pt: &[u8]) -> Result<Vec<u8>> {
if key.len() != 64 {
return Err(Error::InvalidArg {
arg: "key",
reason: format!("AES-256-SIV key must be 64 bytes, got {}", key.len()),
});
}
if nonce.len() != 16 {
return Err(Error::InvalidArg {
arg: "nonce",
reason: format!("AES-256-SIV nonce must be 16 bytes, got {}", nonce.len()),
});
}
let cipher = Aes256SivAead::new_from_slice(key).map_err(|_| Error::AeadFailed {
alg: "aes-256-siv",
op: "init",
})?;
let mut buf = pt.to_vec();
cipher
.encrypt_in_place(Nonce::from_slice(nonce), aad, &mut buf)
.map_err(|_| Error::AeadFailed {
alg: "aes-256-siv",
op: "encrypt",
})?;
Ok(buf)
}
pub fn aes_256_siv_decrypt(key: &[u8], nonce: &[u8], aad: &[u8], ct: &[u8]) -> Result<Vec<u8>> {
if key.len() != 64 {
return Err(Error::InvalidArg {
arg: "key",
reason: format!("AES-256-SIV key must be 64 bytes, got {}", key.len()),
});
}
if nonce.len() != 16 {
return Err(Error::InvalidArg {
arg: "nonce",
reason: format!("AES-256-SIV nonce must be 16 bytes, got {}", nonce.len()),
});
}
let cipher = Aes256SivAead::new_from_slice(key).map_err(|_| Error::AeadFailed {
alg: "aes-256-siv",
op: "init",
})?;
let mut buf = ct.to_vec();
cipher
.decrypt_in_place(Nonce::from_slice(nonce), aad, &mut buf)
.map_err(|_| Error::AeadFailed {
alg: "aes-256-siv",
op: "decrypt",
})?;
Ok(buf)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sha3_256_known_answer() {
assert_eq!(
sha3_256_hex(b""),
"a7ffc6f8bf1ed76651c14756a061d662f580ff4de43b49fa82d80a4b80f8434a"
);
assert_eq!(
sha3_256_hex(b"abc"),
"3a985da74fe225b2045c172d6bd390bd855f086e3e9d525b46bfe24511431532"
);
}
#[test]
fn aes_256_siv_round_trip_with_various_payloads() {
let key = [0x42u8; 64];
let nonce = [0u8; 16]; for pt in [
b"".to_vec(),
b"hello".to_vec(),
vec![0u8; 4096],
(0..1024).map(|i| (i % 251) as u8).collect::<Vec<_>>(),
] {
let aad = b"file.aep";
let ct = aes_256_siv_encrypt(&key, &nonce, aad, &pt).unwrap();
let rt = aes_256_siv_decrypt(&key, &nonce, aad, &ct).unwrap();
assert_eq!(pt, rt, "round-trip identity failed for len {}", pt.len());
}
}
#[test]
fn aes_256_siv_rejects_bad_key_length() {
let err = aes_256_siv_encrypt(&[0u8; 32], b"n", b"a", b"p").unwrap_err();
assert!(err.to_string().contains("64 bytes"), "got {err}");
}
}