Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -238,6 +238,9 @@ test-rust-lib:

# Check capabilities separately so Cargo feature unification cannot mask dependencies.
test-rust-feature-gates:
cargo check -p polyglot-sql --no-default-features --features dialect-hana
cargo check -p polyglot-sql --no-default-features --features generate,dialect-hana
cargo check -p polyglot-sql --no-default-features --features transpile,dialect-hana
cargo check -p polyglot-sql --no-default-features
@for feature in generate transpile builder ast-tools semantic openlineage diff planner time \
function-catalog-clickhouse function-catalog-duckdb function-catalog-all-dialects; do \
Expand Down Expand Up @@ -282,6 +285,7 @@ test-rust-verify-core:
@echo "=== Lib unit tests ==="
@cargo test --lib -p polyglot-sql
@cargo test -p polyglot-sql --test deep_nesting_regression
@cargo test -p polyglot-sql --test dialect_matrix
@echo ""
@echo "=== Generic identity tests ==="
@cargo test --test sqlglot_identity test_sqlglot_identity_all -p polyglot-sql -- --nocapture
Expand Down
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ Release notes are tracked in [`CHANGELOG.md`](CHANGELOG.md).
| MySQL | Oracle | PostgreSQL | Presto | Redshift |
| RisingWave | SingleStore | Snowflake | Solr | Spark |
| SQLite | StarRocks | Tableau | Teradata | TiDB |
| Trino | TSQL | DataFusion | Generic SQL | |
| Trino | TSQL | DataFusion | SAP HANA | Generic SQL |

## Quick Start

