Skip to main content

ahri_tre_lake/
upload.rs

1//! Validation of uploaded tables while they are still private scratch files.
2use crate::{DatasetFileFormat, DatasetFileReadOptions, LakeError};
3use std::{io::Read, path::Path, time::Instant};
4
5#[derive(Debug, Clone)]
6pub struct UploadValidation {
7    pub format: DatasetFileFormat,
8    pub options: DatasetFileReadOptions,
9    pub budgets: ahri_tre_types::DisclosureBudgets,
10    pub deadline: Instant,
11    pub cancelled: std::sync::Arc<std::sync::atomic::AtomicBool>,
12}
13impl UploadValidation {
14    pub fn check_deadline(&self) -> Result<(), LakeError> {
15        if Instant::now() >= self.deadline
16            || self.cancelled.load(std::sync::atomic::Ordering::Acquire)
17        {
18            return Err(invalid());
19        }
20        Ok(())
21    }
22    pub(crate) fn validate(&self, path: &Path) -> Result<(), LakeError> {
23        self.check_deadline()?;
24        if path.metadata()?.len() > self.budgets.decoded_file_bytes {
25            return Err(invalid());
26        }
27        if self.format == DatasetFileFormat::ArrowIpc {
28            let mut file = ahri_tre_tabular::open_ipc_messages(path)?;
29            let mut validator = ahri_tre_tabular::streaming::ArrowValidator::for_upload(
30                self.budgets.rows,
31                self.budgets.decoded_file_bytes,
32            );
33            let mut buffer = [0; 16 * 1024];
34            loop {
35                self.check_deadline()?;
36                let size = file.read(&mut buffer)?;
37                if size == 0 {
38                    break;
39                }
40                validator.push(&buffer[..size])?;
41            }
42            validator.finish()?;
43            return Ok(());
44        }
45        match self.format {
46            DatasetFileFormat::Parquet => validate_parquet(path, self.budgets.decoded_file_bytes)?,
47            DatasetFileFormat::Xlsx => validate_xlsx(path, self.budgets.decoded_file_bytes, self)?,
48            _ => (),
49        }
50        let connection = duckdb::Connection::open_in_memory()?;
51        // Engine memory is separate from transfer buffers. No spill outside the
52        // trusted scratch directory, implicit extension installation or network.
53        connection.execute_batch("SET memory_limit='128MB'; SET threads=1; SET temp_directory=''; SET autoinstall_known_extensions=false; SET autoload_known_extensions=false;")?;
54        if self.format == DatasetFileFormat::Xlsx {
55            connection.execute_batch("LOAD excel")?;
56        }
57        let relation = crate::dataset_loading::file_relation_sql(path, self.format, &self.options)?;
58        let (stop, stopped) = std::sync::mpsc::channel::<()>();
59        let interrupt = connection.interrupt_handle();
60        let deadline = self.deadline;
61        let cancelled = self.cancelled.clone();
62        let watchdog = std::thread::spawn(move || {
63            while stopped
64                .recv_timeout(std::time::Duration::from_millis(25))
65                .is_err()
66            {
67                if Instant::now() >= deadline
68                    || cancelled.load(std::sync::atomic::Ordering::Acquire)
69                {
70                    interrupt.interrupt();
71                }
72            }
73        });
74        let rows = connection.query_row(
75            &format!("SELECT count(*), sum(hash(upload_validation)) FROM ({relation}) AS upload_validation"),
76            [],
77            |row| row.get::<_, u64>(0),
78        );
79        let _ = stop.send(());
80        let _ = watchdog.join();
81        self.check_deadline()?;
82        if rows? > self.budgets.rows {
83            return Err(invalid());
84        }
85        Ok(())
86    }
87}
88fn invalid() -> LakeError {
89    std::io::Error::other("Uploaded table exceeds its validation limits").into()
90}
91
92fn validate_parquet(path: &Path, max_decoded: u64) -> Result<(), LakeError> {
93    use std::io::{Read, Seek, SeekFrom};
94    let mut file = std::fs::File::open(path)?;
95    let length = file.metadata()?.len();
96    if length < 12 {
97        return Err(invalid());
98    }
99    file.seek(SeekFrom::End(-8))?;
100    let mut tail = [0; 8];
101    file.read_exact(&mut tail)?;
102    let footer = u32::from_le_bytes(tail[..4].try_into().unwrap()) as u64;
103    if &tail[4..] != b"PAR1" || footer > 8 * 1024 * 1024 || footer + 12 > length {
104        return Err(invalid());
105    }
106    use parquet::file::reader::FileReader;
107    let reader = parquet::file::reader::SerializedFileReader::new(file)
108        .map_err(ahri_tre_tabular::TabularError::from)?;
109    let mut decoded = 0_u64;
110    for group in reader.metadata().row_groups() {
111        for column in group.columns() {
112            let size = u64::try_from(column.uncompressed_size()).map_err(|_| invalid())?;
113            if size > 8 * 1024 * 1024 {
114                return Err(invalid());
115            }
116            decoded = decoded
117                .checked_add(size)
118                .filter(|n| *n <= max_decoded)
119                .ok_or_else(invalid)?;
120        }
121    }
122    Ok(())
123}
124fn validate_xlsx(
125    path: &Path,
126    max_decoded: u64,
127    validation: &UploadValidation,
128) -> Result<(), LakeError> {
129    use std::io::{Read, Seek, SeekFrom};
130    let mut file = std::fs::File::open(path)?;
131    let length = file.metadata()?.len();
132    let size = length.min(65557) as usize;
133    file.seek(SeekFrom::End(-(size as i64)))?;
134    let mut tail = vec![0; size];
135    file.read_exact(&mut tail)?;
136    let end = tail
137        .windows(4)
138        .rposition(|part| part == b"PK\x05\x06")
139        .ok_or_else(invalid)?;
140    let end = tail.get(end..end + 22).ok_or_else(invalid)?;
141    if u16::from_le_bytes(end[10..12].try_into().unwrap()) > 10000
142        || u32::from_le_bytes(end[12..16].try_into().unwrap()) > 8 * 1024 * 1024
143    {
144        return Err(invalid());
145    }
146    file.seek(SeekFrom::Start(0))?;
147    let mut archive = zip::ZipArchive::new(file).map_err(|_| invalid())?;
148    let mut decoded = 0_u64;
149    let mut buffer = [0; 16 * 1024];
150    for index in 0..archive.len() {
151        let mut entry = archive.by_index(index).map_err(|_| invalid())?;
152        loop {
153            validation.check_deadline()?;
154            let count = entry.read(&mut buffer)?;
155            if count == 0 {
156                break;
157            }
158            decoded = decoded
159                .checked_add(count as u64)
160                .filter(|n| *n <= max_decoded)
161                .ok_or_else(invalid)?;
162        }
163    }
164    Ok(())
165}