1use crate::LakeError;
4use sqlparser::{ast::*, dialect::DuckDbDialect, parser::Parser};
5use std::{
6 collections::{BTreeMap, BTreeSet},
7 ops::ControlFlow,
8};
9use uuid::Uuid;
10
11#[derive(Debug, Clone)]
12pub enum QueryBinding {
13 Dataset(Uuid),
14 View(String),
15}
16
17pub struct RestrictedQueryPlan {
18 sql: String,
19 inputs: BTreeSet<Uuid>,
20}
21impl RestrictedQueryPlan {
22 pub fn input_versions(&self) -> &BTreeSet<Uuid> {
23 &self.inputs
24 }
25 pub fn restricted_sql(&self) -> &str {
27 &self.sql
28 }
29}
30
31pub fn analyze_disclosure_query(
32 sql: &str,
33 bindings: &BTreeMap<String, QueryBinding>,
34) -> Result<RestrictedQueryPlan, LakeError> {
35 if bindings.len() > 128
36 || bindings
37 .keys()
38 .any(|name| name.is_empty() || name.len() > 768 || *name != name.to_ascii_lowercase())
39 {
40 return Err(LakeError::QueryNotAdmissible);
41 }
42 let mut inputs = BTreeSet::new();
43 let mut remaining_bytes = AnalysisBudget {
44 bytes: 1_048_576,
45 tokens: 256,
46 };
47 let query = expand(
48 sql,
49 bindings,
50 &mut Vec::new(),
51 &mut inputs,
52 &mut remaining_bytes,
53 )?;
54 let sql = query.to_string();
55 if sql.len() > 1_048_576 {
56 return Err(LakeError::QueryNotAdmissible);
57 }
58 Ok(RestrictedQueryPlan { sql, inputs })
59}
60
61struct AnalysisBudget {
62 bytes: usize,
63 tokens: usize,
64}
65
66fn expand(
67 sql: &str,
68 bindings: &BTreeMap<String, QueryBinding>,
69 views: &mut Vec<String>,
70 inputs: &mut BTreeSet<Uuid>,
71 remaining_bytes: &mut AnalysisBudget,
72) -> Result<Box<Query>, LakeError> {
73 if sql.len() > 65536 || views.len() > 32 {
74 return Err(LakeError::QueryNotAdmissible);
75 }
76 remaining_bytes.bytes = remaining_bytes
77 .bytes
78 .checked_sub(sql.len().max(1))
79 .ok_or(LakeError::QueryNotAdmissible)?;
80 let tokens = sqlparser::tokenizer::Tokenizer::new(&DuckDbDialect {}, sql)
83 .tokenize()
84 .map_err(|_| LakeError::QueryNotAdmissible)?;
85 remaining_bytes.tokens = remaining_bytes
86 .tokens
87 .checked_sub(
88 tokens
89 .iter()
90 .filter(|token| !matches!(token, sqlparser::tokenizer::Token::Whitespace(_)))
91 .count(),
92 )
93 .ok_or(LakeError::QueryNotAdmissible)?;
94 let mut statements = Parser::new(&DuckDbDialect {})
95 .with_recursion_limit(64)
96 .try_with_sql(sql)
97 .and_then(|mut parser| parser.parse_statements())
98 .map_err(|_| LakeError::QueryNotAdmissible)?;
99 if statements.len() != 1 {
100 return Err(LakeError::QueryNotAdmissible);
101 }
102 let Statement::Query(mut query) = statements.remove(0) else {
103 return Err(LakeError::QueryNotAdmissible);
104 };
105 let mut analyzer = Analyzer {
106 bindings,
107 views,
108 inputs,
109 remaining_bytes,
110 scopes: Vec::new(),
111 nodes: 0,
112 };
113 if VisitMut::visit(&mut query, &mut analyzer).is_break() {
114 return Err(LakeError::QueryNotAdmissible);
115 }
116 Ok(query)
117}
118
119struct Analyzer<'a> {
120 bindings: &'a BTreeMap<String, QueryBinding>,
121 views: &'a mut Vec<String>,
122 inputs: &'a mut BTreeSet<Uuid>,
123 remaining_bytes: &'a mut AnalysisBudget,
124 scopes: Vec<(BTreeSet<String>, Option<With>)>,
125 nodes: usize,
126}
127impl VisitorMut for Analyzer<'_> {
128 type Break = ();
129 fn pre_visit_query(&mut self, query: &mut Query) -> ControlFlow<()> {
130 if !query.locks.is_empty()
131 || query.for_clause.is_some()
132 || query.settings.is_some()
133 || query.format_clause.is_some()
134 || !query.pipe_operators.is_empty()
135 || !safe_set(&query.body)
136 {
137 return ControlFlow::Break(());
138 }
139 self.scopes.push((BTreeSet::new(), None));
140 let mut retained_with = query.with.take();
141 if let Some(with) = &mut retained_with {
142 if with.recursive {
143 return ControlFlow::Break(());
144 }
145 for cte in &mut with.cte_tables {
146 let name = cte.alias.name.value.to_ascii_lowercase();
147 if self.bindings.contains_key(&name)
148 || self.scopes.iter().any(|(names, _)| names.contains(&name))
149 {
150 return ControlFlow::Break(());
151 }
152 if VisitMut::visit(&mut cte.query, self).is_break() {
156 return ControlFlow::Break(());
157 }
158 self.scopes.last_mut().unwrap().0.insert(name);
159 }
160 }
161 self.scopes.last_mut().unwrap().1 = retained_with;
162 ControlFlow::Continue(())
163 }
164 fn post_visit_query(&mut self, query: &mut Query) -> ControlFlow<()> {
165 query.with = self.scopes.pop().unwrap().1;
166 ControlFlow::Continue(())
167 }
168 fn pre_visit_select(&mut self, select: &mut Select) -> ControlFlow<()> {
169 if select.into.is_some()
170 || !select.optimizer_hints.is_empty()
171 || select.select_modifiers.is_some()
172 || select.top.is_some()
173 || !select.lateral_views.is_empty()
174 || select.prewhere.is_some()
175 || !select.connect_by.is_empty()
176 || !select.cluster_by.is_empty()
177 || !select.distribute_by.is_empty()
178 || !select.sort_by.is_empty()
179 || select.value_table_mode.is_some()
180 {
181 return ControlFlow::Break(());
182 }
183 ControlFlow::Continue(())
184 }
185 fn pre_visit_statement(&mut self, statement: &mut Statement) -> ControlFlow<()> {
186 if matches!(statement, Statement::Query(_)) {
187 ControlFlow::Continue(())
188 } else {
189 ControlFlow::Break(())
190 }
191 }
192 fn pre_visit_expr(&mut self, expr: &mut Expr) -> ControlFlow<()> {
193 self.nodes += 1;
194 if self.nodes > 8192 {
195 return ControlFlow::Break(());
196 }
197 if let Expr::CompoundIdentifier(parts) = expr
198 && parts.len() > 2
199 {
200 let prefix = parts[..parts.len() - 1]
201 .iter()
202 .map(|p| p.value.to_ascii_lowercase())
203 .collect::<Vec<_>>()
204 .join("\0");
205 if self.bindings.contains_key(&prefix) {
206 *parts = parts[parts.len() - 2..].to_vec();
207 }
208 }
209 let allowed = match expr {
210 Expr::Identifier(_)
211 | Expr::CompoundIdentifier(_)
212 | Expr::Value(_)
213 | Expr::Nested(_)
214 | Expr::IsNull(_)
215 | Expr::IsNotNull(_)
216 | Expr::IsTrue(_)
217 | Expr::IsFalse(_)
218 | Expr::IsNotTrue(_)
219 | Expr::IsNotFalse(_)
220 | Expr::IsUnknown(_)
221 | Expr::IsNotUnknown(_)
222 | Expr::IsDistinctFrom(..)
223 | Expr::IsNotDistinctFrom(..)
224 | Expr::InList { .. }
225 | Expr::InSubquery { .. }
226 | Expr::Between { .. }
227 | Expr::Like { .. }
228 | Expr::ILike { .. }
229 | Expr::Cast { .. }
230 | Expr::Extract { .. }
231 | Expr::Ceil { .. }
232 | Expr::Floor { .. }
233 | Expr::Position { .. }
234 | Expr::Substring { .. }
235 | Expr::Trim { .. }
236 | Expr::TypedString(_)
237 | Expr::Case { .. }
238 | Expr::Exists { .. }
239 | Expr::Subquery(_)
240 | Expr::Tuple(_)
241 | Expr::Interval(_)
242 | Expr::Wildcard(_)
243 | Expr::QualifiedWildcard(..) => true,
244 Expr::BinaryOp { op, .. } => matches!(
245 op.to_string().as_str(),
246 "+" | "-"
247 | "*"
248 | "/"
249 | "%"
250 | "="
251 | "<>"
252 | "!="
253 | "<"
254 | ">"
255 | "<="
256 | ">="
257 | "AND"
258 | "OR"
259 | "||"
260 ),
261 Expr::UnaryOp { op, .. } => matches!(op.to_string().as_str(), "+" | "-" | "NOT"),
262 Expr::Function(function) => {
263 single_name(&function.name).is_some_and(|name| {
264 matches!(
265 name.as_str(),
266 "abs"
267 | "ceil"
268 | "floor"
269 | "round"
270 | "coalesce"
271 | "nullif"
272 | "length"
273 | "lower"
274 | "upper"
275 | "trim"
276 | "substring"
277 | "count"
278 | "sum"
279 | "avg"
280 | "min"
281 | "max"
282 | "stddev"
283 | "stddev_samp"
284 | "stddev_pop"
285 | "variance"
286 | "row_number"
287 | "rank"
288 | "dense_rank"
289 | "lag"
290 | "lead"
291 | "first_value"
292 | "last_value"
293 )
294 }) && !function.uses_odbc_syntax
295 && matches!(function.parameters, FunctionArguments::None)
296 }
297 _ => false,
298 };
299 if allowed {
300 ControlFlow::Continue(())
301 } else {
302 ControlFlow::Break(())
303 }
304 }
305 fn post_visit_table_factor(&mut self, table: &mut TableFactor) -> ControlFlow<()> {
306 let TableFactor::Table {
307 name,
308 alias,
309 args,
310 with_hints,
311 version,
312 with_ordinality,
313 partitions,
314 json_path,
315 sample,
316 index_hints,
317 } = table
318 else {
319 return if matches!(
320 table,
321 TableFactor::Derived { sample: None, .. } | TableFactor::NestedJoin { .. }
322 ) {
323 ControlFlow::Continue(())
324 } else {
325 ControlFlow::Break(())
326 };
327 };
328 if args.is_some()
329 || !with_hints.is_empty()
330 || version.is_some()
331 || *with_ordinality
332 || !partitions.is_empty()
333 || json_path.is_some()
334 || sample.is_some()
335 || !index_hints.is_empty()
336 {
337 return ControlFlow::Break(());
338 }
339 let Some(key) = relation_name(name) else {
340 return ControlFlow::Break(());
341 };
342 if self
343 .scopes
344 .iter()
345 .rev()
346 .any(|(names, _)| names.contains(&key))
347 {
348 return ControlFlow::Continue(());
349 }
350 let Some(binding) = self.bindings.get(&key) else {
351 return ControlFlow::Break(());
352 };
353 let retained_alias = alias.clone().unwrap_or_else(|| TableAlias {
354 explicit: true,
355 name: Ident::with_quote('"', key.rsplit('\0').next().unwrap().to_string()),
356 columns: Vec::new(),
357 at: None,
358 });
359 match binding {
360 QueryBinding::Dataset(version) => {
361 self.inputs.insert(*version);
362 *name = ObjectName::from(vec![Ident::with_quote(
363 '"',
364 format!("input_{}", version.simple()),
365 )]);
366 *alias = Some(retained_alias);
367 }
368 QueryBinding::View(sql) => {
369 if self.views.contains(&key) {
370 return ControlFlow::Break(());
371 }
372 self.views.push(key);
373 let expanded = expand(
374 sql,
375 self.bindings,
376 self.views,
377 self.inputs,
378 self.remaining_bytes,
379 );
380 self.views.pop();
381 let Ok(subquery) = expanded else {
382 return ControlFlow::Break(());
383 };
384 *table = TableFactor::Derived {
385 lateral: false,
386 subquery,
387 alias: Some(retained_alias),
388 sample: None,
389 };
390 }
391 }
392 ControlFlow::Continue(())
393 }
394}
395fn relation_name(name: &ObjectName) -> Option<String> {
396 if name.0.is_empty() || name.0.len() > 3 {
397 return None;
398 }
399 name.0
400 .iter()
401 .map(|p| match p {
402 ObjectNamePart::Identifier(id) => Some(id.value.to_ascii_lowercase()),
403 _ => None,
404 })
405 .collect::<Option<Vec<_>>>()
406 .map(|parts| parts.join("\0"))
407}
408
409pub fn analyze_dataset_transform(
412 sql: &str,
413 study: ahri_tre_types::StudyId,
414 name: &ahri_tre_types::NcName,
415 version: ahri_tre_types::VersionId,
416 major: i32,
417 minor: i32,
418 patch: i32,
419) -> Result<RestrictedQueryPlan, LakeError> {
420 let relation =
421 crate::dataset_loading::dataset_table_relation(study, name.as_str(), major, minor, patch);
422 let bindings = BTreeMap::from([
423 (
424 name.as_str().to_ascii_lowercase(),
425 QueryBinding::Dataset(version.0),
426 ),
427 (
428 relation.table_name.to_ascii_lowercase(),
429 QueryBinding::Dataset(version.0),
430 ),
431 (
432 format!("{}\0{}", relation.schema_name, relation.table_name).to_ascii_lowercase(),
433 QueryBinding::Dataset(version.0),
434 ),
435 (
436 format!(
437 "{}\0{}\0{}",
438 crate::LAKE_ALIAS,
439 relation.schema_name,
440 relation.table_name
441 )
442 .to_ascii_lowercase(),
443 QueryBinding::Dataset(version.0),
444 ),
445 ]);
446 analyze_disclosure_query(sql, &bindings)
447}
448
449fn single_name(name: &ObjectName) -> Option<String> {
450 match name.0.as_slice() {
451 [ObjectNamePart::Identifier(id)] => Some(id.value.to_ascii_lowercase()),
452 _ => None,
453 }
454}
455fn safe_set(set: &SetExpr) -> bool {
456 match set {
457 SetExpr::Select(_) | SetExpr::Query(_) | SetExpr::Values(_) => true,
458 SetExpr::SetOperation { left, right, .. } => safe_set(left) && safe_set(right),
459 _ => false,
460 }
461}