Expand Down
152 changes: 97 additions & 55 deletions crates/polyglot-sql-ast-derive/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,32 +19,9 @@ pub fn derive_ast_node(input: TokenStream) -> TokenStream {

fn expand_ast_node(input: &DeriveInput) -> proc_macro2::TokenStream {
let name = &input.ident;
let immutable = match &input.data {
Data::Struct(data) => visit_fields(&data.fields, false),
Data::Enum(data) => {
let arms = data.variants.iter().map(|variant| {
let variant_name = &variant.ident;
let (pattern, body) =
visit_variant_fields(name, variant_name, &variant.fields, false);
quote!(#pattern => { #body })
});
quote!(match self { #(#arms),* })
}
Data::Union(_) => quote!(),
};
let mutable = match &input.data {
Data::Struct(data) => visit_fields(&data.fields, true),
Data::Enum(data) => {
let arms = data.variants.iter().map(|variant| {
let variant_name = &variant.ident;
let (pattern, body) =
visit_variant_fields(name, variant_name, &variant.fields, true);
quote!(#pattern => { #body })
});
quote!(match self { #(#arms),* })
}
Data::Union(_) => quote!(),
};
let immutable = node_visitor(input, false, true);
let untracked = node_visitor(input, false, false);
let mutable = node_visitor(input, true, false);
let serialized_variant_names = if name == "Expression" {
if let Data::Enum(data) = &input.data {
let names = data.variants.iter().map(|variant| {
Expand Down Expand Up @@ -81,6 +58,14 @@ fn expand_ast_node(input: &DeriveInput) -> proc_macro2::TokenStream {
#immutable
}

fn visit_syntax_untracked<'ast, F, T>(&'ast self, visitor: &mut F, type_visitor: &mut T)
where
F: FnMut(&'ast crate::expressions::Expression),
T: FnMut(&'ast crate::expressions::DataType),
{
#untracked
}

fn visit_expressions_mut<F>(
&mut self,
visitor: &mut F,
Expand All @@ -96,6 +81,22 @@ fn expand_ast_node(input: &DeriveInput) -> proc_macro2::TokenStream {
}
}

fn node_visitor(input: &DeriveInput, mutable: bool, paths: bool) -> proc_macro2::TokenStream {
let name = &input.ident;
match &input.data {
Data::Struct(data) => visit_fields(&data.fields, mutable, paths),
Data::Enum(data) => {
let arms = data.variants.iter().map(|variant| {
let (pattern, body) =
visit_variant_fields(name, &variant.ident, &variant.fields, mutable, paths);
quote!(#pattern => { #body })
});
quote!(match self { #(#arms),* })
}
Data::Union(_) => quote!(),
}
}

/// Match serde's `rename_all = "snake_case"` behavior for Rust enum variants.
fn serde_snake_case(name: &str) -> String {
let mut snake_case = String::with_capacity(name.len());
Expand All @@ -108,7 +109,7 @@ fn serde_snake_case(name: &str) -> String {
snake_case
}

fn visit_fields(fields: &Fields, mutable: bool) -> proc_macro2::TokenStream {
fn visit_fields(fields: &Fields, mutable: bool, paths: bool) -> proc_macro2::TokenStream {
match fields {
Fields::Named(fields) => {
let visits = fields.named.iter().filter_map(|field| {
Expand All @@ -122,11 +123,15 @@ fn visit_fields(fields: &Fields, mutable: bool) -> proc_macro2::TokenStream {
Some(if mutable {
mutable_visit(&field.ty, access)
} else {
let visit = immutable_visit(&field.ty, immutable_access);
quote! {
path.push(crate::ast_children::ChildPathSegment::Field(#field_name));
#visit
path.pop();
let visit = immutable_visit(&field.ty, immutable_access, paths);
if paths {
quote! {
path.push(crate::ast_children::ChildPathSegment::Field(#field_name));
#visit
path.pop();
}
} else {
visit
}
})
});
Expand All @@ -147,11 +152,15 @@ fn visit_fields(fields: &Fields, mutable: bool) -> proc_macro2::TokenStream {
Some(if mutable {
mutable_visit(&field.ty, access)
} else {
let visit = immutable_visit(&field.ty, immutable_access);
quote! {
path.push(crate::ast_children::ChildPathSegment::Index(#index));
#visit
path.pop();
let visit = immutable_visit(&field.ty, immutable_access, paths);
if paths {
quote! {
path.push(crate::ast_children::ChildPathSegment::Index(#index));
#visit
path.pop();
}
} else {
visit
}
})
});
Expand All @@ -166,6 +175,7 @@ fn visit_variant_fields(
variant_name: &syn::Ident,
fields: &Fields,
mutable: bool,
paths: bool,
) -> (proc_macro2::TokenStream, proc_macro2::TokenStream) {
match fields {
Fields::Named(fields) => {
Expand All @@ -190,11 +200,15 @@ fn visit_variant_fields(
Some(if mutable {
mutable_visit(&field.ty, quote!(#ident))
} else {
let visit = immutable_visit(&field.ty, quote!(#ident));
quote! {
path.push(crate::ast_children::ChildPathSegment::Field(#field_name));
#visit
path.pop();
let visit = immutable_visit(&field.ty, quote!(#ident), paths);
if paths {
quote! {
path.push(crate::ast_children::ChildPathSegment::Field(#field_name));
#visit
path.pop();
}
} else {
visit
}
})
});
Expand All @@ -220,8 +234,8 @@ fn visit_variant_fields(
Some(if mutable {
mutable_visit(&field.ty, quote!(#binding))
} else {
let visit = immutable_visit(&field.ty, quote!(#binding));
if single_expression_payload {
let visit = immutable_visit(&field.ty, quote!(#binding), paths);
if single_expression_payload || !paths {
visit
} else {
quote! {
Expand All @@ -241,19 +255,30 @@ fn visit_variant_fields(
}
}

fn immutable_visit(ty: &Type, access: proc_macro2::TokenStream) -> proc_macro2::TokenStream {
fn immutable_visit(
ty: &Type,
access: proc_macro2::TokenStream,
paths: bool,
) -> proc_macro2::TokenStream {
if is_expression(ty) {
return quote!(visitor(path, #access););
return if paths {
quote!(visitor(path, #access);)
} else {
quote!(visitor(#access);)
};
}
if let Some(inner) = container_inner(ty, "Option") {
let visit = immutable_visit(inner, quote!(value));
let visit = immutable_visit(inner, quote!(value), paths);
return quote!(if let Some(value) = (#access).as_ref() { #visit });
}
if let Some(inner) = container_inner(ty, "Box") {
return immutable_visit(inner, quote!((#access).as_ref()));
return immutable_visit(inner, quote!((#access).as_ref()), paths);
}
if let Some(inner) = container_inner(ty, "Vec") {
let visit = immutable_visit(inner, quote!(value));
let visit = immutable_visit(inner, quote!(value), paths);
if !paths {
return quote!(for value in (#access).iter() { #visit });
}
return quote! {
for (index, value) in (#access).iter().enumerate() {
path.push(crate::ast_children::ChildPathSegment::Index(index));
Expand All @@ -265,19 +290,36 @@ fn immutable_visit(ty: &Type, access: proc_macro2::TokenStream) -> proc_macro2::
if let Type::Tuple(tuple) = ty {
let visits = tuple.elems.iter().enumerate().map(|(index, element)| {
let tuple_index = syn::Index::from(index);
let visit = immutable_visit(element, quote!(&(#access).#tuple_index));
quote! {
path.push(crate::ast_children::ChildPathSegment::Index(#index));
#visit
path.pop();
let visit = immutable_visit(element, quote!(&(#access).#tuple_index), paths);
if paths {
quote! {
path.push(crate::ast_children::ChildPathSegment::Index(#index));
#visit
path.pop();
}
} else {
visit
}
});
return quote!(#(#visits)*);
}
if is_scalar(ty) {
return quote!();
}
quote!(crate::ast_children::AstNode::visit_expressions(#access, path, visitor);)
if paths {
quote!(crate::ast_children::AstNode::visit_expressions(#access, path, visitor);)
} else {
let visit_type = if matches!(ty, Type::Path(path) if path.path.segments.last().is_some_and(|segment| segment.ident == "DataType"))
{
quote!(type_visitor(#access);)
} else {
quote!()
};
quote! {
#visit_type
crate::ast_children::AstNode::visit_syntax_untracked(#access, visitor, type_visitor);
}
}
}

fn mutable_visit(ty: &Type, access: proc_macro2::TokenStream) -> proc_macro2::TokenStream {
Expand Down
3 changes: 2 additions & 1 deletion crates/polyglot-sql-ffi/src/dialects.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@ use polyglot_sql::dialects::DialectType;
use std::os::raw::c_char;
use std::ptr;

const DIALECTS: [DialectType; 34] = [
const DIALECTS: [DialectType; 35] = [
DialectType::Generic,
DialectType::PostgreSQL,
DialectType::MySQL,
Expand Down Expand Up @@ -38,6 +38,7 @@ const DIALECTS: [DialectType; 34] = [
DialectType::Dremio,
DialectType::Exasol,
DialectType::DataFusion,
DialectType::HANA,
];

/// Return supported dialect names as JSON.
Expand Down
23 changes: 22 additions & 1 deletion crates/polyglot-sql-ffi/tests/ffi_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2050,7 +2050,7 @@ fn test_dialect_list_and_count() {
let list: Vec<String> = serde_json::from_str(&json).expect("invalid dialect list json");
let count = polyglot_dialect_count();
assert_eq!(list.len() as i32, count);
assert_eq!(count, 34);
assert_eq!(count, 35);
let unique: BTreeSet<&str> = list.iter().map(String::as_str).collect();
assert_eq!(unique.len(), list.len());
assert!(list.iter().any(|d| d == "generic"));
Expand Down Expand Up @@ -2176,3 +2176,24 @@ fn test_public_api_matches_capability_contract() {
declared_available
);
}

#[test]
fn test_hana_round_trip_and_unsupported_conversion() {
let sql = c("SELECT * FROM t FOR JSON ('arraywrap' = 'NO')");
let hana = c("hana");
let duckdb = c("duckdb");
let (status, output, error) = consume_result(polyglot_transpile(
sql.as_ptr(),
hana.as_ptr(),
hana.as_ptr(),
));
assert_eq!(status, 0, "{error:?}");
assert!(output.unwrap().contains("FOR JSON"));
let (status, _, error) = consume_result(polyglot_transpile(
sql.as_ptr(),
hana.as_ptr(),
duckdb.as_ptr(),
));
assert_ne!(status, 0);
assert!(error.unwrap().contains("HANA"));
}
2 changes: 1 addition & 1 deletion crates/polyglot-sql-python/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -329,7 +329,7 @@ All functions are exported from `polyglot_sql`.

Current dialect names returned by `polyglot_sql.dialects()`:

`athena`, `bigquery`, `clickhouse`, `cockroachdb`, `datafusion`, `databricks`, `doris`, `dremio`, `drill`, `druid`, `duckdb`, `dune`, `exasol`, `fabric`, `generic`, `hive`, `materialize`, `mysql`, `oracle`, `postgres`, `presto`, `redshift`, `risingwave`, `singlestore`, `snowflake`, `solr`, `spark`, `sqlite`, `starrocks`, `tableau`, `teradata`, `tidb`, `trino`, `tsql`.
`athena`, `bigquery`, `clickhouse`, `cockroachdb`, `datafusion`, `databricks`, `doris`, `dremio`, `drill`, `druid`, `duckdb`, `dune`, `exasol`, `fabric`, `generic`, `hana`, `hive`, `materialize`, `mysql`, `oracle`, `postgres`, `presto`, `redshift`, `risingwave`, `singlestore`, `snowflake`, `solr`, `spark`, `sqlite`, `starrocks`, `tableau`, `teradata`, `tidb`, `trino`, `tsql`.

## Error Handling

Expand Down
10 changes: 10 additions & 0 deletions crates/polyglot-sql-python/python/polyglot_sql/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -596,6 +596,11 @@ def dense_rank():
FromTimeZone,
FromUnixtime,
Function,
Upsert,
StorageProperty,
Hierarchy,
ViewParameter,
Call,
FunctionEmits,
GapFill,
Generate,
Expand Down Expand Up @@ -1302,6 +1307,11 @@ def dense_rank():
"Alias",
"Cast",
"Function",
"Upsert",
"StorageProperty",
"Hierarchy",
"ViewParameter",
"Call",
"AggregateFunction",
"WindowFunction",
"From",
Expand Down
5 changes: 5 additions & 0 deletions crates/polyglot-sql-python/python/polyglot_sql/__init__.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -855,6 +855,11 @@ class FromBase(Expression): ...
class FromTimeZone(Expression): ...
class FromUnixtime(Expression): ...
class Function(Expression): ...
class Upsert(Expression): ...
class StorageProperty(Expression): ...
class Hierarchy(Expression): ...
class ViewParameter(Expression): ...
class Call(Expression): ...
class FunctionEmits(Expression): ...
class GapFill(Expression): ...
class Generate(Expression): ...
Expand Down
1 change: 1 addition & 0 deletions crates/polyglot-sql-python/src/dialects.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ const DIALECT_NAMES: &[&str] = &[
"exasol",
"fabric",
"generic",
"hana",
"hive",
"materialize",
"mysql",
Expand Down
5 changes: 5 additions & 0 deletions crates/polyglot-sql-python/src/expr_types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,11 @@ define_expression_subclasses!(
Exists,
MemberOf,
Function,
Upsert,
StorageProperty,
Hierarchy,
ViewParameter,
Call,
AggregateFunction,
WindowFunction,
From,
Expand Down
1 change: 1 addition & 0 deletions crates/polyglot-sql-python/tests/test_dialects.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
"exasol",
"fabric",
"generic",
"hana",
"hive",
"materialize",
"mysql",
Expand Down
Loading