shadow_crypt_core/v2/
header_ops.rs1use 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 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; 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 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}