1use super::*;
2
3use std::convert::TryFrom;
4
5use std::mem::MaybeUninit;
6
7use std::ptr::NonNull;
8
9use libc::c_int;
10use libc::c_uint;
11use libc::c_void;
12
13#[allow(non_camel_case_types)]
14#[repr(transparent)]
15struct EVP_AEAD_CTX {
16 _unused: c_void,
17}
18
19#[derive(Clone)]
20#[repr(C)]
21pub(crate) struct AES_KEY {
22 rd_key: [u32; 4 * (14 + 1)],
23 rounds: c_int,
24}
25
26impl Algorithm {
27 fn get_evp_aead(self) -> *const EVP_AEAD {
28 match self {
29 Algorithm::AES128_GCM => unsafe { EVP_aead_aes_128_gcm_tls13() },
30 Algorithm::AES256_GCM => unsafe { EVP_aead_aes_256_gcm_tls13() },
31 Algorithm::ChaCha20_Poly1305 => unsafe {
32 EVP_aead_chacha20_poly1305()
33 },
34 }
35 }
36}
37
38pub(crate) struct PacketKey {
39 alg: Algorithm,
40
41 ctx: NonNull<EVP_AEAD_CTX>,
42
43 nonce: Vec<u8>,
44}
45
46impl PacketKey {
47 pub fn new(
48 alg: Algorithm, key: Vec<u8>, iv: Vec<u8>, _enc: u32,
49 ) -> Result<Self> {
50 Ok(Self {
51 alg,
52 ctx: make_aead_ctx(alg, &key)?,
53 nonce: iv,
54 })
55 }
56
57 pub fn from_secret(aead: Algorithm, secret: &[u8], enc: u32) -> Result<Self> {
58 let key_len = aead.key_len();
59 let nonce_len = aead.nonce_len();
60
61 let mut key = vec![0; key_len];
62 let mut iv = vec![0; nonce_len];
63
64 derive_pkt_key(aead, secret, &mut key)?;
65 derive_pkt_iv(aead, secret, &mut iv)?;
66
67 let mut pkt_key = Self::new(aead, key, iv, enc)?;
68
69 let _ = pkt_key.seal_with_u64_counter(0, b"", &mut [0_u8; 16], 0, None);
76
77 Ok(pkt_key)
78 }
79
80 pub fn open_with_u64_counter(
81 &self, counter: u64, ad: &[u8], buf: &mut [u8],
82 ) -> Result<usize> {
83 let tag_len = self.alg.tag_len();
84
85 let mut out_len = match buf.len().checked_sub(tag_len) {
86 Some(n) => n,
87 None => return Err(Error::CryptoFail),
88 };
89
90 let max_out_len = out_len;
91
92 let nonce = make_nonce(&self.nonce, counter);
93
94 let rc = unsafe {
95 EVP_AEAD_CTX_open(
96 self.ctx.as_ptr(), buf.as_mut_ptr(), &mut out_len, max_out_len, nonce[..].as_ptr(), nonce.len(), buf.as_ptr(), buf.len(), ad.as_ptr(), ad.len(), )
107 };
108
109 if rc != 1 {
110 return Err(Error::CryptoFail);
111 }
112
113 Ok(out_len)
114 }
115
116 pub fn seal_with_u64_counter(
117 &mut self, counter: u64, ad: &[u8], buf: &mut [u8], in_len: usize,
118 extra_in: Option<&[u8]>,
119 ) -> Result<usize> {
120 let tag_len = self.alg.tag_len();
121
122 let mut out_tag_len = tag_len;
123
124 let (extra_in_ptr, extra_in_len) = match extra_in {
125 Some(v) => (v.as_ptr(), v.len()),
126
127 None => (std::ptr::null(), 0),
128 };
129
130 if in_len + tag_len + extra_in_len > buf.len() {
132 return Err(Error::CryptoFail);
133 }
134
135 let nonce = make_nonce(&self.nonce, counter);
136
137 let rc = unsafe {
138 EVP_AEAD_CTX_seal_scatter(
139 self.ctx.as_ptr(), buf.as_mut_ptr(), buf[in_len..].as_mut_ptr(), &mut out_tag_len, tag_len + extra_in_len, nonce[..].as_ptr(), nonce.len(), buf.as_ptr(), in_len, extra_in_ptr, extra_in_len, ad.as_ptr(), ad.len(), )
153 };
154
155 if rc != 1 {
156 return Err(Error::CryptoFail);
157 }
158
159 Ok(in_len + out_tag_len)
160 }
161}
162
163unsafe impl Send for PacketKey {}
166unsafe impl Sync for PacketKey {}
167
168impl Drop for PacketKey {
169 fn drop(&mut self) {
170 unsafe { EVP_AEAD_CTX_free(self.ctx.as_ptr()) }
171 }
172}
173
174#[derive(Clone)]
175#[allow(clippy::large_enum_variant)]
176pub(crate) enum HeaderProtectionKey {
177 Aes(AES_KEY),
178
179 ChaCha(Vec<u8>),
180}
181
182impl HeaderProtectionKey {
183 pub fn new(alg: Algorithm, hp_key: Vec<u8>) -> Result<Self> {
184 match alg {
185 Algorithm::AES128_GCM | Algorithm::AES256_GCM => unsafe {
186 let key_len_bits = alg.key_len() as u32 * 8;
187
188 let mut aes_key = MaybeUninit::<AES_KEY>::uninit();
189
190 let rc = AES_set_encrypt_key(
191 hp_key.as_ptr(),
192 key_len_bits,
193 aes_key.as_mut_ptr(),
194 );
195
196 if rc != 0 {
197 return Err(Error::CryptoFail);
198 }
199
200 let aes_key = aes_key.assume_init();
201 Ok(Self::Aes(aes_key))
202 },
203
204 Algorithm::ChaCha20_Poly1305 => Ok(Self::ChaCha(hp_key)),
205 }
206 }
207
208 pub fn new_mask(&self, sample: &[u8]) -> Result<HeaderProtectionMask> {
209 match self {
210 Self::Aes(aes_key) => {
211 let mut block = [0_u8; 16];
212
213 unsafe {
214 AES_ecb_encrypt(
215 sample.as_ptr(),
216 block.as_mut_ptr(),
217 aes_key as _,
218 1,
219 )
220 };
221
222 let new_mask =
228 HeaderProtectionMask::try_from(&block[..HP_MASK_LEN])
229 .unwrap();
230 Ok(new_mask)
231 },
232
233 Self::ChaCha(key) => {
234 const PLAINTEXT: &[u8; HP_MASK_LEN] = &[0_u8; HP_MASK_LEN];
235
236 let mut new_mask = HeaderProtectionMask::default();
237
238 let counter = u32::from_le_bytes([
239 sample[0], sample[1], sample[2], sample[3],
240 ]);
241
242 unsafe {
243 CRYPTO_chacha_20(
244 new_mask.as_mut_ptr(),
245 PLAINTEXT.as_ptr(),
246 PLAINTEXT.len(),
247 key.as_ptr(),
248 sample[size_of::<u32>()..].as_ptr(),
249 counter,
250 );
251 };
252
253 Ok(new_mask)
254 },
255 }
256 }
257}
258
259fn make_aead_ctx(alg: Algorithm, key: &[u8]) -> Result<NonNull<EVP_AEAD_CTX>> {
260 let ctx = unsafe {
261 let aead = alg.get_evp_aead();
262
263 EVP_AEAD_CTX_new(aead, key.as_ptr(), alg.key_len(), alg.tag_len())
264 };
265
266 NonNull::new(ctx).ok_or(Error::CryptoFail)
267}
268
269pub(crate) fn hkdf_extract(
270 alg: Algorithm, out: &mut [u8], secret: &[u8], salt: &[u8],
271) -> Result<()> {
272 let mut out_len = out.len();
273
274 let rc = unsafe {
275 HKDF_extract(
276 out.as_mut_ptr(),
277 &mut out_len,
278 alg.get_evp_digest(),
279 secret.as_ptr(),
280 secret.len(),
281 salt.as_ptr(),
282 salt.len(),
283 )
284 };
285
286 if rc != 1 {
287 return Err(Error::CryptoFail);
288 }
289
290 Ok(())
291}
292
293pub(crate) fn hkdf_expand(
294 alg: Algorithm, out: &mut [u8], secret: &[u8], info: &[u8],
295) -> Result<()> {
296 let rc = unsafe {
297 HKDF_expand(
298 out.as_mut_ptr(),
299 out.len(),
300 alg.get_evp_digest(),
301 secret.as_ptr(),
302 secret.len(),
303 info.as_ptr(),
304 info.len(),
305 )
306 };
307
308 if rc != 1 {
309 return Err(Error::CryptoFail);
310 }
311
312 Ok(())
313}
314
315extern "C" {
316 fn EVP_aead_aes_128_gcm_tls13() -> *const EVP_AEAD;
317
318 fn EVP_aead_aes_256_gcm_tls13() -> *const EVP_AEAD;
319
320 fn EVP_aead_chacha20_poly1305() -> *const EVP_AEAD;
321
322 fn HKDF_extract(
324 out_key: *mut u8, out_len: *mut usize, digest: *const EVP_MD,
325 secret: *const u8, secret_len: usize, salt: *const u8, salt_len: usize,
326 ) -> c_int;
327
328 fn HKDF_expand(
329 out_key: *mut u8, out_len: usize, digest: *const EVP_MD, prk: *const u8,
330 prk_len: usize, info: *const u8, info_len: usize,
331 ) -> c_int;
332
333 fn EVP_AEAD_CTX_new(
335 aead: *const EVP_AEAD, key: *const u8, key_len: usize, tag_len: usize,
336 ) -> *mut EVP_AEAD_CTX;
337
338 fn EVP_AEAD_CTX_free(ctx: *mut EVP_AEAD_CTX);
339
340 fn EVP_AEAD_CTX_open(
341 ctx: *const EVP_AEAD_CTX, out: *mut u8, out_len: *mut usize,
342 max_out_len: usize, nonce: *const u8, nonce_len: usize, inp: *const u8,
343 in_len: usize, ad: *const u8, ad_len: usize,
344 ) -> c_int;
345
346 fn EVP_AEAD_CTX_seal_scatter(
347 ctx: *mut EVP_AEAD_CTX, out: *mut u8, out_tag: *mut u8,
348 out_tag_len: *mut usize, max_out_tag_len: usize, nonce: *const u8,
349 nonce_len: usize, inp: *const u8, in_len: usize, extra_in: *const u8,
350 extra_in_len: usize, ad: *const u8, ad_len: usize,
351 ) -> c_int;
352
353 fn AES_set_encrypt_key(
355 key: *const u8, bits: c_uint, aeskey: *mut AES_KEY,
356 ) -> c_int;
357
358 fn AES_ecb_encrypt(
359 inp: *const u8, out: *mut u8, key: *const AES_KEY, enc: c_int,
360 );
361
362 fn CRYPTO_chacha_20(
364 out: *mut u8, inp: *const u8, in_len: usize, key: *const u8,
365 nonce: *const u8, counter: u32,
366 );
367}