1use std::str::FromStr;
21
22use aes_gcm::aead::generic_array::typenum::U12;
23use aes_gcm::aead::rand_core::RngCore;
24use aes_gcm::aead::{Aead, AeadCore, KeyInit, OsRng, Payload};
25use aes_gcm::{Aes128Gcm, Aes256Gcm, AesGcm, Nonce};
26
27type Aes192Gcm = AesGcm<aes_gcm::aes::Aes192, U12>;
30
31use crate::error::invalid_data;
32use crate::sensitive::SensitiveBytes;
33use crate::{Error, ErrorKind, Result};
34
35#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
40pub enum AesKeySize {
41 #[default]
43 Bits128 = 128,
44 Bits192 = 192,
46 Bits256 = 256,
48}
49
50impl AesKeySize {
51 pub fn key_length(&self) -> usize {
53 match self {
54 Self::Bits128 => 16,
55 Self::Bits192 => 24,
56 Self::Bits256 => 32,
57 }
58 }
59
60 pub fn from_key_length(len: usize) -> Result<Self> {
65 match len {
66 16 => Ok(Self::Bits128),
67 24 => Ok(Self::Bits192),
68 32 => Ok(Self::Bits256),
69 _ => Err(invalid_data!(
70 "Invalid data key length: {len} (must be 16, 24, or 32)"
71 )),
72 }
73 }
74}
75
76impl FromStr for AesKeySize {
77 type Err = Error;
78
79 fn from_str(s: &str) -> Result<Self> {
80 match s {
81 "128" | "AES_GCM_128" | "AES128_GCM" => Ok(Self::Bits128),
82 "192" | "AES_GCM_192" | "AES192_GCM" => Ok(Self::Bits192),
83 "256" | "AES_GCM_256" | "AES256_GCM" => Ok(Self::Bits256),
84 _ => Err(invalid_data!("Invalid AES key size: {s}")),
85 }
86 }
87}
88
89#[derive(Clone, Debug, PartialEq, Eq)]
94pub struct SecureKey {
95 key: SensitiveBytes,
96 key_size: AesKeySize,
97}
98
99impl SecureKey {
100 pub fn new(key: &[u8]) -> Result<Self> {
105 let key_size = AesKeySize::from_key_length(key.len())?;
106 Ok(Self {
107 key: SensitiveBytes::new(key),
108 key_size,
109 })
110 }
111
112 pub fn generate(key_size: AesKeySize) -> Self {
114 let mut key = vec![0u8; key_size.key_length()];
115 OsRng.fill_bytes(&mut key);
116 Self {
117 key: SensitiveBytes::new(key),
118 key_size,
119 }
120 }
121
122 pub fn key_size(&self) -> AesKeySize {
124 self.key_size
125 }
126
127 pub fn as_bytes(&self) -> &[u8] {
129 self.key.as_bytes()
130 }
131}
132
133impl TryFrom<SensitiveBytes> for SecureKey {
134 type Error = Error;
135
136 fn try_from(key: SensitiveBytes) -> Result<Self> {
137 let key_size = AesKeySize::from_key_length(key.len())?;
138 Ok(Self { key, key_size })
139 }
140}
141
142pub struct AesGcmCipher {
144 key: SensitiveBytes,
145 key_size: AesKeySize,
146}
147
148impl AesGcmCipher {
149 pub const NONCE_LEN: usize = 12;
151 pub const TAG_LEN: usize = 16;
153
154 pub fn new(key: SecureKey) -> Self {
156 Self {
157 key: SensitiveBytes::new(key.as_bytes()),
158 key_size: key.key_size(),
159 }
160 }
161
162 pub fn encrypt(&self, plaintext: &[u8], aad: Option<&[u8]>) -> Result<Vec<u8>> {
172 match self.key_size {
173 AesKeySize::Bits128 => {
174 encrypt_aes_gcm::<Aes128Gcm>(self.key.as_bytes(), plaintext, aad)
175 }
176 AesKeySize::Bits192 => {
177 encrypt_aes_gcm::<Aes192Gcm>(self.key.as_bytes(), plaintext, aad)
178 }
179 AesKeySize::Bits256 => {
180 encrypt_aes_gcm::<Aes256Gcm>(self.key.as_bytes(), plaintext, aad)
181 }
182 }
183 }
184
185 pub fn decrypt(&self, ciphertext: &[u8], aad: Option<&[u8]>) -> Result<Vec<u8>> {
194 if ciphertext.len() < Self::NONCE_LEN + Self::TAG_LEN {
195 return Err(invalid_data!(
196 "Ciphertext too short: expected at least {} bytes, got {}",
197 Self::NONCE_LEN + Self::TAG_LEN,
198 ciphertext.len()
199 ));
200 }
201
202 match self.key_size {
203 AesKeySize::Bits128 => {
204 decrypt_aes_gcm::<Aes128Gcm>(self.key.as_bytes(), ciphertext, aad)
205 }
206 AesKeySize::Bits192 => {
207 decrypt_aes_gcm::<Aes192Gcm>(self.key.as_bytes(), ciphertext, aad)
208 }
209 AesKeySize::Bits256 => {
210 decrypt_aes_gcm::<Aes256Gcm>(self.key.as_bytes(), ciphertext, aad)
211 }
212 }
213 }
214}
215
216fn encrypt_aes_gcm<C>(key_bytes: &[u8], plaintext: &[u8], aad: Option<&[u8]>) -> Result<Vec<u8>>
217where C: Aead + AeadCore + KeyInit {
218 let cipher = C::new_from_slice(key_bytes)
219 .map_err(|e| invalid_data!("Invalid AES key").with_source(anyhow::anyhow!(e)))?;
220 let nonce = C::generate_nonce(&mut OsRng);
221
222 let ciphertext = if let Some(aad) = aad {
223 cipher.encrypt(&nonce, Payload {
224 msg: plaintext,
225 aad,
226 })
227 } else {
228 cipher.encrypt(&nonce, plaintext.as_ref())
229 }
230 .map_err(|e| {
231 Error::new(ErrorKind::Unexpected, "AES-GCM encryption failed")
232 .with_source(anyhow::anyhow!(e))
233 })?;
234
235 let mut result = Vec::with_capacity(nonce.len() + ciphertext.len());
237 result.extend_from_slice(&nonce);
238 result.extend_from_slice(&ciphertext);
239 Ok(result)
240}
241
242fn decrypt_aes_gcm<C>(key_bytes: &[u8], ciphertext: &[u8], aad: Option<&[u8]>) -> Result<Vec<u8>>
243where C: Aead + AeadCore + KeyInit {
244 let cipher = C::new_from_slice(key_bytes)
245 .map_err(|e| invalid_data!("Invalid AES key").with_source(anyhow::anyhow!(e)))?;
246
247 let nonce = Nonce::from_slice(&ciphertext[..AesGcmCipher::NONCE_LEN]);
248 let encrypted_data = &ciphertext[AesGcmCipher::NONCE_LEN..];
249
250 let plaintext = if let Some(aad) = aad {
251 cipher.decrypt(nonce, Payload {
252 msg: encrypted_data,
253 aad,
254 })
255 } else {
256 cipher.decrypt(nonce, encrypted_data)
257 }
258 .map_err(|e| {
259 Error::new(ErrorKind::Unexpected, "AES-GCM decryption failed")
260 .with_source(anyhow::anyhow!(e))
261 })?;
262
263 Ok(plaintext)
264}
265
266#[cfg(test)]
267mod tests {
268 use super::*;
269
270 #[test]
271 fn test_aes_key_size() {
272 assert_eq!(AesKeySize::Bits128.key_length(), 16);
273 assert_eq!(AesKeySize::Bits192.key_length(), 24);
274 assert_eq!(AesKeySize::Bits256.key_length(), 32);
275
276 assert_eq!(
277 AesKeySize::from_key_length(16).unwrap(),
278 AesKeySize::Bits128
279 );
280 assert_eq!(
281 AesKeySize::from_key_length(24).unwrap(),
282 AesKeySize::Bits192
283 );
284 assert_eq!(
285 AesKeySize::from_key_length(32).unwrap(),
286 AesKeySize::Bits256
287 );
288 assert!(AesKeySize::from_key_length(8).is_err());
289
290 for len in [0, 8, 15, 20, 33] {
291 let err = AesKeySize::from_key_length(len).unwrap_err();
292 assert_eq!(err.kind(), ErrorKind::DataInvalid, "for length {len}");
293 assert_eq!(
294 err.message(),
295 format!("Invalid data key length: {len} (must be 16, 24, or 32)")
296 );
297 }
298
299 assert_eq!(AesKeySize::from_str("128").unwrap(), AesKeySize::Bits128);
300 assert_eq!(
301 AesKeySize::from_str("AES_GCM_128").unwrap(),
302 AesKeySize::Bits128
303 );
304 assert_eq!(
305 AesKeySize::from_str("AES_GCM_256").unwrap(),
306 AesKeySize::Bits256
307 );
308 for size in ["", "127", "AES_GCM_512", "INVALID"] {
309 let err = AesKeySize::from_str(size).unwrap_err();
310 assert_eq!(err.kind(), ErrorKind::DataInvalid, "for size {size}");
311 assert_eq!(err.message(), format!("Invalid AES key size: {size}"));
312 }
313 }
314
315 #[test]
316 fn test_secure_key() {
317 let key1 = SecureKey::generate(AesKeySize::Bits128);
319 assert_eq!(key1.as_bytes().len(), 16);
320 assert_eq!(key1.key_size(), AesKeySize::Bits128);
321
322 let valid_key = [0u8; 16];
324 assert!(SecureKey::new(valid_key.as_slice()).is_ok());
325
326 let invalid_key = [0u8; 33];
327 assert!(SecureKey::new(invalid_key.as_slice()).is_err());
328 }
329
330 #[test]
331 fn test_aes128_gcm_encryption_roundtrip() {
332 let key = SecureKey::generate(AesKeySize::Bits128);
333 let cipher = AesGcmCipher::new(key);
334
335 let plaintext = b"Hello, Iceberg encryption!";
336 let aad = b"additional authenticated data";
337
338 let ciphertext = cipher.encrypt(plaintext, None).unwrap();
340 assert!(ciphertext.len() > plaintext.len() + 12); assert_ne!(&ciphertext[12..], plaintext); let decrypted = cipher.decrypt(&ciphertext, None).unwrap();
344 assert_eq!(decrypted, plaintext);
345
346 let ciphertext = cipher.encrypt(plaintext, Some(aad)).unwrap();
348 let decrypted = cipher.decrypt(&ciphertext, Some(aad)).unwrap();
349 assert_eq!(decrypted, plaintext);
350
351 assert!(cipher.decrypt(&ciphertext, Some(b"wrong aad")).is_err());
353 }
354
355 #[test]
356 fn test_aes192_gcm_encryption_roundtrip() {
357 let key = SecureKey::generate(AesKeySize::Bits192);
358 let cipher = AesGcmCipher::new(key);
359
360 let plaintext = b"Hello, Iceberg encryption!";
361 let aad = b"additional authenticated data";
362
363 let ciphertext = cipher.encrypt(plaintext, None).unwrap();
365 let decrypted = cipher.decrypt(&ciphertext, None).unwrap();
366 assert_eq!(decrypted, plaintext);
367
368 let ciphertext = cipher.encrypt(plaintext, Some(aad)).unwrap();
370 let decrypted = cipher.decrypt(&ciphertext, Some(aad)).unwrap();
371 assert_eq!(decrypted, plaintext);
372
373 assert!(cipher.decrypt(&ciphertext, Some(b"wrong aad")).is_err());
375 }
376
377 #[test]
378 fn test_aes256_gcm_encryption_roundtrip() {
379 let key = SecureKey::generate(AesKeySize::Bits256);
380 let cipher = AesGcmCipher::new(key);
381
382 let plaintext = b"Hello, Iceberg encryption!";
383 let aad = b"additional authenticated data";
384
385 let ciphertext = cipher.encrypt(plaintext, None).unwrap();
387 let decrypted = cipher.decrypt(&ciphertext, None).unwrap();
388 assert_eq!(decrypted, plaintext);
389
390 let ciphertext = cipher.encrypt(plaintext, Some(aad)).unwrap();
392 let decrypted = cipher.decrypt(&ciphertext, Some(aad)).unwrap();
393 assert_eq!(decrypted, plaintext);
394
395 assert!(cipher.decrypt(&ciphertext, Some(b"wrong aad")).is_err());
397 }
398
399 #[test]
400 fn test_cross_key_size_incompatibility() {
401 let plaintext = b"Cross-key test";
402
403 let key128 = SecureKey::generate(AesKeySize::Bits128);
404 let key256 = SecureKey::generate(AesKeySize::Bits256);
405
406 let cipher128 = AesGcmCipher::new(key128);
407 let cipher256 = AesGcmCipher::new(key256);
408
409 let ciphertext = cipher128.encrypt(plaintext, None).unwrap();
411 assert!(cipher256.decrypt(&ciphertext, None).is_err());
412 }
413
414 #[test]
415 fn test_encryption_with_empty_plaintext() {
416 let key = SecureKey::generate(AesKeySize::Bits128);
417 let cipher = AesGcmCipher::new(key);
418
419 let plaintext = b"";
420 let ciphertext = cipher.encrypt(plaintext, None).unwrap();
421
422 assert_eq!(ciphertext.len(), 12 + 16); let decrypted = cipher.decrypt(&ciphertext, None).unwrap();
426 assert_eq!(decrypted, plaintext);
427 }
428
429 #[test]
430 fn test_decryption_with_tampered_ciphertext() {
431 let key = SecureKey::generate(AesKeySize::Bits128);
432 let cipher = AesGcmCipher::new(key);
433
434 let plaintext = b"Sensitive data";
435 let mut ciphertext = cipher.encrypt(plaintext, None).unwrap();
436
437 if ciphertext.len() > 12 {
439 ciphertext[12] ^= 0xFF;
440 }
441
442 assert!(cipher.decrypt(&ciphertext, None).is_err());
444 }
445
446 #[test]
447 fn test_different_keys_produce_different_ciphertexts() {
448 let key1 = SecureKey::generate(AesKeySize::Bits128);
449 let key2 = SecureKey::generate(AesKeySize::Bits128);
450
451 let cipher1 = AesGcmCipher::new(key1);
452 let cipher2 = AesGcmCipher::new(key2);
453
454 let plaintext = b"Same plaintext";
455
456 let ciphertext1 = cipher1.encrypt(plaintext, None).unwrap();
457 let ciphertext2 = cipher2.encrypt(plaintext, None).unwrap();
458
459 assert_ne!(&ciphertext1[12..], &ciphertext2[12..]);
462 }
463
464 #[test]
465 fn test_ciphertext_format_java_compatible() {
466 let key = SecureKey::generate(AesKeySize::Bits128);
468 let cipher = AesGcmCipher::new(key);
469
470 let plaintext = b"Test data";
471 let ciphertext = cipher.encrypt(plaintext, None).unwrap();
472
473 assert_eq!(
475 ciphertext.len(),
476 12 + plaintext.len() + 16,
477 "Ciphertext should be nonce + plaintext + tag length"
478 );
479
480 let nonce = &ciphertext[..12];
482 assert_eq!(nonce.len(), 12, "Nonce should be 12 bytes");
483
484 let encrypted_with_tag = &ciphertext[12..];
486 assert_eq!(
487 encrypted_with_tag.len(),
488 plaintext.len() + 16,
489 "Encrypted portion should be plaintext length + 16-byte tag"
490 );
491 }
492}