1use 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 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}