ahri_tre_lake/
sql_source.rs1use 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 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 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 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 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 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#[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#[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 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 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}