Skip to main content

shadow_crypt_core/v2/
header_ops.rs

1use crate::{errors::HeaderError, v2::key::KeyDerivationParams};
2
3use super::header::{FileHeader, MAGIC, VERSION};
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: [u8; 6] = bytes[0..6].try_into().ok()?;
68    let version = bytes[6];
69
70    // Unlike v1, v2 rejects a wrong magic or version byte at parse time.
71    if magic != MAGIC || version != VERSION {
72        return None;
73    }
74
75    let header_length = u32::from_le_bytes(bytes[7..11].try_into().ok()?);
76    let salt = bytes[11..27].try_into().ok()?;
77    let kdf_memory = u32::from_le_bytes(bytes[27..31].try_into().ok()?);
78    let kdf_iterations = u32::from_le_bytes(bytes[31..35].try_into().ok()?);
79    let kdf_parallelism = u32::from_le_bytes(bytes[35..39].try_into().ok()?);
80    let kdf_key_length = bytes[39];
81    let content_nonce = bytes[40..64].try_into().ok()?;
82    let filename_nonce = bytes[64..88].try_into().ok()?;
83    let filename_ciphertext_length = u16::from_le_bytes(bytes[88..90].try_into().ok()?);
84
85    let expected_length: usize = FileHeader::min_length() + filename_ciphertext_length as usize;
86
87    if header_length != expected_length as u32 {
88        return None;
89    }
90
91    if bytes.len() < expected_length {
92        return None;
93    }
94
95    let filename_ciphertext = bytes[FileHeader::min_length()
96        ..(FileHeader::min_length() + filename_ciphertext_length as usize)]
97        .to_vec();
98
99    Some(FileHeader {
100        magic,
101        version,
102        header_length,
103        salt,
104        kdf_memory,
105        kdf_iterations,
106        kdf_parallelism,
107        kdf_key_length,
108        content_nonce,
109        filename_nonce,
110        filename_ciphertext_length,
111        filename_ciphertext,
112    })
113}
114
115#[cfg(test)]
116mod tests {
117    use super::*;
118    use crate::profile;
119    use crate::v2::key::KeyDerivationParams;
120
121    fn create_test_header() -> FileHeader {
122        let salt = [1u8; 16];
123        let kdf_params = KeyDerivationParams::from(profile::SecurityProfile::Test);
124        let content_nonce = [2u8; 24];
125        let filename_nonce = [3u8; 24];
126        let filename_ciphertext = vec![4, 5, 6, 7, 8];
127
128        FileHeader::new(
129            salt,
130            kdf_params,
131            content_nonce,
132            filename_nonce,
133            filename_ciphertext,
134        )
135        .unwrap()
136    }
137
138    #[test]
139    fn test_round_trip_serialization() {
140        let original = create_test_header();
141        let serialized = serialize(&original);
142        assert_eq!(serialized.len(), original.header_length as usize);
143        assert_eq!(&serialized[0..6], b"SHADOW");
144        assert_eq!(serialized[6], 2);
145
146        let deserialized = try_deserialize(&serialized).unwrap();
147        assert_eq!(deserialized.magic, original.magic);
148        assert_eq!(deserialized.version, original.version);
149        assert_eq!(deserialized.header_length, original.header_length);
150        assert_eq!(deserialized.salt, original.salt);
151        assert_eq!(deserialized.kdf_memory, original.kdf_memory);
152        assert_eq!(deserialized.kdf_iterations, original.kdf_iterations);
153        assert_eq!(deserialized.kdf_parallelism, original.kdf_parallelism);
154        assert_eq!(deserialized.kdf_key_length, original.kdf_key_length);
155        assert_eq!(deserialized.content_nonce, original.content_nonce);
156        assert_eq!(deserialized.filename_nonce, original.filename_nonce);
157        assert_eq!(
158            deserialized.filename_ciphertext_length,
159            original.filename_ciphertext_length
160        );
161        assert_eq!(
162            deserialized.filename_ciphertext,
163            original.filename_ciphertext
164        );
165    }
166
167    #[test]
168    fn test_try_deserialize_rejects_wrong_version() {
169        let header = create_test_header();
170        let mut serialized = serialize(&header);
171        serialized[6] = 1; // claim v1
172
173        assert!(try_deserialize(&serialized).is_err());
174    }
175
176    #[test]
177    fn test_try_deserialize_rejects_wrong_magic() {
178        let header = create_test_header();
179        let mut serialized = serialize(&header);
180        serialized[0..6].copy_from_slice(b"NOTSHD");
181
182        assert!(try_deserialize(&serialized).is_err());
183    }
184
185    #[test]
186    fn test_try_deserialize_insufficient_bytes() {
187        let bytes = vec![0u8; 50];
188        assert!(matches!(
189            try_deserialize(&bytes),
190            Err(HeaderError::InsufficientBytes)
191        ));
192    }
193
194    #[test]
195    fn test_try_deserialize_inconsistent_lengths() {
196        let header = create_test_header();
197        let mut serialized = serialize(&header);
198        // header_length no longer matches min_length + filename_ciphertext_length
199        serialized[7..11].copy_from_slice(&(200u32.to_le_bytes()));
200
201        assert!(try_deserialize(&serialized).is_err());
202    }
203
204    #[test]
205    fn test_empty_filename_ciphertext_round_trip() {
206        let header = FileHeader::new(
207            [1u8; 16],
208            KeyDerivationParams::from(profile::SecurityProfile::Test),
209            [2u8; 24],
210            [3u8; 24],
211            vec![],
212        )
213        .unwrap();
214
215        let deserialized = try_deserialize(&serialize(&header)).unwrap();
216        assert_eq!(deserialized.filename_ciphertext_length, 0);
217        assert!(deserialized.filename_ciphertext.is_empty());
218    }
219}