Skip to main content

ahri_tre_lake/
restricted_query.rs

1//! Dependency expansion runs before optimization. Accepted SQL names only
2//! adapter-owned input tables and must execute in the restricted worker.
3use 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    /// Regenerated from the accepted AST, never the original untrusted text.
26    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    // A left-associative expression may have shallow parser recursion and a
81    // deep AST. Bound tokens before allocating any recursive syntax tree.
82    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                // A nonrecursive CTE sees only earlier CTEs and outer scopes.
153                // Visit here, then hide WITH from the generated traversal to
154                // avoid visiting and rewriting its input relations twice.
155                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
409/// Bind the existing Dataset SQL-transform source using its canonical relation
410/// and logical alias. Any additional or external relation must fail resolution.
411pub 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}