Skip to main content

shadow_crypt_core/v1/
header_ops.rs

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        // Check that the serialized data has the correct length
137        assert_eq!(serialized.len(), header.header_length as usize);
138
139        // Check magic bytes
140        assert_eq!(&serialized[0..6], b"SHADOW");
141
142        // Check version
143        assert_eq!(serialized[6], 1);
144
145        // Check header length (little endian)
146        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        // Check salt
151        assert_eq!(&serialized[11..27], &header.salt);
152
153        // Check KDF parameters
154        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        // Check key length
167        assert_eq!(serialized[39], header.kdf_key_length);
168
169        // Check nonces
170        assert_eq!(&serialized[40..64], &header.content_nonce);
171        assert_eq!(&serialized[64..88], &header.filename_nonce);
172
173        // Check filename ciphertext length
174        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        // Check filename ciphertext
179        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]; // Less than min_length
237
238        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        // Set invalid header length (too small)
250        bytes[7..11].copy_from_slice(&(50u32.to_le_bytes())); // Header length smaller than min
251
252        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]; // FileHeader::min_length() is 90, but we need more for filename
260        // Set up a valid header but with filename_ciphertext_length > 0
261        bytes[0..6].copy_from_slice(b"SHADOW");
262        bytes[6] = 1; // version
263        bytes[7..11].copy_from_slice(&(100u32.to_le_bytes())); // header_length = 100 (90 + 10)
264        // salt, kdf params, nonces, etc. - keep as zeros for simplicity
265        bytes[88..90].copy_from_slice(&(10u16.to_le_bytes())); // filename_ciphertext_length = 10
266        // But we only have 95 bytes total, and we need 100, so insufficient
267
268        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        // Ensure all fields match
286        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![]; // Empty filename
331
332        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]; // Large filename
372
373        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}