rustls-ccm 0.2.0

CCM and CCM-8 cipher suites for rustls (TLS 1.2 and TLS 1.3)
Documentation
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);

        // TLS 1.3 inner plaintext: [payload][content_type:1]
        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 {
        // plaintext + content_type byte + tag
        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;
        // Must have at least tag + 1 byte (content type)
        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)?;

        // Remove tag, leaving [plaintext][content_type][padding]
        payload.truncate(message_len);

        msg.into_tls13_unpadded_message()
    }
}