1use crate::{
2 errors::{FileError, HeaderError},
3 memory::{SecureKey, SecureString},
4 v1::{crypt, key::KeyDerivationParams},
5};
6
7#[derive(Debug, Clone)]
30pub struct FileHeader {
31 salt: [u8; 16],
32 kdf_params: KeyDerivationParams,
33 content_nonce: [u8; 24],
34 filename_nonce: [u8; 24],
35 filename_ciphertext: Vec<u8>,
36}
37
38const MAGIC: [u8; 6] = *b"SHADOW";
39const VERSION: u8 = 1;
40
41impl FileHeader {
42 pub fn new(
43 salt: [u8; 16],
44 kdf_params: KeyDerivationParams,
45 content_nonce: [u8; 24],
46 filename_nonce: [u8; 24],
47 filename_ciphertext: Vec<u8>,
48 ) -> Self {
49 FileHeader {
50 salt,
51 kdf_params,
52 content_nonce,
53 filename_nonce,
54 filename_ciphertext,
55 }
56 }
57
58 pub(crate) const fn min_length() -> usize {
61 6 + 1 + 4 + 16 + 4 + 4 + 4 + 1 + 24 + 24 + 2 }
73
74 pub fn header_length(&self) -> usize {
76 Self::min_length() + self.filename_ciphertext.len()
77 }
78
79 pub fn serialize(&self) -> Vec<u8> {
80 let mut bytes = Vec::with_capacity(self.header_length());
81
82 bytes.extend_from_slice(&MAGIC);
83 bytes.push(VERSION);
84 bytes.extend_from_slice(&(self.header_length() as u32).to_le_bytes());
85 bytes.extend_from_slice(&self.salt);
86 bytes.extend_from_slice(&self.kdf_params.memory_cost.to_le_bytes());
87 bytes.extend_from_slice(&self.kdf_params.time_cost.to_le_bytes());
88 bytes.extend_from_slice(&self.kdf_params.parallelism.to_le_bytes());
89 bytes.push(self.kdf_params.key_size);
90 bytes.extend_from_slice(&self.content_nonce);
91 bytes.extend_from_slice(&self.filename_nonce);
92 bytes.extend_from_slice(&(self.filename_ciphertext.len() as u16).to_le_bytes());
93 bytes.extend_from_slice(&self.filename_ciphertext);
94
95 bytes
96 }
97
98 pub fn try_deserialize(bytes: &[u8]) -> Result<FileHeader, HeaderError> {
99 if bytes.len() < FileHeader::min_length() {
100 return Err(HeaderError::InsufficientBytes);
101 }
102
103 let length = read_header_length(bytes)?;
104
105 if bytes.len() < length as usize {
106 return Err(HeaderError::InsufficientBytes);
107 }
108
109 match Self::deserialize(bytes) {
110 Some(header) => Ok(header),
111 None => Err(HeaderError::InvalidData),
112 }
113 }
114
115 fn deserialize(bytes: &[u8]) -> Option<FileHeader> {
116 if bytes.len() < FileHeader::min_length() {
117 return None;
118 }
119 let header_length = u32::from_le_bytes(bytes[7..11].try_into().ok()?);
122 let salt = bytes[11..27].try_into().ok()?;
123 let kdf_memory = u32::from_le_bytes(bytes[27..31].try_into().ok()?);
124 let kdf_iterations = u32::from_le_bytes(bytes[31..35].try_into().ok()?);
125 let kdf_parallelism = u32::from_le_bytes(bytes[35..39].try_into().ok()?);
126 let kdf_key_length = bytes[39];
127 let content_nonce = bytes[40..64].try_into().ok()?;
128 let filename_nonce = bytes[64..88].try_into().ok()?;
129 let filename_ciphertext_length = u16::from_le_bytes(bytes[88..90].try_into().ok()?);
130
131 let expected_length: usize = FileHeader::min_length() + filename_ciphertext_length as usize;
132
133 if header_length != expected_length as u32 {
134 return None;
135 }
136
137 if bytes.len() < expected_length {
138 return None;
139 }
140
141 let filename_ciphertext = bytes[FileHeader::min_length()..expected_length].to_vec();
142
143 Some(FileHeader {
144 salt,
145 kdf_params: KeyDerivationParams::new(
146 kdf_memory,
147 kdf_iterations,
148 kdf_parallelism,
149 kdf_key_length,
150 ),
151 content_nonce,
152 filename_nonce,
153 filename_ciphertext,
154 })
155 }
156
157 pub fn salt(&self) -> &[u8; 16] {
158 &self.salt
159 }
160
161 pub fn kdf_params(&self) -> &KeyDerivationParams {
163 &self.kdf_params
164 }
165
166 pub fn content_nonce(&self) -> &[u8; 24] {
167 &self.content_nonce
168 }
169
170 pub fn filename_nonce(&self) -> &[u8; 24] {
171 &self.filename_nonce
172 }
173
174 pub fn filename_ciphertext(&self) -> &[u8] {
175 &self.filename_ciphertext
176 }
177
178 pub fn decrypt_content(
180 &self,
181 ciphertext: &[u8],
182 key: &SecureKey,
183 ) -> Result<crate::memory::SecureBytes, FileError> {
184 let (content, _) = crypt::decrypt_bytes(ciphertext, key.as_bytes(), &self.content_nonce)?;
185 Ok(content)
186 }
187
188 pub fn decrypt_filename(&self, key: &SecureKey) -> Result<SecureString, FileError> {
190 let (filename_bytes, _) = crypt::decrypt_bytes(
191 &self.filename_ciphertext,
192 key.as_bytes(),
193 &self.filename_nonce,
194 )?;
195 let filename = String::from_utf8(filename_bytes.as_slice().to_vec())
196 .map_err(|_| FileError::InvalidFilename)?;
197 Ok(SecureString::new(filename))
198 }
199}
200
201fn read_header_length(bytes: &[u8]) -> Result<u32, HeaderError> {
203 if bytes.len() < 11 {
204 return Err(HeaderError::InsufficientBytes);
205 }
206 let length_bytes = &bytes[7..11];
207 let length = u32::from_le_bytes(
208 length_bytes
209 .try_into()
210 .map_err(|_| HeaderError::InvalidData)?,
211 );
212 Ok(length)
213}
214
215#[cfg(test)]
216mod tests {
217 use crate::profile;
218
219 use super::*;
220
221 fn get_test_params() -> KeyDerivationParams {
222 let profile = profile::SecurityProfile::Test;
223 KeyDerivationParams::from(profile)
224 }
225
226 fn create_test_header() -> FileHeader {
227 FileHeader::new(
228 [1u8; 16],
229 get_test_params(),
230 [2u8; 24],
231 [3u8; 24],
232 vec![4, 5, 6, 7, 8],
233 )
234 }
235
236 #[test]
237 fn serialized_magic_and_version_are_correct() {
238 let serialized = create_test_header().serialize();
239 assert_eq!(&serialized[0..6], b"SHADOW");
240 assert_eq!(serialized[6], 1);
241 }
242
243 #[test]
244 fn header_size_is_calculated_correctly() {
245 let filename_ciphertext = vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10];
246 let header = FileHeader::new(
247 [0u8; 16],
248 get_test_params(),
249 [0u8; 24],
250 [0u8; 24],
251 filename_ciphertext.clone(),
252 );
253
254 assert_eq!(header.header_length(), 90 + filename_ciphertext.len());
255 assert_eq!(header.serialize().len(), header.header_length());
256 }
257
258 #[test]
259 fn kdf_params_round_trip_through_header() {
260 let params = get_test_params();
261 let header = FileHeader::new(
262 [0u8; 16],
263 params.clone(),
264 [0u8; 24],
265 [0u8; 24],
266 vec![1, 2, 3],
267 );
268 assert_eq!(header.kdf_params(), ¶ms);
269 }
270
271 #[test]
272 fn test_serialize_field_offsets() {
273 let header = create_test_header();
274 let serialized = header.serialize();
275 let params = header.kdf_params();
276
277 assert_eq!(serialized.len(), header.header_length());
278 assert_eq!(&serialized[0..6], b"SHADOW");
279 assert_eq!(serialized[6], 1);
280 assert_eq!(
281 u32::from_le_bytes(serialized[7..11].try_into().unwrap()) as usize,
282 header.header_length()
283 );
284 assert_eq!(&serialized[11..27], header.salt());
285 assert_eq!(
286 u32::from_le_bytes(serialized[27..31].try_into().unwrap()),
287 params.memory_cost
288 );
289 assert_eq!(
290 u32::from_le_bytes(serialized[31..35].try_into().unwrap()),
291 params.time_cost
292 );
293 assert_eq!(
294 u32::from_le_bytes(serialized[35..39].try_into().unwrap()),
295 params.parallelism
296 );
297 assert_eq!(serialized[39], params.key_size);
298 assert_eq!(&serialized[40..64], header.content_nonce());
299 assert_eq!(&serialized[64..88], header.filename_nonce());
300 assert_eq!(
301 u16::from_le_bytes(serialized[88..90].try_into().unwrap()) as usize,
302 header.filename_ciphertext().len()
303 );
304 assert_eq!(
305 &serialized[FileHeader::min_length()..],
306 header.filename_ciphertext()
307 );
308 }
309
310 #[test]
311 fn test_round_trip_serialization() {
312 let original = create_test_header();
313 let serialized = original.serialize();
314
315 let deserialized = FileHeader::try_deserialize(&serialized).unwrap();
316 assert_eq!(deserialized.salt(), original.salt());
317 assert_eq!(deserialized.kdf_params(), original.kdf_params());
318 assert_eq!(deserialized.content_nonce(), original.content_nonce());
319 assert_eq!(deserialized.filename_nonce(), original.filename_nonce());
320 assert_eq!(
321 deserialized.filename_ciphertext(),
322 original.filename_ciphertext()
323 );
324 }
325
326 #[test]
327 fn test_try_deserialize_insufficient_bytes() {
328 let bytes = vec![0u8; 50];
329 assert!(matches!(
330 FileHeader::try_deserialize(&bytes),
331 Err(HeaderError::InsufficientBytes)
332 ));
333 }
334
335 #[test]
336 fn test_try_deserialize_invalid_data() {
337 let mut bytes = vec![0u8; 100];
338 bytes[7..11].copy_from_slice(&(50u32.to_le_bytes()));
340
341 let result = FileHeader::try_deserialize(&bytes);
342 assert!(matches!(result.unwrap_err(), HeaderError::InvalidData));
343 }
344
345 #[test]
346 fn test_try_deserialize_inconsistent_lengths() {
347 let mut serialized = create_test_header().serialize();
348 serialized[7..11].copy_from_slice(&(200u32.to_le_bytes()));
350
351 assert!(FileHeader::try_deserialize(&serialized).is_err());
352 }
353
354 #[test]
355 fn test_try_deserialize_insufficient_bytes_for_filename() {
356 let mut bytes = vec![0u8; 95];
359 bytes[0..6].copy_from_slice(b"SHADOW");
360 bytes[6] = 1; bytes[7..11].copy_from_slice(&(100u32.to_le_bytes()));
362 bytes[88..90].copy_from_slice(&(10u16.to_le_bytes()));
363
364 let result = FileHeader::try_deserialize(&bytes);
365 assert!(matches!(
366 result.unwrap_err(),
367 HeaderError::InsufficientBytes
368 ));
369 }
370
371 #[test]
372 fn test_empty_filename_ciphertext_round_trip() {
373 let header = FileHeader::new([1u8; 16], get_test_params(), [2u8; 24], [3u8; 24], vec![]);
374
375 let deserialized = FileHeader::try_deserialize(&header.serialize()).unwrap();
376 assert!(deserialized.filename_ciphertext().is_empty());
377 }
378
379 #[test]
380 fn test_large_filename_ciphertext_round_trip() {
381 let header = FileHeader::new(
382 [1u8; 16],
383 get_test_params(),
384 [2u8; 24],
385 [3u8; 24],
386 vec![4u8; 1000],
387 );
388
389 let deserialized = FileHeader::try_deserialize(&header.serialize()).unwrap();
390 assert_eq!(deserialized.filename_ciphertext(), &[4u8; 1000][..]);
391 }
392}