use alloc::boxed::Box;
use ic_core::traits::Aead as _;
use rustls::crypto::cipher::{
make_tls12_aad, make_tls13_aad, AeadKey, InboundOpaqueMessage, InboundPlainMessage, Iv,
KeyBlockShape, MessageDecrypter, MessageEncrypter, Nonce, OutboundOpaqueMessage,
OutboundPlainMessage, PrefixedPayload, Tls12AeadAlgorithm, Tls13AeadAlgorithm,
UnsupportedOperationError,
};
use rustls::{ConnectionTrafficSecrets, ContentType, Error, ProtocolVersion};
const TAG_LEN: usize = 16;
const TLS12_EXPLICIT_NONCE_LEN: usize = 8;
const TLS12_FIXED_IV_LEN: usize = 4;
const TLS12_CHACHA_IV_LEN: usize = 12;
pub(crate) static TLS13_AES_128_GCM: Tls13Aead = Tls13Aead {
suite: Suite::Aes128,
};
pub(crate) static TLS13_AES_256_GCM: Tls13Aead = Tls13Aead {
suite: Suite::Aes256,
};
pub(crate) static TLS13_CHACHA20_POLY1305: Tls13Aead = Tls13Aead {
suite: Suite::ChaCha20,
};
pub(crate) static TLS12_AES_128_GCM: Tls12Aead = Tls12Aead {
suite: Suite::Aes128,
};
pub(crate) static TLS12_AES_256_GCM: Tls12Aead = Tls12Aead {
suite: Suite::Aes256,
};
pub(crate) static TLS12_CHACHA20_POLY1305: Tls12Aead = Tls12Aead {
suite: Suite::ChaCha20,
};
#[derive(Debug, Clone, Copy, PartialEq)]
pub(crate) enum Suite {
Aes128,
Aes256,
ChaCha20,
}
impl Suite {
pub(crate) fn key_len(self) -> usize {
match self {
Self::Aes128 => 16,
Self::Aes256 | Self::ChaCha20 => 32,
}
}
pub(crate) fn cipher(self, key: &[u8]) -> Result<Cipher, Error> {
if key.len() != self.key_len() {
return Err(Error::General(alloc::format!(
"{:?} takes a {}-byte key, not {}",
self,
self.key_len(),
key.len()
)));
}
let bad = |_| Error::General(alloc::format!("{self:?} rejected a key rustls supplied"));
Ok(match self {
Self::Aes128 => Cipher::Aes128(ic_cipher::Aes128Gcm::new(key).map_err(bad)?),
Self::Aes256 => Cipher::Aes256(ic_cipher::Aes256Gcm::new(key).map_err(bad)?),
Self::ChaCha20 => Cipher::ChaCha20(ic_cipher::ChaCha20Poly1305::new(key).map_err(bad)?),
})
}
fn traffic_secrets(self, key: AeadKey, iv: Iv) -> ConnectionTrafficSecrets {
match self {
Self::Aes128 => ConnectionTrafficSecrets::Aes128Gcm { key, iv },
Self::Aes256 => ConnectionTrafficSecrets::Aes256Gcm { key, iv },
Self::ChaCha20 => ConnectionTrafficSecrets::Chacha20Poly1305 { key, iv },
}
}
}
pub(crate) enum Cipher {
Aes128(ic_cipher::Aes128Gcm),
Aes256(ic_cipher::Aes256Gcm),
ChaCha20(ic_cipher::ChaCha20Poly1305),
}
impl Cipher {
pub(crate) fn seal(
&self,
nonce: &[u8],
aad: &[u8],
in_out: &mut [u8],
tag: &mut [u8],
) -> Result<(), ()> {
match self {
Self::Aes128(c) => c.seal_detached(nonce, aad, in_out, tag),
Self::Aes256(c) => c.seal_detached(nonce, aad, in_out, tag),
Self::ChaCha20(c) => c.seal_detached(nonce, aad, in_out, tag),
}
.map_err(|_| ())
}
pub(crate) fn open(
&self,
nonce: &[u8],
aad: &[u8],
in_out: &mut [u8],
tag: &[u8],
) -> Result<(), ()> {
match self {
Self::Aes128(c) => c.open_detached(nonce, aad, in_out, tag),
Self::Aes256(c) => c.open_detached(nonce, aad, in_out, tag),
Self::ChaCha20(c) => c.open_detached(nonce, aad, in_out, tag),
}
.map_err(|_| ())
}
}
pub(crate) struct Tls13Aead {
suite: Suite,
}
impl Tls13AeadAlgorithm for Tls13Aead {
fn encrypter(&self, key: AeadKey, iv: Iv) -> Box<dyn MessageEncrypter> {
Box::new(Tls13Encrypter {
cipher: self
.suite
.cipher(key.as_ref())
.expect("rustls supplied a key this suite does not use"),
iv,
})
}
fn decrypter(&self, key: AeadKey, iv: Iv) -> Box<dyn MessageDecrypter> {
Box::new(Tls13Decrypter {
cipher: self
.suite
.cipher(key.as_ref())
.expect("rustls supplied a key this suite does not use"),
iv,
})
}
fn key_len(&self) -> usize {
self.suite.key_len()
}
fn extract_keys(
&self,
key: AeadKey,
iv: Iv,
) -> Result<ConnectionTrafficSecrets, UnsupportedOperationError> {
Ok(self.suite.traffic_secrets(key, iv))
}
fn fips(&self) -> bool {
false
}
}
struct Tls13Encrypter {
cipher: Cipher,
iv: Iv,
}
impl MessageEncrypter for Tls13Encrypter {
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.to_array());
let nonce = Nonce::new(&self.iv, seq).0;
let aad = make_tls13_aad(total_len);
let mut tag = [0u8; TAG_LEN];
let body = payload.as_mut();
self.cipher
.seal(&nonce, &aad, body, &mut tag)
.map_err(|_| Error::EncryptError)?;
payload.extend_from_slice(&tag);
Ok(OutboundOpaqueMessage::new(
ContentType::ApplicationData,
ProtocolVersion::TLSv1_2,
payload,
))
}
fn encrypted_payload_len(&self, payload_len: usize) -> usize {
payload_len + 1 + TAG_LEN
}
}
struct Tls13Decrypter {
cipher: Cipher,
iv: Iv,
}
impl MessageDecrypter for Tls13Decrypter {
fn decrypt<'a>(
&mut self,
mut msg: InboundOpaqueMessage<'a>,
seq: u64,
) -> Result<InboundPlainMessage<'a>, Error> {
let payload = &mut msg.payload;
if payload.len() < TAG_LEN {
return Err(Error::DecryptError);
}
let nonce = Nonce::new(&self.iv, seq).0;
let aad = make_tls13_aad(payload.len());
let cipher_len = payload.len() - TAG_LEN;
let (body, tag) = payload.split_at_mut(cipher_len);
let tag: [u8; TAG_LEN] = tag.try_into().expect("split at exactly the tag length");
self.cipher
.open(&nonce, &aad, body, &tag)
.map_err(|_| Error::DecryptError)?;
payload.truncate(cipher_len);
msg.into_tls13_unpadded_message()
}
}
pub(crate) struct Tls12Aead {
suite: Suite,
}
impl Tls12AeadAlgorithm for Tls12Aead {
fn encrypter(&self, key: AeadKey, iv: &[u8], extra: &[u8]) -> Box<dyn MessageEncrypter> {
let cipher = self
.suite
.cipher(key.as_ref())
.expect("rustls supplied a key this suite does not use");
if self.suite == Suite::ChaCha20 {
let mut fixed = [0u8; TLS12_CHACHA_IV_LEN];
fixed.copy_from_slice(iv);
return Box::new(Tls12ChaChaEncrypter {
cipher,
iv: Iv::new(fixed),
});
}
let mut nonce = [0u8; 12];
nonce[..TLS12_FIXED_IV_LEN].copy_from_slice(iv);
nonce[TLS12_FIXED_IV_LEN..].copy_from_slice(extra);
Box::new(Tls12Encrypter { cipher, nonce })
}
fn decrypter(&self, key: AeadKey, iv: &[u8]) -> Box<dyn MessageDecrypter> {
let cipher = self
.suite
.cipher(key.as_ref())
.expect("rustls supplied a key this suite does not use");
if self.suite == Suite::ChaCha20 {
let mut fixed = [0u8; TLS12_CHACHA_IV_LEN];
fixed.copy_from_slice(iv);
return Box::new(Tls12ChaChaDecrypter {
cipher,
iv: Iv::new(fixed),
});
}
let mut fixed = [0u8; TLS12_FIXED_IV_LEN];
fixed.copy_from_slice(iv);
Box::new(Tls12Decrypter { cipher, fixed })
}
fn key_block_shape(&self) -> KeyBlockShape {
if self.suite == Suite::ChaCha20 {
return KeyBlockShape {
enc_key_len: self.suite.key_len(),
fixed_iv_len: TLS12_CHACHA_IV_LEN,
explicit_nonce_len: 0,
};
}
KeyBlockShape {
enc_key_len: self.suite.key_len(),
fixed_iv_len: TLS12_FIXED_IV_LEN,
explicit_nonce_len: TLS12_EXPLICIT_NONCE_LEN,
}
}
fn extract_keys(
&self,
key: AeadKey,
iv: &[u8],
explicit: &[u8],
) -> Result<ConnectionTrafficSecrets, UnsupportedOperationError> {
let mut nonce = [0u8; 12];
if iv.len() + explicit.len() != nonce.len() {
return Err(UnsupportedOperationError);
}
nonce[..iv.len()].copy_from_slice(iv);
nonce[iv.len()..].copy_from_slice(explicit);
Ok(self.suite.traffic_secrets(key, Iv::new(nonce)))
}
fn fips(&self) -> bool {
false
}
}
struct Tls12Encrypter {
cipher: Cipher,
nonce: [u8; 12],
}
impl MessageEncrypter for Tls12Encrypter {
fn encrypt(
&mut self,
msg: OutboundPlainMessage<'_>,
seq: u64,
) -> Result<OutboundOpaqueMessage, Error> {
let mut nonce = self.nonce;
for (n, s) in nonce[TLS12_FIXED_IV_LEN..]
.iter_mut()
.zip(seq.to_be_bytes())
{
*n ^= s;
}
let explicit = &nonce[TLS12_FIXED_IV_LEN..];
let total_len = self.encrypted_payload_len(msg.payload.len());
let mut payload = PrefixedPayload::with_capacity(total_len);
payload.extend_from_slice(explicit);
payload.extend_from_chunks(&msg.payload);
let aad = make_tls12_aad(seq, msg.typ, msg.version, msg.payload.len());
let mut tag = [0u8; TAG_LEN];
let body = &mut payload.as_mut()[TLS12_EXPLICIT_NONCE_LEN..];
self.cipher
.seal(&nonce, &aad, body, &mut tag)
.map_err(|_| Error::EncryptError)?;
payload.extend_from_slice(&tag);
Ok(OutboundOpaqueMessage::new(msg.typ, msg.version, payload))
}
fn encrypted_payload_len(&self, payload_len: usize) -> usize {
TLS12_EXPLICIT_NONCE_LEN + payload_len + TAG_LEN
}
}
struct Tls12Decrypter {
cipher: Cipher,
fixed: [u8; TLS12_FIXED_IV_LEN],
}
impl MessageDecrypter for Tls12Decrypter {
fn decrypt<'a>(
&mut self,
mut msg: InboundOpaqueMessage<'a>,
seq: u64,
) -> Result<InboundPlainMessage<'a>, Error> {
let payload = &mut msg.payload;
if payload.len() < TLS12_EXPLICIT_NONCE_LEN + TAG_LEN {
return Err(Error::DecryptError);
}
let mut nonce = [0u8; 12];
nonce[..TLS12_FIXED_IV_LEN].copy_from_slice(&self.fixed);
nonce[TLS12_FIXED_IV_LEN..].copy_from_slice(&payload[..TLS12_EXPLICIT_NONCE_LEN]);
let plain_len = payload.len() - TLS12_EXPLICIT_NONCE_LEN - TAG_LEN;
let aad = make_tls12_aad(seq, msg.typ, msg.version, plain_len);
let tag_at = TLS12_EXPLICIT_NONCE_LEN + plain_len;
let mut tag = [0u8; TAG_LEN];
tag.copy_from_slice(&payload[tag_at..]);
let body = &mut payload[TLS12_EXPLICIT_NONCE_LEN..tag_at];
self.cipher
.open(&nonce, &aad, body, &tag)
.map_err(|_| Error::DecryptError)?;
payload.copy_within(TLS12_EXPLICIT_NONCE_LEN.., 0);
payload.truncate(plain_len);
Ok(msg.into_plain_message())
}
}
struct Tls12ChaChaEncrypter {
cipher: Cipher,
iv: Iv,
}
impl MessageEncrypter for Tls12ChaChaEncrypter {
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);
let nonce = Nonce::new(&self.iv, seq).0;
let aad = make_tls12_aad(seq, msg.typ, msg.version, msg.payload.len());
let mut tag = [0u8; TAG_LEN];
self.cipher
.seal(&nonce, &aad, payload.as_mut(), &mut tag)
.map_err(|_| Error::EncryptError)?;
payload.extend_from_slice(&tag);
Ok(OutboundOpaqueMessage::new(msg.typ, msg.version, payload))
}
fn encrypted_payload_len(&self, payload_len: usize) -> usize {
payload_len + TAG_LEN
}
}
struct Tls12ChaChaDecrypter {
cipher: Cipher,
iv: Iv,
}
impl MessageDecrypter for Tls12ChaChaDecrypter {
fn decrypt<'a>(
&mut self,
mut msg: InboundOpaqueMessage<'a>,
seq: u64,
) -> Result<InboundPlainMessage<'a>, Error> {
let payload = &mut msg.payload;
if payload.len() < TAG_LEN {
return Err(Error::DecryptError);
}
let plain_len = payload.len() - TAG_LEN;
let nonce = Nonce::new(&self.iv, seq).0;
let aad = make_tls12_aad(seq, msg.typ, msg.version, plain_len);
let (body, tag) = payload.split_at_mut(plain_len);
let tag: [u8; TAG_LEN] = tag.try_into().expect("split at exactly the tag length");
self.cipher
.open(&nonce, &aad, body, &tag)
.map_err(|_| Error::DecryptError)?;
payload.truncate(plain_len);
Ok(msg.into_plain_message())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn each_suite_binds_the_key_length_it_advertises() {
for suite in [Suite::Aes128, Suite::Aes256, Suite::ChaCha20] {
let key = alloc::vec![0x42u8; suite.key_len()];
assert!(suite.cipher(&key).is_ok(), "{suite:?} rejected its own key");
for wrong in [suite.key_len() - 1, suite.key_len() + 1] {
assert!(
suite.cipher(&alloc::vec![0u8; wrong]).is_err(),
"{suite:?} accepted a {wrong}-byte key"
);
}
}
assert_eq!(Suite::Aes256.key_len(), Suite::ChaCha20.key_len());
}
#[test]
fn the_tls12_key_block_shape_matches_the_framing() {
let gcm = Tls12Aead {
suite: Suite::Aes128,
}
.key_block_shape();
assert_eq!(gcm.fixed_iv_len, 4);
assert_eq!(gcm.explicit_nonce_len, 8);
assert_eq!(gcm.enc_key_len, 16);
let chacha = Tls12Aead {
suite: Suite::ChaCha20,
}
.key_block_shape();
assert_eq!(
chacha.fixed_iv_len, 12,
"RFC 7905: the whole nonce is implicit"
);
assert_eq!(chacha.explicit_nonce_len, 0, "RFC 7905: nothing is sent");
assert_eq!(chacha.enc_key_len, 32);
for shape in [gcm, chacha] {
assert_eq!(shape.fixed_iv_len + shape.explicit_nonce_len, 12);
}
}
}