1use std::fmt;
21use std::str::FromStr;
22
23use aes_gcm::aead::generic_array::typenum::U12;
24use aes_gcm::aead::rand_core::RngCore;
25use aes_gcm::aead::{Aead, AeadCore, KeyInit, OsRng, Payload};
26use aes_gcm::{Aes128Gcm, Aes256Gcm, AesGcm, Nonce};
27use zeroize::Zeroizing;
28
29type Aes192Gcm = AesGcm<aes_gcm::aes::Aes192, U12>;
32
33use crate::{Error, ErrorKind, Result};
34
35#[derive(Clone, PartialEq, Eq)]
46pub struct SensitiveBytes(Zeroizing<Box<[u8]>>);
47
48impl SensitiveBytes {
49 pub fn new(bytes: impl Into<Box<[u8]>>) -> Self {
51 Self(Zeroizing::new(bytes.into()))
52 }
53
54 pub fn as_bytes(&self) -> &[u8] {
56 &self.0
57 }
58
59 pub fn len(&self) -> usize {
61 self.0.len()
62 }
63
64 pub fn is_empty(&self) -> bool {
66 self.0.is_empty()
67 }
68}
69
70impl fmt::Debug for SensitiveBytes {
71 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
72 write!(f, "[{} bytes REDACTED]", self.0.len())
73 }
74}
75
76impl fmt::Display for SensitiveBytes {
77 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
78 write!(f, "[{} bytes REDACTED]", self.0.len())
79 }
80}
81
82#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
87pub enum AesKeySize {
88 #[default]
90 Bits128 = 128,
91 Bits192 = 192,
93 Bits256 = 256,
95}
96
97impl AesKeySize {
98 pub fn key_length(&self) -> usize {
100 match self {
101 Self::Bits128 => 16,
102 Self::Bits192 => 24,
103 Self::Bits256 => 32,
104 }
105 }
106
107 pub fn from_key_length(len: usize) -> Result<Self> {
112 match len {
113 16 => Ok(Self::Bits128),
114 24 => Ok(Self::Bits192),
115 32 => Ok(Self::Bits256),
116 _ => Err(Error::new(
117 ErrorKind::FeatureUnsupported,
118 format!("Unsupported data key length: {len} (must be 16, 24, or 32)"),
119 )),
120 }
121 }
122}
123
124impl FromStr for AesKeySize {
125 type Err = Error;
126
127 fn from_str(s: &str) -> Result<Self> {
128 match s {
129 "128" | "AES_GCM_128" | "AES128_GCM" => Ok(Self::Bits128),
130 "192" | "AES_GCM_192" | "AES192_GCM" => Ok(Self::Bits192),
131 "256" | "AES_GCM_256" | "AES256_GCM" => Ok(Self::Bits256),
132 _ => Err(Error::new(
133 ErrorKind::FeatureUnsupported,
134 format!("Unsupported AES key size: {s}"),
135 )),
136 }
137 }
138}
139
140#[derive(Clone, Debug, PartialEq, Eq)]
145pub struct SecureKey {
146 key: SensitiveBytes,
147 key_size: AesKeySize,
148}
149
150impl SecureKey {
151 pub fn new(key: &[u8]) -> Result<Self> {
156 let key_size = AesKeySize::from_key_length(key.len())?;
157 Ok(Self {
158 key: SensitiveBytes::new(key),
159 key_size,
160 })
161 }
162
163 pub fn generate(key_size: AesKeySize) -> Self {
165 let mut key = vec![0u8; key_size.key_length()];
166 OsRng.fill_bytes(&mut key);
167 Self {
168 key: SensitiveBytes::new(key),
169 key_size,
170 }
171 }
172
173 pub fn key_size(&self) -> AesKeySize {
175 self.key_size
176 }
177
178 pub fn as_bytes(&self) -> &[u8] {
180 self.key.as_bytes()
181 }
182}
183
184impl TryFrom<SensitiveBytes> for SecureKey {
185 type Error = Error;
186
187 fn try_from(key: SensitiveBytes) -> Result<Self> {
188 let key_size = AesKeySize::from_key_length(key.len())?;
189 Ok(Self { key, key_size })
190 }
191}
192
193pub struct AesGcmCipher {
195 key: SensitiveBytes,
196 key_size: AesKeySize,
197}
198
199impl AesGcmCipher {
200 pub const NONCE_LEN: usize = 12;
202 pub const TAG_LEN: usize = 16;
204
205 pub fn new(key: SecureKey) -> Self {
207 Self {
208 key: SensitiveBytes::new(key.as_bytes()),
209 key_size: key.key_size(),
210 }
211 }
212
213 pub fn encrypt(&self, plaintext: &[u8], aad: Option<&[u8]>) -> Result<Vec<u8>> {
223 match self.key_size {
224 AesKeySize::Bits128 => {
225 encrypt_aes_gcm::<Aes128Gcm>(self.key.as_bytes(), plaintext, aad)
226 }
227 AesKeySize::Bits192 => {
228 encrypt_aes_gcm::<Aes192Gcm>(self.key.as_bytes(), plaintext, aad)
229 }
230 AesKeySize::Bits256 => {
231 encrypt_aes_gcm::<Aes256Gcm>(self.key.as_bytes(), plaintext, aad)
232 }
233 }
234 }
235
236 pub fn decrypt(&self, ciphertext: &[u8], aad: Option<&[u8]>) -> Result<Vec<u8>> {
245 if ciphertext.len() < Self::NONCE_LEN + Self::TAG_LEN {
246 return Err(Error::new(
247 ErrorKind::DataInvalid,
248 format!(
249 "Ciphertext too short: expected at least {} bytes, got {}",
250 Self::NONCE_LEN + Self::TAG_LEN,
251 ciphertext.len()
252 ),
253 ));
254 }
255
256 match self.key_size {
257 AesKeySize::Bits128 => {
258 decrypt_aes_gcm::<Aes128Gcm>(self.key.as_bytes(), ciphertext, aad)
259 }
260 AesKeySize::Bits192 => {
261 decrypt_aes_gcm::<Aes192Gcm>(self.key.as_bytes(), ciphertext, aad)
262 }
263 AesKeySize::Bits256 => {
264 decrypt_aes_gcm::<Aes256Gcm>(self.key.as_bytes(), ciphertext, aad)
265 }
266 }
267 }
268}
269
270fn encrypt_aes_gcm<C>(key_bytes: &[u8], plaintext: &[u8], aad: Option<&[u8]>) -> Result<Vec<u8>>
271where C: Aead + AeadCore + KeyInit {
272 let cipher = C::new_from_slice(key_bytes).map_err(|e| {
273 Error::new(ErrorKind::DataInvalid, "Invalid AES key").with_source(anyhow::anyhow!(e))
274 })?;
275 let nonce = C::generate_nonce(&mut OsRng);
276
277 let ciphertext = if let Some(aad) = aad {
278 cipher.encrypt(&nonce, Payload {
279 msg: plaintext,
280 aad,
281 })
282 } else {
283 cipher.encrypt(&nonce, plaintext.as_ref())
284 }
285 .map_err(|e| {
286 Error::new(ErrorKind::Unexpected, "AES-GCM encryption failed")
287 .with_source(anyhow::anyhow!(e))
288 })?;
289
290 let mut result = Vec::with_capacity(nonce.len() + ciphertext.len());
292 result.extend_from_slice(&nonce);
293 result.extend_from_slice(&ciphertext);
294 Ok(result)
295}
296
297fn decrypt_aes_gcm<C>(key_bytes: &[u8], ciphertext: &[u8], aad: Option<&[u8]>) -> Result<Vec<u8>>
298where C: Aead + AeadCore + KeyInit {
299 let cipher = C::new_from_slice(key_bytes).map_err(|e| {
300 Error::new(ErrorKind::DataInvalid, "Invalid AES key").with_source(anyhow::anyhow!(e))
301 })?;
302
303 let nonce = Nonce::from_slice(&ciphertext[..AesGcmCipher::NONCE_LEN]);
304 let encrypted_data = &ciphertext[AesGcmCipher::NONCE_LEN..];
305
306 let plaintext = if let Some(aad) = aad {
307 cipher.decrypt(nonce, Payload {
308 msg: encrypted_data,
309 aad,
310 })
311 } else {
312 cipher.decrypt(nonce, encrypted_data)
313 }
314 .map_err(|e| {
315 Error::new(ErrorKind::Unexpected, "AES-GCM decryption failed")
316 .with_source(anyhow::anyhow!(e))
317 })?;
318
319 Ok(plaintext)
320}
321
322#[cfg(test)]
323mod tests {
324 use super::*;
325
326 #[test]
327 fn test_aes_key_size() {
328 assert_eq!(AesKeySize::Bits128.key_length(), 16);
329 assert_eq!(AesKeySize::Bits192.key_length(), 24);
330 assert_eq!(AesKeySize::Bits256.key_length(), 32);
331
332 assert_eq!(
333 AesKeySize::from_key_length(16).unwrap(),
334 AesKeySize::Bits128
335 );
336 assert_eq!(
337 AesKeySize::from_key_length(24).unwrap(),
338 AesKeySize::Bits192
339 );
340 assert_eq!(
341 AesKeySize::from_key_length(32).unwrap(),
342 AesKeySize::Bits256
343 );
344 assert!(AesKeySize::from_key_length(8).is_err());
345
346 assert_eq!(AesKeySize::from_str("128").unwrap(), AesKeySize::Bits128);
347 assert_eq!(
348 AesKeySize::from_str("AES_GCM_128").unwrap(),
349 AesKeySize::Bits128
350 );
351 assert_eq!(
352 AesKeySize::from_str("AES_GCM_256").unwrap(),
353 AesKeySize::Bits256
354 );
355 assert!(AesKeySize::from_str("INVALID").is_err());
356 }
357
358 #[test]
359 fn test_secure_key() {
360 let key1 = SecureKey::generate(AesKeySize::Bits128);
362 assert_eq!(key1.as_bytes().len(), 16);
363 assert_eq!(key1.key_size(), AesKeySize::Bits128);
364
365 let valid_key = [0u8; 16];
367 assert!(SecureKey::new(valid_key.as_slice()).is_ok());
368
369 let invalid_key = [0u8; 33];
370 assert!(SecureKey::new(invalid_key.as_slice()).is_err());
371 }
372
373 #[test]
374 fn test_aes128_gcm_encryption_roundtrip() {
375 let key = SecureKey::generate(AesKeySize::Bits128);
376 let cipher = AesGcmCipher::new(key);
377
378 let plaintext = b"Hello, Iceberg encryption!";
379 let aad = b"additional authenticated data";
380
381 let ciphertext = cipher.encrypt(plaintext, None).unwrap();
383 assert!(ciphertext.len() > plaintext.len() + 12); assert_ne!(&ciphertext[12..], plaintext); let decrypted = cipher.decrypt(&ciphertext, None).unwrap();
387 assert_eq!(decrypted, plaintext);
388
389 let ciphertext = cipher.encrypt(plaintext, Some(aad)).unwrap();
391 let decrypted = cipher.decrypt(&ciphertext, Some(aad)).unwrap();
392 assert_eq!(decrypted, plaintext);
393
394 assert!(cipher.decrypt(&ciphertext, Some(b"wrong aad")).is_err());
396 }
397
398 #[test]
399 fn test_aes192_gcm_encryption_roundtrip() {
400 let key = SecureKey::generate(AesKeySize::Bits192);
401 let cipher = AesGcmCipher::new(key);
402
403 let plaintext = b"Hello, Iceberg encryption!";
404 let aad = b"additional authenticated data";
405
406 let ciphertext = cipher.encrypt(plaintext, None).unwrap();
408 let decrypted = cipher.decrypt(&ciphertext, None).unwrap();
409 assert_eq!(decrypted, plaintext);
410
411 let ciphertext = cipher.encrypt(plaintext, Some(aad)).unwrap();
413 let decrypted = cipher.decrypt(&ciphertext, Some(aad)).unwrap();
414 assert_eq!(decrypted, plaintext);
415
416 assert!(cipher.decrypt(&ciphertext, Some(b"wrong aad")).is_err());
418 }
419
420 #[test]
421 fn test_aes256_gcm_encryption_roundtrip() {
422 let key = SecureKey::generate(AesKeySize::Bits256);
423 let cipher = AesGcmCipher::new(key);
424
425 let plaintext = b"Hello, Iceberg encryption!";
426 let aad = b"additional authenticated data";
427
428 let ciphertext = cipher.encrypt(plaintext, None).unwrap();
430 let decrypted = cipher.decrypt(&ciphertext, None).unwrap();
431 assert_eq!(decrypted, plaintext);
432
433 let ciphertext = cipher.encrypt(plaintext, Some(aad)).unwrap();
435 let decrypted = cipher.decrypt(&ciphertext, Some(aad)).unwrap();
436 assert_eq!(decrypted, plaintext);
437
438 assert!(cipher.decrypt(&ciphertext, Some(b"wrong aad")).is_err());
440 }
441
442 #[test]
443 fn test_cross_key_size_incompatibility() {
444 let plaintext = b"Cross-key test";
445
446 let key128 = SecureKey::generate(AesKeySize::Bits128);
447 let key256 = SecureKey::generate(AesKeySize::Bits256);
448
449 let cipher128 = AesGcmCipher::new(key128);
450 let cipher256 = AesGcmCipher::new(key256);
451
452 let ciphertext = cipher128.encrypt(plaintext, None).unwrap();
454 assert!(cipher256.decrypt(&ciphertext, None).is_err());
455 }
456
457 #[test]
458 fn test_encryption_with_empty_plaintext() {
459 let key = SecureKey::generate(AesKeySize::Bits128);
460 let cipher = AesGcmCipher::new(key);
461
462 let plaintext = b"";
463 let ciphertext = cipher.encrypt(plaintext, None).unwrap();
464
465 assert_eq!(ciphertext.len(), 12 + 16); let decrypted = cipher.decrypt(&ciphertext, None).unwrap();
469 assert_eq!(decrypted, plaintext);
470 }
471
472 #[test]
473 fn test_decryption_with_tampered_ciphertext() {
474 let key = SecureKey::generate(AesKeySize::Bits128);
475 let cipher = AesGcmCipher::new(key);
476
477 let plaintext = b"Sensitive data";
478 let mut ciphertext = cipher.encrypt(plaintext, None).unwrap();
479
480 if ciphertext.len() > 12 {
482 ciphertext[12] ^= 0xFF;
483 }
484
485 assert!(cipher.decrypt(&ciphertext, None).is_err());
487 }
488
489 #[test]
490 fn test_different_keys_produce_different_ciphertexts() {
491 let key1 = SecureKey::generate(AesKeySize::Bits128);
492 let key2 = SecureKey::generate(AesKeySize::Bits128);
493
494 let cipher1 = AesGcmCipher::new(key1);
495 let cipher2 = AesGcmCipher::new(key2);
496
497 let plaintext = b"Same plaintext";
498
499 let ciphertext1 = cipher1.encrypt(plaintext, None).unwrap();
500 let ciphertext2 = cipher2.encrypt(plaintext, None).unwrap();
501
502 assert_ne!(&ciphertext1[12..], &ciphertext2[12..]);
505 }
506
507 #[test]
508 fn test_ciphertext_format_java_compatible() {
509 let key = SecureKey::generate(AesKeySize::Bits128);
511 let cipher = AesGcmCipher::new(key);
512
513 let plaintext = b"Test data";
514 let ciphertext = cipher.encrypt(plaintext, None).unwrap();
515
516 assert_eq!(
518 ciphertext.len(),
519 12 + plaintext.len() + 16,
520 "Ciphertext should be nonce + plaintext + tag length"
521 );
522
523 let nonce = &ciphertext[..12];
525 assert_eq!(nonce.len(), 12, "Nonce should be 12 bytes");
526
527 let encrypted_with_tag = &ciphertext[12..];
529 assert_eq!(
530 encrypted_with_tag.len(),
531 plaintext.len() + 16,
532 "Encrypted portion should be plaintext length + 16-byte tag"
533 );
534 }
535}