#[cfg(not(test))]
use rand_chacha::rand_core::RngCore;
#[cfg(all(not(test), not(feature = "std")))]
use rand_chacha::{rand_core::SeedableRng, ChaCha20Rng};
#[cfg(all(not(test), not(feature = "std")))]
use crate::seed;
use crate::{
impl_command_display, impl_command_ops, impl_default, impl_encrypted_message_ops,
impl_message_from_buf, len, AesKey, CommandOps, Error, MessageOps, Result, SequenceCount,
};
use super::{encrypted_index as index, WrappedEncryptedMessage};
#[repr(C)]
#[derive(Clone, Debug, PartialEq, zeroize::Zeroize, zeroize::ZeroizeOnDrop)]
pub struct EncryptedCommand {
buf: [u8; len::ENCRYPTED_COMMAND],
}
impl EncryptedCommand {
pub fn new() -> Self {
let mut msg = Self {
buf: [0u8; len::ENCRYPTED_COMMAND],
};
msg.init();
msg
}
pub fn count(&self) -> SequenceCount {
self.count_buf().into()
}
fn count_buf(&self) -> &[u8] {
self.buf[index::COUNT..index::COUNT_END].as_ref()
}
fn set_count(&mut self, count: SequenceCount) {
self.buf[index::COUNT..index::COUNT_END]
.copy_from_slice(count.as_inner().to_le_bytes().as_ref());
}
pub fn with_count(mut self, count: SequenceCount) -> Self {
self.set_count(count);
self
}
pub fn message_data(&self) -> &[u8] {
let start = self.data_start();
let end = self.data_end();
self.buf[start..end].as_ref()
}
fn data_start(&self) -> usize {
index::DATA
}
fn data_end(&self) -> usize {
self.data_start() + self.data_len()
}
pub fn set_message_data(&mut self, message: &dyn CommandOps) -> Result<()> {
let len = message.data_len();
if message.data().len() != len {
return Err(Error::InvalidDataLength((len, message.data().len())));
}
if (0..=len::MAX_ENCRYPTED_DATA).contains(&len) {
self.set_data_len(len as u8);
let start = self.data_start();
let end = self.data_end();
let data = message.data();
log::trace!("Encrypted data: {data:x?}, length: {len}");
self.buf[start..end].copy_from_slice(data);
Ok(())
} else {
Err(Error::InvalidDataLength((len, len::MAX_ENCRYPTED_DATA)))
}
}
pub fn with_message_data(mut self, message: &dyn CommandOps) -> Result<Self> {
self.set_message_data(message)?;
Ok(self)
}
pub fn packing(&self) -> &[u8] {
let start = self.packing_start();
let end = self.packing_end();
self.buf[start..end].as_ref()
}
#[cfg(not(feature = "std"))]
pub fn set_packing(&mut self) {
if self.packing_len() == 0 {
return;
}
#[cfg(not(test))]
let mut rng = ChaCha20Rng::from_seed(seed(self.buf(), self.count_buf()));
let start = self.packing_start();
let end = self.packing_end();
#[cfg(not(test))]
rng.fill_bytes(&mut self.buf[start..end]);
#[cfg(test)]
self.buf[start..end].copy_from_slice([0; 255][..end - start].as_ref());
}
#[cfg(feature = "std")]
pub fn set_packing(&mut self) {
if self.packing_len() == 0 {
return;
}
#[cfg(not(test))]
let mut rng = rand::thread_rng();
let start = self.packing_start();
let end = self.packing_end();
#[cfg(not(test))]
rng.fill_bytes(&mut self.buf[start..end]);
#[cfg(test)]
self.buf[start..end].copy_from_slice([0; 255][..end - start].as_ref());
}
fn packing_start(&self) -> usize {
self.data_end()
}
fn packing_end(&self) -> usize {
self.packing_start() + self.packing_len()
}
pub fn packing_len(&self) -> usize {
let meta = len::METADATA - 1;
let raw_len = meta + self.data_len();
len::aes_packing_len(raw_len)
}
fn encrypt_data(&mut self) -> &mut [u8] {
let len = self.len();
self.buf[index::LEN..len].as_mut()
}
pub fn encrypt(mut self, key: &AesKey) -> WrappedEncryptedMessage {
use aes::cipher::{BlockEncrypt, KeyInit};
self.set_packing();
if super::sequence_count().as_inner() != 0 {
self.set_count(super::sequence_count());
}
self.calculate_checksum();
let mut enc_msg = WrappedEncryptedMessage::new();
let enc_len = self.len();
enc_msg.set_data_len(enc_len as u8);
log::trace!("Encrypted message: {:x?}", self.buf());
let plain_data = self.encrypt_data();
let cipher_data = enc_msg.data_mut()[1..].as_mut();
let ciph = aes::Aes128::new(key);
for (pchunk, cchunk) in plain_data
.chunks_exact(16)
.zip(cipher_data.chunks_exact_mut(16))
{
ciph.encrypt_block_b2b(pchunk.into(), cchunk.into());
}
enc_msg.calculate_checksum();
if let Err(err) = enc_msg.verify_checksum() {
log::error!("error validating wrapped encrypted checksum: {err}");
}
if let Err(err) = enc_msg.stuff_encrypted_data() {
log::error!("error stuffing encrypted command message: {err}");
}
let count = super::sequence_count();
let next_count = super::increment_sequence_count();
log::trace!("encryption sequence count: {count}");
log::trace!("next encryption sequence count: {next_count}");
enc_msg
}
pub fn decrypt(key: &AesKey, mut message: WrappedEncryptedMessage) -> Self {
use crate::aes;
if let Err(err) = message.unstuff_encrypted_data() {
log::error!("error unstuffing encrypted command message: {err}");
}
let mut dec_msg = Self::new();
dec_msg.set_data_len(message.data_len().saturating_sub(len::ENCRYPTED_METADATA) as u8);
let cipher_data = message.data()[1..].as_ref();
let plain_data = dec_msg.encrypt_data();
if let Err(err) = aes::aes_decrypt_inplace(key.as_ref(), cipher_data, plain_data) {
log::error!("error decrypting message: {err}");
}
super::increment_sequence_count();
dec_msg
}
}
impl_default!(EncryptedCommand);
impl_command_display!(EncryptedCommand);
impl_message_from_buf!(EncryptedCommand);
impl_encrypted_message_ops!(EncryptedCommand);
impl_command_ops!(EncryptedCommand);