1use crate::{
2 errors::{FileError, HeaderError},
3 memory::{SecureKey, SecureString},
4 v2::{crypt, key::KeyDerivationParams},
5};
6
7#[derive(Debug, Clone)]
34pub struct FileHeader {
35 salt: [u8; 16],
36 kdf_params: KeyDerivationParams,
37 content_nonce: [u8; 24],
38 filename_nonce: [u8; 24],
39 filename_ciphertext: Vec<u8>,
40}
41
42pub const MAGIC: [u8; 6] = *b"SHADOW";
43pub const VERSION: u8 = 2;
44
45impl FileHeader {
46 pub fn new(
50 salt: [u8; 16],
51 kdf_params: KeyDerivationParams,
52 content_nonce: [u8; 24],
53 filename_nonce: [u8; 24],
54 filename_ciphertext: Vec<u8>,
55 ) -> Result<Self, HeaderError> {
56 if u16::try_from(filename_ciphertext.len()).is_err() {
57 return Err(HeaderError::FilenameTooLong);
58 }
59
60 Ok(FileHeader {
61 salt,
62 kdf_params,
63 content_nonce,
64 filename_nonce,
65 filename_ciphertext,
66 })
67 }
68
69 pub(crate) const fn min_length() -> usize {
72 6 + 1 + 4 + 16 + 4 + 4 + 4 + 1 + 24 + 24 + 2 }
84
85 pub fn header_length(&self) -> usize {
87 Self::min_length() + self.filename_ciphertext.len()
88 }
89
90 pub fn serialize(&self) -> Vec<u8> {
91 let mut bytes = Vec::with_capacity(self.header_length());
92
93 bytes.extend_from_slice(&MAGIC);
94 bytes.push(VERSION);
95 bytes.extend_from_slice(&(self.header_length() as u32).to_le_bytes());
96 bytes.extend_from_slice(&self.salt);
97 bytes.extend_from_slice(&self.kdf_params.memory_cost.to_le_bytes());
98 bytes.extend_from_slice(&self.kdf_params.time_cost.to_le_bytes());
99 bytes.extend_from_slice(&self.kdf_params.parallelism.to_le_bytes());
100 bytes.push(self.kdf_params.key_size);
101 bytes.extend_from_slice(&self.content_nonce);
102 bytes.extend_from_slice(&self.filename_nonce);
103 bytes.extend_from_slice(&(self.filename_ciphertext.len() as u16).to_le_bytes());
104 bytes.extend_from_slice(&self.filename_ciphertext);
105
106 bytes
107 }
108
109 pub fn try_deserialize(bytes: &[u8]) -> Result<FileHeader, HeaderError> {
110 if bytes.len() < FileHeader::min_length() {
111 return Err(HeaderError::InsufficientBytes);
112 }
113
114 let length = read_header_length(bytes)?;
115
116 if bytes.len() < length as usize {
117 return Err(HeaderError::InsufficientBytes);
118 }
119
120 match Self::deserialize(bytes) {
121 Some(header) => Ok(header),
122 None => Err(HeaderError::InvalidData),
123 }
124 }
125
126 fn deserialize(bytes: &[u8]) -> Option<FileHeader> {
127 if bytes.len() < FileHeader::min_length() {
128 return None;
129 }
130 let magic: [u8; 6] = bytes[0..6].try_into().ok()?;
131 let version = bytes[6];
132
133 if magic != MAGIC || version != VERSION {
135 return None;
136 }
137
138 let header_length = u32::from_le_bytes(bytes[7..11].try_into().ok()?);
139 let salt = bytes[11..27].try_into().ok()?;
140 let kdf_memory = u32::from_le_bytes(bytes[27..31].try_into().ok()?);
141 let kdf_iterations = u32::from_le_bytes(bytes[31..35].try_into().ok()?);
142 let kdf_parallelism = u32::from_le_bytes(bytes[35..39].try_into().ok()?);
143 let kdf_key_length = bytes[39];
144 let content_nonce = bytes[40..64].try_into().ok()?;
145 let filename_nonce = bytes[64..88].try_into().ok()?;
146 let filename_ciphertext_length = u16::from_le_bytes(bytes[88..90].try_into().ok()?);
147
148 let expected_length: usize = FileHeader::min_length() + filename_ciphertext_length as usize;
149
150 if header_length != expected_length as u32 {
151 return None;
152 }
153
154 if bytes.len() < expected_length {
155 return None;
156 }
157
158 let filename_ciphertext = bytes[FileHeader::min_length()..expected_length].to_vec();
159
160 Some(FileHeader {
161 salt,
162 kdf_params: KeyDerivationParams::new(
163 kdf_memory,
164 kdf_iterations,
165 kdf_parallelism,
166 kdf_key_length,
167 ),
168 content_nonce,
169 filename_nonce,
170 filename_ciphertext,
171 })
172 }
173
174 pub fn salt(&self) -> &[u8; 16] {
175 &self.salt
176 }
177
178 pub fn kdf_params(&self) -> &KeyDerivationParams {
180 &self.kdf_params
181 }
182
183 pub fn content_nonce(&self) -> &[u8; 24] {
184 &self.content_nonce
185 }
186
187 pub fn filename_nonce(&self) -> &[u8; 24] {
188 &self.filename_nonce
189 }
190
191 pub fn filename_ciphertext(&self) -> &[u8] {
192 &self.filename_ciphertext
193 }
194
195 pub fn decrypt_content(
198 &self,
199 ciphertext: &[u8],
200 key: &SecureKey,
201 ) -> Result<crate::memory::SecureBytes, FileError> {
202 let (content, _) = crypt::decrypt_bytes(
203 ciphertext,
204 key.as_bytes(),
205 &self.content_nonce,
206 &self.binding().aad(AadPurpose::Content),
207 )?;
208 Ok(content)
209 }
210
211 pub fn decrypt_filename(&self, key: &SecureKey) -> Result<SecureString, FileError> {
214 let (filename_bytes, _) = crypt::decrypt_bytes(
215 &self.filename_ciphertext,
216 key.as_bytes(),
217 &self.filename_nonce,
218 &self.binding().aad(AadPurpose::Filename),
219 )?;
220 let filename = String::from_utf8(filename_bytes.as_slice().to_vec())
221 .map_err(|_| FileError::InvalidFilename)?;
222 Ok(SecureString::new(filename))
223 }
224
225 pub fn binding(&self) -> HeaderBinding<'_> {
228 HeaderBinding {
229 salt: &self.salt,
230 kdf_params: &self.kdf_params,
231 content_nonce: &self.content_nonce,
232 filename_nonce: &self.filename_nonce,
233 }
234 }
235}
236
237fn read_header_length(bytes: &[u8]) -> Result<u32, HeaderError> {
239 if bytes.len() < 11 {
240 return Err(HeaderError::InsufficientBytes);
241 }
242 let length_bytes = &bytes[7..11];
243 let length = u32::from_le_bytes(
244 length_bytes
245 .try_into()
246 .map_err(|_| HeaderError::InvalidData)?,
247 );
248 Ok(length)
249}
250
251#[derive(Debug, Clone, Copy, PartialEq, Eq)]
257pub enum AadPurpose {
258 Filename,
259 Content,
260}
261
262impl AadPurpose {
263 fn domain_tag(self) -> &'static [u8] {
264 match self {
265 AadPurpose::Filename => b"shadow-crypt/v2/filename",
266 AadPurpose::Content => b"shadow-crypt/v2/content",
267 }
268 }
269}
270
271#[derive(Debug, Clone, Copy)]
285pub struct HeaderBinding<'a> {
286 salt: &'a [u8; 16],
287 kdf_params: &'a KeyDerivationParams,
288 content_nonce: &'a [u8; 24],
289 filename_nonce: &'a [u8; 24],
290}
291
292impl<'a> HeaderBinding<'a> {
293 pub fn new(
294 salt: &'a [u8; 16],
295 kdf_params: &'a KeyDerivationParams,
296 content_nonce: &'a [u8; 24],
297 filename_nonce: &'a [u8; 24],
298 ) -> Self {
299 Self {
300 salt,
301 kdf_params,
302 content_nonce,
303 filename_nonce,
304 }
305 }
306
307 pub fn aad(&self, purpose: AadPurpose) -> Vec<u8> {
310 let mut aad = Vec::with_capacity(FileHeader::min_length() + 24);
311 aad.extend_from_slice(&MAGIC);
312 aad.push(VERSION);
313 aad.extend_from_slice(self.salt);
314 aad.extend_from_slice(&self.kdf_params.memory_cost.to_le_bytes());
315 aad.extend_from_slice(&self.kdf_params.time_cost.to_le_bytes());
316 aad.extend_from_slice(&self.kdf_params.parallelism.to_le_bytes());
317 aad.push(self.kdf_params.key_size);
318 aad.extend_from_slice(self.content_nonce);
319 aad.extend_from_slice(self.filename_nonce);
320 aad.extend_from_slice(purpose.domain_tag());
321 aad
322 }
323}
324
325#[cfg(test)]
326mod tests {
327 use super::*;
328 use crate::profile;
329
330 fn get_test_params() -> KeyDerivationParams {
331 KeyDerivationParams::from(profile::SecurityProfile::Test)
332 }
333
334 fn create_test_header() -> FileHeader {
335 FileHeader::new(
336 [1u8; 16],
337 get_test_params(),
338 [2u8; 24],
339 [3u8; 24],
340 vec![4, 5, 6, 7, 8],
341 )
342 .unwrap()
343 }
344
345 #[test]
346 fn serialized_magic_and_version_are_correct() {
347 let serialized = create_test_header().serialize();
348 assert_eq!(&serialized[0..6], b"SHADOW");
349 assert_eq!(serialized[6], 2);
350 }
351
352 #[test]
353 fn header_size_is_calculated_correctly() {
354 let filename_ciphertext = vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10];
355 let header = FileHeader::new(
356 [0u8; 16],
357 get_test_params(),
358 [0u8; 24],
359 [0u8; 24],
360 filename_ciphertext.clone(),
361 )
362 .unwrap();
363
364 assert_eq!(header.header_length(), 90 + filename_ciphertext.len());
365 assert_eq!(header.serialize().len(), header.header_length());
366 }
367
368 #[test]
369 fn oversized_filename_ciphertext_is_rejected() {
370 let filename_ciphertext = vec![0u8; u16::MAX as usize + 1];
371 let result = FileHeader::new(
372 [0u8; 16],
373 get_test_params(),
374 [0u8; 24],
375 [0u8; 24],
376 filename_ciphertext,
377 );
378 assert!(matches!(result, Err(HeaderError::FilenameTooLong)));
379 }
380
381 #[test]
382 fn max_length_filename_ciphertext_is_accepted() {
383 let filename_ciphertext = vec![0u8; u16::MAX as usize];
384 let header = FileHeader::new(
385 [0u8; 16],
386 get_test_params(),
387 [0u8; 24],
388 [0u8; 24],
389 filename_ciphertext,
390 )
391 .unwrap();
392 assert_eq!(header.filename_ciphertext().len(), u16::MAX as usize);
393 }
394
395 #[test]
396 fn aad_differs_by_purpose() {
397 let salt = [1u8; 16];
398 let params = get_test_params();
399 let content_nonce = [2u8; 24];
400 let filename_nonce = [3u8; 24];
401 let binding = HeaderBinding::new(&salt, ¶ms, &content_nonce, &filename_nonce);
402
403 assert_ne!(
404 binding.aad(AadPurpose::Filename),
405 binding.aad(AadPurpose::Content)
406 );
407 }
408
409 #[test]
410 fn header_binding_matches_standalone_binding() {
411 let salt = [1u8; 16];
412 let params = get_test_params();
413 let content_nonce = [2u8; 24];
414 let filename_nonce = [3u8; 24];
415
416 let standalone = HeaderBinding::new(&salt, ¶ms, &content_nonce, &filename_nonce);
417 let header = FileHeader::new(
418 salt,
419 params.clone(),
420 content_nonce,
421 filename_nonce,
422 vec![1, 2, 3],
423 )
424 .unwrap();
425
426 assert_eq!(
427 standalone.aad(AadPurpose::Content),
428 header.binding().aad(AadPurpose::Content)
429 );
430 assert_eq!(
431 standalone.aad(AadPurpose::Filename),
432 header.binding().aad(AadPurpose::Filename)
433 );
434 }
435
436 #[test]
437 fn kdf_params_round_trip_through_header() {
438 let params = get_test_params();
439 let header = FileHeader::new(
440 [0u8; 16],
441 params.clone(),
442 [0u8; 24],
443 [0u8; 24],
444 vec![1, 2, 3],
445 )
446 .unwrap();
447 assert_eq!(header.kdf_params(), ¶ms);
448 }
449
450 #[test]
451 fn test_round_trip_serialization() {
452 let original = create_test_header();
453 let serialized = original.serialize();
454 assert_eq!(serialized.len(), original.header_length());
455
456 let deserialized = FileHeader::try_deserialize(&serialized).unwrap();
457 assert_eq!(deserialized.salt(), original.salt());
458 assert_eq!(deserialized.kdf_params(), original.kdf_params());
459 assert_eq!(deserialized.content_nonce(), original.content_nonce());
460 assert_eq!(deserialized.filename_nonce(), original.filename_nonce());
461 assert_eq!(
462 deserialized.filename_ciphertext(),
463 original.filename_ciphertext()
464 );
465 }
466
467 #[test]
468 fn test_try_deserialize_rejects_wrong_version() {
469 let mut serialized = create_test_header().serialize();
470 serialized[6] = 1; assert!(FileHeader::try_deserialize(&serialized).is_err());
473 }
474
475 #[test]
476 fn test_try_deserialize_rejects_wrong_magic() {
477 let mut serialized = create_test_header().serialize();
478 serialized[0..6].copy_from_slice(b"NOTSHD");
479
480 assert!(FileHeader::try_deserialize(&serialized).is_err());
481 }
482
483 #[test]
484 fn test_try_deserialize_insufficient_bytes() {
485 let bytes = vec![0u8; 50];
486 assert!(matches!(
487 FileHeader::try_deserialize(&bytes),
488 Err(HeaderError::InsufficientBytes)
489 ));
490 }
491
492 #[test]
493 fn test_try_deserialize_inconsistent_lengths() {
494 let mut serialized = create_test_header().serialize();
495 serialized[7..11].copy_from_slice(&(200u32.to_le_bytes()));
497
498 assert!(FileHeader::try_deserialize(&serialized).is_err());
499 }
500
501 #[test]
502 fn test_empty_filename_ciphertext_round_trip() {
503 let header = FileHeader::new(
504 [1u8; 16],
505 KeyDerivationParams::from(profile::SecurityProfile::Test),
506 [2u8; 24],
507 [3u8; 24],
508 vec![],
509 )
510 .unwrap();
511
512 let deserialized = FileHeader::try_deserialize(&header.serialize()).unwrap();
513 assert!(deserialized.filename_ciphertext().is_empty());
514 }
515}