diff --git a/src/dialect_function_map.cpp b/src/dialect_function_map.cpp index 0fb5add..cb6a563 100644 --- a/src/dialect_function_map.cpp +++ b/src/dialect_function_map.cpp @@ -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. @@ -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: diff --git a/src/include/sql_dialect.hpp b/src/include/sql_dialect.hpp index 9ee3b1b..fe3351b 100644 --- a/src/include/sql_dialect.hpp +++ b/src/include/sql_dialect.hpp @@ -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) diff --git a/src/lpts_ast_builder.cpp b/src/lpts_ast_builder.cpp index 693cba9..aca747a 100644 --- a/src/lpts_ast_builder.cpp +++ b/src/lpts_ast_builder.cpp @@ -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> extra_projection_outputs; + vector referenced_bindings; const unique_ptr &FindColumnBinding(const ColumnBinding &binding, const char *context) const { auto it = column_map.find(MappableColumnBinding(binding)); @@ -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); } @@ -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); } @@ -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"; @@ -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"); } @@ -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 { diff --git a/src/lpts_expression_renderer.cpp b/src/lpts_expression_renderer.cpp index cded17d..c92f51e 100644 --- a/src/lpts_expression_renderer.cpp +++ b/src/lpts_expression_renderer.cpp @@ -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: @@ -375,6 +375,35 @@ static string RenderConstantForDialect(const BoundConstantExpression &constant, return "TIMESTAMP '" + EscapeSingleQuotes(value.ToString()) + "'"; case LogicalTypeId::VARCHAR: return "'" + EscapeSingleQuotes(value.GetValue()) + "'"; + case LogicalTypeId::INTERVAL: { + const auto interval = value.GetValue(); + vector 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(); } @@ -832,9 +861,8 @@ string UnquotedDatetimeUnitLiteral(const unique_ptr &arg) { if (constant.value.IsNull() || constant.value.type().id() != LogicalTypeId::VARCHAR) { return string(); } - static const std::set kSparkUnits = {"YEAR", "QUARTER", "MONTH", "WEEK", - "DAY", "DAYOFYEAR", "HOUR", "MINUTE", - "SECOND", "MILLISECOND", "MICROSECOND"}; + static const std::set kSparkUnits = {"YEAR", "QUARTER", "MONTH", "WEEK", "DAY", "DAYOFYEAR", + "HOUR", "MINUTE", "SECOND", "MILLISECOND", "MICROSECOND"}; string unit = StringUtil::Upper(constant.value.GetValue()); if (kSparkUnits.find(unit) == kSparkUnits.end()) { return string(); @@ -1383,6 +1411,12 @@ string LptsExpressionRenderer::ExpressionToAliasedString(const unique_ptr