1use std::fmt;
22
23use aes_gcm::aead::OsRng;
24use aes_gcm::aead::rand_core::RngCore;
25
26use super::{AesKeySize, SecureKey};
27use crate::error::invalid_data;
28use crate::{Error, ErrorKind, Result};
29
30#[derive(Clone, PartialEq, Eq)]
38pub struct StandardKeyMetadata {
39 encryption_key: SecureKey,
40 aad_prefix: Option<Box<[u8]>>,
41 file_length: Option<u64>,
42}
43
44impl fmt::Debug for StandardKeyMetadata {
45 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
46 f.debug_struct("StandardKeyMetadata")
47 .field("encryption_key", &self.encryption_key)
48 .field(
49 "aad_prefix",
50 &self
51 .aad_prefix
52 .as_ref()
53 .map(|b| format!("[{} bytes]", b.len())),
54 )
55 .field("file_length", &self.file_length)
56 .finish()
57 }
58}
59
60impl StandardKeyMetadata {
61 pub fn try_new(encryption_key: &[u8]) -> Result<Self> {
63 Ok(Self::from(SecureKey::new(encryption_key)?))
64 }
65
66 pub(crate) fn generate(key_size: AesKeySize) -> Self {
69 Self::from(SecureKey::generate(key_size)).with_aad_prefix(&generate_aad_prefix())
70 }
71
72 pub fn with_aad_prefix(mut self, aad_prefix: &[u8]) -> Self {
74 self.aad_prefix = Some(aad_prefix.into());
75 self
76 }
77
78 pub fn with_file_length(mut self, length: u64) -> Self {
80 self.file_length = Some(length);
81 self
82 }
83
84 pub fn encryption_key(&self) -> &SecureKey {
86 &self.encryption_key
87 }
88
89 pub fn aad_prefix(&self) -> Option<&[u8]> {
91 self.aad_prefix.as_deref()
92 }
93
94 pub fn file_length(&self) -> Option<u64> {
96 self.file_length
97 }
98
99 pub fn encode(&self) -> Result<Box<[u8]>> {
101 _serde::StandardKeyMetadataV1::try_from(self)?.encode()
102 }
103
104 pub fn decode(bytes: &[u8]) -> Result<Self> {
106 _serde::StandardKeyMetadataV1::decode(bytes).and_then(Self::try_from)
107 }
108}
109
110impl From<SecureKey> for StandardKeyMetadata {
111 fn from(encryption_key: SecureKey) -> Self {
113 Self {
114 encryption_key,
115 aad_prefix: None,
116 file_length: None,
117 }
118 }
119}
120
121const AAD_PREFIX_LENGTH: usize = 16;
123
124fn generate_aad_prefix() -> Box<[u8]> {
125 let mut prefix = vec![0u8; AAD_PREFIX_LENGTH];
126 OsRng.fill_bytes(&mut prefix);
127 prefix.into_boxed_slice()
128}
129
130mod _serde {
131 use std::io::Cursor;
132 use std::sync::{Arc, LazyLock};
133
134 use apache_avro::reader::datum::GenericDatumReader;
135 use apache_avro::writer::datum::GenericDatumWriter;
136 use apache_avro::{Schema as AvroSchema, from_value, to_value};
137 use serde::{Deserialize, Serialize};
138
139 use super::*;
140 use crate::avro::schema_to_avro_schema;
141 use crate::spec::{NestedField, PrimitiveType, Schema, Type};
142
143 pub(super) const V1: u8 = 1;
144
145 pub(super) static AVRO_SCHEMA_V1: LazyLock<AvroSchema> = LazyLock::new(|| {
147 let schema = Schema::builder()
148 .with_fields(vec![
149 Arc::new(NestedField::required(
150 0,
151 "encryption_key",
152 Type::Primitive(PrimitiveType::Binary),
153 )),
154 Arc::new(NestedField::optional(
155 1,
156 "aad_prefix",
157 Type::Primitive(PrimitiveType::Binary),
158 )),
159 Arc::new(NestedField::optional(
160 2,
161 "file_length",
162 Type::Primitive(PrimitiveType::Long),
163 )),
164 ])
165 .build()
166 .expect("Failed to build StandardKeyMetadata Iceberg schema");
167
168 schema_to_avro_schema("StandardKeyMetadata", &schema)
169 .expect("Failed to convert StandardKeyMetadata to Avro schema")
170 });
171
172 #[derive(Serialize, Deserialize)]
175 pub(super) struct StandardKeyMetadataV1 {
176 pub encryption_key: serde_bytes::ByteBuf,
177 pub aad_prefix: Option<serde_bytes::ByteBuf>,
178 pub file_length: Option<i64>,
179 }
180
181 impl StandardKeyMetadataV1 {
182 pub(super) fn encode(&self) -> Result<Box<[u8]>> {
183 let value = to_value(self)
184 .and_then(|v| v.resolve(&AVRO_SCHEMA_V1))
185 .map_err(|e| {
186 Error::new(ErrorKind::Unexpected, "Failed to encode key metadata")
187 .with_source(e)
188 })?;
189
190 let datum = GenericDatumWriter::builder(&AVRO_SCHEMA_V1)
191 .build()
192 .and_then(|writer| writer.write_value_to_vec(value))
193 .map_err(|e| {
194 Error::new(ErrorKind::Unexpected, "Failed to encode key metadata")
195 .with_source(e)
196 })?;
197
198 let mut result = Vec::with_capacity(1 + datum.len());
199 result.push(V1);
200 result.extend_from_slice(&datum);
201 Ok(result.into_boxed_slice())
202 }
203
204 pub(super) fn decode(bytes: &[u8]) -> Result<Self> {
205 if bytes.is_empty() {
206 return Err(invalid_data!("Empty key metadata buffer"));
207 }
208
209 let version = bytes[0];
210 if version != V1 {
211 return Err(Error::new(
212 ErrorKind::FeatureUnsupported,
213 format!("Unsupported key metadata version: {version} (supported: {V1})"),
214 ));
215 }
216
217 let mut reader = Cursor::new(&bytes[1..]);
218 let value = GenericDatumReader::builder(&AVRO_SCHEMA_V1)
219 .build()
220 .and_then(|datum_reader| datum_reader.read_value(&mut reader))
221 .map_err(|e| invalid_data!("Failed to decode key metadata").with_source(e))?;
222
223 from_value(&value)
224 .map_err(|e| invalid_data!("Failed to decode key metadata fields").with_source(e))
225 }
226 }
227
228 impl TryFrom<&StandardKeyMetadata> for StandardKeyMetadataV1 {
229 type Error = Error;
230
231 fn try_from(metadata: &StandardKeyMetadata) -> Result<Self> {
232 let file_length = metadata
233 .file_length
234 .map(i64::try_from)
235 .transpose()
236 .map_err(|e| {
237 Error::new(
238 ErrorKind::DataInvalid,
239 "Key metadata file length exceeds the Avro long range",
240 )
241 .with_source(e)
242 })?;
243 Ok(Self {
244 encryption_key: serde_bytes::ByteBuf::from(metadata.encryption_key.as_bytes()),
245 aad_prefix: metadata
246 .aad_prefix
247 .as_ref()
248 .map(|b| serde_bytes::ByteBuf::from(b.as_ref())),
249 file_length,
250 })
251 }
252 }
253
254 impl TryFrom<StandardKeyMetadataV1> for StandardKeyMetadata {
255 type Error = Error;
256
257 fn try_from(v1: StandardKeyMetadataV1) -> Result<Self> {
258 let encryption_key = SecureKey::new(&v1.encryption_key).map_err(|e| {
259 invalid_data!("Invalid encryption key in key metadata").with_source(e)
260 })?;
261 Ok(Self {
262 encryption_key,
263 aad_prefix: v1.aad_prefix.map(|b| b.into_vec().into_boxed_slice()),
264 file_length: v1.file_length.map(u64::try_from).transpose().map_err(|e| {
265 invalid_data!("Negative file length in key metadata").with_source(e)
266 })?,
267 })
268 }
269 }
270}
271
272#[cfg(test)]
273mod tests {
274 use super::*;
275
276 #[test]
277 fn test_roundtrip() {
278 let key = b"0123456789012345";
279 let aad = b"1234567890123456";
280
281 let metadata = StandardKeyMetadata::try_new(key)
282 .unwrap()
283 .with_aad_prefix(aad);
284 let serialized = metadata.encode().unwrap();
285 let parsed = StandardKeyMetadata::decode(&serialized).unwrap();
286
287 assert_eq!(parsed.encryption_key().as_bytes(), key);
288 assert_eq!(parsed.aad_prefix(), Some(aad.as_slice()));
289 assert_eq!(parsed.file_length(), None);
290 }
291
292 #[test]
293 fn test_roundtrip_with_length() {
294 let key = b"0123456789012345";
295 let aad = b"1234567890123456";
296
297 let file_length = 100_000;
298 let metadata = StandardKeyMetadata::try_new(key)
299 .unwrap()
300 .with_aad_prefix(aad)
301 .with_file_length(file_length);
302 let serialized = metadata.encode().unwrap();
303 let parsed = StandardKeyMetadata::decode(&serialized).unwrap();
304
305 assert_eq!(parsed.encryption_key().as_bytes(), key);
306 assert_eq!(parsed.aad_prefix(), Some(aad.as_slice()));
307 assert_eq!(parsed.file_length(), Some(file_length));
308 }
309
310 #[test]
311 fn test_unsupported_version() {
312 let result = StandardKeyMetadata::decode(&[0x02]);
313 assert!(result.is_err());
314 let err = result.unwrap_err();
315 assert_eq!(err.kind(), ErrorKind::FeatureUnsupported);
316 assert_eq!(
317 err.message(),
318 "Unsupported key metadata version: 2 (supported: 1)"
319 );
320 }
321
322 #[test]
323 fn test_empty_buffer() {
324 let result = StandardKeyMetadata::decode(&[]);
325 assert!(result.is_err());
326 assert_eq!(result.unwrap_err().kind(), ErrorKind::DataInvalid);
327 }
328
329 #[test]
330 fn test_roundtrip_without_aad() {
331 let key = b"0123456789012345";
332 let metadata = StandardKeyMetadata::try_new(key).unwrap();
333 let serialized = metadata.encode().unwrap();
334 let parsed = StandardKeyMetadata::decode(&serialized).unwrap();
335
336 assert_eq!(parsed.encryption_key().as_bytes(), key);
337 assert_eq!(parsed.aad_prefix(), None);
338 }
339
340 #[test]
341 fn test_new_rejects_invalid_key_length() {
342 for len in [16usize, 24, 32] {
344 assert!(StandardKeyMetadata::try_new(&vec![0u8; len]).is_ok());
345 }
346
347 for len in [0usize, 4, 15, 20, 33] {
350 assert!(StandardKeyMetadata::try_new(&vec![0u8; len]).is_err());
351 }
352 }
353
354 #[test]
355 fn test_encode_rejects_file_length_above_long_max() {
356 let metadata = StandardKeyMetadata::try_new(&[0u8; 16])
357 .unwrap()
358 .with_file_length(i64::MAX as u64 + 1);
359
360 let err = metadata.encode().unwrap_err();
361 assert_eq!(err.kind(), ErrorKind::DataInvalid);
362 }
363
364 #[test]
365 fn test_decode_rejects_negative_file_length() {
366 let serialized = _serde::StandardKeyMetadataV1 {
367 encryption_key: serde_bytes::ByteBuf::from(vec![0u8; 16]),
368 aad_prefix: None,
369 file_length: Some(-1),
370 }
371 .encode()
372 .unwrap();
373
374 let err = StandardKeyMetadata::decode(&serialized).unwrap_err();
375 assert_eq!(err.kind(), ErrorKind::DataInvalid);
376 }
377
378 #[test]
379 fn test_decode_rejects_invalid_key_length() {
380 for len in [0usize, 4, 15, 20, 33] {
384 let serialized = _serde::StandardKeyMetadataV1 {
385 encryption_key: serde_bytes::ByteBuf::from(vec![0u8; len]),
386 aad_prefix: None,
387 file_length: None,
388 }
389 .encode()
390 .unwrap();
391
392 let err = StandardKeyMetadata::decode(&serialized).unwrap_err();
393 assert_eq!(err.kind(), ErrorKind::DataInvalid);
394 assert!(
395 err.to_string()
396 .contains("Invalid encryption key in key metadata")
397 );
398 }
399 }
400
401 #[test]
402 fn test_decode_tolerates_trailing_bytes() {
403 let key = b"0123456789012345";
408 let aad = b"1234567890123456";
409 let file_length = 1024;
410
411 let serialized = StandardKeyMetadata::try_new(key)
412 .unwrap()
413 .with_aad_prefix(aad)
414 .with_file_length(file_length)
415 .encode()
416 .unwrap();
417
418 for trailing in [b"\xde\xad\xbe\xef".as_slice(), b"\x02\x08more".as_slice()] {
421 let mut extended = serialized.to_vec();
422 extended.extend_from_slice(trailing);
423
424 let parsed = StandardKeyMetadata::decode(&extended).unwrap();
425
426 assert_eq!(parsed.encryption_key().as_bytes(), key);
427 assert_eq!(parsed.aad_prefix(), Some(aad.as_slice()));
428 assert_eq!(parsed.file_length(), Some(file_length));
429 }
430 }
431}