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()
}
}
#[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");
}
#[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");
}
#[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");
}