use zeroize::{Zeroize, ZeroizeOnDrop};
use crate::bytearray::ByteArray;
use crate::cipherstate::{CipherState, CipherStates};
use crate::error::{CipherError, CipherResult};
use crate::traits::{Cipher, Hash};
#[derive(ZeroizeOnDrop, Zeroize, Clone)]
pub struct SymmetricState<C, H>
where
C: Cipher,
H: Hash,
{
cipherstate: Option<CipherState<C>>,
h: H::Output,
ck: H::Output,
}
impl<C, H> SymmetricState<C, H>
where
C: Cipher,
H: Hash,
{
pub fn new(noise_pattern_name: &str) -> Self {
let pattern_bytes = noise_pattern_name.as_bytes();
let mut h = H::Output::new_zero();
if pattern_bytes.len() <= H::hash_len() {
h.as_mut()[..pattern_bytes.len()].copy_from_slice(pattern_bytes);
} else {
h = H::hash(pattern_bytes);
}
Self {
cipherstate: None,
ck: h.clone(),
h,
}
}
pub fn new_zero() -> Self {
Self {
cipherstate: None,
ck: H::Output::new_zero(),
h: H::Output::new_zero(),
}
}
pub fn mix_hash(&mut self, data: &[u8]) {
let mut h = H::default();
h.input(self.h.as_slice());
h.input(data);
self.h = h.result();
}
pub fn mix_key(&mut self, input_key_material: &[u8]) {
let (ck, temp_k) = H::hkdf(self.ck.as_slice(), input_key_material);
self.ck = ck;
self.cipherstate = Some(CipherState::new(&temp_k.as_slice()[..C::key_len()], 0));
}
pub fn mix_key_and_hash(&mut self, input_key_material: &[u8]) {
let (ck, temp_h, temp_k) = H::hkdf3(self.ck.as_slice(), input_key_material);
self.ck = ck;
self.mix_hash(temp_h.as_slice());
self.cipherstate = Some(CipherState::new(&temp_k.as_slice()[..C::key_len()], 0));
}
pub fn encrypt_and_hash(&mut self, plaintext: &[u8], out: &mut [u8]) -> CipherResult<()> {
if let Some(ref mut c) = self.cipherstate {
c.encrypt_with_ad(self.h.as_slice(), plaintext, out)?;
} else {
out.copy_from_slice(plaintext);
};
self.mix_hash(out);
Ok(())
}
pub fn decrypt_and_hash(&mut self, data: &[u8], out: &mut [u8]) -> CipherResult<()> {
if let Some(ref mut c) = self.cipherstate {
c.decrypt_with_ad(self.h.as_slice(), data, out)?;
} else {
out.copy_from_slice(data)
}
self.mix_hash(data);
Ok(())
}
pub fn split(&self) -> CipherResult<CipherStates<C>> {
if !self.has_key() {
return Err(CipherError::MissingKeyMaterial);
}
let (mut temp_k1, mut temp_k2) = H::hkdf(self.ck.as_slice(), &[]);
let ct = CipherStates {
initiator_to_responder: CipherState::new(&temp_k1.as_slice()[..C::key_len()], 0),
responder_to_initiator: CipherState::new(&temp_k2.as_slice()[..C::key_len()], 0),
};
temp_k1.zeroize();
temp_k2.zeroize();
Ok(ct)
}
pub fn get_hash(&self) -> H::Output {
self.h.clone()
}
pub fn get_chaining_key(&self) -> H::Output {
self.ck.clone()
}
pub fn has_key(&self) -> bool {
self.cipherstate.is_some()
}
}
#[cfg(test)]
mod tests {
use super::SymmetricState;
use crate::crypto::cipher::{AesGcm, ChaChaPoly};
use crate::crypto::hash::{Blake2b, Blake2s, Sha256, Sha512};
use crate::traits::{Cipher, Hash};
impl<C: Cipher, H: Hash> PartialEq for SymmetricState<C, H> {
fn eq(&self, other: &Self) -> bool {
self.h == other.h && self.ck == other.ck
}
}
fn symmetric_suite<C: Cipher, H: Hash>() {
let mut s1 = SymmetricState::<C, H>::new("complex delirium");
let mut s2 = SymmetricState::<C, H>::new("complex delirium");
assert!(!s1.has_key());
assert!(!s2.has_key());
assert!(s1 == s2);
s1.mix_hash(b"all wound up");
s2.mix_hash(b"all wound up");
assert!(s1 == s2);
assert!(!s1.has_key() && !s2.has_key());
s1.mix_key(b"sleep disturbed");
s2.mix_key(b"sleep disturbed");
assert!(s1 == s2);
assert!(s1.has_key() && s2.has_key());
s1.mix_key_and_hash(b"sleep disturbed");
s2.mix_key_and_hash(b"sleep disturbed");
assert!(s1 == s2);
s1.mix_key_and_hash(&[]);
s2.mix_key_and_hash(&[]);
assert!(s1 == s2);
let mut buf1 = [0; 4096];
let mut buf2 = [0; 4096];
let msg = b"caught off guard";
s1.encrypt_and_hash(msg, &mut buf1[..msg.len() + C::tag_len()])
.unwrap();
assert!(s1 != s2);
assert!(msg != &buf1[..msg.len()]);
s2.decrypt_and_hash(&buf1[..msg.len() + C::tag_len()], &mut buf2[..msg.len()])
.unwrap();
assert_eq!(*msg, buf2[..msg.len()]);
assert!(s1 == s2);
let s1_c = s1.split().unwrap();
let s2_c = s2.split().unwrap();
assert!(s1_c.initiator_to_responder.take() == s2_c.initiator_to_responder.take());
assert!(s1_c.responder_to_initiator.take() == s2_c.responder_to_initiator.take());
s1.mix_key_and_hash(b"run");
s2.mix_key_and_hash(b"try to hide");
assert!(s1 != s2);
s1.encrypt_and_hash(msg, &mut buf1[..msg.len() + C::tag_len()])
.unwrap();
assert!(s1 != s2);
assert!(msg != &buf1[..msg.len()]);
assert!(s2
.decrypt_and_hash(&buf1[..msg.len() + C::tag_len()], &mut buf2[..msg.len()])
.is_err());
cant_split_without_key::<C, H>();
}
fn cant_split_without_key<C: Cipher, H: Hash>() {
let mut s1 = SymmetricState::<C, H>::new("complex delirium");
s1.mix_hash(b"all wound up");
assert!(s1.split().is_err());
}
#[test]
fn symmetric_suites() {
symmetric_suite::<ChaChaPoly, Sha256>();
symmetric_suite::<ChaChaPoly, Sha512>();
symmetric_suite::<ChaChaPoly, Blake2b>();
symmetric_suite::<ChaChaPoly, Blake2s>();
symmetric_suite::<AesGcm, Sha256>();
symmetric_suite::<AesGcm, Sha512>();
symmetric_suite::<AesGcm, Blake2b>();
symmetric_suite::<AesGcm, Blake2s>();
}
}