rustls-ccm 0.2.0

CCM and CCM-8 cipher suites for rustls (TLS 1.2 and TLS 1.3)
Documentation
use std::io::{Read, Write};
use std::sync::Arc;

use rustls::pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer};
use rustls::{
    CipherSuite, ClientConfig, ClientConnection, RootCertStore, ServerConfig, ServerConnection,
};

fn ecdsa_cert_and_key() -> (CertificateDer<'static>, PrivateKeyDer<'static>) {
    let kp = rcgen::KeyPair::generate_for(&rcgen::PKCS_ECDSA_P256_SHA256).unwrap();
    let params = rcgen::CertificateParams::new(vec!["localhost".into()]).unwrap();
    let cert = params.self_signed(&kp).unwrap();
    let cert_der = CertificateDer::from(cert.der().to_vec());
    let key_der = PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(kp.serialize_der()));
    (cert_der, key_der)
}

struct TestPair {
    client: ClientConnection,
    server: ServerConnection,
}

impl TestPair {
    fn new(suites: Vec<rustls::SupportedCipherSuite>, tls12_only: bool) -> Self {
        let (cert, key) = ecdsa_cert_and_key();
        let mut roots = RootCertStore::empty();
        roots.add(cert.clone()).unwrap();

        let mut provider = rustls::crypto::aws_lc_rs::default_provider();
        provider.cipher_suites = suites;
        let provider = Arc::new(provider);

        let versions: Vec<&'static rustls::SupportedProtocolVersion> = if tls12_only {
            vec![&rustls::version::TLS12]
        } else {
            rustls::DEFAULT_VERSIONS.to_vec()
        };

        let server_cfg = ServerConfig::builder_with_provider(provider.clone())
            .with_protocol_versions(&versions)
            .unwrap()
            .with_no_client_auth()
            .with_single_cert(vec![cert], key)
            .unwrap();

        let client_cfg = ClientConfig::builder_with_provider(provider)
            .with_protocol_versions(&versions)
            .unwrap()
            .with_root_certificates(roots)
            .with_no_client_auth();

        let client =
            ClientConnection::new(Arc::new(client_cfg), "localhost".try_into().unwrap()).unwrap();
        let server = ServerConnection::new(Arc::new(server_cfg)).unwrap();

        Self { client, server }
    }

    fn handshake(&mut self) {
        let mut buf = Vec::new();
        loop {
            let mut progress = false;

            buf.clear();
            if self.client.write_tls(&mut buf).unwrap() > 0 {
                self.server.read_tls(&mut &buf[..]).unwrap();
                self.server.process_new_packets().unwrap();
                progress = true;
            }

            buf.clear();
            if self.server.write_tls(&mut buf).unwrap() > 0 {
                self.client.read_tls(&mut &buf[..]).unwrap();
                self.client.process_new_packets().unwrap();
                progress = true;
            }

            if !progress {
                break;
            }
        }
    }

    fn send_to_server(&mut self, msg: &[u8]) {
        self.client.writer().write_all(msg).unwrap();
        let mut buf = Vec::new();
        self.client.write_tls(&mut buf).unwrap();
        let mut cursor = &buf[..];
        while !cursor.is_empty() {
            self.server.read_tls(&mut cursor).unwrap();
            self.server.process_new_packets().unwrap();
        }
    }

    fn send_to_client(&mut self, msg: &[u8]) {
        self.server.writer().write_all(msg).unwrap();
        let mut buf = Vec::new();
        self.server.write_tls(&mut buf).unwrap();
        let mut cursor = &buf[..];
        while !cursor.is_empty() {
            self.client.read_tls(&mut cursor).unwrap();
            self.client.process_new_packets().unwrap();
        }
    }

    fn round_trip(&mut self, msg: &[u8]) {
        self.send_to_server(msg);
        let mut received = vec![0u8; msg.len()];
        self.server.reader().read_exact(&mut received).unwrap();
        assert_eq!(&received, msg);

        self.send_to_client(msg);
        let mut echoed = vec![0u8; msg.len()];
        self.client.reader().read_exact(&mut echoed).unwrap();
        assert_eq!(&echoed, msg);
    }

    fn negotiated_suite(&self) -> CipherSuite {
        self.client.negotiated_cipher_suite().unwrap().suite()
    }
}

