Skip to main content

iceberg/encryption/
io.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18//! Encrypted file wrappers for InputFile / OutputFile.
19
20use 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
31/// An AGS1 stream-encrypted input file wrapping a plain [`InputFile`].
32///
33/// Transparently decrypts on read.
34pub struct EncryptedInputFile {
35    inner: InputFile,
36    key_metadata: StandardKeyMetadata,
37}
38
39impl EncryptedInputFile {
40    /// Creates a new encrypted input file.
41    pub fn new(inner: InputFile, key_metadata: StandardKeyMetadata) -> Self {
42        Self {
43            inner,
44            key_metadata,
45        }
46    }
47
48    /// Absolute path of the file.
49    pub fn location(&self) -> &str {
50        self.inner.location()
51    }
52
53    /// Check if file exists.
54    pub async fn exists(&self) -> Result<bool> {
55        self.inner.exists().await
56    }
57
58    /// Returns file metadata from the declared encrypted length without performing I/O.
59    ///
60    /// The returned size is the **plaintext** size.
61    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    /// Read and returns whole content of file (decrypted plaintext).
69    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    /// Creates a reader that transparently decrypts on each read.
76    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    // A storage stat would hide truncation; require the original length from key metadata.
86    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    /// Returns a reference to the file's key metadata.
99    pub fn key_metadata(&self) -> &StandardKeyMetadata {
100        &self.key_metadata
101    }
102
103    /// Consumes self and returns the underlying plain input file.
104    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
117/// An AGS1 stream-encrypted output file wrapping a plain [`OutputFile`].
118///
119/// Transparently encrypts on write.
120pub struct EncryptedOutputFile {
121    inner: OutputFile,
122    key_metadata: StandardKeyMetadata,
123}
124
125impl EncryptedOutputFile {
126    /// Creates a new encrypted output file.
127    pub fn new(inner: OutputFile, key_metadata: StandardKeyMetadata) -> Self {
128        Self {
129            inner,
130            key_metadata,
131        }
132    }
133
134    /// Returns key metadata using the encrypted size returned by [`FileWrite::close`] or [`Self::write`].
135    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    /// Absolute path of the file.
145    pub fn location(&self) -> &str {
146        self.inner.location()
147    }
148
149    /// Creates a writer that transparently encrypts on each write.
150    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    /// Write bytes to the file and return its encrypted size.
160    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    /// Deletes the underlying file.
167    pub async fn delete(&self) -> Result<()> {
168        self.inner.delete().await
169    }
170
171    /// Consumes self and returns the underlying plain output file.
172    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        // A missing path proves the size comes from the key metadata rather than a stat call.
244        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        // A declared length is trusted without a stat, so an inflated one is only caught once a
307        // read runs off the end of the real file. Both a minimal overstatement and one spanning a
308        // whole extra block must fail rather than silently return short plaintext.
309        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            // Not even the bytes that genuinely are on disk can be read back.
324            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}