Skip to main content

ahri_tre_app/
sql_source.rs

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}