use std::marker::PhantomData;
use ccm::aead::array::Array;
use ccm::aead::{AeadInOut, KeyInit, Tag};
use rustls::crypto::cipher::{
AeadKey, InboundOpaqueMessage, InboundPlainMessage, Iv, MessageDecrypter, MessageEncrypter,
Nonce, OutboundOpaqueMessage, OutboundPlainMessage, PrefixedPayload, Tls13AeadAlgorithm,
UnsupportedOperationError, make_tls13_aad,
};
use rustls::{ConnectionTrafficSecrets, ContentType, Error, ProtocolVersion};
use crate::CcmVariant;
pub(crate) struct Tls13CcmAead<V: CcmVariant>(PhantomData<V>);
impl<V: CcmVariant> Tls13CcmAead<V> {
pub(crate) const NEW: Self = Self(PhantomData);
}
impl<V: CcmVariant> Tls13AeadAlgorithm for Tls13CcmAead<V> {
fn encrypter(&self, key: AeadKey, iv: Iv) -> Box<dyn MessageEncrypter> {
let cipher = V::Cipher::new_from_slice(key.as_ref()).unwrap();
Box::new(Tls13CcmEncrypter::<V> {
cipher,
iv,
_v: PhantomData,
})
}
fn decrypter(&self, key: AeadKey, iv: Iv) -> Box<dyn MessageDecrypter> {
let cipher = V::Cipher::new_from_slice(key.as_ref()).unwrap();
Box::new(Tls13CcmDecrypter::<V> {
cipher,
iv,
_v: PhantomData,
})
}
fn key_len(&self) -> usize {
V::KEY_LEN
}
fn extract_keys(
&self,
_key: AeadKey,
_iv: Iv,
) -> Result<ConnectionTrafficSecrets, UnsupportedOperationError> {
Err(UnsupportedOperationError)
}
}
struct Tls13CcmEncrypter<V: CcmVariant> {
cipher: V::Cipher,
iv: Iv,
_v: PhantomData<V>,
}
impl<V: CcmVariant> MessageEncrypter for Tls13CcmEncrypter<V> {
fn encrypt(
&mut self,
msg: OutboundPlainMessage<'_>,
seq: u64,
) -> Result<OutboundOpaqueMessage, Error> {
let total_len = self.encrypted_payload_len(msg.payload.len());
let mut payload = PrefixedPayload::with_capacity(total_len);
payload.extend_from_chunks(&msg.payload);
payload.extend_from_slice(&[msg.typ.into()]);
let nonce = Nonce::new(&self.iv, seq);
let aad = make_tls13_aad(total_len);
let ccm_nonce = Array::from(nonce.0);
let tag = self
.cipher
.encrypt_inout_detached(&ccm_nonce, &aad, payload.as_mut().into())
.map_err(|_| Error::EncryptError)?;
payload.extend_from_slice(tag.as_slice());
Ok(OutboundOpaqueMessage::new(
ContentType::ApplicationData,
ProtocolVersion::TLSv1_2,
payload,
))
}
fn encrypted_payload_len(&self, payload_len: usize) -> usize {
payload_len + 1 + V::TAG_LEN
}
}
struct Tls13CcmDecrypter<V: CcmVariant> {
cipher: V::Cipher,
iv: Iv,
_v: PhantomData<V>,
}
impl<V: CcmVariant> MessageDecrypter for Tls13CcmDecrypter<V> {
fn decrypt<'a>(
&mut self,
mut msg: InboundOpaqueMessage<'a>,
seq: u64,
) -> Result<InboundPlainMessage<'a>, Error> {
let payload = &msg.payload;
if payload.len() < V::TAG_LEN + 1 {
return Err(Error::DecryptError);
}
let nonce = Nonce::new(&self.iv, seq);
let aad = make_tls13_aad(payload.len());
let message_len = payload.len() - V::TAG_LEN;
let tag =
Tag::<V::Cipher>::try_from(&payload[message_len..]).map_err(|_| Error::DecryptError)?;
let payload = &mut msg.payload;
let ccm_nonce = Array::from(nonce.0);
self.cipher
.decrypt_inout_detached(&ccm_nonce, &aad, (&mut payload[..message_len]).into(), &tag)
.map_err(|_| Error::DecryptError)?;
payload.truncate(message_len);
msg.into_tls13_unpadded_message()
}
}