use crate::{
CopyToSlice, Error, Padding,
crypto::{Backend, BackendEncryptor, Scheme, xor},
std::ops::{Deref, DerefMut},
};
pub struct KeyingState<const KEY_SIZE: usize, AlgorithmT>
where
Scheme<KEY_SIZE, AlgorithmT>: Backend<KEY_SIZE>,
{
backend: Scheme<KEY_SIZE, AlgorithmT>,
k1: [u8; KEY_SIZE],
k2: [u8; KEY_SIZE],
}
impl<const KEY_SIZE: usize, AlgorithmT> KeyingState<KEY_SIZE, AlgorithmT>
where
Scheme<KEY_SIZE, AlgorithmT>: Backend<KEY_SIZE>,
{
pub fn generate_cmac_short(&mut self, header: &[u8], data: &[&[u8]]) -> [u8; 8] {
self.generate_cmac(header, data)[..8].try_into().unwrap()
}
pub fn generate_cmac(&mut self, header: &[u8], data: &[&[u8]]) -> [u8; KEY_SIZE] {
let mut working_block = [0u8; KEY_SIZE];
let (mut data, data_len) = {
let data_len = data.iter().fold(header.len(), |n, data| n + data.len());
(
header.iter().chain(data.iter().flat_map(|v| v.iter())),
data_len,
)
};
let padded = !data_len.is_multiple_of(KEY_SIZE);
let mut encryptor = self.encryptor();
let mut data_n = 0;
while let Some(n) = (&mut data)
.take(KEY_SIZE)
.copied()
.copy_to_slice(&mut working_block)
{
data_n += n;
if data_n == data_len {
if padded {
Padding::pad(&mut working_block, n);
xor(&mut working_block, &self.k2);
} else {
xor(&mut working_block, &self.k1);
}
encryptor.encrypt(&mut working_block);
break;
}
encryptor.encrypt(&mut working_block);
}
self.set_iv(working_block);
working_block
}
pub fn validate_cmac<'a, IoBackendErrorT>(
&mut self,
data: &'a [u8],
trailer: Option<&[u8]>,
) -> Result<&'a [u8], Error<IoBackendErrorT>> {
if data.len() < 8 {
return Err(Error::InvalidSignature);
}
let (data, read_cmac) = data.split_at(data.len() - 8);
let computed_cmac = self.generate_cmac_short(data, &[trailer.unwrap_or(&[])]);
if read_cmac != computed_cmac {
return Err(Error::InvalidSignature);
}
Ok(data)
}
}
impl<const KEY_SIZE: usize, AlgorithmT> Deref for KeyingState<KEY_SIZE, AlgorithmT>
where
Scheme<KEY_SIZE, AlgorithmT>: Backend<KEY_SIZE>,
{
type Target = Scheme<KEY_SIZE, AlgorithmT>;
fn deref(&self) -> &Scheme<KEY_SIZE, AlgorithmT> {
&self.backend
}
}
impl<const KEY_SIZE: usize, AlgorithmT> DerefMut for KeyingState<KEY_SIZE, AlgorithmT>
where
Scheme<KEY_SIZE, AlgorithmT>: Backend<KEY_SIZE>,
{
fn deref_mut(&mut self) -> &mut Scheme<KEY_SIZE, AlgorithmT> {
&mut self.backend
}
}
impl KeyingState<8, des::Des> {
pub fn new(key: [u8; 8]) -> Self {
let (k1, k2) = {
let session_key = Scheme::<8, des::Des>::new(key);
session_key.generate_cmac_keys()
};
Self {
k1,
k2,
backend: Scheme::<8, _>::new(key),
}
}
}
impl KeyingState<16, aes::Aes128> {
pub fn new(key: [u8; 16]) -> Self {
let (k1, k2) = {
let session_key = Scheme::<16, aes::Aes128>::new(key);
session_key.generate_cmac_keys()
};
Self {
k1,
k2,
backend: Scheme::<16, _>::new(key),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use hex_literal::hex;
#[test]
fn test_capture() {
let mut ks =
KeyingState::<16, _>::new(hex!("DE 04 17 85 F5 9C 23 F5 C4 EB A7 EE B7 89 78 55"));
let cmac = ks.generate_cmac_short(&hex!("CD 01 03 00 00 1C 00 00"), &[]);
assert_eq!(hex!("EF 93 44 CC 38 E2 A9 F0"), cmac);
let msg = ks
.validate_cmac::<()>(&hex!("00 19 AC EF 56 46 0F CA DB"), None)
.unwrap();
assert_eq!(&hex!("00"), msg);
}
}