use rustls::quic::Version;
use rustls::Side;
fn get_quic_keys(client_dcid: &[u8]) -> rustls::quic::Keys {
let suite = rustls::crypto::ring::cipher_suite::TLS13_AES_128_GCM_SHA256;
suite
.tls13()
.unwrap()
.quic_suite()
.unwrap()
.keys(client_dcid, Side::Server, Version::V1)
}
#[test]
fn test_rustls_initial_keys_match_aioquic() {
let client_dcid = [0x0e, 0x95, 0x59, 0x96, 0x7f, 0xd8, 0xf3, 0xf3];
let keys = get_quic_keys(&client_dcid);
let remote_packet = keys.remote.packet.as_ref();
let remote_header = keys.remote.header.as_ref();
println!("remote packet tag_len: {}", remote_packet.tag_len());
println!("remote header sample_len: {}", remote_header.sample_len());
let plaintext = b"hello world test";
let aad = b"some header aad test";
let pn: u64 = 42;
let mut buf = plaintext.to_vec();
let tag = remote_packet
.encrypt_in_place(pn, aad, &mut buf)
.expect("encrypt");
buf.extend_from_slice(tag.as_ref());
let decrypted = remote_packet
.decrypt_in_place(pn, aad, &mut buf)
.expect("decrypt");
assert_eq!(decrypted, plaintext.as_slice());
println!("Roundtrip encrypt/decrypt OK!");
}
#[test]
fn test_aead_with_real_packet_parameters() {
let client_dcid = [0x0e, 0x95, 0x59, 0x96, 0x7f, 0xd8, 0xf3, 0xf3];
let keys = get_quic_keys(&client_dcid);
let aad = [0xc7, 0x00, 0x00, 0x00, 0x01, 0x08, 0x0e, 0x95, 0x59, 0x96, 0x7f, 0xd8, 0xf3, 0xf3,
0x08, 0x48, 0x21, 0xbe, 0x84, 0x26, 0x0a, 0xd0, 0x0b, 0x00, 0x41, 0xf2, 0x78, 0xaa, 0xdd];
let pn: u64 = 7908061;
let plaintext = b"test plaintext content for verification";
let mut buf = plaintext.to_vec();
let tag = keys.remote.packet.encrypt_in_place(pn, &aad, &mut buf).expect("encrypt");
buf.extend_from_slice(tag.as_ref());
let decrypted = keys.remote.packet.decrypt_in_place(pn, &aad, &mut buf).expect("decrypt");
assert_eq!(decrypted, plaintext.as_slice());
println!("Encrypt/decrypt roundtrip with pn={} OK", pn);
let mut buf2 = plaintext.to_vec();
let tag2 = keys.remote.packet.encrypt_in_place(pn, &aad, &mut buf2).expect("encrypt");
buf2.extend_from_slice(tag2.as_ref());
let bad_pn = pn + 1;
let result = keys.remote.packet.decrypt_in_place(bad_pn, &aad, &mut buf2);
assert!(result.is_err(), "Wrong pn should fail");
println!("Wrong pn correctly fails");
let mut buf3 = plaintext.to_vec();
let tag3 = keys.remote.packet.encrypt_in_place(pn, &aad, &mut buf3).expect("encrypt");
buf3.extend_from_slice(tag3.as_ref());
let bad_aad = [0u8; 29];
let result = keys.remote.packet.decrypt_in_place(pn, &bad_aad, &mut buf3);
assert!(result.is_err(), "Wrong AAD should fail");
println!("Wrong AAD correctly fails");
}