1use std::sync::Arc;
21
22use bytes::Bytes;
23
24use super::crypto::AesGcmCipher;
25use super::key_metadata::StandardKeyMetadata;
26use super::stream::{AesGcmFileRead, AesGcmFileWrite, MIN_STREAM_LENGTH};
27use crate::Result;
28use crate::error::invalid_data;
29use crate::io::{FileMetadata, FileRead, FileWrite, InputFile, OutputFile};
30
31pub struct EncryptedInputFile {
35 inner: InputFile,
36 key_metadata: StandardKeyMetadata,
37}
38
39impl EncryptedInputFile {
40 pub fn new(inner: InputFile, key_metadata: StandardKeyMetadata) -> Self {
42 Self {
43 inner,
44 key_metadata,
45 }
46 }
47
48 pub fn location(&self) -> &str {
50 self.inner.location()
51 }
52
53 pub async fn exists(&self) -> Result<bool> {
55 self.inner.exists().await
56 }
57
58 pub fn metadata(&self) -> Result<FileMetadata> {
62 let plaintext_size = AesGcmFileRead::calculate_plaintext_length(self.encrypted_length()?)?;
63 Ok(FileMetadata {
64 size: plaintext_size,
65 })
66 }
67
68 pub async fn read(&self) -> Result<Bytes> {
70 let meta = self.metadata()?;
71 let reader = self.reader().await?;
72 reader.read(0..meta.size).await
73 }
74
75 pub async fn reader(&self) -> Result<Box<dyn FileRead>> {
77 let encrypted_length = self.encrypted_length()?;
78 let raw_reader = self.inner.reader().await?;
79 let cipher = build_cipher(&self.key_metadata)?;
80 let aad_prefix: Box<[u8]> = self.key_metadata.aad_prefix().unwrap_or_default().into();
81 let decrypting = AesGcmFileRead::new(raw_reader, cipher, aad_prefix, encrypted_length)?;
82 Ok(Box::new(decrypting))
83 }
84
85 fn encrypted_length(&self) -> Result<u64> {
87 let length = self.key_metadata.file_length().ok_or_else(|| {
88 invalid_data!("AGS1 key metadata is missing the encrypted file length")
89 })?;
90 if length < u64::from(MIN_STREAM_LENGTH) {
91 return Err(invalid_data!(
92 "Invalid encrypted file length: {length} is less than {MIN_STREAM_LENGTH}"
93 ));
94 }
95 Ok(length)
96 }
97
98 pub fn key_metadata(&self) -> &StandardKeyMetadata {
100 &self.key_metadata
101 }
102
103 pub fn into_inner(self) -> InputFile {
105 self.inner
106 }
107}
108
109impl std::fmt::Debug for EncryptedInputFile {
110 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
111 f.debug_struct("EncryptedInputFile")
112 .field("path", &self.inner.location())
113 .finish_non_exhaustive()
114 }
115}
116
117pub struct EncryptedOutputFile {
121 inner: OutputFile,
122 key_metadata: StandardKeyMetadata,
123}
124
125impl EncryptedOutputFile {
126 pub fn new(inner: OutputFile, key_metadata: StandardKeyMetadata) -> Self {
128 Self {
129 inner,
130 key_metadata,
131 }
132 }
133
134 pub fn key_metadata_with_saved_file_metadata(
136 &self,
137 file_metadata: &FileMetadata,
138 ) -> StandardKeyMetadata {
139 self.key_metadata
140 .clone()
141 .with_file_length(file_metadata.size)
142 }
143
144 pub fn location(&self) -> &str {
146 self.inner.location()
147 }
148
149 pub async fn writer(&self) -> Result<Box<dyn FileWrite>> {
151 let raw_writer = self.inner.writer().await?;
152 let cipher = build_cipher(&self.key_metadata)?;
153 let aad_prefix: Box<[u8]> = self.key_metadata.aad_prefix().unwrap_or_default().into();
154 Ok(Box::new(AesGcmFileWrite::new(
155 raw_writer, cipher, aad_prefix,
156 )))
157 }
158
159 pub async fn write(&self, bs: Bytes) -> Result<FileMetadata> {
161 let mut writer = self.writer().await?;
162 writer.write(bs).await?;
163 writer.close().await
164 }
165
166 pub async fn delete(&self) -> Result<()> {
168 self.inner.delete().await
169 }
170
171 pub fn into_inner(self) -> OutputFile {
173 self.inner
174 }
175}
176
177impl std::fmt::Debug for EncryptedOutputFile {
178 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
179 f.debug_struct("EncryptedOutputFile")
180 .field("path", &self.inner.location())
181 .finish_non_exhaustive()
182 }
183}
184
185fn build_cipher(metadata: &StandardKeyMetadata) -> Result<Arc<AesGcmCipher>> {
186 let key = metadata.encryption_key().clone();
187 Ok(Arc::new(AesGcmCipher::new(key)))
188}
189
190#[cfg(test)]
191mod tests {
192 use super::*;
193 use crate::ErrorKind;
194 use crate::encryption::stream::{
195 CIPHER_BLOCK_SIZE, GCM_STREAM_HEADER_LENGTH, PLAIN_BLOCK_SIZE,
196 };
197 use crate::io::FileIO;
198
199 fn key_metadata() -> StandardKeyMetadata {
200 StandardKeyMetadata::try_new(b"0123456789abcdef")
201 .unwrap()
202 .with_aad_prefix(b"test-aad-prefix!")
203 }
204
205 #[tokio::test]
206 async fn test_write_read_roundtrip() {
207 let fileio = FileIO::new_with_memory();
208 let path = "memory:///test/io_roundtrip.bin";
209 let plaintext = b"Hello from EncryptedInputFile/EncryptedOutputFile!";
210
211 let output = EncryptedOutputFile::new(fileio.new_output(path).unwrap(), key_metadata());
212 let file_metadata = output.write(Bytes::from(plaintext.to_vec())).await.unwrap();
213
214 let input = EncryptedInputFile::new(
215 fileio.new_input(path).unwrap(),
216 output.key_metadata_with_saved_file_metadata(&file_metadata),
217 );
218 let content = input.read().await.unwrap();
219 assert_eq!(&content[..], plaintext);
220 }
221
222 #[tokio::test]
223 async fn test_metadata_returns_plaintext_size() {
224 let fileio = FileIO::new_with_memory();
225 let path = "memory:///test/io_metadata.bin";
226 let plaintext = b"some bytes to measure";
227
228 let output = EncryptedOutputFile::new(fileio.new_output(path).unwrap(), key_metadata());
229 let file_metadata = output.write(Bytes::from(plaintext.to_vec())).await.unwrap();
230
231 let raw_size = fileio
232 .new_input(path)
233 .unwrap()
234 .metadata()
235 .await
236 .unwrap()
237 .size;
238 assert!(
239 raw_size > plaintext.len() as u64,
240 "encrypted file should be larger than plaintext (header + nonce + tag)"
241 );
242
243 let input = EncryptedInputFile::new(
245 fileio.new_input("memory:///does-not-exist").unwrap(),
246 output.key_metadata_with_saved_file_metadata(&file_metadata),
247 );
248 let meta = input.metadata().unwrap();
249 assert_eq!(meta.size, plaintext.len() as u64);
250 }
251
252 #[tokio::test]
253 async fn test_missing_file_length_is_rejected() {
254 let fileio = FileIO::new_with_memory();
255 let path = "memory:///test/missing_length.bin";
256 let output = EncryptedOutputFile::new(fileio.new_output(path).unwrap(), key_metadata());
257 output.write(Bytes::from_static(b"data")).await.unwrap();
258 let input = EncryptedInputFile::new(fileio.new_input(path).unwrap(), key_metadata());
259
260 for err in [
261 input.metadata().err().unwrap(),
262 input.reader().await.err().unwrap(),
263 input.read().await.unwrap_err(),
264 ] {
265 assert_eq!(err.kind(), ErrorKind::DataInvalid);
266 assert!(
267 err.to_string()
268 .contains("missing the encrypted file length")
269 );
270 }
271 }
272
273 #[tokio::test]
274 async fn test_invalid_file_length_is_rejected() {
275 let fileio = FileIO::new_with_memory();
276 for length in [
277 0,
278 u64::from(GCM_STREAM_HEADER_LENGTH),
279 u64::from(MIN_STREAM_LENGTH - 1),
280 ] {
281 let input = EncryptedInputFile::new(
282 fileio
283 .new_input("memory:///test/invalid_length.bin")
284 .unwrap(),
285 key_metadata().with_file_length(length),
286 );
287 assert_eq!(
288 input.metadata().err().unwrap().kind(),
289 ErrorKind::DataInvalid
290 );
291 assert_eq!(
292 input.reader().await.err().unwrap().kind(),
293 ErrorKind::DataInvalid
294 );
295 }
296 }
297
298 #[tokio::test]
299 async fn test_oversized_file_length_is_rejected() {
300 let fileio = FileIO::new_with_memory();
301 let path = "memory:///test/oversized_length.bin";
302 let plaintext = Bytes::from_static(b"some bytes to measure");
303 let output = EncryptedOutputFile::new(fileio.new_output(path).unwrap(), key_metadata());
304 let file_metadata = output.write(plaintext.clone()).await.unwrap();
305
306 for excess in [1, u64::from(CIPHER_BLOCK_SIZE)] {
310 let input = EncryptedInputFile::new(
311 fileio.new_input(path).unwrap(),
312 key_metadata().with_file_length(file_metadata.size + excess),
313 );
314
315 let inflated_size = input.metadata().unwrap().size;
316 assert!(inflated_size > plaintext.len() as u64);
317
318 assert_eq!(
319 input.read().await.unwrap_err().kind(),
320 ErrorKind::DataInvalid
321 );
322
323 let reader = input.reader().await.unwrap();
325 assert_eq!(
326 reader
327 .read(0..plaintext.len() as u64)
328 .await
329 .unwrap_err()
330 .kind(),
331 ErrorKind::DataInvalid
332 );
333 }
334 }
335
336 #[tokio::test]
337 async fn test_truncated_file_is_rejected() {
338 let fileio = FileIO::new_with_memory();
339 let path = "memory:///test/truncated.bin";
340 let plaintext = Bytes::from(vec![42; 2 * PLAIN_BLOCK_SIZE as usize + 17]);
341 let output = EncryptedOutputFile::new(fileio.new_output(path).unwrap(), key_metadata());
342 let file_metadata = output.write(plaintext.clone()).await.unwrap();
343 let metadata = key_metadata().with_file_length(file_metadata.size);
344 let ciphertext = fileio.new_input(path).unwrap().read().await.unwrap();
345 let truncated_length = (GCM_STREAM_HEADER_LENGTH + CIPHER_BLOCK_SIZE) as usize;
346 fileio
347 .new_output(path)
348 .unwrap()
349 .write(ciphertext.slice(..truncated_length))
350 .await
351 .unwrap();
352
353 let input = EncryptedInputFile::new(fileio.new_input(path).unwrap(), metadata);
354 assert_eq!(input.metadata().unwrap().size, plaintext.len() as u64);
355 let reader = input.reader().await.unwrap();
356 assert_eq!(
357 reader.read(0..u64::from(PLAIN_BLOCK_SIZE)).await.unwrap(),
358 plaintext.slice(..PLAIN_BLOCK_SIZE as usize)
359 );
360 assert_eq!(
361 input.read().await.unwrap_err().kind(),
362 ErrorKind::DataInvalid
363 );
364 }
365
366 #[tokio::test]
367 async fn test_truncated_empty_file_is_rejected() {
368 let fileio = FileIO::new_with_memory();
369 let path = "memory:///test/truncated_empty.bin";
370 let output = EncryptedOutputFile::new(fileio.new_output(path).unwrap(), key_metadata());
371 let file_metadata = output.write(Bytes::new()).await.unwrap();
372 let metadata = key_metadata().with_file_length(file_metadata.size);
373 let ciphertext = fileio.new_input(path).unwrap().read().await.unwrap();
374 fileio
375 .new_output(path)
376 .unwrap()
377 .write(ciphertext.slice(..GCM_STREAM_HEADER_LENGTH as usize))
378 .await
379 .unwrap();
380 let input = EncryptedInputFile::new(fileio.new_input(path).unwrap(), metadata);
381 assert_eq!(
382 input.read().await.unwrap_err().kind(),
383 ErrorKind::DataInvalid
384 );
385 }
386
387 #[tokio::test]
388 async fn test_close_returns_encrypted_size() {
389 let fileio = FileIO::new_with_memory();
390 let path = "memory:///test/streaming.bin";
391 let output = EncryptedOutputFile::new(fileio.new_output(path).unwrap(), key_metadata());
392 for plaintext in [
393 Bytes::from(vec![42; PLAIN_BLOCK_SIZE as usize + 17]),
394 Bytes::new(),
395 ] {
396 let mut writer = output.writer().await.unwrap();
397 for chunk in plaintext.chunks(1024) {
398 writer.write(Bytes::copy_from_slice(chunk)).await.unwrap();
399 }
400 let metadata = writer.close().await.unwrap();
401 let size = fileio
402 .new_input(path)
403 .unwrap()
404 .metadata()
405 .await
406 .unwrap()
407 .size;
408 assert_eq!(metadata.size, size);
409 }
410 }
411}