use crate::constants::*;
use crate::sponge::Sponge;
use argon2::{self, Argon2};
use subtle::ConstantTimeEq;
use zeroize::{Zeroize, ZeroizeOnDrop};
const REKEY_INTERVAL: usize = 64 * 1024; const PARALLEL_THRESHOLD: usize = 1024 * 512;
const PROCESSING_CHUNK_SIZE: usize = 1024;
#[derive(Debug, PartialEq, Eq)]
pub enum Error {
InvalidTag,
InvalidState,
Argon2Error(String),
}
#[derive(Zeroize, ZeroizeOnDrop)]
pub struct Quasor {
#[cfg(test)]
pub master_key: [u8; 32],
#[cfg(not(test))]
master_key: [u8; 32],
}
impl Quasor {
pub fn new(password: &[u8], salt: &[u8]) -> Result<Self, Error> {
let mut master_key = [0u8; 32];
let argon2 = Argon2::default();
argon2
.hash_password_into(password, salt, &mut master_key)
.map_err(|e| Error::Argon2Error(e.to_string()))?;
Ok(Self { master_key })
}
pub fn from_raw_key(key: [u8; 32]) -> Self {
Self { master_key: key }
}
pub fn encrypt(&self, plaintext: &[u8], ad: &[u8]) -> (Vec<u8>, Vec<u8>, Vec<u8>) {
let nonce = self.derive_nonce(plaintext, ad);
let mut enc_state = EncryptState::new(&self.master_key, &nonce, ad);
let ciphertext = enc_state.process(plaintext);
let tag = enc_state.finalize();
(ciphertext, tag.to_vec(), nonce.to_vec())
}
pub fn decrypt(
&self,
nonce: &[u8],
ciphertext: &[u8],
ad: &[u8],
tag: &[u8],
) -> Result<Vec<u8>, Error> {
let mut dec_state = DecryptState::new(&self.master_key, nonce, ad);
let plaintext = dec_state.process(ciphertext);
dec_state.verify_tag(tag)?;
let derived_nonce = self.derive_nonce(&plaintext, ad);
if !bool::from(derived_nonce.as_slice().ct_eq(nonce)) {
return Err(Error::InvalidTag);
}
Ok(plaintext)
}
pub fn create_encrypt_stream<'a>(
&'a self,
plaintext: &'a [u8],
ad: &'a [u8],
) -> QuasorEncryptStream<'a> {
QuasorEncryptStream::new(&self.master_key, plaintext, ad)
}
pub fn create_decrypt_stream<'a>(
&'a self,
nonce: &'a [u8],
ad: &'a [u8],
tag: &'a [u8],
) -> QuasorDecryptStream<'a> {
QuasorDecryptStream::new(&self.master_key, nonce, ad, tag)
}
pub fn derive_nonce(&self, plaintext: &[u8], ad: &[u8]) -> [u8; 16] {
let mut hasher = blake3::Hasher::new_keyed(&self.master_key);
hasher.update(&(ad.len() as u64).to_le_bytes());
hasher.update(ad);
hasher.update(&(plaintext.len() as u64).to_le_bytes());
if plaintext.len() > PARALLEL_THRESHOLD {
hasher.update_rayon(plaintext);
} else {
hasher.update(plaintext);
}
let mut output = [0u8; 16];
let mut output_reader = hasher.finalize_xof();
output_reader.fill(&mut output);
output
}
}
pub struct QuasorEncryptStream<'a> {
master_key: &'a [u8; 32],
plaintext: &'a [u8],
ad: &'a [u8],
enc_state: Option<EncryptState>,
nonce: Option<[u8; 16]>,
tag: Option<[u8; TAG_SIZE]>,
}
impl<'a> QuasorEncryptStream<'a> {
fn new(master_key: &'a [u8; 32], plaintext: &'a [u8], ad: &'a [u8]) -> Self {
Self {
master_key,
plaintext,
ad,
enc_state: None,
nonce: None,
tag: None,
}
}
pub fn init(&mut self) {
let mut hasher = blake3::Hasher::new_keyed(self.master_key);
hasher.update(&(self.ad.len() as u64).to_le_bytes());
hasher.update(self.ad);
hasher.update(&(self.plaintext.len() as u64).to_le_bytes());
if self.plaintext.len() > PARALLEL_THRESHOLD {
hasher.update_rayon(self.plaintext);
} else {
hasher.update(self.plaintext);
}
let mut output = [0u8; 16];
let mut output_reader = hasher.finalize_xof();
output_reader.fill(&mut output);
self.nonce = Some(output);
self.enc_state = Some(EncryptState::new(self.master_key, &output, self.ad));
}
pub fn get_nonce(&self) -> Result<&[u8; 16], Error> {
self.nonce.as_ref().ok_or(Error::InvalidState)
}
pub fn process(&mut self, plaintext_chunk: &[u8]) -> Result<Vec<u8>, Error> {
let state = self.enc_state.as_mut().ok_or(Error::InvalidState)?;
Ok(state.process(plaintext_chunk))
}
pub fn finalize(&mut self) -> Result<[u8; TAG_SIZE], Error> {
if self.tag.is_some() {
return Ok(self.tag.unwrap());
}
let state = self.enc_state.as_mut().ok_or(Error::InvalidState)?;
let tag = state.finalize();
self.tag = Some(tag);
Ok(tag)
}
}
pub struct QuasorDecryptStream<'a> {
dec_state: DecryptState,
master_key: &'a [u8; 32],
nonce: &'a [u8],
ad: &'a [u8],
tag: &'a [u8],
plaintext_chunks: Vec<Vec<u8>>,
}
impl<'a> QuasorDecryptStream<'a> {
fn new(master_key: &'a [u8; 32], nonce: &'a [u8], ad: &'a [u8], tag: &'a [u8]) -> Self {
Self {
dec_state: DecryptState::new(master_key, nonce, ad),
master_key,
nonce,
ad,
tag,
plaintext_chunks: Vec::new(),
}
}
pub fn process(&mut self, ciphertext_chunk: &[u8]) -> Vec<u8> {
let plaintext_chunk = self.dec_state.process(ciphertext_chunk);
self.plaintext_chunks.push(plaintext_chunk.clone());
plaintext_chunk
}
pub fn finalize(&mut self) -> Result<(), Error> {
self.dec_state.verify_tag(self.tag)?;
let plaintext: Vec<u8> = self.plaintext_chunks.concat();
let mut hasher = blake3::Hasher::new_keyed(self.master_key);
hasher.update(&(self.ad.len() as u64).to_le_bytes());
hasher.update(self.ad);
hasher.update(&(plaintext.len() as u64).to_le_bytes());
if plaintext.len() > PARALLEL_THRESHOLD {
hasher.update_rayon(&plaintext);
} else {
hasher.update(&plaintext);
}
let mut derived_nonce = [0u8; 16];
let mut output_reader = hasher.finalize_xof();
output_reader.fill(&mut derived_nonce);
if !bool::from(derived_nonce.as_slice().ct_eq(self.nonce)) {
return Err(Error::InvalidTag);
}
Ok(())
}
}
pub struct EncryptState {
sponge: Sponge,
rekey_counter: u64,
bytes_processed_since_rekey: usize,
encryption_started: bool,
}
impl EncryptState {
pub fn new(key: &[u8], nonce: &[u8], ad: &[u8]) -> Self {
let mut sponge = Sponge::new();
sponge.absorb(DOMAIN_INIT);
sponge.absorb(key);
sponge.absorb(nonce);
sponge.absorb(ad);
Self {
sponge,
rekey_counter: 0,
bytes_processed_since_rekey: 0,
encryption_started: false,
}
}
pub fn process(&mut self, plaintext: &[u8]) -> Vec<u8> {
let mut ciphertext = Vec::with_capacity(plaintext.len());
if !self.encryption_started {
self.sponge.absorb(DOMAIN_ENCRYPT);
self.encryption_started = true;
}
for chunk in plaintext.chunks(PROCESSING_CHUNK_SIZE) {
let mut squeeze_sponge = self.sponge.fork();
let keystream = squeeze_sponge.squeeze(chunk.len());
let mut cipher_chunk = Vec::with_capacity(chunk.len());
for (p, k) in chunk.iter().zip(keystream.iter()) {
cipher_chunk.push(p ^ k);
}
self.sponge.absorb(chunk);
ciphertext.extend_from_slice(&cipher_chunk);
self.bytes_processed_since_rekey += chunk.len();
if self.bytes_processed_since_rekey >= REKEY_INTERVAL {
self.perform_rekey();
self.bytes_processed_since_rekey = 0;
}
}
ciphertext
}
fn perform_rekey(&mut self) {
self.sponge.absorb(DOMAIN_REKEY);
self.sponge.absorb(&self.rekey_counter.to_le_bytes());
let mut squeeze_sponge = self.sponge.fork();
let new_key = squeeze_sponge.squeeze(MASTER_KEY_SIZE);
self.sponge.absorb(&new_key);
self.rekey_counter += 1;
}
pub fn finalize(&mut self) -> [u8; TAG_SIZE] {
self.sponge.absorb(DOMAIN_AUTH);
let tag_vec = self.sponge.squeeze(TAG_SIZE);
let mut tag = [0u8; TAG_SIZE];
tag.copy_from_slice(&tag_vec);
tag
}
}
pub struct DecryptState {
sponge: Sponge,
rekey_counter: u64,
bytes_processed_since_rekey: usize,
encryption_started: bool,
}
impl DecryptState {
pub fn new(key: &[u8], nonce: &[u8], ad: &[u8]) -> Self {
let mut sponge = Sponge::new();
sponge.absorb(DOMAIN_INIT);
sponge.absorb(key);
sponge.absorb(nonce);
sponge.absorb(ad);
Self {
sponge,
rekey_counter: 0,
bytes_processed_since_rekey: 0,
encryption_started: false,
}
}
pub fn process(&mut self, ciphertext: &[u8]) -> Vec<u8> {
let mut plaintext = Vec::with_capacity(ciphertext.len());
if !self.encryption_started {
self.sponge.absorb(DOMAIN_ENCRYPT);
self.encryption_started = true;
}
for chunk in ciphertext.chunks(PROCESSING_CHUNK_SIZE) {
let mut squeeze_sponge = self.sponge.fork();
let keystream = squeeze_sponge.squeeze(chunk.len());
let mut plain_chunk = Vec::with_capacity(chunk.len());
for (c, k) in chunk.iter().zip(keystream.iter()) {
plain_chunk.push(c ^ k);
}
self.sponge.absorb(&plain_chunk);
plaintext.extend_from_slice(&plain_chunk);
self.bytes_processed_since_rekey += chunk.len();
if self.bytes_processed_since_rekey >= REKEY_INTERVAL {
self.perform_rekey();
self.bytes_processed_since_rekey = 0;
}
}
plaintext
}
fn perform_rekey(&mut self) {
self.sponge.absorb(DOMAIN_REKEY);
self.sponge.absorb(&self.rekey_counter.to_le_bytes());
let mut squeeze_sponge = self.sponge.fork();
let new_key = squeeze_sponge.squeeze(MASTER_KEY_SIZE);
self.sponge.absorb(&new_key);
self.rekey_counter += 1;
}
pub fn verify_tag(&mut self, tag: &[u8]) -> Result<(), Error> {
self.sponge.absorb(DOMAIN_AUTH);
let expected_tag_vec = self.sponge.squeeze(TAG_SIZE);
if bool::from(expected_tag_vec.as_slice().ct_eq(tag)) {
Ok(())
} else {
Err(Error::InvalidTag)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_forward_secrecy_rekeying_unit() {
let key = [99; 32];
let quasor = Quasor::from_raw_key(key);
let ad = b"forward_secrecy_test";
let block1_size = REKEY_INTERVAL - 100;
let block2_size = 200;
let total_size = block1_size + block2_size;
let plaintext = vec![1u8; total_size];
let nonce = quasor.derive_nonce(&plaintext, ad);
let mut enc_state = EncryptState::new(&quasor.master_key, &nonce, ad);
let c_block1 = enc_state.process(&plaintext[0..block1_size]);
enc_state.process(&plaintext[block1_size..total_size]);
let state_after_rekey = enc_state.sponge.clone();
let mut attack_dec_state = DecryptState {
sponge: state_after_rekey,
rekey_counter: 0,
bytes_processed_since_rekey: 0,
encryption_started: true, };
let recovered_p_block1 = attack_dec_state.process(&c_block1);
assert_ne!(
recovered_p_block1,
&plaintext[0..block1_size],
"Forward Secrecy FAILED: Post-rekey state was able to decrypt pre-rekey data!"
);
}
}