Skip to main content

ahri_tre_lake/
sql_source.rs

1//! Closed read-query planning for one selected external database.
2use crate::{LakeError, QueryBinding, analyze_disclosure_query};
3use sqlparser::{ast::*, dialect::DuckDbDialect, parser::Parser};
4use std::{collections::BTreeMap, ops::ControlFlow};
5
6pub struct SqlSourcePlan {
7    relations: Vec<Vec<String>>,
8    sql: String,
9}
10
11impl SqlSourcePlan {
12    pub fn parse(sql: &str) -> Result<Self, LakeError> {
13        // Bound parsing before constructing a recursive AST, as for disclosure.
14        if sql.len() > 65536 {
15            return Err(LakeError::QueryNotAdmissible);
16        }
17        let dialect = DuckDbDialect {};
18        let tokens = sqlparser::tokenizer::Tokenizer::new(&dialect, sql)
19            .tokenize()
20            .map_err(|_| LakeError::QueryNotAdmissible)?;
21        if tokens
22            .iter()
23            .filter(|t| !matches!(t, sqlparser::tokenizer::Token::Whitespace(_)))
24            .count()
25            > 256
26        {
27            return Err(LakeError::QueryNotAdmissible);
28        }
29        let statements = Parser::new(&dialect)
30            .with_recursion_limit(64)
31            .try_with_sql(sql)
32            .and_then(|mut parser| parser.parse_statements())
33            .map_err(|_| LakeError::QueryNotAdmissible)?;
34        let mut names = SourceNames::default();
35        if Visit::visit(&statements, &mut names).is_break() {
36            return Err(LakeError::QueryNotAdmissible);
37        }
38        names.relations.retain(|key, _| !names.ctes.contains(key));
39        let bindings = names
40            .relations
41            .keys()
42            .map(|key| (key.clone(), QueryBinding::Dataset(uuid::Uuid::new_v4())))
43            .collect();
44        // Use the shared closed-query admission rules without adopting its
45        // Dataset alias rewriting: external SQL identifiers retain engine rules.
46        analyze_disclosure_query(sql, &bindings)?;
47        Ok(Self {
48            relations: names
49                .relations
50                .into_values()
51                .map(|parts| parts.into_iter().map(|id| id.value).collect())
52                .collect(),
53            sql: statements[0].to_string(),
54        })
55    }
56
57    pub fn relations(&self) -> &[Vec<String>] {
58        &self.relations
59    }
60
61    /// Validated and regenerated SQL. Unqualified tables use the connector's
62    /// fixed default schema; three-part database retargeting is never accepted.
63    pub fn query(&self) -> &str {
64        &self.sql
65    }
66
67    pub fn in_catalog(&self, catalog: &str, default_schema: &str) -> Result<String, LakeError> {
68        let mut query = Parser::parse_sql(&DuckDbDialect {}, &self.sql)
69            .map_err(|_| LakeError::QueryNotAdmissible)?;
70        let mut names = SourceNames::default();
71        if Visit::visit(&query, &mut names).is_break() {
72            return Err(LakeError::QueryNotAdmissible);
73        }
74        names.relations.retain(|key, _| !names.ctes.contains(key));
75        let targets = names
76            .relations
77            .into_iter()
78            .map(|(key, parts)| {
79                let mut target = vec![catalog.to_string()];
80                if parts.len() == 1 {
81                    target.push(default_schema.to_string());
82                }
83                // PostgreSQL folds unquoted identifiers; quoted identifiers keep
84                // their exact spelling. DuckDB/SQLite are case insensitive here.
85                target.extend(parts.into_iter().map(|id| {
86                    if id.quote_style.is_some() {
87                        id.value
88                    } else {
89                        id.value.to_ascii_lowercase()
90                    }
91                }));
92                (key, target)
93            })
94            .collect();
95        let _ = VisitMut::visit(&mut query, &mut SourceTargets(targets));
96        Ok(query[0].to_string())
97    }
98}
99
100#[derive(Default)]
101struct SourceNames {
102    relations: BTreeMap<String, Vec<Ident>>,
103    ctes: std::collections::BTreeSet<String>,
104}
105impl Visitor for SourceNames {
106    type Break = ();
107    fn pre_visit_query(&mut self, query: &Query) -> ControlFlow<()> {
108        if let Some(with) = &query.with {
109            self.ctes.extend(
110                with.cte_tables
111                    .iter()
112                    .map(|cte| cte.alias.name.value.to_ascii_lowercase()),
113            );
114        }
115        ControlFlow::Continue(())
116    }
117    fn pre_visit_relation(&mut self, name: &ObjectName) -> ControlFlow<()> {
118        let parts = name
119            .0
120            .iter()
121            .map(|part| match part {
122                ObjectNamePart::Identifier(id) => Some(id.clone()),
123                _ => None,
124            })
125            .collect::<Option<Vec<_>>>();
126        let Some(parts) = parts.filter(|parts| (1..=2).contains(&parts.len())) else {
127            return ControlFlow::Break(());
128        };
129        let key = parts
130            .iter()
131            .map(|id| id.value.to_ascii_lowercase())
132            .collect::<Vec<_>>()
133            .join("\0");
134        // The shared admission binder is case insensitive. Refuse a collision
135        // instead of silently binding distinct source relations to one input.
136        if self
137            .relations
138            .get(&key)
139            .is_some_and(|prior| prior != &parts)
140        {
141            return ControlFlow::Break(());
142        }
143        self.relations.insert(key, parts);
144        ControlFlow::Continue(())
145    }
146}
147
148struct SourceTargets(BTreeMap<String, Vec<String>>);
149impl VisitorMut for SourceTargets {
150    type Break = ();
151    fn pre_visit_relation(&mut self, name: &mut ObjectName) -> ControlFlow<()> {
152        let key = name
153            .0
154            .iter()
155            .map(|part| match part {
156                ObjectNamePart::Identifier(id) => id.value.to_ascii_lowercase(),
157                _ => String::new(),
158            })
159            .collect::<Vec<_>>()
160            .join("\0");
161        if let Some(parts) = self.0.get(&key) {
162            *name = ObjectName::from(
163                parts
164                    .iter()
165                    .map(|part| Ident::with_quote('"', part))
166                    .collect::<Vec<_>>(),
167            );
168        }
169        ControlFlow::Continue(())
170    }
171}
172
173/// Local source authority stays in the acquiring process and its private worker.
174/// This type is never part of a Trusted request or a persisted source record.
175#[derive(serde::Serialize, serde::Deserialize)]
176pub enum SqlReadSource {
177    Duckdb { path: std::path::PathBuf },
178    Sqlite { path: std::path::PathBuf },
179    Postgresql { endpoint: SqlDatabaseEndpoint },
180    Mssql { endpoint: SqlDatabaseEndpoint },
181}
182
183/// One policy-approved socket and its original TLS identity. No connection
184/// string, database override, initialization SQL, or ambient login is accepted.
185#[derive(serde::Serialize, serde::Deserialize)]
186pub struct SqlDatabaseEndpoint {
187    pub address: std::net::IpAddr,
188    pub host: String,
189    pub port: u16,
190    pub database: String,
191    pub ca_certificates: Option<String>,
192}
193
194#[derive(serde::Serialize, serde::Deserialize)]
195pub struct SqlSourceRead {
196    pub source: SqlReadSource,
197    pub sql: String,
198    pub budgets: ahri_tre_types::DisclosureBudgets,
199    #[serde(skip)]
200    pub credentials: Option<ahri_tre_secrets::SourceCredentials>,
201}
202impl SqlSourceRead {
203    #[doc(hidden)]
204    pub fn read_private_input(mut reader: impl std::io::Read) -> Result<Self, LakeError> {
205        let mut length = [0; 4];
206        reader.read_exact(&mut length)?;
207        let length = u32::from_be_bytes(length) as usize;
208        if length > 131072 {
209            return Err(LakeError::QueryNotAdmissible);
210        }
211        let mut bytes = vec![0; length];
212        reader.read_exact(&mut bytes)?;
213        let mut request: Self =
214            serde_json::from_slice(&bytes).map_err(|_| LakeError::QueryNotAdmissible)?;
215        let kind = match request.source {
216            SqlReadSource::Postgresql { .. } => ahri_tre_types::SourceAuthentication::PostgreSql,
217            SqlReadSource::Mssql { .. } => ahri_tre_types::SourceAuthentication::MsSql,
218            _ => ahri_tre_types::SourceAuthentication::None,
219        };
220        request.credentials = Some(
221            ahri_tre_secrets::SourceCredentials::read(kind, reader)
222                .map_err(|_| LakeError::QueryNotAdmissible)?,
223        );
224        Ok(request)
225    }
226    /// Return a verified finite Parquet file only after the worker and its source
227    /// connection have ended. Failure drops the anonymous partial output.
228    pub fn read(
229        self,
230        executable: &std::path::Path,
231        directory: &std::path::Path,
232        deadline: std::time::Instant,
233        cancelled: std::sync::Arc<std::sync::atomic::AtomicBool>,
234    ) -> Result<std::fs::File, LakeError> {
235        use std::{
236            io::{Seek, Write},
237            process::{Command, Stdio},
238            sync::atomic::Ordering,
239            time::{Duration, Instant},
240        };
241        SqlSourcePlan::parse(&self.sql)?;
242        if !self.budgets.within(&self.budgets)
243            || Instant::now() >= deadline
244            || cancelled.load(Ordering::Acquire)
245        {
246            return Err(LakeError::QueryNotAdmissible);
247        }
248        let encoded = serde_json::to_vec(&self).map_err(|_| LakeError::QueryNotAdmissible)?;
249        if encoded.len() > 131072 {
250            return Err(LakeError::QueryNotAdmissible);
251        }
252        let mut output = tempfile::tempfile_in(directory)?;
253        let mut command = Command::new(executable);
254        crate::disclosure_inputs::parent_bound(&mut command)?;
255        command.env_clear();
256        // Only the engine installation and loader paths survive. No provider,
257        // PostgreSQL, credential, proxy, or client configuration environment.
258        for key in ["HOME", "LD_LIBRARY_PATH", "DYLD_LIBRARY_PATH"] {
259            if let Some(value) = std::env::var_os(key) {
260                command.env(key, value);
261            }
262        }
263        let mut child = command
264            .arg("--acquire-sql")
265            .stdin(Stdio::piped())
266            .stdout(output.try_clone()?)
267            .stderr(Stdio::null())
268            .spawn()?;
269        let mut input = child.stdin.take().ok_or(LakeError::QueryNotAdmissible)?;
270        let success = std::thread::scope(|scope| {
271            let sent = scope.spawn(move || {
272                input.write_all(&(encoded.len() as u32).to_be_bytes())?;
273                input.write_all(&encoded)?;
274                if let Some(credentials) = self.credentials {
275                    credentials.expose(|bytes| input.write_all(bytes))?;
276                }
277                Ok::<_, std::io::Error>(())
278            });
279            let status = loop {
280                match child.try_wait() {
281                    Ok(Some(status)) => break status.success(),
282                    Ok(None) if Instant::now() < deadline && !cancelled.load(Ordering::Acquire) => {
283                        std::thread::sleep(Duration::from_millis(10))
284                    }
285                    _ => {
286                        let _ = child.kill();
287                        let _ = child.wait();
288                        break false;
289                    }
290                }
291            };
292            status && sent.join().is_ok_and(|result| result.is_ok())
293        });
294        if !success || Instant::now() >= deadline || cancelled.load(Ordering::Acquire) {
295            return Err(LakeError::QueryNotAdmissible);
296        }
297        output.rewind()?;
298        Ok(output)
299    }
300}