use core::fmt;
use tc_block_cipher::{
BlockCipher, BlockCipherInit, BlockError, CipherDirection, InitError, KeyParams,
};
use tc_zeroize::Zeroize;
use crate::cipher::{MAX_ROUND_KEYS, RoundKeys};
use crate::{ALGO_NAME, BLOCK_BYTES, KEY_BYTES, cipher};
pub struct AriaTableEngine {
round_keys: RoundKeys,
rounds: usize,
initialised: bool,
}
impl AriaTableEngine {
pub const fn new() -> Self {
Self {
round_keys: [[0; BLOCK_BYTES]; MAX_ROUND_KEYS],
rounds: 0,
initialised: false,
}
}
}
impl Default for AriaTableEngine {
fn default() -> Self {
Self::new()
}
}
impl fmt::Display for AriaTableEngine {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(ALGO_NAME)
}
}
impl Drop for AriaTableEngine {
fn drop(&mut self) {
self.round_keys.zeroize();
}
}
impl BlockCipher for AriaTableEngine {
type Error = BlockError;
fn block_size(&self) -> usize {
BLOCK_BYTES
}
fn process_block(&mut self, input: &[u8], output: &mut [u8]) -> Result<usize, BlockError> {
if !self.initialised {
return Err(BlockError::NotInitialised);
}
let (Some(input), Some(output)) = (
input.first_chunk::<BLOCK_BYTES>(),
output.first_chunk_mut::<BLOCK_BYTES>(),
) else {
return Err(BlockError::BufferTooShort);
};
cipher::process_block(&self.round_keys, self.rounds, input, output);
Ok(BLOCK_BYTES)
}
}
impl<P: KeyParams + ?Sized> BlockCipherInit<P> for AriaTableEngine {
type Error = InitError;
fn init(&mut self, direction: CipherDirection, params: &P) -> Result<(), InitError> {
let key = params.key();
if !KEY_BYTES.contains(&key.len()) {
return Err(InitError::InvalidKeyLength(key.len()));
}
let for_encryption = direction == CipherDirection::Encrypt;
(self.round_keys, self.rounds) = cipher::key_schedule(for_encryption, key);
self.initialised = true;
Ok(())
}
}