use core::fmt;
use super::AesBlock;
use crate::incremental::{check_ad_length, tags_match, MessageLength, Quarantine};
use crate::wipe::wipe_value;
pub use crate::Error;
pub type Key = [u8; 16];
pub type Nonce = [u8; 16];
#[repr(transparent)]
#[derive(Debug, Clone, Copy)]
struct State {
blocks: [AesBlock; 8],
}
impl State {
#[inline(always)]
fn update(&mut self, d1: AesBlock, d2: AesBlock) {
let blocks = &mut self.blocks;
let tmp = blocks[7];
blocks[7] = blocks[6].round(blocks[7]);
blocks[6] = blocks[5].round(blocks[6]);
blocks[5] = blocks[4].round(blocks[5]);
blocks[4] = blocks[3].round(blocks[4]).xor(d2);
blocks[3] = blocks[2].round(blocks[3]);
blocks[2] = blocks[1].round(blocks[2]);
blocks[1] = blocks[0].round(blocks[1]);
blocks[0] = tmp.round(blocks[0]).xor(d1);
}
#[inline(always)]
pub fn new(key: &Key, nonce: &Nonce) -> Self {
let c0 = AesBlock::from_bytes(&[
0x00, 0x01, 0x01, 0x02, 0x03, 0x05, 0x08, 0x0d, 0x15, 0x22, 0x37, 0x59, 0x90, 0xe9,
0x79, 0x62,
]);
let c1 = AesBlock::from_bytes(&[
0xdb, 0x3d, 0x18, 0x55, 0x6d, 0xc2, 0x2f, 0xf1, 0x20, 0x11, 0x31, 0x42, 0x73, 0xb5,
0x28, 0xdd,
]);
let key_block = AesBlock::from_bytes(key);
let nonce_block = AesBlock::from_bytes(nonce);
let blocks: [AesBlock; 8] = [
key_block.xor(nonce_block),
c1,
c0,
c1,
key_block.xor(nonce_block),
key_block.xor(c0),
key_block.xor(c1),
key_block.xor(c0),
];
let mut state = State { blocks };
for _ in 0..10 {
state.update(nonce_block, key_block);
}
state
}
#[inline(always)]
fn absorb(&mut self, src: &[u8; 32]) {
let msg0 = AesBlock::from_bytes(&src[..16]);
let msg1 = AesBlock::from_bytes(&src[16..32]);
self.update(msg0, msg1);
}
fn absorb_ad(&mut self, ad: &[u8]) {
let mut src = [0u8; 32];
let mut i = 0;
while i + 32 <= ad.len() {
src.copy_from_slice(&ad[i..][..32]);
self.absorb(&src);
i += 32;
}
if ad.len() % 32 != 0 {
src.fill(0);
src[..ad.len() % 32].copy_from_slice(&ad[i..]);
self.absorb(&src);
}
}
#[inline(always)]
fn keystream(&self) -> (AesBlock, AesBlock) {
let blocks = &self.blocks;
(
blocks[6].xor(blocks[1]).xor(blocks[2].and(blocks[3])),
blocks[2].xor(blocks[5]).xor(blocks[6].and(blocks[7])),
)
}
#[inline(always)]
fn enc(&mut self, dst: &mut [u8; 32], src: &[u8; 32]) {
let (z0, z1) = self.keystream();
let msg0 = AesBlock::from_bytes(&src[..16]);
let msg1 = AesBlock::from_bytes(&src[16..32]);
let c0 = msg0.xor(z0);
let c1 = msg1.xor(z1);
dst[..16].copy_from_slice(&c0.to_bytes());
dst[16..32].copy_from_slice(&c1.to_bytes());
self.update(msg0, msg1);
}
#[inline(always)]
fn dec(&mut self, dst: &mut [u8; 32], src: &[u8; 32]) {
let (z0, z1) = self.keystream();
let msg0 = AesBlock::from_bytes(&src[0..16]).xor(z0);
let msg1 = AesBlock::from_bytes(&src[16..32]).xor(z1);
dst[..16].copy_from_slice(&msg0.to_bytes());
dst[16..32].copy_from_slice(&msg1.to_bytes());
self.update(msg0, msg1);
}
#[inline(always)]
fn squeeze_keystream(&self, dst: &mut [u8; 32]) {
let (z0, z1) = self.keystream();
dst[..16].copy_from_slice(&z0.to_bytes());
dst[16..32].copy_from_slice(&z1.to_bytes());
}
#[inline(always)]
fn dec_partial(&mut self, dst: &mut [u8; 32], src: &[u8]) {
let len = src.len();
let mut src_padded = [0u8; 32];
src_padded[..len].copy_from_slice(src);
let (z0, z1) = self.keystream();
let msg_padded0 = AesBlock::from_bytes(&src_padded[0..16]).xor(z0);
let msg_padded1 = AesBlock::from_bytes(&src_padded[16..32]).xor(z1);
dst[..16].copy_from_slice(&msg_padded0.to_bytes());
dst[16..32].copy_from_slice(&msg_padded1.to_bytes());
dst[len..].fill(0);
let msg0 = AesBlock::from_bytes(&dst[0..16]);
let msg1 = AesBlock::from_bytes(&dst[16..32]);
self.update(msg0, msg1);
}
#[inline(always)]
fn mac<const TAG_BYTES: usize>(&mut self, adlen: u64, mlen: u64) -> Tag<TAG_BYTES> {
let tmp = {
let blocks = &self.blocks;
let mut sizes = [0u8; 16];
sizes[..8].copy_from_slice(&(adlen * 8).to_le_bytes());
sizes[8..16].copy_from_slice(&(mlen * 8).to_le_bytes());
AesBlock::from_bytes(&sizes).xor(blocks[2])
};
for _ in 0..7 {
self.update(tmp, tmp);
}
let blocks = &self.blocks;
let mut tag = [0u8; TAG_BYTES];
match TAG_BYTES {
16 => tag.copy_from_slice(
&blocks[0]
.xor(blocks[1])
.xor(blocks[2])
.xor(blocks[3])
.xor(blocks[4])
.xor(blocks[5])
.xor(blocks[6])
.to_bytes(),
),
32 => {
tag[..16].copy_from_slice(
&blocks[0]
.xor(blocks[1])
.xor(blocks[2])
.xor(blocks[3])
.to_bytes(),
);
tag[16..].copy_from_slice(
&blocks[4]
.xor(blocks[5])
.xor(blocks[6])
.xor(blocks[7])
.to_bytes(),
);
}
_ => unreachable!(),
}
tag
}
#[inline(always)]
fn mac_finalize<const TAG_BYTES: usize>(&mut self, data_len: usize) -> Tag<TAG_BYTES> {
let tmp = {
let blocks = &self.blocks;
let mut sizes = [0u8; 16];
sizes[..8].copy_from_slice(&(data_len as u64 * 8).to_le_bytes());
sizes[8..16].copy_from_slice(&(TAG_BYTES as u64 * 8).to_le_bytes());
AesBlock::from_bytes(&sizes).xor(blocks[2])
};
for _ in 0..7 {
self.update(tmp, tmp);
}
let blocks = &self.blocks;
let mut tag = [0u8; TAG_BYTES];
match TAG_BYTES {
16 => tag.copy_from_slice(
&blocks[0]
.xor(blocks[1])
.xor(blocks[2])
.xor(blocks[3])
.xor(blocks[4])
.xor(blocks[5])
.xor(blocks[6])
.to_bytes(),
),
32 => {
tag[..16].copy_from_slice(
&blocks[0]
.xor(blocks[1])
.xor(blocks[2])
.xor(blocks[3])
.to_bytes(),
);
tag[16..].copy_from_slice(
&blocks[4]
.xor(blocks[5])
.xor(blocks[6])
.xor(blocks[7])
.to_bytes(),
);
}
_ => unreachable!(),
}
tag
}
}
#[repr(transparent)]
pub struct Aegis128L<const TAG_BYTES: usize>(State);
pub type Tag<const TAG_BYTES: usize> = [u8; TAG_BYTES];
impl<const TAG_BYTES: usize> Aegis128L<TAG_BYTES> {
pub fn new(key: &Key, nonce: &Nonce) -> Self {
assert!(
TAG_BYTES == 16 || TAG_BYTES == 32,
"Invalid tag length, must be 16 or 32"
);
Aegis128L(State::new(key, nonce))
}
#[cfg(feature = "std")]
pub fn encrypt(mut self, m: &[u8], ad: &[u8]) -> (Vec<u8>, Tag<TAG_BYTES>) {
let state = &mut self.0;
let mlen = m.len();
let adlen = ad.len();
let mut c = Vec::with_capacity(mlen);
let mut src = [0u8; 32];
let mut dst = [0u8; 32];
state.absorb_ad(ad);
let mut i = 0;
while i + 32 <= mlen {
src.copy_from_slice(&m[i..][..32]);
state.enc(&mut dst, &src);
c.extend_from_slice(&dst);
i += 32;
}
if mlen % 32 != 0 {
src.fill(0);
src[..mlen % 32].copy_from_slice(&m[i..]);
state.enc(&mut dst, &src);
c.extend_from_slice(&dst[..mlen % 32]);
}
let tag = state.mac::<TAG_BYTES>(adlen as u64, mlen as u64);
(c, tag)
}
pub fn encrypt_in_place(mut self, mc: &mut [u8], ad: &[u8]) -> Tag<TAG_BYTES> {
let state = &mut self.0;
let mclen = mc.len();
let adlen = ad.len();
let mut src = [0u8; 32];
let mut dst = [0u8; 32];
state.absorb_ad(ad);
let mut i = 0;
while i + 32 <= mclen {
src.copy_from_slice(&mc[i..][..32]);
state.enc(&mut dst, &src);
mc[i..][..32].copy_from_slice(&dst);
i += 32;
}
if mclen % 32 != 0 {
src.fill(0);
src[..mclen % 32].copy_from_slice(&mc[i..]);
state.enc(&mut dst, &src);
mc[i..].copy_from_slice(&dst[..mclen % 32]);
}
state.mac::<TAG_BYTES>(adlen as u64, mclen as u64)
}
#[cfg(feature = "std")]
pub fn decrypt(mut self, c: &[u8], tag: &Tag<TAG_BYTES>, ad: &[u8]) -> Result<Vec<u8>, Error> {
let state = &mut self.0;
let clen = c.len();
let adlen = ad.len();
let mut m = Vec::with_capacity(clen);
let mut src = [0u8; 32];
let mut dst = [0u8; 32];
state.absorb_ad(ad);
let mut i = 0;
while i + 32 <= clen {
src.copy_from_slice(&c[i..][..32]);
state.dec(&mut dst, &src);
m.extend_from_slice(&dst);
i += 32;
}
if clen % 32 != 0 {
state.dec_partial(&mut dst, &c[i..]);
m.extend_from_slice(&dst[0..clen % 32]);
}
let tag2 = state.mac::<TAG_BYTES>(adlen as u64, clen as u64);
let mut acc = 0;
for (a, b) in tag.iter().zip(tag2.iter()) {
acc |= a ^ b;
}
if acc != 0 {
m.fill(0xaa);
return Err(Error::InvalidTag);
}
Ok(m)
}
pub fn decrypt_in_place(
mut self,
mc: &mut [u8],
tag: &Tag<TAG_BYTES>,
ad: &[u8],
) -> Result<(), Error> {
let state = &mut self.0;
let mclen = mc.len();
let adlen = ad.len();
let mut src = [0u8; 32];
let mut dst = [0u8; 32];
state.absorb_ad(ad);
let mut i = 0;
while i + 32 <= mclen {
src.copy_from_slice(&mc[i..][..32]);
state.dec(&mut dst, &src);
mc[i..][..32].copy_from_slice(&dst);
i += 32;
}
if mclen % 32 != 0 {
state.dec_partial(&mut dst, &mc[i..]);
mc[i..].copy_from_slice(&dst[0..mclen % 32]);
}
let tag2 = state.mac::<TAG_BYTES>(adlen as u64, mclen as u64);
let mut acc = 0;
for (a, b) in tag.iter().zip(tag2.iter()) {
acc |= a ^ b;
}
if acc != 0 {
mc.fill(0xaa);
return Err(Error::InvalidTag);
}
Ok(())
}
pub fn encryptor(&self, associated_data: &[u8]) -> Encryptor<TAG_BYTES> {
Encryptor {
inner: IncrementalState::new(&self.0, associated_data),
}
}
pub fn decryptor<'a>(
&self,
associated_data: &[u8],
plaintext: &'a mut [u8],
) -> Decryptor<'a, TAG_BYTES> {
Decryptor {
inner: IncrementalState::new(&self.0, associated_data),
plaintext: Quarantine::new(plaintext),
}
}
}
impl<const TAG_BYTES: usize> fmt::Debug for Aegis128L<TAG_BYTES> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Aegis128L").finish_non_exhaustive()
}
}
struct IncrementalState {
state: State,
buf: [u8; 32],
pos: usize,
adlen: u64,
mlen: MessageLength,
}
impl IncrementalState {
fn new(cipher_state: &State, ad: &[u8]) -> Self {
check_ad_length(ad);
let mut state = *cipher_state;
state.absorb_ad(ad);
IncrementalState {
state,
buf: [0u8; 32],
pos: 0,
adlen: ad.len() as u64,
mlen: MessageLength::new(),
}
}
fn transform<const DECRYPT: bool>(&mut self, mc: &mut [u8]) {
let mut offset = 0;
if self.pos != 0 {
let n = mc.len().min(32 - self.pos);
for j in 0..n {
let input = mc[j];
let output = input ^ self.buf[self.pos + j];
self.buf[self.pos + j] = if DECRYPT { output } else { input };
mc[j] = output;
}
self.pos += n;
offset = n;
if self.pos < 32 {
return;
}
let buf = self.buf;
self.state.absorb(&buf);
self.pos = 0;
}
let mut src = [0u8; 32];
let mut dst = [0u8; 32];
while offset + 32 <= mc.len() {
src.copy_from_slice(&mc[offset..][..32]);
if DECRYPT {
self.state.dec(&mut dst, &src);
} else {
self.state.enc(&mut dst, &src);
}
mc[offset..][..32].copy_from_slice(&dst);
offset += 32;
}
let left = mc.len() - offset;
if left != 0 {
self.state.squeeze_keystream(&mut self.buf);
for j in 0..left {
let input = mc[offset + j];
let output = input ^ self.buf[j];
self.buf[j] = if DECRYPT { output } else { input };
mc[offset + j] = output;
}
self.pos = left;
}
}
fn tag<const TAG_BYTES: usize>(&mut self) -> Tag<TAG_BYTES> {
if self.pos != 0 {
let mut tmp = [0u8; 32];
tmp[..self.pos].copy_from_slice(&self.buf[..self.pos]);
self.state.absorb(&tmp);
}
self.state.mac::<TAG_BYTES>(self.adlen, self.mlen.get())
}
}
pub struct Encryptor<const TAG_BYTES: usize> {
inner: IncrementalState,
}
impl<const TAG_BYTES: usize> Encryptor<TAG_BYTES> {
pub fn update(&mut self, plaintext: &[u8], ciphertext: &mut [u8]) {
assert_eq!(
plaintext.len(),
ciphertext.len(),
"plaintext and ciphertext chunks must have the same length"
);
ciphertext.copy_from_slice(plaintext);
self.update_in_place(ciphertext);
}
pub fn update_in_place(&mut self, buffer: &mut [u8]) {
self.inner.mlen.add(buffer.len());
self.inner.transform::<false>(buffer);
}
pub fn finalize(mut self) -> Tag<TAG_BYTES> {
self.inner.tag::<TAG_BYTES>()
}
#[cfg(test)]
pub(crate) fn set_consumed_length_for_tests(&mut self, mlen: u64) {
self.inner.mlen.set_for_tests(mlen);
}
}
impl<const TAG_BYTES: usize> fmt::Debug for Encryptor<TAG_BYTES> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("aegis128l::Encryptor")
.finish_non_exhaustive()
}
}
impl<const TAG_BYTES: usize> Drop for Encryptor<TAG_BYTES> {
fn drop(&mut self) {
wipe_value(&mut self.inner);
}
}
pub struct Decryptor<'a, const TAG_BYTES: usize> {
inner: IncrementalState,
plaintext: Quarantine<'a>,
}
impl<'a, const TAG_BYTES: usize> Decryptor<'a, TAG_BYTES> {
pub fn update(&mut self, ciphertext: &[u8]) -> Result<(), Error> {
self.plaintext.fits(ciphertext.len())?;
self.inner.mlen.try_add(ciphertext.len())?;
let plaintext = self.plaintext.next_chunk(ciphertext.len());
plaintext.copy_from_slice(ciphertext);
self.inner.transform::<true>(plaintext);
Ok(())
}
pub fn finalize(mut self, tag: &Tag<TAG_BYTES>) -> Result<&'a mut [u8], Error> {
let computed = self.inner.tag::<TAG_BYTES>();
if !tags_match(tag, &computed) {
return Err(Error::InvalidTag);
}
Ok(self.plaintext.release())
}
#[cfg(test)]
pub(crate) fn set_consumed_length_for_tests(&mut self, mlen: u64) {
self.inner.mlen.set_for_tests(mlen);
}
}
impl<const TAG_BYTES: usize> fmt::Debug for Decryptor<'_, TAG_BYTES> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("aegis128l::Decryptor")
.finish_non_exhaustive()
}
}
impl<const TAG_BYTES: usize> Drop for Decryptor<'_, TAG_BYTES> {
fn drop(&mut self) {
wipe_value(&mut self.inner);
}
}
#[derive(Clone)]
pub struct Aegis128LMac<const TAG_BYTES: usize> {
state: State,
buf: [u8; 32],
buf_len: usize,
msg_len: usize,
}
impl<const TAG_BYTES: usize> fmt::Debug for Aegis128LMac<TAG_BYTES> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Aegis128LMac").finish_non_exhaustive()
}
}
impl<const TAG_BYTES: usize> Aegis128LMac<TAG_BYTES> {
pub fn new(key: &Key) -> Self {
let nonce = [0u8; 16];
Self::new_with_nonce(key, &nonce)
}
pub fn new_with_nonce(key: &Key, nonce: &Nonce) -> Self {
assert!(
TAG_BYTES == 16 || TAG_BYTES == 32,
"Invalid tag length, must be 16 or 32"
);
Aegis128LMac {
state: State::new(key, nonce),
buf: [0u8; 32],
buf_len: 0,
msg_len: 0,
}
}
pub fn update(&mut self, data: &[u8]) {
self.msg_len += data.len();
let mut offset = 0;
if self.buf_len > 0 {
let needed = 32 - self.buf_len;
if data.len() < needed {
self.buf[self.buf_len..self.buf_len + data.len()].copy_from_slice(data);
self.buf_len += data.len();
return;
}
self.buf[self.buf_len..].copy_from_slice(&data[..needed]);
self.state.absorb(&self.buf);
self.buf_len = 0;
offset = needed;
}
while offset + 32 <= data.len() {
let mut block = [0u8; 32];
block.copy_from_slice(&data[offset..offset + 32]);
self.state.absorb(&block);
offset += 32;
}
if offset < data.len() {
self.buf_len = data.len() - offset;
self.buf[..self.buf_len].copy_from_slice(&data[offset..]);
}
}
pub fn finalize(mut self) -> Tag<TAG_BYTES> {
if self.buf_len > 0 || self.msg_len == 0 {
self.buf[self.buf_len..].fill(0);
self.state.absorb(&self.buf);
}
self.state.mac_finalize::<TAG_BYTES>(self.msg_len)
}
pub fn verify(self, expected: &Tag<TAG_BYTES>) -> Result<(), Error> {
let tag = self.finalize();
let mut acc = 0u8;
for (a, b) in tag.iter().zip(expected.iter()) {
acc |= a ^ b;
}
if acc != 0 {
return Err(Error::InvalidTag);
}
Ok(())
}
}