1use std::path::PathBuf;
2
3use ahri_tre_pgmeta::{PgMetadataAdapter, PgMetadataConnection};
4use ahri_tre_types::{TreValueType, ValueTypeId};
5use duckdb::Connection;
6
7use crate::{
8 AnalyzeSqlSourceRequest, AppError, OpenDuckDbSourceRequest, OpenMsSqlSourceRequest,
9 OpenPostgresSourceRequest, OpenSqliteSourceRequest, SqlSourceAnalysis, SqlSourceColumnAnalysis,
10 SqlSourceVocabularyAnalysis, SqlSourceVocabularyItemAnalysis,
11};
12
13const SOURCE_ALIAS: &str = "__tre_sql_source";
14
15pub enum SqlSourceConnection {
16 DuckDb {
17 path: PathBuf,
18 connection: Connection,
19 description: String,
20 },
21 Sqlite {
22 path: PathBuf,
23 description: String,
24 },
25 Postgres {
26 connection_info: String,
27 description: String,
28 schema: Option<String>,
29 connection: Box<PgMetadataConnection>,
30 },
31 MsSql {
32 connection_string: String,
33 description: String,
34 schema: Option<String>,
35 },
36}
37
38impl std::fmt::Debug for SqlSourceConnection {
39 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
40 f.debug_struct("SqlSourceConnection")
41 .field("kind", &self.kind())
42 .field("description", &self.description())
43 .finish_non_exhaustive()
44 }
45}
46
47impl SqlSourceConnection {
48 pub fn open_duckdb(request: OpenDuckDbSourceRequest) -> Result<Self, AppError> {
49 let connection = Connection::open(&request.path)
50 .map_err(|error| infrastructure_error("open DuckDB SQL source", error))?;
51 Ok(Self::DuckDb {
52 description: format!("duckdb:{}", request.path.display()),
53 path: request.path,
54 connection,
55 })
56 }
57
58 pub fn open_sqlite(request: OpenSqliteSourceRequest) -> Result<Self, AppError> {
59 if !request.path.is_file() {
60 return Err(AppError::Validation(format!(
61 "SQLite SQL source not found: {}",
62 request.path.display()
63 )));
64 }
65 Ok(Self::Sqlite {
66 description: format!("sqlite:{}", request.path.display()),
67 path: request.path,
68 })
69 }
70
71 pub fn open_postgres(request: OpenPostgresSourceRequest) -> Result<Self, AppError> {
72 let adapter = PgMetadataAdapter::direct(request.config.clone());
73 let connection = adapter
74 .connect()
75 .map_err(|error| infrastructure_error("open PostgreSQL SQL source", error))?;
76 Ok(Self::Postgres {
77 connection_info: request.config.conninfo(),
78 description: request.config.safe_conninfo(),
79 schema: request.schema,
80 connection: Box::new(connection),
81 })
82 }
83
84 pub fn open_mssql(request: OpenMsSqlSourceRequest) -> Result<Self, AppError> {
85 if request.connection_string.trim().is_empty() {
86 return Err(AppError::Validation(
87 "MSSQL source connection string must not be empty".to_string(),
88 ));
89 }
90 Ok(Self::MsSql {
91 description: request
92 .safe_description
93 .unwrap_or_else(|| "mssql:<redacted>".to_string()),
94 connection_string: request.connection_string,
95 schema: request.schema,
96 })
97 }
98
99 pub fn kind(&self) -> &'static str {
100 match self {
101 Self::DuckDb { .. } => "duckdb",
102 Self::Sqlite { .. } => "sqlite",
103 Self::Postgres { .. } => "postgresql",
104 Self::MsSql { .. } => "mssql",
105 }
106 }
107
108 pub fn description(&self) -> &str {
109 match self {
110 Self::DuckDb { description, .. }
111 | Self::Sqlite { description, .. }
112 | Self::Postgres { description, .. }
113 | Self::MsSql { description, .. } => description,
114 }
115 }
116
117 pub fn analyze(
118 &mut self,
119 request: AnalyzeSqlSourceRequest,
120 ) -> Result<SqlSourceAnalysis, AppError> {
121 let sql = read_only_select(&request.sql)?;
122 let shape = QueryShape::parse(&sql);
123 let analysis_connection = Connection::open_in_memory()
124 .map_err(|error| infrastructure_error("open SQL analysis DuckDB", error))?;
125 let prepared = self.prepare_query_connection(&analysis_connection, &shape)?;
126 let columns = describe_query(&analysis_connection, &sql)?;
127 let mut warnings = shape.warnings.clone();
128 warnings.extend(prepared.warnings);
129 let mut analyzed = Vec::with_capacity(columns.len());
130 let single_table = shape.tables.len() == 1;
131
132 for described in columns {
133 let lineage = shape.lineage_for(&described.name, single_table);
134 let mut value_type_id = value_type_id_for_sql(&described.sql_type);
135 let mut value_type = value_type_name(value_type_id)?.to_string();
136 let mut vocabulary = None;
137 if let Some(lineage) = &lineage {
138 match infer_vocabulary(
139 &analysis_connection,
140 self,
141 &lineage.table_name,
142 &described.name,
143 ) {
144 Ok(Some(inferred)) => {
145 value_type_id = TreValueType::Enumeration.value_type_id();
146 value_type = TreValueType::Enumeration.datastore_name().to_string();
147 vocabulary = Some(inferred);
148 }
149 Ok(None) => {}
150 Err(message) => warnings.push(message),
151 }
152 } else {
153 warnings.push(format!(
154 "metadata lineage for column {} is uncertain; registered variable will use name and value type only",
155 described.name
156 ));
157 }
158
159 let description = lineage.as_ref().and_then(|lineage| {
160 self.column_description(&lineage.table_name, &lineage.column_name)
161 .ok()
162 .flatten()
163 });
164 analyzed.push(SqlSourceColumnAnalysis {
165 name: described.name,
166 sql_type: described.sql_type,
167 value_type_id,
168 value_type,
169 description,
170 source_table: lineage.as_ref().map(|lineage| lineage.table_name.clone()),
171 source_column: lineage.as_ref().map(|lineage| lineage.column_name.clone()),
172 vocabulary,
173 });
174 }
175
176 Ok(SqlSourceAnalysis {
177 source_kind: self.kind().to_string(),
178 source_description: self.description().to_string(),
179 columns: analyzed,
180 warnings,
181 })
182 }
183
184 pub fn prepare_lake_query(
185 &self,
186 connection: &Connection,
187 sql: &str,
188 ) -> Result<PreparedSqlSourceQuery, AppError> {
189 let sql = read_only_select(sql)?;
190 let shape = QueryShape::parse(&sql);
191 let prepared = self.prepare_query_connection(connection, &shape)?;
192 let mut warnings = shape.warnings;
193 warnings.extend(prepared.warnings);
194 Ok(PreparedSqlSourceQuery { sql, warnings })
195 }
196
197 fn prepare_query_connection(
198 &self,
199 connection: &Connection,
200 shape: &QueryShape,
201 ) -> Result<PreparedSqlSourceQuery, AppError> {
202 self.attach(connection)?;
203 let mut warnings = Vec::new();
204 for table in &shape.tables {
205 if table.name.contains('.') {
206 warnings.push(format!(
207 "query table {} is schema-qualified; automatic temporary view creation is skipped",
208 table.name
209 ));
210 continue;
211 }
212 let view_sql = quote_identifier(&table.name);
213 let source_sql = source_relation_sql(SOURCE_ALIAS, &table.name);
214 connection
215 .execute_batch(&format!(
216 "CREATE OR REPLACE TEMP VIEW {view_sql} AS SELECT * FROM {source_sql};"
217 ))
218 .map_err(|error| {
219 infrastructure_error("prepare SQL source temporary view", error)
220 })?;
221 }
222 Ok(PreparedSqlSourceQuery {
223 sql: String::new(),
224 warnings,
225 })
226 }
227
228 fn attach(&self, connection: &Connection) -> Result<(), AppError> {
229 match self {
230 Self::DuckDb { path, .. } => connection
231 .execute_batch(&format!(
232 "ATTACH OR REPLACE {} AS {} (READ_ONLY);",
233 sql_literal(&path.display().to_string()),
234 quote_identifier(SOURCE_ALIAS)
235 ))
236 .map_err(|error| infrastructure_error("attach DuckDB SQL source", error))?,
237 Self::Sqlite { path, .. } => connection
238 .execute_batch(&format!(
239 "INSTALL sqlite; LOAD sqlite; ATTACH OR REPLACE {} AS {} (TYPE sqlite, READ_ONLY);",
240 sql_literal(&path.display().to_string()),
241 quote_identifier(SOURCE_ALIAS)
242 ))
243 .map_err(|error| infrastructure_error("attach SQLite SQL source", error))?,
244 Self::Postgres {
245 connection_info,
246 schema,
247 ..
248 } => {
249 let schema_sql = schema
250 .as_deref()
251 .filter(|value| !value.trim().is_empty())
252 .map(|value| format!(", SCHEMA {}", sql_literal(value)))
253 .unwrap_or_default();
254 connection
255 .execute_batch(&format!(
256 "INSTALL postgres; LOAD postgres; ATTACH OR REPLACE {} AS {} (TYPE postgres, READ_ONLY{});",
257 sql_literal(connection_info),
258 quote_identifier(SOURCE_ALIAS),
259 schema_sql
260 ))
261 .map_err(|error| infrastructure_error("attach PostgreSQL SQL source", error))?;
262 }
263 Self::MsSql {
264 connection_string,
265 schema,
266 ..
267 } => {
268 let schema_sql = schema
269 .as_deref()
270 .filter(|value| !value.trim().is_empty())
271 .map(|value| format!(", SCHEMA {}", sql_literal(value)))
272 .unwrap_or_default();
273 connection
274 .execute_batch(&format!(
275 "INSTALL sqlserver FROM community; LOAD sqlserver; ATTACH OR REPLACE {} AS {} (TYPE sqlserver, READ_ONLY{});",
276 sql_literal(connection_string),
277 quote_identifier(SOURCE_ALIAS),
278 schema_sql
279 ))
280 .map_err(|error| infrastructure_error("attach MSSQL SQL source", error))?;
281 }
282 }
283 Ok(())
284 }
285
286 fn column_description(
287 &mut self,
288 table_name: &str,
289 column_name: &str,
290 ) -> Result<Option<String>, AppError> {
291 match self {
292 Self::Postgres {
293 connection, schema, ..
294 } => {
295 let schema = schema.as_deref().unwrap_or("public");
296 let row = connection
297 .client()
298 .query_opt(
299 "
300 SELECT pg_catalog.col_description(
301 format('%I.%I', table_schema, table_name)::regclass::oid,
302 ordinal_position
303 ) AS description
304 FROM information_schema.columns
305 WHERE table_schema = $1
306 AND table_name = $2
307 AND column_name = $3
308 ",
309 &[&schema, &table_name, &column_name],
310 )
311 .map_err(|error| {
312 infrastructure_error("read PostgreSQL column description", error)
313 })?;
314 Ok(row.and_then(|row| row.get::<_, Option<String>>("description")))
315 }
316 _ => Ok(None),
317 }
318 }
319}
320
321#[derive(Debug, Clone, PartialEq, Eq)]
322pub struct PreparedSqlSourceQuery {
323 pub sql: String,
324 pub warnings: Vec<String>,
325}
326
327#[derive(Debug, Clone)]
328struct DescribedColumn {
329 name: String,
330 sql_type: String,
331}
332
333#[derive(Debug, Clone)]
334struct QueryShape {
335 tables: Vec<QueryTable>,
336 columns: Vec<QueryColumn>,
337 warnings: Vec<String>,
338}
339
340#[derive(Debug, Clone)]
341struct QueryTable {
342 name: String,
343 alias: Option<String>,
344}
345
346#[derive(Debug, Clone)]
347struct QueryColumn {
348 output_name: String,
349 source_name: Option<String>,
350 qualifier: Option<String>,
351 wildcard: bool,
352}
353
354#[derive(Debug, Clone)]
355struct ColumnLineage {
356 table_name: String,
357 column_name: String,
358}
359
360impl QueryShape {
361 fn parse(sql: &str) -> Self {
362 let mut warnings = Vec::new();
363 let tokens = tokenize(sql);
364 let from_index = tokens
365 .iter()
366 .position(|token| token.eq_ignore_ascii_case("from"));
367 let Some(from_index) = from_index else {
368 return Self {
369 tables: Vec::new(),
370 columns: Vec::new(),
371 warnings: vec![
372 "query metadata discovery could not find a top-level FROM clause".to_string(),
373 ],
374 };
375 };
376 let select_text = sql
377 .get(6..sql.to_ascii_lowercase().find(" from ").unwrap_or(sql.len()))
378 .unwrap_or("")
379 .trim();
380 let columns = parse_select_columns(select_text, &mut warnings);
381 let mut tables = Vec::new();
382 let mut index = from_index;
383 while index < tokens.len() {
384 if tokens[index].eq_ignore_ascii_case("from")
385 || tokens[index].eq_ignore_ascii_case("join")
386 {
387 index += 1;
388 if index >= tokens.len() {
389 break;
390 }
391 let name = tokens[index].clone();
392 if name == "(" {
393 warnings.push(
394 "query contains a nested source; rich metadata discovery is limited"
395 .to_string(),
396 );
397 index += 1;
398 continue;
399 }
400 let alias = tokens.get(index + 1).and_then(|candidate| {
401 if is_table_alias_stop(candidate) {
402 None
403 } else if candidate.eq_ignore_ascii_case("as") {
404 tokens.get(index + 2).cloned()
405 } else {
406 Some(candidate.clone())
407 }
408 });
409 tables.push(QueryTable { name, alias });
410 }
411 index += 1;
412 }
413 Self {
414 tables,
415 columns,
416 warnings,
417 }
418 }
419
420 fn lineage_for(&self, column_name: &str, single_table: bool) -> Option<ColumnLineage> {
421 if single_table {
422 return self.tables.first().map(|table| ColumnLineage {
423 table_name: table.name.clone(),
424 column_name: column_name.to_string(),
425 });
426 }
427 self.columns.iter().find_map(|column| {
428 if column.wildcard || column.output_name != column_name {
429 return None;
430 }
431 let qualifier = column.qualifier.as_deref()?;
432 let table = self.tables.iter().find(|table| {
433 table.alias.as_deref() == Some(qualifier) || table.name == qualifier
434 })?;
435 Some(ColumnLineage {
436 table_name: table.name.clone(),
437 column_name: column
438 .source_name
439 .clone()
440 .unwrap_or_else(|| column_name.to_string()),
441 })
442 })
443 }
444}
445
446fn read_only_select(sql: &str) -> Result<String, AppError> {
447 let trimmed = sql.trim().trim_end_matches(';').trim();
448 let lowered = trimmed.to_ascii_lowercase();
449 let forbidden = [
450 ";",
451 "insert ",
452 "update ",
453 "delete ",
454 "drop ",
455 "create ",
456 "alter ",
457 "merge ",
458 "truncate ",
459 "grant ",
460 "revoke ",
461 "exec ",
462 "execute ",
463 "call ",
464 "copy ",
465 "attach ",
466 "detach ",
467 ];
468 if trimmed.is_empty()
469 || !lowered.starts_with("select ")
470 || forbidden.iter().any(|needle| lowered.contains(needle))
471 {
472 return Err(AppError::Validation(
473 "sql-to-dataset requires a single read-only SELECT statement".to_string(),
474 ));
475 }
476 Ok(trimmed.to_string())
477}
478
479fn describe_query(connection: &Connection, sql: &str) -> Result<Vec<DescribedColumn>, AppError> {
480 let mut statement = connection
481 .prepare(&format!("DESCRIBE SELECT * FROM ({sql}) src"))
482 .map_err(|error| infrastructure_error("describe SQL source query", error))?;
483 let rows = statement
484 .query_map([], |row| {
485 Ok(DescribedColumn {
486 name: row.get::<_, String>(0)?,
487 sql_type: row.get::<_, String>(1)?,
488 })
489 })
490 .map_err(|error| infrastructure_error("read SQL source query description", error))?;
491 let mut columns = Vec::new();
492 for row in rows {
493 columns.push(
494 row.map_err(|error| {
495 infrastructure_error("decode SQL source query description", error)
496 })?,
497 );
498 }
499 Ok(columns)
500}
501
502fn infer_vocabulary(
503 connection: &Connection,
504 source: &SqlSourceConnection,
505 _source_table: &str,
506 column_name: &str,
507) -> Result<Option<SqlSourceVocabularyAnalysis>, String> {
508 let candidates = code_table_candidates(column_name);
509 for candidate in candidates {
510 let relation = source_relation_sql(SOURCE_ALIAS, &candidate);
511 let count =
512 match connection.query_row(&format!("SELECT count(*) FROM {relation}"), [], |row| {
513 row.get::<_, i64>(0)
514 }) {
515 Ok(count) => count,
516 Err(_) => continue,
517 };
518 if !(0..500).contains(&count) {
519 continue;
520 }
521 let columns = describe_relation_columns(connection, &relation)
522 .map_err(|error| format!("could not inspect code table {candidate}: {error}"))?;
523 let Some(value_column) = columns.first().map(|column| column.name.clone()) else {
524 continue;
525 };
526 let code_column = columns
527 .iter()
528 .find(|column| matches!(column.name.as_str(), "label" | "name" | "code"))
529 .map(|column| column.name.clone())
530 .unwrap_or_else(|| value_column.clone());
531 let description_column = columns
532 .iter()
533 .find(|column| column.name == "description")
534 .map(|column| column.name.clone());
535 let description_sql = description_column
536 .as_ref()
537 .map(|column| quote_identifier(column))
538 .unwrap_or_else(|| "NULL".to_string());
539 let mut statement = connection
540 .prepare(&format!(
541 "SELECT {}, CAST({} AS VARCHAR), CAST({} AS VARCHAR) FROM {relation} ORDER BY 1",
542 quote_identifier(&value_column),
543 quote_identifier(&code_column),
544 description_sql
545 ))
546 .map_err(|error| format!("could not read code table {candidate}: {error}"))?;
547 let rows = statement
548 .query_map([], |row| {
549 Ok(SqlSourceVocabularyItemAnalysis {
550 value: row.get::<_, i32>(0)?,
551 code: row.get::<_, String>(1)?,
552 description: row.get::<_, Option<String>>(2)?,
553 })
554 })
555 .map_err(|error| format!("could not query code table {candidate}: {error}"))?;
556 let mut items = Vec::new();
557 for row in rows {
558 items.push(
559 row.map_err(|error| format!("could not decode code table {candidate}: {error}"))?,
560 );
561 }
562 if !items.is_empty() {
563 return Ok(Some(SqlSourceVocabularyAnalysis {
564 name: candidate,
565 description: Some(format!("Inferred from {} source code table", source.kind())),
566 items,
567 }));
568 }
569 }
570 Ok(None)
571}
572
573fn describe_relation_columns(
574 connection: &Connection,
575 relation: &str,
576) -> Result<Vec<DescribedColumn>, duckdb::Error> {
577 let mut statement = connection.prepare(&format!("DESCRIBE SELECT * FROM {relation}"))?;
578 let rows = statement.query_map([], |row| {
579 Ok(DescribedColumn {
580 name: row.get::<_, String>(0)?,
581 sql_type: row.get::<_, String>(1)?,
582 })
583 })?;
584 let mut columns = Vec::new();
585 for row in rows {
586 columns.push(row?);
587 }
588 Ok(columns)
589}
590
591fn code_table_candidates(column_name: &str) -> Vec<String> {
592 let mut candidates = vec![column_name.to_string()];
593 if !column_name.ends_with('s') {
594 candidates.push(format!("{column_name}s"));
595 }
596 if let Some(stem) = column_name.strip_suffix('y') {
597 candidates.push(format!("{stem}ies"));
598 }
599 candidates
600}
601
602fn parse_select_columns(select_text: &str, warnings: &mut Vec<String>) -> Vec<QueryColumn> {
603 split_top_level_commas(select_text)
604 .into_iter()
605 .map(|item| {
606 let item = item.trim();
607 if item == "*" {
608 return QueryColumn {
609 output_name: "*".to_string(),
610 source_name: None,
611 qualifier: None,
612 wildcard: true,
613 };
614 }
615 let before_alias = item
616 .split_once(" as ")
617 .map(|(left, _)| left)
618 .unwrap_or(item)
619 .trim();
620 let output_name = item
621 .split_once(" as ")
622 .map(|(_, right)| right.trim().trim_matches('"').to_string())
623 .unwrap_or_else(|| {
624 before_alias
625 .rsplit('.')
626 .next()
627 .unwrap_or(before_alias)
628 .trim_matches('"')
629 .to_string()
630 });
631 if before_alias.contains('(') {
632 warnings.push(format!(
633 "query column {output_name} is computed; rich metadata discovery is limited"
634 ));
635 return QueryColumn {
636 output_name,
637 source_name: None,
638 qualifier: None,
639 wildcard: false,
640 };
641 }
642 let (qualifier, source_name) = before_alias
643 .rsplit_once('.')
644 .map(|(qualifier, name)| {
645 (
646 Some(qualifier.trim_matches('"').to_string()),
647 Some(name.trim_matches('"').to_string()),
648 )
649 })
650 .unwrap_or((None, Some(before_alias.trim_matches('"').to_string())));
651 QueryColumn {
652 output_name,
653 source_name,
654 qualifier,
655 wildcard: false,
656 }
657 })
658 .collect()
659}
660
661fn split_top_level_commas(value: &str) -> Vec<String> {
662 let mut items = Vec::new();
663 let mut depth = 0_i32;
664 let mut start = 0_usize;
665 for (index, character) in value.char_indices() {
666 match character {
667 '(' => depth += 1,
668 ')' => depth -= 1,
669 ',' if depth == 0 => {
670 items.push(value[start..index].to_string());
671 start = index + 1;
672 }
673 _ => {}
674 }
675 }
676 if start < value.len() {
677 items.push(value[start..].to_string());
678 }
679 items
680}
681
682fn tokenize(sql: &str) -> Vec<String> {
683 sql.replace(',', " ")
684 .replace('(', " ( ")
685 .replace(')', " ) ")
686 .split_whitespace()
687 .map(|token| token.trim_matches('"').to_string())
688 .collect()
689}
690
691fn is_table_alias_stop(token: &str) -> bool {
692 matches!(
693 token.to_ascii_lowercase().as_str(),
694 "on" | "where"
695 | "join"
696 | "left"
697 | "right"
698 | "inner"
699 | "outer"
700 | "full"
701 | "cross"
702 | "group"
703 | "order"
704 | "limit"
705 | "offset"
706 | "union"
707 )
708}
709
710fn value_type_id_for_sql(sql_type: &str) -> ValueTypeId {
711 let lowered = sql_type.to_ascii_lowercase();
712 if lowered.contains("int") {
713 ValueTypeId(1)
714 } else if lowered.contains("float")
715 || lowered.contains("double")
716 || lowered.contains("real")
717 || lowered.contains("decimal")
718 || lowered.contains("numeric")
719 {
720 ValueTypeId(2)
721 } else if lowered == "date" {
722 ValueTypeId(4)
723 } else if lowered.contains("timestamp") || lowered.contains("datetime") {
724 ValueTypeId(5)
725 } else if lowered == "time" || lowered.starts_with("time(") {
726 ValueTypeId(6)
727 } else if lowered.contains("bool") || lowered == "bit" {
728 ValueTypeId(1)
729 } else {
730 ValueTypeId(3)
731 }
732}
733
734pub fn value_type_name(value_type_id: ValueTypeId) -> Result<&'static str, AppError> {
735 TreValueType::from_value_type_id(value_type_id)
736 .map(TreValueType::datastore_name)
737 .map_err(|error| AppError::Validation(error.to_string()))
738}
739
740fn source_relation_sql(alias: &str, table_name: &str) -> String {
741 format!(
742 "{}.{}",
743 quote_identifier(alias),
744 quote_identifier(table_name)
745 )
746}
747
748fn quote_identifier(value: &str) -> String {
749 format!("\"{}\"", value.replace('"', "\"\""))
750}
751
752fn sql_literal(value: &str) -> String {
753 format!("'{}'", value.replace('\'', "''"))
754}
755
756fn infrastructure_error(operation: &str, error: impl std::fmt::Display) -> AppError {
757 AppError::Infrastructure(format!("{operation} failed: {error}"))
758}