// -- TLS 1.3 -----------------------------------------------------------------

#[test]
fn tls13_ccm() {
    let mut pair = TestPair::new(vec![*rustls_ccm::TLS13_AES_128_CCM_SHA256], false);
    pair.handshake();
    assert_eq!(
        pair.negotiated_suite(),
        CipherSuite::TLS13_AES_128_CCM_SHA256
    );
    pair.round_trip(b"hello ccm");
}

#[test]
fn tls13_ccm8() {
    let mut pair = TestPair::new(vec![*rustls_ccm::TLS13_AES_128_CCM_8_SHA256], false);
    pair.handshake();
    assert_eq!(
        pair.negotiated_suite(),
        CipherSuite::TLS13_AES_128_CCM_8_SHA256
    );
    pair.round_trip(b"hello ccm-8");
}

// -- TLS 1.2 -----------------------------------------------------------------

#[test]
fn tls12_ecdhe_ecdsa_aes128_ccm() {
    let mut pair = TestPair::new(vec![*rustls_ccm::TLS_ECDHE_ECDSA_WITH_AES_128_CCM], true);
    pair.handshake();
    assert_eq!(
        pair.negotiated_suite(),
        CipherSuite::TLS_ECDHE_ECDSA_WITH_AES_128_CCM
    );
    pair.round_trip(b"hello 128-ccm");
}

#[test]
fn tls12_ecdhe_ecdsa_aes256_ccm() {
    let mut pair = TestPair::new(vec![*rustls_ccm::TLS_ECDHE_ECDSA_WITH_AES_256_CCM], true);
    pair.handshake();
    assert_eq!(
        pair.negotiated_suite(),
        CipherSuite::TLS_ECDHE_ECDSA_WITH_AES_256_CCM
    );
    pair.round_trip(b"hello 256-ccm");
}

#[test]
fn tls12_ecdhe_ecdsa_aes128_ccm8() {
    let mut pair = TestPair::new(vec![*rustls_ccm::TLS_ECDHE_ECDSA_WITH_AES_128_CCM_8], true);
    pair.handshake();
    assert_eq!(
        pair.negotiated_suite(),
        CipherSuite::TLS_ECDHE_ECDSA_WITH_AES_128_CCM_8
    );
    pair.round_trip(b"hello 128-ccm-8");
}

#[test]
fn tls12_ecdhe_ecdsa_aes256_ccm8() {
    let mut pair = TestPair::new(vec![*rustls_ccm::TLS_ECDHE_ECDSA_WITH_AES_256_CCM_8], true);
    pair.handshake();
    assert_eq!(
        pair.negotiated_suite(),
        CipherSuite::TLS_ECDHE_ECDSA_WITH_AES_256_CCM_8
    );
    pair.round_trip(b"hello 256-ccm-8");
}

// -- Misc --------------------------------------------------------------------

#[test]
fn large_payload() {
    let mut pair = TestPair::new(vec![*rustls_ccm::TLS13_AES_128_CCM_SHA256], false);
    pair.handshake();
    pair.round_trip(&vec![0xABu8; 4096]);
}

#[test]
fn multiple_messages() {
    let mut pair = TestPair::new(vec![*rustls_ccm::TLS_ECDHE_ECDSA_WITH_AES_128_CCM_8], true);
    pair.handshake();
    for i in 0..50 {
        pair.round_trip(format!("message {i}").as_bytes());
    }
}

#[test]
fn full_provider_negotiates_ccm() {
    let ccm_suites: Vec<_> = rustls_ccm::all_suites().iter().map(|s| s.suite()).collect();
    let mut all = rustls_ccm::all_suites().to_vec();
    all.extend(rustls::crypto::aws_lc_rs::default_provider().cipher_suites);
    let mut pair = TestPair::new(all, false);
    pair.handshake();
    assert!(
        ccm_suites.contains(&pair.negotiated_suite()),
        "expected a CCM suite, got {:?}",
        pair.negotiated_suite()
    );
    pair.round_trip(b"negotiated something");
}