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
6 changes: 4 additions & 2 deletions src/dialect_function_map.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -28,8 +28,8 @@ string RemapForPostgres(const string &name) {
// target will reject as an unknown function — which reads like an engine
// limitation instead of a translation one.
if (name == "arg_min" || name == "argmin" || name == "arg_max" || name == "argmax") {
ThrowLptsNotImplemented("LPTS_UNSUPPORTED_FUNCTION", SqlDialect::POSTGRES, "aggregate", name,
"AGGREGATE", "postgres has no arg_min/arg_max equivalent");
ThrowLptsNotImplemented("LPTS_UNSUPPORTED_FUNCTION", SqlDialect::POSTGRES, "aggregate", name, "AGGREGATE",
"postgres has no arg_min/arg_max equivalent");
}
if (name == "string_split" || name == "str_split") {
// PostgreSQL uses string_to_array(string, delimiter) for the same purpose.
Expand Down Expand Up @@ -195,6 +195,8 @@ string RemapFunctionNameForDialect(const string &duckdb_name, SqlDialect dialect
switch (dialect) {
case SqlDialect::POSTGRES:
return RemapForPostgres(duckdb_name);
case SqlDialect::FELDERA:
return duckdb_name;
case SqlDialect::SPARK:
return RemapForSpark(duckdb_name);
case SqlDialect::HIVE:
Expand Down
1 change: 1 addition & 0 deletions src/include/sql_dialect.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ namespace duckdb {
enum class SqlDialect {
DUCKDB, ///< Default. Uses DuckDB-specific syntax (fully-qualified table refs, "ident" quoting, etc.)
POSTGRES, ///< PostgreSQL-compatible syntax (unqualified table refs, "ident" quoting, etc.)
FELDERA, ///< Feldera SQL (unqualified refs, ANSI quoting, native ARG_MIN/ARG_MAX)
SPARK, ///< Apache Spark SQL syntax (catalog.schema.table qualified, `ident` backtick quoting, ROWS/RANGE-only
///< windows)
HIVE, ///< Apache Hive SQL syntax (schema.table qualified, `ident` backtick quoting)
Expand Down
14 changes: 11 additions & 3 deletions src/lpts_ast_builder.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -475,6 +475,7 @@ class AstBuilder {
/// arguments bound to a pre-projection column; without carrying that binding upward, LPTS
/// can emit a stale CTE column name in the aggregate SELECT list.
unordered_map<const LogicalOperator *, vector<ColumnBinding>> extra_projection_outputs;
vector<ColumnBinding> referenced_bindings;

const unique_ptr<ColStruct> &FindColumnBinding(const ColumnBinding &binding, const char *context) const {
auto it = column_map.find(MappableColumnBinding(binding));
Expand Down Expand Up @@ -548,6 +549,7 @@ class AstBuilder {
auto *child = op->children[0].get();
const auto child_bindings = child->GetColumnBindings();
for (const auto &ref : refs) {
AddUniqueBinding(referenced_bindings, ref);
if (!HasBinding(child_bindings, ref)) {
EnsureBindingAvailableFrom(child, ref);
}
Expand All @@ -571,6 +573,7 @@ class AstBuilder {
auto *child = op->children[0].get();
const auto child_bindings = child->GetColumnBindings();
for (const auto &ref : refs) {
AddUniqueBinding(referenced_bindings, ref);
if (!HasBinding(child_bindings, ref)) {
EnsureBindingAvailableFrom(child, ref);
}
Expand Down Expand Up @@ -1224,6 +1227,11 @@ class AstBuilder {
// (optimizer removed unused columns), the binding index may
// differ from the loop index.
const idx_t col_id_idx = cb.column_index;
if (dialect != SqlDialect::DUCKDB && col_id_idx < col_ids.size() &&
col_ids[col_id_idx].IsVirtualColumn() && !HasBinding(referenced_bindings, cb)) {
all_columns_native = false;
continue;
}
if (col_id_idx >= col_ids.size()) {
all_columns_native = false;
string col_name = "rowid";
Expand Down Expand Up @@ -1335,7 +1343,7 @@ class AstBuilder {
cte_column_names.clear();
column_is_expr.clear();
column_names.push_back("1");
column_is_expr.push_back(false);
column_is_expr.push_back(true);
cte_column_names.push_back("t" + std::to_string(table_index) + "_dummy");
}

Expand Down Expand Up @@ -1749,8 +1757,8 @@ class AstBuilder {
if (agg_name != "count_star") {
agg_name = RemapFunctionNameForDialect(agg_name, dialect);
}
const bool render_count_star = agg_name == "count_star" && ba.children.empty() &&
!is_export_state && dialect != SqlDialect::DUCKDB;
const bool render_count_star = agg_name == "count_star" && ba.children.empty() && !is_export_state &&
dialect != SqlDialect::DUCKDB;
if (render_count_star) {
agg_str << "count(*";
} else {
Expand Down
45 changes: 40 additions & 5 deletions src/lpts_expression_renderer.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -338,7 +338,7 @@ static string RenderCastTargetType(const LogicalType &type, SqlDialect dialect)
case LogicalTypeId::DECIMAL:
return type.ToString();
case LogicalTypeId::VARCHAR:
return "VARCHAR";
return dialect == SqlDialect::SPARK ? "STRING" : "VARCHAR";
case LogicalTypeId::DATE:
return "DATE";
case LogicalTypeId::TIME:
Expand Down Expand Up @@ -375,6 +375,35 @@ static string RenderConstantForDialect(const BoundConstantExpression &constant,
return "TIMESTAMP '" + EscapeSingleQuotes(value.ToString()) + "'";
case LogicalTypeId::VARCHAR:
return "'" + EscapeSingleQuotes(value.GetValue<string>()) + "'";
case LogicalTypeId::INTERVAL: {
const auto interval = value.GetValue<interval_t>();
vector<string> parts;
if (interval.months != 0) {
parts.push_back("INTERVAL '" + std::to_string(interval.months) + "' MONTH");
}
if (interval.days != 0) {
parts.push_back("INTERVAL '" + std::to_string(interval.days) + "' DAY");
}
if (interval.micros != 0) {
const bool negative = interval.micros < 0;
const uint64_t magnitude = negative ? uint64_t(-(interval.micros + 1)) + 1 : uint64_t(interval.micros);
const uint64_t seconds = magnitude / Interval::MICROS_PER_SEC;
const uint64_t micros = magnitude % Interval::MICROS_PER_SEC;
string seconds_text = (negative ? "-" : "") + std::to_string(seconds);
if (micros != 0) {
string fraction = StringUtil::Format("%06llu", (unsigned long long)micros);
while (fraction.back() == '0') {
fraction.pop_back();
}
seconds_text += "." + fraction;
}
parts.push_back("INTERVAL '" + seconds_text + "' SECOND");
}
if (parts.empty()) {
return "INTERVAL '0' SECOND";
}
return parts.size() == 1 ? parts[0] : "(" + StringUtil::Join(parts, " + ") + ")";
}
default:
return value.ToSQLString();
}
Expand Down Expand Up @@ -832,9 +861,8 @@ string UnquotedDatetimeUnitLiteral(const unique_ptr<Expression> &arg) {
if (constant.value.IsNull() || constant.value.type().id() != LogicalTypeId::VARCHAR) {
return string();
}
static const std::set<string> kSparkUnits = {"YEAR", "QUARTER", "MONTH", "WEEK",
"DAY", "DAYOFYEAR", "HOUR", "MINUTE",
"SECOND", "MILLISECOND", "MICROSECOND"};
static const std::set<string> kSparkUnits = {"YEAR", "QUARTER", "MONTH", "WEEK", "DAY", "DAYOFYEAR",
"HOUR", "MINUTE", "SECOND", "MILLISECOND", "MICROSECOND"};
string unit = StringUtil::Upper(constant.value.GetValue<string>());
if (kSparkUnits.find(unit) == kSparkUnits.end()) {
return string();
Expand Down Expand Up @@ -1383,6 +1411,12 @@ string LptsExpressionRenderer::ExpressionToAliasedString(const unique_ptr<Expres
break;
}
ValidateFunctionForDialect(func_expr, dialect);
if (dialect == SqlDialect::FELDERA && func_expr.children.size() == 1 &&
(func_expr.function.name == "year" || func_expr.function.name == "month")) {
expr_str << "EXTRACT(" << StringUtil::Upper(func_expr.function.name) << " FROM "
<< ExpressionToAliasedString(func_expr.children[0]) << ")";
break;
}
// Dialect-specific function name remapping (see dialect_function_map.hpp).
string func_name = RemapFunctionNameForDialect(func_expr.function.name, dialect);
// For lambda functions, only serialize non-lambda, non-capture children
Expand Down Expand Up @@ -1425,7 +1459,8 @@ string LptsExpressionRenderer::ExpressionToAliasedString(const unique_ptr<Expres
// function resolver see the function call directly and bypass the
// keyword-syntax path.
string emit_name = func_name;
if (func_name == "position" || func_name == "substring" || func_name == "overlay" || func_name == "trim") {
if (IsDuckDBDialect(dialect) && (func_name == "position" || func_name == "substring" ||
func_name == "overlay" || func_name == "trim")) {
emit_name = "\"" + func_name + "\"";
}
// Spark's datediff takes the unit as a bare keyword, not a string:
Expand Down
9 changes: 7 additions & 2 deletions src/sql_dialect.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,9 @@ SqlDialect ParseSqlDialectSetting(const string &value, const string &setting_nam
if (normalized == "postgres" || normalized == "postgresql") {
return SqlDialect::POSTGRES;
}
if (normalized == "feldera") {
return SqlDialect::FELDERA;
}
if (normalized == "spark") {
return SqlDialect::SPARK;
}
Expand All @@ -35,7 +38,7 @@ SqlDialect ParseSqlDialectSetting(const string &value, const string &setting_nam
return SqlDialect::MYSQL_MARIADB;
}
throw InvalidInputException(
"Unknown %s '%s'. Valid values: 'duckdb', 'postgres', 'spark', 'hive', 'trino', 'presto', "
"Unknown %s '%s'. Valid values: 'duckdb', 'postgres', 'feldera', 'spark', 'hive', 'trino', 'presto', "
"'snowflake', 'bigquery', 'redshift', 'mysql', 'mariadb'",
setting_name, value);
}
Expand All @@ -50,6 +53,8 @@ string SqlDialectToString(SqlDialect dialect) {
return "duckdb";
case SqlDialect::POSTGRES:
return "postgres";
case SqlDialect::FELDERA:
return "feldera";
case SqlDialect::SPARK:
return "spark";
case SqlDialect::HIVE:
Expand All @@ -75,7 +80,7 @@ bool DialectUsesBacktickQuotedIdentifiers(SqlDialect dialect) {
}

bool DialectUsesUnqualifiedTableNames(SqlDialect dialect) {
return dialect == SqlDialect::POSTGRES || dialect == SqlDialect::REDSHIFT;
return dialect == SqlDialect::POSTGRES || dialect == SqlDialect::FELDERA || dialect == SqlDialect::REDSHIFT;
}

bool DialectUsesSchemaQualifiedTableNames(SqlDialect dialect) {
Expand Down
33 changes: 33 additions & 0 deletions test/sql/dialect_hardening.test
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,39 @@

require lpts

# Feldera is PostgreSQL-like but supports ARG_MIN/ARG_MAX directly. Its
# DataFusion-backed ad-hoc evaluator also needs standard EXTRACT syntax.
statement ok
CREATE TABLE feldera_edges (id INTEGER, amount DECIMAL(12,2), ts TIMESTAMP);

statement ok
SET lpts_dialect = 'feldera';

query III
SELECT sql LIKE '%arg_min(%' AS arg_min_preserved,
sql LIKE '%arg_max(%' AS arg_max_preserved,
sql NOT LIKE '%LPTS_UNSUPPORTED_FUNCTION%' AS translated
FROM lpts_query('SELECT arg_min(id, amount), arg_max(id, amount) FROM feldera_edges');
----
true true true

query II
SELECT sql LIKE '%EXTRACT(YEAR FROM%' AS extract_year,
sql LIKE '%EXTRACT(MONTH FROM%' AS extract_month
FROM lpts_query('SELECT extract(year FROM ts), extract(month FROM ts) FROM feldera_edges');
----
true true

query II
SELECT sql LIKE '%INTERVAL ''30'' DAY%' AS ansi_interval,
sql NOT LIKE '%::INTERVAL%' AS no_duckdb_interval_cast
FROM lpts_query('SELECT ts + INTERVAL ''30'' DAY FROM feldera_edges');
----
true true

statement ok
SET lpts_dialect = 'duckdb';

statement ok
CREATE TABLE users (id INTEGER, name VARCHAR, ts TIMESTAMP, vals INTEGER[]);

Expand Down
27 changes: 27 additions & 0 deletions test/sql/dialect_spark.test
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,9 @@ CREATE TABLE events (id INTEGER, ts TIMESTAMP, vals INTEGER[]);
statement ok
INSERT INTO events VALUES (1, TIMESTAMP '2024-01-15 00:00:00', [1, 2, 3]);

statement ok
CREATE TABLE compiler_edges (d DECIMAL(4,4), name VARCHAR);

# ============================================================
# Setup: enable SPARK dialect
# ============================================================
Expand Down Expand Up @@ -127,6 +130,30 @@ SELECT sql LIKE '%COALESCE(%' AS coalesce_passthrough FROM lpts_query('SELECT co
----
true

# Cross-engine compiler-bench regressions: do not expose DuckDB's synthetic
# rowid for COUNT(*), and emit Spark-native casts/function calls.
query III
SELECT sql NOT LIKE '%rowid%' AS count_star_has_no_rowid,
sql LIKE '%SELECT 1%' AS count_star_uses_dummy,
sql LIKE '%count(*)%' AS count_star_preserved
FROM lpts_query('SELECT count(*) FROM compiler_edges');
----
true true true

query II
SELECT sql LIKE '%CAST(d AS STRING)%' AS varchar_cast_is_string,
sql NOT LIKE '%"substring"(%' AS substring_is_not_quoted
FROM lpts_query('SELECT CAST(d AS VARCHAR), substring(name, 1, 2) FROM compiler_edges');
----
true true

query II
SELECT sql LIKE '%INTERVAL ''30'' DAY%' AS spark_interval,
sql NOT LIKE '%::INTERVAL%' AS no_postgres_interval_cast
FROM lpts_query('SELECT ts + INTERVAL ''30'' DAY FROM events');
----
true true

# ============================================================
# Window-frame SPARK gate: EXCLUDE clauses → NotImplementedException
# (Spark SQL supports only ROWS / RANGE without exclusion.)
Expand Down
Loading