File
Blob: firmware/crates/esp32-radio/src/platform/crypto/suites.rs
| 1 | //! Preserve upstream cipher metadata and ordering; replace AES operations only. |
| 2 | use super::{key::Key, native}; |
| 3 | use dimpl::{ |
| 4 | CryptoError, CryptoOperation, HashAlgorithm, |
| 5 | crypto::{ |
| 6 | Aad, Buf, Cipher, CryptoProvider, Dtls12CipherSuite, Dtls13CipherSuite, Nonce, |
| 7 | SupportedDtls12CipherSuite, SupportedDtls13CipherSuite, TmpBuf, |
| 8 | }, |
| 9 | }; |
| 10 | use std::sync::OnceLock; |
| 11 | |
| 12 | #[derive(Debug)] |
| 13 | struct Gcm(Key); |
| 14 | fn cipher(key: &[u8], expected: usize) -> Result<Box<dyn Cipher>, CryptoError> { |
| 15 | if key.len() != expected { |
| 16 | return Err(CryptoError::OperationFailed(CryptoOperation::Encrypt)); |
| 17 | } |
| 18 | let key = Key::new(key).ok_or(CryptoError::OperationFailed(CryptoOperation::Encrypt))?; |
| 19 | Ok(Box::new(Gcm(key))) |
| 20 | } |
| 21 | impl Cipher for Gcm { |
| 22 | fn encrypt(&mut self, data: &mut Buf, aad: Aad, nonce: Nonce) -> Result<(), CryptoError> { |
| 23 | let length = data.len(); |
| 24 | if length > 16384 { |
| 25 | return Err(CryptoError::OperationFailed(CryptoOperation::Encrypt)); |
| 26 | } |
| 27 | data.resize(length + 16, 0); |
| 28 | let iv: &[u8; 12] = nonce[..12] |
| 29 | .try_into() |
| 30 | .map_err(|_| CryptoError::InvalidNonce)?; |
| 31 | if native::gcm_in_place(false, self.0.bytes(), iv, &aad, data, length).is_err() { |
| 32 | data.clear(); |
| 33 | return Err(CryptoError::OperationFailed(CryptoOperation::Encrypt)); |
| 34 | } |
| 35 | Ok(()) |
| 36 | } |
| 37 | fn decrypt(&mut self, data: &mut TmpBuf, aad: Aad, nonce: Nonce) -> Result<(), CryptoError> { |
| 38 | let length = data.len(); |
| 39 | let clear = length |
| 40 | .checked_sub(16) |
| 41 | .ok_or(CryptoError::OperationFailed(CryptoOperation::Decrypt))?; |
| 42 | let iv: &[u8; 12] = nonce[..12] |
| 43 | .try_into() |
| 44 | .map_err(|_| CryptoError::InvalidNonce)?; |
| 45 | native::gcm_in_place(true, self.0.bytes(), iv, &aad, data.as_mut(), length) |
| 46 | .map_err(|_| CryptoError::OperationFailed(CryptoOperation::Decrypt))?; |
| 47 | data.truncate(clear); |
| 48 | Ok(()) |
| 49 | } |
| 50 | } |
| 51 | |
| 52 | #[derive(Debug)] |
| 53 | struct Suite12(&'static dyn SupportedDtls12CipherSuite); |
| 54 | impl SupportedDtls12CipherSuite for Suite12 { |
| 55 | fn suite(&self) -> Dtls12CipherSuite { |
| 56 | self.0.suite() |
| 57 | } |
| 58 | fn hash_algorithm(&self) -> HashAlgorithm { |
| 59 | self.0.hash_algorithm() |
| 60 | } |
| 61 | fn key_lengths(&self) -> (usize, usize, usize) { |
| 62 | self.0.key_lengths() |
| 63 | } |
| 64 | fn explicit_nonce_len(&self) -> usize { |
| 65 | self.0.explicit_nonce_len() |
| 66 | } |
| 67 | fn tag_len(&self) -> usize { |
| 68 | self.0.tag_len() |
| 69 | } |
| 70 | fn min_protected_fragment_len(&self) -> usize { |
| 71 | self.0.min_protected_fragment_len() |
| 72 | } |
| 73 | fn create_cipher(&self, key: &[u8]) -> Result<Box<dyn Cipher>, CryptoError> { |
| 74 | match self.suite() { |
| 75 | Dtls12CipherSuite::ECDHE_ECDSA_AES128_GCM_SHA256 |
| 76 | | Dtls12CipherSuite::ECDHE_ECDSA_AES256_GCM_SHA384 => { |
| 77 | cipher(key, self.0.key_lengths().1) |
| 78 | } |
| 79 | _ => self.0.create_cipher(key), |
| 80 | } |
| 81 | } |
| 82 | } |
| 83 | |
| 84 | #[derive(Debug)] |
| 85 | struct Suite13(&'static dyn SupportedDtls13CipherSuite); |
| 86 | impl Suite13 { |
| 87 | fn aes(&self) -> bool { |
| 88 | matches!( |
| 89 | self.0.suite(), |
| 90 | Dtls13CipherSuite::AES_128_GCM_SHA256 | Dtls13CipherSuite::AES_256_GCM_SHA384 |
| 91 | ) |
| 92 | } |
| 93 | } |
| 94 | impl SupportedDtls13CipherSuite for Suite13 { |
| 95 | fn suite(&self) -> Dtls13CipherSuite { |
| 96 | self.0.suite() |
| 97 | } |
| 98 | fn hash_algorithm(&self) -> HashAlgorithm { |
| 99 | self.0.hash_algorithm() |
| 100 | } |
| 101 | fn key_len(&self) -> usize { |
| 102 | self.0.key_len() |
| 103 | } |
| 104 | fn iv_len(&self) -> usize { |
| 105 | self.0.iv_len() |
| 106 | } |
| 107 | fn tag_len(&self) -> usize { |
| 108 | self.0.tag_len() |
| 109 | } |
| 110 | fn min_protected_fragment_len(&self) -> usize { |
| 111 | self.0.min_protected_fragment_len() |
| 112 | } |
| 113 | fn create_cipher(&self, key: &[u8]) -> Result<Box<dyn Cipher>, CryptoError> { |
| 114 | if self.aes() { |
| 115 | cipher(key, self.0.key_len()) |
| 116 | } else { |
| 117 | self.0.create_cipher(key) |
| 118 | } |
| 119 | } |
| 120 | fn encrypt_sn(&self, key: &[u8], sample: &[u8; 16]) -> [u8; 16] { |
| 121 | assert_eq!(key.len(), self.key_len(), "invalid DTLS record key length"); |
| 122 | let mut out = [0; 16]; |
| 123 | if self.aes() && native::ecb(key, sample, &mut out).is_ok() { |
| 124 | out |
| 125 | } else { |
| 126 | self.0.encrypt_sn(key, sample) |
| 127 | } |
| 128 | } |
| 129 | } |
| 130 | |
| 131 | pub(super) fn install(provider: &mut CryptoProvider) { |
| 132 | // Upstream requires 'static factory references. These immutable descriptors |
| 133 | // contain only references to upstream factories, never keys or session state. |
| 134 | static VALUES12: OnceLock<Vec<Suite12>> = OnceLock::new(); |
| 135 | static REFS12: OnceLock<Vec<&'static dyn SupportedDtls12CipherSuite>> = OnceLock::new(); |
| 136 | static VALUES13: OnceLock<Vec<Suite13>> = OnceLock::new(); |
| 137 | static REFS13: OnceLock<Vec<&'static dyn SupportedDtls13CipherSuite>> = OnceLock::new(); |
| 138 | let values = VALUES12.get_or_init(|| { |
| 139 | provider |
| 140 | .cipher_suites |
| 141 | .iter() |
| 142 | .copied() |
| 143 | .map(Suite12) |
| 144 | .collect() |
| 145 | }); |
| 146 | provider.cipher_suites = REFS12.get_or_init(|| { |
| 147 | values |
| 148 | .iter() |
| 149 | .map(|v| v as &dyn SupportedDtls12CipherSuite) |
| 150 | .collect() |
| 151 | }); |
| 152 | let values = VALUES13.get_or_init(|| { |
| 153 | provider |
| 154 | .dtls13_cipher_suites |
| 155 | .iter() |
| 156 | .copied() |
| 157 | .map(Suite13) |
| 158 | .collect() |
| 159 | }); |
| 160 | provider.dtls13_cipher_suites = REFS13.get_or_init(|| { |
| 161 | values |
| 162 | .iter() |
| 163 | .map(|v| v as &dyn SupportedDtls13CipherSuite) |
| 164 | .collect() |
| 165 | }); |
| 166 | } |