Skip to content
File

Blob: firmware/crates/esp32-radio/src/platform/crypto/suites.rs

rust167 lines
1//! Preserve upstream cipher metadata and ordering; replace AES operations only.
2use super::{key::Key, native};
3use dimpl::{
4 CryptoError, CryptoOperation, HashAlgorithm,
5 crypto::{
6 Aad, Buf, Cipher, CryptoProvider, Dtls12CipherSuite, Dtls13CipherSuite, Nonce,
7 SupportedDtls12CipherSuite, SupportedDtls13CipherSuite, TmpBuf,
8 },
9};
10use std::sync::OnceLock;
11 
12#[derive(Debug)]
13struct Gcm(Key);
14fn 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}
21impl 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)]
53struct Suite12(&'static dyn SupportedDtls12CipherSuite);
54impl 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)]
85struct Suite13(&'static dyn SupportedDtls13CipherSuite);
86impl 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}
94impl 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 
131pub(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}