1use std::cell::RefCell;
4use std::io::{Read, Seek, Write};
5
6use corecrypto::cipher::Ciphers;
7use corecrypto::header::{Header, HeaderType};
8use corecrypto::key::decrypt_master_key;
9use corecrypto::primitives::Mode;
10use corecrypto::protected::Protected;
11use corecrypto::stream::DecryptionStreams;
12
13#[derive(Debug)]
14pub enum Error {
15 InitializeChiphers,
16 InitializeStreams,
17 DeserializeHeader,
18 ReadEncryptedData,
19 DecryptMasterKey,
20 DecryptData,
21 WriteData,
22 RewindDataReader,
23}
24
25impl std::fmt::Display for Error {
26 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
27 match self {
28 Error::InitializeChiphers => f.write_str("Cannot initialize chiphers"),
29 Error::InitializeStreams => f.write_str("Cannot initialize streams"),
30 Error::DeserializeHeader => f.write_str("Cannot deserialize header"),
31 Error::ReadEncryptedData => f.write_str("Unable to read encrypted data"),
32 Error::DecryptMasterKey => f.write_str("Cannot decrypt master key"),
33 Error::DecryptData => f.write_str("Unable to decrypt data"),
34 Error::WriteData => f.write_str("Unable to write data"),
35 Error::RewindDataReader => f.write_str("Unable to rewind the reader"),
36 }
37 }
38}
39
40impl std::error::Error for Error {}
41
42pub type OnDecryptedHeaderFn = Box<dyn FnOnce(&HeaderType)>;
43
44pub struct Request<'a, R, W>
45where
46 R: Read + Seek,
47 W: Write + Seek,
48{
49 pub header_reader: Option<&'a RefCell<R>>,
50 pub reader: &'a RefCell<R>,
51 pub writer: &'a RefCell<W>,
52 pub raw_key: Protected<Vec<u8>>,
53 pub on_decrypted_header: Option<OnDecryptedHeaderFn>,
54}
55
56pub fn execute<R, W>(req: Request<'_, R, W>) -> Result<(), Error>
57where
58 R: Read + Seek,
59 W: Write + Seek,
60{
61 let (header, aad) = match req.header_reader {
62 Some(header_reader) => {
63 let (header, aad) = Header::deserialize(&mut *header_reader.borrow_mut())
64 .map_err(|_| Error::DeserializeHeader)?;
65
66 #[allow(clippy::cast_possible_truncation)]
68 let mut header_bytes = vec![0u8; header.get_size() as usize];
69
70 req.reader
71 .borrow_mut()
72 .read_exact(&mut header_bytes)
73 .or_else(|e| {
74 if e.kind() == std::io::ErrorKind::UnexpectedEof {
75 Ok(())
76 } else {
77 Err(e)
78 }
79 })
80 .map_err(|_| Error::ReadEncryptedData)?;
81
82 if !header_bytes.into_iter().all(|b| b == 0) {
83 req.reader
85 .borrow_mut()
86 .rewind()
87 .map_err(|_| Error::RewindDataReader)?;
88 }
89
90 (header, aad)
91 }
92 None => Header::deserialize(&mut *req.reader.borrow_mut())
93 .map_err(|_| Error::DeserializeHeader)?,
94 };
95
96 if let Some(cb) = req.on_decrypted_header {
97 cb(&header.header_type);
98 }
99
100 match header.header_type.mode {
101 Mode::MemoryMode => {
102 let mut encrypted_data = Vec::new();
103 req.reader
104 .borrow_mut()
105 .read_to_end(&mut encrypted_data)
106 .map_err(|_| Error::ReadEncryptedData)?;
107
108 let master_key =
109 decrypt_master_key(req.raw_key, &header).map_err(|_| Error::DecryptMasterKey)?;
110
111 let ciphers = Ciphers::initialize(master_key, &header.header_type.algorithm)
112 .map_err(|_| Error::InitializeChiphers)?;
113
114 let payload = corecrypto::Payload {
115 aad: &aad,
116 msg: &encrypted_data,
117 };
118
119 let decrypted_bytes = ciphers
120 .decrypt(&header.nonce, payload)
121 .map_err(|_| Error::DecryptData)?;
122
123 req.writer
124 .borrow_mut()
125 .write_all(&decrypted_bytes)
126 .map_err(|_| Error::WriteData)?;
127 }
128 Mode::StreamMode => {
129 let master_key =
130 decrypt_master_key(req.raw_key, &header).map_err(|_| Error::DecryptMasterKey)?;
131
132 let streams = DecryptionStreams::initialize(
133 master_key,
134 &header.nonce,
135 &header.header_type.algorithm,
136 )
137 .map_err(|_| Error::InitializeStreams)?;
138
139 streams
140 .decrypt_file(
141 &mut *req.reader.borrow_mut(),
142 &mut *req.writer.borrow_mut(),
143 &aad,
144 )
145 .map_err(|_| Error::DecryptData)?;
146 }
147 }
148
149 Ok(())
150}
151
152#[cfg(test)]
153mod tests {
154 use super::*;
155 use std::io::Cursor;
156
157 use crate::encrypt::tests::{
158 PASSWORD, V4_ENCRYPTED_CONTENT, V5_ENCRYPTED_CONTENT, V5_ENCRYPTED_DETACHED_CONTENT,
159 V5_ENCRYPTED_DETACHED_HEADER, V5_ENCRYPTED_FULL_DETACHED_CONTENT,
160 };
161
162 #[test]
163 fn should_decrypt_encrypted_content_with_v4_version() {
164 let mut input_content = V4_ENCRYPTED_CONTENT.to_vec();
165 let input_cur = RefCell::new(Cursor::new(&mut input_content));
166
167 let mut output_content = vec![];
168 let output_cur = RefCell::new(Cursor::new(&mut output_content));
169
170 let req = Request {
171 header_reader: None,
172 reader: &input_cur,
173 writer: &output_cur,
174 raw_key: Protected::new(PASSWORD.to_vec()),
175 on_decrypted_header: None,
176 };
177
178 match execute(req) {
179 Ok(_) => {
180 assert_eq!(output_content, "Hello world".as_bytes().to_vec());
181 }
182 _ => unreachable!(),
183 }
184 }
185
186 #[test]
187 fn should_decrypt_encrypted_content_with_v5_version() {
188 let mut input_content = V5_ENCRYPTED_CONTENT.to_vec();
189 let input_cur = RefCell::new(Cursor::new(&mut input_content));
190
191 let mut output_content = vec![];
192 let output_cur = RefCell::new(Cursor::new(&mut output_content));
193
194 let req = Request {
195 header_reader: None,
196 reader: &input_cur,
197 writer: &output_cur,
198 raw_key: Protected::new(PASSWORD.to_vec()),
199 on_decrypted_header: None,
200 };
201
202 match execute(req) {
203 Ok(_) => {
204 assert_eq!(output_content, "Hello world".as_bytes().to_vec());
205 }
206 _ => unreachable!(),
207 }
208 }
209
210 #[test]
211 fn should_decrypt_encrypted_detached_header_and_content_with_v5_version() {
212 let mut input_content = V5_ENCRYPTED_DETACHED_CONTENT.to_vec();
213 let input_cur = RefCell::new(Cursor::new(&mut input_content));
214
215 let mut input_header = V5_ENCRYPTED_DETACHED_HEADER.to_vec();
216 let header_cur = RefCell::new(Cursor::new(&mut input_header));
217
218 let mut output_content = vec![];
219 let output_cur = RefCell::new(Cursor::new(&mut output_content));
220
221 let req = Request {
222 header_reader: Some(&header_cur),
223 reader: &input_cur,
224 writer: &output_cur,
225 raw_key: Protected::new(PASSWORD.to_vec()),
226 on_decrypted_header: None,
227 };
228
229 match execute(req) {
230 Ok(_) => {
231 assert_eq!(output_content, "Hello world".as_bytes().to_vec());
232 }
233 _ => unreachable!(),
234 }
235 }
236
237 #[test]
238 fn should_decrypt_encrypted_full_detached_header_and_content_with_v5_version() {
239 let mut input_content = V5_ENCRYPTED_FULL_DETACHED_CONTENT.to_vec();
240 let input_cur = RefCell::new(Cursor::new(&mut input_content));
241
242 let mut input_header = V5_ENCRYPTED_DETACHED_HEADER.to_vec();
243 let header_cur = RefCell::new(Cursor::new(&mut input_header));
244
245 let mut output_content = vec![];
246 let output_cur = RefCell::new(Cursor::new(&mut output_content));
247
248 let req = Request {
249 header_reader: Some(&header_cur),
250 reader: &input_cur,
251 writer: &output_cur,
252 raw_key: Protected::new(PASSWORD.to_vec()),
253 on_decrypted_header: None,
254 };
255
256 match execute(req) {
257 Ok(_) => {
258 assert_eq!(output_content, "Hello world".as_bytes().to_vec());
259 }
260 _ => unreachable!(),
261 }
262 }
263}