1use crate::{errors::HeaderError, v1::key::KeyDerivationParams};
2
3use super::header::FileHeader;
4
5pub fn serialize(header: &FileHeader) -> Vec<u8> {
6 let mut bytes = Vec::new();
7
8 bytes.extend_from_slice(header.magic.as_slice());
9 bytes.push(header.version);
10 bytes.extend_from_slice(header.header_length.to_le_bytes().as_slice());
11 bytes.extend_from_slice(header.salt.as_slice());
12 bytes.extend_from_slice(header.kdf_memory.to_le_bytes().as_slice());
13 bytes.extend_from_slice(header.kdf_iterations.to_le_bytes().as_slice());
14 bytes.extend_from_slice(header.kdf_parallelism.to_le_bytes().as_slice());
15 bytes.push(header.kdf_key_length);
16 bytes.extend_from_slice(header.content_nonce.as_slice());
17 bytes.extend_from_slice(header.filename_nonce.as_slice());
18 bytes.extend_from_slice(header.filename_ciphertext_length.to_le_bytes().as_slice());
19 bytes.extend_from_slice(header.filename_ciphertext.as_slice());
20
21 bytes
22}
23
24pub fn get_length_from_bytes(bytes: &[u8]) -> Result<u32, HeaderError> {
25 if bytes.len() < 11 {
26 return Err(HeaderError::InsufficientBytes);
27 }
28 let length_bytes = &bytes[7..11];
29 let length = u32::from_le_bytes(
30 length_bytes
31 .try_into()
32 .map_err(|_| HeaderError::InvalidData)?,
33 );
34 Ok(length)
35}
36
37pub fn get_kdf_params(header: &FileHeader) -> KeyDerivationParams {
38 KeyDerivationParams {
39 memory_cost: header.kdf_memory,
40 time_cost: header.kdf_iterations,
41 parallelism: header.kdf_parallelism,
42 key_size: header.kdf_key_length,
43 }
44}
45
46pub fn try_deserialize(bytes: &[u8]) -> Result<FileHeader, HeaderError> {
47 if bytes.len() < FileHeader::min_length() {
48 return Err(HeaderError::InsufficientBytes);
49 }
50
51 let length: u32 = get_length_from_bytes(bytes)?;
52
53 if bytes.len() < length as usize {
54 return Err(HeaderError::InsufficientBytes);
55 }
56
57 match deserialize(bytes) {
58 Some(header) => Ok(header),
59 None => Err(HeaderError::InvalidData),
60 }
61}
62
63fn deserialize(bytes: &[u8]) -> Option<FileHeader> {
64 if bytes.len() < FileHeader::min_length() {
65 return None;
66 }
67 let magic = bytes[0..6].try_into().ok()?;
68 let version = bytes[6];
69 let header_length = u32::from_le_bytes(bytes[7..11].try_into().ok()?);
70 let salt = bytes[11..27].try_into().ok()?;
71 let kdf_memory = u32::from_le_bytes(bytes[27..31].try_into().ok()?);
72 let kdf_iterations = u32::from_le_bytes(bytes[31..35].try_into().ok()?);
73 let kdf_parallelism = u32::from_le_bytes(bytes[35..39].try_into().ok()?);
74 let kdf_key_length = bytes[39];
75 let content_nonce = bytes[40..64].try_into().ok()?;
76 let filename_nonce = bytes[64..88].try_into().ok()?;
77 let filename_ciphertext_length = u16::from_le_bytes(bytes[88..90].try_into().ok()?);
78
79 let expected_length: usize = FileHeader::min_length() + filename_ciphertext_length as usize;
80
81 if header_length != expected_length as u32 {
82 return None;
83 }
84
85 if bytes.len() < expected_length {
86 return None;
87 }
88
89 let filename_ciphertext = bytes[FileHeader::min_length()
90 ..(FileHeader::min_length() + filename_ciphertext_length as usize)]
91 .to_vec();
92
93 Some(FileHeader {
94 magic,
95 version,
96 header_length,
97 salt,
98 kdf_memory,
99 kdf_iterations,
100 kdf_parallelism,
101 kdf_key_length,
102 content_nonce,
103 filename_nonce,
104 filename_ciphertext_length,
105 filename_ciphertext,
106 })
107}
108
109#[cfg(test)]
110mod tests {
111 use super::*;
112 use crate::profile;
113 use crate::v1::key::KeyDerivationParams;
114
115 fn create_test_header() -> FileHeader {
116 let salt = [1u8; 16];
117 let kdf_params = KeyDerivationParams::from(profile::SecurityProfile::Test);
118 let content_nonce = [2u8; 24];
119 let filename_nonce = [3u8; 24];
120 let filename_ciphertext = vec![4, 5, 6, 7, 8];
121
122 FileHeader::new(
123 salt,
124 kdf_params,
125 content_nonce,
126 filename_nonce,
127 filename_ciphertext,
128 )
129 }
130
131 #[test]
132 fn test_serialize() {
133 let header = create_test_header();
134 let serialized = serialize(&header);
135
136 assert_eq!(serialized.len(), header.header_length as usize);
138
139 assert_eq!(&serialized[0..6], b"SHADOW");
141
142 assert_eq!(serialized[6], 1);
144
145 let header_len_bytes = &serialized[7..11];
147 let header_len = u32::from_le_bytes(header_len_bytes.try_into().unwrap());
148 assert_eq!(header_len, header.header_length);
149
150 assert_eq!(&serialized[11..27], &header.salt);
152
153 let kdf_memory_bytes = &serialized[27..31];
155 let kdf_memory = u32::from_le_bytes(kdf_memory_bytes.try_into().unwrap());
156 assert_eq!(kdf_memory, header.kdf_memory);
157
158 let kdf_iterations_bytes = &serialized[31..35];
159 let kdf_iterations = u32::from_le_bytes(kdf_iterations_bytes.try_into().unwrap());
160 assert_eq!(kdf_iterations, header.kdf_iterations);
161
162 let kdf_parallelism_bytes = &serialized[35..39];
163 let kdf_parallelism = u32::from_le_bytes(kdf_parallelism_bytes.try_into().unwrap());
164 assert_eq!(kdf_parallelism, header.kdf_parallelism);
165
166 assert_eq!(serialized[39], header.kdf_key_length);
168
169 assert_eq!(&serialized[40..64], &header.content_nonce);
171 assert_eq!(&serialized[64..88], &header.filename_nonce);
172
173 let filename_len_bytes = &serialized[88..90];
175 let filename_len = u16::from_le_bytes(filename_len_bytes.try_into().unwrap());
176 assert_eq!(filename_len, header.filename_ciphertext_length);
177
178 let filename_start = FileHeader::min_length();
180 let filename_end = filename_start + header.filename_ciphertext.len();
181 assert_eq!(
182 &serialized[filename_start..filename_end],
183 &header.filename_ciphertext[..]
184 );
185 }
186
187 #[test]
188 fn test_try_deserialize_valid() {
189 let original_header = create_test_header();
190 let serialized = serialize(&original_header);
191
192 let result = try_deserialize(&serialized);
193 assert!(result.is_ok());
194
195 let deserialized_header = result.unwrap();
196 assert_eq!(deserialized_header.magic, original_header.magic);
197 assert_eq!(deserialized_header.version, original_header.version);
198 assert_eq!(
199 deserialized_header.header_length,
200 original_header.header_length
201 );
202 assert_eq!(deserialized_header.salt, original_header.salt);
203 assert_eq!(deserialized_header.kdf_memory, original_header.kdf_memory);
204 assert_eq!(
205 deserialized_header.kdf_iterations,
206 original_header.kdf_iterations
207 );
208 assert_eq!(
209 deserialized_header.kdf_parallelism,
210 original_header.kdf_parallelism
211 );
212 assert_eq!(
213 deserialized_header.kdf_key_length,
214 original_header.kdf_key_length
215 );
216 assert_eq!(
217 deserialized_header.content_nonce,
218 original_header.content_nonce
219 );
220 assert_eq!(
221 deserialized_header.filename_nonce,
222 original_header.filename_nonce
223 );
224 assert_eq!(
225 deserialized_header.filename_ciphertext_length,
226 original_header.filename_ciphertext_length
227 );
228 assert_eq!(
229 deserialized_header.filename_ciphertext,
230 original_header.filename_ciphertext
231 );
232 }
233
234 #[test]
235 fn test_try_deserialize_insufficient_bytes() {
236 let bytes = vec![0u8; 50]; let result = try_deserialize(&bytes);
239 assert!(result.is_err());
240 assert!(matches!(
241 result.unwrap_err(),
242 HeaderError::InsufficientBytes
243 ));
244 }
245
246 #[test]
247 fn test_try_deserialize_invalid_data() {
248 let mut bytes = vec![0u8; 100];
249 bytes[7..11].copy_from_slice(&(50u32.to_le_bytes())); let result = try_deserialize(&bytes);
253 assert!(result.is_err());
254 assert!(matches!(result.unwrap_err(), HeaderError::InvalidData));
255 }
256
257 #[test]
258 fn test_try_deserialize_insufficient_bytes_for_filename() {
259 let mut bytes = vec![0u8; 95]; bytes[0..6].copy_from_slice(b"SHADOW");
262 bytes[6] = 1; bytes[7..11].copy_from_slice(&(100u32.to_le_bytes())); bytes[88..90].copy_from_slice(&(10u16.to_le_bytes())); let result = try_deserialize(&bytes);
269 assert!(result.is_err());
270 assert!(matches!(
271 result.unwrap_err(),
272 HeaderError::InsufficientBytes
273 ));
274 }
275
276 #[test]
277 fn test_round_trip_serialization() {
278 let original_header = create_test_header();
279 let serialized = serialize(&original_header);
280 let deserialized_result = try_deserialize(&serialized);
281
282 assert!(deserialized_result.is_ok());
283 let deserialized_header = deserialized_result.unwrap();
284
285 assert_eq!(original_header.magic, deserialized_header.magic);
287 assert_eq!(original_header.version, deserialized_header.version);
288 assert_eq!(
289 original_header.header_length,
290 deserialized_header.header_length
291 );
292 assert_eq!(original_header.salt, deserialized_header.salt);
293 assert_eq!(original_header.kdf_memory, deserialized_header.kdf_memory);
294 assert_eq!(
295 original_header.kdf_iterations,
296 deserialized_header.kdf_iterations
297 );
298 assert_eq!(
299 original_header.kdf_parallelism,
300 deserialized_header.kdf_parallelism
301 );
302 assert_eq!(
303 original_header.kdf_key_length,
304 deserialized_header.kdf_key_length
305 );
306 assert_eq!(
307 original_header.content_nonce,
308 deserialized_header.content_nonce
309 );
310 assert_eq!(
311 original_header.filename_nonce,
312 deserialized_header.filename_nonce
313 );
314 assert_eq!(
315 original_header.filename_ciphertext_length,
316 deserialized_header.filename_ciphertext_length
317 );
318 assert_eq!(
319 original_header.filename_ciphertext,
320 deserialized_header.filename_ciphertext
321 );
322 }
323
324 #[test]
325 fn test_empty_filename_ciphertext() {
326 let salt = [1u8; 16];
327 let kdf_params = KeyDerivationParams::from(profile::SecurityProfile::Test);
328 let content_nonce = [2u8; 24];
329 let filename_nonce = [3u8; 24];
330 let filename_ciphertext = vec![]; let header = FileHeader::new(
333 salt,
334 kdf_params,
335 content_nonce,
336 filename_nonce,
337 filename_ciphertext,
338 );
339
340 let serialized = serialize(&header);
341 let deserialized_result = try_deserialize(&serialized);
342
343 assert!(deserialized_result.is_ok());
344 let deserialized_header = deserialized_result.unwrap();
345 assert_eq!(header.magic, deserialized_header.magic);
346 assert_eq!(header.version, deserialized_header.version);
347 assert_eq!(header.header_length, deserialized_header.header_length);
348 assert_eq!(header.salt, deserialized_header.salt);
349 assert_eq!(header.kdf_memory, deserialized_header.kdf_memory);
350 assert_eq!(header.kdf_iterations, deserialized_header.kdf_iterations);
351 assert_eq!(header.kdf_parallelism, deserialized_header.kdf_parallelism);
352 assert_eq!(header.kdf_key_length, deserialized_header.kdf_key_length);
353 assert_eq!(header.content_nonce, deserialized_header.content_nonce);
354 assert_eq!(header.filename_nonce, deserialized_header.filename_nonce);
355 assert_eq!(
356 header.filename_ciphertext_length,
357 deserialized_header.filename_ciphertext_length
358 );
359 assert_eq!(
360 header.filename_ciphertext,
361 deserialized_header.filename_ciphertext
362 );
363 }
364
365 #[test]
366 fn test_large_filename_ciphertext() {
367 let salt = [1u8; 16];
368 let kdf_params = KeyDerivationParams::from(profile::SecurityProfile::Test);
369 let content_nonce = [2u8; 24];
370 let filename_nonce = [3u8; 24];
371 let filename_ciphertext = vec![4u8; 1000]; let header = FileHeader::new(
374 salt,
375 kdf_params,
376 content_nonce,
377 filename_nonce,
378 filename_ciphertext,
379 );
380
381 let serialized = serialize(&header);
382 let deserialized_result = try_deserialize(&serialized);
383
384 assert!(deserialized_result.is_ok());
385 let deserialized_header = deserialized_result.unwrap();
386 assert_eq!(header.magic, deserialized_header.magic);
387 assert_eq!(header.version, deserialized_header.version);
388 assert_eq!(header.header_length, deserialized_header.header_length);
389 assert_eq!(header.salt, deserialized_header.salt);
390 assert_eq!(header.kdf_memory, deserialized_header.kdf_memory);
391 assert_eq!(header.kdf_iterations, deserialized_header.kdf_iterations);
392 assert_eq!(header.kdf_parallelism, deserialized_header.kdf_parallelism);
393 assert_eq!(header.kdf_key_length, deserialized_header.kdf_key_length);
394 assert_eq!(header.content_nonce, deserialized_header.content_nonce);
395 assert_eq!(header.filename_nonce, deserialized_header.filename_nonce);
396 assert_eq!(
397 header.filename_ciphertext_length,
398 deserialized_header.filename_ciphertext_length
399 );
400 assert_eq!(
401 header.filename_ciphertext,
402 deserialized_header.filename_ciphertext
403 );
404 }
405}