diff --git a/README.md b/README.md index 67ca057..52fcfb1 100644 --- a/README.md +++ b/README.md @@ -122,7 +122,13 @@ Options: -h, --help Print help ``` -Supported dialects: `generic`, `ansi`, `postgresql`, `mysql`, `hive`, `databricks`, `snowflake`, `bigquery`. +Supported dialects — one per dialect that [sqlparser](https://crates.io/crates/sqlparser) exposes: + +`generic`, `ansi`, `postgresql`, `mysql`, `hive`, `databricks`, `snowflake`, `bigquery`, `duckdb`, `redshift`, `spark`, `clickhouse`, `sqlite`, `mssql`, `oracle`, `teradata`. + +The aliases `postgres`, `sparksql`, `tsql`, and `sqlserver` are also accepted. Names are case-insensitive. + +Dialects without a dedicated `sqlparser` implementation — Trino/Presto, for example — generally parse with `generic`, which accepts a superset of most grammars. ## CatalogProvider diff --git a/sqllineage-python/sqllineage.pyi b/sqllineage-python/sqllineage.pyi index 04f025b..c6ccfc9 100644 --- a/sqllineage-python/sqllineage.pyi +++ b/sqllineage-python/sqllineage.pyi @@ -53,7 +53,9 @@ def analyze( Args: sql: One or more SQL statements separated by ``;``. dialect: SQL dialect (generic, ansi, postgresql, mysql, hive, - databricks, snowflake, bigquery). + databricks, snowflake, bigquery, duckdb, redshift, spark, + clickhouse, sqlite, mssql, oracle, teradata). Also accepts + the aliases postgres, sparksql, tsql, and sqlserver. catalog: Optional object implementing ``list_columns(table: TableRef) -> list[str] | None`` and ``resolve_column(column: str, candidates: list[TableRef]) -> TableRef | None``. diff --git a/sqllineage-python/src/lib.rs b/sqllineage-python/src/lib.rs index 337dfc1..01e3128 100644 --- a/sqllineage-python/src/lib.rs +++ b/sqllineage-python/src/lib.rs @@ -267,7 +267,11 @@ impl sqllineage_core::CatalogProvider for PyCatalog { /// /// Args: /// sql: One or more SQL statements (separated by `;`). -/// dialect: SQL dialect name (default: "generic"). +/// dialect: SQL dialect name (default: "generic"). One of generic, ansi, +/// postgresql, mysql, hive, databricks, snowflake, bigquery, +/// duckdb, redshift, spark, clickhouse, sqlite, mssql, oracle, +/// teradata; the aliases postgres, sparksql, tsql, and sqlserver +/// are also accepted. /// catalog: Optional object with `list_columns(table) -> list[str] | None` /// and `resolve_column(column, candidates) -> TableRef | None`. /// normalize_case: Lowercase unquoted identifiers (default: True). @@ -282,21 +286,12 @@ fn analyze( catalog: Option>, normalize_case: bool, ) -> PyResult> { - let d = match dialect.to_lowercase().as_str() { - "generic" => sqllineage_core::Dialect::Generic, - "ansi" => sqllineage_core::Dialect::Ansi, - "postgresql" | "postgres" => sqllineage_core::Dialect::PostgreSql, - "mysql" => sqllineage_core::Dialect::MySql, - "hive" => sqllineage_core::Dialect::Hive, - "databricks" => sqllineage_core::Dialect::Databricks, - "snowflake" => sqllineage_core::Dialect::Snowflake, - "bigquery" => sqllineage_core::Dialect::BigQuery, - other => { - return Err(pyo3::exceptions::PyValueError::new_err(format!( - "unknown dialect: '{other}'" - ))); - } - }; + let d: sqllineage_core::Dialect = + dialect + .parse() + .map_err(|e: sqllineage_core::UnknownDialect| { + pyo3::exceptions::PyValueError::new_err(e.to_string()) + })?; let catalog_box: Option> = catalog.map(|obj| Box::new(PyCatalog { obj }) as Box); diff --git a/sqllineage/src/bin/sqllineage.rs b/sqllineage/src/bin/sqllineage.rs index 49445c6..112d450 100644 --- a/sqllineage/src/bin/sqllineage.rs +++ b/sqllineage/src/bin/sqllineage.rs @@ -30,13 +30,10 @@ struct Cli { fn main() { let cli = Cli::parse(); - let dialect = match parse_dialect(&cli.dialect) { - Some(d) => d, - None => { - eprintln!( - "error: unknown dialect '{}'. valid: generic, ansi, postgresql, mysql, hive, databricks, snowflake, bigquery", - cli.dialect - ); + let dialect: Dialect = match cli.dialect.parse() { + Ok(d) => d, + Err(e) => { + eprintln!("error: {e}"); process::exit(1); } }; @@ -72,20 +69,6 @@ fn main() { } } -fn parse_dialect(s: &str) -> Option { - match s.to_lowercase().as_str() { - "generic" => Some(Dialect::Generic), - "ansi" => Some(Dialect::Ansi), - "postgresql" | "postgres" => Some(Dialect::PostgreSql), - "mysql" => Some(Dialect::MySql), - "hive" => Some(Dialect::Hive), - "databricks" => Some(Dialect::Databricks), - "snowflake" => Some(Dialect::Snowflake), - "bigquery" => Some(Dialect::BigQuery), - _ => None, - } -} - fn format_json(result: &AnalyzeResult, columns: bool) -> String { if columns { serde_json::to_string_pretty(result).unwrap_or_default() diff --git a/sqllineage/src/dialect.rs b/sqllineage/src/dialect.rs index e6f1d1f..817d0eb 100644 --- a/sqllineage/src/dialect.rs +++ b/sqllineage/src/dialect.rs @@ -1,7 +1,8 @@ use crate::types::Dialect; use sqlparser::dialect::{ - self, AnsiDialect, BigQueryDialect, DatabricksDialect, GenericDialect, HiveDialect, - MySqlDialect, PostgreSqlDialect, SnowflakeDialect, + self, AnsiDialect, BigQueryDialect, ClickHouseDialect, DatabricksDialect, DuckDbDialect, + GenericDialect, HiveDialect, MsSqlDialect, MySqlDialect, OracleDialect, PostgreSqlDialect, + RedshiftSqlDialect, SQLiteDialect, SnowflakeDialect, SparkSqlDialect, TeradataDialect, }; impl Dialect { @@ -15,6 +16,14 @@ impl Dialect { Dialect::Databricks => Box::new(DatabricksDialect), Dialect::Snowflake => Box::new(SnowflakeDialect), Dialect::BigQuery => Box::new(BigQueryDialect), + Dialect::DuckDb => Box::new(DuckDbDialect {}), + Dialect::Redshift => Box::new(RedshiftSqlDialect {}), + Dialect::Spark => Box::new(SparkSqlDialect {}), + Dialect::ClickHouse => Box::new(ClickHouseDialect {}), + Dialect::SQLite => Box::new(SQLiteDialect {}), + Dialect::MsSql => Box::new(MsSqlDialect {}), + Dialect::Oracle => Box::new(OracleDialect {}), + Dialect::Teradata => Box::new(TeradataDialect {}), } } } diff --git a/sqllineage/src/types.rs b/sqllineage/src/types.rs index fbdc051..94d11f7 100644 --- a/sqllineage/src/types.rs +++ b/sqllineage/src/types.rs @@ -194,7 +194,17 @@ impl Default for AnalyzeOptions { } /// Supported SQL dialects (maps to sqlparser dialects). -#[derive(Debug, Clone, Copy, Default)] +/// +/// Every dialect that the pinned `sqlparser` release exposes has a variant +/// here. Parsing a statement with the closest dialect matters for lineage: +/// `generic` accepts a superset of most grammars, but it does not apply +/// dialect-specific rules such as `BigQuery`'s backtick quoting or T-SQL's +/// bracket quoting. +/// +/// Marked `#[non_exhaustive]`: `sqlparser` gains dialects over time, and +/// adding one here should not be a breaking change for downstream matches. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash)] +#[non_exhaustive] pub enum Dialect { #[default] Generic, @@ -205,8 +215,130 @@ pub enum Dialect { Databricks, Snowflake, BigQuery, + DuckDb, + Redshift, + Spark, + ClickHouse, + SQLite, + /// Microsoft SQL Server (T-SQL). + MsSql, + Oracle, + Teradata, } +impl Dialect { + /// Every supported dialect, in the order used for help and error text. + pub const ALL: &'static [Self] = &[ + Self::Generic, + Self::Ansi, + Self::PostgreSql, + Self::MySql, + Self::Hive, + Self::Databricks, + Self::Snowflake, + Self::BigQuery, + Self::DuckDb, + Self::Redshift, + Self::Spark, + Self::ClickHouse, + Self::SQLite, + Self::MsSql, + Self::Oracle, + Self::Teradata, + ]; + + /// The canonical lowercase name, as accepted by [`Dialect::from_str`] and + /// printed by [`Display`]. + /// + /// [`Display`]: std::fmt::Display + pub const fn name(self) -> &'static str { + match self { + Self::Generic => "generic", + Self::Ansi => "ansi", + Self::PostgreSql => "postgresql", + Self::MySql => "mysql", + Self::Hive => "hive", + Self::Databricks => "databricks", + Self::Snowflake => "snowflake", + Self::BigQuery => "bigquery", + Self::DuckDb => "duckdb", + Self::Redshift => "redshift", + Self::Spark => "spark", + Self::ClickHouse => "clickhouse", + Self::SQLite => "sqlite", + Self::MsSql => "mssql", + Self::Oracle => "oracle", + Self::Teradata => "teradata", + } + } + + /// A comma-separated list of every canonical name, for help and error text. + pub fn names() -> String { + Self::ALL + .iter() + .map(|d| d.name()) + .collect::>() + .join(", ") + } +} + +impl fmt::Display for Dialect { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.name()) + } +} + +impl std::str::FromStr for Dialect { + type Err = UnknownDialect; + + /// Parse a dialect name, case-insensitively. + /// + /// Accepts each canonical name from [`Dialect::name`] plus a few common + /// spellings: `postgres`, `sparksql`, `tsql`, and `sqlserver`. + fn from_str(s: &str) -> Result { + match s.to_lowercase().as_str() { + "generic" => Ok(Self::Generic), + "ansi" => Ok(Self::Ansi), + "postgresql" | "postgres" => Ok(Self::PostgreSql), + "mysql" => Ok(Self::MySql), + "hive" => Ok(Self::Hive), + "databricks" => Ok(Self::Databricks), + "snowflake" => Ok(Self::Snowflake), + "bigquery" => Ok(Self::BigQuery), + "duckdb" => Ok(Self::DuckDb), + "redshift" => Ok(Self::Redshift), + "spark" | "sparksql" => Ok(Self::Spark), + "clickhouse" => Ok(Self::ClickHouse), + "sqlite" => Ok(Self::SQLite), + "mssql" | "tsql" | "sqlserver" => Ok(Self::MsSql), + "oracle" => Ok(Self::Oracle), + "teradata" => Ok(Self::Teradata), + _ => Err(UnknownDialect { + name: s.to_string(), + }), + } + } +} + +/// Error returned when a dialect name is not recognized. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct UnknownDialect { + pub name: String, +} + +impl fmt::Display for UnknownDialect { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!( + f, + "unknown dialect '{}'. valid: {}", + self.name, + Dialect::names() + ) + } +} + +impl std::error::Error for UnknownDialect {} + /// Error returned when SQL parsing fails. #[derive(Debug, Clone)] pub struct ParseError { diff --git a/sqllineage/tests/dialect.rs b/sqllineage/tests/dialect.rs new file mode 100644 index 0000000..96a2f7b --- /dev/null +++ b/sqllineage/tests/dialect.rs @@ -0,0 +1,39 @@ +use sqllineage::{AnalyzeOptions, Dialect, TableRef, analyze}; + +#[test] +fn every_dialect_parses_a_basic_query() { + for &dialect in Dialect::ALL { + let results = analyze( + "SELECT a FROM t", + AnalyzeOptions { + dialect, + ..AnalyzeOptions::default() + }, + ) + .unwrap_or_else(|e| panic!("{dialect} failed to parse: {e}")); + + assert_eq!( + results[0].tables.inputs, + vec![TableRef::new("t")], + "{dialect}" + ); + } +} + +/// `name` and `from_str` are separate tables; this keeps them agreeing. +#[test] +fn canonical_names_round_trip() { + for &dialect in Dialect::ALL { + assert_eq!(dialect.name().parse::(), Ok(dialect)); + assert_eq!(dialect.to_string(), dialect.name()); + } +} + +#[test] +fn aliases_and_mixed_case_resolve() { + assert_eq!("postgres".parse(), Ok(Dialect::PostgreSql)); + assert_eq!("sparksql".parse(), Ok(Dialect::Spark)); + assert_eq!("tsql".parse(), Ok(Dialect::MsSql)); + assert_eq!("sqlserver".parse(), Ok(Dialect::MsSql)); + assert_eq!("DuckDB".parse(), Ok(Dialect::DuckDb)); +}