diff --git a/CMakeLists.txt b/CMakeLists.txt index a9b5997ef..99bad1474 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -481,6 +481,7 @@ set(DUCKDB_SRC_FILES src/duckdb/ub_src_planner_expression_binder.cpp src/duckdb/ub_src_planner_filter.cpp src/duckdb/ub_src_planner_operator.cpp + src/duckdb/ub_src_planner_sql_export.cpp src/duckdb/ub_src_planner_subquery.cpp src/duckdb/ub_src_storage.cpp src/duckdb/ub_src_storage_buffer.cpp diff --git a/src/duckdb/extension/core_functions/scalar/date/date_part.cpp b/src/duckdb/extension/core_functions/scalar/date/date_part.cpp index 7d3e6fb3d..2a1b8426c 100644 --- a/src/duckdb/extension/core_functions/scalar/date/date_part.cpp +++ b/src/duckdb/extension/core_functions/scalar/date/date_part.cpp @@ -1,3 +1,4 @@ +#include "duckdb/parser/expression/function_expression.hpp" #include "duckdb/common/vector/struct_vector.hpp" #include "core_functions/scalar/date_functions.hpp" #include "duckdb/common/case_insensitive_map.hpp" @@ -2716,9 +2717,23 @@ ScalarFunctionSet JulianDayFun::GetFunctions() { return operator_set; } +//! Binding date_part with a constant part replaces it with the unary part function (year, month, ...), so the bound +//! call has one argument and must be rendered under the replacement's own name +static unique_ptr DatePartUnbind(FunctionUnbindInput &input) { + auto &function = input.expression.Function(); + if (input.children.size() == 1) { + return make_uniq(function.GetQualifiedName(), std::move(input.children)); + } + if (input.children.size() != 2) { + return nullptr; + } + return make_uniq(function.GetDefinition()->GetQualifiedName(), std::move(input.children)); +} + // Names the "part,ts" pair shared by date_part's per-type overloads. static ScalarFunction NamePartTsArguments(ScalarFunction fun, const LogicalType &type) { fun.GetSignature().AddParameter("part", LogicalType::VARCHAR).AddParameter("ts", type); + fun.SetUnbindCallback(DatePartUnbind); return fun; } diff --git a/src/duckdb/extension/core_functions/scalar/struct/struct_update.cpp b/src/duckdb/extension/core_functions/scalar/struct/struct_update.cpp index 7b3a030a2..121c279d9 100644 --- a/src/duckdb/extension/core_functions/scalar/struct/struct_update.cpp +++ b/src/duckdb/extension/core_functions/scalar/struct/struct_update.cpp @@ -1,3 +1,4 @@ +#include "duckdb/parser/expression/function_expression.hpp" #include "duckdb/common/vector/map_vector.hpp" #include "duckdb/common/vector/struct_vector.hpp" #include "core_functions/scalar/struct_functions.hpp" @@ -140,11 +141,25 @@ static unique_ptr StructUpdateStats(ClientContext &context, Func return new_stats.ToUnique(); } +static unique_ptr StructUpdateUnbind(FunctionUnbindInput &input) { + vector arguments; + for (idx_t i = 0; i < input.children.size(); i++) { + auto name = i == 0 ? Identifier() : input.expression.GetChildren()[i]->GetAlias(); + if (i > 0 && name.empty()) { + return nullptr; + } + arguments.emplace_back(std::move(name), std::move(input.children[i])); + } + return make_uniq(input.expression.Function().GetDefinition()->GetQualifiedName(), + std::move(arguments)); +} + ScalarFunction StructUpdateFun::GetFunction() { ScalarFunction fun({}, LogicalTypeId::STRUCT, StructUpdateFunction, StructUpdateBind, StructUpdateStats); fun.SetNullHandling(FunctionNullHandling::SPECIAL_HANDLING); fun.GetSignature().AddParameter("struct", LogicalType::ANY).AddKwargsParameter("kwargs", LogicalType::ANY); fun.GetProperties().SetRequiresExpressionNames(true); + fun.SetUnbindCallback(StructUpdateUnbind); fun.SetSerializeCallback(VariableReturnBindData::Serialize); fun.SetDeserializeCallback(VariableReturnBindData::Deserialize); return fun; diff --git a/src/duckdb/extension/json/json_functions/read_single_json_file.cpp b/src/duckdb/extension/json/json_functions/read_single_json_file.cpp index 558cd8069..684004028 100644 --- a/src/duckdb/extension/json/json_functions/read_single_json_file.cpp +++ b/src/duckdb/extension/json/json_functions/read_single_json_file.cpp @@ -333,8 +333,9 @@ TableFunction JSONFunctions::GetJSONTableFunction(Identifier name, shared_ptr function_info) { diff --git a/src/duckdb/extension/parquet/column_reader.cpp b/src/duckdb/extension/parquet/column_reader.cpp index 7d0d116e6..0eab2daad 100644 --- a/src/duckdb/extension/parquet/column_reader.cpp +++ b/src/duckdb/extension/parquet/column_reader.cpp @@ -1012,6 +1012,20 @@ static unique_ptr CreateDecimalReader(const ParquetReader &reader, } } +static_assert(ParquetTimestampLogicalType(ParquetExtraTypeInfo::UNIT_NS) != LogicalTypeId::TIMESTAMP); +static_assert(ParquetTimestampTzLogicalType(ParquetExtraTypeInfo::UNIT_NS) != LogicalTypeId::TIMESTAMP_TZ); + +static_assert(ParquetTimestampLogicalType(ParquetExtraTypeInfo::IMPALA_TIMESTAMP) != LogicalTypeId::TIMESTAMP_NS); +static_assert(ParquetTimestampTzLogicalType(ParquetExtraTypeInfo::IMPALA_TIMESTAMP) != LogicalTypeId::TIMESTAMP_TZ_NS); +static_assert(ParquetTimestampLogicalType(ParquetExtraTypeInfo::UNIT_MS) != LogicalTypeId::TIMESTAMP_NS); +static_assert(ParquetTimestampTzLogicalType(ParquetExtraTypeInfo::UNIT_MS) != LogicalTypeId::TIMESTAMP_TZ_NS); +static_assert(ParquetTimestampLogicalType(ParquetExtraTypeInfo::UNIT_MICROS) != LogicalTypeId::TIMESTAMP_NS); +static_assert(ParquetTimestampTzLogicalType(ParquetExtraTypeInfo::UNIT_MICROS) != LogicalTypeId::TIMESTAMP_TZ_NS); + +static_assert(ParquetTimeLogicalType(ParquetExtraTypeInfo::UNIT_NS) != LogicalTypeId::TIME); +static_assert(ParquetTimeLogicalType(ParquetExtraTypeInfo::UNIT_MS) != LogicalTypeId::TIME_NS); +static_assert(ParquetTimeLogicalType(ParquetExtraTypeInfo::UNIT_MICROS) != LogicalTypeId::TIME_NS); + unique_ptr ColumnReader::CreateReader(const ParquetReader &reader, const ParquetColumnSchema &schema) { switch (schema.type.id()) { case LogicalTypeId::BOOLEAN: @@ -1052,22 +1066,12 @@ unique_ptr ColumnReader::CreateReader(const ParquetReader &reader, case ParquetExtraTypeInfo::UNIT_MICROS: return make_uniq>(reader, schema); - case ParquetExtraTypeInfo::UNIT_NS: - return make_uniq>(reader, schema); default: throw InternalException("TIMESTAMP requires type info"); } case LogicalTypeId::TIMESTAMP_NS: case LogicalTypeId::TIMESTAMP_TZ_NS: switch (schema.type_info) { - case ParquetExtraTypeInfo::IMPALA_TIMESTAMP: - return make_uniq>(reader, schema); - case ParquetExtraTypeInfo::UNIT_MS: - return make_uniq>(reader, - schema); - case ParquetExtraTypeInfo::UNIT_MICROS: - return make_uniq>(reader, - schema); case ParquetExtraTypeInfo::UNIT_NS: return make_uniq>(reader, schema); @@ -1082,17 +1086,11 @@ unique_ptr ColumnReader::CreateReader(const ParquetReader &reader, return make_uniq>(reader, schema); case ParquetExtraTypeInfo::UNIT_MICROS: return make_uniq>(reader, schema); - case ParquetExtraTypeInfo::UNIT_NS: - return make_uniq>(reader, schema); default: throw InternalException("TIME requires type info"); } case LogicalTypeId::TIME_NS: switch (schema.type_info) { - case ParquetExtraTypeInfo::UNIT_MS: - return make_uniq>(reader, schema); - case ParquetExtraTypeInfo::UNIT_MICROS: - return make_uniq>(reader, schema); case ParquetExtraTypeInfo::UNIT_NS: return make_uniq>(reader, schema); default: diff --git a/src/duckdb/extension/parquet/include/parquet_column_schema.hpp b/src/duckdb/extension/parquet/include/parquet_column_schema.hpp index 2850ab3ab..9a9fdeae0 100644 --- a/src/duckdb/extension/parquet/include/parquet_column_schema.hpp +++ b/src/duckdb/extension/parquet/include/parquet_column_schema.hpp @@ -45,6 +45,54 @@ enum class ParquetExtraTypeInfo { FLOAT16 }; +constexpr LogicalTypeId ParquetTimestampLogicalType(ParquetExtraTypeInfo type_info) { + switch (type_info) { + case ParquetExtraTypeInfo::IMPALA_TIMESTAMP: + case ParquetExtraTypeInfo::UNIT_MS: + case ParquetExtraTypeInfo::UNIT_MICROS: + return LogicalTypeId::TIMESTAMP; + case ParquetExtraTypeInfo::UNIT_NS: + return LogicalTypeId::TIMESTAMP_NS; + default: + return LogicalTypeId::INVALID; + } +} + +constexpr LogicalTypeId ParquetTimestampTzLogicalType(ParquetExtraTypeInfo type_info) { + switch (type_info) { + case ParquetExtraTypeInfo::UNIT_NS: + return LogicalTypeId::TIMESTAMP_TZ_NS; + case ParquetExtraTypeInfo::UNIT_MS: + case ParquetExtraTypeInfo::UNIT_MICROS: + return LogicalTypeId::TIMESTAMP_TZ; + default: + return LogicalTypeId::INVALID; + } +} + +constexpr LogicalTypeId ParquetTimeLogicalType(ParquetExtraTypeInfo type_info) { + switch (type_info) { + case ParquetExtraTypeInfo::UNIT_NS: + return LogicalTypeId::TIME_NS; + case ParquetExtraTypeInfo::UNIT_MS: + case ParquetExtraTypeInfo::UNIT_MICROS: + return LogicalTypeId::TIME; + default: + return LogicalTypeId::INVALID; + } +} + +constexpr LogicalTypeId ParquetTimeTzLogicalType(ParquetExtraTypeInfo type_info) { + switch (type_info) { + case ParquetExtraTypeInfo::UNIT_MS: + case ParquetExtraTypeInfo::UNIT_MICROS: + case ParquetExtraTypeInfo::UNIT_NS: + return LogicalTypeId::TIME_TZ; + default: + return LogicalTypeId::INVALID; + } +} + struct ParquetColumnSchema { public: ParquetColumnSchema() = default; diff --git a/src/duckdb/extension/parquet/include/parquet_timestamp.hpp b/src/duckdb/extension/parquet/include/parquet_timestamp.hpp index bc6067807..14cce5d9d 100644 --- a/src/duckdb/extension/parquet/include/parquet_timestamp.hpp +++ b/src/duckdb/extension/parquet/include/parquet_timestamp.hpp @@ -22,24 +22,17 @@ struct Int96 { }; timestamp_t ImpalaTimestampToTimestamp(const Int96 &raw_ts); -timestamp_ns_t ImpalaTimestampToTimestampNS(const Int96 &raw_ts); Int96 TimestampToImpalaTimestamp(timestamp_t &ts); timestamp_t ParquetTimestampMicrosToTimestamp(const int64_t &raw_ts); timestamp_t ParquetTimestampMsToTimestamp(const int64_t &raw_ts); -timestamp_t ParquetTimestampNsToTimestamp(const int64_t &raw_ts); -timestamp_ns_t ParquetTimestampMsToTimestampNs(const int64_t &raw_ms); -timestamp_ns_t ParquetTimestampUsToTimestampNs(const int64_t &raw_us); timestamp_ns_t ParquetTimestampNsToTimestampNs(const int64_t &raw_ns); date_t ParquetIntToDate(const int32_t &raw_date); dtime_t ParquetMsIntToTime(const int32_t &raw_millis); dtime_t ParquetIntToTime(const int64_t &raw_micros); -dtime_t ParquetNsIntToTime(const int64_t &raw_nanos); -dtime_ns_t ParquetMsIntToTimeNs(const int32_t &raw_millis); -dtime_ns_t ParquetUsIntToTimeNs(const int64_t &raw_micros); dtime_ns_t ParquetIntToTimeNs(const int64_t &raw_nanos); dtime_tz_t ParquetIntToTimeMsTZ(const int32_t &raw_millis); diff --git a/src/duckdb/extension/parquet/parquet_reader.cpp b/src/duckdb/extension/parquet/parquet_reader.cpp index ae8c51f40..b650b7289 100644 --- a/src/duckdb/extension/parquet/parquet_reader.cpp +++ b/src/duckdb/extension/parquet/parquet_reader.cpp @@ -73,6 +73,14 @@ namespace duckdb { +static_assert(ParquetTimestampTzLogicalType(ParquetExtraTypeInfo::UNIT_MS) == LogicalTypeId::TIMESTAMP_TZ); +static_assert(ParquetTimestampTzLogicalType(ParquetExtraTypeInfo::UNIT_MICROS) == LogicalTypeId::TIMESTAMP_TZ); +static_assert(ParquetTimestampTzLogicalType(ParquetExtraTypeInfo::UNIT_NS) == LogicalTypeId::TIMESTAMP_TZ_NS); + +static_assert(ParquetTimeTzLogicalType(ParquetExtraTypeInfo::UNIT_MS) == LogicalTypeId::TIME_TZ); +static_assert(ParquetTimeTzLogicalType(ParquetExtraTypeInfo::UNIT_MICROS) == LogicalTypeId::TIME_TZ); +static_assert(ParquetTimeTzLogicalType(ParquetExtraTypeInfo::UNIT_NS) == LogicalTypeId::TIME_TZ); + const char *ParquetPrefetchStrategyToString(ParquetPrefetchStrategy strategy) { switch (strategy) { case ParquetPrefetchStrategy::WHOLE_GROUP: @@ -422,14 +430,9 @@ LogicalType ParquetReader::DeriveLogicalType(const SchemaElement &s_ele, const P throw NotImplementedException("Unimplemented TIMESTAMP encoding - missing UNIT"); } if (s_ele.logicalType.TIMESTAMP.isAdjustedToUTC) { - if (s_ele.logicalType.TIMESTAMP.unit.__isset.NANOS) { - return LogicalType::TIMESTAMP_TZ_NS; - } - return LogicalType::TIMESTAMP_TZ; - } else if (s_ele.logicalType.TIMESTAMP.unit.__isset.NANOS) { - return LogicalType::TIMESTAMP_NS; + return LogicalType(ParquetTimestampTzLogicalType(schema.type_info)); } - return LogicalType::TIMESTAMP; + return LogicalType(ParquetTimestampLogicalType(schema.type_info)); } else if (s_ele.logicalType.__isset.TIME) { if (s_ele.logicalType.TIME.unit.__isset.MILLIS) { schema.type_info = ParquetExtraTypeInfo::UNIT_MS; @@ -441,11 +444,9 @@ LogicalType ParquetReader::DeriveLogicalType(const SchemaElement &s_ele, const P throw NotImplementedException("Unimplemented TIME encoding - missing UNIT"); } if (s_ele.logicalType.TIME.isAdjustedToUTC) { - return LogicalType::TIME_TZ; - } else if (s_ele.logicalType.TIME.unit.__isset.NANOS) { - return LogicalType::TIME_NS; + return LogicalType(ParquetTimeTzLogicalType(schema.type_info)); } - return LogicalType::TIME; + return LogicalType(ParquetTimeLogicalType(schema.type_info)); } } if (s_ele.__isset.converted_type) { @@ -511,14 +512,14 @@ LogicalType ParquetReader::DeriveLogicalType(const SchemaElement &s_ele, const P case ConvertedType::TIMESTAMP_MICROS: schema.type_info = ParquetExtraTypeInfo::UNIT_MICROS; if (s_ele.type == Type::INT64) { - return LogicalType::TIMESTAMP; + return LogicalType(ParquetTimestampLogicalType(schema.type_info)); } else { throw IOException("TIMESTAMP converted type can only be set for value of Type::INT64"); } case ConvertedType::TIMESTAMP_MILLIS: schema.type_info = ParquetExtraTypeInfo::UNIT_MS; if (s_ele.type == Type::INT64) { - return LogicalType::TIMESTAMP; + return LogicalType(ParquetTimestampLogicalType(schema.type_info)); } else { throw IOException("TIMESTAMP converted type can only be set for value of Type::INT64"); } @@ -559,14 +560,14 @@ LogicalType ParquetReader::DeriveLogicalType(const SchemaElement &s_ele, const P case ConvertedType::TIME_MILLIS: schema.type_info = ParquetExtraTypeInfo::UNIT_MS; if (s_ele.type == Type::INT32) { - return LogicalType::TIME; + return LogicalType(ParquetTimeLogicalType(schema.type_info)); } else { throw IOException("TIME_MILLIS converted type can only be set for value of Type::INT32"); } case ConvertedType::TIME_MICROS: schema.type_info = ParquetExtraTypeInfo::UNIT_MICROS; if (s_ele.type == Type::INT64) { - return LogicalType::TIME; + return LogicalType(ParquetTimeLogicalType(schema.type_info)); } else { throw IOException("TIME_MICROS converted type can only be set for value of Type::INT64"); } @@ -593,7 +594,7 @@ LogicalType ParquetReader::DeriveLogicalType(const SchemaElement &s_ele, const P return LogicalType::BIGINT; case Type::INT96: // always a timestamp it would seem schema.type_info = ParquetExtraTypeInfo::IMPALA_TIMESTAMP; - return LogicalType::TIMESTAMP; + return LogicalType(ParquetTimestampLogicalType(schema.type_info)); case Type::FLOAT: return LogicalType::FLOAT; case Type::DOUBLE: diff --git a/src/duckdb/extension/parquet/parquet_statistics.cpp b/src/duckdb/extension/parquet/parquet_statistics.cpp index b53a329d5..5743c0251 100644 --- a/src/duckdb/extension/parquet/parquet_statistics.cpp +++ b/src/duckdb/extension/parquet/parquet_statistics.cpp @@ -253,31 +253,19 @@ Value ParquetStatisticsUtils::ConvertValueInternal(const LogicalType &type, cons switch (schema_ele.type_info) { case ParquetExtraTypeInfo::UNIT_MS: return Value::TIME(Time::FromTimeMs(val)); - case ParquetExtraTypeInfo::UNIT_NS: - return Value::TIME(Time::FromTimeNs(val)); case ParquetExtraTypeInfo::UNIT_MICROS: default: return Value::TIME(dtime_t(val)); } } case LogicalTypeId::TIME_NS: { - int64_t val; - if (stats.size() == sizeof(int32_t)) { - val = Load(stats_data); - } else if (stats.size() == sizeof(int64_t)) { - val = Load(stats_data); - } else { + if (stats.size() != sizeof(int64_t)) { throw InvalidInputException("Incorrect stats size for type TIME_NS"); } - switch (schema_ele.type_info) { - case ParquetExtraTypeInfo::UNIT_MS: - return Value::TIME_NS(ParquetMsIntToTimeNs(NumericCast(val))); - case ParquetExtraTypeInfo::UNIT_NS: - return Value::TIME_NS(ParquetIntToTimeNs(val)); - case ParquetExtraTypeInfo::UNIT_MICROS: - default: - return Value::TIME_NS(dtime_ns_t(val)); + if (schema_ele.type_info != ParquetExtraTypeInfo::UNIT_NS) { + throw InternalException("TIME_NS requires nanosecond type info"); } + return Value::TIME_NS(ParquetIntToTimeNs(Load(stats_data))); } case LogicalTypeId::TIME_TZ: { int64_t val; @@ -315,9 +303,6 @@ Value ParquetStatisticsUtils::ConvertValueInternal(const LogicalType &type, cons case ParquetExtraTypeInfo::UNIT_MS: timestamp_value = ParquetTimestampMsToTimestamp(val); break; - case ParquetExtraTypeInfo::UNIT_NS: - timestamp_value = ParquetTimestampNsToTimestamp(val); - break; case ParquetExtraTypeInfo::UNIT_MICROS: default: timestamp_value = timestamp_t(val); @@ -331,30 +316,13 @@ Value ParquetStatisticsUtils::ConvertValueInternal(const LogicalType &type, cons } case LogicalTypeId::TIMESTAMP_TZ_NS: case LogicalTypeId::TIMESTAMP_NS: { - timestamp_ns_t timestamp_value; - if (schema_ele.type_info == ParquetExtraTypeInfo::IMPALA_TIMESTAMP) { - if (stats.size() != sizeof(Int96)) { - throw InvalidInputException("Incorrect stats size for type TIMESTAMP_NS"); - } - timestamp_value = ImpalaTimestampToTimestampNS(Load(stats_data)); - } else { - if (stats.size() != sizeof(int64_t)) { - throw InvalidInputException("Incorrect stats size for type TIMESTAMP_NS"); - } - auto val = Load(stats_data); - switch (schema_ele.type_info) { - case ParquetExtraTypeInfo::UNIT_MS: - timestamp_value = ParquetTimestampMsToTimestampNs(val); - break; - case ParquetExtraTypeInfo::UNIT_NS: - timestamp_value = ParquetTimestampNsToTimestampNs(val); - break; - case ParquetExtraTypeInfo::UNIT_MICROS: - default: - timestamp_value = ParquetTimestampUsToTimestampNs(val); - break; - } + if (stats.size() != sizeof(int64_t)) { + throw InvalidInputException("Incorrect stats size for type TIMESTAMP_NS"); + } + if (schema_ele.type_info != ParquetExtraTypeInfo::UNIT_NS) { + throw InternalException("TIMESTAMP_NS requires nanosecond type info"); } + auto timestamp_value = ParquetTimestampNsToTimestampNs(Load(stats_data)); if (type.id() == LogicalTypeId::TIMESTAMP_TZ_NS) { return Value::TIMESTAMPTZNS(timestamp_tz_ns_t(timestamp_value)); } diff --git a/src/duckdb/extension/parquet/parquet_timestamp.cpp b/src/duckdb/extension/parquet/parquet_timestamp.cpp index c95b34a11..5caba4dd8 100644 --- a/src/duckdb/extension/parquet/parquet_timestamp.cpp +++ b/src/duckdb/extension/parquet/parquet_timestamp.cpp @@ -14,7 +14,6 @@ static constexpr int64_t JULIAN_TO_UNIX_EPOCH_DAYS = 2440588LL; static constexpr int64_t MILLISECONDS_PER_DAY = 86400000LL; static constexpr int64_t MICROSECONDS_PER_DAY = MILLISECONDS_PER_DAY * 1000LL; static constexpr int64_t NANOSECONDS_PER_MICRO = 1000LL; -static constexpr int64_t NANOSECONDS_PER_DAY = MICROSECONDS_PER_DAY * 1000LL; static inline int64_t ImpalaTimestampToDays(const Int96 &impala_timestamp) { return impala_timestamp.value[2] - JULIAN_TO_UNIX_EPOCH_DAYS; @@ -27,18 +26,6 @@ static int64_t ImpalaTimestampToMicroseconds(const Int96 &impala_timestamp) { return days_since_epoch * MICROSECONDS_PER_DAY + microseconds; } -static int64_t ImpalaTimestampToNanoseconds(const Int96 &impala_timestamp) { - int64_t days_since_epoch = ImpalaTimestampToDays(impala_timestamp); - auto nanoseconds = Load(const_data_ptr_cast(impala_timestamp.value)); - return days_since_epoch * NANOSECONDS_PER_DAY + nanoseconds; -} - -timestamp_ns_t ImpalaTimestampToTimestampNS(const Int96 &raw_ts) { - timestamp_ns_t result; - result.value = ImpalaTimestampToNanoseconds(raw_ts); - return result; -} - timestamp_t ImpalaTimestampToTimestamp(const Int96 &raw_ts) { auto impala_us = ImpalaTimestampToMicroseconds(raw_ts); return Timestamp::FromEpochMicroSeconds(impala_us); @@ -70,38 +57,12 @@ timestamp_t ParquetTimestampMsToTimestamp(const int64_t &raw_ts) { return Timestamp::FromEpochMs(raw_ts); } -timestamp_ns_t ParquetTimestampMsToTimestampNs(const int64_t &raw_ms) { - timestamp_ns_t input; - input.value = raw_ms; - if (!input.IsFinite()) { - return input; - } - return Timestamp::TimestampNsFromEpochMillis(raw_ms); -} - -timestamp_ns_t ParquetTimestampUsToTimestampNs(const int64_t &raw_us) { - timestamp_ns_t input; - input.value = raw_us; - if (!input.IsFinite()) { - return input; - } - return Timestamp::TimestampNsFromEpochMicros(raw_us); -} - timestamp_ns_t ParquetTimestampNsToTimestampNs(const int64_t &raw_ns) { timestamp_ns_t result; result.value = raw_ns; return result; } -timestamp_t ParquetTimestampNsToTimestamp(const int64_t &raw_ts) { - timestamp_t input(raw_ts); - if (!input.IsFinite()) { - return input; - } - return Timestamp::FromEpochNanoSeconds(raw_ts); -} - date_t ParquetIntToDate(const int32_t &raw_date) { return date_t(raw_date); } @@ -124,18 +85,6 @@ dtime_t ParquetIntToTime(const int64_t &raw_micros) { return dtime_t(raw_micros); } -dtime_t ParquetNsIntToTime(const int64_t &raw_nanos) { - return Time::FromTimeNs(raw_nanos); -} - -dtime_ns_t ParquetMsIntToTimeNs(const int32_t &raw_millis) { - return dtime_ns_t(Interval::NANOS_PER_MSEC * raw_millis); -} - -dtime_ns_t ParquetUsIntToTimeNs(const int64_t &raw_micros) { - return dtime_ns_t(raw_micros * Interval::NANOS_PER_MICRO); -} - dtime_ns_t ParquetIntToTimeNs(const int64_t &raw_nanos) { return dtime_ns_t(raw_nanos); } diff --git a/src/duckdb/generated_extension_loader_package_build.cpp b/src/duckdb/generated_extension_loader_package_build.cpp index a2dfed398..0615ecaea 100644 --- a/src/duckdb/generated_extension_loader_package_build.cpp +++ b/src/duckdb/generated_extension_loader_package_build.cpp @@ -41,7 +41,7 @@ int32_t duckdb_extension_core_functions_describe(duckdb_extension_descriptor *de descriptor->version = 1; descriptor->name = "core_functions"; descriptor->extension_version = EXT_VERSION_CORE_FUNCTIONS; - descriptor->api_version = "v2.0.0-alpha43385"; + descriptor->api_version = "v2.0.0-alpha43586"; descriptor->entry_cpp = (void (*)(void))core_functions_duckdb_cpp_init; return 0; } @@ -74,7 +74,7 @@ int32_t duckdb_extension_parquet_describe(duckdb_extension_descriptor *descripto descriptor->version = 1; descriptor->name = "parquet"; descriptor->extension_version = EXT_VERSION_PARQUET; - descriptor->api_version = "v2.0.0-alpha43385"; + descriptor->api_version = "v2.0.0-alpha43586"; descriptor->entry_cpp = (void (*)(void))parquet_duckdb_cpp_init; return 0; } @@ -107,7 +107,7 @@ int32_t duckdb_extension_icu_describe(duckdb_extension_descriptor *descriptor) { descriptor->version = 1; descriptor->name = "icu"; descriptor->extension_version = EXT_VERSION_ICU; - descriptor->api_version = "v2.0.0-alpha43385"; + descriptor->api_version = "v2.0.0-alpha43586"; descriptor->entry_cpp = (void (*)(void))icu_duckdb_cpp_init; return 0; } @@ -140,7 +140,7 @@ int32_t duckdb_extension_json_describe(duckdb_extension_descriptor *descriptor) descriptor->version = 1; descriptor->name = "json"; descriptor->extension_version = EXT_VERSION_JSON; - descriptor->api_version = "v2.0.0-alpha43385"; + descriptor->api_version = "v2.0.0-alpha43586"; descriptor->entry_cpp = (void (*)(void))json_duckdb_cpp_init; return 0; } diff --git a/src/duckdb/src/catalog/catalog_entry/sequence_catalog_entry.cpp b/src/duckdb/src/catalog/catalog_entry/sequence_catalog_entry.cpp index acbd079b0..fec16d5d2 100644 --- a/src/duckdb/src/catalog/catalog_entry/sequence_catalog_entry.cpp +++ b/src/duckdb/src/catalog/catalog_entry/sequence_catalog_entry.cpp @@ -56,14 +56,15 @@ int64_t SequenceCatalogEntry::NextValue(DuckTransaction &transaction) { lock_guard seqlock(lock); int64_t result; result = data.counter; - bool overflow = !TryAddOperator::Operation(data.counter, data.increment, data.counter); + int64_t next_counter; + bool overflow = !TryAddOperator::Operation(data.counter, data.increment, next_counter); if (data.cycle) { if (overflow) { - data.counter = data.increment < 0 ? data.max_value : data.min_value; - } else if (data.counter < data.min_value) { - data.counter = data.max_value; - } else if (data.counter > data.max_value) { - data.counter = data.min_value; + next_counter = data.increment < 0 ? data.max_value : data.min_value; + } else if (next_counter < data.min_value) { + next_counter = data.max_value; + } else if (next_counter > data.max_value) { + next_counter = data.min_value; } } else { if (result < data.min_value || (overflow && data.increment < 0)) { @@ -73,6 +74,7 @@ int64_t SequenceCatalogEntry::NextValue(DuckTransaction &transaction) { throw SequenceException("nextval: reached maximum value of sequence \"%s\" (%lld)", name, data.max_value); } } + data.counter = next_counter; data.last_value = result; data.usage_count++; if (!temporary) { diff --git a/src/duckdb/src/common/enum_util.cpp b/src/duckdb/src/common/enum_util.cpp index efab5c457..39fcb2755 100644 --- a/src/duckdb/src/common/enum_util.cpp +++ b/src/duckdb/src/common/enum_util.cpp @@ -224,6 +224,7 @@ #include "duckdb/planner/bound_result_modifier.hpp" #include "duckdb/planner/filter/table_filter_functions.hpp" #include "duckdb/planner/logical_operator_repeatability.hpp" +#include "duckdb/planner/logical_plan_verification_result.hpp" #include "duckdb/planner/table_filter.hpp" #include "duckdb/storage/buffer/block_handle.hpp" #include "duckdb/storage/buffer/buffer_pool_reservation.hpp" @@ -1747,19 +1748,21 @@ const StringUtil::EnumStringLiteral *GetDebugStatementVerificationValues() { { static_cast(DebugStatementVerification::REPARSE_STATEMENT), "REPARSE_STATEMENT" }, { static_cast(DebugStatementVerification::SERIALIZE_STATEMENT), "SERIALIZE_STATEMENT" }, { static_cast(DebugStatementVerification::PREPARED_STATEMENT), "PREPARED_STATEMENT" }, - { static_cast(DebugStatementVerification::EXPLAIN_STATEMENT), "EXPLAIN_STATEMENT" } + { static_cast(DebugStatementVerification::EXPLAIN_STATEMENT), "EXPLAIN_STATEMENT" }, + { static_cast(DebugStatementVerification::EXPLAIN_SQL), "EXPLAIN_SQL" }, + { static_cast(DebugStatementVerification::EXPLAIN_SQL_STRICT), "EXPLAIN_SQL_STRICT" } }; return values; } template<> const char* EnumUtil::ToChars(DebugStatementVerification value) { - return StringUtil::EnumToString(GetDebugStatementVerificationValues(), 6, "DebugStatementVerification", static_cast(value)); + return StringUtil::EnumToString(GetDebugStatementVerificationValues(), 8, "DebugStatementVerification", static_cast(value)); } template<> DebugStatementVerification EnumUtil::FromString(const char *value) { - return static_cast(StringUtil::StringToEnum(GetDebugStatementVerificationValues(), 6, "DebugStatementVerification", value)); + return static_cast(StringUtil::StringToEnum(GetDebugStatementVerificationValues(), 8, "DebugStatementVerification", value)); } const StringUtil::EnumStringLiteral *GetDebugVectorVerificationValues() { @@ -2121,19 +2124,20 @@ ExplainOutputType EnumUtil::FromString(const char *value) { const StringUtil::EnumStringLiteral *GetExplainTypeValues() { static constexpr StringUtil::EnumStringLiteral values[] { { static_cast(ExplainType::EXPLAIN_STANDARD), "EXPLAIN_STANDARD" }, - { static_cast(ExplainType::EXPLAIN_ANALYZE), "EXPLAIN_ANALYZE" } + { static_cast(ExplainType::EXPLAIN_ANALYZE), "EXPLAIN_ANALYZE" }, + { static_cast(ExplainType::EXPLAIN_SQL), "EXPLAIN_SQL" } }; return values; } template<> const char* EnumUtil::ToChars(ExplainType value) { - return StringUtil::EnumToString(GetExplainTypeValues(), 2, "ExplainType", static_cast(value)); + return StringUtil::EnumToString(GetExplainTypeValues(), 3, "ExplainType", static_cast(value)); } template<> ExplainType EnumUtil::FromString(const char *value) { - return static_cast(StringUtil::StringToEnum(GetExplainTypeValues(), 2, "ExplainType", value)); + return static_cast(StringUtil::StringToEnum(GetExplainTypeValues(), 3, "ExplainType", value)); } const StringUtil::EnumStringLiteral *GetExponentTypeValues() { @@ -3634,6 +3638,70 @@ LogicalOperatorType EnumUtil::FromString(const char *value) return static_cast(StringUtil::StringToEnum(GetLogicalOperatorTypeValues(), 68, "LogicalOperatorType", value)); } +const StringUtil::EnumStringLiteral *GetLogicalPlanVerificationIssueCodeValues() { + static constexpr StringUtil::EnumStringLiteral values[] { + { static_cast(LogicalPlanVerificationIssueCode::INVALID_BINDING), "INVALID_BINDING" }, + { static_cast(LogicalPlanVerificationIssueCode::TYPE_MISMATCH), "TYPE_MISMATCH" }, + { static_cast(LogicalPlanVerificationIssueCode::UNSUPPORTED_OPERATOR), "UNSUPPORTED_OPERATOR" }, + { static_cast(LogicalPlanVerificationIssueCode::UNSUPPORTED_EXPRESSION), "UNSUPPORTED_EXPRESSION" }, + { static_cast(LogicalPlanVerificationIssueCode::UNSUPPORTED_FUNCTION), "UNSUPPORTED_FUNCTION" }, + { static_cast(LogicalPlanVerificationIssueCode::UNSUPPORTED_SOURCE), "UNSUPPORTED_SOURCE" }, + { static_cast(LogicalPlanVerificationIssueCode::UNSUPPORTED_EXTENSION), "UNSUPPORTED_EXTENSION" }, + { static_cast(LogicalPlanVerificationIssueCode::MALFORMED_EXTENSION_RESULT), "MALFORMED_EXTENSION_RESULT" }, + { static_cast(LogicalPlanVerificationIssueCode::UNSUPPORTED_EXPORT_FEATURE), "UNSUPPORTED_EXPORT_FEATURE" }, + { static_cast(LogicalPlanVerificationIssueCode::INTERNAL_INVARIANT), "INTERNAL_INVARIANT" } + }; + return values; +} + +template<> +const char* EnumUtil::ToChars(LogicalPlanVerificationIssueCode value) { + return StringUtil::EnumToString(GetLogicalPlanVerificationIssueCodeValues(), 10, "LogicalPlanVerificationIssueCode", static_cast(value)); +} + +template<> +LogicalPlanVerificationIssueCode EnumUtil::FromString(const char *value) { + return static_cast(StringUtil::StringToEnum(GetLogicalPlanVerificationIssueCodeValues(), 10, "LogicalPlanVerificationIssueCode", value)); +} + +const StringUtil::EnumStringLiteral *GetLogicalPlanVerificationPathComponentTypeValues() { + static constexpr StringUtil::EnumStringLiteral values[] { + { static_cast(LogicalPlanVerificationPathComponentType::OPERATOR_CHILD), "OPERATOR_CHILD" }, + { static_cast(LogicalPlanVerificationPathComponentType::OPERATOR_EXPRESSION), "OPERATOR_EXPRESSION" }, + { static_cast(LogicalPlanVerificationPathComponentType::EXPRESSION_CHILD), "EXPRESSION_CHILD" } + }; + return values; +} + +template<> +const char* EnumUtil::ToChars(LogicalPlanVerificationPathComponentType value) { + return StringUtil::EnumToString(GetLogicalPlanVerificationPathComponentTypeValues(), 3, "LogicalPlanVerificationPathComponentType", static_cast(value)); +} + +template<> +LogicalPlanVerificationPathComponentType EnumUtil::FromString(const char *value) { + return static_cast(StringUtil::StringToEnum(GetLogicalPlanVerificationPathComponentTypeValues(), 3, "LogicalPlanVerificationPathComponentType", value)); +} + +const StringUtil::EnumStringLiteral *GetLogicalPlanVerificationPhaseValues() { + static constexpr StringUtil::EnumStringLiteral values[] { + { static_cast(LogicalPlanVerificationPhase::VERIFY), "VERIFY" }, + { static_cast(LogicalPlanVerificationPhase::EXPRESSION_EXPORT), "EXPRESSION_EXPORT" }, + { static_cast(LogicalPlanVerificationPhase::PLAN_EXPORT), "PLAN_EXPORT" } + }; + return values; +} + +template<> +const char* EnumUtil::ToChars(LogicalPlanVerificationPhase value) { + return StringUtil::EnumToString(GetLogicalPlanVerificationPhaseValues(), 3, "LogicalPlanVerificationPhase", static_cast(value)); +} + +template<> +LogicalPlanVerificationPhase EnumUtil::FromString(const char *value) { + return static_cast(StringUtil::StringToEnum(GetLogicalPlanVerificationPhaseValues(), 3, "LogicalPlanVerificationPhase", value)); +} + const StringUtil::EnumStringLiteral *GetLogicalTypeIdValues() { static constexpr StringUtil::EnumStringLiteral values[] { { static_cast(LogicalTypeId::INVALID), "INVALID" }, diff --git a/src/duckdb/src/common/hive_partitioning.cpp b/src/duckdb/src/common/hive_partitioning.cpp index 594b6ff2f..b52f1ddf6 100644 --- a/src/duckdb/src/common/hive_partitioning.cpp +++ b/src/duckdb/src/common/hive_partitioning.cpp @@ -2,6 +2,7 @@ #include "duckdb/common/uhugeint.hpp" #include "duckdb/execution/expression_executor.hpp" +#include "duckdb/function/table_function.hpp" #include "duckdb/planner/expression/bound_columnref_expression.hpp" #include "duckdb/planner/expression/bound_constant_expression.hpp" #include "duckdb/planner/expression/bound_reference_expression.hpp" @@ -252,7 +253,8 @@ void HivePartitioning::ApplyFiltersToFileList(ClientContext &context, vector 0) { + geometries_remaining--; const auto le = reader.Read() == 1; const auto meta = reader.Read(le); @@ -957,7 +964,8 @@ WKBAnalysis AnalyzeWKB(BlobReader &reader) { case 5: // MULTILINESTRING case 6: // MULTIPOLYGON case 7: { // GEOMETRYCOLLECTION - reader.Skip(sizeof(uint32_t)); + const auto part_count = reader.Read(le); + geometries_remaining += part_count; result.size += sizeof(uint32_t); // part count } break; default: { @@ -966,6 +974,7 @@ WKBAnalysis AnalyzeWKB(BlobReader &reader) { } } } + result.has_trailing_data = !reader.IsAtEnd(); return result; } @@ -1051,7 +1060,23 @@ constexpr const idx_t Geometry::MAX_RECURSION_DEPTH; bool Geometry::FromBinary(const string_t &wkb, string_t &result, StringHeap &heap, bool strict) { BlobReader reader(wkb.GetData(), static_cast(wkb.GetSize())); - const auto analysis = AnalyzeWKB(reader); + WKBAnalysis analysis; + try { + analysis = AnalyzeWKB(reader); + } catch (InvalidInputException &) { + // Truncated input (e.g. a collection declaring more children than are present) must + // uphold the non-strict contract: report failure instead of throwing. + if (strict) { + throw; + } + return false; + } + if (analysis.has_trailing_data) { + if (strict) { + throw InvalidInputException("Unexpected trailing data at position %zu", reader.GetPosition()); + } + return false; + } if (analysis.any_unknown) { if (strict) { throw InvalidInputException("Unsupported geometry type in WKB"); diff --git a/src/duckdb/src/common/types/type_manager.cpp b/src/duckdb/src/common/types/type_manager.cpp index 228c5ba8e..2373d366e 100644 --- a/src/duckdb/src/common/types/type_manager.cpp +++ b/src/duckdb/src/common/types/type_manager.cpp @@ -12,6 +12,10 @@ CastFunctionSet &TypeManager::GetCastFunctions() { return *cast_functions; } +const CastFunctionSet &TypeManager::GetCastFunctions() const { + return *cast_functions; +} + static LogicalType TransformStringToUnboundType(const string &str, const ParserOptions &options) { if (StringUtil::Lower(str) == "null") { return LogicalType::SQLNULL; diff --git a/src/duckdb/src/common/types/vector.cpp b/src/duckdb/src/common/types/vector.cpp index 7f4ba5fb6..8d37df736 100644 --- a/src/duckdb/src/common/types/vector.cpp +++ b/src/duckdb/src/common/types/vector.cpp @@ -521,55 +521,16 @@ void Vector::Serialize(Serializer &serializer, bool compressed_serialization) { if (!serializer.ShouldSerialize(StorageVersion::V1_3_0)) { compressed_serialization = false; } - if (compressed_serialization) { - auto vtype = GetVectorType(); - if (vtype == VectorType::DICTIONARY_VECTOR && DictionaryVector::DictionarySize(*this).IsValid()) { - auto dict = Vector::Ref(DictionaryVector::Child(*this)); - if (dict.GetVectorType() == VectorType::FLAT_VECTOR) { - idx_t dict_count = DictionaryVector::DictionarySize(*this).GetIndex(); - auto old_sel = DictionaryVector::SelVector(*this); - SelectionVector new_sel(count), used_sel(count), map_sel(dict_count); - - // dictionaries may be large (row-group level). A vector may use only a small part. - // So, restrict dict to the used_sel subset & remap old_sel into new_sel to the new dict positions - sel_t CODE_UNSEEN = static_cast(dict_count); - for (sel_t i = 0; i < dict_count; ++i) { - map_sel[i] = CODE_UNSEEN; // initialize with unused marker - } - idx_t used_count = 0; - for (idx_t i = 0; i < count; ++i) { - auto pos = old_sel[i]; - if (map_sel[pos] == CODE_UNSEEN) { - map_sel[pos] = static_cast(used_count); - used_sel[used_count++] = pos; - } - new_sel[i] = map_sel[pos]; - } - if (used_count * 2 < count) { // only serialize as a dict vector if that makes things smaller - auto sel_data = reinterpret_cast(new_sel.data()); - dict.Slice(used_sel, used_count); - serializer.WriteProperty(90, "vector_type", VectorType::DICTIONARY_VECTOR); - serializer.WriteProperty(91, "sel_vector", sel_data, sizeof(sel_t) * count); - serializer.WriteProperty(92, "dict_count", used_count); - return dict.Serialize(serializer, false); - } - } - } else if (vtype == VectorType::CONSTANT_VECTOR && count >= 1) { - serializer.WriteProperty(90, "vector_type", VectorType::CONSTANT_VECTOR); - // Resize to 1 so that size() == count == 1 during the recursive call, then restore - FlatVector::SetSize(*this, 1); - Vector::Serialize(serializer, false); // just serialize one value - FlatVector::SetSize(*this, count); - return; - } else if (vtype == VectorType::SEQUENCE_VECTOR) { - serializer.WriteProperty(90, "vector_type", VectorType::SEQUENCE_VECTOR); - auto &sequence = buffer->Cast(); - serializer.WriteProperty(91, "seq_start", sequence.start); - serializer.WriteProperty(92, "seq_increment", sequence.increment); - return; // for sequence vectors we do not serialize anything else - } else { - // TODO: other compressed vector types (SHREDDED, FSST) - } + if (compressed_serialization && GetVectorType() == VectorType::CONSTANT_VECTOR && count >= 1) { + serializer.WriteProperty(90, "vector_type", VectorType::CONSTANT_VECTOR); + // Resize to 1 so that size() == count == 1 during the recursive call, then restore + FlatVector::SetSize(*this, 1); + Vector::Serialize(serializer, false); // just serialize one value + FlatVector::SetSize(*this, count); + return; + } + if (buffer && buffer->TrySerialize(serializer, logical_type, compressed_serialization)) { + return; } ToUnifiedFormat(vdata); @@ -740,17 +701,8 @@ void Vector::Deserialize(Deserializer &deserializer, idx_t count) { Vector::Deserialize(deserializer, 1); // read a vector of size 1 Vector::SetVectorType(VectorType::CONSTANT_VECTOR); return; - } else if (vtype == VectorType::DICTIONARY_VECTOR) { - SelectionVector sel(count); - deserializer.ReadProperty(91, "sel_vector", reinterpret_cast(sel.data()), sizeof(sel_t) * count); - const auto dict_count = deserializer.ReadProperty(92, "dict_count"); - Vector::Deserialize(deserializer, dict_count); // deserialize the dictionary in this vector - Vector::Slice(sel, count); // will create a dictionary vector - return; - } else if (vtype == VectorType::SEQUENCE_VECTOR) { - const int64_t seq_start = deserializer.ReadProperty(91, "seq_start"); - const int64_t seq_increment = deserializer.ReadProperty(92, "seq_increment"); - Vector::Sequence(seq_start, seq_increment, count); + } else if (vtype != VectorType::FLAT_VECTOR) { + buffer = VectorBuffer::Deserialize(deserializer, vtype, logical_type, count); return; } diff --git a/src/duckdb/src/common/types/vector_buffer.cpp b/src/duckdb/src/common/types/vector_buffer.cpp index b4a8ecaa0..637c2fa10 100644 --- a/src/duckdb/src/common/types/vector_buffer.cpp +++ b/src/duckdb/src/common/types/vector_buffer.cpp @@ -1,15 +1,18 @@ #include "duckdb/common/vector/array_vector.hpp" #include "duckdb/common/vector/constant_vector.hpp" +#include "duckdb/common/vector/dictionary_vector.hpp" #include "duckdb/common/vector/flat_vector.hpp" #include "duckdb/common/vector/fsst_vector.hpp" #include "duckdb/common/vector/list_vector.hpp" #include "duckdb/common/vector/map_vector.hpp" +#include "duckdb/common/vector/sequence_vector.hpp" #include "duckdb/common/vector/shredded_vector.hpp" #include "duckdb/common/vector/string_vector.hpp" #include "duckdb/common/vector/struct_vector.hpp" #include "duckdb/common/types/vector_buffer.hpp" #include "duckdb/common/assert.hpp" +#include "duckdb/common/enum_util.hpp" #include "duckdb/common/types/vector.hpp" #include "duckdb/common/vector_operations/vector_operations.hpp" #include "duckdb/storage/buffer/buffer_handle.hpp" @@ -145,6 +148,24 @@ void VectorBuffer::ToUnifiedFormat(UnifiedVectorFormat &format) const { throw InternalException("ToUnifiedFormat not supported for this buffer type - flatten first"); } +bool VectorBuffer::TrySerialize(Serializer &serializer, const LogicalType &type, bool compressed_serialization) const { + return false; +} + +buffer_ptr VectorBuffer::Deserialize(Deserializer &deserializer, VectorType vector_type, + const LogicalType &type, idx_t count) { + switch (vector_type) { + case VectorType::DICTIONARY_VECTOR: + return DictionaryBuffer::Deserialize(deserializer, type, count); + case VectorType::SEQUENCE_VECTOR: + return SequenceBuffer::Deserialize(deserializer, type, count); + case VectorType::SHREDDED_VECTOR: + return ShreddedVectorBuffer::Deserialize(deserializer, type, count); + default: + throw SerializationException("Unsupported vector type %s for deserialization", EnumUtil::ToString(vector_type)); + } +} + buffer_ptr VectorBuffer::Flatten(const LogicalType &type) const { auto result = FlattenSliceInternal(type, *FlatVector::IncrementalSelectionVector(), Size()); if (result && (result->Size() != Size())) { diff --git a/src/duckdb/src/common/value_operations/comparison_operations.cpp b/src/duckdb/src/common/value_operations/comparison_operations.cpp index a89f890ad..aa4f313c4 100644 --- a/src/duckdb/src/common/value_operations/comparison_operations.cpp +++ b/src/duckdb/src/common/value_operations/comparison_operations.cpp @@ -187,7 +187,7 @@ static bool TemplatedBooleanOperation(const Value &left, const Value &right) { return false; } } - return true; + return ValuePositionComparator::TieBreak(left_children.size(), right_children.size()); } default: throw InternalException("Unimplemented type for value comparison"); diff --git a/src/duckdb/src/common/vector/dictionary_vector.cpp b/src/duckdb/src/common/vector/dictionary_vector.cpp index b0cc8a492..bb0c8e40a 100644 --- a/src/duckdb/src/common/vector/dictionary_vector.cpp +++ b/src/duckdb/src/common/vector/dictionary_vector.cpp @@ -3,6 +3,8 @@ #include "duckdb/common/types/uuid.hpp" #include "duckdb/common/vector_operations/vector_operations.hpp" #include "duckdb/common/types/sel_cache.hpp" +#include "duckdb/common/serializer/deserializer.hpp" +#include "duckdb/common/serializer/serializer.hpp" namespace duckdb { @@ -192,4 +194,58 @@ const Vector &DictionaryVector::GetCachedHashes(const Vector &input) { return *entry.cached_hashes; } +bool DictionaryBuffer::TrySerialize(Serializer &serializer, const LogicalType &type, + bool compressed_serialization) const { + auto dictionary_size = GetDictionarySize(); + if (!compressed_serialization || !dictionary_size.IsValid()) { + return false; + } + auto dict = Vector::Ref(entry->data); + if (dict.GetVectorType() != VectorType::FLAT_VECTOR) { + return false; + } + auto count = Size(); + auto dict_count = dictionary_size.GetIndex(); + SelectionVector new_sel(count), used_sel(count), map_sel(dict_count); + + // dictionaries may be large (row-group level). A vector may use only a small part. + // So, restrict dict to the used_sel subset & remap old_sel into new_sel to the new dict positions + sel_t CODE_UNSEEN = static_cast(dict_count); + for (sel_t i = 0; i < dict_count; ++i) { + map_sel[i] = CODE_UNSEEN; // initialize with unused marker + } + idx_t used_count = 0; + for (idx_t i = 0; i < count; ++i) { + auto pos = sel_vector[i]; + if (map_sel[pos] == CODE_UNSEEN) { + map_sel[pos] = static_cast(used_count); + used_sel[used_count++] = pos; + } + new_sel[i] = map_sel[pos]; + } + if (used_count * 2 >= count) { + // only serialize as a dict vector if that makes things smaller + return false; + } + auto sel_data = reinterpret_cast(new_sel.data()); + dict.Slice(used_sel, used_count); + serializer.WriteProperty(90, "vector_type", VectorType::DICTIONARY_VECTOR); + serializer.WriteProperty(91, "sel_vector", sel_data, sizeof(sel_t) * count); + serializer.WriteProperty(92, "dict_count", used_count); + dict.Serialize(serializer, false); + return true; +} + +buffer_ptr DictionaryBuffer::Deserialize(Deserializer &deserializer, const LogicalType &type, + idx_t count) { + SelectionVector sel(count); + deserializer.ReadProperty(91, "sel_vector", reinterpret_cast(sel.data()), sizeof(sel_t) * count); + const auto dict_count = deserializer.ReadProperty(92, "dict_count"); + Vector dict(type, MaxValue(dict_count, STANDARD_VECTOR_SIZE)); + dict.Deserialize(deserializer, dict_count); + FlatVector::SetSize(dict, dict_count); + dict.Slice(sel, count); + return dict.GetBufferRef(); +} + } // namespace duckdb diff --git a/src/duckdb/src/common/vector/sequence_vector.cpp b/src/duckdb/src/common/vector/sequence_vector.cpp index 049eac486..f2894c1cd 100644 --- a/src/duckdb/src/common/vector/sequence_vector.cpp +++ b/src/duckdb/src/common/vector/sequence_vector.cpp @@ -1,5 +1,7 @@ #include "duckdb/common/vector/sequence_vector.hpp" #include "duckdb/common/vector_operations/vector_operations.hpp" +#include "duckdb/common/serializer/deserializer.hpp" +#include "duckdb/common/serializer/serializer.hpp" namespace duckdb { @@ -49,4 +51,21 @@ void SequenceVector::GetSequence(const Vector &vector, int64_t &start, int64_t & increment = data.increment; } +bool SequenceBuffer::TrySerialize(Serializer &serializer, const LogicalType &type, + bool compressed_serialization) const { + if (!compressed_serialization) { + return false; + } + serializer.WriteProperty(90, "vector_type", VectorType::SEQUENCE_VECTOR); + serializer.WriteProperty(91, "seq_start", start); + serializer.WriteProperty(92, "seq_increment", increment); + return true; +} + +buffer_ptr SequenceBuffer::Deserialize(Deserializer &deserializer, const LogicalType &type, idx_t count) { + const auto seq_start = deserializer.ReadProperty(91, "seq_start"); + const auto seq_increment = deserializer.ReadProperty(92, "seq_increment"); + return make_buffer(seq_start, seq_increment, count_t(count)); +} + } // namespace duckdb diff --git a/src/duckdb/src/common/vector/shredded_vector.cpp b/src/duckdb/src/common/vector/shredded_vector.cpp index 89c4096f5..5a621dced 100644 --- a/src/duckdb/src/common/vector/shredded_vector.cpp +++ b/src/duckdb/src/common/vector/shredded_vector.cpp @@ -1,5 +1,8 @@ #include "duckdb/common/vector/shredded_vector.hpp" #include "duckdb/common/vector/struct_vector.hpp" +#include "duckdb/common/vector/flat_vector.hpp" +#include "duckdb/common/serializer/deserializer.hpp" +#include "duckdb/common/serializer/serializer.hpp" #include "duckdb/function/scalar/variant_utils.hpp" namespace duckdb { @@ -39,6 +42,36 @@ void ShreddedVectorBuffer::SetVectorType(VectorType new_vector_type) { throw InternalException("ShreddedVectorBuffer::SetVectorType is not implemented and shouldn't be reached"); } +bool ShreddedVectorBuffer::TrySerialize(Serializer &serializer, const LogicalType &type, + bool compressed_serialization) const { + // older versions cannot read shredded vectors - these are unshredded prior to serialization instead + if (!serializer.ShouldSerialize(StorageVersion::V2_0_0)) { + return false; + } + serializer.WriteProperty(90, "vector_type", VectorType::SHREDDED_VECTOR); + serializer.WriteProperty(91, "shredded_type", shredded_data->GetType()); + serializer.WriteObject(92, "shredded_data", [&](Serializer &object) { + auto serialized_vector = Vector::Ref(*shredded_data); + serialized_vector.Serialize(object, compressed_serialization); + }); + return true; +} + +buffer_ptr ShreddedVectorBuffer::Deserialize(Deserializer &deserializer, const LogicalType &type, + idx_t count) { + if (type.id() != LogicalTypeId::VARIANT) { + throw SerializationException("Shredded vectors can only be deserialized as VARIANT vectors"); + } + auto shredded_type = deserializer.ReadProperty(91, "shredded_type"); + if (shredded_type.id() != LogicalTypeId::STRUCT || StructType::GetChildCount(shredded_type) != 2) { + throw SerializationException("Shredded vector data must be a struct with two children"); + } + Vector shredded_data(shredded_type, MaxValue(count, STANDARD_VECTOR_SIZE)); + deserializer.ReadObject(92, "shredded_data", [&](Deserializer &obj) { shredded_data.Deserialize(obj, count); }); + FlatVector::SetSize(shredded_data, count); + return make_buffer(shredded_data, count_t(count)); +} + Value ShreddedVectorBuffer::GetValue(const LogicalType &type, idx_t index) const { // FIXME: this is extremely inefficient auto &shredded = StructVector::GetEntries(*shredded_data)[1]; diff --git a/src/duckdb/src/execution/mark_join_row_comparison.cpp b/src/duckdb/src/execution/mark_join_row_comparison.cpp index d21407ad2..e7de97107 100644 --- a/src/duckdb/src/execution/mark_join_row_comparison.cpp +++ b/src/duckdb/src/execution/mark_join_row_comparison.cpp @@ -3,7 +3,6 @@ #include "duckdb/common/vector/constant_vector.hpp" #include "duckdb/common/vector/struct_vector.hpp" #include "duckdb/common/vector/flat_vector.hpp" -#include "duckdb/planner/joinside.hpp" #include "duckdb/common/vector_operations/vector_operations.hpp" namespace duckdb { @@ -104,86 +103,4 @@ void MarkJoinRowComparison::Compare(const Vector &left, const Vector &right, Exp } } -MarkJoinRowComparison::MarkJoinRowComparison(const DataChunk &left) : comparison(LogicalType::BOOLEAN) { - left_reference.Initialize(Allocator::DefaultAllocator(), left.GetTypes()); -} - -void MarkJoinRowComparison::CompareConjunction(DataChunk &left, idx_t left_row, DataChunk &right, - const vector &conditions, Vector &result) { - D_ASSERT(left_row < left.size()); - D_ASSERT(right.size() <= STANDARD_VECTOR_SIZE); - D_ASSERT(left.ColumnCount() == conditions.size()); - D_ASSERT(right.ColumnCount() == conditions.size()); - left_reference.Reset(); - bool pair_is_false[STANDARD_VECTOR_SIZE] = {false}; - bool pair_is_unknown[STANDARD_VECTOR_SIZE] = {false}; - for (idx_t condition_idx = 0; condition_idx < conditions.size(); condition_idx++) { - const auto type = conditions[condition_idx].GetComparisonType(); - if (left.data[condition_idx].GetType().id() == LogicalTypeId::TUPLE && - (type == ExpressionType::COMPARE_EQUAL || type == ExpressionType::COMPARE_NOTEQUAL)) { - bool is_false[STANDARD_VECTOR_SIZE] = {false}; - bool is_unknown[STANDARD_VECTOR_SIZE] = {false}; - CompareEquality(left.data[condition_idx], left_row, left.size(), right.data[condition_idx], right.size(), - is_false, is_unknown); - comparison.SetVectorType(VectorType::FLAT_VECTOR); - FlatVector::ValidityMutable(comparison).Reset(right.size()); - auto writer = FlatVector::Writer(comparison, right.size()); - for (idx_t row = 0; row < right.size(); row++) { - if (is_unknown[row]) { - writer.WriteNull(); - } else { - writer.WriteValue(type == ExpressionType::COMPARE_EQUAL ? !is_false[row] : is_false[row]); - } - } - } else { - ConstantVector::Reference(left_reference.data[condition_idx], count_t(right.size()), - left.data[condition_idx], left_row, left.size()); - Compare(left_reference.data[condition_idx], right.data[condition_idx], type, comparison); - } - auto entries = comparison.Values(); - for (idx_t right_row = 0; right_row < right.size(); right_row++) { - auto entry = entries[right_row]; - if (!entry.IsValid()) { - pair_is_unknown[right_row] = true; - } else if (!entry.GetValue()) { - pair_is_false[right_row] = true; - } - } - } - result.SetVectorType(VectorType::FLAT_VECTOR); - FlatVector::ValidityMutable(result).Reset(right.size()); - auto writer = FlatVector::Writer(result, right.size()); - for (idx_t right_row = 0; right_row < right.size(); right_row++) { - if (pair_is_false[right_row]) { - writer.WriteValue(false); - } else if (pair_is_unknown[right_row]) { - writer.WriteNull(); - } else { - writer.WriteValue(true); - } - } -} - -void MarkJoinRowComparison::Perform(DataChunk &left, DataChunk &right, bool found_match[], - const vector &conditions, optional_ptr found_unknown) { - Vector comparison(LogicalType::BOOLEAN); - MarkJoinRowComparison comparer(left); - for (idx_t left_row = 0; left_row < left.size(); left_row++) { - if (found_match[left_row]) { - continue; - } - comparer.CompareConjunction(left, left_row, right, conditions, comparison); - for (auto entry : comparison.Values()) { - if (entry.IsValid()) { - if (entry.GetValue()) { - found_match[left_row] = true; - break; - } - } else if (found_unknown) { - found_unknown.get()[left_row] = true; - } - } - } -} - } // namespace duckdb diff --git a/src/duckdb/src/execution/nested_loop_join/nested_loop_join_mark.cpp b/src/duckdb/src/execution/nested_loop_join/nested_loop_join_mark.cpp index 785b3d7fc..e5cb1499c 100644 --- a/src/duckdb/src/execution/nested_loop_join/nested_loop_join_mark.cpp +++ b/src/duckdb/src/execution/nested_loop_join/nested_loop_join_mark.cpp @@ -183,6 +183,7 @@ static void MarkJoinComparisonSwitch(const Vector &left, const Vector &right, id void NestedLoopJoinMark::Perform(DataChunk &left, ColumnDataCollection &right, bool found_match[], const vector &conditions, optional_ptr found_unknown) { + D_ASSERT(conditions.size() == 1); // initialize a new temporary selection vector for the left chunk // loop over all chunks in the RHS ColumnDataScanState scan_state; @@ -192,11 +193,6 @@ void NestedLoopJoinMark::Perform(DataChunk &left, ColumnDataCollection &right, b right.InitializeScanChunk(scan_chunk); while (right.Scan(scan_state, scan_chunk)) { - if (conditions.size() > 1) { - MarkJoinRowComparison::Perform(left, scan_chunk, found_match, conditions, found_unknown); - continue; - } - D_ASSERT(conditions.size() == 1); MarkJoinComparisonSwitch(left.data[0], scan_chunk.data[0], left.size(), scan_chunk.size(), found_match, conditions[0].GetComparisonType(), found_unknown); } diff --git a/src/duckdb/src/execution/operator/join/physical_hash_join.cpp b/src/duckdb/src/execution/operator/join/physical_hash_join.cpp index aa50361fa..2083b0878 100644 --- a/src/duckdb/src/execution/operator/join/physical_hash_join.cpp +++ b/src/duckdb/src/execution/operator/join/physical_hash_join.cpp @@ -51,6 +51,7 @@ PhysicalHashJoin::PhysicalHashJoin(PhysicalPlan &physical_plan, LogicalOperator : PhysicalComparisonJoin(physical_plan, op, PhysicalOperatorType::HASH_JOIN, std::move(conds), join_type, estimated_cardinality), delim_types(std::move(delim_types)) { + D_ASSERT(join_type != JoinType::MARK || !predicate); filter_pushdown = std::move(pushdown_info_p); children.push_back(left); diff --git a/src/duckdb/src/execution/operator/join/physical_nested_loop_join.cpp b/src/duckdb/src/execution/operator/join/physical_nested_loop_join.cpp index 749e8c856..4797ca49c 100644 --- a/src/duckdb/src/execution/operator/join/physical_nested_loop_join.cpp +++ b/src/duckdb/src/execution/operator/join/physical_nested_loop_join.cpp @@ -19,11 +19,10 @@ PhysicalNestedLoopJoin::PhysicalNestedLoopJoin(PhysicalPlan &physical_plan, Logi : PhysicalComparisonJoin(physical_plan, op, PhysicalOperatorType::NESTED_LOOP_JOIN, std::move(conds), join_type, estimated_cardinality) { filter_pushdown = std::move(pushdown_info_p); + D_ASSERT(join_type != JoinType::MARK || (conditions.size() == 1 && !predicate)); track_unknown = - join_type == JoinType::MARK && - (predicate || conditions.size() > 1 || - (conditions.size() == 1 && (conditions[0].GetLHS().GetReturnType().id() == LogicalTypeId::TUPLE || - conditions[0].GetComparisonType() == ExpressionType::COMPARE_DISTINCT_FROM))); + join_type == JoinType::MARK && (conditions[0].GetLHS().GetReturnType().id() == LogicalTypeId::TUPLE || + conditions[0].GetComparisonType() == ExpressionType::COMPARE_DISTINCT_FROM); children.push_back(left); children.push_back(right); if (join_type == JoinType::MARK) { @@ -460,13 +459,14 @@ OperatorResultType PhysicalNestedLoopJoin::ExecuteInternal(ExecutionContext &con } } -static void ResolveSimpleJoinPredicate(const vector &conditions, JoinType join_type, DataChunk &input, +static void ResolveSimpleJoinPredicate(const vector &conditions, DataChunk &input, PhysicalNestedLoopJoinState &state, NestedLoopJoinGlobalState &gstate, - bool found_match[], bool found_unknown[]) { + bool found_match[]) { + D_ASSERT(conditions.size() == 1); gstate.right_condition_data.InitializeScan(state.condition_scan_state); gstate.right_payload_data.InitializeScan(state.payload_scan_state); Vector comparison(LogicalType::BOOLEAN); - MarkJoinRowComparison comparer(state.left_condition); + Vector left_reference(state.left_condition.data[0].GetType()); VectorCache predicate_cache(state.pred_executor.GetAllocator(), LogicalType::BOOLEAN); Vector predicate_result(predicate_cache); while (gstate.right_condition_data.Scan(state.condition_scan_state, state.right_condition)) { @@ -478,12 +478,15 @@ static void ResolveSimpleJoinPredicate(const vector &conditions, if (found_match[left_row]) { continue; } - comparer.CompareConjunction(state.left_condition, left_row, state.right_condition, conditions, comparison); + ConstantVector::Reference(left_reference, count_t(state.right_condition.size()), + state.left_condition.data[0], left_row, state.left_condition.size()); + MarkJoinRowComparison::Compare(left_reference, state.right_condition.data[0], + conditions[0].GetComparisonType(), comparison); auto comparisons = comparison.Values(); idx_t candidate_count = 0; for (idx_t right_row = 0; right_row < state.right_condition.size(); right_row++) { auto entry = comparisons[right_row]; - if (entry.IsValid() ? entry.GetValue() : join_type == JoinType::MARK) { + if (entry.IsValid() && entry.GetValue()) { state.pred_matches.set_index(candidate_count++, right_row); } } @@ -502,14 +505,10 @@ static void ResolveSimpleJoinPredicate(const vector &conditions, auto predicates = predicate_result.Values(); for (idx_t candidate = 0; candidate < candidate_count; candidate++) { auto predicate = predicates[candidate]; - if (predicate.IsValid() && !predicate.GetValue()) { - continue; - } - if (predicate.IsValid() && comparisons[state.pred_matches.get_index(candidate)].IsValid()) { + if (predicate.IsValid() && predicate.GetValue()) { found_match[left_row] = true; break; } - found_unknown[left_row] = true; } } } @@ -528,7 +527,8 @@ void PhysicalNestedLoopJoin::ResolveSimpleJoin(ExecutionContext &context, DataCh bool found_unknown[STANDARD_VECTOR_SIZE] = {false}; if (predicate) { - ResolveSimpleJoinPredicate(conditions, join_type, input, state, gstate, found_match, found_unknown); + D_ASSERT(join_type == JoinType::SEMI || join_type == JoinType::ANTI); + ResolveSimpleJoinPredicate(conditions, input, state, gstate, found_match); } else { NestedLoopJoinMark::Perform(state.left_condition, gstate.right_condition_data, found_match, conditions, track_unknown ? optional_ptr(found_unknown) : nullptr); diff --git a/src/duckdb/src/function/aggregate/distributive/minmax.cpp b/src/duckdb/src/function/aggregate/distributive/minmax.cpp index ae8c2664b..7978c2d55 100644 --- a/src/duckdb/src/function/aggregate/distributive/minmax.cpp +++ b/src/duckdb/src/function/aggregate/distributive/minmax.cpp @@ -1,3 +1,4 @@ +#include "duckdb/parser/expression/function_expression.hpp" #include "duckdb/catalog/catalog_entry/aggregate_function_catalog_entry.hpp" #include "duckdb/catalog/catalog.hpp" #include "duckdb/function/aggregate_state_layout.hpp" @@ -408,11 +409,25 @@ unique_ptr BindMinMax(BindAggregateFunctionInput &input) { return std::move(expr->BindInfoMutable()); } +//! BindMinMax replaces a collated min/max with arg_min/arg_max over the collation key, so the bound call has one +//! more argument than the definition and must be rendered under the replacement's own name +static unique_ptr MinMaxUnbind(AggregateFunctionUnbindInput &input) { + auto &function = input.expression.Function(); + auto name = function.GetDefinition()->GetQualifiedName(); + if (input.children.size() == function.GetLogicalArguments().size() + 1) { + name = function.GetQualifiedName(); + } else if (input.children.size() != function.GetLogicalArguments().size()) { + return nullptr; + } + return make_uniq(name, std::move(input.children)); +} + template AggregateFunction GetMinMaxOperator(const string &name) { AggregateFunction fun(Identifier(name), {}, LogicalType::ANY, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, BindMinMax); fun.GetSignature().AddParameter("arg", LogicalTypeId::ANY); + fun.SetUnbindCallback(MinMaxUnbind); return fun; } @@ -583,6 +598,7 @@ AggregateFunction GetMinMaxNFunction() { AggregateFunction fun({}, LogicalType::LIST(LogicalType::ANY), nullptr, nullptr, nullptr, nullptr, nullptr, FunctionNullHandling::DEFAULT_NULL_HANDLING, nullptr, MinMaxNBind, nullptr); fun.GetSignature().AddParameter("arg", LogicalTypeId::ANY).AddParameter("n", LogicalType::BIGINT); + fun.SetUnbindCallback(MinMaxUnbind); return fun; } diff --git a/src/duckdb/src/function/aggregate_function.cpp b/src/duckdb/src/function/aggregate_function.cpp index 257281b02..6fec60cf4 100644 --- a/src/duckdb/src/function/aggregate_function.cpp +++ b/src/duckdb/src/function/aggregate_function.cpp @@ -1,3 +1,4 @@ +#include "duckdb/parser/expression/function_expression.hpp" #include "duckdb/function/aggregate_function.hpp" #include "duckdb/execution/operator/aggregate/aggregate_object.hpp" @@ -8,6 +9,14 @@ namespace duckdb { +AggregateFunctionUnbindInput::AggregateFunctionUnbindInput(const BoundAggregateExpression &expression_p, + vector> children_p) + : expression(expression_p), children(std::move(children_p)) { +} + +AggregateFunctionUnbindInput::~AggregateFunctionUnbindInput() { +} + unique_ptr AggregateFunction::PropagateInputValueStats(ClientContext &context, BoundAggregateExpression &expr, AggregateStatisticsInput &input) { @@ -85,7 +94,7 @@ bool AggregateFunctionCallbacks::operator==(const AggregateFunctionCallbacks &rh combine == rhs.combine && finalize == rhs.finalize && init_local_state_finalize == rhs.init_local_state_finalize && cluster_update == rhs.cluster_update && window == rhs.window && window_init == rhs.window_init && window_batch == rhs.window_batch && - bind == rhs.bind && destructor == rhs.destructor && statistics == rhs.statistics && + bind == rhs.bind && unbind == rhs.unbind && destructor == rhs.destructor && statistics == rhs.statistics && serialize == rhs.serialize && deserialize == rhs.deserialize && direct_rewrite == rhs.direct_rewrite && rewrite == rhs.rewrite && rewrite_policy == rhs.rewrite_policy && rewrite_optimizer_type == rhs.rewrite_optimizer_type && rewrite_cost == rhs.rewrite_cost && diff --git a/src/duckdb/src/function/cast/cast_function_set.cpp b/src/duckdb/src/function/cast/cast_function_set.cpp index 4fa392db3..665a066db 100644 --- a/src/duckdb/src/function/cast/cast_function_set.cpp +++ b/src/duckdb/src/function/cast/cast_function_set.cpp @@ -39,6 +39,10 @@ CastFunctionSet &CastFunctionSet::Get(ClientContext &context) { return DBConfig::GetConfig(context).GetCastFunctions(); } +const CastFunctionSet &CastFunctionSet::Get(const ClientContext &context) { + return DBConfig::GetConfig(context).GetCastFunctions(); +} + CollationBinding &CollationBinding::Get(ClientContext &context) { return DBConfig::GetConfig(context).GetCollationBinding(); } @@ -58,6 +62,11 @@ BoundCastInfo CastFunctionSet::GetCastFunction(const LogicalType &source, const result.SetStatisticsCallback(CastStatistics::Propagate); return result; } + if (registration_probe) { + // VARIANT selects additional casts from runtime values. + registered_cast_found |= + source.id() == LogicalTypeId::VARIANT || registration_probe->HasRegisteredCast(source, target); + } auto bind_cast = [&](BindCastFunction &bind_function) { BindCastInput input(*this, bind_function.info.get(), get_input.context); input.query_location = get_input.query_location; @@ -98,8 +107,8 @@ struct MapCastNode { int64_t implicit_cast_cost; }; -template -static auto RelaxedTypeMatch(type_map_t &map, const LogicalType &type) -> decltype(map.find(type)) { +template +static auto RelaxedTypeMatch(MAP &map, const LogicalType &type) -> decltype(map.find(type)) { D_ASSERT(map.find(type) == map.end()); // we shouldn't be here switch (type.id()) { case LogicalTypeId::LIST: @@ -137,7 +146,7 @@ static auto RelaxedTypeMatch(type_map_t &map, const LogicalType struct MapCastInfo : public BindCastInfo { public: - const optional_ptr GetEntry(const LogicalType &source, const LogicalType &target) { + optional_ptr GetEntry(const LogicalType &source, const LogicalType &target) const { auto source_type_id_entry = casts.find(source.id()); if (source_type_id_entry == casts.end()) { source_type_id_entry = casts.find(LogicalTypeId::ANY); @@ -184,8 +193,29 @@ struct MapCastInfo : public BindCastInfo { type_id_map_t>>> casts; }; +bool CastFunctionSet::HasRegisteredCast(const LogicalType &source, const LogicalType &target) const { + return map_info && map_info->GetEntry(source, target); +} + +bool CastFunctionSet::CanOverrideDefaultCast(const LogicalType &source, const LogicalType &target) const { + if (!map_info || source == target) { + return false; + } + if (HasRegisteredCast(source, target)) { + return true; + } + CastFunctionSet probe; + probe.registration_probe = this; + GetCastFunctionInput input; + probe.GetCastFunction(source, target, input); + return probe.registered_cast_found; +} + int64_t CastFunctionSet::ImplicitCastCost(optional_ptr context, const LogicalType &source, const LogicalType &target) { + if (registration_probe && registration_probe->HasRegisteredCast(source, target)) { + registered_cast_found = true; + } // check if a cast has been registered if (map_info) { auto entry = map_info->GetEntry(source, target); diff --git a/src/duckdb/src/function/cast/time_casts.cpp b/src/duckdb/src/function/cast/time_casts.cpp index d8a44ff06..e29cdab9e 100644 --- a/src/duckdb/src/function/cast/time_casts.cpp +++ b/src/duckdb/src/function/cast/time_casts.cpp @@ -153,7 +153,8 @@ BoundCastInfo DefaultCasts::TimestampTzCastSwitch(BindCastInput &input, const Lo // timestamp with time zone to timestamp (us) return ReinterpretCast; case LogicalTypeId::TIMESTAMP_NS: - // timestamptz (us) to timestamp (ns) + case LogicalTypeId::TIMESTAMP_TZ_NS: + // timestamptz (us) to timestamp [with time zone] (ns) return BoundCastInfo( &VectorCastHelpers::TryCastErrorLoop); case LogicalTypeId::TIMESTAMP_MS: diff --git a/src/duckdb/src/function/scalar/compressed_materialization_utils.cpp b/src/duckdb/src/function/scalar/compressed_materialization_utils.cpp index 2320279d7..f3d2949d9 100644 --- a/src/duckdb/src/function/scalar/compressed_materialization_utils.cpp +++ b/src/duckdb/src/function/scalar/compressed_materialization_utils.cpp @@ -42,6 +42,17 @@ CMExpressionType CMUtils::GetExpressionType(const BoundFunctionExpression &expre return type; } +optional_ptr CMUtils::GetWrappedInput(const Expression &expression) { + if (expression.GetExpressionClass() != ExpressionClass::BOUND_FUNCTION) { + return nullptr; + } + auto &function = expression.Cast(); + if (function.GetChildren().empty() || GetExpressionType(function) == CMExpressionType::NONE) { + return nullptr; + } + return function.GetChildren()[0].get(); +} + const vector CMUtils::IntegralTypes() { return {LogicalType::UTINYINT, LogicalType::USMALLINT, LogicalType::UINTEGER, LogicalType::UBIGINT}; } diff --git a/src/duckdb/src/function/scalar/struct/struct_pack.cpp b/src/duckdb/src/function/scalar/struct/struct_pack.cpp index 7b611e420..7ebeee975 100644 --- a/src/duckdb/src/function/scalar/struct/struct_pack.cpp +++ b/src/duckdb/src/function/scalar/struct/struct_pack.cpp @@ -1,3 +1,4 @@ +#include "duckdb/parser/expression/function_expression.hpp" #include "duckdb/common/vector/map_vector.hpp" #include "duckdb/common/vector/struct_vector.hpp" #include "duckdb/function/scalar/nested_functions.hpp" @@ -76,6 +77,18 @@ static unique_ptr StructPackStats(ClientContext &context, Functi return struct_stats.ToUnique(); } +template +static unique_ptr StructPackUnbind(FunctionUnbindInput &input) { + auto &expression = input.expression; + vector arguments; + for (idx_t i = 0; i < input.children.size(); i++) { + auto name = IS_STRUCT_PACK ? StructType::GetChildName(expression.GetReturnType(), i) : Identifier(); + arguments.emplace_back(std::move(name), std::move(input.children[i])); + } + return make_uniq(expression.Function().GetDefinition()->GetQualifiedName(), + std::move(arguments)); +} + template static ScalarFunction GetStructPackFunction() { ScalarFunction fun(IS_STRUCT_PACK ? "struct_pack" : "row", {}, @@ -96,6 +109,7 @@ static ScalarFunction GetStructPackFunction() { fun.SetNullHandling(FunctionNullHandling::SPECIAL_HANDLING); fun.SetSerializeCallback(VariableReturnBindData::Serialize); fun.SetDeserializeCallback(VariableReturnBindData::Deserialize); + fun.SetUnbindCallback(StructPackUnbind); return fun; } diff --git a/src/duckdb/src/function/scalar/system/aggregate_export.cpp b/src/duckdb/src/function/scalar/system/aggregate_export.cpp index b6737d24e..c45423190 100644 --- a/src/duckdb/src/function/scalar/system/aggregate_export.cpp +++ b/src/duckdb/src/function/scalar/system/aggregate_export.cpp @@ -20,6 +20,11 @@ #include "duckdb/planner/expression/bound_constant_expression.hpp" #include "duckdb/planner/expression/bound_function_expression.hpp" #include "duckdb/function/aggregate/distributive_functions.hpp" +#include "duckdb/parser/expression/cast_expression.hpp" +#include "duckdb/parser/expression/constant_expression.hpp" +#include "duckdb/parser/expression/function_expression.hpp" +#include "duckdb/common/type_visitor.hpp" +#include "duckdb/parser/expression/type_expression.hpp" namespace duckdb { @@ -771,6 +776,55 @@ void ToAggregateStateFunction(DataChunk &input, ExpressionState &state, Vector & } // namespace +unique_ptr ExportAggregateFunction::StateToSQL(const LogicalType &type, + unique_ptr value) { + auto info = type.GetExtensionInfo(); + if (!type.IsAggregateState() || !info) { + return nullptr; + } + auto name = info->properties.find("function_name"); + auto parameters = info->properties.find("parameters"); + const bool has_function_name = + name != info->properties.end() && !name->second.IsNull() && name->second.type().id() == LogicalTypeId::VARCHAR; + const bool has_parameters = parameters != info->properties.end() && !parameters->second.IsNull() && + parameters->second.type().id() == LogicalTypeId::LIST; + if (!has_function_name || !has_parameters) { + return nullptr; + } + vector types; + map constants; + ParseStateParameters(parameters->second, types, constants); + vector signature; + vector> constant_arguments; + for (idx_t i = 0; i < types.size(); i++) { + if (!TypeExpression::CanRepresent(types[i]) || TypeVisitor::Contains(types[i], [](const LogicalType &child) { + return child.id() == LogicalTypeId::VARCHAR && !StringType::GetCollation(child).empty(); + })) { + return nullptr; + } + signature.emplace_back(types[i].ToString()); + auto entry = constants.find(i); + unique_ptr constant; + try { + constant = ConstantExpression::FromValue(entry == constants.end() ? Value() : entry->second); + } catch (const NotImplementedException &) { + return nullptr; + } + constant_arguments.push_back(make_uniq(LogicalType::VARIANT(), std::move(constant))); + } + vector> arguments; + arguments.push_back(std::move(value)); + arguments.push_back(ConstantExpression::FromValue(name->second)); + arguments.push_back(ConstantExpression::FromValue(Value::LIST(LogicalType::VARCHAR, std::move(signature)))); + arguments.push_back( + make_uniq(QualifiedName("system", "main", "list_value"), std::move(constant_arguments))); + auto orders = info->properties.find("order_bys"); + if (orders != info->properties.end()) { + arguments.push_back(ConstantExpression::FromValue(orders->second)); + } + return make_uniq(QualifiedName("system", "main", "to_aggregate_state"), std::move(arguments)); +} + void ExportAggregateFunction::SetStateExport(BoundAggregateExpression &aggregate, LogicalType state_layout) { auto &bound_function = aggregate.FunctionMutable(); // functions with an explicit export callback use it as the finalize; others use the field-based serialization diff --git a/src/duckdb/src/function/scalar/system/write_log.cpp b/src/duckdb/src/function/scalar/system/write_log.cpp index 225023104..af6a55c31 100644 --- a/src/duckdb/src/function/scalar/system/write_log.cpp +++ b/src/duckdb/src/function/scalar/system/write_log.cpp @@ -1,3 +1,4 @@ +#include "duckdb/parser/expression/function_expression.hpp" #include "duckdb/function/scalar/system_functions.hpp" #include "duckdb/execution/expression_executor.hpp" #include "duckdb/main/client_data.hpp" @@ -148,6 +149,21 @@ void WriteLogFunction(DataChunk &args, ExpressionState &state, Vector &result) { } } +//! write_log captures its named arguments as child aliases rather than as named parameters, so they are reattached +//! as names when the call is rendered +static unique_ptr WriteLogUnbind(FunctionUnbindInput &input) { + vector arguments; + for (idx_t i = 0; i < input.children.size(); i++) { + auto name = i == 0 ? Identifier() : input.expression.GetChildren()[i]->GetAlias(); + if (i > 0 && name.empty()) { + return nullptr; + } + arguments.emplace_back(std::move(name), std::move(input.children[i])); + } + return make_uniq(input.expression.Function().GetDefinition()->GetQualifiedName(), + std::move(arguments)); +} + } // namespace ScalarFunctionSet WriteLogFun::GetFunctions() { @@ -157,6 +173,7 @@ ScalarFunctionSet WriteLogFun::GetFunctions() { nullptr, LogicalType(LogicalTypeId::INVALID), FunctionStability::VOLATILE); fun.GetSignature().AddKwargsParameter("kwargs", LogicalType::ANY); fun.GetProperties().SetRequiresExpressionNames(true); + fun.SetUnbindCallback(WriteLogUnbind); set.AddFunction(std::move(fun)); return set; diff --git a/src/duckdb/src/function/table/system/duckdb_keywords.cpp b/src/duckdb/src/function/table/system/duckdb_keywords.cpp index ea5d55244..704c63884 100644 --- a/src/duckdb/src/function/table/system/duckdb_keywords.cpp +++ b/src/duckdb/src/function/table/system/duckdb_keywords.cpp @@ -56,8 +56,9 @@ void DuckDBKeywordsFunction(ClientContext &context, TableFunctionInput &data_p, } void DuckDBKeywordsFun::RegisterFunction(BuiltinFunctions &set) { - set.AddFunction( - TableFunction("duckdb_keywords", {}, DuckDBKeywordsFunction, DuckDBKeywordsBind, DuckDBKeywordsInit)); + TableFunction duckdb_keywords("duckdb_keywords", {}, DuckDBKeywordsFunction, DuckDBKeywordsBind, + DuckDBKeywordsInit); + set.AddFunction(duckdb_keywords); } } // namespace duckdb diff --git a/src/duckdb/src/function/table/system/duckdb_optimizers.cpp b/src/duckdb/src/function/table/system/duckdb_optimizers.cpp index 719e10514..c0b6b09b3 100644 --- a/src/duckdb/src/function/table/system/duckdb_optimizers.cpp +++ b/src/duckdb/src/function/table/system/duckdb_optimizers.cpp @@ -51,8 +51,9 @@ void DuckDBOptimizersFunction(ClientContext &context, TableFunctionInput &data_p } void DuckDBOptimizersFun::RegisterFunction(BuiltinFunctions &set) { - set.AddFunction( - TableFunction("duckdb_optimizers", {}, DuckDBOptimizersFunction, DuckDBOptimizersBind, DuckDBOptimizersInit)); + TableFunction duckdb_optimizers("duckdb_optimizers", {}, DuckDBOptimizersFunction, DuckDBOptimizersBind, + DuckDBOptimizersInit); + set.AddFunction(duckdb_optimizers); } } // namespace duckdb diff --git a/src/duckdb/src/function/table/system/duckdb_sequences.cpp b/src/duckdb/src/function/table/system/duckdb_sequences.cpp index 089dfd73d..af589fc83 100644 --- a/src/duckdb/src/function/table/system/duckdb_sequences.cpp +++ b/src/duckdb/src/function/table/system/duckdb_sequences.cpp @@ -156,8 +156,9 @@ void DuckDBSequencesFunction(ClientContext &context, TableFunctionInput &data_p, } void DuckDBSequencesFun::RegisterFunction(BuiltinFunctions &set) { - set.AddFunction( - TableFunction("duckdb_sequences", {}, DuckDBSequencesFunction, DuckDBSequencesBind, DuckDBSequencesInit)); + auto function = + TableFunction("duckdb_sequences", {}, DuckDBSequencesFunction, DuckDBSequencesBind, DuckDBSequencesInit); + set.AddFunction(std::move(function)); } } // namespace duckdb diff --git a/src/duckdb/src/function/table/system/duckdb_tables.cpp b/src/duckdb/src/function/table/system/duckdb_tables.cpp index 257b403aa..fbfa53d11 100644 --- a/src/duckdb/src/function/table/system/duckdb_tables.cpp +++ b/src/duckdb/src/function/table/system/duckdb_tables.cpp @@ -175,7 +175,8 @@ void DuckDBTablesFunction(ClientContext &context, TableFunctionInput &data_p, Da } void DuckDBTablesFun::RegisterFunction(BuiltinFunctions &set) { - set.AddFunction(TableFunction("duckdb_tables", {}, DuckDBTablesFunction, DuckDBTablesBind, DuckDBTablesInit)); + auto function = TableFunction("duckdb_tables", {}, DuckDBTablesFunction, DuckDBTablesBind, DuckDBTablesInit); + set.AddFunction(std::move(function)); } } // namespace duckdb diff --git a/src/duckdb/src/function/table/system/pragma_user_agent.cpp b/src/duckdb/src/function/table/system/pragma_user_agent.cpp index 82b04eca2..05705d5c0 100644 --- a/src/duckdb/src/function/table/system/pragma_user_agent.cpp +++ b/src/duckdb/src/function/table/system/pragma_user_agent.cpp @@ -41,8 +41,9 @@ void PragmaUserAgentFunction(ClientContext &context, TableFunctionInput &data_p, } void PragmaUserAgent::RegisterFunction(BuiltinFunctions &set) { - set.AddFunction( - TableFunction("pragma_user_agent", {}, PragmaUserAgentFunction, PragmaUserAgentBind, PragmaUserAgentInit)); + TableFunction pragma_user_agent("pragma_user_agent", {}, PragmaUserAgentFunction, PragmaUserAgentBind, + PragmaUserAgentInit); + set.AddFunction(pragma_user_agent); } } // namespace duckdb diff --git a/src/duckdb/src/function/table/table_scan.cpp b/src/duckdb/src/function/table/table_scan.cpp index 7b47f9b46..0e8131533 100644 --- a/src/duckdb/src/function/table/table_scan.cpp +++ b/src/duckdb/src/function/table/table_scan.cpp @@ -40,6 +40,7 @@ #include "duckdb/storage/table/data_table_info.hpp" #include "duckdb/storage/table/scan_state.hpp" #include "duckdb/planner/expression_iterator.hpp" +#include "duckdb/parser/tableref/basetableref.hpp" #include "duckdb/transaction/duck_transaction_manager.hpp" #include "duckdb/main/profiler/profiling_node.hpp" @@ -1227,6 +1228,53 @@ void SetPartitionsToScan(vector partition_indices, optional_ptr>(partition_indices.begin(), partition_indices.end()); } +static string TableScanToSQLGuard(const LogicalGet &get, bool has_input) { + auto &data = get.bind_data->Cast(); + if (has_input) { + return "input_child"; + } + if (get.ordinality_idx.IsValid()) { + return "ordinality"; + } + if (!get.scan_partition_indices.empty()) { + return "scan_partitions"; + } + if (data.is_create_index) { + return "create_index"; + } + if (data.partitions_to_scan) { + return "bound_scan_partitions"; + } + const bool has_filters = + get.table_filters.HasFilters() || get.table_filters.HasMultiColumnFilters() || get.dynamic_filters; + for (auto &index : get.GetColumnIds()) { + if (index.IsRowNumberColumn() && has_filters) { + return "row_number_with_filters"; + } + if (index.IsVirtualColumn() && !index.IsRowIdColumn() && !index.IsRowNumberColumn()) { + return "virtual_column"; + } + if (!index.IsVirtualColumn()) { + auto &definition = data.table.GetColumn(index.ToLogical()); + if (!index.IsPushdownExtract() && definition.Type() != get.GetColumnType(index)) { + return "column_type"; + } + } + } + return string(); +} + +static TableFunctionToSQLResult TableScanToSQL(ClientContext &, const LogicalGet &get) { + auto guard = TableScanToSQLGuard(get, !get.children.empty()); + if (!guard.empty()) { + return {nullptr, std::move(guard)}; + } + auto table = make_uniq(); + auto entry = get.GetTable(); + table->SetQualifiedName(entry->schema.GetQualifiedName(entry->name)); + return {std::move(table), {}}; +} + TableFunction TableScanFunction::GetFunction() { TableFunction scan_function("seq_scan", {}, TableScanFunc); scan_function.init_local = TableScanInitLocal; @@ -1237,6 +1285,7 @@ TableFunction TableScanFunction::GetFunction() { scan_function.get_metrics = TableScanGetMetrics; scan_function.pushdown_complex_filter = nullptr; scan_function.to_string = TableScanToString; + scan_function.to_sql = TableScanToSQL; scan_function.table_scan_progress = TableScanProgress; scan_function.get_partition_data = TableScanGetPartitionData; scan_function.get_partition_stats = TableScanGetPartitionStats; diff --git a/src/duckdb/src/function/table/version/pragma_version.cpp b/src/duckdb/src/function/table/version/pragma_version.cpp index e9e27542c..f3ad38db6 100644 --- a/src/duckdb/src/function/table/version/pragma_version.cpp +++ b/src/duckdb/src/function/table/version/pragma_version.cpp @@ -1,5 +1,5 @@ #ifndef DUCKDB_PATCH_VERSION -#define DUCKDB_PATCH_VERSION "0-alpha43385" +#define DUCKDB_PATCH_VERSION "0-alpha43586" #endif #ifndef DUCKDB_MINOR_VERSION #define DUCKDB_MINOR_VERSION 0 @@ -8,10 +8,10 @@ #define DUCKDB_MAJOR_VERSION 2 #endif #ifndef DUCKDB_VERSION -#define DUCKDB_VERSION "v2.0.0-alpha43385" +#define DUCKDB_VERSION "v2.0.0-alpha43586" #endif #ifndef DUCKDB_SOURCE_ID -#define DUCKDB_SOURCE_ID "ca15f79c32" +#define DUCKDB_SOURCE_ID "09eb7f7004" #endif #include "duckdb/function/table/system_functions.hpp" #include "duckdb/main/database.hpp" diff --git a/src/duckdb/src/function/table_function.cpp b/src/duckdb/src/function/table_function.cpp index b4955fe82..28801bbd6 100644 --- a/src/duckdb/src/function/table_function.cpp +++ b/src/duckdb/src/function/table_function.cpp @@ -72,7 +72,7 @@ bool TableFunction::operator==(const TableFunction &rhs) const { in_out_function_final == rhs.in_out_function_final && statistics == rhs.statistics && dependency == rhs.dependency && cardinality == rhs.cardinality && pushdown_complex_filter == rhs.pushdown_complex_filter && pushdown_expression == rhs.pushdown_expression && - to_string == rhs.to_string && table_scan_progress == rhs.table_scan_progress && + to_string == rhs.to_string && to_sql == rhs.to_sql && table_scan_progress == rhs.table_scan_progress && get_partition_data == rhs.get_partition_data && get_bind_info == rhs.get_bind_info && projection_expression_pushdown == rhs.projection_expression_pushdown && get_multi_file_reader == rhs.get_multi_file_reader && supports_pushdown_type == rhs.supports_pushdown_type && diff --git a/src/duckdb/src/include/duckdb/common/enum_util.hpp b/src/duckdb/src/include/duckdb/common/enum_util.hpp index 124c0d8e6..19a737091 100644 --- a/src/duckdb/src/include/duckdb/common/enum_util.hpp +++ b/src/duckdb/src/include/duckdb/common/enum_util.hpp @@ -338,6 +338,12 @@ enum class LogicalOperatorRepeatability : uint8_t; enum class LogicalOperatorType : uint8_t; +enum class LogicalPlanVerificationIssueCode : int32_t; + +enum class LogicalPlanVerificationPathComponentType : int32_t; + +enum class LogicalPlanVerificationPhase : int32_t; + enum class LogicalTypeId : uint8_t; enum class LookupResultType : uint8_t; @@ -1114,6 +1120,15 @@ const char* EnumUtil::ToChars(LogicalOperatorRepea template<> const char* EnumUtil::ToChars(LogicalOperatorType value); +template<> +const char* EnumUtil::ToChars(LogicalPlanVerificationIssueCode value); + +template<> +const char* EnumUtil::ToChars(LogicalPlanVerificationPathComponentType value); + +template<> +const char* EnumUtil::ToChars(LogicalPlanVerificationPhase value); + template<> const char* EnumUtil::ToChars(LogicalTypeId value); @@ -2048,6 +2063,15 @@ LogicalOperatorRepeatability EnumUtil::FromString( template<> LogicalOperatorType EnumUtil::FromString(const char *value); +template<> +LogicalPlanVerificationIssueCode EnumUtil::FromString(const char *value); + +template<> +LogicalPlanVerificationPathComponentType EnumUtil::FromString(const char *value); + +template<> +LogicalPlanVerificationPhase EnumUtil::FromString(const char *value); + template<> LogicalTypeId EnumUtil::FromString(const char *value); diff --git a/src/duckdb/src/include/duckdb/common/enums/database_modification_type.hpp b/src/duckdb/src/include/duckdb/common/enums/database_modification_type.hpp index faafd4dea..040f5ab62 100644 --- a/src/duckdb/src/include/duckdb/common/enums/database_modification_type.hpp +++ b/src/duckdb/src/include/duckdb/common/enums/database_modification_type.hpp @@ -36,6 +36,10 @@ struct DatabaseModificationType { return *this; } + bool Contains(DatabaseModificationType other) const { + return (value & other.value) == other.value; + } + bool InsertData() const { return value & INSERT_DATA; } diff --git a/src/duckdb/src/include/duckdb/common/enums/debug_statement_verification.hpp b/src/duckdb/src/include/duckdb/common/enums/debug_statement_verification.hpp index 0b7617f9d..c79345cac 100644 --- a/src/duckdb/src/include/duckdb/common/enums/debug_statement_verification.hpp +++ b/src/duckdb/src/include/duckdb/common/enums/debug_statement_verification.hpp @@ -18,7 +18,9 @@ enum class DebugStatementVerification : uint8_t { REPARSE_STATEMENT, SERIALIZE_STATEMENT, PREPARED_STATEMENT, - EXPLAIN_STATEMENT + EXPLAIN_STATEMENT, + EXPLAIN_SQL, + EXPLAIN_SQL_STRICT }; } // namespace duckdb diff --git a/src/duckdb/src/include/duckdb/common/extra_operator_info.hpp b/src/duckdb/src/include/duckdb/common/extra_operator_info.hpp index cb0688c0b..ae452cc44 100644 --- a/src/duckdb/src/include/duckdb/common/extra_operator_info.hpp +++ b/src/duckdb/src/include/duckdb/common/extra_operator_info.hpp @@ -21,6 +21,9 @@ namespace duckdb { class ExtraOperatorInfo { public: + //! Table index of retained file-filter bindings, reserved above any binder-generated index + static constexpr idx_t FILE_FILTER_TABLE_INDEX = DConstants::INVALID_INDEX - 1; + ExtraOperatorInfo() : file_filters(""), sample_options(nullptr) { } ExtraOperatorInfo(ExtraOperatorInfo &&extra_info) noexcept = default; diff --git a/src/duckdb/src/include/duckdb/common/string_util.hpp b/src/duckdb/src/include/duckdb/common/string_util.hpp index bee930182..b6e889cb2 100644 --- a/src/duckdb/src/include/duckdb/common/string_util.hpp +++ b/src/duckdb/src/include/duckdb/common/string_util.hpp @@ -152,6 +152,21 @@ class StringUtil { DUCKDB_API static string Join(const vector &input, const string &separator); DUCKDB_API static string Join(const set &input, const string &separator); + //! Join container elements transformed to strings using the given separator + template + static string Join(const CONTAINER &input, const string &separator, const FUNC &f) { + string result; + bool first = true; + for (const auto &entry : input) { + if (!first) { + result += separator; + } + result += f(entry); + first = false; + } + return result; + } + //! Encode special URL characters in a string DUCKDB_API static string URLEncode(const string &str, bool encode_slash = true); DUCKDB_API static idx_t URLEncodeSize(const char *input, idx_t input_size, bool encode_slash = true); diff --git a/src/duckdb/src/include/duckdb/common/types/type_manager.hpp b/src/duckdb/src/include/duckdb/common/types/type_manager.hpp index 2347bba75..e38def071 100644 --- a/src/duckdb/src/include/duckdb/common/types/type_manager.hpp +++ b/src/duckdb/src/include/duckdb/common/types/type_manager.hpp @@ -18,6 +18,7 @@ class TypeManager { //! Get the CastFunctionSet from the TypeManager CastFunctionSet &GetCastFunctions(); + const CastFunctionSet &GetCastFunctions() const; //! Try to parse and bind a logical type from a string. Throws an exception if the type could not be parsed. LogicalType ParseLogicalType(const string &type_str, ClientContext &context) const; diff --git a/src/duckdb/src/include/duckdb/common/types/vector_buffer.hpp b/src/duckdb/src/include/duckdb/common/types/vector_buffer.hpp index b974841d3..b10b9f667 100644 --- a/src/duckdb/src/include/duckdb/common/types/vector_buffer.hpp +++ b/src/duckdb/src/include/duckdb/common/types/vector_buffer.hpp @@ -26,6 +26,8 @@ class VectorBuffer; class Vector; struct ValidityMask; struct SelCache; +class Serializer; +class Deserializer; enum class VectorBufferType : uint8_t { STANDARD_BUFFER, // VectorType::FLAT/CONSTANT - Fixed-Size Type - Holds a single array of data @@ -190,6 +192,11 @@ class VectorBuffer : public enable_shared_from_this { const SelectionVector &sel, idx_t count); //! Create a UnifiedVectorFormat from the buffer's data virtual void ToUnifiedFormat(UnifiedVectorFormat &format) const; + //! Serialize the buffer as-is (including the vector_type), returns false if the vector must be serialized flat + virtual bool TrySerialize(Serializer &serializer, const LogicalType &type, bool compressed_serialization) const; + //! Deserialize a buffer written by TrySerialize + static buffer_ptr Deserialize(Deserializer &deserializer, VectorType vector_type, + const LogicalType &type, idx_t count); protected: //! Slice a constant vector with a specific count diff --git a/src/duckdb/src/include/duckdb/common/vector/dictionary_vector.hpp b/src/duckdb/src/include/duckdb/common/vector/dictionary_vector.hpp index 5afa22ef3..0a181002f 100644 --- a/src/duckdb/src/include/duckdb/common/vector/dictionary_vector.hpp +++ b/src/duckdb/src/include/duckdb/common/vector/dictionary_vector.hpp @@ -87,6 +87,8 @@ class DictionaryBuffer : public VectorBuffer { Value GetValue(const LogicalType &type, idx_t index) const override; buffer_ptr SliceWithCache(SelCache &cache, const LogicalType &type, const SelectionVector &sel, idx_t count) override; + bool TrySerialize(Serializer &serializer, const LogicalType &type, bool compressed_serialization) const override; + static buffer_ptr Deserialize(Deserializer &deserializer, const LogicalType &type, idx_t count); protected: buffer_ptr SliceInternal(const LogicalType &type, idx_t offset, idx_t end) override; diff --git a/src/duckdb/src/include/duckdb/common/vector/sequence_vector.hpp b/src/duckdb/src/include/duckdb/common/vector/sequence_vector.hpp index dc77d5e9b..39876ceb6 100644 --- a/src/duckdb/src/include/duckdb/common/vector/sequence_vector.hpp +++ b/src/duckdb/src/include/duckdb/common/vector/sequence_vector.hpp @@ -27,6 +27,8 @@ class SequenceBuffer : public VectorBuffer { idx_t GetAllocationSize() const override; string ToString(const LogicalType &type, idx_t count) const override; Value GetValue(const LogicalType &type, idx_t index) const override; + bool TrySerialize(Serializer &serializer, const LogicalType &type, bool compressed_serialization) const override; + static buffer_ptr Deserialize(Deserializer &deserializer, const LogicalType &type, idx_t count); protected: buffer_ptr FlattenSliceInternal(const LogicalType &type, const SelectionVector &sel, diff --git a/src/duckdb/src/include/duckdb/common/vector/shredded_vector.hpp b/src/duckdb/src/include/duckdb/common/vector/shredded_vector.hpp index 852b48afe..0a6fd12b8 100644 --- a/src/duckdb/src/include/duckdb/common/vector/shredded_vector.hpp +++ b/src/duckdb/src/include/duckdb/common/vector/shredded_vector.hpp @@ -31,6 +31,8 @@ class ShreddedVectorBuffer : public VectorBuffer { string ToString(const LogicalType &type, idx_t count) const override; Value GetValue(const LogicalType &type, idx_t index) const override; void SetVectorType(VectorType new_vector_type) override; + bool TrySerialize(Serializer &serializer, const LogicalType &type, bool compressed_serialization) const override; + static buffer_ptr Deserialize(Deserializer &deserializer, const LogicalType &type, idx_t count); protected: buffer_ptr SliceInternal(const LogicalType &type, idx_t offset, idx_t end) override; diff --git a/src/duckdb/src/include/duckdb/execution/mark_join_row_comparison.hpp b/src/duckdb/src/include/duckdb/execution/mark_join_row_comparison.hpp index 11cd76c06..52ac55a23 100644 --- a/src/duckdb/src/include/duckdb/execution/mark_join_row_comparison.hpp +++ b/src/duckdb/src/include/duckdb/execution/mark_join_row_comparison.hpp @@ -13,22 +13,10 @@ namespace duckdb { -struct JoinCondition; - struct MarkJoinRowComparison { - explicit MarkJoinRowComparison(const DataChunk &left); - static void Compare(const Vector &left, const Vector &right, ExpressionType comparison_type, Vector &result); - void CompareConjunction(DataChunk &left, idx_t left_row, DataChunk &right, const vector &conditions, - Vector &result); - static void Perform(DataChunk &left, DataChunk &right, bool found_match[], const vector &conditions, - optional_ptr found_unknown); static void CompareEquality(const Vector &left, idx_t left_row, idx_t left_count, const Vector &right, idx_t right_count, bool row_is_false[], bool row_is_unknown[]); - -private: - DataChunk left_reference; - Vector comparison; }; } // namespace duckdb diff --git a/src/duckdb/src/include/duckdb/function/aggregate_function.hpp b/src/duckdb/src/include/duckdb/function/aggregate_function.hpp index b02e818c5..7d05e37e3 100644 --- a/src/duckdb/src/include/duckdb/function/aggregate_function.hpp +++ b/src/duckdb/src/include/duckdb/function/aggregate_function.hpp @@ -77,6 +77,19 @@ class BindAggregateFunctionInput : public BindFunctionInput { BoundAggregateFunction &bound_function; }; +class FunctionExpression; + +struct AggregateFunctionUnbindInput { + AggregateFunctionUnbindInput(const BoundAggregateExpression &expression_p, + vector> children_p); + ~AggregateFunctionUnbindInput(); + + const BoundAggregateExpression &expression; + vector> children; +}; + +typedef unique_ptr (*aggregate_function_unbind_t)(AggregateFunctionUnbindInput &input); + //! The type used for sizing hashed aggregate function states typedef idx_t (*aggregate_size_t)(AggregateStateInput &input); //! The type used for initializing hashed aggregate function states (batched: initializes `count` states) @@ -194,6 +207,9 @@ class AggregateFunctionCallbacks { bool HasBindCallback() const { return bind != nullptr; } bind_aggregate_function_t GetBindCallback() const { return bind; } void SetBindCallback(bind_aggregate_function_t callback) { bind = callback; } + bool HasUnbindCallback() const { return unbind != nullptr; } + aggregate_function_unbind_t GetUnbindCallback() const { return unbind; } + void SetUnbindCallback(aggregate_function_unbind_t callback) { unbind = callback; } bool HasStateInitCallback() const { return initialize != nullptr; } aggregate_initialize_t GetStateInitCallback() const { return initialize; } @@ -296,6 +312,7 @@ class AggregateFunctionCallbacks { //! The bind function (may be null) bind_aggregate_function_t bind = nullptr; + aggregate_function_unbind_t unbind = nullptr; //! The destructor method (may be null) aggregate_destructor_t destructor = nullptr; @@ -393,6 +410,9 @@ class BaseAggregateFunction { auto HasBindCallback() const -> bool { return callbacks.bind != nullptr; } auto GetBindCallback() const -> bind_aggregate_function_t { return callbacks.bind; } auto SetBindCallback(bind_aggregate_function_t callback) -> void { callbacks.bind = callback; } + auto HasUnbindCallback() const -> bool { return callbacks.unbind != nullptr; } + auto GetUnbindCallback() const -> aggregate_function_unbind_t { return callbacks.unbind; } + auto SetUnbindCallback(aggregate_function_unbind_t callback) -> void { callbacks.unbind = callback; } auto HasStateInitCallback() const -> bool { return callbacks.initialize != nullptr; } auto GetStateInitCallback() const -> aggregate_initialize_t { return callbacks.initialize; } diff --git a/src/duckdb/src/include/duckdb/function/cast/cast_function_set.hpp b/src/duckdb/src/include/duckdb/function/cast/cast_function_set.hpp index 97e539441..c1dcc8ff9 100644 --- a/src/duckdb/src/include/duckdb/function/cast/cast_function_set.hpp +++ b/src/duckdb/src/include/duckdb/function/cast/cast_function_set.hpp @@ -45,6 +45,7 @@ class CastFunctionSet { public: DUCKDB_API static CastFunctionSet &Get(ClientContext &context); + DUCKDB_API static const CastFunctionSet &Get(const ClientContext &context); DUCKDB_API static CastFunctionSet &Get(DatabaseInstance &db); //! Returns a cast function (from source -> target) @@ -64,6 +65,9 @@ class CastFunctionSet { int64_t implicit_cast_cost = -1); DUCKDB_API void RegisterCastFunction(const LogicalType &source, const LogicalType &target, bind_cast_function_t bind, int64_t implicit_cast_cost = -1); + //! Whether a registered cast can change this default conversion, including nested casts. + //! Data-dependent VARIANT conversions are treated conservatively. + DUCKDB_API bool CanOverrideDefaultCast(const LogicalType &source, const LogicalType &target) const; //! Register a combine rule for LogicalType::TryGetMaxLogicalType, consulted before previously registered rules //! and the built-in rules @@ -85,9 +89,13 @@ class CastFunctionSet { vector combine_rules; //! If any custom cast functions have been defined using RegisterCastFunction, this holds the map optional_ptr map_info; + //! A default-only binding probe observes registrations without invoking them. + optional_ptr registration_probe; + bool registered_cast_found = false; private: void RegisterCastFunction(const LogicalType &source, const LogicalType &target, MapCastNode node); + bool HasRegisteredCast(const LogicalType &source, const LogicalType &target) const; }; } // namespace duckdb diff --git a/src/duckdb/src/include/duckdb/function/scalar/compressed_materialization_utils.hpp b/src/duckdb/src/include/duckdb/function/scalar/compressed_materialization_utils.hpp index cc87c047f..ff8fb362b 100644 --- a/src/duckdb/src/include/duckdb/function/scalar/compressed_materialization_utils.hpp +++ b/src/duckdb/src/include/duckdb/function/scalar/compressed_materialization_utils.hpp @@ -23,6 +23,8 @@ struct CMUtils { static unique_ptr Bind(BindScalarFunctionInput &input); static CMExpressionType GetExpressionType(const BoundFunctionExpression &expression); + //! The input wrapped by a compressed materialization expression, or nullptr for any other expression + static optional_ptr GetWrappedInput(const Expression &expression); static void MarkCast(BoundFunctionExpression &expression); private: diff --git a/src/duckdb/src/include/duckdb/function/table_function.hpp b/src/duckdb/src/include/duckdb/function/table_function.hpp index ccfffd14d..fe85651e1 100644 --- a/src/duckdb/src/include/duckdb/function/table_function.hpp +++ b/src/duckdb/src/include/duckdb/function/table_function.hpp @@ -52,6 +52,7 @@ enum class OrderByStatistics : uint8_t; struct RowGroupOrderOptions; class LogicalOperator; class Binder; +class QueryNode; struct TableFunctionInfo { DUCKDB_API virtual ~TableFunctionInfo(); @@ -426,6 +427,15 @@ typedef unique_ptr (*table_function_combine_schema_t)(ClientContex vector &names); typedef InsertionOrderPreservingMap (*table_function_to_string_t)(TableFunctionToStringInput &input); +struct TableFunctionToSQLResult { + unique_ptr source; + string unsupported_reason; +}; + +//! Reconstruct the unprojected source; the exporter applies scan projections, predicates and ordinality. +//! Return an owned table reference without executing the source, or a reason why reconstruction is unsupported. +typedef TableFunctionToSQLResult (*table_function_to_sql_t)(ClientContext &context, const LogicalGet &get); + typedef void (*table_function_serialize_t)(Serializer &serializer, const optional_ptr bind_data, const TableFunction &function); typedef unique_ptr (*table_function_deserialize_t)(Deserializer &deserializer, TableFunction &function); @@ -557,6 +567,9 @@ class TableFunction : public SimpleNamedParameterFunction { // NOLINT: work-arou table_function_schedule_io_t schedule_io; //! (Optional) function for rendering the operator to a string in explain/profiling output (invoked pre-execution) table_function_to_string_t to_string; + //! (Optional) reconstruct the source's SQL-visible state without retaining native objects. + //! Must not execute the source or perform effects during export. + table_function_to_sql_t to_sql = nullptr; //! (Optional) return how much of the table we have scanned up to this point (% of the data) table_function_progress_t table_scan_progress; //! (Optional) returns the partition info of the current scan operator diff --git a/src/duckdb/src/include/duckdb/main/config.hpp b/src/duckdb/src/include/duckdb/main/config.hpp index 6f79a7031..6e0f8852f 100644 --- a/src/duckdb/src/include/duckdb/main/config.hpp +++ b/src/duckdb/src/include/duckdb/main/config.hpp @@ -289,6 +289,7 @@ struct DBConfig { bool operator!=(const DBConfig &other); DUCKDB_API CastFunctionSet &GetCastFunctions(); + DUCKDB_API const CastFunctionSet &GetCastFunctions() const; DUCKDB_API TypeManager &GetTypeManager(); DUCKDB_API CollationBinding &GetCollationBinding(); DUCKDB_API IndexTypeSet &GetIndexTypes(); diff --git a/src/duckdb/src/include/duckdb/main/prepared_statement_data.hpp b/src/duckdb/src/include/duckdb/main/prepared_statement_data.hpp index 4a1e10baa..1802ba3dd 100644 --- a/src/duckdb/src/include/duckdb/main/prepared_statement_data.hpp +++ b/src/duckdb/src/include/duckdb/main/prepared_statement_data.hpp @@ -62,4 +62,7 @@ class PreparedStatementData { DUCKDB_API bool TryGetType(const Identifier &identifier, LogicalType &result); }; +DUCKDB_API bool CheckCatalogIdentity(ClientContext &context, const Identifier &catalog_name, + StatementProperties::CatalogIdentity catalog_identity); + } // namespace duckdb diff --git a/src/duckdb/src/include/duckdb/optimizer/type_pushdown.hpp b/src/duckdb/src/include/duckdb/optimizer/type_pushdown.hpp index 4a37be3d6..44a21605f 100644 --- a/src/duckdb/src/include/duckdb/optimizer/type_pushdown.hpp +++ b/src/duckdb/src/include/duckdb/optimizer/type_pushdown.hpp @@ -125,6 +125,7 @@ unique_ptr PushdownOptimize(ClientContext &context, unique_ptr< } TableFunctionProjectionExpressionInput input {analysis.get, *expr, column_index}; if (analysis.get.function.projection_expression_pushdown(context, input)) { + analysis.get.has_pushed_projection = true; analysis.get.returned_types[analysis.StorageIndex(column_index)] = expr->GetReturnType(); if (!analysis.get.types.empty()) { analysis.get.ResolveOperatorTypes(); diff --git a/src/duckdb/src/include/duckdb/parser/expression/constant_expression.hpp b/src/duckdb/src/include/duckdb/parser/expression/constant_expression.hpp index e8b60ded4..ec754eb8b 100644 --- a/src/duckdb/src/include/duckdb/parser/expression/constant_expression.hpp +++ b/src/duckdb/src/include/duckdb/parser/expression/constant_expression.hpp @@ -42,6 +42,8 @@ class ConstantExpression : public ParsedExpression { //! Builds the parsed expression for a value: a literal when it re-binds to the same value, the constructor //! call the parser produces for a nested value, or a cast otherwise DUCKDB_API static unique_ptr FromValue(const Value &value); + //! Whether SQL needs a value constructor to preserve the type's metadata + DUCKDB_API static bool RequiresTypeWitness(const LogicalType &type); public: string ToString() const override; diff --git a/src/duckdb/src/include/duckdb/parser/expression/type_expression.hpp b/src/duckdb/src/include/duckdb/parser/expression/type_expression.hpp index 0e398c547..ffe595d13 100644 --- a/src/duckdb/src/include/duckdb/parser/expression/type_expression.hpp +++ b/src/duckdb/src/include/duckdb/parser/expression/type_expression.hpp @@ -29,6 +29,10 @@ class TypeExpression : public ParsedExpression { //! parameterised built-in is written out with its parameters, and a type that carries an alias (a //! user-defined type) is named by that alias. Throws for types that have no SQL spelling. DUCKDB_API static unique_ptr FromLogicalType(const LogicalType &type); + //! Whether the type ID has a SQL spelling + DUCKDB_API static bool IsSQLType(LogicalTypeId id); + //! Whether the complete type can be written as a SQL type expression + DUCKDB_API static bool CanRepresent(const LogicalType &type); const QualifiedName &GetQualifiedName() const { return qualified_name; diff --git a/src/duckdb/src/include/duckdb/parser/statement/explain_statement.hpp b/src/duckdb/src/include/duckdb/parser/statement/explain_statement.hpp index fcc9423de..a9a8a1d44 100644 --- a/src/duckdb/src/include/duckdb/parser/statement/explain_statement.hpp +++ b/src/duckdb/src/include/duckdb/parser/statement/explain_statement.hpp @@ -13,7 +13,7 @@ namespace duckdb { -enum class ExplainType : uint8_t { EXPLAIN_STANDARD, EXPLAIN_ANALYZE }; +enum class ExplainType : uint8_t { EXPLAIN_STANDARD, EXPLAIN_ANALYZE, EXPLAIN_SQL }; class ExplainStatement : public SQLStatement { public: @@ -25,6 +25,8 @@ class ExplainStatement : public SQLStatement { unique_ptr stmt; ExplainType explain_type; + //! Internal verification may request an empty result for unsupported SQL export. + bool allow_unsupported_sql = false; ProfilerPrintFormat format = ProfilerPrintFormat::Default(); protected: diff --git a/src/duckdb/src/include/duckdb/planner/bound_expression_sql_exporter.hpp b/src/duckdb/src/include/duckdb/planner/bound_expression_sql_exporter.hpp index 19cacf637..2999a80fc 100644 --- a/src/duckdb/src/include/duckdb/planner/bound_expression_sql_exporter.hpp +++ b/src/duckdb/src/include/duckdb/planner/bound_expression_sql_exporter.hpp @@ -12,6 +12,7 @@ #include "duckdb/common/optional.hpp" #include "duckdb/common/types.hpp" #include "duckdb/parser/parsed_expression.hpp" +#include "duckdb/parser/expression/function_expression.hpp" #include "duckdb/planner/column_binding.hpp" #include "duckdb/planner/logical_plan_verification_result.hpp" @@ -20,10 +21,17 @@ namespace duckdb { class Expression; +class ClientContext; +class BoundWindowExpression; +class BoundAggregateExpression; +class BoundUnnestExpression; struct ResolvedSQLColumnReference { vector names; + //! Type represented by the exported SQL expression. LogicalType type; + //! Optimizer-selected type accepted from the bound plan, when different. + optional optimizer_type; }; using BoundExpressionSQLBindingResolver = @@ -31,6 +39,10 @@ using BoundExpressionSQLBindingResolver = struct BoundExpressionSQLExportContext { BoundExpressionSQLBindingResolver resolve_binding; + //! The context that will bind the exported expression, when known + optional_ptr client_context; + //! Internal pruning expressions may be discarded only while exporting their complete logical plan. + bool discard_optimizer_metadata = false; }; //! Reconstructs SQL from bound/optimized logical expressions before physical planning lowers their structure. @@ -43,6 +55,19 @@ class BoundExpressionSQLExporter { DUCKDB_API static LogicalPlanVerificationResult> ExportAtPath(const Expression &expression, const BoundExpressionSQLExportContext &context, const LogicalPlanVerificationPath &path); + + //! Only for root expressions placed by their owning logical relation. + DUCKDB_API static LogicalPlanVerificationResult> + ExportAggregateCallAtPath(const BoundAggregateExpression &expression, + const BoundExpressionSQLExportContext &context, const LogicalPlanVerificationPath &path); + + DUCKDB_API static LogicalPlanVerificationResult> + ExportWindowAtPath(const BoundWindowExpression &expression, const BoundExpressionSQLExportContext &context, + const LogicalPlanVerificationPath &path); + + DUCKDB_API static LogicalPlanVerificationResult> + ExportUnnestAtPath(const BoundUnnestExpression &expression, const BoundExpressionSQLExportContext &context, + const LogicalPlanVerificationPath &path); }; } // namespace duckdb diff --git a/src/duckdb/src/include/duckdb/planner/expression/bound_window_expression.hpp b/src/duckdb/src/include/duckdb/planner/expression/bound_window_expression.hpp index 9de76e9c2..2757cf2fd 100644 --- a/src/duckdb/src/include/duckdb/planner/expression/bound_window_expression.hpp +++ b/src/duckdb/src/include/duckdb/planner/expression/bound_window_expression.hpp @@ -11,6 +11,7 @@ #include "duckdb/parser/expression/window_expression.hpp" #include "duckdb/function/function.hpp" #include "duckdb/planner/expression.hpp" +#include "duckdb/planner/expression/window_range_info.hpp" #include "duckdb/parser/parsed_expression.hpp" #include "duckdb/planner/bound_result_modifier.hpp" @@ -123,6 +124,24 @@ class BoundWindowExpression : public Expression { const LogicalType &SQLRangeOrderType() const { return sql_range_order_type; } + const unique_ptr &SQLRangeStartBoundary() const { + return sql_range_start_boundary; + } + unique_ptr &SQLRangeStartBoundaryMutable() { + return sql_range_start_boundary; + } + const unique_ptr &SQLRangeEndBoundary() const { + return sql_range_end_boundary; + } + unique_ptr &SQLRangeEndBoundaryMutable() { + return sql_range_end_boundary; + } + const vector &SQLRangeOrderCasts() const { + return sql_range_order_casts; + } + vector &SQLRangeOrderCastsMutable() { + return sql_range_order_casts; + } const vector &ArgOrders() const { return arg_orders; } @@ -206,6 +225,10 @@ class BoundWindowExpression : public Expression { unique_ptr sql_range_end; LogicalType sql_range_order_type = LogicalType::INVALID; + unique_ptr sql_range_start_boundary; + unique_ptr sql_range_end_boundary; + vector sql_range_order_casts; + //! The set of argument ordering clauses //! These are distinct from the frame ordering clauses e.g., the "x" in //! FIRST_VALUE(a ORDER BY x) OVER (PARTITION BY p ORDER BY s) diff --git a/src/duckdb/src/include/duckdb/planner/expression/window_range_info.hpp b/src/duckdb/src/include/duckdb/planner/expression/window_range_info.hpp new file mode 100644 index 000000000..e740f0ac0 --- /dev/null +++ b/src/duckdb/src/include/duckdb/planner/expression/window_range_info.hpp @@ -0,0 +1,53 @@ +//===----------------------------------------------------------------------===// +// DuckDB +// +// duckdb/planner/expression/window_range_info.hpp +// +// +//===----------------------------------------------------------------------===// + +#pragma once + +#include "duckdb/common/types.hpp" +#include "duckdb/common/optional_ptr.hpp" +#include "duckdb/parser/qualified_name.hpp" +#include "duckdb/parser/expression/window_expression.hpp" + +namespace duckdb { +class Expression; + +//! A coercion inserted while binding a RANGE boundary, without executable bind data. +struct WindowRangeCast { + LogicalType source_type; + LogicalType target_type; + bool try_cast; + bool default_cast; + + void Serialize(Serializer &serializer) const; + static WindowRangeCast Deserialize(Deserializer &deserializer); + + static bool Capture(const Expression &expression, optional_ptr input, + vector &casts); + static optional_ptr Match(const Expression &expression, const vector &casts); +}; + +//! Identifies the binder-produced call and its live operands; it owns no column bindings or expressions. +struct WindowRangeBoundary { + WindowBoundary boundary = WindowBoundary::INVALID; + OrderType direction = OrderType::INVALID; + QualifiedName function_name; + vector arguments; + LogicalType return_type; + LogicalType order_type; + LogicalType offset_type; + vector order_casts; + vector result_casts; + + void Serialize(Serializer &serializer) const; + static unique_ptr Deserialize(Deserializer &deserializer); + + static unique_ptr Capture(const Expression &expression, optional_ptr order, + optional_ptr offset); + optional_ptr Match(const Expression &expression, const Expression &order) const; +}; +} // namespace duckdb diff --git a/src/duckdb/src/include/duckdb/planner/logical_operator.hpp b/src/duckdb/src/include/duckdb/planner/logical_operator.hpp index 239b08425..a65bdc553 100644 --- a/src/duckdb/src/include/duckdb/planner/logical_operator.hpp +++ b/src/duckdb/src/include/duckdb/planner/logical_operator.hpp @@ -23,6 +23,12 @@ namespace duckdb { class LogicalPlanVerifier; +struct LogicalPlanSQLExportRelation; +struct LogicalPlanVerificationPath; +template +class LogicalPlanVerificationResult; +class LogicalPlanSQLExportContext; +using LogicalPlanSQLExportResult = LogicalPlanVerificationResult; //! LogicalOperator is the base class of the logical operators present in the //! logical query tree @@ -47,6 +53,8 @@ class LogicalOperator { bool has_estimated_cardinality; public: + virtual LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path); virtual vector GetColumnBindings(); virtual TableIndex GetRootIndex(); static string ColumnBindingsToString(const vector &bindings); diff --git a/src/duckdb/src/include/duckdb/planner/logical_plan_sql_export_context.hpp b/src/duckdb/src/include/duckdb/planner/logical_plan_sql_export_context.hpp new file mode 100644 index 000000000..5483d476e --- /dev/null +++ b/src/duckdb/src/include/duckdb/planner/logical_plan_sql_export_context.hpp @@ -0,0 +1,83 @@ +#pragma once + +#include "duckdb/planner/logical_plan_sql_exporter.hpp" +#include "duckdb/parser/query_node/select_node.hpp" + +namespace duckdb { +class Expression; +class LogicalOperator; +class LogicalComparisonJoin; +struct LogicalExtensionOperator; +class ClientContext; + +struct LogicalPlanSQLExportedChild { + LogicalPlanSQLExportRelation relation; + Identifier relation_alias; +}; + +struct LogicalPlanSQLExportSource { + reference op; + Identifier name; + LogicalPlanSQLExportRelation relation; +}; + +class LogicalPlanSQLExportContext { +public: + explicit LogicalPlanSQLExportContext(ClientContext &context_p); + LogicalPlanSQLExportResult Export(LogicalOperator &op, const LogicalPlanVerificationPath &path); + +public: + ClientContext &GetClientContext() const { + return context; + } + + Identifier NextRelationAlias(const Identifier &preferred = Identifier()); + LogicalPlanVerificationResult ExportChild(LogicalOperator &child, + const LogicalPlanVerificationPath &path); + LogicalPlanVerificationResult + ExportChild(LogicalOperator &child, const LogicalPlanVerificationPath &path, + const vector &sources); + LogicalPlanVerificationResult> + ExportExpression(const LogicalOperator &op, const vector> &expressions, + idx_t expression_ordinal, const BoundExpressionSQLExportContext &expression_context, + const LogicalPlanVerificationPath &path); + unique_ptr CreateNamedSource(const Identifier &name, const vector &fields, + bool recurring = false); + + LogicalPlanVerificationResult + ExportNamedProducer(LogicalOperator &op, const LogicalPlanVerificationPath &path, const Identifier &name); + + LogicalPlanSQLExportResult ExportRow(unique_ptr select, unique_ptr value, + vector fields); + unique_ptr ForwardFields(const LogicalPlanSQLExportedChild &child, + const vector &fields, + optional_ptr plain = nullptr); + + //! Operators from the root down to the operator currently being exported + const vector> &Ancestors() const { + return ancestors; + } + //! Make a CTE relation referenceable while its consumers are exported + void PushNamedRelation(TableIndex index, const Identifier &name, bool is_recurring); + //! Remove the innermost named relation and return how often it was referenced + idx_t PopNamedRelation(); + //! Reference the innermost named relation with this index; empty when it is out of scope + optional ReferenceNamedRelation(TableIndex index, bool is_recurring); + +private: + struct SourceScope; + optional_ptr source_scope; + struct NamedRelation { + TableIndex index; + Identifier name; + bool is_recurring; + idx_t references; + }; + ClientContext &context; + idx_t next_relation_ordinal = 0; + identifier_set_t relation_aliases; + vector> ancestors; + vector named_relations; +}; + +} // namespace duckdb diff --git a/src/duckdb/src/include/duckdb/planner/logical_plan_sql_exporter.hpp b/src/duckdb/src/include/duckdb/planner/logical_plan_sql_exporter.hpp new file mode 100644 index 000000000..8e42fa673 --- /dev/null +++ b/src/duckdb/src/include/duckdb/planner/logical_plan_sql_exporter.hpp @@ -0,0 +1,52 @@ +//===----------------------------------------------------------------------===// +// DuckDB +// +// duckdb/planner/logical_plan_sql_exporter.hpp +// +// +//===----------------------------------------------------------------------===// + +#pragma once + +#include "duckdb/common/identifier.hpp" +#include "duckdb/common/optional.hpp" +#include "duckdb/parser/parsed_expression.hpp" +#include "duckdb/parser/query_node.hpp" +#include "duckdb/parser/tableref.hpp" +#include "duckdb/planner/column_binding.hpp" +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/logical_plan_verification_result.hpp" + +namespace duckdb { + +class ClientContext; +class LogicalOperator; +struct LogicalExtensionOperator; + +struct LogicalPlanSQLExportField { + ColumnBinding source_binding; + //! Type represented by the exported SQL expression. + LogicalType type; + //! Optimizer-selected type accepted from the bound plan, when different. + optional optimizer_type; +}; + +struct LogicalPlanSQLExportRelation { + unique_ptr query; + vector fields; +}; + +using LogicalPlanSQLExportResult = LogicalPlanVerificationResult; + +struct LogicalPlanSQLExportOptions { + optional> output_names; +}; + +class LogicalPlanSQLExporter { +public: + //! Export a verified plan to an owned query with positional binding/type fields. + DUCKDB_API static LogicalPlanVerificationResult + Export(ClientContext &context, LogicalOperator &root, const LogicalPlanSQLExportOptions &options = {}); +}; + +} // namespace duckdb diff --git a/src/duckdb/src/include/duckdb/planner/logical_plan_verification_result.hpp b/src/duckdb/src/include/duckdb/planner/logical_plan_verification_result.hpp index f100b8fcd..7d5eb3ff5 100644 --- a/src/duckdb/src/include/duckdb/planner/logical_plan_verification_result.hpp +++ b/src/duckdb/src/include/duckdb/planner/logical_plan_verification_result.hpp @@ -18,7 +18,7 @@ namespace duckdb { enum class LogicalPlanVerificationPathRoot { LOGICAL_PLAN, STANDALONE_EXPRESSION }; -enum class LogicalPlanVerificationPathComponentType { OPERATOR_CHILD, OPERATOR_EXPRESSION, EXPRESSION_CHILD }; +enum class LogicalPlanVerificationPathComponentType : int32_t { OPERATOR_CHILD, OPERATOR_EXPRESSION, EXPRESSION_CHILD }; struct LogicalPlanVerificationPathComponent { LogicalPlanVerificationPathComponentType type; @@ -92,7 +92,7 @@ struct LogicalPlanVerificationConstructIdentity { DUCKDB_API bool operator<(const LogicalPlanVerificationConstructIdentity &other) const; }; -enum class LogicalPlanVerificationIssueCode { +enum class LogicalPlanVerificationIssueCode : int32_t { INVALID_BINDING, TYPE_MISMATCH, UNSUPPORTED_OPERATOR, @@ -105,7 +105,7 @@ enum class LogicalPlanVerificationIssueCode { INTERNAL_INVARIANT }; -enum class LogicalPlanVerificationPhase { VERIFY, EXPRESSION_EXPORT, PLAN_EXPORT }; +enum class LogicalPlanVerificationPhase : int32_t { VERIFY, EXPRESSION_EXPORT, PLAN_EXPORT }; struct LogicalPlanVerificationIssue { LogicalPlanVerificationIssueCode code = LogicalPlanVerificationIssueCode::INTERNAL_INVARIANT; @@ -134,6 +134,13 @@ class LogicalPlanVerificationResult { return LogicalPlanVerificationResult(optional(std::move(value)), {}); } + //! Propagate the normalized issues of a failed result carrying a different value type + template + static LogicalPlanVerificationResult Failure(const LogicalPlanVerificationResult &failed) { + D_ASSERT(failed.HasError()); + return LogicalPlanVerificationResult(optional(), failed.GetIssues()); + } + static LogicalPlanVerificationResult Failure(vector issues) { for (auto &issue : issues) { if (!issue.IsValid()) { diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_aggregate.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_aggregate.hpp index edbd30030..39574d026 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_aggregate.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_aggregate.hpp @@ -20,6 +20,9 @@ namespace duckdb { //! operator. class LogicalAggregate : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_AGGREGATE_AND_GROUP_BY; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_column_data_get.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_column_data_get.hpp index 02621f16d..fd4665cda 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_column_data_get.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_column_data_get.hpp @@ -19,6 +19,9 @@ class ManagedResultSet; //! LogicalColumnDataGet represents a scan operation from a ColumnDataCollection class LogicalColumnDataGet : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_CHUNK_GET; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_comparison_join.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_comparison_join.hpp index 8a24c75b4..a66580f9c 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_comparison_join.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_comparison_join.hpp @@ -69,6 +69,7 @@ class LogicalComparisonJoin : public LogicalJoin { bool HasEquality(idx_t &range_count) const; bool HasArbitraryConditions() const; + bool TryGetMarkJoinGroupTypes(vector &group_types) const; }; } // namespace duckdb diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_cteref.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_cteref.hpp index d77a064b1..34f7d25e1 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_cteref.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_cteref.hpp @@ -16,6 +16,9 @@ namespace duckdb { //! LogicalCTERef represents a reference to a recursive CTE class LogicalCTERef : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_CTE_REF; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_distinct.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_distinct.hpp index c09a6aeac..a6fc5d3cc 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_distinct.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_distinct.hpp @@ -16,6 +16,9 @@ namespace duckdb { //! LogicalDistinct filters duplicate entries from its child operator class LogicalDistinct : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_DISTINCT; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_dummy_scan.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_dummy_scan.hpp index 9b29954d5..e800c69b5 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_dummy_scan.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_dummy_scan.hpp @@ -15,6 +15,9 @@ namespace duckdb { //! LogicalDummyScan represents a dummy scan returning a single row class LogicalDummyScan : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_DUMMY_SCAN; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_empty_result.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_empty_result.hpp index 253f1838f..dbced54d3 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_empty_result.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_empty_result.hpp @@ -18,6 +18,9 @@ class LogicalEmptyResult : public LogicalOperator { LogicalEmptyResult(); public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_EMPTY_RESULT; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_explain.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_explain.hpp index 338261103..35e4088af 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_explain.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_explain.hpp @@ -28,11 +28,15 @@ class LogicalExplain : public LogicalOperator { string physical_plan; string logical_plan_unopt; string logical_plan_opt; + vector sql_output_names; + bool allow_unsupported_sql = false; public: void Serialize(Serializer &serializer) const override; static unique_ptr Deserialize(Deserializer &deserializer); + unique_ptr CreateSQLResult(ClientContext &context, TableIndex table_index); + idx_t EstimateCardinality(ClientContext &context) override; bool SupportSerialization() const override; diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_expression_get.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_expression_get.hpp index 19d893421..d707a1996 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_expression_get.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_expression_get.hpp @@ -12,9 +12,14 @@ namespace duckdb { +struct LogicalPlanSQLExportField; + //! LogicalExpressionGet represents a scan operation over a set of to-be-executed expressions class LogicalExpressionGet : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_EXPRESSION_GET; public: @@ -44,6 +49,11 @@ class LogicalExpressionGet : public LogicalOperator { vector GetTableIndex() const override; string GetName() const override; +private: + LogicalPlanVerificationResult + ExportSQLInput(LogicalPlanSQLExportContext &context, const LogicalPlanVerificationPath &path, + vector fields); + protected: void ResolveTypes() override { // types are resolved in the constructor diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_extension_operator.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_extension_operator.hpp index 7eb3c3bde..5946618b2 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_extension_operator.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_extension_operator.hpp @@ -17,6 +17,9 @@ class ColumnBindingResolver; struct LogicalExtensionOperator : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_EXTENSION_OPERATOR; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_filter.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_filter.hpp index b594d8980..ce2f31eaa 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_filter.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_filter.hpp @@ -15,6 +15,9 @@ namespace duckdb { //! LogicalFilter represents a filter operation (e.g. WHERE or HAVING clause) class LogicalFilter : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_FILTER; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_get.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_get.hpp index f5d597464..fc69ede99 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_get.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_get.hpp @@ -12,6 +12,7 @@ #include "duckdb/common/enums/ordinality_request_type.hpp" #include "duckdb/planner/logical_operator.hpp" #include "duckdb/planner/table_filter_set.hpp" +#include "duckdb/planner/tableref/bound_at_clause.hpp" #include "duckdb/common/extra_operator_info.hpp" #include "duckdb/storage/table/row_group_order_options.hpp" @@ -19,10 +20,14 @@ namespace duckdb { class TableCatalogEntry; class DynamicTableFilterSet; +struct LogicalPlanSQLExportField; //! LogicalGet represents a scan operation from a data source class LogicalGet : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_GET; public: @@ -64,6 +69,10 @@ class LogicalGet : public LogicalOperator { //! pushed down into the table scan //! Stored so the can be included in explain output ExtraOperatorInfo extra_info; + //! The scan consumed a projection whose source expression is no longer retained. + bool has_pushed_projection = false; + //! The effective AT clause the table was looked up with, retained for SQL reconstruction + unique_ptr at_clause; //! Contains a reference to dynamically generated table filters (through e.g. a join up in the tree) shared_ptr dynamic_filters; //! Information for WITH ORDINALITY @@ -119,6 +128,10 @@ class LogicalGet : public LogicalOperator { void ResolveTypes() override; private: + friend class LogicalWindow; + LogicalPlanVerificationResult + ExportSQLSource(LogicalPlanSQLExportContext &context, const LogicalPlanVerificationPath &path, + optional_ptr ordinality = nullptr); LogicalGet(); private: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_join.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_join.hpp index 666a1b3b4..7a67a26d9 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_join.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_join.hpp @@ -17,6 +17,9 @@ namespace duckdb { //! LogicalJoin represents a join between two relations class LogicalJoin : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_INVALID; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_limit.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_limit.hpp index 9624e6458..9579f27da 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_limit.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_limit.hpp @@ -16,6 +16,9 @@ namespace duckdb { //! LogicalLimit represents a LIMIT clause class LogicalLimit : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_LIMIT; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_materialized_cte.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_materialized_cte.hpp index d4c6369ec..8fabccb71 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_materialized_cte.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_materialized_cte.hpp @@ -28,6 +28,9 @@ class LogicalMaterializedCTE : public LogicalCTE { } public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_MATERIALIZED_CTE; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_order.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_order.hpp index b601033d5..42169f345 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_order.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_order.hpp @@ -18,6 +18,9 @@ namespace duckdb { //! LogicalOrder represents an ORDER BY clause, sorting the data class LogicalOrder : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_ORDER_BY; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_pivot.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_pivot.hpp index b124e2126..644bd161d 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_pivot.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_pivot.hpp @@ -18,6 +18,9 @@ namespace duckdb { class LogicalPivot : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_PIVOT; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_projection.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_projection.hpp index 357128c61..6121fe3cb 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_projection.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_projection.hpp @@ -15,6 +15,9 @@ namespace duckdb { //! LogicalProjection represents the projection list in a SELECT clause class LogicalProjection : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_PROJECTION; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_recursive_cte.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_recursive_cte.hpp index c6e3b631e..6e7b79f95 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_recursive_cte.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_recursive_cte.hpp @@ -18,6 +18,9 @@ class LogicalRecursiveCTE : public LogicalCTE { LogicalRecursiveCTE(); public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_RECURSIVE_CTE; public: @@ -45,6 +48,12 @@ class LogicalRecursiveCTE : public LogicalCTE { vector GetTableIndex() const override; string GetName() const override; +private: + friend class LogicalPlanSQLExportContext; + LogicalPlanVerificationResult + ExportSQLDefinition(LogicalPlanSQLExportContext &context, const LogicalPlanVerificationPath &path, + const Identifier &name); + protected: void ResolveTypes() override; }; diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_sample.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_sample.hpp index 7bce81409..490b06fec 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_sample.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_sample.hpp @@ -16,6 +16,9 @@ namespace duckdb { //! LogicalSample represents a SAMPLE clause class LogicalSample : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_SAMPLE; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_secure_view.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_secure_view.hpp index 20236566a..b9546dbc0 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_secure_view.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_secure_view.hpp @@ -20,6 +20,9 @@ class BoundAtClause; //! acts as an optimization barrier that prevents the optimizer from pushing anything into the view. class LogicalSecureView : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_SECURE_VIEW; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_set_operation.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_set_operation.hpp index ce256bbba..342a8c6aa 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_set_operation.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_set_operation.hpp @@ -17,6 +17,9 @@ class LogicalSetOperation : public LogicalOperator { bool allow_out_of_order); public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_INVALID; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_top_n.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_top_n.hpp index aaeb99d8b..0348034c9 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_top_n.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_top_n.hpp @@ -17,6 +17,9 @@ struct DynamicFilterData; //! LogicalTopN represents a comibination of ORDER BY and LIMIT clause, using Min/Max Heap class LogicalTopN : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_TOP_N; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_unconditional_join.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_unconditional_join.hpp index c25509bbf..7635bc64b 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_unconditional_join.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_unconditional_join.hpp @@ -16,6 +16,9 @@ namespace duckdb { //! where the join condition is implicit (cross product, position, etc.) class LogicalUnconditionalJoin : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_INVALID; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_unnest.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_unnest.hpp index 7398b2fb4..779e3d7a8 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_unnest.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_unnest.hpp @@ -15,6 +15,9 @@ namespace duckdb { //! LogicalUnnest represents the logical UNNEST operator. class LogicalUnnest : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_UNNEST; public: diff --git a/src/duckdb/src/include/duckdb/planner/operator/logical_window.hpp b/src/duckdb/src/include/duckdb/planner/operator/logical_window.hpp index 6dcf5f02d..2f73c5245 100644 --- a/src/duckdb/src/include/duckdb/planner/operator/logical_window.hpp +++ b/src/duckdb/src/include/duckdb/planner/operator/logical_window.hpp @@ -16,6 +16,9 @@ namespace duckdb { //! operator. class LogicalWindow : public LogicalOperator { public: + LogicalPlanSQLExportResult ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) override; + static constexpr const LogicalOperatorType TYPE = LogicalOperatorType::LOGICAL_WINDOW; public: diff --git a/src/duckdb/src/include/duckdb/planner/planner.hpp b/src/duckdb/src/include/duckdb/planner/planner.hpp index 380f53961..9d0e0914a 100644 --- a/src/duckdb/src/include/duckdb/planner/planner.hpp +++ b/src/duckdb/src/include/duckdb/planner/planner.hpp @@ -39,6 +39,8 @@ class Planner { public: void CreatePlan(unique_ptr statement); + //! Apply the execution planner's optimizer policy, including mandatory rewrites. + void Optimize(); static void VerifyPlan(ClientContext &context, unique_ptr &op, optional_ptr map = nullptr); diff --git a/src/duckdb/src/include/duckdb/planner/sql_export/bound_expression_sql_exporter_internal.hpp b/src/duckdb/src/include/duckdb/planner/sql_export/bound_expression_sql_exporter_internal.hpp new file mode 100644 index 000000000..c2b7c45a0 --- /dev/null +++ b/src/duckdb/src/include/duckdb/planner/sql_export/bound_expression_sql_exporter_internal.hpp @@ -0,0 +1,171 @@ +//===----------------------------------------------------------------------===// +// DuckDB +// +// duckdb/planner/sql_export/bound_expression_sql_exporter_internal.hpp +// +// +//===----------------------------------------------------------------------===// + +#pragma once + +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" +#include "duckdb/function/function.hpp" +#include "duckdb/parser/qualified_name.hpp" + +namespace duckdb { + +class BoundAggregateExpression; +class BoundCaseExpression; +class BoundColumnRefExpression; +class BoundConjunctionExpression; +class BoundConstantExpression; +class BoundFunctionExpression; +class BoundLambdaExpression; +class BoundOperatorExpression; +class BoundReferenceExpression; +class BoundUnnestExpression; +class BoundWindowExpression; + +using BoundExpressionSQLExportResult = LogicalPlanVerificationResult>; +using BoundAggregateSQLExportResult = LogicalPlanVerificationResult>; + +class BoundExpressionSQLExportState { +public: + static LogicalPlanVerificationIssue + InternalInvariant(optional path, string message, + optional construct = {}); + + static LogicalPlanVerificationIssue InternalExpressionInvariant(const LogicalPlanVerificationPath &path, + const Expression &expression, string message); + + static LogicalPlanVerificationIssue UnsupportedFeature(const LogicalPlanVerificationPath &path, string feature, + string message); + + static LogicalPlanVerificationIssue UnsupportedFunction(const LogicalPlanVerificationPath &path, + LogicalPlanVerificationFunctionIdentity identity, + string message); + + static bool HasNestedCollation(const LogicalType &type); + + static BoundExpressionSQLExportResult PreserveCollation(const LogicalType &type, + BoundExpressionSQLExportResult result, + const LogicalPlanVerificationPath &path); + + template + static LogicalPlanVerificationFunctionIdentity DefinitionFunctionIdentity(const FUNCTION &definition, + const vector &arguments, + const LogicalType &return_type) { + LogicalPlanVerificationFunctionIdentity identity; + identity.catalog = definition.GetCatalogName().GetIdentifierName(); + identity.schema = definition.GetSchemaName().GetIdentifierName(); + identity.name = definition.GetName().GetIdentifierName(); + identity.arguments = arguments; + identity.return_type = return_type; + return identity; + } + + template + static optional RebindableFunctionName(const FUNCTION &definition) { + auto name = definition.GetQualifiedName(); + if (name.Catalog().empty()) { + name = name.WithCatalog(Identifier::SystemCatalog()); + } + if (name.Schema().empty()) { + name = QualifiedName(name.Catalog(), Identifier::DefaultSchema(), name.Name()); + } + if (name.Path().empty()) { + return {}; + } + for (auto &component : name.Path()) { + if (component.empty()) { + return {}; + } + } + return name; + } + + static LogicalType SQLCastType(const LogicalType &type); + + static unique_ptr SQLCast(const LogicalType &type, unique_ptr child, + bool try_cast = false); + + explicit BoundExpressionSQLExportState(const BoundExpressionSQLExportContext &context_p); + BoundExpressionSQLExportResult Export(const Expression &expression, const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportWindow(const BoundWindowExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportUnnest(const BoundUnnestExpression &expression, + const LogicalPlanVerificationPath &path); + BoundAggregateSQLExportResult ExportAggregateCall(const BoundAggregateExpression &expression, + const LogicalPlanVerificationPath &path); + +private: + struct ChildExpression { + explicit ChildExpression(optional_ptr expression_p, + optional expected_type_p = {}) + : expression(expression_p), expected_type(std::move(expected_type_p)) { + } + + optional_ptr expression; + optional expected_type; + }; + BoundExpressionSQLExportResult ExportInternal(const Expression &expression, + const LogicalPlanVerificationPath &path); + template + BoundExpressionSQLExportResult ExportWindowFunction(const BoundWindowExpression &expression, + const FUNCTION &function, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult RestoreResultType(const LogicalType &type, unique_ptr result, + const LogicalPlanVerificationPath &path); + static bool RequiresConstantConstructor(const LogicalType &type); + BoundExpressionSQLExportResult CastToConstructedType(const LogicalType &type, unique_ptr child, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportConstant(const BoundConstantExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportReference(const BoundReferenceExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportLambda(const BoundLambdaExpression &lambda, + const BoundFunctionExpression &function, idx_t logical_argument_count, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportColumnRef(const BoundColumnRefExpression &expression, + const LogicalPlanVerificationPath &path); + LogicalPlanVerificationIssue InvalidBinding(const LogicalPlanVerificationPath &path, const ColumnBinding &binding, + string message); + BoundExpressionSQLExportResult ExportFunction(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportCast(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportComparison(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportBetween(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportConjunction(const BoundConjunctionExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportCase(const BoundCaseExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportOperator(const BoundOperatorExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult CompressedMaterializationFailure(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path, + string message); + optional + TryExportCompressedMaterialization(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportScalarFunction(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path); + BoundAggregateSQLExportResult BuildAggregateCall(const BoundAggregateExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportAggregate(const BoundAggregateExpression &expression, + const LogicalPlanVerificationPath &path); + BoundExpressionSQLExportResult ExportChild(const Expression &expression, const LogicalPlanVerificationPath &path, + idx_t child_index); + void ExportChildren(const vector> &source, const LogicalPlanVerificationPath &path, + vector> &result, vector &issues, + const optional &expected_type = {}); + void ExportChildren(const vector &source, const LogicalPlanVerificationPath &path, + vector> &result, vector &issues); + const BoundExpressionSQLExportContext &context; + vector>> lambda_reference_scopes; +}; + +} // namespace duckdb diff --git a/src/duckdb/src/include/duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp b/src/duckdb/src/include/duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp new file mode 100644 index 000000000..2ff9927a9 --- /dev/null +++ b/src/duckdb/src/include/duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp @@ -0,0 +1,97 @@ +//===----------------------------------------------------------------------===// +// DuckDB +// +// duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp +// +// +//===----------------------------------------------------------------------===// + +#pragma once + +#include "duckdb/planner/logical_plan_sql_export_context.hpp" +#include "duckdb/planner/column_binding_map.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" +#include "duckdb/parser/query_node/select_node.hpp" + +namespace duckdb { + +class Expression; +class LogicalOperator; +class LogicalMaterializedCTE; +class LogicalCTERef; +class LogicalRecursiveCTE; +class LogicalComparisonJoin; +class LogicalColumnDataGet; +class LogicalGet; +class LogicalExpressionGet; +class LogicalFilter; +class LogicalProjection; +class LogicalSecureView; +class LogicalSample; +class LogicalPivot; +class LogicalLimit; +class LogicalSetOperation; +class LogicalAggregate; +struct LogicalExtensionOperator; +class ClientContext; + +using LogicalPlanSQLFieldResult = LogicalPlanVerificationResult>; +struct SQLSourceQueryResult { + unique_ptr query; + string unsupported_reason; +}; + +//! Helpers shared by the logical operators' SQL reconstruction +struct LogicalPlanSQLExportHelpers { + static SQLSourceQueryResult ReconstructSQLSource(ClientContext &context, const LogicalGet &get, + unique_ptr input, const Identifier &relation_alias, + bool source_ordinality); + + static LogicalPlanVerificationPath PlanChildPath(const LogicalPlanVerificationPath &path, idx_t ordinal); + + static LogicalPlanVerificationPath PlanExpressionPath(const LogicalPlanVerificationPath &path, idx_t ordinal); + + static LogicalPlanVerificationResult> + ExportTypedNull(const LogicalType &type, const LogicalPlanVerificationPath &path); + + static LogicalPlanVerificationIssue PlanUnsupportedFeature(const LogicalPlanVerificationPath &path, string feature, + string message); + + static LogicalPlanVerificationFunctionIdentity LogicalSourceIdentity(); + + static LogicalPlanVerificationFunctionIdentity LogicalSourceIdentity(const LogicalGet &get); + + static LogicalPlanVerificationIssue UnsupportedSource(const LogicalPlanVerificationPath &path, + LogicalPlanVerificationFunctionIdentity source, string guard); + + static Identifier FieldIdentifier(idx_t ordinal); + + static LogicalPlanSQLFieldResult CreateFields(LogicalOperator &op, const LogicalPlanVerificationPath &path); + + static BoundExpressionSQLExportContext + CreateBindingContext(ClientContext &context, const vector> &children, + const vector> &plain_scopes = {}); + + static void PropagateSemanticTypes(vector &fields, + const vector> &children); + + static unique_ptr CreateSubquery(LogicalPlanSQLExportedChild child); + + static unique_ptr ChildColumn(const LogicalPlanSQLExportedChild &child, idx_t field_index, + optional_ptr plain = nullptr); + + static vector> CollectExpressions(const LogicalOperator &op); + + static bool HasEffectfulExpressions(const LogicalOperator &op); + + static bool CollectScopeAliases(const TableRef &table, identifier_set_t &aliases); + + static optional_ptr PlainScope(const QueryNode &query); + + static void SetChildScope(SelectNode &select, LogicalPlanSQLExportedChild child, + optional_ptr plain); + + static bool IsIdentityProjection(const LogicalProjection &projection, + const vector &fields); +}; +} // namespace duckdb diff --git a/src/duckdb/src/include/duckdb/planner/sql_export_helpers.hpp b/src/duckdb/src/include/duckdb/planner/sql_export_helpers.hpp index adb6665cf..84e7d68cc 100644 --- a/src/duckdb/src/include/duckdb/planner/sql_export_helpers.hpp +++ b/src/duckdb/src/include/duckdb/planner/sql_export_helpers.hpp @@ -1,76 +1,111 @@ +//===----------------------------------------------------------------------===// +// DuckDB +// +// duckdb/planner/sql_export_helpers.hpp +// +// +//===----------------------------------------------------------------------===// + #pragma once #include "duckdb/common/identifier.hpp" +#include "duckdb/parser/expression/cast_expression.hpp" +#include "duckdb/parser/expression/collate_expression.hpp" +#include "duckdb/parser/expression/conjunction_expression.hpp" +#include "duckdb/parser/expression/function_expression.hpp" #include "duckdb/common/type_visitor.hpp" -#include "duckdb/common/unordered_set.hpp" #include "duckdb/planner/logical_plan_verification_result.hpp" namespace duckdb { -namespace SQLExportHelpers { +struct SQLExportHelpers { + template + static unique_ptr SystemFunction(Identifier name, vector arguments) { + return make_uniq( + QualifiedName(Identifier::SystemCatalog(), Identifier::DefaultSchema(), std::move(name)), + std::move(arguments)); + } + + static unique_ptr Conjoin(unique_ptr left, unique_ptr right) { + if (!left) { + return right; + } + if (!right) { + return left; + } + return make_uniq(ExpressionType::CONJUNCTION_AND, std::move(left), std::move(right)); + } -inline LogicalPlanVerificationPath -ChildPath(const LogicalPlanVerificationPath &path, idx_t ordinal, - LogicalPlanVerificationPathComponentType type = LogicalPlanVerificationPathComponentType::EXPRESSION_CHILD) { - auto result = path; - result.components.push_back({type, ordinal}); - return result; -} + static unique_ptr Conjoin(vector> predicates) { + if (predicates.empty()) { + return nullptr; + } + if (predicates.size() == 1) { + return std::move(predicates[0]); + } + return make_uniq(ExpressionType::CONJUNCTION_AND, std::move(predicates)); + } -inline LogicalPlanVerificationIssue MakeIssue(LogicalPlanVerificationIssueCode code, LogicalPlanVerificationPhase phase, - optional path, - optional construct, - string message) { - LogicalPlanVerificationIssue issue; - issue.code = code; - issue.phase = phase; - issue.path = std::move(path); - issue.construct = std::move(construct); - issue.message = std::move(message); - return issue; -} + static LogicalPlanVerificationPath ChildPath( + const LogicalPlanVerificationPath &path, idx_t ordinal, + LogicalPlanVerificationPathComponentType type = LogicalPlanVerificationPathComponentType::EXPRESSION_CHILD) { + auto result = path; + result.components.push_back({type, ordinal}); + return result; + } -inline bool IsValidIdentifier(const Identifier &identifier) { - auto &name = identifier.GetIdentifierName(); - return !name.empty() && name.find('\0') == string::npos && Value::StringIsValid(name); -} + static LogicalPlanVerificationIssue MakeIssue(LogicalPlanVerificationIssueCode code, + LogicalPlanVerificationPhase phase, + optional path, + optional construct, + string message) { + LogicalPlanVerificationIssue issue; + issue.code = code; + issue.phase = phase; + issue.path = std::move(path); + issue.construct = std::move(construct); + issue.message = std::move(message); + return issue; + } -inline bool IsSQLExportType(LogicalTypeId id) { - static const auto admitted_ids = [] { - // SQL export follows AllTypes' value-type coverage, with SQLNULL included and TUPLE excluded. - unordered_set ids {LogicalTypeId::SQLNULL}; - for (auto &type : LogicalType::AllTypes()) { - if (type.id() != LogicalTypeId::TUPLE) { - ids.insert(type.id()); + static unique_ptr OrderExpression(const LogicalType &type, + unique_ptr expression) { + if (type.id() == LogicalTypeId::VARCHAR) { + if (type.HasAlias()) { + expression = make_uniq(LogicalType::VARCHAR, std::move(expression)); } + return make_uniq("C", std::move(expression)); } - return ids; - }(); - return admitted_ids.count(id) != 0; -} + return expression; + } -inline bool IsSQLRepresentableType(const LogicalType &type) { - return type.IsComplete() && - !TypeVisitor::Contains(type, [](const LogicalType &child) { return !IsSQLExportType(child.id()); }); -} + static bool IsValidIdentifier(const Identifier &identifier) { + auto &name = identifier.GetIdentifierName(); + return !name.empty() && name.find('\0') == string::npos && Value::StringIsValid(name); + } -inline bool SQLTypesMatch(const LogicalType &left, const LogicalType &right) { - if (left != right) { - return false; + static bool IsSQLExportType(LogicalTypeId id) { + return TypeExpression::IsSQLType(id); + } + + static bool IsSQLRepresentableType(const LogicalType &type) { + return TypeExpression::CanRepresent(type); + } + + static bool IsSQLValueType(const LogicalType &type) { + return type.IsComplete() && + !TypeVisitor::Contains(type, [](const LogicalType &child) { return !IsSQLExportType(child.id()); }); } - vector collations; - TypeVisitor::Contains(left, [&](const LogicalType &child) { - if (child.id() == LogicalTypeId::VARCHAR) { - collations.push_back(StringType::GetCollation(child)); - } - return false; - }); - idx_t index = 0; - const auto mismatch = TypeVisitor::Contains(right, [&](const LogicalType &child) { - return child.id() == LogicalTypeId::VARCHAR && - (index >= collations.size() || collations[index++] != StringType::GetCollation(child)); - }); - return !mismatch && index == collations.size(); -} -} // namespace SQLExportHelpers + static string TypeCollationSignature(const LogicalType &type) { + string result; + TypeVisitor::Contains(type, [&](const LogicalType &child) { + if (child.id() == LogicalTypeId::VARCHAR) { + const auto collation = StringType::GetCollation(child); + result += std::to_string(collation.size()) + ":" + collation + ";"; + } + return false; + }); + return result; + } +}; } // namespace duckdb diff --git a/src/duckdb/src/include/duckdb/planner/tableref/bound_at_clause.hpp b/src/duckdb/src/include/duckdb/planner/tableref/bound_at_clause.hpp index b5e27efda..039a1335c 100644 --- a/src/duckdb/src/include/duckdb/planner/tableref/bound_at_clause.hpp +++ b/src/duckdb/src/include/duckdb/planner/tableref/bound_at_clause.hpp @@ -12,6 +12,9 @@ namespace duckdb { +class Serializer; +class Deserializer; + //! The AT clause specifies which version of a table to read class BoundAtClause { public: @@ -26,6 +29,9 @@ class BoundAtClause { return val; } + void Serialize(Serializer &serializer) const; + static unique_ptr Deserialize(Deserializer &deserializer); + private: //! The unit (e.g. TIMESTAMP or VERSION) Identifier unit; diff --git a/src/duckdb/src/include/duckdb/storage/compression/chimp/chimp_fetch.hpp b/src/duckdb/src/include/duckdb/storage/compression/chimp/chimp_fetch.hpp index 6898f1ebb..8261fcd75 100644 --- a/src/duckdb/src/include/duckdb/storage/compression/chimp/chimp_fetch.hpp +++ b/src/duckdb/src/include/duckdb/storage/compression/chimp/chimp_fetch.hpp @@ -24,16 +24,16 @@ void ChimpFetchRow(ColumnSegment &segment, ColumnFetchState &state, row_t row_id auto &buffer_manager = BufferManager::GetBufferManager(segment.GetDatabase()); auto handle = buffer_manager.Pin(state.context, segment.GetBlockHandle()); - ChimpScanState scan_state(std::move(handle), segment); - scan_state.Skip(segment, UnsafeNumericCast(row_id)); + auto scan_state = make_uniq>(std::move(handle), segment); + scan_state->Skip(segment, UnsafeNumericCast(row_id)); auto result_data = FlatVector::GetDataMutableUnsafe(result); - if (scan_state.GroupFinished() && scan_state.total_value_count < scan_state.segment_count) { - scan_state.LoadGroup(scan_state.group_state.values); + if (scan_state->GroupFinished() && scan_state->total_value_count < scan_state->segment_count) { + scan_state->LoadGroup(scan_state->group_state.values); } - scan_state.group_state.Scan(&result_data[result_idx], 1); + scan_state->group_state.Scan(&result_data[result_idx], 1); - scan_state.total_value_count++; + scan_state->total_value_count++; } } // namespace duckdb diff --git a/src/duckdb/src/main/client_context.cpp b/src/duckdb/src/main/client_context.cpp index 8cbb81d22..2f701c6ca 100644 --- a/src/duckdb/src/main/client_context.cpp +++ b/src/duckdb/src/main/client_context.cpp @@ -317,7 +317,7 @@ void ClientContext::BeginQueryInternal(ClientContextLock &lock, const SQLStateme throw ErrorManager::InvalidatedDatabase(*this, ValidChecker::InvalidatedMessage(db_inst)); } active_query = make_uniq(); - if (transaction.IsAutoCommit()) { + if (transaction.IsAutoCommit() && !transaction.HasActiveTransaction()) { transaction.BeginTransaction(); } @@ -494,6 +494,9 @@ shared_ptr ClientContext::CreatePreparedStatementInternal D_ASSERT(logical_planner.plan || !logical_planner.properties.bound_all_parameters); } + if (logical_planner.properties.bound_all_parameters) { + logical_planner.Optimize(); + } auto logical_plan = std::move(logical_planner.plan); // extract the result column names from the plan result->properties = logical_planner.properties; @@ -504,31 +507,6 @@ shared_ptr ClientContext::CreatePreparedStatementInternal // not all parameters were bound - return return result; } -#ifdef DEBUG - logical_plan->Verify(*this); -#endif - bool optimize = Settings::Get(*this); - if (Settings::Get(*this)) { - // verify disable optimizer - disable EXCEPT for explain, otherwise every single EXPLAIN query breaks - if (logical_plan->type != LogicalOperatorType::LOGICAL_EXPLAIN) { - optimize = false; - } - } - if (logical_plan->RequireOptimizer()) { - { - auto optimizer_timer = profiler.StartTimer(); - Optimizer optimizer(*logical_planner.binder, *this); - if (optimize) { - logical_plan = optimizer.Optimize(std::move(logical_plan)); - } else { - logical_plan = optimizer.LowerMandatoryAggregateRewrites(std::move(logical_plan)); - } - D_ASSERT(logical_plan); - } -#ifdef DEBUG - logical_plan->Verify(*this); -#endif - } // Convert the logical query plan into a physical query plan. { diff --git a/src/duckdb/src/main/client_verify.cpp b/src/duckdb/src/main/client_verify.cpp index 31fb8ab3e..b4c6b78d0 100644 --- a/src/duckdb/src/main/client_verify.cpp +++ b/src/duckdb/src/main/client_verify.cpp @@ -239,7 +239,13 @@ void ClientContext::StatementVerification(ClientContextLock &lock, unique_ptrnamed_values = std::move(prep_verifier.values); ReplaceStatement(statement, std::move(execute)); - } else if (verification == DebugStatementVerification::EXPLAIN_STATEMENT) { + } else if (verification == DebugStatementVerification::EXPLAIN_STATEMENT || + verification == DebugStatementVerification::EXPLAIN_SQL || + verification == DebugStatementVerification::EXPLAIN_SQL_STRICT) { + const bool export_sql = verification != DebugStatementVerification::EXPLAIN_STATEMENT; + if (export_sql && statement->type != StatementType::SELECT_STATEMENT) { + return; + } if (statement->type == StatementType::EXPLAIN_STATEMENT) { // don't explain explain... return; @@ -252,20 +258,64 @@ void ClientContext::StatementVerification(ClientContextLock &lock, unique_ptr(statement->Copy()); - // Disable the profiler during the verification EXPLAIN to prevent it from consuming the profiler context - // (which would lose parser timing captured before StatementVerification was called) and from - // overwriting the profiling output file with the EXPLAIN's profiling data. - auto &client_config = ClientConfig::GetConfig(*this); - bool saved_profiler = client_config.enable_profiler; - ScopedConfigSetting suppress_profiling( - client_config, [](ClientConfig &config) { config.enable_profiler = false; }, - [saved_profiler](ClientConfig &config) { config.enable_profiler = saved_profiler; }); - auto explain_result = RunStatementInternal(lock, std::move(explain_stmt), query_parameters); - if (explain_result->HasError()) { - explain_result->ThrowError(); + auto explain_and_replace = [&]() { + // deliberately left without source text: the EXPLAIN runs as a nested statement, and giving it a + // query would let it consume the error location that belongs to the statement we are verifying + auto explain_stmt = make_uniq( + statement->Copy(), export_sql ? ExplainType::EXPLAIN_SQL : ExplainType::EXPLAIN_STANDARD); + explain_stmt->allow_unsupported_sql = verification == DebugStatementVerification::EXPLAIN_SQL; + // Disable the profiler during the verification EXPLAIN to prevent it from consuming the profiler context + // (which would lose parser timing captured before StatementVerification was called) and from + // overwriting the profiling output file with the EXPLAIN's profiling data. + auto &client_config = ClientConfig::GetConfig(*this); + bool saved_profiler = client_config.enable_profiler; + ScopedConfigSetting suppress_profiling( + client_config, [](ClientConfig &config) { config.enable_profiler = false; }, + [saved_profiler](ClientConfig &config) { config.enable_profiler = saved_profiler; }); + auto explain_result = RunStatementInternal(lock, std::move(explain_stmt), query_parameters, false); + if (explain_result->HasError()) { + explain_result->ThrowError(); + } + if (!export_sql) { + return; + } + auto chunk = explain_result->Fetch(); + if (!chunk || chunk->size() == 0) { + D_ASSERT(verification == DebugStatementVerification::EXPLAIN_SQL); + return; + } + D_ASSERT(chunk && chunk->size() == 1 && chunk->ColumnCount() == 2); + auto sql = chunk->GetValue(1, 0).GetValue(); + auto parser_options = GetParserOptions(); + parser_options.identifier_case_mode = IdentifierCaseMode::PRESERVE_CASE; + Parser parser(parser_options); + parser.ParseQuery(sql); + if (parser.statements.size() != 1 || parser.statements[0]->type != StatementType::SELECT_STATEMENT || + !parser.statements[0]->named_param_map.empty()) { + throw InternalException("SQL export verification did not produce one parameter-free SELECT"); + } + ReplaceStatement(statement, std::move(parser.statements[0])); + }; + // The generated SQL encodes what the optimizer derived under the planning snapshot (statistics, exact + // cardinalities), so it must execute in the transaction the EXPLAIN planned it in + const bool shared_transaction = export_sql && transaction.IsAutoCommit(); + if (shared_transaction) { + transaction.SetAutoCommit(false); + } + try { + explain_and_replace(); + } catch (...) { + if (shared_transaction) { + if (transaction.HasActiveTransaction()) { + transaction.Rollback(nullptr); + } + transaction.SetAutoCommit(true); + } + throw; + } + if (shared_transaction) { + // the transaction stays open and the verified statement commits it + transaction.SetAutoCommit(true); } } } diff --git a/src/duckdb/src/main/config.cpp b/src/duckdb/src/main/config.cpp index 549c707cb..83a5660a4 100644 --- a/src/duckdb/src/main/config.cpp +++ b/src/duckdb/src/main/config.cpp @@ -593,6 +593,11 @@ CastFunctionSet &DBConfig::GetCastFunctions() { return type_manager->GetCastFunctions(); } +const CastFunctionSet &DBConfig::GetCastFunctions() const { + const auto &manager = *type_manager; + return manager.GetCastFunctions(); +} + TypeManager &DBConfig::GetTypeManager() { return *type_manager; } diff --git a/src/duckdb/src/optimizer/join_order/relation_manager.cpp b/src/duckdb/src/optimizer/join_order/relation_manager.cpp index 7d83d9cef..f963cb088 100644 --- a/src/duckdb/src/optimizer/join_order/relation_manager.cpp +++ b/src/duckdb/src/optimizer/join_order/relation_manager.cpp @@ -264,6 +264,22 @@ static bool JoinIsReorderable(LogicalOperator &op) { return false; } +static bool IsProjectedRecursiveCTERef(const LogicalOperator &op, const JoinOrderOptimizer &optimizer) { + optional_ptr source = op; + while (source->children.size() == 1) { + if (OperatorIsNonReorderable(source->type) || + (OperatorNeedsRelation(source->type) && source->type != LogicalOperatorType::LOGICAL_PROJECTION)) { + return false; + } + source = source->children[0].get(); + } + if (source->type != LogicalOperatorType::LOGICAL_CTE_REF) { + return false; + } + auto &cte_ref = source->Cast(); + return optimizer.recursive_cte_indexes.find(cte_ref.cte_index) != optimizer.recursive_cte_indexes.end(); +} + static bool RecursiveCTERefCanReorder(optional_ptr parent) { if (!parent || parent->type != LogicalOperatorType::LOGICAL_COMPARISON_JOIN) { return false; @@ -584,7 +600,11 @@ bool RelationManager::ExtractJoinRelations(JoinOrderOptimizer &optimizer, Logica } proj.SetEstimatedCardinality(proj_stats->cardinality); ModifyStatsIfLimit(limit_op.get(), *proj_stats); - return AddRelation(input_op, parent, *proj_stats); + if (!AddRelation(input_op, parent, *proj_stats)) { + return false; + } + // A projection must preserve the recursive reference's join-order restriction. + return !IsProjectedRecursiveCTERef(proj, optimizer) || RecursiveCTERefCanReorder(parent); } case LogicalOperatorType::LOGICAL_EMPTY_RESULT: { // optimize the child and copy the stats diff --git a/src/duckdb/src/optimizer/remove_unused_columns.cpp b/src/duckdb/src/optimizer/remove_unused_columns.cpp index ab3c4dfbd..4af9f46b1 100644 --- a/src/duckdb/src/optimizer/remove_unused_columns.cpp +++ b/src/duckdb/src/optimizer/remove_unused_columns.cpp @@ -36,7 +36,6 @@ #include "duckdb/optimizer/optimizer.hpp" #include "duckdb/planner/logical_operator_repeatability.hpp" #include "duckdb/planner/subquery/column_binding_layout.hpp" -#include "duckdb/planner/sql_export_helpers.hpp" #include "duckdb/planner/expression_iterator.hpp" #include @@ -468,7 +467,7 @@ void RemoveUnusedColumns::VisitSecureView(LogicalSecureView &view) { unique_ptr source_expression; for (idx_t i = 0; i < view.output_bindings.size(); i++) { if (view.output_bindings[i] != new_bindings[output_idx] || - !SQLExportHelpers::SQLTypesMatch(view.output_expressions[i]->GetReturnType(), new_types[output_idx])) { + !view.output_expressions[i]->GetReturnType().EqualsIncludingCollation(new_types[output_idx])) { continue; } if (source_expression && !source_expression->Equals(*view.output_expressions[i])) { diff --git a/src/duckdb/src/parser/expression/star_expression.cpp b/src/duckdb/src/parser/expression/star_expression.cpp index 30848415f..8603ae4bc 100644 --- a/src/duckdb/src/parser/expression/star_expression.cpp +++ b/src/duckdb/src/parser/expression/star_expression.cpp @@ -26,42 +26,21 @@ string StarExpression::ToString() const { result += relation_name.empty() ? "*" : SQLIdentifier(relation_name) + ".*"; if (!exclude_list.empty()) { result += " EXCLUDE ("; - bool first_entry = true; - for (auto &entry : exclude_list) { - if (!first_entry) { - result += ", "; - } - result += entry.ToString(); - first_entry = false; - } + result += StringUtil::Join(exclude_list, ", ", [](const auto &entry) { return entry.ToString(); }); result += ")"; } if (!replace_list.empty()) { result += " REPLACE ("; - bool first_entry = true; - for (auto &entry : replace_list) { - if (!first_entry) { - result += ", "; - } - result += entry.second->ToString(); - result += " AS "; - result += SQLIdentifier(entry.first); - first_entry = false; - } + result += StringUtil::Join(replace_list, ", ", [](const auto &entry) { + return entry.second->ToString() + " AS " + SQLIdentifier(entry.first); + }); result += ")"; } if (!rename_list.empty()) { result += " RENAME ("; - bool first_entry = true; - for (auto &entry : rename_list) { - if (!first_entry) { - result += ", "; - } - result += entry.first.ToString(); - result += " AS "; - result += SQLIdentifier(entry.second); - first_entry = false; - } + result += StringUtil::Join(rename_list, ", ", [](const auto &entry) { + return entry.first.ToString() + " AS " + SQLIdentifier(entry.second); + }); result += ")"; } if (columns) { diff --git a/src/duckdb/src/parser/expression/type_expression.cpp b/src/duckdb/src/parser/expression/type_expression.cpp index 1844ec6e8..21ce59632 100644 --- a/src/duckdb/src/parser/expression/type_expression.cpp +++ b/src/duckdb/src/parser/expression/type_expression.cpp @@ -7,8 +7,30 @@ #include "duckdb/common/string_util.hpp" #include "duckdb/parser/expression/constant_expression.hpp" #include "duckdb/common/types/hash.hpp" +#include "duckdb/common/type_visitor.hpp" +#include "duckdb/common/unordered_set.hpp" namespace duckdb { +bool TypeExpression::IsSQLType(LogicalTypeId id) { + static const auto admitted_ids = [] { + // TYPE has SQL syntax but is not included in AllTypes. + unordered_set ids {LogicalTypeId::SQLNULL, LogicalTypeId::TYPE}; + for (auto &type : LogicalType::AllTypes()) { + ids.insert(type.id()); + } + return ids; + }(); + return admitted_ids.count(id) != 0; +} + +bool TypeExpression::CanRepresent(const LogicalType &type) { + return type.IsComplete() && !TypeVisitor::Contains(type, [](const LogicalType &child) { + const bool empty_tuple = child.id() == LogicalTypeId::TUPLE && StructType::GetChildCount(child) == 0; + const bool empty_enum = child.id() == LogicalTypeId::ENUM && EnumType::GetSize(child) == 0; + return !IsSQLType(child.id()) || empty_tuple || empty_enum; + }); +} + TypeExpression::TypeExpression(QualifiedName qualified_name_p, vector> children_p) : ParsedExpression(ExpressionType::TYPE, ExpressionClass::TYPE), qualified_name(std::move(qualified_name_p)), children(std::move(children_p)) { diff --git a/src/duckdb/src/parser/expression/value_expression.cpp b/src/duckdb/src/parser/expression/value_expression.cpp index 7c89645fd..43890a99a 100644 --- a/src/duckdb/src/parser/expression/value_expression.cpp +++ b/src/duckdb/src/parser/expression/value_expression.cpp @@ -1,4 +1,6 @@ #include "duckdb/common/string_util.hpp" +#include "duckdb/common/types/geometry_crs.hpp" +#include "duckdb/parser/expression/case_expression.hpp" #include "duckdb/common/types/value.hpp" #include "duckdb/common/value_operations/value_operations.hpp" #include "duckdb/parser/expression/cast_expression.hpp" @@ -6,6 +8,12 @@ #include "duckdb/parser/expression/function_expression.hpp" #include "duckdb/parser/expression/type_expression.hpp" +#include "duckdb/common/type_visitor.hpp" +#include "duckdb/common/types/variant_iterator.hpp" +#include "duckdb/common/types/vector.hpp" +#include "duckdb/function/scalar/generic_common.hpp" +#include "duckdb/parser/expression/collate_expression.hpp" + #include namespace duckdb { @@ -64,7 +72,7 @@ static unique_ptr StringCast(const Value &value) { } static unique_ptr ListValueExpression(vector> children) { - return make_uniq("list_value", std::move(children)); + return make_uniq(QualifiedName("system", "main", "list_value"), std::move(children)); } static vector> ChildExpressions(const vector &values) { @@ -83,14 +91,15 @@ static unique_ptr NamedArgument(const Identifier &name, const child->SetAlias(name); vector arguments; arguments.emplace_back(name, std::move(child)); - return make_uniq(Identifier(function_name), std::move(arguments)); + return make_uniq(QualifiedName("system", "main", Identifier(function_name)), + std::move(arguments)); } static unique_ptr StructExpression(const Value &value) { auto &type = value.type(); auto &children = StructValue::GetChildren(value); if (StructType::IsUnnamed(type)) { - return make_uniq("row", ChildExpressions(children)); + return make_uniq(QualifiedName("system", "main", "row"), ChildExpressions(children)); } vector arguments; for (idx_t i = 0; i < children.size(); i++) { @@ -99,7 +108,7 @@ static unique_ptr StructExpression(const Value &value) { child->SetAlias(name); arguments.emplace_back(name, std::move(child)); } - return make_uniq("struct_pack", std::move(arguments)); + return make_uniq(QualifiedName("system", "main", "struct_pack"), std::move(arguments)); } static unique_ptr MapExpression(const Value &value) { @@ -114,11 +123,204 @@ static unique_ptr MapExpression(const Value &value) { arguments.push_back(ListValueExpression(std::move(keys))); arguments.push_back(ListValueExpression(std::move(values))); // the map constructor cannot reproduce empty or NULL-only key/value types - return CastTo(value.type(), make_uniq("map", std::move(arguments))); + return CastTo(value.type(), + make_uniq(QualifiedName("system", "main", "map"), std::move(arguments))); +} + +static unique_ptr ValueFunction(const string &name, vector> arguments) { + return make_uniq(QualifiedName("system", "main", Identifier(name)), std::move(arguments)); +} + +static unique_ptr IntervalExpression(const interval_t &value) { + vector> parts; + for (auto &part : vector> {{"to_months", Value::INTEGER(value.months)}, + {"to_days", Value::INTEGER(value.days)}, + {"to_microseconds", Value::BIGINT(value.micros)}}) { + vector> arguments; + arguments.push_back(ConstantExpression::FromValue(part.second)); + parts.push_back(ValueFunction(part.first, std::move(arguments))); + } + vector> sum; + sum.push_back(std::move(parts[0])); + sum.push_back(std::move(parts[1])); + auto months_and_days = ValueFunction("add", std::move(sum)); + vector> result; + result.push_back(std::move(months_and_days)); + result.push_back(std::move(parts[2])); + return ValueFunction("add", std::move(result)); +} + +static unique_ptr GeometryExpression(const Value &value) { + auto &type = value.type(); + unique_ptr result; + if (value.IsNull()) { + auto geometry = GeoType::HasCRS(type) ? Value("GEOMETRYCOLLECTION EMPTY") : Value(); + result = CastTo(LogicalType::GEOMETRY(), ConstantExpression::FromValue(geometry)); + } else { + vector> arguments; + arguments.push_back(ConstantExpression::FromValue(Value::BLOB_RAW(StringValue::Get(value)))); + result = ValueFunction("st_geomfromwkb", std::move(arguments)); + } + if (!GeoType::HasCRS(type)) { + return result; + } + vector> arguments; + arguments.push_back(std::move(result)); + arguments.push_back(ConstantExpression::String(GeoType::GetCRS(type).GetDefinition())); + result = ValueFunction("st_setcrs", std::move(arguments)); + if (value.IsNull()) { + auto typed_null = make_uniq(); + typed_null->CaseChecksMutable().push_back({ConstantExpression::Boolean(false), std::move(result)}); + typed_null->ElseMutable() = ConstantExpression::Null(); + result = std::move(typed_null); + } + return result; +} + +static bool HasUnsupportedVariantKeys(const VariantNode &node) { + // Struct literals require nonempty, case-insensitively unique keys; inspect before the lossy STRUCT conversion. + if (node.GetTypeId() == VariantLogicalType::OBJECT) { + identifier_set_t names; + for (auto &child : node.GetObjectChildren()) { + if (child.key.GetSize() == 0 || !names.insert(Identifier(child.key.GetString())).second || + HasUnsupportedVariantKeys(child.value)) { + return true; + } + } + } else if (node.GetTypeId() == VariantLogicalType::ARRAY) { + for (auto child : node.GetArrayChildren()) { + if (HasUnsupportedVariantKeys(child)) { + return true; + } + } + } + return false; +} + +bool ConstantExpression::RequiresTypeWitness(const LogicalType &type) { + return TypeVisitor::Contains(type, [&](const LogicalType &child) { + return child.IsAggregateState() || child.id() == LogicalTypeId::TYPE || + (child.id() == LogicalTypeId::GEOMETRY && GeoType::HasCRS(child)) || + (type.id() != LogicalTypeId::VARCHAR && child.id() == LogicalTypeId::VARCHAR && + !StringType::GetCollation(child).empty()); + }); +} + +static unique_ptr NestedValueExpression(const LogicalType &type, optional_ptr value) { + if (type.IsAggregateState() || type.id() == LogicalTypeId::GEOMETRY || type.id() == LogicalTypeId::TYPE || + (type.id() != LogicalTypeId::TUPLE && TypeExpression::CanRepresent(type) && + !ConstantExpression::RequiresTypeWitness(type))) { + return ConstantExpression::FromValue(value ? value->WithType(type) : Value(type)); + } + const bool empty_list = + value && !value->IsNull() && type.id() == LogicalTypeId::LIST && ListValue::GetChildren(*value).empty(); + if (value && (value->IsNull() || empty_list)) { + auto result = make_uniq(); + result->CaseChecksMutable().push_back( + {ConstantExpression::Boolean(false), NestedValueExpression(type, nullptr)}); + result->ElseMutable() = empty_list ? ListValueExpression({}) : ConstantExpression::Null(); + return std::move(result); + } + if (type.id() == LogicalTypeId::MAP) { + vector keys, values; + if (value) { + for (auto &entry : MapValue::GetChildren(*value)) { + auto &children = StructValue::GetChildren(entry); + keys.push_back(children[0]); + values.push_back(children[1]); + } + } + vector> arguments; + arguments.push_back(ConstantExpression::FromValue(Value::LIST(MapType::KeyType(type), std::move(keys)))); + arguments.push_back(ConstantExpression::FromValue(Value::LIST(MapType::ValueType(type), std::move(values)))); + auto result = ValueFunction("map", std::move(arguments)); + return type.HasAlias() ? CastTo(type, std::move(result)) : std::move(result); + } + vector arguments; + child_list_t child_types; + optional_ptr> children; + string function; + switch (type.id()) { + case LogicalTypeId::TUPLE: + case LogicalTypeId::STRUCT: + function = type.id() == LogicalTypeId::TUPLE || StructType::IsUnnamed(type) ? "row" : "struct_pack"; + child_types = StructType::GetChildTypes(type); + if (value) { + children = StructValue::GetChildren(*value); + } + break; + case LogicalTypeId::LIST: + function = "list_value"; + if (value) { + children = ListValue::GetChildren(*value); + } + child_types.resize(children ? children->size() : 1, {Identifier(), ListType::GetChildType(type)}); + break; + case LogicalTypeId::ARRAY: + function = "array_value"; + if (value) { + children = ArrayValue::GetChildren(*value); + } + child_types.resize(ArrayType::GetSize(type), {Identifier(), ArrayType::GetChildType(type)}); + break; + default: + throw NotImplementedException("The nested value type has no SQL constructor"); + } + for (idx_t i = 0; i < child_types.size(); i++) { + auto child = NestedValueExpression(child_types[i].second, children ? &(*children)[i] : nullptr); + auto name = function == "struct_pack" ? child_types[i].first : Identifier(); + child->SetAlias(name); + arguments.emplace_back(name, std::move(child)); + } + unique_ptr result = + make_uniq(QualifiedName("system", "main", Identifier(function)), std::move(arguments)); + return type.HasAlias() ? CastTo(type, std::move(result)) : std::move(result); } unique_ptr ConstantExpression::FromValue(const Value &value) { auto &type = value.type(); + if (type.IsAggregateState()) { + auto storage_type = type.WithAlias("").WithExtensionInfo(nullptr); + Vector source(value, count_t(1)); + Vector storage(storage_type, 1); + storage.Reinterpret(source); + auto result = ExportAggregateFunction::StateToSQL(type, FromValue(storage.GetValue(0))); + if (!result) { + throw NotImplementedException("Aggregate state SQL parameters are not representable"); + } + return result; + } + if (type.id() == LogicalTypeId::VARCHAR && !StringType::GetCollation(type).empty()) { + auto result = FromValue(value.WithType(LogicalType::VARCHAR)); + if (type.HasAlias()) { + result = CastTo(type, std::move(result)); + } + return make_uniq(StringType::GetCollation(type), std::move(result)); + } + if (type.id() == LogicalTypeId::TYPE) { + if (value.IsNull()) { + return CastTo(type, Null()); + } + vector> arguments; + arguments.push_back(FromValue(Value(TypeValue::GetType(value)))); + return ValueFunction("get_type", std::move(arguments)); + } + if (type.id() == LogicalTypeId::GEOMETRY) { + auto result = GeometryExpression(value); + return type.HasAlias() ? CastTo(type, std::move(result)) : std::move(result); + } + if (TypeVisitor::Contains(type, [](const LogicalType &child) { + return child.id() == LogicalTypeId::ENUM && EnumType::GetSize(child) == 0; + })) { + throw NotImplementedException("The nested value type has no SQL constructor"); + } + if (RequiresTypeWitness(type) || type.id() == LogicalTypeId::TUPLE) { + return NestedValueExpression(type, value); + } + if (!value.IsNull() && type.id() == LogicalTypeId::INTERVAL) { + auto result = IntervalExpression(IntervalValue::Get(value)); + return type.HasAlias() ? CastTo(type, std::move(result)) : std::move(result); + } if (value.IsNull()) { if (type.id() == LogicalTypeId::SQLNULL) { return Null(); @@ -172,6 +374,11 @@ unique_ptr ConstantExpression::FromValue(const Value &value) { case LogicalTypeId::STRUCT: return StructExpression(value); case LogicalTypeId::VARIANT: { + Vector vector(value, count_t(1)); + VariantIterator iterator(vector); + if (HasUnsupportedVariantKeys(iterator.Root(0))) { + throw NotImplementedException("VARIANT object keys cannot be represented by a struct literal"); + } auto payload = VariantValue::GetValue(value); return CastTo(type, CastTo(payload.type(), FromValue(payload))); } @@ -187,12 +394,11 @@ unique_ptr ConstantExpression::FromValue(const Value &value) { if (children.empty()) { return CastTo(type, ListValueExpression({})); } - return CastTo(type, make_uniq("array_value", ChildExpressions(children))); + return CastTo(type, make_uniq(QualifiedName("system", "main", "array_value"), + ChildExpressions(children))); } case LogicalTypeId::MAP: return MapExpression(value); - case LogicalTypeId::TYPE: - return TypeExpression::FromLogicalType(TypeValue::GetType(value)); case LogicalTypeId::POINTER: return FromLiteral(Literal::Pointer(value.GetPointer())); case LogicalTypeId::UNION: { diff --git a/src/duckdb/src/parser/peg/transformer/transform_explain.cpp b/src/duckdb/src/parser/peg/transformer/transform_explain.cpp index ab65b9d7a..37224d83c 100644 --- a/src/duckdb/src/parser/peg/transformer/transform_explain.cpp +++ b/src/duckdb/src/parser/peg/transformer/transform_explain.cpp @@ -18,6 +18,7 @@ unique_ptr PEGTransformerFactory::TransformExplainStatement( const optional> &explain_option_list, unique_ptr explainable_statements) { auto explain_type = analyze_keyword ? ExplainType::EXPLAIN_ANALYZE : ExplainType::EXPLAIN_STANDARD; bool format_is_set = false; + bool sql_is_set = false; auto format = ProfilerPrintFormat::Default(); if (explain_option_list) { for (auto option : *explain_option_list) { @@ -38,6 +39,11 @@ unique_ptr PEGTransformerFactory::TransformExplainStatement( format = ParseProfilerPrintFormat(option.children[0]); } format_is_set = true; + } else if (option_name == "sql") { + if (sql_is_set || !option.children.empty() || option.expression) { + throw InvalidInputException("SQL must be provided once without arguments"); + } + sql_is_set = true; } else if (option_name == "analyze") { explain_type = ExplainType::EXPLAIN_ANALYZE; } else { @@ -45,6 +51,13 @@ unique_ptr PEGTransformerFactory::TransformExplainStatement( } } } + if (sql_is_set) { + if (format_is_set || explain_type == ExplainType::EXPLAIN_ANALYZE) { + throw InvalidInputException("EXPLAIN (SQL) cannot be combined with ANALYZE or FORMAT"); + } + transformer.PivotEntryCheck("EXPLAIN (SQL) statement"); + explain_type = ExplainType::EXPLAIN_SQL; + } auto statement = std::move(explainable_statements); return make_uniq(std::move(statement), explain_type, format); } diff --git a/src/duckdb/src/parser/statement/explain_statement.cpp b/src/duckdb/src/parser/statement/explain_statement.cpp index dc18e4296..4d1e0f2cd 100644 --- a/src/duckdb/src/parser/statement/explain_statement.cpp +++ b/src/duckdb/src/parser/statement/explain_statement.cpp @@ -10,7 +10,8 @@ ExplainStatement::ExplainStatement(unique_ptr stmt, ExplainType ex } ExplainStatement::ExplainStatement(const ExplainStatement &other) - : SQLStatement(other), stmt(other.stmt->Copy()), explain_type(other.explain_type), format(other.format) { + : SQLStatement(other), stmt(other.stmt->Copy()), explain_type(other.explain_type), + allow_unsupported_sql(other.allow_unsupported_sql), format(other.format) { } unique_ptr ExplainStatement::Copy() const { @@ -18,6 +19,9 @@ unique_ptr ExplainStatement::Copy() const { } string ExplainStatement::OptionsToString() const { + if (explain_type == ExplainType::EXPLAIN_SQL) { + return "(SQL)"; + } string options; if (explain_type == ExplainType::EXPLAIN_ANALYZE) { options += "("; diff --git a/src/duckdb/src/parser/statement/external_resource_statement.cpp b/src/duckdb/src/parser/statement/external_resource_statement.cpp index 3d67ed736..c0f90e494 100644 --- a/src/duckdb/src/parser/statement/external_resource_statement.cpp +++ b/src/duckdb/src/parser/statement/external_resource_statement.cpp @@ -33,12 +33,12 @@ string ExternalResourceStatement::ToString() const { result += " AS " + SQLIdentifier(name); } if (!options.empty()) { - vector stringified; - for (auto &opt : options) { - stringified.push_back( - StringUtil::Format("%s %s", SQLIdentifier(opt.first).ToString(opt.first), opt.second->ToString())); - } - result += " (" + StringUtil::Join(stringified, ", ") + ")"; + result += " ("; + result += StringUtil::Join(options, ", ", [](const auto &opt) { + return StringUtil::Format("%s %s", SQLIdentifier(opt.first).ToString(opt.first), + opt.second->ToString()); + }); + result += ")"; } break; case ExternalResourceOperation::REGISTER: diff --git a/src/duckdb/src/parser/statement/multi_statement.cpp b/src/duckdb/src/parser/statement/multi_statement.cpp index 2821084cd..bf4be6dfe 100644 --- a/src/duckdb/src/parser/statement/multi_statement.cpp +++ b/src/duckdb/src/parser/statement/multi_statement.cpp @@ -16,11 +16,7 @@ unique_ptr MultiStatement::Copy() const { } string MultiStatement::ToString() const { - vector stringified; - for (auto &stmt : statements) { - stringified.push_back(stmt->ToString()); - } - return StringUtil::Join(stringified, ";") + ";"; + return StringUtil::Join(statements, ";", [](const auto &stmt) { return stmt->ToString(); }) + ";"; } } // namespace duckdb diff --git a/src/duckdb/src/planner/binder/expression/bind_window_expression.cpp b/src/duckdb/src/planner/binder/expression/bind_window_expression.cpp index 4fd95ebea..91cebd084 100644 --- a/src/duckdb/src/planner/binder/expression/bind_window_expression.cpp +++ b/src/duckdb/src/planner/binder/expression/bind_window_expression.cpp @@ -47,7 +47,8 @@ static bool IsRangeType(const LogicalType &type) { } static LogicalType BindRangeExpression(ClientContext &context, const string &name, unique_ptr &bound, - unique_ptr &bound_order) { + unique_ptr &bound_order, + optional_ptr> origin = nullptr) { vector> children; D_ASSERT(bound_order); @@ -59,6 +60,8 @@ static LogicalType BindRangeExpression(ClientContext &context, const string &nam throw BinderException(error_context, "Window RANGE expressions cannot be NULL"); } children.emplace_back(std::move(bound)); + optional_ptr order_input = children[0].get(); + optional_ptr offset_input = children[1].get(); ErrorData error; FunctionBinder function_binder(context); @@ -72,6 +75,9 @@ static LogicalType BindRangeExpression(ClientContext &context, const string &nam if (!IsRangeType(function->GetReturnType())) { throw BinderException(error_context, "Invalid type for Window RANGE expression"); } + if (origin) { + *origin = WindowRangeBoundary::Capture(*function, order_input, offset_input); + } bound = std::move(function); return bound->GetReturnType(); } @@ -315,13 +321,15 @@ BindResult BaseSelectBinder::BindWindowExpression(WindowExpression &window, idx_ D_ASSERT(window.OrderBy().size() == 1); range_sense = config.ResolveOrder(context, window.OrderByMutable()[0].type); const auto range_name = (range_sense == OrderType::ASCENDING) ? "-" : "+"; - start_type = BindRangeExpression(context, range_name, bound_start, bound_orders[0]); + start_type = BindRangeExpression(context, range_name, bound_start, bound_orders[0], + &result->SQLRangeStartBoundaryMutable()); } else if (window.WindowStart() == WindowBoundary::EXPR_FOLLOWING_RANGE) { D_ASSERT(window.OrderBy().size() == 1); range_sense = config.ResolveOrder(context, window.OrderByMutable()[0].type); const auto range_name = (range_sense == OrderType::ASCENDING) ? "+" : "-"; - start_type = BindRangeExpression(context, range_name, bound_start, bound_orders[0]); + start_type = BindRangeExpression(context, range_name, bound_start, bound_orders[0], + &result->SQLRangeStartBoundaryMutable()); } LogicalType end_type = LogicalType::BIGINT; @@ -329,13 +337,24 @@ BindResult BaseSelectBinder::BindWindowExpression(WindowExpression &window, idx_ D_ASSERT(window.OrderBy().size() == 1); range_sense = config.ResolveOrder(context, window.OrderByMutable()[0].type); const auto range_name = (range_sense == OrderType::ASCENDING) ? "-" : "+"; - end_type = BindRangeExpression(context, range_name, bound_end, bound_orders[0]); + end_type = + BindRangeExpression(context, range_name, bound_end, bound_orders[0], &result->SQLRangeEndBoundaryMutable()); } else if (window.WindowEnd() == WindowBoundary::EXPR_FOLLOWING_RANGE) { D_ASSERT(window.OrderBy().size() == 1); range_sense = config.ResolveOrder(context, window.OrderByMutable()[0].type); const auto range_name = (range_sense == OrderType::ASCENDING) ? "+" : "-"; - end_type = BindRangeExpression(context, range_name, bound_end, bound_orders[0]); + end_type = + BindRangeExpression(context, range_name, bound_end, bound_orders[0], &result->SQLRangeEndBoundaryMutable()); + } + + if (result->SQLRangeStartBoundary()) { + result->SQLRangeStartBoundaryMutable()->boundary = window.WindowStart(); + result->SQLRangeStartBoundaryMutable()->direction = range_sense; + } + if (result->SQLRangeEndBoundary()) { + result->SQLRangeEndBoundaryMutable()->boundary = window.WindowEnd(); + result->SQLRangeEndBoundaryMutable()->direction = range_sense; } // Cast ORDER and boundary expressions to the same type @@ -353,7 +372,12 @@ BindResult BaseSelectBinder::BindWindowExpression(WindowExpression &window, idx_ } // Cast all three to match + optional_ptr original_order = bound_order.get(); bound_order = BoundCastExpression::AddCastToType(context, std::move(bound_order), order_type); + if (!WindowRangeCast::Capture(*bound_order, original_order, result->SQLRangeOrderCastsMutable())) { + result->SQLRangeStartBoundaryMutable().reset(); + result->SQLRangeEndBoundaryMutable().reset(); + } start_type = end_type = order_type; } @@ -373,8 +397,20 @@ BindResult BaseSelectBinder::BindWindowExpression(WindowExpression &window, idx_ } result->FilterMutable() = CastWindowExpression(std::move(bound_filter), LogicalType::BOOLEAN); + optional_ptr original_start = bound_start.get(); + optional_ptr original_end = bound_end.get(); result->StartExprMutable() = CastWindowExpression(std::move(bound_start), start_type); result->EndExprMutable() = CastWindowExpression(std::move(bound_end), end_type); + if (result->SQLRangeStartBoundary() && + !WindowRangeCast::Capture(*result->StartExpr(), original_start, + result->SQLRangeStartBoundaryMutable()->result_casts)) { + result->SQLRangeStartBoundaryMutable().reset(); + } + if (result->SQLRangeEndBoundary() && + !WindowRangeCast::Capture(*result->EndExpr(), original_end, + result->SQLRangeEndBoundaryMutable()->result_casts)) { + result->SQLRangeEndBoundaryMutable().reset(); + } result->WindowStartMutable() = window.WindowStart(); result->WindowEndMutable() = window.WindowEnd(); result->WindowExcludeMutable() = window.WindowExclude(); diff --git a/src/duckdb/src/planner/binder/statement/bind_explain.cpp b/src/duckdb/src/planner/binder/statement/bind_explain.cpp index 6ed696907..4f9457dc3 100644 --- a/src/duckdb/src/planner/binder/statement/bind_explain.cpp +++ b/src/duckdb/src/planner/binder/statement/bind_explain.cpp @@ -7,13 +7,16 @@ namespace duckdb { BoundStatement Binder::Bind(ExplainStatement &stmt) { BoundStatement result; + if (stmt.explain_type == ExplainType::EXPLAIN_SQL && stmt.stmt->type != StatementType::SELECT_STATEMENT) { + throw NotImplementedException("EXPLAIN (SQL) supports SELECT, VALUES and WITH queries only"); + } // bind the underlying statement auto plan = Bind(*stmt.stmt); // render the unoptimized logical plan, but only when it will be shown: a plain EXPLAIN in a multi-plan format. // (it is unused for EXPLAIN ANALYZE, and single-plan formats like FORMAT WEB render only the final plan) string logical_plan_unopt; - if (stmt.explain_type != ExplainType::EXPLAIN_ANALYZE) { + if (stmt.explain_type == ExplainType::EXPLAIN_STANDARD) { auto renderer = TreeRenderer::CreateRenderer(context, stmt.format); if (!renderer || !renderer->RendersSinglePlan()) { logical_plan_unopt = plan.plan->ToString(context, stmt.format); @@ -21,6 +24,10 @@ BoundStatement Binder::Bind(ExplainStatement &stmt) { } auto explain = make_uniq(std::move(plan.plan), stmt.explain_type, stmt.format); explain->logical_plan_unopt = logical_plan_unopt; + if (stmt.explain_type == ExplainType::EXPLAIN_SQL) { + explain->sql_output_names = std::move(plan.names); + explain->allow_unsupported_sql = stmt.allow_unsupported_sql; + } result.plan = std::move(explain); result.names = {"explain_key", "explain_value"}; diff --git a/src/duckdb/src/planner/binder/tableref/bind_basetableref.cpp b/src/duckdb/src/planner/binder/tableref/bind_basetableref.cpp index bffa396db..7fc21e99e 100644 --- a/src/duckdb/src/planner/binder/tableref/bind_basetableref.cpp +++ b/src/duckdb/src/planner/binder/tableref/bind_basetableref.cpp @@ -285,6 +285,9 @@ BoundStatement Binder::Bind(BaseTableRef &ref) { auto logical_get = make_uniq(table_index, scan_function, std::move(bind_data), std::move(return_types), std::move(return_names), std::move(virtual_columns)); + if (entry_at_clause) { + logical_get->at_clause = make_uniq(entry_at_clause->Unit(), entry_at_clause->GetValue()); + } auto table_entry = logical_get->GetTable(); auto &col_ids = logical_get->GetMutableColumnIds(); if (!table_entry) { diff --git a/src/duckdb/src/planner/binder/tableref/plan_joinref.cpp b/src/duckdb/src/planner/binder/tableref/plan_joinref.cpp index baa5c7525..c6b18f6b4 100644 --- a/src/duckdb/src/planner/binder/tableref/plan_joinref.cpp +++ b/src/duckdb/src/planner/binder/tableref/plan_joinref.cpp @@ -365,6 +365,29 @@ unique_ptr LogicalComparisonJoin::CreateJoin(ClientContext &con std::move(conditions)); } +static void PlanMarkJoin(LogicalOperator &op) { + if (op.type == LogicalOperatorType::LOGICAL_COMPARISON_JOIN) { + auto &join = op.Cast(); + if (join.TryGetMarkJoinGroupTypes(join.mark_types)) { + return; + } + bool all_equal = !join.conditions.empty(); + bool all_null_safe = all_equal; + for (auto &condition : join.conditions) { + if (!condition.IsComparison()) { + all_equal = all_null_safe = false; + break; + } + all_equal &= condition.GetComparisonType() == ExpressionType::COMPARE_EQUAL; + all_null_safe &= condition.GetComparisonType() == ExpressionType::COMPARE_NOT_DISTINCT_FROM; + } + if (all_equal || all_null_safe || (join.conditions.size() == 1 && join.conditions[0].IsComparison())) { + return; + } + } + throw NotImplementedException("Unsupported explicit MARK join conditions"); +} + unique_ptr Binder::CreatePlan(BoundJoinRef &ref) { auto old_is_outside_flattened = is_outside_flattened; // Plan laterals from outermost to innermost @@ -384,6 +407,9 @@ unique_ptr Binder::CreatePlan(BoundJoinRef &ref) { } if (ref.lateral) { + if (ref.type == JoinType::MARK) { + throw NotImplementedException("Unsupported explicit MARK join conditions"); + } auto new_plan = PlanLateralJoin(std::move(left), std::move(right), ref.correlated_columns, ref.type, std::move(ref.condition)); return new_plan; @@ -420,6 +446,9 @@ unique_ptr Binder::CreatePlan(BoundJoinRef &ref) { return std::move(filter); } if (has_dependent_condition) { + if (ref.type == JoinType::MARK) { + throw NotImplementedException("Unsupported explicit MARK join conditions"); + } auto join = make_uniq(ref.type); join->children.push_back(std::move(left)); join->children.push_back(std::move(right)); @@ -439,6 +468,7 @@ unique_ptr Binder::CreatePlan(BoundJoinRef &ref) { if (ref.type == JoinType::MARK) { join->Cast().mark_index = ref.mark_index; + PlanMarkJoin(*join); } if (!ref.duplicate_eliminated_columns.empty()) { D_ASSERT(join->type == LogicalOperatorType::LOGICAL_COMPARISON_JOIN); diff --git a/src/duckdb/src/planner/bound_expression_sql_exporter.cpp b/src/duckdb/src/planner/bound_expression_sql_exporter.cpp deleted file mode 100644 index eb504ed13..000000000 --- a/src/duckdb/src/planner/bound_expression_sql_exporter.cpp +++ /dev/null @@ -1,805 +0,0 @@ -#include "duckdb/planner/bound_expression_sql_exporter.hpp" - -#include "duckdb/planner/sql_export_helpers.hpp" -#include "duckdb/common/types/variant_iterator.hpp" -#include "duckdb/parser/expression/between_expression.hpp" -#include "duckdb/parser/expression/case_expression.hpp" -#include "duckdb/parser/expression/cast_expression.hpp" -#include "duckdb/parser/expression/columnref_expression.hpp" -#include "duckdb/parser/expression/comparison_expression.hpp" -#include "duckdb/parser/expression/conjunction_expression.hpp" -#include "duckdb/parser/expression/constant_expression.hpp" -#include "duckdb/parser/expression/function_expression.hpp" -#include "duckdb/parser/expression/operator_expression.hpp" -#include "duckdb/planner/bound_result_modifier.hpp" -#include "duckdb/planner/expression/bound_aggregate_expression.hpp" -#include "duckdb/planner/expression/bound_between_expression.hpp" -#include "duckdb/planner/expression/bound_case_expression.hpp" -#include "duckdb/planner/expression/bound_cast_expression.hpp" -#include "duckdb/planner/expression/bound_columnref_expression.hpp" -#include "duckdb/planner/expression/bound_comparison_expression.hpp" -#include "duckdb/planner/expression/bound_conjunction_expression.hpp" -#include "duckdb/planner/expression/bound_constant_expression.hpp" -#include "duckdb/planner/expression/bound_function_expression.hpp" -#include "duckdb/planner/expression/bound_operator_expression.hpp" - -namespace duckdb { - -using BoundExpressionSQLExportResult = LogicalPlanVerificationResult>; - -using SQLExportHelpers::ChildPath; -using SQLExportHelpers::IsSQLRepresentableType; -using SQLExportHelpers::IsValidIdentifier; - -static bool HasUnsupportedVariantKeys(const VariantNode &node) { - // Struct literals require nonempty, case-insensitively unique keys; inspect before the lossy STRUCT conversion. - if (node.GetTypeId() == VariantLogicalType::OBJECT) { - identifier_set_t names; - for (auto &child : node.GetObjectChildren()) { - if (child.key.GetSize() == 0 || !names.insert(Identifier(child.key.GetString())).second || - HasUnsupportedVariantKeys(child.value)) { - return true; - } - } - } else if (node.GetTypeId() == VariantLogicalType::ARRAY) { - for (auto child : node.GetArrayChildren()) { - if (HasUnsupportedVariantKeys(child)) { - return true; - } - } - } - return false; -} - -static bool HasUnsupportedVariantKeys(const Value &value) { - if (value.IsNull()) { - return false; - } - optional_ptr> children; - switch (value.type().id()) { - case LogicalTypeId::VARIANT: { - Vector vector(value, count_t(1)); - VariantIterator iterator(vector); - return HasUnsupportedVariantKeys(iterator.Root(0)); - } - case LogicalTypeId::STRUCT: - children = StructValue::GetChildren(value); - break; - case LogicalTypeId::LIST: - children = ListValue::GetChildren(value); - break; - case LogicalTypeId::ARRAY: - children = ArrayValue::GetChildren(value); - break; - case LogicalTypeId::MAP: - children = MapValue::GetChildren(value); - break; - case LogicalTypeId::UNION: - return HasUnsupportedVariantKeys(UnionValue::GetValue(value)); - default: - return false; - } - for (auto &child : *children) { - if (HasUnsupportedVariantKeys(child)) { - return true; - } - } - return false; -} - -static bool IsExpressionRootPath(const LogicalPlanVerificationPath &path) { - if (!path.IsValid()) { - return false; - } - if (path.root == LogicalPlanVerificationPathRoot::STANDALONE_EXPRESSION) { - return true; - } - for (auto &component : path.components) { - if (component.type == LogicalPlanVerificationPathComponentType::OPERATOR_EXPRESSION) { - return true; - } - } - return false; -} - -static LogicalPlanVerificationIssue -InternalInvariant(optional path, string message, - optional construct = {}) { - return SQLExportHelpers::MakeIssue(LogicalPlanVerificationIssueCode::INTERNAL_INVARIANT, - LogicalPlanVerificationPhase::EXPRESSION_EXPORT, std::move(path), - std::move(construct), std::move(message)); -} - -static LogicalPlanVerificationIssue InternalExpressionInvariant(const LogicalPlanVerificationPath &path, - const Expression &expression, string message) { - return InternalInvariant(path, std::move(message), - LogicalPlanVerificationConstructIdentity::Expression(expression.GetExpressionClass())); -} - -static LogicalPlanVerificationIssue UnsupportedExpression(const LogicalPlanVerificationPath &path, - ExpressionClass expression_class) { - return SQLExportHelpers::MakeIssue( - LogicalPlanVerificationIssueCode::UNSUPPORTED_EXPRESSION, LogicalPlanVerificationPhase::EXPRESSION_EXPORT, path, - LogicalPlanVerificationConstructIdentity::Expression(expression_class), - "The bound expression class does not have a SQL AST representation in this exporter"); -} - -static LogicalPlanVerificationIssue UnsupportedFeature(const LogicalPlanVerificationPath &path, string feature, - string message) { - return SQLExportHelpers::MakeIssue( - LogicalPlanVerificationIssueCode::UNSUPPORTED_EXPORT_FEATURE, LogicalPlanVerificationPhase::EXPRESSION_EXPORT, - path, LogicalPlanVerificationConstructIdentity::ExportFeature(std::move(feature)), std::move(message)); -} - -static LogicalPlanVerificationIssue UnsupportedFunction(const LogicalPlanVerificationPath &path, - LogicalPlanVerificationFunctionIdentity identity, - string message) { - return SQLExportHelpers::MakeIssue( - LogicalPlanVerificationIssueCode::UNSUPPORTED_FUNCTION, LogicalPlanVerificationPhase::EXPRESSION_EXPORT, path, - LogicalPlanVerificationConstructIdentity::Function(std::move(identity)), std::move(message)); -} - -static BoundExpressionSQLExportResult Failure(LogicalPlanVerificationIssue issue) { - vector issues; - issues.push_back(std::move(issue)); - return BoundExpressionSQLExportResult::Failure(std::move(issues)); -} - -static bool ChildrenAreConsistentWithArguments(const vector> &children, - const vector &arguments) { - if (children.size() != arguments.size()) { - return false; - } - for (idx_t child_index = 0; child_index < children.size(); child_index++) { - if (!children[child_index] || !children[child_index]->GetReturnType().IsComplete() || - (arguments[child_index].IsComplete() && children[child_index]->GetReturnType() != arguments[child_index])) { - return false; - } - } - return true; -} - -template -static LogicalPlanVerificationFunctionIdentity DefinitionFunctionIdentity(const FUNCTION &definition, - const vector &arguments, - const LogicalType &return_type) { - LogicalPlanVerificationFunctionIdentity identity; - identity.catalog = definition.GetCatalogName().GetIdentifierName(); - identity.schema = definition.GetSchemaName().GetIdentifierName(); - identity.name = definition.GetName().GetIdentifierName(); - identity.arguments = arguments; - identity.return_type = return_type; - return identity; -} - -class BoundExpressionSQLExportState { -public: - explicit BoundExpressionSQLExportState(const BoundExpressionSQLExportContext &context_p) : context(context_p) { - } - - BoundExpressionSQLExportResult Export(const Expression &expression, const LogicalPlanVerificationPath &path) { - switch (expression.GetExpressionClass()) { - case ExpressionClass::BOUND_CONSTANT: - return ExportConstant(expression.Cast(), path); - case ExpressionClass::BOUND_COLUMN_REF: - return ExportColumnRef(expression.Cast(), path); - case ExpressionClass::BOUND_FUNCTION: - return ExportFunction(expression.Cast(), path); - case ExpressionClass::BOUND_CONJUNCTION: - return ExportConjunction(expression.Cast(), path); - case ExpressionClass::BOUND_CASE: - return ExportCase(expression.Cast(), path); - case ExpressionClass::BOUND_OPERATOR: - return ExportOperator(expression.Cast(), path); - case ExpressionClass::BOUND_AGGREGATE: - return ExportAggregate(expression.Cast(), path); - case ExpressionClass::BOUND_DEFAULT: - case ExpressionClass::BOUND_PARAMETER: - case ExpressionClass::BOUND_REF: - case ExpressionClass::BOUND_SUBQUERY: - case ExpressionClass::BOUND_WINDOW: - case ExpressionClass::BOUND_UNNEST: - case ExpressionClass::BOUND_LAMBDA: - case ExpressionClass::BOUND_LAMBDA_REF: - case ExpressionClass::LEGACY_BOUND_CAST: - case ExpressionClass::LEGACY_BOUND_COMPARISON: - case ExpressionClass::LEGACY_BOUND_BETWEEN: - return Failure(UnsupportedExpression(path, expression.GetExpressionClass())); - case ExpressionClass::BOUND_EXPANDED: - case ExpressionClass::AGGREGATE: - case ExpressionClass::CASE: - case ExpressionClass::CAST: - case ExpressionClass::COLUMN_REF: - case ExpressionClass::COMPARISON: - case ExpressionClass::CONJUNCTION: - case ExpressionClass::CONSTANT: - case ExpressionClass::DEFAULT: - case ExpressionClass::FUNCTION: - case ExpressionClass::OPERATOR: - case ExpressionClass::STAR: - case ExpressionClass::SUBQUERY: - case ExpressionClass::WINDOW: - case ExpressionClass::PARAMETER: - case ExpressionClass::COLLATE: - case ExpressionClass::LAMBDA: - case ExpressionClass::POSITIONAL_REFERENCE: - case ExpressionClass::BETWEEN: - case ExpressionClass::LAMBDA_REF: - case ExpressionClass::TYPE: - return Failure( - InternalExpressionInvariant(path, expression, "Expression export requires a final bound class")); - case ExpressionClass::INVALID: - return Failure(InternalInvariant(path, "Expression export received an invalid expression class")); - } - return Failure(InternalInvariant(path, "Expression export received an unknown expression class")); - } - -private: - BoundExpressionSQLExportResult ExportConstant(const BoundConstantExpression &expression, - const LogicalPlanVerificationPath &path) { - if (expression.GetExpressionType() != ExpressionType::VALUE_CONSTANT) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound constant has an invalid expression type")); - } - auto &return_type = expression.GetReturnType(); - auto &value = expression.GetValue(); - if (!IsSQLRepresentableType(return_type) || !IsSQLRepresentableType(value.type())) { - return Failure(InternalExpressionInvariant(path, expression, "Bound constant has an unexportable type")); - } - if (return_type != value.type()) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound constant value and return types differ")); - } - if (TypeVisitor::Contains(return_type, [](const LogicalType &type) { return type.IsAggregateState(); })) { - return Failure(UnsupportedFeature(path, "aggregate_state_literal", - "Aggregate-state values do not have a SQL literal representation")); - } - if (TypeVisitor::Contains(return_type, LogicalTypeId::VARIANT) && HasUnsupportedVariantKeys(value)) { - return Failure(UnsupportedFeature(path, "variant_literal", - "VARIANT object keys cannot be represented by a struct literal")); - } - auto result = ConstantExpression::FromValue(value); - if (return_type.id() != LogicalTypeId::SQLNULL) { - result = make_uniq(return_type, std::move(result)); - } - return BoundExpressionSQLExportResult::Success(std::move(result)); - } - - BoundExpressionSQLExportResult ExportColumnRef(const BoundColumnRefExpression &expression, - const LogicalPlanVerificationPath &path) { - if (expression.GetExpressionType() != ExpressionType::BOUND_COLUMN_REF) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound column reference has an invalid expression type")); - } - auto &binding = expression.Binding(); - if (!binding.table_index.IsValid() || !binding.column_index.IsValid()) { - return Failure(InvalidBinding(path, binding, "Bound column reference has an incomplete binding")); - } - if (!IsSQLRepresentableType(expression.GetReturnType())) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound column reference has an incomplete type")); - } - if (expression.Depth() != 0) { - auto issue = UnsupportedFeature(path, "correlated_column_reference", - "Correlated column references require an owning query export context"); - issue.facts.emplace_back("depth", Value::UBIGINT(expression.Depth())); - return Failure(std::move(issue)); - } - if (!context.resolve_binding) { - return Failure(InvalidBinding(path, binding, "No SQL column binding resolver was provided")); - } - auto resolved = context.resolve_binding(binding); - if (!resolved) { - return Failure(InvalidBinding(path, binding, "The SQL column binding resolver has no matching entry")); - } - if (resolved->names.empty()) { - return Failure(InvalidBinding(path, binding, "The resolved SQL column name is empty")); - } - for (auto &name : resolved->names) { - if (!IsValidIdentifier(name)) { - return Failure( - InvalidBinding(path, binding, "The resolved SQL column name contains an invalid identifier")); - } - } - if (!IsSQLRepresentableType(resolved->type)) { - return Failure(InternalExpressionInvariant(path, expression, "The resolved SQL column type is incomplete")); - } - if (resolved->type != expression.GetReturnType()) { - LogicalPlanVerificationIssue issue; - issue.code = LogicalPlanVerificationIssueCode::TYPE_MISMATCH; - issue.phase = LogicalPlanVerificationPhase::EXPRESSION_EXPORT; - issue.path = path; - issue.construct = LogicalPlanVerificationConstructIdentity::BindingTypeMismatch(resolved->type, - expression.GetReturnType()); - issue.message = "The resolved SQL column type differs from the bound expression type"; - return Failure(std::move(issue)); - } - return BoundExpressionSQLExportResult::Success(make_uniq(std::move(resolved->names))); - } - - LogicalPlanVerificationIssue InvalidBinding(const LogicalPlanVerificationPath &path, const ColumnBinding &binding, - string message) { - LogicalPlanVerificationIssue issue; - issue.code = LogicalPlanVerificationIssueCode::INVALID_BINDING; - issue.phase = LogicalPlanVerificationPhase::EXPRESSION_EXPORT; - issue.path = path; - issue.facts.emplace_back("column_index", Value::UBIGINT(binding.column_index.GetIndexUnsafe())); - issue.facts.emplace_back("table_index", Value::UBIGINT(binding.table_index.index)); - issue.message = std::move(message); - return issue; - } - - BoundExpressionSQLExportResult ExportFunction(const BoundFunctionExpression &expression, - const LogicalPlanVerificationPath &path) { - switch (expression.GetExpressionType()) { - case ExpressionType::BOUND_FUNCTION: - return ExportScalarFunction(expression, path); - case ExpressionType::OPERATOR_CAST: - return ExportCast(expression, path); - case ExpressionType::COMPARE_EQUAL: - case ExpressionType::COMPARE_NOTEQUAL: - case ExpressionType::COMPARE_LESSTHAN: - case ExpressionType::COMPARE_GREATERTHAN: - case ExpressionType::COMPARE_LESSTHANOREQUALTO: - case ExpressionType::COMPARE_GREATERTHANOREQUALTO: - case ExpressionType::COMPARE_DISTINCT_FROM: - case ExpressionType::COMPARE_NOT_DISTINCT_FROM: - return ExportComparison(expression, path); - case ExpressionType::COMPARE_BETWEEN: - return ExportBetween(expression, path); - default: - return Failure( - InternalExpressionInvariant(path, expression, "Bound function has an invalid expression type")); - } - } - - BoundExpressionSQLExportResult ExportCast(const BoundFunctionExpression &expression, - const LogicalPlanVerificationPath &path) { - if (expression.GetExpressionType() != ExpressionType::OPERATOR_CAST || expression.GetChildren().size() != 1 || - !expression.GetChildren()[0] || !BoundCastExpression::HasValidBindData(expression) || - !ChildrenAreConsistentWithArguments(expression.GetChildren(), expression.Function().GetArguments()) || - expression.GetReturnType() != expression.Function().GetReturnType() || - !IsSQLRepresentableType(expression.GetReturnType()) || - !IsSQLRepresentableType(expression.GetChildren()[0]->GetReturnType())) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound cast has malformed type, data, or arity")); - } - if (BoundCastExpression::IsDefaultCast(expression)) { - return Failure(UnsupportedFeature(path, "default_cast_binding", - "A default-only bound cast cannot be reconstructed through SQL binding")); - } - auto child = ExportChild(*expression.GetChildren()[0], path, 0); - if (child.HasError()) { - return child; - } - return BoundExpressionSQLExportResult::Success(make_uniq( - expression.GetReturnType(), std::move(child.GetValue()), BoundCastExpression::IsTryCast(expression))); - } - - BoundExpressionSQLExportResult ExportComparison(const BoundFunctionExpression &expression, - const LogicalPlanVerificationPath &path) { - if (!BoundComparisonExpression::IsComparison(expression.GetExpressionType()) || - expression.GetReturnType() != LogicalType::BOOLEAN || expression.GetChildren().size() != 2 || - !expression.GetChildren()[0] || !expression.GetChildren()[1] || expression.BindInfo() || - !ChildrenAreConsistentWithArguments(expression.GetChildren(), expression.Function().GetArguments()) || - expression.GetReturnType() != expression.Function().GetReturnType() || - !IsSQLRepresentableType(expression.GetChildren()[0]->GetReturnType()) || - expression.GetChildren()[0]->GetReturnType() != expression.GetChildren()[1]->GetReturnType()) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound comparison has malformed type or arity")); - } - vector> children; - vector issues; - ExportChildren(expression.GetChildren(), path, children, issues); - if (!issues.empty()) { - return BoundExpressionSQLExportResult::Failure(std::move(issues)); - } - return BoundExpressionSQLExportResult::Success(make_uniq( - expression.GetExpressionType(), std::move(children[0]), std::move(children[1]))); - } - - BoundExpressionSQLExportResult ExportBetween(const BoundFunctionExpression &expression, - const LogicalPlanVerificationPath &path) { - if (expression.GetExpressionType() != ExpressionType::COMPARE_BETWEEN || - expression.GetReturnType() != LogicalType::BOOLEAN || expression.GetChildren().size() != 3 || - !expression.GetChildren()[0] || !expression.GetChildren()[1] || !expression.GetChildren()[2] || - !BoundBetweenExpression::HasValidBindData(expression) || - !ChildrenAreConsistentWithArguments(expression.GetChildren(), expression.Function().GetArguments()) || - expression.GetReturnType() != expression.Function().GetReturnType() || - !IsSQLRepresentableType(expression.GetChildren()[0]->GetReturnType()) || - expression.GetChildren()[0]->GetReturnType() != expression.GetChildren()[1]->GetReturnType() || - expression.GetChildren()[0]->GetReturnType() != expression.GetChildren()[2]->GetReturnType()) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound BETWEEN has malformed type, data, or arity")); - } - vector> children; - vector issues; - ExportChildren(expression.GetChildren(), path, children, issues); - if (!issues.empty()) { - return BoundExpressionSQLExportResult::Failure(std::move(issues)); - } - auto lower_inclusive = BoundBetweenExpression::LowerInclusive(expression); - auto upper_inclusive = BoundBetweenExpression::UpperInclusive(expression); - if (lower_inclusive && upper_inclusive) { - return BoundExpressionSQLExportResult::Success( - make_uniq(std::move(children[0]), std::move(children[1]), std::move(children[2]))); - } - if (expression.GetChildren()[0]->IsVolatile()) { - return Failure(UnsupportedFeature( - path, "exclusive_between_input_evaluation", - "An exclusive BETWEEN cannot duplicate a volatile input while preserving evaluation semantics")); - } - auto lower = make_uniq(BoundBetweenExpression::LowerComparisonType(expression), - children[0]->Copy(), std::move(children[1])); - auto upper = make_uniq(BoundBetweenExpression::UpperComparisonType(expression), - std::move(children[0]), std::move(children[2])); - return BoundExpressionSQLExportResult::Success( - make_uniq(ExpressionType::CONJUNCTION_AND, std::move(lower), std::move(upper))); - } - - BoundExpressionSQLExportResult ExportConjunction(const BoundConjunctionExpression &expression, - const LogicalPlanVerificationPath &path) { - if ((expression.GetExpressionType() != ExpressionType::CONJUNCTION_AND && - expression.GetExpressionType() != ExpressionType::CONJUNCTION_OR) || - expression.GetReturnType() != LogicalType::BOOLEAN || expression.GetChildren().size() < 2) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound conjunction has malformed type or arity")); - } - vector> children; - vector issues; - ExportChildren(expression.GetChildren(), path, children, issues, LogicalType::BOOLEAN); - if (!issues.empty()) { - return BoundExpressionSQLExportResult::Failure(std::move(issues)); - } - auto result = make_uniq(expression.GetExpressionType()); - result->GetChildrenMutable() = std::move(children); - return BoundExpressionSQLExportResult::Success(std::move(result)); - } - - BoundExpressionSQLExportResult ExportCase(const BoundCaseExpression &expression, - const LogicalPlanVerificationPath &path) { - if (expression.GetExpressionType() != ExpressionType::CASE_EXPR) { - return Failure(InternalExpressionInvariant(path, expression, "Bound CASE has an invalid expression type")); - } - if (!IsSQLRepresentableType(expression.GetReturnType()) || expression.CaseChecks().empty()) { - return Failure(InternalExpressionInvariant(path, expression, "Bound CASE has malformed type or arity")); - } - vector> source_children; - vector> expected_types; - for (auto &check : expression.CaseChecks()) { - source_children.push_back(check.when_expr.get()); - expected_types.push_back(LogicalType::BOOLEAN); - source_children.push_back(check.then_expr.get()); - expected_types.push_back(expression.GetReturnType()); - } - source_children.push_back(expression.ElseExpression().get()); - expected_types.push_back(expression.GetReturnType()); - - vector> children; - vector issues; - ExportChildren(source_children, path, children, issues, expected_types); - if (!issues.empty()) { - return BoundExpressionSQLExportResult::Failure(std::move(issues)); - } - auto result = make_uniq(); - for (idx_t check_index = 0; check_index < expression.CaseChecks().size(); check_index++) { - CaseCheck check; - check.when_expr = std::move(children[check_index * 2]); - check.then_expr = std::move(children[check_index * 2 + 1]); - result->CaseChecksMutable().push_back(std::move(check)); - } - result->ElseMutable() = std::move(children.back()); - return BoundExpressionSQLExportResult::Success(std::move(result)); - } - - BoundExpressionSQLExportResult ExportOperator(const BoundOperatorExpression &expression, - const LogicalPlanVerificationPath &path) { - optional expected_type; - auto child_count = expression.GetChildren().size(); - switch (expression.GetExpressionType()) { - case ExpressionType::OPERATOR_NOT: - if (child_count != 1 || expression.GetReturnType() != LogicalType::BOOLEAN) { - return Failure(InternalExpressionInvariant(path, expression, "Bound NOT has malformed type or arity")); - } - expected_type = LogicalType::BOOLEAN; - break; - case ExpressionType::OPERATOR_IS_NULL: - case ExpressionType::OPERATOR_IS_NOT_NULL: - if (child_count != 1 || expression.GetReturnType() != LogicalType::BOOLEAN) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound NULL test has malformed type or arity")); - } - break; - case ExpressionType::COMPARE_IN: - case ExpressionType::COMPARE_NOT_IN: - if (child_count < 2 || expression.GetReturnType() != LogicalType::BOOLEAN || !expression.GetChildren()[0]) { - return Failure(InternalExpressionInvariant(path, expression, "Bound IN has malformed type or arity")); - } - expected_type = expression.GetChildren()[0]->GetReturnType(); - break; - case ExpressionType::OPERATOR_COALESCE: - if (child_count < 2 || !IsSQLRepresentableType(expression.GetReturnType())) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound COALESCE has malformed type or arity")); - } - expected_type = expression.GetReturnType(); - break; - case ExpressionType::OPERATOR_TRY: - if (child_count != 1 || !expression.GetChildren()[0] || - !IsSQLRepresentableType(expression.GetReturnType())) { - return Failure(InternalExpressionInvariant(path, expression, "Bound TRY has malformed type or arity")); - } - if (expression.GetChildren()[0]->IsVolatile()) { - return Failure(UnsupportedFeature(path, "try_volatile_child", - "TRY cannot be rebound around a volatile expression")); - } - expected_type = expression.GetReturnType(); - break; - case ExpressionType::OPERATOR_UNPACK: - case ExpressionType::OPERATOR_NULLIF: - case ExpressionType::GROUPING_FUNCTION: - case ExpressionType::ARRAY_EXTRACT: - case ExpressionType::ARRAY_SLICE: - case ExpressionType::STRUCT_EXTRACT: - case ExpressionType::ARRAY_CONSTRUCTOR: - case ExpressionType::ARROW: - return Failure( - UnsupportedFeature(path, "bound_operator", "The bound operator has no admitted parsed SQL AST form")); - default: - return Failure( - InternalExpressionInvariant(path, expression, "Bound operator has an invalid expression type")); - } - vector> children; - vector issues; - ExportChildren(expression.GetChildren(), path, children, issues, expected_type); - if (!issues.empty()) { - return BoundExpressionSQLExportResult::Failure(std::move(issues)); - } - return BoundExpressionSQLExportResult::Success( - make_uniq(expression.GetExpressionType(), std::move(children))); - } - - BoundExpressionSQLExportResult ExportScalarFunction(const BoundFunctionExpression &expression, - const LogicalPlanVerificationPath &path) { - auto &function = expression.Function(); - for (auto &child : expression.GetChildren()) { - if (!child) { - return Failure(InternalExpressionInvariant(path, expression, "Bound scalar function has a null child")); - } - } - auto &definition = function.GetDefinition(); - if (!definition) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound scalar function has no retained definition")); - } - auto identity = - DefinitionFunctionIdentity(*definition, function.GetLogicalArguments(), function.GetLogicalReturnType()); - if (!identity.IsValid()) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound scalar function identity is incomplete")); - } - QualifiedName name(definition->GetCatalogName(), definition->GetSchemaName(), definition->GetName()); - if (!IsValidIdentifier(name.Catalog()) || !IsValidIdentifier(name.Schema()) || - !IsValidIdentifier(name.Name()) || !IsSQLRepresentableType(expression.GetReturnType())) { - return Failure(UnsupportedFunction(path, std::move(identity), - "The retained scalar function definition is not representable as SQL")); - } - for (auto &argument : identity.arguments) { - if (!IsSQLRepresentableType(argument)) { - return Failure(UnsupportedFunction(path, std::move(identity), - "The bound scalar function uses an internal argument type")); - } - } - if (definition->GetProperties().GetCaptureArgumentAliases() || - definition->GetProperties().RequiresExpressionNames()) { - return Failure(UnsupportedFunction( - path, std::move(identity), - "The bound scalar function requires expression names that SQL export does not preserve")); - } - if (!ChildrenAreConsistentWithArguments(expression.GetChildren(), function.GetArguments()) || - expression.GetReturnType() != function.GetReturnType() || - !ChildrenAreConsistentWithArguments(expression.GetChildren(), function.GetLogicalArguments()) || - expression.GetReturnType() != function.GetLogicalReturnType()) { - return Failure(UnsupportedFunction( - path, std::move(identity), "The bound scalar function no longer represents its logical SQL signature")); - } - vector> children; - vector issues; - ExportChildren(expression.GetChildren(), path, children, issues); - if (!issues.empty()) { - return BoundExpressionSQLExportResult::Failure(std::move(issues)); - } - unique_ptr result = - make_uniq(name, std::move(children), nullptr, nullptr, false, false, false); - // Aggregate-state schemas are inferred from the arguments, not expressible as SQL cast targets. - if (!function.GetLogicalReturnType().IsAggregateState() && definition->HasBindCallback() && - definition->GetReturnType() != function.GetLogicalReturnType()) { - result = make_uniq(function.GetLogicalReturnType(), std::move(result)); - } - return BoundExpressionSQLExportResult::Success(std::move(result)); - } - - BoundExpressionSQLExportResult ExportAggregate(const BoundAggregateExpression &expression, - const LogicalPlanVerificationPath &path) { - for (auto &child : expression.GetChildren()) { - if (!child) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound aggregate function has a null child")); - } - } - if (expression.GetExpressionType() != ExpressionType::BOUND_AGGREGATE) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound aggregate has an invalid expression type")); - } - auto &function = expression.Function(); - auto &definition = function.GetDefinition(); - if (!definition) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound aggregate function has no retained definition")); - } - auto identity = - DefinitionFunctionIdentity(*definition, function.GetLogicalArguments(), function.GetLogicalReturnType()); - if (!identity.IsValid()) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound aggregate function identity is incomplete")); - } - QualifiedName name(definition->GetCatalogName(), definition->GetSchemaName(), definition->GetName()); - if (!IsValidIdentifier(name.Catalog()) || !IsValidIdentifier(name.Schema()) || - !IsValidIdentifier(name.Name()) || !IsSQLRepresentableType(expression.GetReturnType())) { - return Failure(UnsupportedFunction( - path, std::move(identity), "The retained aggregate function definition is not representable as SQL")); - } - if (expression.GetAggregateType() != AggregateType::NON_DISTINCT && - expression.GetAggregateType() != AggregateType::DISTINCT) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound aggregate has an invalid distinct mode")); - } - if (expression.StateExportMode() != AggregateStateExportMode::NONE && - expression.StateExportMode() != AggregateStateExportMode::STATE_EXPORT) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound aggregate has an invalid state export mode")); - } - for (auto &argument : identity.arguments) { - if (!IsSQLRepresentableType(argument)) { - return Failure(UnsupportedFunction(path, std::move(identity), - "The bound aggregate uses an internal argument type")); - } - } - if (definition->GetProperties().GetCaptureArgumentAliases() || - definition->GetProperties().RequiresExpressionNames()) { - return Failure( - UnsupportedFunction(path, std::move(identity), - "The bound aggregate requires expression names that SQL export does not preserve")); - } - if (!ChildrenAreConsistentWithArguments(expression.GetChildren(), function.GetLogicalArguments())) { - return Failure(UnsupportedFunction(path, std::move(identity), - "The bound aggregate no longer represents its logical SQL arguments")); - } - if (expression.StateExportMode() == AggregateStateExportMode::NONE && - expression.GetReturnType() != function.GetLogicalReturnType()) { - return Failure(UnsupportedFunction(path, std::move(identity), - "The bound aggregate no longer represents its logical SQL return type")); - } - vector> source_children; - vector> expected_types; - for (idx_t child_index = 0; child_index < expression.GetChildren().size(); child_index++) { - source_children.push_back(expression.GetChildren()[child_index].get()); - expected_types.push_back(identity.arguments[child_index]); - } - if (expression.GetFilter()) { - source_children.push_back(expression.GetFilter().get()); - expected_types.push_back(LogicalType::BOOLEAN); - } - if (expression.GetOrderBys()) { - for (auto &order : expression.GetOrderBys()->orders) { - if ((order.type != OrderType::ASCENDING && order.type != OrderType::DESCENDING) || - (order.null_order != OrderByNullType::NULLS_FIRST && - order.null_order != OrderByNullType::NULLS_LAST)) { - return Failure( - InternalExpressionInvariant(path, expression, "Bound aggregate has an invalid ordering mode")); - } - source_children.push_back(order.expression.get()); - expected_types.emplace_back(); - } - } - - vector> children; - vector issues; - ExportChildren(source_children, path, children, issues, expected_types); - if (!issues.empty()) { - return BoundExpressionSQLExportResult::Failure(std::move(issues)); - } - vector> arguments; - for (idx_t child_index = 0; child_index < expression.GetChildren().size(); child_index++) { - arguments.push_back(std::move(children[child_index])); - } - idx_t child_index = expression.GetChildren().size(); - unique_ptr filter; - if (expression.GetFilter()) { - filter = std::move(children[child_index++]); - } - unique_ptr order_bys; - if (expression.GetOrderBys()) { - order_bys = make_uniq(); - for (auto &order : expression.GetOrderBys()->orders) { - order_bys->orders.emplace_back(order.type, order.null_order, std::move(children[child_index++])); - } - } - unique_ptr result = make_uniq( - name, std::move(arguments), std::move(filter), std::move(order_bys), expression.IsDistinct(), false, - expression.StateExportMode() == AggregateStateExportMode::STATE_EXPORT); - if (expression.StateExportMode() == AggregateStateExportMode::NONE && - !function.GetLogicalReturnType().IsAggregateState() && definition->HasBindCallback() && - definition->GetReturnType() != function.GetLogicalReturnType()) { - result = make_uniq(function.GetLogicalReturnType(), std::move(result)); - } - return BoundExpressionSQLExportResult::Success(std::move(result)); - } - - BoundExpressionSQLExportResult ExportChild(const Expression &expression, const LogicalPlanVerificationPath &path, - idx_t child_index) { - return Export(expression, ChildPath(path, child_index)); - } - - void ExportChildren(const vector> &source, const LogicalPlanVerificationPath &path, - vector> &result, vector &issues, - const optional &expected_type = {}) { - vector> source_refs; - vector> expected_types; - for (auto &child : source) { - source_refs.push_back(child.get()); - expected_types.push_back(expected_type); - } - ExportChildren(source_refs, path, result, issues, expected_types); - } - - void ExportChildren(const vector> &source, const LogicalPlanVerificationPath &path, - vector> &result, vector &issues, - const vector> &expected_types) { - D_ASSERT(source.size() == expected_types.size()); - result.resize(source.size()); - for (idx_t child_index = 0; child_index < source.size(); child_index++) { - auto child_path = ChildPath(path, child_index); - if (!source[child_index]) { - issues.push_back(InternalInvariant(child_path, "Bound expression has a null child")); - continue; - } - if (expected_types[child_index].has_value() && - source[child_index]->GetReturnType() != expected_types[child_index].value()) { - issues.push_back(InternalExpressionInvariant(child_path, *source[child_index], - "Bound expression child has an unexpected type")); - continue; - } - auto child = Export(*source[child_index], child_path); - if (child.HasError()) { - for (auto &issue : child.GetIssues()) { - issues.push_back(issue); - } - } else { - result[child_index] = std::move(child.GetValue()); - } - } - } - -private: - const BoundExpressionSQLExportContext &context; -}; - -LogicalPlanVerificationResult> -BoundExpressionSQLExporter::Export(const Expression &expression, const BoundExpressionSQLExportContext &context) { - LogicalPlanVerificationPath path; - path.root = LogicalPlanVerificationPathRoot::STANDALONE_EXPRESSION; - return ExportAtPath(expression, context, path); -} - -LogicalPlanVerificationResult> -BoundExpressionSQLExporter::ExportAtPath(const Expression &expression, const BoundExpressionSQLExportContext &context, - const LogicalPlanVerificationPath &path) { - if (!IsExpressionRootPath(path)) { - return Failure(InternalInvariant({}, "Expression export requires an already-valid expression root path")); - } - BoundExpressionSQLExportState state(context); - return state.Export(expression, path); -} - -} // namespace duckdb diff --git a/src/duckdb/src/planner/expression/bound_window_expression.cpp b/src/duckdb/src/planner/expression/bound_window_expression.cpp index 32918f6ab..1771cc977 100644 --- a/src/duckdb/src/planner/expression/bound_window_expression.cpp +++ b/src/duckdb/src/planner/expression/bound_window_expression.cpp @@ -211,6 +211,11 @@ unique_ptr BoundWindowExpression::Copy() const { new_window->sql_range_start = sql_range_start ? sql_range_start->Copy() : nullptr; new_window->sql_range_end = sql_range_end ? sql_range_end->Copy() : nullptr; new_window->sql_range_order_type = sql_range_order_type; + new_window->sql_range_start_boundary = + sql_range_start_boundary ? make_uniq(*sql_range_start_boundary) : nullptr; + new_window->sql_range_end_boundary = + sql_range_end_boundary ? make_uniq(*sql_range_end_boundary) : nullptr; + new_window->sql_range_order_casts = sql_range_order_casts; new_window->ignore_nulls = ignore_nulls; new_window->distinct = distinct; @@ -293,6 +298,11 @@ void BoundWindowExpression::Serialize(Serializer &serializer) const { serializer.WritePropertyWithDefault(216, "sql_range_end", sql_range_end, unique_ptr()); serializer.WritePropertyWithDefault(217, "sql_range_order_type", sql_range_order_type, LogicalType::INVALID); + serializer.WritePropertyWithDefault(218, "sql_range_start_boundary", sql_range_start_boundary, + unique_ptr()); + serializer.WritePropertyWithDefault(219, "sql_range_end_boundary", sql_range_end_boundary, + unique_ptr()); + serializer.WritePropertyWithDefault>(220, "sql_range_order_casts", sql_range_order_casts); } unique_ptr BoundWindowExpression::Deserialize(Deserializer &deserializer) { @@ -344,6 +354,12 @@ unique_ptr BoundWindowExpression::Deserialize(Deserializer &deserial unique_ptr()); deserializer.ReadPropertyWithExplicitDefault(217, "sql_range_order_type", result->sql_range_order_type, LogicalType::INVALID); + deserializer.ReadPropertyWithExplicitDefault(218, "sql_range_start_boundary", result->sql_range_start_boundary, + unique_ptr()); + deserializer.ReadPropertyWithExplicitDefault(219, "sql_range_end_boundary", result->sql_range_end_boundary, + unique_ptr()); + deserializer.ReadPropertyWithExplicitDefault>(220, "sql_range_order_casts", + result->sql_range_order_casts, {}); // Builtin window functions didn't used to be serialized, so we need to look them up in the system catalog if (!result->aggregate && !result->window) { diff --git a/src/duckdb/src/planner/expression/window_range_info.cpp b/src/duckdb/src/planner/expression/window_range_info.cpp new file mode 100644 index 000000000..cfe8da279 --- /dev/null +++ b/src/duckdb/src/planner/expression/window_range_info.cpp @@ -0,0 +1,150 @@ +#include "duckdb/planner/expression/window_range_info.hpp" +#include "duckdb/planner/expression/bound_cast_expression.hpp" +#include "duckdb/planner/expression/bound_function_expression.hpp" +#include "duckdb/common/serializer/serializer.hpp" +#include "duckdb/common/serializer/deserializer.hpp" + +namespace duckdb { + +static bool HasDirectRangeArguments(const BoundScalarFunction &function) { + const auto &definition = function.GetDefinition(); + if (!definition || function.HasBindExpressionCallback() || definition->HasUnbindCallback()) { + return false; + } + auto &properties = definition->GetProperties(); + return !properties.GetCaptureArgumentAliases() && !properties.RequiresExpressionNames(); +} + +bool WindowRangeCast::Capture(const Expression &expression, optional_ptr input, + vector &casts) { + optional_ptr current = expression; + while (current.get() != input.get()) { + if (!BoundCastExpression::IsCast(*current)) { + return false; + } + auto &cast = current->Cast(); + if (!BoundCastExpression::HasValidBindData(cast)) { + return false; + } + casts.push_back({BoundCastExpression::SourceType(cast), cast.GetReturnType(), + BoundCastExpression::IsTryCast(cast), BoundCastExpression::IsDefaultCast(cast)}); + current = &BoundCastExpression::Child(cast); + } + return true; +} + +optional_ptr WindowRangeCast::Match(const Expression &expression, + const vector &casts) { + optional_ptr current = expression; + for (auto &origin : casts) { + if (!BoundCastExpression::IsCast(*current)) { + return nullptr; + } + auto &cast = current->Cast(); + if (!BoundCastExpression::HasValidBindData(cast) || + !BoundCastExpression::SourceType(cast).EqualsIncludingCollation(origin.source_type) || + !cast.GetReturnType().EqualsIncludingCollation(origin.target_type) || + BoundCastExpression::IsTryCast(cast) != origin.try_cast || + BoundCastExpression::IsDefaultCast(cast) != origin.default_cast) { + return nullptr; + } + current = &BoundCastExpression::Child(cast); + } + return current; +} + +unique_ptr WindowRangeBoundary::Capture(const Expression &expression, + optional_ptr order, + optional_ptr offset) { + if (expression.GetExpressionClass() != ExpressionClass::BOUND_FUNCTION) { + return nullptr; + } + auto &call = expression.Cast(); + auto &function = call.Function(); + const auto &definition = function.GetDefinition(); + if (!HasDirectRangeArguments(function) || call.GetChildren().size() != 2) { + return nullptr; + } + auto result = make_uniq(); + vector offset_casts; + if (!WindowRangeCast::Capture(*call.GetChildren()[0], order, result->order_casts) || + !WindowRangeCast::Capture(*call.GetChildren()[1], offset, offset_casts)) { + return nullptr; + } + result->function_name = definition->GetQualifiedName(); + result->arguments = function.GetLogicalArguments(); + result->return_type = function.GetLogicalReturnType(); + result->order_type = call.GetChildren()[0]->GetReturnType(); + result->offset_type = call.GetChildren()[1]->GetReturnType(); + return result; +} + +optional_ptr WindowRangeBoundary::Match(const Expression &expression, const Expression &order) const { + auto endpoint = WindowRangeCast::Match(expression, result_casts); + if (!endpoint || endpoint->GetExpressionClass() != ExpressionClass::BOUND_FUNCTION) { + return nullptr; + } + auto &call = endpoint->Cast(); + auto &function = call.Function(); + const auto &definition = function.GetDefinition(); + if (!HasDirectRangeArguments(function) || call.GetChildren().size() != 2) { + return nullptr; + } + const bool matches_signature = definition->GetQualifiedName() == function_name && + function.GetLogicalArguments() == arguments && + function.GetLogicalReturnType() == return_type; + const bool matches_types = call.GetReturnType().EqualsIncludingCollation(return_type) && + call.GetChildren()[0]->GetReturnType().EqualsIncludingCollation(order_type) && + call.GetChildren()[1]->GetReturnType().EqualsIncludingCollation(offset_type); + if (!matches_signature || !matches_types) { + return nullptr; + } + auto input_order = WindowRangeCast::Match(*call.GetChildren()[0], order_casts); + if (!input_order || !Expression::Equals(*input_order, order)) { + return nullptr; + } + return call.GetChildren()[1].get(); +} + +void WindowRangeCast::Serialize(Serializer &serializer) const { + serializer.WriteProperty(100, "source_type", source_type); + serializer.WriteProperty(101, "target_type", target_type); + serializer.WriteProperty(102, "try_cast", try_cast); + serializer.WriteProperty(103, "default_cast", default_cast); +} + +WindowRangeCast WindowRangeCast::Deserialize(Deserializer &deserializer) { + WindowRangeCast result; + deserializer.ReadProperty(100, "source_type", result.source_type); + deserializer.ReadProperty(101, "target_type", result.target_type); + deserializer.ReadProperty(102, "try_cast", result.try_cast); + deserializer.ReadProperty(103, "default_cast", result.default_cast); + return result; +} + +void WindowRangeBoundary::Serialize(Serializer &serializer) const { + serializer.WriteProperty(100, "function_name", function_name); + serializer.WriteProperty(101, "arguments", arguments); + serializer.WriteProperty(102, "return_type", return_type); + serializer.WriteProperty(103, "order_type", order_type); + serializer.WriteProperty(104, "offset_type", offset_type); + serializer.WriteProperty(105, "order_casts", order_casts); + serializer.WriteProperty(106, "result_casts", result_casts); + serializer.WriteProperty(107, "boundary", boundary); + serializer.WriteProperty(108, "direction", direction); +} + +unique_ptr WindowRangeBoundary::Deserialize(Deserializer &deserializer) { + auto result = make_uniq(); + deserializer.ReadProperty(100, "function_name", result->function_name); + deserializer.ReadProperty(101, "arguments", result->arguments); + deserializer.ReadProperty(102, "return_type", result->return_type); + deserializer.ReadProperty(103, "order_type", result->order_type); + deserializer.ReadProperty(104, "offset_type", result->offset_type); + deserializer.ReadProperty(105, "order_casts", result->order_casts); + deserializer.ReadProperty(106, "result_casts", result->result_casts); + deserializer.ReadProperty(107, "boundary", result->boundary); + deserializer.ReadProperty(108, "direction", result->direction); + return result; +} +} // namespace duckdb diff --git a/src/duckdb/src/planner/logical_plan_verifier.cpp b/src/duckdb/src/planner/logical_plan_verifier.cpp index 5dd934c2d..872ed14a1 100644 --- a/src/duckdb/src/planner/logical_plan_verifier.cpp +++ b/src/duckdb/src/planner/logical_plan_verifier.cpp @@ -186,6 +186,29 @@ struct LogicalPlanVerificationState { issues.push_back(std::move(issue)); } + void AddNullOperatorChild(LogicalOperator &op, const LogicalPlanVerificationPath &path, idx_t child_index) { + LogicalPlanVerificationIssue issue; + issue.code = LogicalPlanVerificationIssueCode::INTERNAL_INVARIANT; + issue.path = path; + issue.construct = GetOperatorConstruct(op); + issue.facts.emplace_back("invariant", Value("null_operator_child")); + issue.facts.emplace_back("child_index", Value::UBIGINT(child_index)); + issue.message = StringUtil::Format("Logical operator child %llu is null", child_index); + issues.push_back(std::move(issue)); + } + + void AddNullOperatorExpression(LogicalOperator &op, const LogicalPlanVerificationPath &path, + idx_t expression_index) { + LogicalPlanVerificationIssue issue; + issue.code = LogicalPlanVerificationIssueCode::INTERNAL_INVARIANT; + issue.path = path; + issue.construct = GetOperatorConstruct(op); + issue.facts.emplace_back("invariant", Value("null_operator_expression")); + issue.facts.emplace_back("expression_index", Value::UBIGINT(expression_index)); + issue.message = StringUtil::Format("Logical operator expression %llu is null", expression_index); + issues.push_back(std::move(issue)); + } + bool HasResolvedInputs(LogicalOperator &op) const { return resolved_inputs.find(reference(op)) != resolved_inputs.end(); } @@ -194,6 +217,18 @@ struct LogicalPlanVerificationState { return resolved_outputs.find(reference(op)) != resolved_outputs.end(); } + bool HasNullSlots(LogicalOperator &op) const { + for (auto &child : op.children) { + if (!child) { + return true; + } + } + bool has_null = false; + LogicalOperatorVisitor::EnumerateExpressions( + op, [&](unique_ptr *expression) { has_null |= !expression || !*expression; }); + return has_null; + } + private: static void AddBindingFacts(LogicalPlanVerificationIssue &issue, const ColumnBinding &binding) { issue.facts.emplace_back("table_index", Value::UBIGINT(binding.table_index.index)); @@ -230,11 +265,19 @@ struct LogicalPlanVerificationState { auto expression_path = path; expression_path.components.push_back( {LogicalPlanVerificationPathComponentType::OPERATOR_EXPRESSION, expression_index++}); + if (!expression || !*expression) { + AddNullOperatorExpression(op, expression_path, expression_index - 1); + return; + } IndexExpression(**expression, expression_path); }); for (idx_t child_index = 0; child_index < op.children.size(); child_index++) { auto child_path = path; child_path.components.push_back({LogicalPlanVerificationPathComponentType::OPERATOR_CHILD, child_index}); + if (!op.children[child_index]) { + AddNullOperatorChild(op, child_path, child_index); + continue; + } IndexOperator(*op.children[child_index], child_path); } } @@ -286,11 +329,15 @@ bool LogicalPlanVerifier::ResolveOperatorTypes(LogicalOperator &op, LogicalPlanV op.types.clear(); bool children_resolved = true; for (auto &child : op.children) { + if (!child) { + children_resolved = false; + continue; + } if (!ResolveOperatorTypes(*child, verification_state)) { children_resolved = false; } } - if (!children_resolved) { + if (!children_resolved || verification_state.HasNullSlots(op)) { return false; } verification_state.resolved_inputs.insert(reference(op)); @@ -357,14 +404,21 @@ void LogicalPlanVerifier::VerifyColumnBindings(LogicalOperator &op, LogicalPlanV return; } for (auto &child : op.children) { - VerifyColumnBindings(*child, verification_state); + if (child) { + VerifyColumnBindings(*child, verification_state); + } } } static void VerifyTableIndexes(LogicalOperator &op, LogicalPlanVerificationState &verification_state, unordered_map &seen_indexes) { for (auto &child : op.children) { - VerifyTableIndexes(*child, verification_state, seen_indexes); + if (child) { + VerifyTableIndexes(*child, verification_state, seen_indexes); + } + } + if (verification_state.HasNullSlots(op)) { + return; } auto table_indexes = op.GetTableIndex(); for (idx_t table_index_ordinal = 0; table_index_ordinal < table_indexes.size(); table_index_ordinal++) { diff --git a/src/duckdb/src/planner/operator/logical_comparison_join.cpp b/src/duckdb/src/planner/operator/logical_comparison_join.cpp index 3e8e1ef19..a6afc8b77 100644 --- a/src/duckdb/src/planner/operator/logical_comparison_join.cpp +++ b/src/duckdb/src/planner/operator/logical_comparison_join.cpp @@ -1,6 +1,7 @@ #include "duckdb/planner/operator/logical_comparison_join.hpp" #include "duckdb/planner/expression/bound_comparison_expression.hpp" #include "duckdb/common/enum_util.hpp" +#include "duckdb/common/type_visitor.hpp" namespace duckdb { @@ -69,4 +70,75 @@ bool LogicalComparisonJoin::HasArbitraryConditions() const { return false; } +static bool MarkJoinTypesMatch(const LogicalType &left, const LogicalType &right) { + if (left != right) { + return false; + } + vector collations; + TypeVisitor::Contains(left, [&](const LogicalType &type) { + if (type.id() == LogicalTypeId::VARCHAR) { + collations.push_back(StringType::GetCollation(type)); + } + return false; + }); + idx_t index = 0; + const auto mismatch = TypeVisitor::Contains(right, [&](const LogicalType &type) { + if (type.id() != LogicalTypeId::VARCHAR) { + return false; + } + if (index >= collations.size()) { + return true; + } + const auto &collation = collations[index]; + index++; + return collation != StringType::GetCollation(type); + }); + return !mismatch && index == collations.size(); +} + +bool LogicalComparisonJoin::TryGetMarkJoinGroupTypes(vector &group_types) const { + group_types.clear(); + if (join_type != JoinType::MARK || conditions.size() < 2) { + return false; + } + vector result; + for (idx_t i = 0; i < conditions.size(); i++) { + auto &condition = conditions[i]; + // Reconstruct matching bound types only; SQL binding may first insert casts. + if (!condition.IsComparison() || condition.GetLHS().GetReturnType() != condition.GetRHS().GetReturnType()) { + return false; + } + if (i + 1 < conditions.size()) { + if (condition.GetComparisonType() != ExpressionType::COMPARE_NOT_DISTINCT_FROM || + !MarkJoinTypesMatch(condition.GetLHS().GetReturnType(), condition.GetRHS().GetReturnType())) { + return false; + } + result.push_back(condition.GetLHS().GetReturnType()); + continue; + } + auto type_id = condition.GetLHS().GetReturnType().id(); + // Native UNION quantifiers are supported, but grouped SQL reconstruction is not yet supported. + if (type_id == LogicalTypeId::TUPLE || type_id == LogicalTypeId::UNION) { + return false; + } + switch (condition.GetComparisonType()) { + case ExpressionType::COMPARE_EQUAL: + case ExpressionType::COMPARE_NOTEQUAL: + break; + case ExpressionType::COMPARE_LESSTHAN: + case ExpressionType::COMPARE_GREATERTHAN: + case ExpressionType::COMPARE_LESSTHANOREQUALTO: + case ExpressionType::COMPARE_GREATERTHANOREQUALTO: + if (condition.GetLHS().GetReturnType().IsNested()) { + return false; + } + break; + default: + return false; + } + } + group_types = std::move(result); + return true; +} + } // namespace duckdb diff --git a/src/duckdb/src/planner/operator/logical_explain.cpp b/src/duckdb/src/planner/operator/logical_explain.cpp index 9c771f906..06773871d 100644 --- a/src/duckdb/src/planner/operator/logical_explain.cpp +++ b/src/duckdb/src/planner/operator/logical_explain.cpp @@ -1,4 +1,8 @@ #include "duckdb/planner/operator/logical_explain.hpp" +#include "duckdb/planner/logical_plan_sql_exporter.hpp" +#include "duckdb/planner/operator/logical_column_data_get.hpp" +#include "duckdb/common/types/column/column_data_collection.hpp" +#include "duckdb/common/types/data_chunk.hpp" namespace duckdb { @@ -8,6 +12,53 @@ LogicalExplain::LogicalExplain(unique_ptr plan, ExplainType exp children.push_back(std::move(plan)); } +unique_ptr LogicalExplain::CreateSQLResult(ClientContext &context, TableIndex table_index) { + D_ASSERT(explain_type == ExplainType::EXPLAIN_SQL); + D_ASSERT(children.size() == 1); + LogicalPlanSQLExportOptions options; + options.output_names = sql_output_names; + auto exported = LogicalPlanSQLExporter::Export(context, *children[0], options); + if (exported.HasError()) { + auto &issue = exported.GetIssues()[0]; + auto message = "EXPLAIN (SQL) cannot render this query: " + issue.message; + const bool is_source_function = + issue.construct && issue.construct->type == LogicalPlanVerificationConstructType::SOURCE_FUNCTION; + const bool has_source_name = + is_source_function && issue.construct->function && issue.construct->function->name != "logical_source"; + if (has_source_name) { + message = StringUtil::Format("EXPLAIN (SQL) cannot render table function \"%s\".", + issue.construct->function->name); + } + for (auto &entry : exported.GetIssues()) { + switch (entry.code) { + case LogicalPlanVerificationIssueCode::UNSUPPORTED_OPERATOR: + case LogicalPlanVerificationIssueCode::UNSUPPORTED_EXPRESSION: + case LogicalPlanVerificationIssueCode::UNSUPPORTED_FUNCTION: + case LogicalPlanVerificationIssueCode::UNSUPPORTED_SOURCE: + case LogicalPlanVerificationIssueCode::UNSUPPORTED_EXTENSION: + case LogicalPlanVerificationIssueCode::UNSUPPORTED_EXPORT_FEATURE: + break; + default: + throw InternalException("EXPLAIN (SQL) cannot render this query: " + entry.message); + } + } + if (!allow_unsupported_sql) { + throw NotImplementedException({{"sql_export_unsupported", "true"}}, message); + } + } + vector result_types {LogicalType::VARCHAR, LogicalType::VARCHAR}; + auto collection = + make_uniq(context, result_types, ColumnDataAllocatorType::IN_MEMORY_ALLOCATOR); + if (exported.IsSuccess()) { + DataChunk chunk; + chunk.Initialize(Allocator::Get(context), result_types); + chunk.data[0].Append(Value("sql")); + chunk.data[1].Append(Value(exported.GetValue().query->ToString())); + collection->Append(chunk); + } + return make_uniq(table_index, std::move(result_types), std::move(collection)); +} + idx_t LogicalExplain::EstimateCardinality(ClientContext &context) { return 3; } diff --git a/src/duckdb/src/planner/operator/logical_get.cpp b/src/duckdb/src/planner/operator/logical_get.cpp index 9515d9ceb..6715034f9 100644 --- a/src/duckdb/src/planner/operator/logical_get.cpp +++ b/src/duckdb/src/planner/operator/logical_get.cpp @@ -362,6 +362,8 @@ void LogicalGet::Serialize(Serializer &serializer) const { serializer.WritePropertyWithDefault(216, "source_ordinality", source_ordinality, OrdinalityType::WITHOUT_ORDINALITY); serializer.WriteProperty(217, "table_filters_by_projection", true); + serializer.WritePropertyWithDefault(218, "has_pushed_projection", has_pushed_projection, false); + serializer.WritePropertyWithDefault(219, "at_clause", at_clause); } unique_ptr LogicalGet::Deserialize(Deserializer &deserializer) { @@ -399,6 +401,9 @@ unique_ptr LogicalGet::Deserialize(Deserializer &deserializer) 216, "source_ordinality", OrdinalityType::WITHOUT_ORDINALITY); auto table_filters_by_projection = deserializer.ReadPropertyWithExplicitDefault(217, "table_filters_by_projection", false); + result->has_pushed_projection = + deserializer.ReadPropertyWithExplicitDefault(218, "has_pushed_projection", false); + result->at_clause = deserializer.ReadPropertyWithDefault>(219, "at_clause"); if (!legacy_column_ids.empty()) { if (!result->column_ids.empty()) { throw SerializationException( diff --git a/src/duckdb/src/planner/planner.cpp b/src/duckdb/src/planner/planner.cpp index ba4bad57b..1b5881abb 100644 --- a/src/duckdb/src/planner/planner.cpp +++ b/src/duckdb/src/planner/planner.cpp @@ -26,6 +26,7 @@ #include "duckdb/planner/operator_extension.hpp" #include "duckdb/planner/planner_extension.hpp" #include "duckdb/planner/logical_plan_verifier.hpp" +#include "duckdb/planner/operator/logical_explain.hpp" #include "duckdb/optimizer/optimizer.hpp" namespace duckdb { @@ -33,6 +34,49 @@ namespace duckdb { Planner::Planner(ClientContext &context) : binder(Binder::CreateBinder(context)), context(context) { } +static void RenderSQLExplain(unique_ptr &op, ClientContext &context, Binder &binder) { + if (op->type == LogicalOperatorType::LOGICAL_PREPARE || op->type == LogicalOperatorType::LOGICAL_EXECUTE) { + for (auto &child : op->children) { + RenderSQLExplain(child, context, binder); + } + } else if (op->type == LogicalOperatorType::LOGICAL_EXPLAIN) { + auto &explain = op->Cast(); + if (explain.explain_type == ExplainType::EXPLAIN_SQL) { + op = explain.CreateSQLResult(context, binder.GenerateTableIndex()); + } + } +} + +void Planner::Optimize() { + auto &profiler = QueryProfiler::Get(context); +#ifdef DEBUG + plan->Verify(context); +#endif + bool optimize = Settings::Get(context); + if (Settings::Get(context)) { + // verify disable optimizer - disable EXCEPT for explain, otherwise every single EXPLAIN query breaks + if (plan->type != LogicalOperatorType::LOGICAL_EXPLAIN) { + optimize = false; + } + } + if (plan->RequireOptimizer()) { + { + auto optimizer_timer = profiler.StartTimer(); + Optimizer optimizer(*binder, context); + if (optimize) { + plan = optimizer.Optimize(std::move(plan)); + } else { + plan = optimizer.LowerMandatoryAggregateRewrites(std::move(plan)); + } + D_ASSERT(plan); + } +#ifdef DEBUG + plan->Verify(context); +#endif + } + RenderSQLExplain(plan, context, *binder); +} + // Pre-decorrelation pass: replace LogicalTrigger with LogicalDependentJoin so the standard // FlattenDependentJoins machinery can decorrelate the trigger body. static void RewriteTriggersToDependent(Binder &binder, LogicalOperator &op) { diff --git a/src/duckdb/src/planner/sql_export/bound_expression_sql_exporter.cpp b/src/duckdb/src/planner/sql_export/bound_expression_sql_exporter.cpp new file mode 100644 index 000000000..ccb09f928 --- /dev/null +++ b/src/duckdb/src/planner/sql_export/bound_expression_sql_exporter.cpp @@ -0,0 +1,664 @@ +#include "duckdb/planner/sql_export/bound_expression_sql_exporter_internal.hpp" +#include "duckdb/common/extension_type_info.hpp" +#include "duckdb/function/scalar/compressed_materialization_utils.hpp" +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" +#include "duckdb/parser/expression/between_expression.hpp" +#include "duckdb/parser/expression/case_expression.hpp" +#include "duckdb/parser/expression/cast_expression.hpp" +#include "duckdb/parser/expression/columnref_expression.hpp" +#include "duckdb/parser/expression/collate_expression.hpp" +#include "duckdb/parser/expression/comparison_expression.hpp" +#include "duckdb/parser/expression/conjunction_expression.hpp" +#include "duckdb/parser/expression/function_expression.hpp" +#include "duckdb/parser/expression/operator_expression.hpp" +#include "duckdb/planner/expression/bound_aggregate_expression.hpp" +#include "duckdb/planner/expression/bound_between_expression.hpp" +#include "duckdb/planner/expression/bound_case_expression.hpp" +#include "duckdb/planner/expression/bound_cast_expression.hpp" +#include "duckdb/planner/expression/bound_columnref_expression.hpp" +#include "duckdb/planner/expression/bound_comparison_expression.hpp" +#include "duckdb/planner/expression/bound_conjunction_expression.hpp" +#include "duckdb/planner/expression/bound_constant_expression.hpp" +#include "duckdb/planner/expression/bound_function_expression.hpp" +#include "duckdb/planner/expression/bound_operator_expression.hpp" +#include "duckdb/planner/expression/bound_reference_expression.hpp" +#include "duckdb/function/cast/cast_function_set.hpp" +#include "duckdb/function/scalar/generic_common.hpp" +#include "duckdb/planner/expression/bound_window_expression.hpp" +#include "duckdb/planner/expression/bound_unnest_expression.hpp" + +namespace duckdb { + +static inline bool IsExpressionRootPath(const LogicalPlanVerificationPath &path) { + if (!path.IsValid()) { + return false; + } + if (path.root == LogicalPlanVerificationPathRoot::STANDALONE_EXPRESSION) { + return true; + } + for (auto &component : path.components) { + if (component.type == LogicalPlanVerificationPathComponentType::OPERATOR_EXPRESSION) { + return true; + } + } + return false; +} + +static LogicalPlanVerificationIssue UnsupportedExpression(const LogicalPlanVerificationPath &path, + ExpressionClass expression_class) { + return SQLExportHelpers::MakeIssue( + LogicalPlanVerificationIssueCode::UNSUPPORTED_EXPRESSION, LogicalPlanVerificationPhase::EXPRESSION_EXPORT, path, + LogicalPlanVerificationConstructIdentity::Expression(expression_class), + "The bound expression class does not have a SQL AST representation in this exporter"); +} + +static inline bool ChildrenAreConsistentWithArguments(const vector> &children, + const vector &arguments) { + if (children.size() != arguments.size()) { + return false; + } + for (idx_t child_index = 0; child_index < children.size(); child_index++) { + if (!children[child_index] || !children[child_index]->GetReturnType().IsComplete()) { + return false; + } + if (arguments[child_index].IsComplete() && children[child_index]->GetReturnType() != arguments[child_index]) { + return false; + } + } + return true; +} + +LogicalPlanVerificationIssue +BoundExpressionSQLExportState::InternalInvariant(optional path, string message, + optional construct) { + return SQLExportHelpers::MakeIssue(LogicalPlanVerificationIssueCode::INTERNAL_INVARIANT, + LogicalPlanVerificationPhase::EXPRESSION_EXPORT, std::move(path), + std::move(construct), std::move(message)); +} + +LogicalPlanVerificationIssue +BoundExpressionSQLExportState::InternalExpressionInvariant(const LogicalPlanVerificationPath &path, + const Expression &expression, string message) { + return BoundExpressionSQLExportState::InternalInvariant( + path, std::move(message), + LogicalPlanVerificationConstructIdentity::Expression(expression.GetExpressionClass())); +} + +LogicalPlanVerificationIssue BoundExpressionSQLExportState::UnsupportedFeature(const LogicalPlanVerificationPath &path, + string feature, string message) { + return SQLExportHelpers::MakeIssue( + LogicalPlanVerificationIssueCode::UNSUPPORTED_EXPORT_FEATURE, LogicalPlanVerificationPhase::EXPRESSION_EXPORT, + path, LogicalPlanVerificationConstructIdentity::ExportFeature(std::move(feature)), std::move(message)); +} + +LogicalPlanVerificationIssue +BoundExpressionSQLExportState::UnsupportedFunction(const LogicalPlanVerificationPath &path, + LogicalPlanVerificationFunctionIdentity identity, string message) { + return SQLExportHelpers::MakeIssue( + LogicalPlanVerificationIssueCode::UNSUPPORTED_FUNCTION, LogicalPlanVerificationPhase::EXPRESSION_EXPORT, path, + LogicalPlanVerificationConstructIdentity::Function(std::move(identity)), std::move(message)); +} + +bool BoundExpressionSQLExportState::HasNestedCollation(const LogicalType &type) { + return type.id() != LogicalTypeId::VARCHAR && TypeVisitor::Contains(type, [](const LogicalType &child) { + return child.id() == LogicalTypeId::VARCHAR && !StringType::GetCollation(child).empty(); + }); +} + +BoundExpressionSQLExportResult +BoundExpressionSQLExportState::PreserveCollation(const LogicalType &type, BoundExpressionSQLExportResult result, + const LogicalPlanVerificationPath &path) { + if (result.HasError()) { + return result; + } + if (type.id() == LogicalTypeId::VARCHAR) { + auto collation = StringType::GetCollation(type); + if (!collation.empty()) { + result.GetValue() = make_uniq(std::move(collation), std::move(result.GetValue())); + } + } else if (BoundExpressionSQLExportState::HasNestedCollation(type)) { + auto issue = BoundExpressionSQLExportState::UnsupportedFeature( + path, "nested_result_collation", "Nested result collations require a typed SQL representation"); + issue.facts.emplace_back("logical_type", Value(type.ToString())); + issue.facts.emplace_back("varchar_collations", Value(SQLExportHelpers::TypeCollationSignature(type))); + return BoundExpressionSQLExportResult::Failure({std::move(issue)}); + } + return result; +} + +LogicalType BoundExpressionSQLExportState::SQLCastType(const LogicalType &type) { + // Scalar collations are applied by COLLATE, outside the cast's type expression. + return type.id() == LogicalTypeId::VARCHAR && !type.HasAlias() ? LogicalType::VARCHAR : type; +} + +unique_ptr BoundExpressionSQLExportState::SQLCast(const LogicalType &type, + unique_ptr child, bool try_cast) { + return make_uniq(BoundExpressionSQLExportState::SQLCastType(type), std::move(child), try_cast); +} + +BoundExpressionSQLExportState::BoundExpressionSQLExportState(const BoundExpressionSQLExportContext &context_p) + : context(context_p) { +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::Export(const Expression &expression, + const LogicalPlanVerificationPath &path) { + auto result = ExportInternal(expression, path); + const bool preserves_collation = expression.GetExpressionClass() == ExpressionClass::BOUND_CONSTANT || + (expression.GetExpressionClass() == ExpressionClass::BOUND_FUNCTION && + expression.GetReturnType().id() != LogicalTypeId::VARCHAR); + if (expression.GetExpressionClass() == ExpressionClass::BOUND_COLUMN_REF || preserves_collation) { + return result; + } + return BoundExpressionSQLExportState::PreserveCollation(expression.GetReturnType(), std::move(result), path); +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::ExportInternal(const Expression &expression, + const LogicalPlanVerificationPath &path) { + switch (expression.GetExpressionClass()) { + case ExpressionClass::BOUND_CONSTANT: + return ExportConstant(expression.Cast(), path); + case ExpressionClass::BOUND_COLUMN_REF: + return ExportColumnRef(expression.Cast(), path); + case ExpressionClass::BOUND_REF: + return ExportReference(expression.Cast(), path); + case ExpressionClass::BOUND_FUNCTION: + return ExportFunction(expression.Cast(), path); + case ExpressionClass::BOUND_CONJUNCTION: + return ExportConjunction(expression.Cast(), path); + case ExpressionClass::BOUND_CASE: + return ExportCase(expression.Cast(), path); + case ExpressionClass::BOUND_OPERATOR: + return ExportOperator(expression.Cast(), path); + case ExpressionClass::BOUND_AGGREGATE: + return ExportAggregate(expression.Cast(), path); + case ExpressionClass::BOUND_DEFAULT: + case ExpressionClass::BOUND_PARAMETER: + case ExpressionClass::BOUND_SUBQUERY: + case ExpressionClass::BOUND_WINDOW: + case ExpressionClass::BOUND_UNNEST: + case ExpressionClass::BOUND_LAMBDA: + case ExpressionClass::BOUND_LAMBDA_REF: + case ExpressionClass::LEGACY_BOUND_CAST: + case ExpressionClass::LEGACY_BOUND_COMPARISON: + case ExpressionClass::LEGACY_BOUND_BETWEEN: + return BoundExpressionSQLExportResult::Failure({UnsupportedExpression(path, expression.GetExpressionClass())}); + case ExpressionClass::BOUND_EXPANDED: + case ExpressionClass::AGGREGATE: + case ExpressionClass::CASE: + case ExpressionClass::CAST: + case ExpressionClass::COLUMN_REF: + case ExpressionClass::COMPARISON: + case ExpressionClass::CONJUNCTION: + case ExpressionClass::CONSTANT: + case ExpressionClass::DEFAULT: + case ExpressionClass::FUNCTION: + case ExpressionClass::OPERATOR: + case ExpressionClass::STAR: + case ExpressionClass::SUBQUERY: + case ExpressionClass::WINDOW: + case ExpressionClass::PARAMETER: + case ExpressionClass::COLLATE: + case ExpressionClass::LAMBDA: + case ExpressionClass::POSITIONAL_REFERENCE: + case ExpressionClass::BETWEEN: + case ExpressionClass::LAMBDA_REF: + case ExpressionClass::TYPE: + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::InternalExpressionInvariant( + path, expression, "Expression export requires a final bound class")}); + case ExpressionClass::INVALID: + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::InternalInvariant( + path, "Expression export received an invalid expression class")}); + } + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::InternalInvariant( + path, "Expression export received an unknown expression class")}); +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::ExportUnnest(const BoundUnnestExpression &expression, + const LogicalPlanVerificationPath &path) { + D_ASSERT(expression.Child()); + auto child = ExportChild(*expression.Child(), path, 0); + if (child.HasError()) { + return child; + } + vector> arguments; + arguments.push_back(std::move(child.GetValue())); + return BoundExpressionSQLExportState::PreserveCollation( + expression.GetReturnType(), + BoundExpressionSQLExportResult::Success(make_uniq("unnest", std::move(arguments))), path); +} + +BoundExpressionSQLExportResult +BoundExpressionSQLExportState::ExportReference(const BoundReferenceExpression &expression, + const LogicalPlanVerificationPath &path) { + if (expression.GetExpressionType() != ExpressionType::BOUND_REF || lambda_reference_scopes.empty() || + expression.Index() >= lambda_reference_scopes.back().size()) { + return BoundExpressionSQLExportResult::Failure({UnsupportedExpression(path, expression.GetExpressionClass())}); + } + return BoundExpressionSQLExportResult::Success(lambda_reference_scopes.back()[expression.Index()]->Copy()); +} + +BoundExpressionSQLExportResult +BoundExpressionSQLExportState::ExportColumnRef(const BoundColumnRefExpression &expression, + const LogicalPlanVerificationPath &path) { + D_ASSERT(expression.GetExpressionType() == ExpressionType::BOUND_COLUMN_REF); + auto &binding = expression.Binding(); + if (!binding.table_index.IsValid() || !binding.column_index.IsValid()) { + return BoundExpressionSQLExportResult::Failure( + {InvalidBinding(path, binding, "Bound column reference has an incomplete binding")}); + } + if (!SQLExportHelpers::IsSQLValueType(expression.GetReturnType())) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::InternalExpressionInvariant( + path, expression, "Bound column reference has an incomplete type")}); + } + if (expression.Depth() != 0) { + auto issue = BoundExpressionSQLExportState::UnsupportedFeature( + path, "correlated_column_reference", "Correlated column references require an owning query export context"); + issue.facts.emplace_back("depth", Value::UBIGINT(expression.Depth())); + return BoundExpressionSQLExportResult::Failure({std::move(issue)}); + } + if (!context.resolve_binding) { + return BoundExpressionSQLExportResult::Failure( + {InvalidBinding(path, binding, "No SQL column binding resolver was provided")}); + } + auto resolved = context.resolve_binding(binding); + if (!resolved) { + return BoundExpressionSQLExportResult::Failure( + {InvalidBinding(path, binding, "The SQL column binding resolver has no matching entry")}); + } + if (resolved->names.empty()) { + return BoundExpressionSQLExportResult::Failure( + {InvalidBinding(path, binding, "The resolved SQL column name is empty")}); + } + for (auto &name : resolved->names) { + if (!SQLExportHelpers::IsValidIdentifier(name)) { + return BoundExpressionSQLExportResult::Failure( + {InvalidBinding(path, binding, "The resolved SQL column name contains an invalid identifier")}); + } + } + if (!SQLExportHelpers::IsSQLValueType(resolved->type)) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::InternalExpressionInvariant( + path, expression, "The resolved SQL column type is incomplete")}); + } + auto optimizer_type_match = resolved->optimizer_type && *resolved->optimizer_type == expression.GetReturnType(); + if (resolved->type != expression.GetReturnType() && !optimizer_type_match) { + LogicalPlanVerificationIssue issue; + issue.code = LogicalPlanVerificationIssueCode::TYPE_MISMATCH; + issue.phase = LogicalPlanVerificationPhase::EXPRESSION_EXPORT; + issue.path = path; + issue.construct = + LogicalPlanVerificationConstructIdentity::BindingTypeMismatch(resolved->type, expression.GetReturnType()); + issue.message = "The resolved SQL column type differs from the bound expression type"; + return BoundExpressionSQLExportResult::Failure({std::move(issue)}); + } + auto result = BoundExpressionSQLExportResult::Success(make_uniq(std::move(resolved->names))); + if (optimizer_type_match) { + return result; + } + if (!expression.GetReturnType().EqualsIncludingCollation(resolved->type)) { + if (expression.GetReturnType().id() != LogicalTypeId::VARCHAR) { + auto issue = BoundExpressionSQLExportState::UnsupportedFeature( + path, "nested_result_collation", + "Changing nested input collations requires a typed SQL representation"); + issue.facts.emplace_back("logical_type", Value(expression.GetReturnType().ToString())); + issue.facts.emplace_back("input_logical_type", Value(resolved->type.ToString())); + issue.facts.emplace_back("input_varchar_collations", + Value(SQLExportHelpers::TypeCollationSignature(resolved->type))); + issue.facts.emplace_back("varchar_collations", + Value(SQLExportHelpers::TypeCollationSignature(expression.GetReturnType()))); + return BoundExpressionSQLExportResult::Failure({std::move(issue)}); + } + if (StringType::GetCollation(expression.GetReturnType()).empty()) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "column_collation_reset", "Clearing an input collation requires a SQL representation")}); + } + return BoundExpressionSQLExportState::PreserveCollation(expression.GetReturnType(), std::move(result), path); + } + return result; +} + +LogicalPlanVerificationIssue BoundExpressionSQLExportState::InvalidBinding(const LogicalPlanVerificationPath &path, + const ColumnBinding &binding, + string message) { + LogicalPlanVerificationIssue issue; + issue.code = LogicalPlanVerificationIssueCode::INVALID_BINDING; + issue.phase = LogicalPlanVerificationPhase::EXPRESSION_EXPORT; + issue.path = path; + issue.facts.emplace_back("column_index", Value::UBIGINT(binding.column_index.GetIndexUnsafe())); + issue.facts.emplace_back("table_index", Value::UBIGINT(binding.table_index.index)); + issue.message = std::move(message); + return issue; +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::ExportFunction(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path) { + switch (expression.GetExpressionType()) { + case ExpressionType::BOUND_FUNCTION: + return ExportScalarFunction(expression, path); + case ExpressionType::OPERATOR_CAST: + return ExportCast(expression, path); + case ExpressionType::COMPARE_EQUAL: + case ExpressionType::COMPARE_NOTEQUAL: + case ExpressionType::COMPARE_LESSTHAN: + case ExpressionType::COMPARE_GREATERTHAN: + case ExpressionType::COMPARE_LESSTHANOREQUALTO: + case ExpressionType::COMPARE_GREATERTHANOREQUALTO: + case ExpressionType::COMPARE_DISTINCT_FROM: + case ExpressionType::COMPARE_NOT_DISTINCT_FROM: + return ExportComparison(expression, path); + case ExpressionType::COMPARE_BETWEEN: + return ExportBetween(expression, path); + default: + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::InternalExpressionInvariant( + path, expression, "Bound function has an invalid expression type")}); + } +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::ExportCast(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path) { + D_ASSERT(expression.GetExpressionType() == ExpressionType::OPERATOR_CAST); + D_ASSERT(expression.GetChildren().size() == 1 && expression.GetChildren()[0]); + D_ASSERT(BoundCastExpression::HasValidBindData(expression)); + D_ASSERT(ChildrenAreConsistentWithArguments(expression.GetChildren(), expression.Function().GetArguments())); + D_ASSERT(expression.GetReturnType() == expression.Function().GetReturnType()); + if (!SQLExportHelpers::IsSQLRepresentableType(expression.GetReturnType()) || + !SQLExportHelpers::IsSQLRepresentableType(expression.GetChildren()[0]->GetReturnType())) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "cast_type", "The cast type has no SQL type representation")}); + } + if (context.discard_optimizer_metadata && CMUtils::GetExpressionType(expression) == CMExpressionType::CAST) { + if (!BoundCastExpression::IsDefaultCast(expression)) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "compressed_materialization_cast", + "The compressed materialization projection contains a non-default cast")}); + } + return ExportChild(*expression.GetChildren()[0], path, 0); + } + if (BoundCastExpression::IsDefaultCast(expression)) { + if (!context.client_context || + CastFunctionSet::Get(*context.client_context) + .CanOverrideDefaultCast(expression.GetChildren()[0]->GetReturnType(), expression.GetReturnType())) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "default_cast_binding", + "A default-only bound cast cannot be reconstructed through this SQL binding")}); + } + } + auto child = ExportChild(*expression.GetChildren()[0], path, 0); + if (child.HasError()) { + return child; + } + if (expression.GetReturnType().IsAggregateState()) { + auto storage_type = expression.GetReturnType().WithAlias("").WithExtensionInfo(nullptr); + auto &source_type = expression.GetChildren()[0]->GetReturnType(); + if (BoundCastExpression::IsTryCast(expression) || !context.client_context || + CastFunctionSet::Get(*context.client_context) + .CanOverrideDefaultCast(source_type, expression.GetReturnType()) || + CastFunctionSet::Get(*context.client_context).CanOverrideDefaultCast(source_type, storage_type)) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "aggregate_state_try_cast", + "Aggregate state TRY_CAST or custom casts require a SQL representation")}); + } + auto result = ExportAggregateFunction::StateToSQL( + expression.GetReturnType(), + BoundExpressionSQLExportState::SQLCast(storage_type, std::move(child.GetValue()))); + if (!result) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "aggregate_state_parameters", "Aggregate state SQL parameters are not representable")}); + } + return BoundExpressionSQLExportResult::Success(std::move(result)); + } + if (BoundExpressionSQLExportState::HasNestedCollation(expression.GetReturnType())) { + if (BoundCastExpression::IsTryCast(expression) || + expression.GetReturnType() == expression.GetChildren()[0]->GetReturnType()) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "nested_result_collation", "This cast cannot preserve nested collations through SQL")}); + } + return CastToConstructedType(expression.GetReturnType(), std::move(child.GetValue()), path); + } + if (RequiresConstantConstructor(expression.GetReturnType()) && !BoundCastExpression::IsTryCast(expression)) { + return CastToConstructedType(expression.GetReturnType(), std::move(child.GetValue()), path); + } + return BoundExpressionSQLExportResult::Success(BoundExpressionSQLExportState::SQLCast( + expression.GetReturnType(), std::move(child.GetValue()), BoundCastExpression::IsTryCast(expression))); +} + +BoundExpressionSQLExportResult +BoundExpressionSQLExportState::ExportComparison(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path) { + D_ASSERT(BoundComparisonExpression::IsComparison(expression.GetExpressionType())); + D_ASSERT(expression.GetReturnType() == LogicalType::BOOLEAN); + D_ASSERT(expression.GetChildren().size() == 2 && expression.GetChildren()[0] && expression.GetChildren()[1]); + D_ASSERT(!expression.BindInfo()); + D_ASSERT(ChildrenAreConsistentWithArguments(expression.GetChildren(), expression.Function().GetArguments())); + D_ASSERT(expression.GetReturnType() == expression.Function().GetReturnType()); + D_ASSERT(expression.GetChildren()[0]->GetReturnType() == expression.GetChildren()[1]->GetReturnType()); + vector> children; + vector issues; + ExportChildren(expression.GetChildren(), path, children, issues); + if (!issues.empty()) { + return BoundExpressionSQLExportResult::Failure(std::move(issues)); + } + return BoundExpressionSQLExportResult::Success(make_uniq( + expression.GetExpressionType(), std::move(children[0]), std::move(children[1]))); +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::ExportBetween(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path) { + D_ASSERT(expression.GetExpressionType() == ExpressionType::COMPARE_BETWEEN); + D_ASSERT(expression.GetReturnType() == LogicalType::BOOLEAN); + D_ASSERT(expression.GetChildren().size() == 3 && expression.GetChildren()[0] && expression.GetChildren()[1] && + expression.GetChildren()[2]); + D_ASSERT(BoundBetweenExpression::HasValidBindData(expression)); + D_ASSERT(ChildrenAreConsistentWithArguments(expression.GetChildren(), expression.Function().GetArguments())); + D_ASSERT(expression.GetReturnType() == expression.Function().GetReturnType()); + D_ASSERT(expression.GetChildren()[0]->GetReturnType() == expression.GetChildren()[1]->GetReturnType()); + D_ASSERT(expression.GetChildren()[0]->GetReturnType() == expression.GetChildren()[2]->GetReturnType()); + vector> children; + vector issues; + ExportChildren(expression.GetChildren(), path, children, issues); + if (!issues.empty()) { + return BoundExpressionSQLExportResult::Failure(std::move(issues)); + } + auto lower_inclusive = BoundBetweenExpression::LowerInclusive(expression); + auto upper_inclusive = BoundBetweenExpression::UpperInclusive(expression); + if (lower_inclusive && upper_inclusive) { + return BoundExpressionSQLExportResult::Success( + make_uniq(std::move(children[0]), std::move(children[1]), std::move(children[2]))); + } + if (expression.GetChildren()[0]->IsVolatile()) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "exclusive_between_input_evaluation", + "An exclusive BETWEEN cannot duplicate a volatile input while preserving evaluation semantics")}); + } + auto lower = make_uniq(BoundBetweenExpression::LowerComparisonType(expression), + children[0]->Copy(), std::move(children[1])); + auto upper = make_uniq(BoundBetweenExpression::UpperComparisonType(expression), + std::move(children[0]), std::move(children[2])); + return BoundExpressionSQLExportResult::Success( + make_uniq(ExpressionType::CONJUNCTION_AND, std::move(lower), std::move(upper))); +} + +BoundExpressionSQLExportResult +BoundExpressionSQLExportState::ExportConjunction(const BoundConjunctionExpression &expression, + const LogicalPlanVerificationPath &path) { + D_ASSERT(expression.GetExpressionType() == ExpressionType::CONJUNCTION_AND || + expression.GetExpressionType() == ExpressionType::CONJUNCTION_OR); + D_ASSERT(expression.GetReturnType() == LogicalType::BOOLEAN); + D_ASSERT(expression.GetChildren().size() >= 2); + vector> children; + vector issues; + ExportChildren(expression.GetChildren(), path, children, issues, LogicalType::BOOLEAN); + if (!issues.empty()) { + return BoundExpressionSQLExportResult::Failure(std::move(issues)); + } + auto result = make_uniq(expression.GetExpressionType()); + result->GetChildrenMutable() = std::move(children); + return BoundExpressionSQLExportResult::Success(std::move(result)); +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::ExportCase(const BoundCaseExpression &expression, + const LogicalPlanVerificationPath &path) { + D_ASSERT(expression.GetExpressionType() == ExpressionType::CASE_EXPR); + D_ASSERT(expression.GetReturnType().IsComplete()); + D_ASSERT(!expression.CaseChecks().empty()); + vector source_children; + for (auto &check : expression.CaseChecks()) { + source_children.emplace_back(check.when_expr.get(), LogicalType::BOOLEAN); + source_children.emplace_back(check.then_expr.get(), expression.GetReturnType()); + } + source_children.emplace_back(expression.ElseExpression().get(), expression.GetReturnType()); + + vector> children; + vector issues; + ExportChildren(source_children, path, children, issues); + if (!issues.empty()) { + return BoundExpressionSQLExportResult::Failure(std::move(issues)); + } + auto result = make_uniq(); + for (idx_t check_index = 0; check_index < expression.CaseChecks().size(); check_index++) { + CaseCheck check; + check.when_expr = std::move(children[check_index * 2]); + check.then_expr = std::move(children[check_index * 2 + 1]); + result->CaseChecksMutable().push_back(std::move(check)); + } + result->ElseMutable() = std::move(children.back()); + return BoundExpressionSQLExportResult::Success(std::move(result)); +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::ExportOperator(const BoundOperatorExpression &expression, + const LogicalPlanVerificationPath &path) { + optional expected_type; + switch (expression.GetExpressionType()) { + case ExpressionType::OPERATOR_NOT: + D_ASSERT(expression.GetChildren().size() == 1 && expression.GetReturnType() == LogicalType::BOOLEAN); + expected_type = LogicalType::BOOLEAN; + break; + case ExpressionType::OPERATOR_IS_NULL: + case ExpressionType::OPERATOR_IS_NOT_NULL: + D_ASSERT(expression.GetChildren().size() == 1 && expression.GetReturnType() == LogicalType::BOOLEAN); + break; + case ExpressionType::COMPARE_IN: + case ExpressionType::COMPARE_NOT_IN: + D_ASSERT(expression.GetChildren().size() >= 2 && expression.GetReturnType() == LogicalType::BOOLEAN && + expression.GetChildren()[0]); + expected_type = expression.GetChildren()[0]->GetReturnType(); + break; + case ExpressionType::OPERATOR_COALESCE: + D_ASSERT(expression.GetChildren().size() >= 2 && expression.GetReturnType().IsComplete()); + expected_type = expression.GetReturnType(); + break; + case ExpressionType::OPERATOR_TRY: + D_ASSERT(expression.GetChildren().size() == 1 && expression.GetChildren()[0] && + expression.GetReturnType().IsComplete()); + if (expression.GetChildren()[0]->IsVolatile()) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "try_volatile_child", "TRY cannot be rebound around a volatile expression")}); + } + expected_type = expression.GetReturnType(); + break; + case ExpressionType::OPERATOR_UNPACK: + case ExpressionType::OPERATOR_NULLIF: + case ExpressionType::GROUPING_FUNCTION: + case ExpressionType::ARRAY_EXTRACT: + case ExpressionType::ARRAY_SLICE: + case ExpressionType::STRUCT_EXTRACT: + case ExpressionType::ARRAY_CONSTRUCTOR: + case ExpressionType::ARROW: + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "bound_operator", "The bound operator has no admitted parsed SQL AST form")}); + default: + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::InternalExpressionInvariant( + path, expression, "Bound operator has an invalid expression type")}); + } + vector> children; + vector issues; + ExportChildren(expression.GetChildren(), path, children, issues, expected_type); + if (!issues.empty()) { + return BoundExpressionSQLExportResult::Failure(std::move(issues)); + } + return BoundExpressionSQLExportResult::Success( + make_uniq(expression.GetExpressionType(), std::move(children))); +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::ExportChild(const Expression &expression, + const LogicalPlanVerificationPath &path, + idx_t child_index) { + return Export(expression, SQLExportHelpers::ChildPath(path, child_index)); +} + +void BoundExpressionSQLExportState::ExportChildren(const vector> &source, + const LogicalPlanVerificationPath &path, + vector> &result, + vector &issues, + const optional &expected_type) { + vector source_children; + for (auto &child : source) { + source_children.emplace_back(child.get(), expected_type); + } + ExportChildren(source_children, path, result, issues); +} + +void BoundExpressionSQLExportState::ExportChildren(const vector &source, + const LogicalPlanVerificationPath &path, + vector> &result, + vector &issues) { + result.resize(source.size()); + for (idx_t child_index = 0; child_index < source.size(); child_index++) { + auto child_path = SQLExportHelpers::ChildPath(path, child_index); + auto &input = source[child_index]; + D_ASSERT(input.expression); + D_ASSERT(!input.expected_type || input.expression->GetReturnType() == *input.expected_type); + auto child = Export(*input.expression, child_path); + if (child.HasError()) { + for (auto &issue : child.GetIssues()) { + issues.push_back(issue); + } + } else { + result[child_index] = std::move(child.GetValue()); + } + } +} + +LogicalPlanVerificationResult> +BoundExpressionSQLExporter::Export(const Expression &expression, const BoundExpressionSQLExportContext &context) { + LogicalPlanVerificationPath path; + path.root = LogicalPlanVerificationPathRoot::STANDALONE_EXPRESSION; + return ExportAtPath(expression, context, path); +} + +LogicalPlanVerificationResult> +BoundExpressionSQLExporter::ExportAtPath(const Expression &expression, const BoundExpressionSQLExportContext &context, + const LogicalPlanVerificationPath &path) { + D_ASSERT(IsExpressionRootPath(path)); + BoundExpressionSQLExportState state(context); + return state.Export(expression, path); +} + +LogicalPlanVerificationResult> +BoundExpressionSQLExporter::ExportAggregateCallAtPath(const BoundAggregateExpression &expression, + const BoundExpressionSQLExportContext &context, + const LogicalPlanVerificationPath &path) { + D_ASSERT(IsExpressionRootPath(path)); + BoundExpressionSQLExportState state(context); + return state.ExportAggregateCall(expression, path); +} + +LogicalPlanVerificationResult> +BoundExpressionSQLExporter::ExportWindowAtPath(const BoundWindowExpression &expression, + const BoundExpressionSQLExportContext &context, + const LogicalPlanVerificationPath &path) { + D_ASSERT(IsExpressionRootPath(path)); + BoundExpressionSQLExportState state(context); + return state.ExportWindow(expression, path); +} + +LogicalPlanVerificationResult> +BoundExpressionSQLExporter::ExportUnnestAtPath(const BoundUnnestExpression &expression, + const BoundExpressionSQLExportContext &context, + const LogicalPlanVerificationPath &path) { + D_ASSERT(IsExpressionRootPath(path)); + BoundExpressionSQLExportState state(context); + return state.ExportUnnest(expression, path); +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/logical_plan_sql_exporter.cpp b/src/duckdb/src/planner/sql_export/logical_plan_sql_exporter.cpp new file mode 100644 index 000000000..be925187c --- /dev/null +++ b/src/duckdb/src/planner/sql_export/logical_plan_sql_exporter.cpp @@ -0,0 +1,264 @@ +#include "duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp" +#include "duckdb/function/scalar/compressed_materialization_utils.hpp" +#include "duckdb/planner/logical_plan_sql_exporter.hpp" +#include "duckdb/parser/expression/columnref_expression.hpp" +#include "duckdb/parser/query_node/select_node.hpp" +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/operator/logical_get.hpp" +#include "duckdb/planner/expression/bound_aggregate_expression.hpp" +#include "duckdb/planner/expression/bound_columnref_expression.hpp" +#include "duckdb/planner/expression/bound_constant_expression.hpp" +#include "duckdb/planner/expression/bound_function_expression.hpp" +#include "duckdb/planner/logical_operator.hpp" +#include "duckdb/planner/logical_plan_verifier.hpp" +#include "duckdb/planner/operator/logical_aggregate.hpp" +#include "duckdb/planner/operator/logical_expression_get.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" +#include "duckdb/planner/operator/logical_projection.hpp" + +namespace duckdb { + +static LogicalPlanVerificationIssue UnsupportedOperator(const LogicalPlanVerificationPath &path, + LogicalOperatorType type) { + return SQLExportHelpers::MakeIssue(LogicalPlanVerificationIssueCode::UNSUPPORTED_OPERATOR, + LogicalPlanVerificationPhase::PLAN_EXPORT, path, + LogicalPlanVerificationConstructIdentity::LogicalOperator(type), + "The logical operator does not have a SQL AST representation in this exporter"); +} + +static optional ConstantSQLInput(const Expression &expression, LogicalOperator &input) { + if (expression.GetExpressionClass() == ExpressionClass::BOUND_CONSTANT) { + return expression.Cast().GetValue(); + } + if (auto wrapped = CMUtils::GetWrappedInput(expression)) { + return ConstantSQLInput(*wrapped, input); + } + if (expression.GetExpressionClass() != ExpressionClass::BOUND_COLUMN_REF || input.children.size() != 1) { + return {}; + } + auto &column = expression.Cast(); + if (column.Depth() != 0) { + return {}; + } + switch (input.type) { + case LogicalOperatorType::LOGICAL_PROJECTION: { + auto &projection = input.Cast(); + if (column.Binding().table_index != projection.table_index) { + return {}; + } + return ConstantSQLInput(*projection.expressions[column.Binding().column_index], *input.children[0]); + } + case LogicalOperatorType::LOGICAL_AGGREGATE_AND_GROUP_BY: { + auto &aggregate = input.Cast(); + auto index = column.Binding().column_index; + if (column.Binding().table_index != aggregate.group_index) { + return {}; + } + for (auto &grouping_set : aggregate.grouping_sets) { + if (!grouping_set.count(index)) { + return {}; + } + } + return ConstantSQLInput(*aggregate.groups[index], *input.children[0]); + } + case LogicalOperatorType::LOGICAL_FILTER: + case LogicalOperatorType::LOGICAL_ORDER_BY: + case LogicalOperatorType::LOGICAL_TOP_N: + case LogicalOperatorType::LOGICAL_LIMIT: + case LogicalOperatorType::LOGICAL_DISTINCT: + return ConstantSQLInput(expression, *input.children[0]); + default: + return {}; + } +} + +static LogicalPlanSQLExportResult ApplyOutputNames(LogicalPlanSQLExportResult result, + const vector &output_names) { + if (result.HasError()) { + return result; + } + if (output_names.size() != result.GetValue().fields.size()) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + LogicalPlanVerificationPath(), "output_names", "Output name count does not match the exported plan")}); + } + if (result.GetValue().query->type == QueryNodeType::SELECT_NODE) { + auto &select = result.GetValue().query->Cast(); + D_ASSERT(select.select_list.size() == output_names.size()); + for (idx_t i = 0; i < output_names.size(); i++) { + select.select_list[i]->SetAlias(output_names[i]); + } + return result; + } + auto fields = result.GetValue().fields; + LogicalPlanSQLExportedChild child {std::move(result.GetValue()), Identifier("exported_query")}; + auto select = make_uniq(); + for (idx_t i = 0; i < output_names.size(); i++) { + auto expression = + make_uniq(LogicalPlanSQLExportHelpers::FieldIdentifier(i), child.relation_alias); + expression->SetAlias(output_names[i]); + select->select_list.push_back(std::move(expression)); + } + select->from_table = LogicalPlanSQLExportHelpers::CreateSubquery(std::move(child)); + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(fields)}); +} + +struct LogicalPlanSQLExportContext::SourceScope { + SourceScope(LogicalPlanSQLExportContext &context_p, const vector &sources_p) + : context(context_p), sources(sources_p), parent(context.source_scope) { + context.source_scope = this; + } + ~SourceScope() { + context.source_scope = parent; + } + + LogicalPlanSQLExportContext &context; + const vector &sources; + optional_ptr parent; +}; + +LogicalPlanSQLExportContext::LogicalPlanSQLExportContext(ClientContext &context_p) : context(context_p) { +} + +LogicalPlanSQLExportResult LogicalPlanSQLExportContext::Export(LogicalOperator &op, + const LogicalPlanVerificationPath &path) { + for (auto scope = source_scope; scope; scope = scope->parent) { + for (auto &source : scope->sources) { + if (&source.op.get() == &op) { + return LogicalPlanSQLExportResult::Success( + {CreateNamedSource(source.name, source.relation.fields), source.relation.fields}); + } + } + } + ancestors.push_back(op); + auto result = op.ToSQL(*this, path); + ancestors.pop_back(); + return result; +} + +void LogicalPlanSQLExportContext::PushNamedRelation(TableIndex index, const Identifier &name, bool is_recurring) { + named_relations.push_back({index, name, is_recurring, 0}); +} + +idx_t LogicalPlanSQLExportContext::PopNamedRelation() { + D_ASSERT(!named_relations.empty()); + auto references = named_relations.back().references; + named_relations.pop_back(); + return references; +} + +optional LogicalPlanSQLExportContext::ReferenceNamedRelation(TableIndex index, bool is_recurring) { + for (idx_t i = named_relations.size(); i > 0; i--) { + auto &relation = named_relations[i - 1]; + if (relation.index == index && relation.is_recurring == is_recurring) { + relation.references++; + return relation.name; + } + } + return optional(); +} + +Identifier LogicalPlanSQLExportContext::NextRelationAlias(const Identifier &preferred) { + if (!preferred.empty() && relation_aliases.insert(preferred).second) { + return preferred; + } + while (true) { + auto name = Identifier("r" + to_string(next_relation_ordinal++)); + if (relation_aliases.insert(name).second) { + return name; + } + } +} + +LogicalPlanVerificationResult +LogicalPlanSQLExportContext::ExportChild(LogicalOperator &child, const LogicalPlanVerificationPath &path) { + auto exported = Export(child, path); + if (exported.HasError()) { + return LogicalPlanVerificationResult::Failure(exported); + } + LogicalPlanSQLExportedChild result {std::move(exported.GetValue()), NextRelationAlias()}; + return LogicalPlanVerificationResult::Success(std::move(result)); +} + +LogicalPlanVerificationResult +LogicalPlanSQLExportContext::ExportChild(LogicalOperator &child, const LogicalPlanVerificationPath &path, + const vector &sources) { + SourceScope scope(*this, sources); + return ExportChild(child, path); +} + +LogicalPlanVerificationResult> LogicalPlanSQLExportContext::ExportExpression( + const LogicalOperator &op, const vector> &expressions, idx_t expression_ordinal, + const BoundExpressionSQLExportContext &expression_context, const LogicalPlanVerificationPath &path) { + D_ASSERT(expression_ordinal < expressions.size()); + auto &expression = expressions[expression_ordinal].get(); + unique_ptr restored; + if (expression.GetExpressionClass() == ExpressionClass::BOUND_AGGREGATE && op.children.size() == 1) { + auto &arguments = expression.Cast().GetChildren(); + for (idx_t i = 0; i < arguments.size(); i++) { + if (arguments[i]->GetExpressionClass() == ExpressionClass::BOUND_CONSTANT) { + continue; + } + auto value = ConstantSQLInput(*arguments[i], *op.children[0]); + if (value && value->type().EqualsIncludingCollation(arguments[i]->GetReturnType())) { + if (!restored) { + restored = expression.Copy(); + } + restored->Cast().GetChildrenMutable()[i] = + make_uniq(*value); + } + } + } + return BoundExpressionSQLExporter::ExportAtPath( + restored ? *restored : expression, expression_context, + LogicalPlanSQLExportHelpers::PlanExpressionPath(path, expression_ordinal)); +} + +unique_ptr LogicalPlanSQLExportContext::ForwardFields(const LogicalPlanSQLExportedChild &child, + const vector &fields, + optional_ptr plain) { + auto select = make_uniq(); + for (auto &field : fields) { + bool found = false; + for (idx_t i = 0; i < child.relation.fields.size(); i++) { + if (field.source_binding != child.relation.fields[i].source_binding) { + continue; + } + auto expression = + plain ? plain->select_list[i]->Copy() : LogicalPlanSQLExportHelpers::ChildColumn(child, i); + expression->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(select->select_list.size())); + select->select_list.push_back(std::move(expression)); + found = true; + break; + } + D_ASSERT(found); + (void)found; + } + return select; +} + +LogicalPlanVerificationResult +LogicalOperator::ToSQL(LogicalPlanSQLExportContext &, const LogicalPlanVerificationPath &path) { + if (type == LogicalOperatorType::LOGICAL_DELIM_GET) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::UnsupportedSource( + path, LogicalPlanSQLExportHelpers::LogicalSourceIdentity(), "delim_get")}); + } + D_ASSERT(type != LogicalOperatorType::LOGICAL_INVALID); + return LogicalPlanSQLExportResult::Failure({UnsupportedOperator(path, type)}); +} + +LogicalPlanVerificationResult +LogicalPlanSQLExporter::Export(ClientContext &context, LogicalOperator &root, + const LogicalPlanSQLExportOptions &options) { + auto verification = LogicalPlanVerifier::VerifyAlways(root); + if (verification.HasError()) { + return LogicalPlanVerificationResult::Failure(verification); + } + LogicalPlanSQLExportContext state(context); + auto result = state.Export(root, LogicalPlanVerificationPath()); + if (!options.output_names) { + return result; + } + return ApplyOutputNames(std::move(result), *options.output_names); +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/sql_export_constants.cpp b/src/duckdb/src/planner/sql_export/sql_export_constants.cpp new file mode 100644 index 000000000..e098cfee3 --- /dev/null +++ b/src/duckdb/src/planner/sql_export/sql_export_constants.cpp @@ -0,0 +1,71 @@ +#include "duckdb/planner/sql_export/bound_expression_sql_exporter_internal.hpp" +#include "duckdb/common/error_data.hpp" +#include "duckdb/parser/expression/cast_expression.hpp" +#include "duckdb/parser/expression/constant_expression.hpp" +#include "duckdb/parser/expression/function_expression.hpp" +#include "duckdb/planner/expression/bound_constant_expression.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" + +namespace duckdb { + +bool BoundExpressionSQLExportState::RequiresConstantConstructor(const LogicalType &type) { + return ConstantExpression::RequiresTypeWitness(type); +} + +BoundExpressionSQLExportResult +BoundExpressionSQLExportState::CastToConstructedType(const LogicalType &type, unique_ptr child, + const LogicalPlanVerificationPath &path) { + vector> arguments; + arguments.push_back(std::move(child)); + try { + arguments.push_back(ConstantExpression::FromValue(Value(type))); + } catch (const NotImplementedException &ex) { + return BoundExpressionSQLExportResult::Failure( + {BoundExpressionSQLExportState::UnsupportedFeature(path, "constant_type", ErrorData(ex).RawMessage())}); + } + return BoundExpressionSQLExportResult::Success( + SQLExportHelpers::SystemFunction("cast_to_type", std::move(arguments))); +} + +BoundExpressionSQLExportResult +BoundExpressionSQLExportState::RestoreResultType(const LogicalType &type, unique_ptr result, + const LogicalPlanVerificationPath &path) { + if (RequiresConstantConstructor(type)) { + return CastToConstructedType(type, std::move(result), path); + } + return BoundExpressionSQLExportResult::Success(BoundExpressionSQLExportState::SQLCast(type, std::move(result))); +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::ExportConstant(const BoundConstantExpression &expression, + const LogicalPlanVerificationPath &path) { + D_ASSERT(expression.GetExpressionType() == ExpressionType::VALUE_CONSTANT); + auto &return_type = expression.GetReturnType(); + auto &value = expression.GetValue(); + D_ASSERT(return_type == value.type()); + if (!SQLExportHelpers::IsSQLValueType(return_type)) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::InternalExpressionInvariant( + path, expression, "Bound constant has an unexportable type")}); + } + unique_ptr result; + try { + result = ConstantExpression::FromValue(value); + } catch (const NotImplementedException &ex) { + return BoundExpressionSQLExportResult::Failure( + {BoundExpressionSQLExportState::UnsupportedFeature(path, "constant_value", ErrorData(ex).RawMessage())}); + } + if (RequiresConstantConstructor(return_type) || !SQLExportHelpers::IsSQLRepresentableType(return_type)) { + return BoundExpressionSQLExportResult::Success(std::move(result)); + } + if (return_type.id() != LogicalTypeId::SQLNULL) { + const bool has_result_cast = + result->GetExpressionClass() == ExpressionClass::CAST && + result->Cast().TargetType().Equals( + *TypeExpression::FromLogicalType(BoundExpressionSQLExportState::SQLCastType(return_type))); + if (!has_result_cast) { + result = BoundExpressionSQLExportState::SQLCast(return_type, std::move(result)); + } + } + return BoundExpressionSQLExportResult::Success(std::move(result)); +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/sql_export_cte.cpp b/src/duckdb/src/planner/sql_export/sql_export_cte.cpp new file mode 100644 index 000000000..d21a0e459 --- /dev/null +++ b/src/duckdb/src/planner/sql_export/sql_export_cte.cpp @@ -0,0 +1,253 @@ +#include "duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp" +#include "duckdb/planner/logical_plan_sql_exporter.hpp" +#include "duckdb/main/settings.hpp" +#include "duckdb/optimizer/optimizer.hpp" +#include "duckdb/parser/query_node/recursive_cte_node.hpp" +#include "duckdb/planner/operator/logical_recursive_cte.hpp" +#include "duckdb/parser/expression/columnref_expression.hpp" +#include "duckdb/parser/expression/constant_expression.hpp" +#include "duckdb/parser/common_table_expression_info.hpp" +#include "duckdb/parser/tableref/basetableref.hpp" +#include "duckdb/parser/query_node/select_node.hpp" +#include "duckdb/parser/query_node/set_operation_node.hpp" +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/expression/bound_aggregate_expression.hpp" +#include "duckdb/planner/expression/bound_columnref_expression.hpp" +#include "duckdb/planner/logical_operator.hpp" +#include "duckdb/planner/operator/logical_materialized_cte.hpp" +#include "duckdb/planner/operator/logical_cteref.hpp" +#include "duckdb/planner/operator/logical_projection.hpp" + +namespace duckdb { + +LogicalPlanSQLExportResult LogicalMaterializedCTE::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &cte = *this; + D_ASSERT(cte.children.size() == 2); + auto fields = LogicalPlanSQLExportHelpers::CreateFields(cte, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + auto name = export_context.NextRelationAlias(cte.ctename); + auto producer = + export_context.ExportNamedProducer(*cte.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0), name); + if (producer.HasError()) { + return LogicalPlanSQLExportResult::Failure(producer); + } + export_context.PushNamedRelation(cte.table_index, name, false); + auto consumer = export_context.ExportChild(*cte.children[1], LogicalPlanSQLExportHelpers::PlanChildPath(path, 1)); + auto references = export_context.PopNamedRelation(); + if (consumer.HasError()) { + return LogicalPlanSQLExportResult::Failure(consumer); + } + if (references == 0) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "cte_unreferenced_evaluation", "SQL would discard the unreferenced CTE producer")}); + } + auto info = make_uniq(); + for (idx_t i = 0; i < producer.GetValue().relation.fields.size(); i++) { + info->aliases.push_back(LogicalPlanSQLExportHelpers::FieldIdentifier(i)); + } + if (producer.GetValue().relation.query->type == QueryNodeType::RECURSIVE_CTE_NODE) { + for (auto &key : producer.GetValue().relation.query->Cast().key_targets) { + info->key_targets.push_back(key->Copy()); + } + } + info->query_node = std::move(producer.GetValue().relation.query); + info->materialized = CTEMaterialize::CTE_MATERIALIZE_ALWAYS; + auto select = export_context.ForwardFields(consumer.GetValue(), fields.GetValue()); + select->from_table = LogicalPlanSQLExportHelpers::CreateSubquery(std::move(consumer.GetValue())); + select->cte_map.map.insert(name, std::move(info)); + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(fields.GetValue())}); +} + +LogicalPlanSQLExportResult LogicalCTERef::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &ref = *this; + D_ASSERT(ref.children.empty()); + auto name = export_context.ReferenceNamedRelation(ref.cte_index, ref.is_recurring); + if (!name) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "cte_reference_scope", "The referenced CTE is not in the exported relation scope")}); + } + auto fields = LogicalPlanSQLExportHelpers::CreateFields(ref, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + return LogicalPlanSQLExportResult::Success( + {export_context.CreateNamedSource(*name, fields.GetValue(), ref.is_recurring), std::move(fields.GetValue())}); +} + +unique_ptr LogicalPlanSQLExportContext::CreateNamedSource(const Identifier &name, + const vector &fields, + bool recurring) { + auto table = make_uniq(); + table->SetTable(name); + if (recurring) { + table->SetQualifiedName(Identifier(), Identifier("recurring"), name); + } + table->alias = NextRelationAlias(); + auto select = make_uniq(); + for (idx_t i = 0; i < fields.size(); i++) { + table->column_name_alias.push_back(LogicalPlanSQLExportHelpers::FieldIdentifier(i)); + select->select_list.push_back( + make_uniq(LogicalPlanSQLExportHelpers::FieldIdentifier(i), table->alias)); + } + select->from_table = std::move(table); + return select; +} + +LogicalPlanSQLExportResult LogicalRecursiveCTE::ExportSQLDefinition(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path, + const Identifier &name) { + auto &cte = *this; + D_ASSERT(cte.children.size() == 2); + auto fields = LogicalPlanSQLExportHelpers::CreateFields(cte, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + auto seed = export_context.Export(*cte.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + if (seed.HasError()) { + return LogicalPlanSQLExportResult::Failure(seed); + } + export_context.PushNamedRelation(cte.table_index, name, false); + export_context.PushNamedRelation(cte.table_index, name, true); + auto step = export_context.Export(*cte.children[1], LogicalPlanSQLExportHelpers::PlanChildPath(path, 1)); + auto recurring_references = export_context.PopNamedRelation(); + auto references = export_context.PopNamedRelation(); + if (step.HasError()) { + return LogicalPlanSQLExportResult::Failure(step); + } + if ((references == 0 && recurring_references == 0) || (cte.ref_recurring && recurring_references == 0)) { + // Binding still needs a self reference when optimization removed the recursive scan. + auto empty = export_context.CreateNamedSource(name, fields.GetValue(), cte.ref_recurring); + empty->select_list.clear(); + for (auto &type : cte.internal_types) { + auto value = LogicalPlanSQLExportHelpers::ExportTypedNull(type, path); + if (value.HasError()) { + return LogicalPlanSQLExportResult::Failure(value); + } + empty->select_list.push_back(std::move(value.GetValue())); + } + empty->where_clause = ConstantExpression::FromValue(Value::BOOLEAN(false)); + auto recursive_step = make_uniq(); + recursive_step->setop_type = SetOperationType::UNION; + recursive_step->setop_all = true; + recursive_step->children.push_back(std::move(step.GetValue().query)); + recursive_step->children.push_back(std::move(empty)); + step.GetValue().query = std::move(recursive_step); + } + auto query = make_uniq(); + query->ctename = name; + query->union_all = cte.union_all; + for (idx_t i = 0; i < fields.GetValue().size(); i++) { + query->aliases.push_back(LogicalPlanSQLExportHelpers::FieldIdentifier(i)); + } + BoundExpressionSQLExportContext key_context; + key_context.client_context = &export_context.GetClientContext(); + key_context.resolve_binding = [&](const ColumnBinding &binding) -> optional { + if (binding.table_index != cte.table_index || binding.column_index.GetIndex() >= cte.internal_types.size()) { + return {}; + } + return ResolvedSQLColumnReference { + {LogicalPlanSQLExportHelpers::FieldIdentifier(binding.column_index.GetIndex())}, + cte.internal_types[binding.column_index.GetIndex()]}; + }; + auto expressions = LogicalPlanSQLExportHelpers::CollectExpressions(cte); + idx_t expression_ordinal = 0; + unordered_set key_columns; + for (auto &key : cte.key_targets) { + auto exported = export_context.ExportExpression(cte, expressions, expression_ordinal++, key_context, path); + if (exported.HasError()) { + return LogicalPlanSQLExportResult::Failure(exported); + } + key_columns.insert(key->Cast().Binding().column_index); + query->key_targets.push_back(std::move(exported.GetValue())); + } + if (!cte.key_targets.empty()) { + for (idx_t column = 0; column < cte.internal_types.size(); column++) { + if (key_columns.count(ProjectionIndex(column))) { + continue; + } + auto ordinal = expression_ordinal++; + auto &aggregate = expressions[ordinal].get().Cast(); + const bool has_order = aggregate.GetOrderBys() && !aggregate.GetOrderBys()->orders.empty(); + const bool has_modifiers = aggregate.IsDistinct() || aggregate.GetFilter() || has_order; + if (has_modifiers || aggregate.StateExportMode() != AggregateStateExportMode::NONE) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + LogicalPlanSQLExportHelpers::PlanExpressionPath(path, ordinal), "recursive_payload_modifiers", + "The recursive payload clause cannot preserve these aggregate modifiers")}); + } + auto exported = BoundExpressionSQLExporter::ExportAggregateCallAtPath( + aggregate, key_context, LogicalPlanSQLExportHelpers::PlanExpressionPath(path, ordinal)); + if (exported.HasError()) { + return LogicalPlanSQLExportResult::Failure(exported); + } + exported.GetValue()->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(column)); + query->key_targets.push_back(std::move(exported.GetValue())); + } + } + query->left = std::move(seed.GetValue().query); + query->right = std::move(step.GetValue().query); + return LogicalPlanSQLExportResult::Success({std::move(query), std::move(fields.GetValue())}); +} + +LogicalPlanVerificationResult +LogicalPlanSQLExportContext::ExportNamedProducer(LogicalOperator &op, const LogicalPlanVerificationPath &path, + const Identifier &name) { + if (op.type == LogicalOperatorType::LOGICAL_PROJECTION) { + D_ASSERT(op.children.size() == 1); + auto fields = LogicalPlanSQLExportHelpers::CreateFields(op, path); + auto child_path = LogicalPlanSQLExportHelpers::PlanChildPath(path, 0); + auto child_fields = LogicalPlanSQLExportHelpers::CreateFields(*op.children[0], child_path); + if (fields.IsSuccess() && child_fields.IsSuccess() && + LogicalPlanSQLExportHelpers::IsIdentityProjection(op.Cast(), child_fields.GetValue())) { + ancestors.push_back(op); + auto exported = ExportNamedProducer(*op.children[0], child_path, name); + ancestors.pop_back(); + if (exported.IsSuccess()) { + exported.GetValue().relation.fields = std::move(fields.GetValue()); + } + return exported; + } + } + if (op.type != LogicalOperatorType::LOGICAL_RECURSIVE_CTE) { + return ExportChild(op, path); + } + ancestors.push_back(op); + auto exported = op.Cast().ExportSQLDefinition(*this, path, name); + ancestors.pop_back(); + if (exported.HasError()) { + return LogicalPlanVerificationResult::Failure(exported); + } + return LogicalPlanVerificationResult::Success( + {std::move(exported.GetValue()), NextRelationAlias()}); +} + +LogicalPlanSQLExportResult LogicalRecursiveCTE::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &cte = *this; + if (Optimizer::OptimizerDisabled(export_context.GetClientContext(), OptimizerType::CTE_INLINING) || + Settings::Get(export_context.GetClientContext())) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "recursive_cte_materialization", + "The SQL wrapper requires CTE inlining to preserve recursive evaluation")}); + } + auto name = export_context.NextRelationAlias(cte.ctename); + auto recursive = ExportSQLDefinition(export_context, path, name); + if (recursive.HasError()) { + return recursive; + } + auto info = make_uniq(); + info->aliases = recursive.GetValue().query->Cast().aliases; + for (auto &key : recursive.GetValue().query->Cast().key_targets) { + info->key_targets.push_back(key->Copy()); + } + info->query_node = std::move(recursive.GetValue().query); + info->materialized = CTEMaterialize::CTE_MATERIALIZE_NEVER; + auto select = export_context.CreateNamedSource(name, recursive.GetValue().fields); + select->cte_map.map.insert(name, std::move(info)); + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(recursive.GetValue().fields)}); +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/sql_export_functions.cpp b/src/duckdb/src/planner/sql_export/sql_export_functions.cpp new file mode 100644 index 000000000..2218de9c7 --- /dev/null +++ b/src/duckdb/src/planner/sql_export/sql_export_functions.cpp @@ -0,0 +1,382 @@ +#include "duckdb/planner/sql_export/bound_expression_sql_exporter_internal.hpp" +#include "duckdb/function/scalar/compressed_materialization_utils.hpp" +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" +#include "duckdb/parser/expression/columnref_expression.hpp" +#include "duckdb/parser/expression/constant_expression.hpp" +#include "duckdb/parser/expression/function_expression.hpp" +#include "duckdb/parser/expression/lambda_expression.hpp" +#include "duckdb/planner/expression/bound_aggregate_expression.hpp" +#include "duckdb/planner/expression/bound_columnref_expression.hpp" +#include "duckdb/planner/expression/bound_function_expression.hpp" +#include "duckdb/planner/expression/bound_lambda_expression.hpp" +#include "duckdb/planner/filter/table_filter_functions.hpp" + +namespace duckdb { + +template +static bool IsOptimizerFunctionQualification(const FUNCTION &function) { + if (function.GetCatalogName().empty() && function.GetSchemaName().empty()) { + return true; + } + return function.GetCatalogName() == "system" && function.GetSchemaName() == "main"; +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::ExportLambda(const BoundLambdaExpression &lambda, + const BoundFunctionExpression &function, + idx_t logical_argument_count, + const LogicalPlanVerificationPath &path) { + const bool has_lambda_body = lambda.GetExpressionType() == ExpressionType::LAMBDA && + lambda.GetReturnType() == LogicalType::LAMBDA && lambda.LambdaExpr(); + const bool has_parameter_names = + lambda.ParameterCount() > 0 && lambda.ParameterNames().size() == lambda.ParameterCount(); + if (!has_lambda_body || !lambda.Captures().empty() || !has_parameter_names) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "lambda_binding", "The bound lambda does not retain its SQL parameter binding")}); + } + + vector> references; + for (idx_t index = 0; index < lambda.ParameterCount(); index++) { + auto ¶meter = lambda.ParameterNames()[lambda.ParameterCount() - index - 1]; + references.push_back(make_uniq(parameter)); + } + for (idx_t index = logical_argument_count; index < function.GetChildren().size(); index++) { + auto reference = Export(*function.GetChildren()[index], SQLExportHelpers::ChildPath(path, index)); + if (reference.HasError()) { + return reference; + } + references.push_back(std::move(reference.GetValue())); + } + + lambda_reference_scopes.push_back(std::move(references)); + try { + auto body = Export(*lambda.LambdaExpr(), SQLExportHelpers::ChildPath(path, 0)); + lambda_reference_scopes.pop_back(); + if (body.HasError()) { + return body; + } + vector parameter_names; + for (auto ¶meter : lambda.ParameterNames()) { + parameter_names.push_back(parameter.GetIdentifierName()); + } + return BoundExpressionSQLExportResult::Success( + make_uniq(std::move(parameter_names), std::move(body.GetValue()))); + } catch (...) { + lambda_reference_scopes.pop_back(); + throw; + } +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::CompressedMaterializationFailure( + const BoundFunctionExpression &expression, const LogicalPlanVerificationPath &path, string message) { + auto &function = expression.Function(); + auto &definition = function.GetDefinition(); + D_ASSERT(definition); + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, + BoundExpressionSQLExportState::DefinitionFunctionIdentity(*definition, function.GetLogicalArguments(), + function.GetLogicalReturnType()), + std::move(message))}); +} + +optional +BoundExpressionSQLExportState::TryExportCompressedMaterialization(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path) { + auto &function = expression.Function(); + auto &definition = function.GetDefinition(); + if (!definition) { + return {}; + } + auto compress = CMUtils::GetExpressionType(expression) == CMExpressionType::COMPRESS; + auto decompress = CMUtils::GetExpressionType(expression) == CMExpressionType::DECOMPRESS; + if (!compress && !decompress) { + return {}; + } + if (!IsOptimizerFunctionQualification(*definition) || !IsOptimizerFunctionQualification(function) || + expression.GetChildren().empty()) { + return CompressedMaterializationFailure(expression, path, + "The compressed materialization expression is malformed"); + } + if (!context.discard_optimizer_metadata) { + return CompressedMaterializationFailure(expression, path, + "Compressed materialization requires its complete logical plan"); + } + if (decompress && expression.GetChildren()[0]->GetExpressionClass() == ExpressionClass::BOUND_COLUMN_REF) { + auto &column = expression.GetChildren()[0]->Cast(); + auto resolved = context.resolve_binding ? context.resolve_binding(column.Binding()) + : optional(); + if (!resolved || resolved->type != expression.GetReturnType()) { + return CompressedMaterializationFailure(expression, path, + "The decompression input has no equivalent SQL representation"); + } + } + return ExportChild(*expression.GetChildren()[0], path, 0); +} + +BoundExpressionSQLExportResult +BoundExpressionSQLExportState::ExportScalarFunction(const BoundFunctionExpression &expression, + const LogicalPlanVerificationPath &path) { + auto &function = expression.Function(); + auto &definition = function.GetDefinition(); + D_ASSERT(definition); + if (function.GetBindCallback() == TableFilterFunctions::Bind && + TableFilterFunctions::IsTableFilterFunction(function.GetName())) { + auto identity = BoundExpressionSQLExportState::DefinitionFunctionIdentity( + *definition, function.GetLogicalArguments(), function.GetLogicalReturnType()); + if (!context.discard_optimizer_metadata) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "Internal table filters require their complete logical plan")}); + } + D_ASSERT(expression.GetReturnType() == LogicalType::BOOLEAN && expression.GetChildren().size() == 1); + return BoundExpressionSQLExportResult::Success(ConstantExpression::FromValue(Value::BOOLEAN(true))); + } + auto compressed = TryExportCompressedMaterialization(expression, path); + if (compressed) { + return std::move(*compressed); + } + auto identity = BoundExpressionSQLExportState::DefinitionFunctionIdentity( + *definition, function.GetLogicalArguments(), function.GetLogicalReturnType()); + if (!identity.IsValid()) { + identity.arguments.clear(); + for (auto &child : expression.GetChildren()) { + identity.arguments.push_back(child->GetReturnType()); + } + identity.return_type = expression.GetReturnType(); + if (!identity.IsValid()) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::InternalExpressionInvariant( + path, expression, "Bound scalar function identity is incomplete")}); + } + } + optional_idx lambda_index; + for (idx_t index = 0; index < MinValue(expression.GetChildren().size(), function.GetLogicalArguments().size()); + index++) { + if (expression.GetChildren()[index]->GetExpressionClass() == ExpressionClass::BOUND_LAMBDA) { + if (lambda_index.IsValid()) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The scalar function retains multiple SQL lambda arguments")}); + } + lambda_index = index; + } + } + if (function.HasBindLambdaCallback() != lambda_index.IsValid()) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The scalar function does not retain its SQL lambda argument")}); + } + const auto logical_argument_count = function.GetLogicalArguments().size(); + const auto child_count = expression.GetChildren().size(); + const bool retained_variadic_arguments = definition->HasVarArgs() && child_count >= logical_argument_count; + const bool has_scalar_arguments = + child_count == logical_argument_count || retained_variadic_arguments || definition->HasUnbindCallback(); + const bool has_expected_arguments = + lambda_index.IsValid() ? child_count >= logical_argument_count : has_scalar_arguments; + if (!has_expected_arguments) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The scalar function does not retain every SQL argument")}); + } + auto name = BoundExpressionSQLExportState::RebindableFunctionName(*definition); + if (!name || !SQLExportHelpers::IsSQLValueType(expression.GetReturnType())) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The retained scalar function definition is not representable as SQL")}); + } + const bool can_reconstruct_argument_names = definition->HasUnbindCallback(); + if (definition->GetProperties().GetCaptureArgumentAliases() && !can_reconstruct_argument_names) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The bound scalar function does not expose its SQL argument names")}); + } + if (definition->GetProperties().RequiresExpressionNames() && !can_reconstruct_argument_names) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The bound scalar function requires expression names that are not retained")}); + } + vector argument_names; + if (!definition->HasUnbindCallback() && !function.GetNamedArguments().empty()) { + auto positional_count = function.GetPositionalArgumentCount(); + if (positional_count + function.GetNamedArguments().size() != expression.GetChildren().size()) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The named SQL arguments are incomplete")}); + } + argument_names.resize(expression.GetChildren().size()); + for (idx_t i = 0; i < function.GetNamedArguments().size(); i++) { + argument_names[positional_count + i] = function.GetNamedArguments()[i]; + } + } + vector> children; + auto sql_argument_count = + lambda_index.IsValid() ? function.GetLogicalArguments().size() : expression.GetChildren().size(); + for (idx_t child_index = 0; child_index < sql_argument_count; child_index++) { + auto child = + lambda_index == child_index + ? ExportLambda(expression.GetChildren()[child_index]->Cast(), expression, + sql_argument_count, SQLExportHelpers::ChildPath(path, child_index)) + : Export(*expression.GetChildren()[child_index], SQLExportHelpers::ChildPath(path, child_index)); + if (child.HasError()) { + return child; + } + children.push_back(std::move(child.GetValue())); + } + unique_ptr result; + if (definition->HasUnbindCallback()) { + FunctionUnbindInput input(expression, std::move(children)); + result = definition->GetUnbindCallback()(input); + if (!result) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The function cannot reconstruct its bound invocation")}); + } + } else if (!argument_names.empty()) { + vector arguments; + for (idx_t argument_index = 0; argument_index < children.size(); argument_index++) { + arguments.emplace_back(argument_names[argument_index], std::move(children[argument_index])); + } + result = make_uniq(*name, std::move(arguments), nullptr, nullptr, false, false, false); + } else { + result = make_uniq(*name, std::move(children), nullptr, nullptr, false, false, false); + } + // Restore result types when binding or optimization changed argument types. + const bool can_restore_result_type = SQLExportHelpers::IsSQLRepresentableType(expression.GetReturnType()) && + !expression.GetReturnType().IsAggregateState(); + const bool has_specialized_result_type = + (definition->HasBindCallback() || definition->GetReturnType().id() == LogicalTypeId::SQLNULL) && + definition->GetReturnType() != expression.GetReturnType(); + if (can_restore_result_type && has_specialized_result_type) { + return RestoreResultType(expression.GetReturnType(), std::move(result), path); + } + return BoundExpressionSQLExportResult::Success(std::move(result)); +} + +BoundAggregateSQLExportResult +BoundExpressionSQLExportState::BuildAggregateCall(const BoundAggregateExpression &expression, + const LogicalPlanVerificationPath &path) { + D_ASSERT(expression.GetExpressionType() == ExpressionType::BOUND_AGGREGATE); + auto &function = expression.Function(); + auto &definition = function.GetDefinition(); + D_ASSERT(definition); + auto identity = BoundExpressionSQLExportState::DefinitionFunctionIdentity( + *definition, function.GetLogicalArguments(), function.GetLogicalReturnType()); + if (!identity.IsValid()) { + return BoundAggregateSQLExportResult::Failure({BoundExpressionSQLExportState::InternalExpressionInvariant( + path, expression, "Bound aggregate function identity is incomplete")}); + } + auto name = BoundExpressionSQLExportState::RebindableFunctionName(*definition); + if (!definition->HasUnbindCallback() && expression.GetChildren().size() != function.GetLogicalArguments().size()) { + return BoundAggregateSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The aggregate does not retain every SQL argument")}); + } + if (!name || !SQLExportHelpers::IsSQLValueType(expression.GetReturnType())) { + return BoundAggregateSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The retained aggregate function definition is not representable as SQL")}); + } + D_ASSERT(expression.GetAggregateType() == AggregateType::NON_DISTINCT || + expression.GetAggregateType() == AggregateType::DISTINCT); + D_ASSERT(expression.StateExportMode() == AggregateStateExportMode::NONE || + expression.StateExportMode() == AggregateStateExportMode::STATE_EXPORT); + if (!definition->HasUnbindCallback() && (definition->GetProperties().GetCaptureArgumentAliases() || + definition->GetProperties().RequiresExpressionNames())) { + return BoundAggregateSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The bound aggregate requires argument aliases that are not retained")}); + } + vector source_children; + for (idx_t child_index = 0; child_index < expression.GetChildren().size(); child_index++) { + source_children.emplace_back(expression.GetChildren()[child_index].get()); + } + if (expression.GetFilter()) { + source_children.emplace_back(expression.GetFilter().get(), LogicalType::BOOLEAN); + } + if (expression.GetOrderBys()) { + for (auto &order : expression.GetOrderBys()->orders) { + D_ASSERT(order.type == OrderType::ASCENDING || order.type == OrderType::DESCENDING); + D_ASSERT(order.null_order == OrderByNullType::NULLS_FIRST || + order.null_order == OrderByNullType::NULLS_LAST); + source_children.emplace_back(order.expression.get()); + } + } + + vector> children; + vector issues; + ExportChildren(source_children, path, children, issues); + if (!issues.empty()) { + return BoundAggregateSQLExportResult::Failure(std::move(issues)); + } + unique_ptr result; + if (definition->HasUnbindCallback()) { + vector> arguments; + for (idx_t i = 0; i < expression.GetChildren().size(); i++) { + arguments.push_back(std::move(children[i])); + } + AggregateFunctionUnbindInput input(expression, std::move(arguments)); + result = definition->GetUnbindCallback()(input); + if (!result) { + return BoundAggregateSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The aggregate cannot reconstruct its bound invocation")}); + } + } else { + auto &named_arguments = function.GetNamedArguments(); + auto positional_count = function.GetPositionalArgumentCount(); + if (!named_arguments.empty() && positional_count + named_arguments.size() != expression.GetChildren().size()) { + return BoundAggregateSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The named SQL arguments are incomplete")}); + } + vector arguments; + for (idx_t i = 0; i < expression.GetChildren().size(); i++) { + auto argument_name = !named_arguments.empty() && i >= positional_count + ? named_arguments[i - positional_count] + : Identifier(); + arguments.emplace_back(std::move(argument_name), std::move(children[i])); + } + result = make_uniq(*name, std::move(arguments)); + } + idx_t child_index = expression.GetChildren().size(); + unique_ptr filter; + if (expression.GetFilter()) { + filter = std::move(children[child_index++]); + } + auto order_bys = make_uniq(); + if (expression.GetOrderBys()) { + for (auto &order : expression.GetOrderBys()->orders) { + order_bys->orders.emplace_back(order.type, order.null_order, + SQLExportHelpers::OrderExpression(order.expression->GetReturnType(), + std::move(children[child_index++]))); + } + } + result->FilterMutable() = std::move(filter); + result->OrderByMutable() = std::move(order_bys); + result->DistinctMutable() = expression.IsDistinct(); + result->ExportStateMutable() = expression.StateExportMode() == AggregateStateExportMode::STATE_EXPORT; + return BoundAggregateSQLExportResult::Success(std::move(result)); +} + +BoundExpressionSQLExportResult +BoundExpressionSQLExportState::ExportAggregate(const BoundAggregateExpression &expression, + const LogicalPlanVerificationPath &path) { + auto call = BuildAggregateCall(expression, path); + if (call.HasError()) { + return BoundExpressionSQLExportResult::Failure(call); + } + auto &function = expression.Function(); + auto &definition = function.GetDefinition(); + unique_ptr result = std::move(call.GetValue()); + const bool can_restore_result_type = expression.StateExportMode() == AggregateStateExportMode::NONE && + SQLExportHelpers::IsSQLRepresentableType(expression.GetReturnType()) && + !expression.GetReturnType().IsAggregateState(); + const bool has_specialized_result_type = + definition->HasBindCallback() && definition->GetReturnType() != expression.GetReturnType(); + if (can_restore_result_type && has_specialized_result_type) { + return RestoreResultType(expression.GetReturnType(), std::move(result), path); + } + return BoundExpressionSQLExportResult::Success(std::move(result)); +} + +BoundAggregateSQLExportResult +BoundExpressionSQLExportState::ExportAggregateCall(const BoundAggregateExpression &expression, + const LogicalPlanVerificationPath &path) { + auto call = BuildAggregateCall(expression, path); + if (call.HasError()) { + return call; + } + if (!expression.GetReturnType().EqualsIncludingCollation(expression.Function().GetLogicalReturnType())) { + return BoundAggregateSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "aggregate_call_result_type", + "A bare aggregate call cannot preserve the bound expression's logical result type")}); + } + return call; +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/sql_export_joins.cpp b/src/duckdb/src/planner/sql_export/sql_export_joins.cpp new file mode 100644 index 000000000..f88adcbb0 --- /dev/null +++ b/src/duckdb/src/planner/sql_export/sql_export_joins.cpp @@ -0,0 +1,235 @@ +#include "duckdb/planner/operator/logical_unconditional_join.hpp" +#include "duckdb/planner/operator/logical_join.hpp" +#include "duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp" +#include "duckdb/parser/expression/columnref_expression.hpp" +#include "duckdb/parser/expression/comparison_expression.hpp" +#include "duckdb/parser/expression/constant_expression.hpp" +#include "duckdb/parser/query_node/select_node.hpp" +#include "duckdb/parser/tableref/joinref.hpp" +#include "duckdb/planner/logical_operator.hpp" +#include "duckdb/planner/operator/logical_comparison_join.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" + +namespace duckdb { + +static string MarkConditionUnsupportedReason(const LogicalComparisonJoin &join) { + bool comparisons_only = !join.conditions.empty(); + bool all_equal = true; + bool all_null_safe = true; + for (auto &condition : join.conditions) { + if (!condition.IsComparison()) { + comparisons_only = false; + continue; + } + all_equal &= condition.GetComparisonType() == ExpressionType::COMPARE_EQUAL; + all_null_safe &= condition.GetComparisonType() == ExpressionType::COMPARE_NOT_DISTINCT_FROM; + } + const bool has_supported_conjunction = join.conditions.size() == 1 || all_equal || all_null_safe; + if (!comparisons_only || !has_supported_conjunction) { + return "The MARK condition requires conjunction execution semantics"; + } + for (auto &condition : join.conditions) { + switch (condition.GetComparisonType()) { + case ExpressionType::COMPARE_DISTINCT_FROM: + return "MARK DISTINCT FROM comparisons cannot preserve NULL semantics"; + case ExpressionType::COMPARE_LESSTHAN: + case ExpressionType::COMPARE_GREATERTHAN: + case ExpressionType::COMPARE_LESSTHANOREQUALTO: + case ExpressionType::COMPARE_GREATERTHANOREQUALTO: + if (condition.GetLHS().GetReturnType().IsNested() || condition.GetRHS().GetReturnType().IsNested()) { + return "MARK ordering comparisons on nested types cannot preserve NULL semantics"; + } + break; + default: + break; + } + } + return string(); +} + +static bool RequiresMarkGroupMetadata(const LogicalComparisonJoin &join) { + if (join.join_type != JoinType::MARK || join.mark_types.empty()) { + return false; + } + idx_t comparison_count = 0; + bool tuple_comparison = false; + bool all_equal = true; + for (auto &condition : join.conditions) { + if (!condition.IsComparison()) { + continue; + } + comparison_count++; + tuple_comparison |= condition.GetLHS().GetReturnType().id() == LogicalTypeId::TUPLE; + all_equal &= condition.GetComparisonType() == ExpressionType::COMPARE_EQUAL; + } + if (comparison_count == join.mark_types.size() + 1) { + return true; + } + // Retained types also disable uncorrelated row-equality NULL handling. + return (comparison_count > 1 || tuple_comparison) && all_equal; +} + +static LogicalPlanVerificationResult> +ExportJoinCondition(LogicalJoin &op, LogicalPlanSQLExportContext &context, + const BoundExpressionSQLExportContext &expression_context, + const LogicalPlanVerificationPath &path) { + using Result = LogicalPlanVerificationResult>; + unique_ptr predicate; + auto expressions = LogicalPlanSQLExportHelpers::CollectExpressions(op); + if (op.type == LogicalOperatorType::LOGICAL_ANY_JOIN) { + if (op.join_type == JoinType::MARK) { + return Result::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "mark_condition_semantics", "The MARK condition requires conjunction execution semantics")}); + } + auto exported = context.ExportExpression(op, expressions, 0, expression_context, path); + if (exported.HasError()) { + return Result::Failure(exported); + } + predicate = std::move(exported.GetValue()); + } else { + auto &comparison = op.Cast(); + if (!comparison.duplicate_eliminated_columns.empty()) { + return Result::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "join_delim_state", "The join requires a duplicate-eliminated input scope")}); + } + if (comparison.join_type == JoinType::MARK) { + auto reason = MarkConditionUnsupportedReason(comparison); + if (!reason.empty()) { + return Result::Failure( + {LogicalPlanSQLExportHelpers::PlanUnsupportedFeature(path, "mark_condition_semantics", reason)}); + } + } + if (RequiresMarkGroupMetadata(comparison)) { + return Result::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "mark_group_null_semantics", "The MARK join requires its group-specific NULL semantics")}); + } + idx_t ordinal = 0; + for (auto &condition : comparison.conditions) { + auto lhs = context.ExportExpression(op, expressions, ordinal++, expression_context, path); + if (lhs.HasError()) { + return Result::Failure(lhs); + } + auto conjunct = std::move(lhs.GetValue()); + if (condition.IsComparison()) { + auto rhs = context.ExportExpression(op, expressions, ordinal++, expression_context, path); + if (rhs.HasError()) { + return Result::Failure(rhs); + } + conjunct = make_uniq(condition.GetComparisonType(), std::move(conjunct), + std::move(rhs.GetValue())); + } + predicate = SQLExportHelpers::Conjoin(std::move(predicate), std::move(conjunct)); + } + if (!predicate) { + predicate = ConstantExpression::FromValue(Value::BOOLEAN(true)); + } + } + return Result::Success(std::move(predicate)); +} + +static LogicalPlanSQLExportResult ExportJoin(LogicalOperator &op, LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) { + D_ASSERT(op.children.size() == 2); + auto fields = LogicalPlanSQLExportHelpers::CreateFields(op, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + auto left = context.ExportChild(*op.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + if (left.HasError()) { + return LogicalPlanSQLExportResult::Failure(left); + } + auto right = context.ExportChild(*op.children[1], LogicalPlanSQLExportHelpers::PlanChildPath(path, 1)); + if (right.HasError()) { + return LogicalPlanSQLExportResult::Failure(right); + } + vector> children {left.GetValue(), right.GetValue()}; + LogicalPlanSQLExportHelpers::PropagateSemanticTypes(fields.GetValue(), children); + auto left_plain = LogicalPlanSQLExportHelpers::PlainScope(*left.GetValue().relation.query); + auto right_plain = LogicalPlanSQLExportHelpers::PlainScope(*right.GetValue().relation.query); + if (left_plain && left_plain->where_clause) { + left_plain = nullptr; + } + if (right_plain && right_plain->where_clause) { + right_plain = nullptr; + } + identifier_set_t left_aliases, right_aliases; + if (left_plain) { + LogicalPlanSQLExportHelpers::CollectScopeAliases(*left_plain->from_table, left_aliases); + if (left_aliases.count(right.GetValue().relation_alias)) { + left_plain = nullptr; + } + } + if (right_plain) { + LogicalPlanSQLExportHelpers::CollectScopeAliases(*right_plain->from_table, right_aliases); + if (right_aliases.count(left.GetValue().relation_alias)) { + right_plain = nullptr; + } + } + if (left_plain && right_plain) { + for (auto &alias : left_aliases) { + if (right_aliases.count(alias)) { + right_plain = nullptr; + break; + } + } + } + auto expression_context = LogicalPlanSQLExportHelpers::CreateBindingContext(context.GetClientContext(), children, + {left_plain, right_plain}); + auto join = make_uniq(); + optional_ptr logical_join; + if (op.type == LogicalOperatorType::LOGICAL_CROSS_PRODUCT) { + join->ref_type = JoinRefType::CROSS; + } else if (op.type == LogicalOperatorType::LOGICAL_POSITIONAL_JOIN) { + join->ref_type = JoinRefType::POSITIONAL; + } else { + logical_join = op.Cast(); + join->type = logical_join->join_type; + if (op.type == LogicalOperatorType::LOGICAL_ASOF_JOIN) { + join->ref_type = JoinRefType::ASOF; + } + auto condition = ExportJoinCondition(*logical_join, context, expression_context, path); + if (condition.HasError()) { + return LogicalPlanSQLExportResult::Failure(condition); + } + join->condition = std::move(condition.GetValue()); + } + auto select = make_uniq(); + for (auto &field : fields.GetValue()) { + unique_ptr expression; + if (logical_join && logical_join->join_type == JoinType::MARK && + field.source_binding.table_index == logical_join->mark_index) { + expression = make_uniq(Identifier("__mark_join_marker")); + } else { + auto resolved = expression_context.resolve_binding(field.source_binding); + D_ASSERT(resolved && resolved->type == field.type); + expression = make_uniq(resolved->names); + } + expression->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(select->select_list.size())); + select->select_list.push_back(std::move(expression)); + } + join->left = left_plain ? std::move(left.GetValue().relation.query->Cast().from_table) + : LogicalPlanSQLExportHelpers::CreateSubquery(std::move(left.GetValue())); + join->right = right_plain ? std::move(right.GetValue().relation.query->Cast().from_table) + : LogicalPlanSQLExportHelpers::CreateSubquery(std::move(right.GetValue())); + select->from_table = std::move(join); + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(fields.GetValue())}); +} + +LogicalPlanVerificationResult +LogicalJoin::ToSQL(LogicalPlanSQLExportContext &context, const LogicalPlanVerificationPath &path) { + if (type != LogicalOperatorType::LOGICAL_COMPARISON_JOIN && type != LogicalOperatorType::LOGICAL_ANY_JOIN && + type != LogicalOperatorType::LOGICAL_ASOF_JOIN) { + return LogicalOperator::ToSQL(context, path); + } + return ExportJoin(*this, context, path); +} + +LogicalPlanVerificationResult +LogicalUnconditionalJoin::ToSQL(LogicalPlanSQLExportContext &context, const LogicalPlanVerificationPath &path) { + if (type != LogicalOperatorType::LOGICAL_CROSS_PRODUCT && type != LogicalOperatorType::LOGICAL_POSITIONAL_JOIN) { + return LogicalOperator::ToSQL(context, path); + } + return ExportJoin(*this, context, path); +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/sql_export_limit.cpp b/src/duckdb/src/planner/sql_export/sql_export_limit.cpp new file mode 100644 index 000000000..0496d0ec9 --- /dev/null +++ b/src/duckdb/src/planner/sql_export/sql_export_limit.cpp @@ -0,0 +1,301 @@ +#include "duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp" +#include "duckdb/parser/expression/columnref_expression.hpp" +#include "duckdb/parser/expression/constant_expression.hpp" +#include "duckdb/parser/expression/subquery_expression.hpp" +#include "duckdb/parser/parsed_expression_iterator.hpp" +#include "duckdb/parser/common_table_expression_info.hpp" +#include "duckdb/parser/tableref/basetableref.hpp" +#include "duckdb/parser/query_node/select_node.hpp" +#include "duckdb/parser/statement/select_statement.hpp" +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/expression/bound_columnref_expression.hpp" +#include "duckdb/planner/logical_operator.hpp" +#include "duckdb/planner/operator/logical_aggregate.hpp" +#include "duckdb/planner/operator/logical_materialized_cte.hpp" +#include "duckdb/planner/operator/logical_cteref.hpp" +#include "duckdb/planner/operator/logical_limit.hpp" +#include "duckdb/planner/operator/logical_expression_get.hpp" + +namespace duckdb { + +namespace { + +struct LimitSQLExport { + unique_ptr modifier; + vector sources; +}; + +class LimitSQLExporter { +public: + explicit LimitSQLExporter(LogicalPlanSQLExportContext &context_p) : context(context_p) { + } + LogicalPlanVerificationResult Export(LogicalLimit &limit, const LogicalPlanVerificationPath &path); + +private: + using LimitExpressionResult = LogicalPlanVerificationResult>; + static bool ProducesOneRow(const LogicalOperator &op, const vector &single_row_ctes = {}); + static LimitExpressionResult LimitBindingFailure(const LogicalPlanVerificationPath &path); + LimitExpressionResult ResolveLimitColumn(const ColumnBinding &binding, LogicalOperator &input, + const LogicalPlanVerificationPath &path); + LimitExpressionResult ExportLimitExpression(const Expression &expression, LogicalOperator &input, + const LogicalPlanVerificationPath &path, + const LogicalPlanVerificationPath &expression_path); + +private: + LogicalPlanSQLExportContext &context; + vector sources; +}; + +bool LimitSQLExporter::ProducesOneRow(const LogicalOperator &op, const vector &single_row_ctes) { + if (op.type == LogicalOperatorType::LOGICAL_DUMMY_SCAN) { + return true; + } + if (op.type == LogicalOperatorType::LOGICAL_AGGREGATE_AND_GROUP_BY) { + auto &aggregate = op.Cast(); + return aggregate.groups.empty() && aggregate.grouping_sets.size() <= 1; + } + if (op.type == LogicalOperatorType::LOGICAL_PROJECTION) { + return ProducesOneRow(*op.children[0], single_row_ctes); + } + if (op.type == LogicalOperatorType::LOGICAL_EXPRESSION_GET) { + return op.Cast().expressions.size() == 1 && + ProducesOneRow(*op.children[0], single_row_ctes); + } + if (op.type == LogicalOperatorType::LOGICAL_CROSS_PRODUCT) { + return ProducesOneRow(*op.children[0], single_row_ctes) && ProducesOneRow(*op.children[1], single_row_ctes); + } + if (op.type == LogicalOperatorType::LOGICAL_MATERIALIZED_CTE) { + auto ctes = single_row_ctes; + if (ProducesOneRow(*op.children[0], ctes)) { + ctes.push_back(op.Cast().table_index); + } + return ProducesOneRow(*op.children[1], ctes); + } + if (op.type == LogicalOperatorType::LOGICAL_CTE_REF) { + auto index = op.Cast().cte_index; + return std::find(single_row_ctes.begin(), single_row_ctes.end(), index) != single_row_ctes.end(); + } + return false; +} + +LimitSQLExporter::LimitExpressionResult LimitSQLExporter::LimitBindingFailure(const LogicalPlanVerificationPath &path) { + return LimitExpressionResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "limit_binding", "SQL LIMIT requires an independent, single-row scalar input")}); +} + +LimitSQLExporter::LimitExpressionResult LimitSQLExporter::ResolveLimitColumn(const ColumnBinding &binding, + LogicalOperator &input, + const LogicalPlanVerificationPath &path) { + auto bindings = input.GetColumnBindings(); + auto found = std::find(bindings.begin(), bindings.end(), binding); + if (found == bindings.end()) { + return LimitBindingFailure(path); + } + auto column = NumericCast(found - bindings.begin()); + if (input.type == LogicalOperatorType::LOGICAL_PROJECTION) { + return ExportLimitExpression(*input.expressions[column], *input.children[0], + LogicalPlanSQLExportHelpers::PlanChildPath(path, 0), + LogicalPlanSQLExportHelpers::PlanExpressionPath(path, column)); + } + if (input.type == LogicalOperatorType::LOGICAL_ORDER_BY) { + return ResolveLimitColumn(binding, *input.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + } + if (input.type != LogicalOperatorType::LOGICAL_CROSS_PRODUCT) { + return LimitBindingFailure(path); + } + for (idx_t child_index = 0; child_index < input.children.size(); child_index++) { + auto &child = *input.children[child_index]; + auto child_bindings = child.GetColumnBindings(); + auto child_column = std::find(child_bindings.begin(), child_bindings.end(), binding); + if (child_column == child_bindings.end()) { + continue; + } + auto child_path = LogicalPlanSQLExportHelpers::PlanChildPath(path, child_index); + if (!ProducesOneRow(child)) { + return ResolveLimitColumn(binding, child, child_path); + } + idx_t source_index = 0; + for (; source_index < sources.size(); source_index++) { + if (&sources[source_index].op.get() == &child) { + break; + } + } + if (source_index == sources.size()) { + auto exported = context.ExportChild(child, child_path, sources); + if (exported.HasError()) { + return LimitExpressionResult::Failure(exported); + } + sources.push_back( + {child, std::move(exported.GetValue().relation_alias), std::move(exported.GetValue().relation)}); + } + auto scalar = make_uniq(); + auto table = make_uniq(); + table->SetTable(sources[source_index].name); + scalar->from_table = std::move(table); + scalar->select_list.push_back(make_uniq( + LogicalPlanSQLExportHelpers::FieldIdentifier(NumericCast(child_column - child_bindings.begin())))); + auto subquery = make_uniq(); + subquery->SubqueryMutable() = make_uniq(); + subquery->SubqueryMutable()->node = std::move(scalar); + subquery->GetSubqueryTypeMutable() = SubqueryType::SCALAR; + return LimitExpressionResult::Success(std::move(subquery)); + } + return LimitBindingFailure(path); +} + +LimitSQLExporter::LimitExpressionResult +LimitSQLExporter::ExportLimitExpression(const Expression &expression, LogicalOperator &input, + const LogicalPlanVerificationPath &path, + const LogicalPlanVerificationPath &expression_path) { + if (expression.GetExpressionClass() == ExpressionClass::BOUND_COLUMN_REF) { + auto &column = expression.Cast(); + return column.Depth() == 0 ? ResolveLimitColumn(column.Binding(), input, path) : LimitBindingFailure(path); + } + if (expression.IsVolatile()) { + return LimitBindingFailure(path); + } + vector, unique_ptr>> replacements; + vector issues; + BoundExpressionSQLExportContext expression_context; + expression_context.client_context = context.GetClientContext(); + expression_context.resolve_binding = [&](const ColumnBinding &binding) -> optional { + auto resolved = ResolveLimitColumn(binding, input, path); + if (resolved.HasError()) { + issues = resolved.GetIssues(); + return {}; + } + auto bindings = input.GetColumnBindings(); + auto found = std::find(bindings.begin(), bindings.end(), binding); + D_ASSERT(found != bindings.end()); + vector names {context.NextRelationAlias(), LogicalPlanSQLExportHelpers::FieldIdentifier(0)}; + replacements.emplace_back(names, std::move(resolved.GetValue())); + return ResolvedSQLColumnReference { + std::move(names), input.types[NumericCast(found - bindings.begin())], {}}; + }; + auto result = BoundExpressionSQLExporter::ExportAtPath(expression, expression_context, expression_path); + if (!issues.empty()) { + return LimitExpressionResult::Failure(std::move(issues)); + } + if (result.HasError()) { + return result; + } + std::function &)> replace = [&](unique_ptr &expr) { + if (expr->GetExpressionClass() == ExpressionClass::COLUMN_REF) { + for (auto &entry : replacements) { + if (expr->Cast().ColumnNames() == entry.first) { + expr = entry.second->Copy(); + return; + } + } + } + ParsedExpressionIterator::EnumerateChildren(*expr, replace); + }; + replace(result.GetValue()); + return result; +} + +LogicalPlanVerificationResult LimitSQLExporter::Export(LogicalLimit &limit, + const LogicalPlanVerificationPath &path) { + auto modifier = make_uniq(); + idx_t expression_ordinal = 0; + for (idx_t i = 0; i < 2; i++) { + auto &value = i == 0 ? limit.limit_val : limit.offset_val; + auto &target = i == 0 ? modifier->limit : modifier->offset; + switch (value.Type()) { + case LimitNodeType::UNSET: + break; + case LimitNodeType::CONSTANT_VALUE: + target = ConstantExpression::FromValue(Value::BIGINT(NumericCast(value.GetConstantValue()))); + break; + case LimitNodeType::CONSTANT_PERCENTAGE: + modifier->limit_type = LimitValueType::PERCENTAGE; + target = ConstantExpression::FromValue(Value::DOUBLE(value.GetConstantPercentage())); + break; + case LimitNodeType::EXPRESSION_VALUE: + case LimitNodeType::EXPRESSION_PERCENTAGE: { + if (!value.GetExpression()->IsScalar()) { + auto expression = ExportLimitExpression( + *value.GetExpression(), *limit.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0), + LogicalPlanSQLExportHelpers::PlanExpressionPath(path, expression_ordinal++)); + if (expression.HasError()) { + return LogicalPlanVerificationResult::Failure(expression); + } + target = std::move(expression.GetValue()); + if (value.Type() == LimitNodeType::EXPRESSION_PERCENTAGE) { + modifier->limit_type = LimitValueType::PERCENTAGE; + } + break; + } + if (value.Type() == LimitNodeType::EXPRESSION_PERCENTAGE) { + modifier->limit_type = LimitValueType::PERCENTAGE; + } + auto expression = BoundExpressionSQLExporter::ExportAtPath( + *value.GetExpression(), {}, + LogicalPlanSQLExportHelpers::PlanExpressionPath(path, expression_ordinal++)); + if (expression.HasError()) { + return LogicalPlanVerificationResult::Failure(expression); + } + target = std::move(expression.GetValue()); + break; + } + default: + D_ASSERT(false); + } + } + return LogicalPlanVerificationResult::Success({std::move(modifier), std::move(sources)}); +} + +} // namespace + +LogicalPlanSQLExportResult LogicalLimit::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &limit = *this; + D_ASSERT(limit.children.size() == 1); + auto fields = LogicalPlanSQLExportHelpers::CreateFields(limit, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + LimitSQLExporter modifier_exporter(export_context); + auto exported = modifier_exporter.Export(limit, path); + if (exported.HasError()) { + return LogicalPlanSQLExportResult::Failure(exported); + } + auto &modifier = exported.GetValue().modifier; + auto &sources = exported.GetValue().sources; + auto child = + export_context.ExportChild(*limit.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0), sources); + if (child.HasError()) { + return LogicalPlanSQLExportResult::Failure(child); + } + LogicalPlanSQLExportHelpers::PropagateSemanticTypes(fields.GetValue(), {child.GetValue()}); + if (limit.unpruned_offset.IsValid()) { + modifier->offset = + ConstantExpression::FromValue(Value::BIGINT(NumericCast(limit.unpruned_offset.GetIndex()))); + } + // Keep the ordering and its LIMIT consumer in the same query scope. + bool has_limit = false; + for (auto &existing : child.GetValue().relation.query->modifiers) { + has_limit |= existing->type == ResultModifierType::LIMIT_MODIFIER; + } + unique_ptr query; + if (has_limit) { + auto select = export_context.ForwardFields(child.GetValue(), fields.GetValue()); + select->from_table = LogicalPlanSQLExportHelpers::CreateSubquery(std::move(child.GetValue())); + query = std::move(select); + } else { + query = std::move(child.GetValue().relation.query); + } + query->modifiers.push_back(std::move(modifier)); + for (auto &source : sources) { + auto info = make_uniq(); + for (idx_t i = 0; i < source.relation.fields.size(); i++) { + info->aliases.push_back(LogicalPlanSQLExportHelpers::FieldIdentifier(i)); + } + info->query_node = std::move(source.relation.query); + info->materialized = CTEMaterialize::CTE_MATERIALIZE_ALWAYS; + query->cte_map.map.insert(source.name, std::move(info)); + } + return LogicalPlanSQLExportResult::Success({std::move(query), std::move(fields.GetValue())}); +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/sql_export_pivot.cpp b/src/duckdb/src/planner/sql_export/sql_export_pivot.cpp new file mode 100644 index 000000000..6edb00337 --- /dev/null +++ b/src/duckdb/src/planner/sql_export/sql_export_pivot.cpp @@ -0,0 +1,215 @@ +#include "duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp" +#include "duckdb/parser/expression/case_expression.hpp" +#include "duckdb/parser/expression/columnref_expression.hpp" +#include "duckdb/parser/expression/constant_expression.hpp" +#include "duckdb/parser/expression/operator_expression.hpp" +#include "duckdb/parser/expression/function_expression.hpp" +#include "duckdb/parser/expression/lambda_expression.hpp" +#include "duckdb/parser/common_table_expression_info.hpp" +#include "duckdb/parser/tableref/basetableref.hpp" +#include "duckdb/parser/query_node/select_node.hpp" +#include "duckdb/parser/tableref/emptytableref.hpp" +#include "duckdb/parser/tableref/joinref.hpp" +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/expression/bound_aggregate_expression.hpp" +#include "duckdb/planner/expression/bound_columnref_expression.hpp" +#include "duckdb/planner/expression/bound_constant_expression.hpp" +#include "duckdb/planner/expression_iterator.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" +#include "duckdb/planner/operator/logical_pivot.hpp" + +namespace duckdb { + +static LogicalPlanVerificationResult> +ExportPivotDefault(ClientContext &context, const BoundAggregateExpression &aggregate, + const LogicalPlanVerificationPath &path) { + if (aggregate.Function().GetStability() == FunctionStability::VOLATILE || + aggregate.Function().GetErrorMode() == FunctionErrors::CAN_THROW_RUNTIME_ERROR) { + return LogicalPlanVerificationResult>::Failure( + {LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "pivot_empty_aggregate", + "A volatile or fallible PIVOT default must be evaluated before query execution")}); + } + if (aggregate.StateExportMode() != AggregateStateExportMode::NONE) { + return LogicalPlanVerificationResult>::Failure( + {LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "pivot_empty_aggregate", "The PIVOT default requires an ordinary aggregate invocation")}); + } + auto copy = aggregate.Copy(); + bool outer_reference = false; + std::function &)> replace = [&](unique_ptr &expression) { + if (expression->GetExpressionClass() == ExpressionClass::BOUND_COLUMN_REF) { + if (expression->Cast().Depth() != 0) { + outer_reference = true; + return; + } + expression = make_uniq(Value(expression->GetReturnType())); + return; + } + ExpressionIterator::EnumerateChildren(*expression, replace); + }; + replace(copy); + if (outer_reference) { + return LogicalPlanVerificationResult>::Failure( + {LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "pivot_empty_aggregate", "The PIVOT default contains an unresolved outer reference")}); + } + auto result = BoundExpressionSQLExporter::ExportAggregateCallAtPath( + copy->Cast(), LogicalPlanSQLExportHelpers::CreateBindingContext(context, {}), path); + if (result.HasError()) { + return LogicalPlanVerificationResult>::Failure(result); + } + return LogicalPlanVerificationResult>::Success(std::move(result.GetValue())); +} + +static inline bool HasConsistentPivotLayout(const LogicalPivot &pivot) { + auto &info = pivot.bound_pivot; + auto aggregate_count = info.aggregates.size(); + if (aggregate_count == 0 || info.group_count > pivot.types.size() || info.types != pivot.types || + info.pivot_values.size() != pivot.types.size() - info.group_count || + info.pivot_values.size() % aggregate_count != 0) { + return false; + } + auto &child_types = pivot.children[0]->types; + if (child_types.size() != info.group_count + aggregate_count + 1) { + return false; + } + for (idx_t group_idx = 0; group_idx < info.group_count; group_idx++) { + if (!child_types[group_idx].EqualsIncludingCollation(info.types[group_idx])) { + return false; + } + } + for (idx_t aggregate_idx = 0; aggregate_idx < aggregate_count; aggregate_idx++) { + auto &aggregate = info.aggregates[aggregate_idx]; + auto &list_type = child_types[info.group_count + aggregate_idx]; + if (!aggregate || aggregate->GetExpressionClass() != ExpressionClass::BOUND_AGGREGATE || + list_type.id() != LogicalTypeId::LIST || + !ListType::GetChildType(list_type).EqualsIncludingCollation(aggregate->GetReturnType())) { + return false; + } + for (idx_t target_idx = 0; target_idx < info.pivot_values.size(); target_idx += aggregate_count) { + if (info.pivot_values[target_idx + aggregate_idx] != info.pivot_values[target_idx] || + !info.types[info.group_count + target_idx + aggregate_idx].EqualsIncludingCollation( + aggregate->GetReturnType())) { + return false; + } + } + } + auto &key_type = child_types.back(); + return key_type.id() == LogicalTypeId::LIST && ListType::GetChildType(key_type).id() == LogicalTypeId::VARCHAR; +} + +LogicalPlanSQLExportResult LogicalPivot::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &pivot = *this; + D_ASSERT(pivot.children.size() == 1); + D_ASSERT(HasConsistentPivotLayout(pivot)); + auto fields = LogicalPlanSQLExportHelpers::CreateFields(pivot, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + auto &info = pivot.bound_pivot; + auto aggregate_count = info.aggregates.size(); + auto target_count = (fields.GetValue().size() - info.group_count) / aggregate_count; + if (target_count == 0) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "pivot_layout", "The PIVOT has no output targets")}); + } + + auto defaults = make_uniq(); + defaults->from_table = make_uniq(); + defaults->where_clause = ConstantExpression::FromValue(Value::BOOLEAN(false)); + auto defaults_name = export_context.NextRelationAlias(); + for (idx_t aggregate_idx = 0; aggregate_idx < aggregate_count; aggregate_idx++) { + auto &expression = info.aggregates[aggregate_idx]; + auto &aggregate = expression->Cast(); + auto value = ExportPivotDefault(export_context.GetClientContext(), aggregate, + LogicalPlanSQLExportHelpers::PlanExpressionPath(path, aggregate_idx)); + if (value.HasError()) { + return LogicalPlanSQLExportResult::Failure(value); + } + value.GetValue()->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(aggregate_idx)); + defaults->select_list.push_back(std::move(value.GetValue())); + } + auto child = export_context.ExportChild(*pivot.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + if (child.HasError()) { + return LogicalPlanSQLExportResult::Failure(child); + } + auto call = [](const char *name, unique_ptr first, + unique_ptr second = nullptr) -> unique_ptr { + vector> arguments; + arguments.push_back(std::move(first)); + if (second) { + arguments.push_back(std::move(second)); + } + return SQLExportHelpers::SystemFunction(name, std::move(arguments)); + }; + auto key_name = export_context.NextRelationAlias(); + auto encoded_keys = + call("list_transform", + LogicalPlanSQLExportHelpers::ChildColumn(child.GetValue(), info.group_count + aggregate_count), + make_uniq(vector {key_name.GetIdentifierName()}, + call("encode", make_uniq(key_name)))); + auto reversed_keys = call("list_reverse", std::move(encoded_keys)); + auto length = + call("len", LogicalPlanSQLExportHelpers::ChildColumn(child.GetValue(), info.group_count + aggregate_count)); + auto select = make_uniq(); + for (idx_t group_idx = 0; group_idx < info.group_count; group_idx++) { + auto expression = LogicalPlanSQLExportHelpers::ChildColumn(child.GetValue(), group_idx); + expression->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(group_idx)); + select->select_list.push_back(std::move(expression)); + } + unordered_set seen_keys; + for (idx_t target_idx = 0; target_idx < target_count; target_idx++) { + auto &key = info.pivot_values[target_idx * aggregate_count]; + unique_ptr position = ConstantExpression::FromValue(Value(LogicalType::BIGINT)); + auto first_target = seen_keys.insert(key).second; + if (first_target) { + position = + call("list_position", reversed_keys->Copy(), call("encode", ConstantExpression::FromValue(Value(key)))); + } + auto index = + call("-", call("+", length->Copy(), ConstantExpression::FromValue(Value::BIGINT(1))), position->Copy()); + for (idx_t aggregate_idx = 0; aggregate_idx < aggregate_count; aggregate_idx++) { + auto expression = + call("list_extract", + LogicalPlanSQLExportHelpers::ChildColumn(child.GetValue(), info.group_count + aggregate_idx), + index->Copy()); + auto fallback = make_uniq(LogicalPlanSQLExportHelpers::FieldIdentifier(aggregate_idx), + defaults_name); + if (!first_target) { + expression = std::move(fallback); + } else { + auto result = make_uniq(); + CaseCheck check; + check.when_expr = make_uniq(ExpressionType::OPERATOR_IS_NULL, position->Copy()); + check.then_expr = std::move(fallback); + result->CaseChecksMutable().push_back(std::move(check)); + result->ElseMutable() = std::move(expression); + expression = std::move(result); + } + expression->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(select->select_list.size())); + select->select_list.push_back(std::move(expression)); + } + } + auto source = make_uniq(); + source->SetTable(defaults_name); + source->alias = defaults_name; + auto join = make_uniq(JoinRefType::CROSS); + join->left = LogicalPlanSQLExportHelpers::CreateSubquery(std::move(child.GetValue())); + join->right = std::move(source); + select->from_table = std::move(join); + auto default_info = make_uniq(); + default_info->query_node = std::move(defaults); + default_info->materialized = CTEMaterialize::CTE_MATERIALIZE_ALWAYS; + select->cte_map.map.insert(defaults_name, std::move(default_info)); + vector row; + for (idx_t i = 0; i < select->select_list.size(); i++) { + row.emplace_back(LogicalPlanSQLExportHelpers::FieldIdentifier(i), std::move(select->select_list[i])); + } + auto packed_row = SQLExportHelpers::SystemFunction("struct_pack", std::move(row)); + select->select_list.clear(); + return export_context.ExportRow(std::move(select), std::move(packed_row), std::move(fields.GetValue())); +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/sql_export_relations.cpp b/src/duckdb/src/planner/sql_export/sql_export_relations.cpp new file mode 100644 index 000000000..59a9d82d5 --- /dev/null +++ b/src/duckdb/src/planner/sql_export/sql_export_relations.cpp @@ -0,0 +1,509 @@ +#include "duckdb/planner/operator/logical_unnest.hpp" +#include "duckdb/planner/operator/logical_window.hpp" +#include "duckdb/planner/operator/logical_distinct.hpp" +#include "duckdb/planner/operator/logical_top_n.hpp" +#include "duckdb/planner/operator/logical_order.hpp" +#include "duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp" +#include "duckdb/function/scalar/compressed_materialization_utils.hpp" +#include "duckdb/planner/logical_plan_sql_exporter.hpp" +#include "duckdb/common/limits.hpp" +#include "duckdb/parser/expression/constant_expression.hpp" +#include "duckdb/parser/expression/operator_expression.hpp" +#include "duckdb/parser/query_node/select_node.hpp" +#include "duckdb/parser/query_node/set_operation_node.hpp" +#include "duckdb/parser/tableref/emptytableref.hpp" +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/operator/logical_get.hpp" +#include "duckdb/planner/expression/bound_columnref_expression.hpp" +#include "duckdb/planner/expression/bound_function_expression.hpp" +#include "duckdb/planner/logical_operator.hpp" +#include "duckdb/planner/operator/logical_aggregate.hpp" +#include "duckdb/planner/operator/logical_empty_result.hpp" +#include "duckdb/planner/operator/logical_set_operation.hpp" +#include "duckdb/planner/operator/logical_filter.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" +#include "duckdb/planner/operator/logical_projection.hpp" +#include "duckdb/planner/operator/logical_sample.hpp" +#include "duckdb/planner/expression/bound_window_expression.hpp" +#include "duckdb/planner/expression/bound_unnest_expression.hpp" + +namespace duckdb { + +static LogicalType SemanticExpressionType(const Expression &expression, + const BoundExpressionSQLExportContext &context) { + if (expression.GetExpressionClass() == ExpressionClass::BOUND_COLUMN_REF && context.resolve_binding) { + auto &column = expression.Cast(); + auto resolved = context.resolve_binding(column.Binding()); + if (resolved) { + const bool matches_optimizer_type = + resolved->optimizer_type && *resolved->optimizer_type == column.GetReturnType(); + if (resolved->type == column.GetReturnType() || matches_optimizer_type) { + return resolved->type; + } + } + } + if (context.discard_optimizer_metadata) { + if (auto wrapped = CMUtils::GetWrappedInput(expression)) { + return SemanticExpressionType(*wrapped, context); + } + } + return expression.GetReturnType(); +} + +static void ApplySemanticType(LogicalPlanSQLExportField &field, const Expression &expression, + const BoundExpressionSQLExportContext &context) { + auto semantic_type = SemanticExpressionType(expression, context); + if (semantic_type == field.type) { + return; + } + field.optimizer_type = field.type; + field.type = std::move(semantic_type); +} + +static bool HasSafePredicates(const LogicalOperator &op) { + if (op.type == LogicalOperatorType::LOGICAL_EXPRESSION_GET) { + return !LogicalPlanSQLExportHelpers::HasEffectfulExpressions(op) && + (op.children[0]->type == LogicalOperatorType::LOGICAL_DUMMY_SCAN || HasSafePredicates(*op.children[0])); + } + if (op.type != LogicalOperatorType::LOGICAL_FILTER && op.type != LogicalOperatorType::LOGICAL_PROJECTION) { + return false; + } + for (auto &expression : op.expressions) { + if (expression->IsVolatile() || expression->CanThrow()) { + return false; + } + } + return HasSafePredicates(*op.children[0]); +} + +template +static LogicalPlanSQLExportResult ExportContextExpressions(LogicalOperator &op, LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path, + vector fields, + EXPORTER exporter) { + D_ASSERT(op.children.size() == 1); + auto child = context.ExportChild(*op.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + if (child.HasError()) { + return LogicalPlanSQLExportResult::Failure(child); + } + LogicalPlanSQLExportHelpers::PropagateSemanticTypes(fields, {child.GetValue()}); + auto select = make_uniq(); + for (idx_t i = 0; i < child.GetValue().relation.fields.size(); i++) { + select->select_list.push_back(LogicalPlanSQLExportHelpers::ChildColumn(child.GetValue(), i)); + } + auto expression_context = + LogicalPlanSQLExportHelpers::CreateBindingContext(context.GetClientContext(), {child.GetValue()}); + for (idx_t i = 0; i < op.expressions.size(); i++) { + auto expression_path = LogicalPlanSQLExportHelpers::PlanExpressionPath(path, i); + auto expression = exporter(op.expressions[i]->Cast(), expression_context, expression_path); + if (expression.HasError()) { + return LogicalPlanSQLExportResult::Failure(expression); + } + expression.GetValue()->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(select->select_list.size())); + select->select_list.push_back(std::move(expression.GetValue())); + } + select->from_table = LogicalPlanSQLExportHelpers::CreateSubquery(std::move(child.GetValue())); + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(fields)}); +} + +LogicalPlanSQLExportResult LogicalFilter::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &filter = *this; + D_ASSERT(filter.children.size() == 1); +#ifdef D_ASSERT_IS_ENABLED + for (auto &expression : filter.expressions) { + D_ASSERT(expression && expression->GetReturnType() == LogicalType::BOOLEAN); + } +#endif + auto fields = LogicalPlanSQLExportHelpers::CreateFields(filter, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + auto child = export_context.ExportChild(*filter.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + if (child.HasError()) { + return LogicalPlanSQLExportResult::Failure(child); + } + LogicalPlanSQLExportHelpers::PropagateSemanticTypes(fields.GetValue(), {child.GetValue()}); + auto plain = LogicalPlanSQLExportHelpers::PlainScope(*child.GetValue().relation.query); + if (plain && plain->where_clause && !HasSafePredicates(*filter.children[0])) { + plain = nullptr; + } + if (plain && plain->where_clause) { + for (auto &predicate : filter.expressions) { + if (predicate->IsVolatile() || predicate->CanThrow()) { + plain = nullptr; + break; + } + } + } + vector> child_references {child.GetValue()}; + auto expression_context = + LogicalPlanSQLExportHelpers::CreateBindingContext(export_context.GetClientContext(), child_references, {plain}); + auto expressions = LogicalPlanSQLExportHelpers::CollectExpressions(filter); + vector> predicates; + for (idx_t expression_index = 0; expression_index < filter.expressions.size(); expression_index++) { + auto predicate = + export_context.ExportExpression(filter, expressions, expression_index, expression_context, path); + if (predicate.HasError()) { + return LogicalPlanSQLExportResult::Failure(predicate); + } + predicates.push_back(std::move(predicate.GetValue())); + } + auto select = make_uniq(); + for (idx_t field_index = 0; field_index < fields.GetValue().size(); field_index++) { + auto child_field_index = + filter.projection_map.empty() ? field_index : filter.projection_map[field_index].GetIndexUnsafe(); + auto expression = LogicalPlanSQLExportHelpers::ChildColumn(child.GetValue(), child_field_index, plain); + expression->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(field_index)); + select->select_list.push_back(std::move(expression)); + } + select->where_clause = SQLExportHelpers::Conjoin(std::move(predicates)); + LogicalPlanSQLExportHelpers::SetChildScope(*select, std::move(child.GetValue()), plain); + LogicalPlanSQLExportRelation relation {std::move(select), std::move(fields.GetValue())}; + return LogicalPlanSQLExportResult::Success(std::move(relation)); +} + +LogicalPlanSQLExportResult LogicalProjection::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &projection = *this; + D_ASSERT(projection.children.size() == 1 && !projection.expressions.empty()); + auto fields = LogicalPlanSQLExportHelpers::CreateFields(projection, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + D_ASSERT(projection.expressions.size() == fields.GetValue().size()); + bool fromless = projection.children[0]->type == LogicalOperatorType::LOGICAL_DUMMY_SCAN; + auto child = + export_context.ExportChild(*projection.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + if (child.HasError()) { + return LogicalPlanSQLExportResult::Failure(child); + } + if (LogicalPlanSQLExportHelpers::IsIdentityProjection(projection, child.GetValue().relation.fields)) { + return LogicalPlanSQLExportResult::Success( + {std::move(child.GetValue().relation.query), std::move(fields.GetValue())}); + } + auto plain = LogicalPlanSQLExportHelpers::PlainScope(*child.GetValue().relation.query); + vector> child_references {child.GetValue()}; + auto expression_context = + LogicalPlanSQLExportHelpers::CreateBindingContext(export_context.GetClientContext(), child_references, {plain}); + auto expressions = LogicalPlanSQLExportHelpers::CollectExpressions(projection); + auto select = make_uniq(); + for (idx_t expression_index = 0; expression_index < projection.expressions.size(); expression_index++) { + fromless = fromless && projection.expressions[expression_index]->IsScalar(); + ApplySemanticType(fields.GetValue()[expression_index], *projection.expressions[expression_index], + expression_context); + auto expression = + export_context.ExportExpression(projection, expressions, expression_index, expression_context, path); + if (expression.HasError()) { + return LogicalPlanSQLExportResult::Failure(expression); + } + expression.GetValue()->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(expression_index)); + select->select_list.push_back(std::move(expression.GetValue())); + } + if (fromless) { + select->from_table = make_uniq(); + } else { + LogicalPlanSQLExportHelpers::SetChildScope(*select, std::move(child.GetValue()), plain); + } + LogicalPlanSQLExportRelation relation {std::move(select), std::move(fields.GetValue())}; + return LogicalPlanSQLExportResult::Success(std::move(relation)); +} + +static LogicalPlanVerificationResult> +ExportOrderModifier(LogicalOperator &op, LogicalPlanSQLExportContext &context, const LogicalPlanVerificationPath &path, + const vector &orders, const BoundExpressionSQLExportContext &expression_context, + idx_t expression_ordinal) { + using Result = LogicalPlanVerificationResult>; + auto expressions = LogicalPlanSQLExportHelpers::CollectExpressions(op); + auto modifier = make_uniq(); + for (auto &order : orders) { + auto expression = context.ExportExpression(op, expressions, expression_ordinal++, expression_context, path); + if (expression.HasError()) { + return Result::Failure(expression); + } + modifier->orders.emplace_back( + order.type, order.null_order, + SQLExportHelpers::OrderExpression(SemanticExpressionType(*order.expression, expression_context), + std::move(expression.GetValue()))); + } + return Result::Success(std::move(modifier)); +} + +static LogicalPlanSQLExportResult ExportOrderedRelation(LogicalOperator &op, LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path, + const vector &orders) { + D_ASSERT(op.children.size() == 1); + auto fields = LogicalPlanSQLExportHelpers::CreateFields(op, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + auto child = context.ExportChild(*op.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + if (child.HasError()) { + return LogicalPlanSQLExportResult::Failure(child); + } + LogicalPlanSQLExportHelpers::PropagateSemanticTypes(fields.GetValue(), {child.GetValue()}); + auto expression_context = + LogicalPlanSQLExportHelpers::CreateBindingContext(context.GetClientContext(), {child.GetValue()}); + auto select = context.ForwardFields(child.GetValue(), fields.GetValue()); + auto modifier = ExportOrderModifier(op, context, path, orders, expression_context, 0); + if (modifier.HasError()) { + return LogicalPlanSQLExportResult::Failure(modifier); + } + select->modifiers.push_back(std::move(modifier.GetValue())); + select->from_table = LogicalPlanSQLExportHelpers::CreateSubquery(std::move(child.GetValue())); + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(fields.GetValue())}); +} + +LogicalPlanSQLExportResult LogicalSample::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &sample = *this; + D_ASSERT(sample.children.size() == 1 && sample.sample_options); + auto &sampling = *sample.sample_options; + if (sampling.seed.IsValid() != sampling.repeatable) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "sample_repeatability", "SQL sampling seeds imply repeatable sampling")}); + } + if (sampling.seed.IsValid() && sampling.seed.GetIndex() > idx_t(NumericLimits::Maximum())) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "sample_seed", "The sampling seed has no SQL spelling")}); + } + auto fields = LogicalPlanSQLExportHelpers::CreateFields(sample, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + auto child = export_context.ExportChild(*sample.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + if (child.HasError()) { + return LogicalPlanSQLExportResult::Failure(child); + } + auto select = export_context.ForwardFields(child.GetValue(), fields.GetValue()); + select->sample = sampling.Copy(); + select->from_table = LogicalPlanSQLExportHelpers::CreateSubquery(std::move(child.GetValue())); + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(fields.GetValue())}); +} + +LogicalPlanSQLExportResult LogicalSetOperation::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &op = *this; + D_ASSERT(!op.children.empty()); + D_ASSERT(op.type == LogicalOperatorType::LOGICAL_UNION || op.children.size() == 2); + auto fields = LogicalPlanSQLExportHelpers::CreateFields(op, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + auto query = make_uniq(); + query->setop_all = op.setop_all; + query->setop_type = op.type == LogicalOperatorType::LOGICAL_UNION ? SetOperationType::UNION + : op.type == LogicalOperatorType::LOGICAL_EXCEPT ? SetOperationType::EXCEPT + : SetOperationType::INTERSECT; + for (idx_t i = 0; i < op.children.size(); i++) { + auto child = export_context.Export(*op.children[i], LogicalPlanSQLExportHelpers::PlanChildPath(path, i)); + if (child.HasError()) { + return LogicalPlanSQLExportResult::Failure(child); + } + D_ASSERT(child.GetValue().fields.size() == fields.GetValue().size()); + query->children.push_back(std::move(child.GetValue().query)); + } + if (op.children.size() == 1) { + // Retain the UNION boundary when optimization removed every other arm. + LogicalEmptyResult empty(op.types, op.GetColumnBindings()); + empty.ResolveOperatorTypes(); + auto exported = empty.ToSQL(export_context, path); + if (exported.HasError()) { + return exported; + } + query->children.push_back(std::move(exported.GetValue().query)); + } + return LogicalPlanSQLExportResult::Success({std::move(query), std::move(fields.GetValue())}); +} + +LogicalPlanSQLExportResult LogicalAggregate::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &aggregate = *this; + D_ASSERT(aggregate.children.size() == 1); +#ifdef D_ASSERT_IS_ENABLED + for (auto &expression : aggregate.expressions) { + D_ASSERT(expression && expression->GetExpressionClass() == ExpressionClass::BOUND_AGGREGATE); + } +#endif + auto fields = LogicalPlanSQLExportHelpers::CreateFields(aggregate, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + D_ASSERT(aggregate.groups.size() + aggregate.expressions.size() + aggregate.grouping_functions.size() == + fields.GetValue().size()); + auto child = + export_context.ExportChild(*aggregate.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + if (child.HasError()) { + return LogicalPlanSQLExportResult::Failure(child); + } + auto plain = LogicalPlanSQLExportHelpers::PlainScope(*child.GetValue().relation.query); + vector> child_references {child.GetValue()}; + auto expression_context = + LogicalPlanSQLExportHelpers::CreateBindingContext(export_context.GetClientContext(), child_references, {plain}); + auto expressions = LogicalPlanSQLExportHelpers::CollectExpressions(aggregate); + for (idx_t group_index = 0; group_index < aggregate.groups.size(); group_index++) { + ApplySemanticType(fields.GetValue()[group_index], *aggregate.groups[group_index], expression_context); + } + auto input_field_count = child.GetValue().relation.fields.size(); + if (!aggregate.groups.empty()) { + // Equal group expressions can still occupy distinct grouping-set positions. + auto input = make_uniq(); + auto input_fields = child.GetValue().relation.fields; + for (idx_t i = 0; i < input_field_count; i++) { + auto expression = LogicalPlanSQLExportHelpers::ChildColumn(child.GetValue(), i, plain); + expression->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(i)); + input->select_list.push_back(std::move(expression)); + } + for (idx_t i = 0; i < aggregate.groups.size(); i++) { + auto expression = export_context.ExportExpression(aggregate, expressions, i, expression_context, path); + if (expression.HasError()) { + return LogicalPlanSQLExportResult::Failure(expression); + } + expression.GetValue()->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(input_fields.size())); + input->select_list.push_back(std::move(expression.GetValue())); + input_fields.push_back(fields.GetValue()[i]); + } + LogicalPlanSQLExportHelpers::SetChildScope(*input, std::move(child.GetValue()), plain); + child.GetValue() = {{std::move(input), std::move(input_fields)}, export_context.NextRelationAlias()}; + plain = nullptr; + expression_context = + LogicalPlanSQLExportHelpers::CreateBindingContext(export_context.GetClientContext(), {child.GetValue()}); + } + auto select = make_uniq(); + for (idx_t expression_index = 0; expression_index < expressions.size(); expression_index++) { + unique_ptr result; + if (expression_index < aggregate.groups.size()) { + result = LogicalPlanSQLExportHelpers::ChildColumn(child.GetValue(), input_field_count + expression_index); + select->groups.group_expressions.push_back(result->Copy()); + } else { + auto expression = + export_context.ExportExpression(aggregate, expressions, expression_index, expression_context, path); + if (expression.HasError()) { + return LogicalPlanSQLExportResult::Failure(expression); + } + result = std::move(expression.GetValue()); + } + result->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(expression_index)); + select->select_list.push_back(std::move(result)); + } + select->groups.grouping_sets = aggregate.grouping_sets; + if (select->groups.grouping_sets.empty() && !aggregate.groups.empty()) { + GroupingSet all_groups; + for (idx_t i = 0; i < aggregate.groups.size(); i++) { + all_groups.insert(ProjectionIndex(i)); + } + select->groups.grouping_sets.push_back(std::move(all_groups)); + } + for (auto &grouping : aggregate.grouping_functions) { + vector> arguments; + for (auto index : grouping) { + arguments.push_back(select->groups.group_expressions[index]->Copy()); + } + auto expression = make_uniq(ExpressionType::GROUPING_FUNCTION, std::move(arguments)); + expression->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(select->select_list.size())); + select->select_list.push_back(std::move(expression)); + } + LogicalPlanSQLExportHelpers::SetChildScope(*select, std::move(child.GetValue()), plain); + LogicalPlanSQLExportRelation relation {std::move(select), std::move(fields.GetValue())}; + return LogicalPlanSQLExportResult::Success(std::move(relation)); +} + +LogicalPlanSQLExportResult LogicalDistinct::ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) { + auto &op = *this; + D_ASSERT(op.children.size() == 1); + auto fields = LogicalPlanSQLExportHelpers::CreateFields(op, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + auto child = context.ExportChild(*op.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + if (child.HasError()) { + return LogicalPlanSQLExportResult::Failure(child); + } + LogicalPlanSQLExportHelpers::PropagateSemanticTypes(fields.GetValue(), {child.GetValue()}); + auto expression_context = + LogicalPlanSQLExportHelpers::CreateBindingContext(context.GetClientContext(), {child.GetValue()}); + auto select = context.ForwardFields(child.GetValue(), fields.GetValue()); + auto expressions = LogicalPlanSQLExportHelpers::CollectExpressions(op); + auto modifier = make_uniq(); + idx_t expression_ordinal = 0; + // DISTINCT targets include optimizer-pruned full-row DISTINCT. + for (idx_t i = 0; i < distinct_targets.size(); i++) { + auto expression = context.ExportExpression(op, expressions, expression_ordinal++, expression_context, path); + if (expression.HasError()) { + return LogicalPlanSQLExportResult::Failure(expression); + } + modifier->distinct_on_targets.push_back(std::move(expression.GetValue())); + } + select->modifiers.push_back(std::move(modifier)); + if (order_by) { + auto order = ExportOrderModifier(op, context, path, order_by->orders, expression_context, expression_ordinal); + if (order.HasError()) { + return LogicalPlanSQLExportResult::Failure(order); + } + select->modifiers.push_back(std::move(order.GetValue())); + } + select->from_table = LogicalPlanSQLExportHelpers::CreateSubquery(std::move(child.GetValue())); + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(fields.GetValue())}); +} + +LogicalPlanSQLExportResult LogicalOrder::ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) { + return ExportOrderedRelation(*this, context, path, orders); +} + +LogicalPlanSQLExportResult LogicalTopN::ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) { + auto result = ExportOrderedRelation(*this, context, path, orders); + if (result.HasError()) { + return result; + } + auto modifier = make_uniq(); + modifier->limit = ConstantExpression::FromValue(Value::BIGINT(NumericCast(limit))); + auto original_offset = unpruned_offset.IsValid() ? unpruned_offset.GetIndex() : offset; + modifier->offset = ConstantExpression::FromValue(Value::BIGINT(NumericCast(original_offset))); + result.GetValue().query->modifiers.push_back(std::move(modifier)); + return result; +} + +LogicalPlanSQLExportResult LogicalWindow::ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) { + auto &op = *this; + D_ASSERT(children.size() == 1); + auto fields = LogicalPlanSQLExportHelpers::CreateFields(op, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + if (op.children[0]->type == LogicalOperatorType::LOGICAL_GET) { + auto &get = op.children[0]->Cast(); + if (get.source_ordinality == OrdinalityType::WITH_ORDINALITY && !get.ordinality_idx.IsValid()) { + bool supported = + op.expressions.size() == 1 && !get.table_filters.HasFilters() && !get.extra_info.sample_options; + if (supported) { + auto &window = op.expressions[0]->Cast(); + supported = window.GetExpressionType() == ExpressionType::WINDOW_ROW_NUMBER && + window.Partitions().empty() && window.OrderBy().empty() && window.GetChildren().empty(); + } + if (!supported) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "ordinality_window", "The source ordinality cannot be reconstructed through this window")}); + } + return get.ExportSQLSource(context, LogicalPlanSQLExportHelpers::PlanChildPath(path, 0), + &fields.GetValue().back()); + } + } + return ExportContextExpressions(*this, context, path, std::move(fields.GetValue()), + BoundExpressionSQLExporter::ExportWindowAtPath); +} + +LogicalPlanSQLExportResult LogicalUnnest::ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) { + auto fields = LogicalPlanSQLExportHelpers::CreateFields(*this, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + return ExportContextExpressions(*this, context, path, std::move(fields.GetValue()), + BoundExpressionSQLExporter::ExportUnnestAtPath); +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/sql_export_scope.cpp b/src/duckdb/src/planner/sql_export/sql_export_scope.cpp new file mode 100644 index 000000000..e85e55d95 --- /dev/null +++ b/src/duckdb/src/planner/sql_export/sql_export_scope.cpp @@ -0,0 +1,270 @@ +#include "duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp" +#include "duckdb/planner/logical_plan_sql_exporter.hpp" +#include "duckdb/parser/expression/columnref_expression.hpp" +#include "duckdb/parser/query_node/select_node.hpp" +#include "duckdb/parser/statement/select_statement.hpp" +#include "duckdb/parser/tableref/joinref.hpp" +#include "duckdb/parser/tableref/subqueryref.hpp" +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/column_binding_map.hpp" +#include "duckdb/planner/operator/logical_get.hpp" +#include "duckdb/planner/expression/bound_columnref_expression.hpp" +#include "duckdb/planner/expression/bound_constant_expression.hpp" +#include "duckdb/planner/logical_operator.hpp" +#include "duckdb/planner/logical_operator_visitor.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" +#include "duckdb/planner/operator/logical_projection.hpp" + +namespace duckdb { + +LogicalPlanVerificationPath LogicalPlanSQLExportHelpers::PlanChildPath(const LogicalPlanVerificationPath &path, + idx_t ordinal) { + return SQLExportHelpers::ChildPath(path, ordinal, LogicalPlanVerificationPathComponentType::OPERATOR_CHILD); +} + +LogicalPlanVerificationPath LogicalPlanSQLExportHelpers::PlanExpressionPath(const LogicalPlanVerificationPath &path, + idx_t ordinal) { + return SQLExportHelpers::ChildPath(path, ordinal, LogicalPlanVerificationPathComponentType::OPERATOR_EXPRESSION); +} + +LogicalPlanVerificationResult> +LogicalPlanSQLExportHelpers::ExportTypedNull(const LogicalType &type, const LogicalPlanVerificationPath &path) { + auto result = BoundExpressionSQLExporter::Export(BoundConstantExpression(Value(type)), {}); + if (!result.HasError()) { + return result; + } + auto issues = result.GetIssues(); + for (auto &issue : issues) { + issue.path = path; + issue.phase = LogicalPlanVerificationPhase::PLAN_EXPORT; + } + return LogicalPlanVerificationResult>::Failure(std::move(issues)); +} + +LogicalPlanVerificationIssue +LogicalPlanSQLExportHelpers::PlanUnsupportedFeature(const LogicalPlanVerificationPath &path, string feature, + string message) { + return SQLExportHelpers::MakeIssue( + LogicalPlanVerificationIssueCode::UNSUPPORTED_EXPORT_FEATURE, LogicalPlanVerificationPhase::PLAN_EXPORT, path, + LogicalPlanVerificationConstructIdentity::ExportFeature(std::move(feature)), std::move(message)); +} + +LogicalPlanVerificationFunctionIdentity LogicalPlanSQLExportHelpers::LogicalSourceIdentity() { + LogicalPlanVerificationFunctionIdentity source; + source.name = "logical_source"; + source.return_type = LogicalType::TABLE; + return source; +} + +LogicalPlanVerificationFunctionIdentity LogicalPlanSQLExportHelpers::LogicalSourceIdentity(const LogicalGet &get) { + LogicalPlanVerificationFunctionIdentity source; + source.catalog = get.function.GetCatalogName().GetIdentifierName(); + source.schema = get.function.GetSchemaName().GetIdentifierName(); + source.name = get.function.GetName().GetIdentifierName(); + for (auto ¶meter : get.parameters) { + source.arguments.push_back(parameter.type()); + } + for (auto ¶meter : get.named_parameters) { + source.arguments.push_back(parameter.second.type()); + } + source.return_type = LogicalType::TABLE; + return source; +} + +LogicalPlanVerificationIssue +LogicalPlanSQLExportHelpers::UnsupportedSource(const LogicalPlanVerificationPath &path, + LogicalPlanVerificationFunctionIdentity source, string guard) { + auto issue = SQLExportHelpers::MakeIssue( + LogicalPlanVerificationIssueCode::UNSUPPORTED_SOURCE, LogicalPlanVerificationPhase::PLAN_EXPORT, path, + LogicalPlanVerificationConstructIdentity::SourceFunction(std::move(source)), + "The logical source does not expose structural SQL export semantics"); + issue.facts.emplace_back("guard", Value(std::move(guard))); + return issue; +} + +Identifier LogicalPlanSQLExportHelpers::FieldIdentifier(idx_t ordinal) { + return Identifier("c" + to_string(ordinal)); +} + +LogicalPlanSQLFieldResult LogicalPlanSQLExportHelpers::CreateFields(LogicalOperator &op, + const LogicalPlanVerificationPath &path) { + auto bindings = op.GetColumnBindings(); + D_ASSERT(bindings.size() == op.types.size()); + vector fields; + for (idx_t i = 0; i < bindings.size(); i++) { + D_ASSERT(bindings[i].table_index.IsValid() && bindings[i].column_index.IsValid()); + if (!SQLExportHelpers::IsSQLValueType(op.types[i])) { + auto issue = LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "output_type", "Output type cannot be represented in SQL"); + issue.facts.emplace_back("column_index", Value::UBIGINT(i)); + issue.facts.emplace_back("logical_type", Value(op.types[i].ToString())); + issue.facts.emplace_back("varchar_collations", + Value(SQLExportHelpers::TypeCollationSignature(op.types[i]))); + return LogicalPlanSQLFieldResult::Failure({std::move(issue)}); + } + fields.push_back({bindings[i], op.types[i]}); + } + return LogicalPlanSQLFieldResult::Success(std::move(fields)); +} + +BoundExpressionSQLExportContext +LogicalPlanSQLExportHelpers::CreateBindingContext(ClientContext &context, + const vector> &children, + const vector> &plain_scopes) { + D_ASSERT(plain_scopes.empty() || plain_scopes.size() == children.size()); + column_binding_map_t entries; + for (idx_t child_index = 0; child_index < children.size(); child_index++) { + auto &child = children[child_index].get(); + auto plain = plain_scopes.empty() ? nullptr : plain_scopes[child_index]; + for (idx_t i = 0; i < child.relation.fields.size(); i++) { + auto &field = child.relation.fields[i]; + auto names = + plain ? plain->select_list[i]->Cast().ColumnNames() + : vector {child.relation_alias, LogicalPlanSQLExportHelpers::FieldIdentifier(i)}; + entries.emplace(field.source_binding, + ResolvedSQLColumnReference {std::move(names), field.type, field.optimizer_type}); + } + } + BoundExpressionSQLExportContext result; + result.client_context = &context; + result.discard_optimizer_metadata = true; + result.resolve_binding = + [entries = std::move(entries)](const ColumnBinding &binding) -> optional { + auto entry = entries.find(binding); + if (entry != entries.end()) { + return entry->second; + } + return {}; + }; + return result; +} + +void LogicalPlanSQLExportHelpers::PropagateSemanticTypes( + vector &fields, const vector> &children) { + for (auto &field : fields) { + for (auto &child : children) { + for (auto &child_field : child.get().relation.fields) { + if (field.source_binding != child_field.source_binding || field.type == child_field.type) { + continue; + } + field.optimizer_type = field.type; + field.type = child_field.type; + } + } + } +} + +unique_ptr LogicalPlanSQLExportHelpers::CreateSubquery(LogicalPlanSQLExportedChild child) { + auto statement = make_uniq(); + statement->node = std::move(child.relation.query); + auto result = make_uniq(std::move(statement), std::move(child.relation_alias)); + for (idx_t i = 0; i < child.relation.fields.size(); i++) { + result->column_name_alias.push_back(LogicalPlanSQLExportHelpers::FieldIdentifier(i)); + } + return std::move(result); +} + +unique_ptr LogicalPlanSQLExportHelpers::ChildColumn(const LogicalPlanSQLExportedChild &child, + idx_t field_index, + optional_ptr plain) { + D_ASSERT(field_index < child.relation.fields.size()); + if (plain) { + return plain->select_list[field_index]->Copy(); + } + return make_uniq(LogicalPlanSQLExportHelpers::FieldIdentifier(field_index), + child.relation_alias); +} + +vector> LogicalPlanSQLExportHelpers::CollectExpressions(const LogicalOperator &op) { + vector> expressions; + LogicalOperatorVisitor::EnumerateExpressions(op, [&](const unique_ptr *expression) { + D_ASSERT(expression && *expression); + expressions.push_back(reference(**expression)); + }); + return expressions; +} + +bool LogicalPlanSQLExportHelpers::HasEffectfulExpressions(const LogicalOperator &op) { + for (auto &expression : LogicalPlanSQLExportHelpers::CollectExpressions(op)) { + if (expression.get().IsVolatile() || expression.get().CanThrow()) { + return true; + } + } + return false; +} + +bool LogicalPlanSQLExportHelpers::CollectScopeAliases(const TableRef &table, identifier_set_t &aliases) { + if (table.sample) { + return false; + } + if (!table.alias.empty()) { + return aliases.insert(table.alias).second; + } + if (table.type != TableReferenceType::JOIN) { + return false; + } + auto &join = table.Cast(); + return join.left && join.right && LogicalPlanSQLExportHelpers::CollectScopeAliases(*join.left, aliases) && + LogicalPlanSQLExportHelpers::CollectScopeAliases(*join.right, aliases); +} + +optional_ptr LogicalPlanSQLExportHelpers::PlainScope(const QueryNode &query) { + if (query.type != QueryNodeType::SELECT_NODE || !query.modifiers.empty() || !query.cte_map.map.empty()) { + return nullptr; + } + auto &select = query.Cast(); + if (!select.from_table || select.sample || select.from_table->sample) { + return nullptr; + } + const bool has_groups = !select.groups.group_expressions.empty() || !select.groups.grouping_sets.empty(); + const bool has_aggregate_handling = select.aggregate_handling != AggregateHandling::STANDARD_HANDLING; + if (has_groups || select.having || select.qualify || has_aggregate_handling) { + return nullptr; + } + identifier_set_t aliases; + if (!LogicalPlanSQLExportHelpers::CollectScopeAliases(*select.from_table, aliases)) { + return nullptr; + } + for (auto &expression : select.select_list) { + if (expression->GetExpressionClass() != ExpressionClass::COLUMN_REF) { + return nullptr; + } + auto &names = expression->Cast().ColumnNames(); + if (names.size() != 2 || !aliases.count(names[0])) { + return nullptr; + } + } + return select; +} + +void LogicalPlanSQLExportHelpers::SetChildScope(SelectNode &select, LogicalPlanSQLExportedChild child, + optional_ptr plain) { + if (!plain) { + select.from_table = LogicalPlanSQLExportHelpers::CreateSubquery(std::move(child)); + return; + } + auto &source = child.relation.query->Cast(); + select.from_table = std::move(source.from_table); + select.where_clause = SQLExportHelpers::Conjoin(std::move(source.where_clause), std::move(select.where_clause)); +} + +bool LogicalPlanSQLExportHelpers::IsIdentityProjection(const LogicalProjection &projection, + const vector &fields) { + if (projection.expressions.size() != fields.size()) { + return false; + } + for (idx_t i = 0; i < fields.size(); i++) { + auto &expression = *projection.expressions[i]; + if (expression.GetExpressionClass() != ExpressionClass::BOUND_COLUMN_REF) { + return false; + } + auto &column = expression.Cast(); + if (column.Depth() != 0 || column.Binding() != fields[i].source_binding || + !column.GetReturnType().EqualsIncludingCollation(fields[i].type)) { + return false; + } + } + return true; +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/sql_export_sources.cpp b/src/duckdb/src/planner/sql_export/sql_export_sources.cpp new file mode 100644 index 000000000..6a92ee17c --- /dev/null +++ b/src/duckdb/src/planner/sql_export/sql_export_sources.cpp @@ -0,0 +1,260 @@ +#include "duckdb/planner/operator/logical_extension_operator.hpp" +#include "duckdb/planner/operator/logical_get.hpp" +#include "duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp" +#include "duckdb/planner/logical_plan_sql_exporter.hpp" +#include "duckdb/common/limits.hpp" +#include "duckdb/parser/expression/conjunction_expression.hpp" +#include "duckdb/parser/expression/constant_expression.hpp" +#include "duckdb/parser/tableref/basetableref.hpp" +#include "duckdb/parser/tableref/at_clause.hpp" +#include "duckdb/parser/query_node/select_node.hpp" +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/filter/expression_filter.hpp" +#include "duckdb/planner/operator/logical_get.hpp" +#include "duckdb/planner/expression/bound_columnref_expression.hpp" +#include "duckdb/planner/operator/logical_limit.hpp" +#include "duckdb/planner/operator/logical_top_n.hpp" +#include "duckdb/planner/operator/logical_extension_operator.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" +#include "duckdb/planner/operator/logical_secure_view.hpp" + +namespace duckdb { + +static LogicalPlanVerificationIssue ExtensionIssue(LogicalPlanVerificationIssueCode code, + const LogicalPlanVerificationPath &path, const string &identifier, + string message) { + return SQLExportHelpers::MakeIssue(code, LogicalPlanVerificationPhase::PLAN_EXPORT, path, + LogicalPlanVerificationConstructIdentity::Extension(identifier), + std::move(message)); +} + +LogicalPlanSQLExportResult LogicalGet::ExportSQLSource(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path, + optional_ptr ordinality) { + auto &get = *this; + D_ASSERT(get.children.size() <= 1); + if (get.has_pushed_projection) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::UnsupportedSource( + path, LogicalPlanSQLExportHelpers::LogicalSourceIdentity(get), "pushed_projection")}); + } + auto fields = LogicalPlanSQLExportHelpers::CreateFields(get, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + if ((!get.extra_info.file_filters.empty() || get.extra_info.total_files.IsValid()) && + !get.extra_info.file_filter_expressions) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "file_filter_residual", "The source does not retain the SQL predicate used for file pruning")}); + } + if (get.row_group_order_options && + (get.row_group_order_options->row_group_offset || get.row_group_order_options->leading_null_group_offset)) { + bool has_unpruned_offset = false; + for (idx_t i = export_context.Ancestors().size(); i > 0; i--) { + auto &ancestor = export_context.Ancestors()[i - 1].get(); + if (ancestor.type == LogicalOperatorType::LOGICAL_LIMIT) { + has_unpruned_offset = ancestor.Cast().unpruned_offset.IsValid(); + break; + } else if (ancestor.type == LogicalOperatorType::LOGICAL_TOP_N) { + has_unpruned_offset = ancestor.Cast().unpruned_offset.IsValid(); + break; + } + } + if (!has_unpruned_offset) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "pruned_offset", "The row group pruning does not retain its original SQL offset")}); + } + } + + vector scan_fields; + for (idx_t i = 0; i < get.GetColumnIds().size(); i++) { + auto &type = get.GetColumnType(get.GetColumnIds()[i]); + if (!SQLExportHelpers::IsSQLValueType(type)) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "scan_type", "Scan type cannot be represented in SQL")}); + } + scan_fields.push_back({ColumnBinding(get.table_index, ProjectionIndex(i)), type}); + } + unique_ptr input; + if (!get.children.empty()) { + auto child = export_context.ExportChild(*get.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + if (child.HasError()) { + return LogicalPlanSQLExportResult::Failure(child); + } + for (auto index : get.projected_input) { + scan_fields.push_back(child.GetValue().relation.fields[index]); + } + input = LogicalPlanSQLExportHelpers::CreateSubquery(std::move(child.GetValue())); + } + if (ordinality) { + fields.GetValue().push_back(*ordinality); + scan_fields.push_back(*ordinality); + } + auto relation_alias = export_context.NextRelationAlias(); + auto source_sql = LogicalPlanSQLExportHelpers::ReconstructSQLSource( + export_context.GetClientContext(), get, std::move(input), relation_alias, bool(ordinality)); + if (!source_sql.query) { + auto guard = source_sql.unsupported_reason.empty() ? "to_sql_callback_declined" : source_sql.unsupported_reason; + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::UnsupportedSource( + path, LogicalPlanSQLExportHelpers::LogicalSourceIdentity(get), std::move(guard))}); + } + auto query = std::move(source_sql.query); + if (get.extra_info.sample_options) { + auto sampling = get.extra_info.sample_options->Copy(); + if (!sampling->repeatable && sampling->seed.IsValid()) { + sampling->seed = optional_idx::Invalid(); + } + if (sampling->repeatable && !sampling->seed.IsValid()) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "sample_repeatability", "SQL sampling seeds imply repeatable sampling")}); + } + if (sampling->seed.IsValid() && sampling->seed.GetIndex() > idx_t(NumericLimits::Maximum())) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "sample_seed", "The sampling seed has no SQL spelling")}); + } + sampling->sample_rate = -1.0; + LogicalPlanSQLExportedChild unsampled {{std::move(query), scan_fields}, export_context.NextRelationAlias()}; + auto sampled = export_context.ForwardFields(unsampled, scan_fields); + sampled->sample = std::move(sampling); + sampled->from_table = LogicalPlanSQLExportHelpers::CreateSubquery(std::move(unsampled)); + query = std::move(sampled); + } + LogicalPlanSQLExportedChild source {{std::move(query), std::move(scan_fields)}, std::move(relation_alias)}; + auto plain = LogicalPlanSQLExportHelpers::PlainScope(*source.relation.query); + if (plain && (plain->where_clause || plain->select_list.size() != source.relation.fields.size())) { + plain = nullptr; + } + auto binding_context = + LogicalPlanSQLExportHelpers::CreateBindingContext(export_context.GetClientContext(), {source}, {plain}); + auto select = export_context.ForwardFields(source, fields.GetValue(), plain); + vector> predicates; + for (auto &entry : get.table_filters) { + if (ExpressionFilter::IsOptionalFilter(entry.Filter())) { + continue; + } + auto &field = source.relation.fields[entry.GetIndex().GetIndex()]; + BoundColumnRefExpression column(field.type, field.source_binding); + predicates.push_back(entry.Filter().ToExpression(column)); + } + auto conjunction = make_uniq(ExpressionType::CONJUNCTION_AND); + auto ordinal = LogicalPlanSQLExportHelpers::CollectExpressions(get).size(); + for (auto &predicate : predicates) { + auto exported = BoundExpressionSQLExporter::ExportAtPath( + *predicate, binding_context, LogicalPlanSQLExportHelpers::PlanExpressionPath(path, ordinal++)); + if (exported.HasError()) { + return LogicalPlanSQLExportResult::Failure(exported); + } + conjunction->AddExpression(std::move(exported.GetValue())); + } + if (conjunction->GetChildren().size() == 1) { + select->where_clause = std::move(conjunction->GetChildrenMutable()[0]); + } else if (!conjunction->GetChildren().empty()) { + select->where_clause = std::move(conjunction); + } + LogicalPlanSQLExportHelpers::SetChildScope(*select, std::move(source), plain); + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(fields.GetValue())}); +} + +LogicalPlanSQLExportResult LogicalSecureView::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &view = *this; + auto fields = LogicalPlanSQLExportHelpers::CreateFields(view, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + if (!view.has_source || view.source_name.Path().empty() || view.source_types.empty()) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "secure_view_source", "The secure view does not retain its qualified source metadata")}); + } + for (auto &component : view.source_name.Path()) { + if (component.empty()) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "secure_view_source", "The secure view source name cannot be represented in SQL")}); + } + } + if (view.output_bindings.size() != fields.GetValue().size() || + view.output_expressions.size() != fields.GetValue().size()) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "secure_view_output", "The secure view output mapping is incomplete")}); + } + for (auto &type : view.source_types) { + if (!SQLExportHelpers::IsSQLRepresentableType(type)) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "secure_view_source", "A secure view source type cannot be represented in SQL")}); + } + } + + if (view.source_filters.size() != view.pushed_filters.size()) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "secure_view_filter", "The secure view does not retain every caller predicate")}); + } + + auto source_alias = export_context.NextRelationAlias(); + auto table = make_uniq(); + table->SetQualifiedName(view.source_name); + table->alias = source_alias; + for (idx_t i = 0; i < view.source_types.size(); i++) { + table->column_name_alias.push_back(LogicalPlanSQLExportHelpers::FieldIdentifier(i)); + } + if (view.has_at_clause) { + table->at_clause = make_uniq(view.at_unit, ConstantExpression::FromValue(view.at_value)); + } + + BoundExpressionSQLExportContext expression_context; + expression_context.client_context = &export_context.GetClientContext(); + expression_context.resolve_binding = [source_alias, source_types = view.source_types]( + const ColumnBinding &binding) -> optional { + if (binding.table_index != TableIndex(0) || binding.column_index.GetIndex() >= source_types.size()) { + return {}; + } + auto index = binding.column_index.GetIndex(); + return ResolvedSQLColumnReference {{source_alias, LogicalPlanSQLExportHelpers::FieldIdentifier(index)}, + source_types[index]}; + }; + + auto select = make_uniq(); + select->from_table = std::move(table); + for (idx_t i = 0; i < fields.GetValue().size(); i++) { + if (view.output_bindings[i] != fields.GetValue()[i].source_binding || + !view.output_expressions[i]->GetReturnType().EqualsIncludingCollation(fields.GetValue()[i].type)) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "secure_view_output", "The secure view output mapping does not match its current schema")}); + } + auto expression = BoundExpressionSQLExporter::ExportAtPath( + *view.output_expressions[i], expression_context, LogicalPlanSQLExportHelpers::PlanExpressionPath(path, i)); + if (expression.HasError()) { + return LogicalPlanSQLExportResult::Failure(expression); + } + expression.GetValue()->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(i)); + select->select_list.push_back(std::move(expression.GetValue())); + } + for (idx_t i = 0; i < view.source_filters.size(); i++) { + if (!view.source_filters[i]) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "secure_view_filter", "The secure view caller predicate has no complete source mapping")}); + } + auto predicate = BoundExpressionSQLExporter::ExportAtPath( + *view.source_filters[i], expression_context, + LogicalPlanSQLExportHelpers::PlanExpressionPath(path, fields.GetValue().size() + i)); + if (predicate.HasError()) { + return LogicalPlanSQLExportResult::Failure(predicate); + } + select->where_clause = + SQLExportHelpers::Conjoin(std::move(select->where_clause), std::move(predicate.GetValue())); + } + + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(fields.GetValue())}); +} + +LogicalPlanVerificationResult LogicalGet::ToSQL(LogicalPlanSQLExportContext &context, + const LogicalPlanVerificationPath &path) { + return ExportSQLSource(context, path); +} + +LogicalPlanVerificationResult +LogicalExtensionOperator::ToSQL(LogicalPlanSQLExportContext &context, const LogicalPlanVerificationPath &path) { + return LogicalPlanSQLExportResult::Failure( + {ExtensionIssue(LogicalPlanVerificationIssueCode::UNSUPPORTED_EXTENSION, path, GetExtensionName(), + "The extension operator does not implement SQL reconstruction")}); +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/sql_export_values.cpp b/src/duckdb/src/planner/sql_export/sql_export_values.cpp new file mode 100644 index 000000000..87ac58933 --- /dev/null +++ b/src/duckdb/src/planner/sql_export/sql_export_values.cpp @@ -0,0 +1,454 @@ +#include "duckdb/planner/operator/logical_empty_result.hpp" +#include "duckdb/planner/operator/logical_dummy_scan.hpp" +#include "duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp" +#include "duckdb/planner/operator/logical_join.hpp" +#include "duckdb/function/scalar/compressed_materialization_utils.hpp" +#include "duckdb/planner/logical_plan_sql_exporter.hpp" +#include "duckdb/main/settings.hpp" +#include "duckdb/parser/expression/case_expression.hpp" +#include "duckdb/parser/expression/columnref_expression.hpp" +#include "duckdb/parser/expression/comparison_expression.hpp" +#include "duckdb/parser/expression/constant_expression.hpp" +#include "duckdb/parser/expression/operator_expression.hpp" +#include "duckdb/parser/expression/function_expression.hpp" +#include "duckdb/parser/query_node/select_node.hpp" +#include "duckdb/parser/statement/select_statement.hpp" +#include "duckdb/parser/tableref/emptytableref.hpp" +#include "duckdb/parser/tableref/expressionlistref.hpp" +#include "duckdb/parser/tableref/joinref.hpp" +#include "duckdb/parser/tableref/subqueryref.hpp" +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/expression/bound_columnref_expression.hpp" +#include "duckdb/planner/expression/bound_constant_expression.hpp" +#include "duckdb/planner/expression/bound_function_expression.hpp" +#include "duckdb/planner/expression/bound_subquery_expression.hpp" +#include "duckdb/planner/expression/bound_window_expression.hpp" +#include "duckdb/planner/expression/bound_aggregate_expression.hpp" +#include "duckdb/function/lambda_functions.hpp" +#include "duckdb/planner/expression_iterator.hpp" +#include "duckdb/planner/logical_operator.hpp" +#include "duckdb/planner/operator/logical_column_data_get.hpp" +#include "duckdb/planner/operator/logical_empty_result.hpp" +#include "duckdb/planner/operator/logical_distinct.hpp" +#include "duckdb/planner/operator/logical_order.hpp" +#include "duckdb/planner/operator/logical_top_n.hpp" +#include "duckdb/planner/operator/logical_expression_get.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" + +namespace duckdb { + +//! Whether the expression is volatile or can throw, ignoring compressed materialization wrappers +static bool IsEffectfulSemanticExpression(const Expression &expression) { + if (auto wrapped = CMUtils::GetWrappedInput(expression)) { + return IsEffectfulSemanticExpression(*wrapped); + } + bool effectful = false; + ExpressionIterator::EnumerateChildren( + expression, [&](const Expression &child) { effectful = effectful || IsEffectfulSemanticExpression(child); }); + if (effectful) { + return true; + } + switch (expression.GetExpressionClass()) { + case ExpressionClass::BOUND_FUNCTION: { + auto &bound_function = expression.Cast(); + auto &function = bound_function.Function(); + if (function.GetStability() == FunctionStability::VOLATILE || + function.GetErrorMode() == FunctionErrors::CAN_THROW_RUNTIME_ERROR) { + return true; + } + if (!function.HasBindLambdaCallback()) { + return false; + } + auto lambda = bound_function.BindInfo()->Cast().GetLambdaExpression(); + return lambda && (lambda->IsVolatile() || lambda->CanThrow()); + } + case ExpressionClass::BOUND_AGGREGATE: + return expression.Cast().Function().GetStability() == FunctionStability::VOLATILE; + case ExpressionClass::BOUND_WINDOW: { + auto &window = expression.Cast(); + auto stability = window.AggregateFunction() ? window.AggregateFunction()->GetStability() + : window.WindowFunction()->GetStability(); + return stability == FunctionStability::VOLATILE; + } + case ExpressionClass::BOUND_SUBQUERY: { + auto &plan = expression.Cast().Subquery().plan; + return plan && plan->HasVolatileExpressions(); + } + default: + return false; + } +} + +static bool HasEffectfulExpressionSubtree(const LogicalOperator &op) { + for (auto &expression : LogicalPlanSQLExportHelpers::CollectExpressions(op)) { + if (IsEffectfulSemanticExpression(expression.get())) { + return true; + } + } + for (auto &child : op.children) { + if (HasEffectfulExpressionSubtree(*child)) { + return true; + } + } + return false; +} + +static bool OrderDistinguishesValues(ClientContext &context, const LogicalType &type) { + auto default_collation = Settings::Get(context); + return !TypeVisitor::Contains(type, [&](const LogicalType &child) { + if (child.IsIntegral()) { + return false; + } + switch (child.id()) { + case LogicalTypeId::BOOLEAN: + case LogicalTypeId::DECIMAL: + case LogicalTypeId::DATE: + case LogicalTypeId::TIME: + case LogicalTypeId::TIMESTAMP: + case LogicalTypeId::TIMESTAMP_SEC: + case LogicalTypeId::TIMESTAMP_MS: + case LogicalTypeId::TIMESTAMP_NS: + case LogicalTypeId::TIMESTAMP_TZ: + case LogicalTypeId::BLOB: + case LogicalTypeId::BIT: + case LogicalTypeId::UUID: + case LogicalTypeId::ENUM: + case LogicalTypeId::LIST: + case LogicalTypeId::ARRAY: + case LogicalTypeId::STRUCT: + case LogicalTypeId::MAP: + return false; + case LogicalTypeId::VARCHAR: + return !StringType::GetCollation(child).empty() || !default_collation.empty(); + default: + return true; + } + }); +} + +static bool OrdersAllFields(ClientContext &context, const vector &orders, LogicalOperator &op) { + auto bindings = op.GetColumnBindings(); + for (idx_t i = 0; i < bindings.size(); i++) { + if (!OrderDistinguishesValues(context, op.types[i])) { + return false; + } + bool found = false; + for (auto &order : orders) { + if (order.expression->GetExpressionClass() == ExpressionClass::BOUND_COLUMN_REF && + order.expression->Cast().Binding() == bindings[i]) { + found = true; + break; + } + } + if (!found) { + return false; + } + } + return true; +} + +static bool HasChunkSensitiveConsumer(ClientContext &context, LogicalOperator &op, bool single_join_errors, + bool order_required = false) { + if (LogicalPlanSQLExportHelpers::HasEffectfulExpressions(op)) { + return true; + } + switch (op.type) { + case LogicalOperatorType::LOGICAL_COMPARISON_JOIN: + case LogicalOperatorType::LOGICAL_ANY_JOIN: + case LogicalOperatorType::LOGICAL_ASOF_JOIN: + case LogicalOperatorType::LOGICAL_DELIM_JOIN: + if (order_required || (single_join_errors && op.Cast().join_type == JoinType::SINGLE)) { + return true; + } + break; + case LogicalOperatorType::LOGICAL_CROSS_PRODUCT: + case LogicalOperatorType::LOGICAL_POSITIONAL_JOIN: + if (order_required) { + return true; + } + break; + case LogicalOperatorType::LOGICAL_LIMIT: + order_required = true; + break; + case LogicalOperatorType::LOGICAL_TOP_N: + order_required = !OrdersAllFields(context, op.Cast().orders, op); + break; + case LogicalOperatorType::LOGICAL_ORDER_BY: + if (OrdersAllFields(context, op.Cast().orders, op)) { + order_required = false; + } + break; + case LogicalOperatorType::LOGICAL_DISTINCT: + order_required |= op.Cast().distinct_type == DistinctType::DISTINCT_ON; + break; + case LogicalOperatorType::LOGICAL_GET: + if (!op.children.empty()) { + return true; + } + break; + case LogicalOperatorType::LOGICAL_PROJECTION: + case LogicalOperatorType::LOGICAL_FILTER: + case LogicalOperatorType::LOGICAL_EXPRESSION_GET: + case LogicalOperatorType::LOGICAL_UNION: + case LogicalOperatorType::LOGICAL_EXCEPT: + case LogicalOperatorType::LOGICAL_INTERSECT: + case LogicalOperatorType::LOGICAL_UNNEST: + case LogicalOperatorType::LOGICAL_CHUNK_GET: + case LogicalOperatorType::LOGICAL_DUMMY_SCAN: + case LogicalOperatorType::LOGICAL_EMPTY_RESULT: + break; + default: + return true; + } + for (auto &child : op.children) { + if (HasChunkSensitiveConsumer(context, *child, single_join_errors, order_required)) { + return true; + } + } + return false; +} + +LogicalPlanSQLExportResult LogicalColumnDataGet::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &get = *this; + D_ASSERT(get.children.empty() && get.collection); + if (!get.collection.is_owned()) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::UnsupportedSource( + path, LogicalPlanSQLExportHelpers::LogicalSourceIdentity(), "borrowed_chunk_collection")}); + } + if (get.collection->Count() == 0) { + LogicalEmptyResult empty(get.types, get.GetColumnBindings()); + empty.ResolveOperatorTypes(); + return empty.ToSQL(export_context, path); + } + auto fields = LogicalPlanSQLExportHelpers::CreateFields(get, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + auto values = make_uniq(); + values->alias = export_context.NextRelationAlias(); + values->expected_types = get.types; + for (idx_t i = 0; i < fields.GetValue().size(); i++) { + values->expected_names.push_back(LogicalPlanSQLExportHelpers::FieldIdentifier(i)); + } + idx_t remaining = get.collection->Count(); + bool repacked = false; + for (auto &chunk : get.collection->Chunks(get.GetColumnIds())) { + repacked |= chunk.size() != MinValue(remaining, STANDARD_VECTOR_SIZE); + remaining -= chunk.size(); + for (idx_t row = 0; row < chunk.size(); row++) { + vector> exported_row; + for (idx_t column = 0; column < chunk.ColumnCount(); column++) { + auto value = + BoundExpressionSQLExporter::Export(BoundConstantExpression(chunk.GetValue(column, row)), {}); + if (value.HasError()) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "chunk_value", "The materialized value cannot be represented in SQL")}); + } + exported_row.push_back(std::move(value.GetValue())); + } + values->values.push_back(std::move(exported_row)); + } + } + if (repacked && HasChunkSensitiveConsumer( + export_context.GetClientContext(), export_context.Ancestors().front().get(), + Settings::Get(export_context.GetClientContext()))) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "chunk_consumer_evaluation", "SQL cannot retain source chunks for an effectful consumer")}); + } + auto select = make_uniq(); + for (idx_t i = 0; i < fields.GetValue().size(); i++) { + select->select_list.push_back( + make_uniq(LogicalPlanSQLExportHelpers::FieldIdentifier(i), values->alias)); + } + select->from_table = std::move(values); + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(fields.GetValue())}); +} + +LogicalPlanSQLExportResult LogicalExpressionGet::ToSQL(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path) { + auto &get = *this; + D_ASSERT(get.children.size() == 1 && get.children[0]); + D_ASSERT(!get.expressions.empty() && !get.expressions[0].empty()); +#ifdef D_ASSERT_IS_ENABLED + auto column_count = get.expressions[0].size(); + D_ASSERT(get.expr_types.size() == column_count); + for (auto &row : get.expressions) { + D_ASSERT(row.size() == column_count); + for (idx_t i = 0; i < column_count; i++) { + D_ASSERT(row[i] && row[i]->GetReturnType() == get.expr_types[i]); + } + } +#endif + auto fields = LogicalPlanSQLExportHelpers::CreateFields(get, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + if (get.children[0]->type != LogicalOperatorType::LOGICAL_DUMMY_SCAN) { + return ExportSQLInput(export_context, path, std::move(fields.GetValue())); + } + + auto values = make_uniq(); + values->alias = export_context.NextRelationAlias(); + values->expected_types = get.expr_types; + for (idx_t i = 0; i < fields.GetValue().size(); i++) { + values->expected_names.push_back(LogicalPlanSQLExportHelpers::FieldIdentifier(i)); + } + BoundExpressionSQLExportContext expression_context; + expression_context.client_context = &export_context.GetClientContext(); + auto expressions = LogicalPlanSQLExportHelpers::CollectExpressions(get); + idx_t expression_ordinal = 0; + for (auto &row : get.expressions) { + vector> exported_row; + for (idx_t column_index = 0; column_index < row.size(); column_index++) { + auto expression = + export_context.ExportExpression(get, expressions, expression_ordinal++, expression_context, path); + if (expression.HasError()) { + return LogicalPlanSQLExportResult::Failure(expression); + } + exported_row.push_back(std::move(expression.GetValue())); + } + values->values.push_back(std::move(exported_row)); + } + auto values_alias = values->alias; + auto select = make_uniq(); + for (idx_t i = 0; i < fields.GetValue().size(); i++) { + auto expression = make_uniq(LogicalPlanSQLExportHelpers::FieldIdentifier(i), values_alias); + select->select_list.push_back(std::move(expression)); + } + select->from_table = std::move(values); + LogicalPlanSQLExportRelation relation {std::move(select), std::move(fields.GetValue())}; + return LogicalPlanSQLExportResult::Success(std::move(relation)); +} + +LogicalPlanSQLExportResult LogicalExpressionGet::ExportSQLInput(LogicalPlanSQLExportContext &export_context, + const LogicalPlanVerificationPath &path, + vector fields) { + auto &get = *this; + if (HasEffectfulExpressionSubtree(get)) { + return LogicalPlanSQLExportResult::Failure({LogicalPlanSQLExportHelpers::PlanUnsupportedFeature( + path, "values_expression_evaluation", "VALUES with input requires nonvolatile, nonthrowing expressions")}); + } + auto child = export_context.ExportChild(*get.children[0], LogicalPlanSQLExportHelpers::PlanChildPath(path, 0)); + if (child.HasError()) { + return LogicalPlanSQLExportResult::Failure(child); + } + auto expression_context = + LogicalPlanSQLExportHelpers::CreateBindingContext(export_context.GetClientContext(), {child.GetValue()}); + auto expressions = LogicalPlanSQLExportHelpers::CollectExpressions(get); + auto select = make_uniq(); + select->from_table = LogicalPlanSQLExportHelpers::CreateSubquery(std::move(child.GetValue())); + + auto row_alias = export_context.NextRelationAlias(); + auto cases = make_uniq(); + auto rows = make_uniq(); + rows->alias = row_alias; + rows->expected_names.push_back(LogicalPlanSQLExportHelpers::FieldIdentifier(0)); + rows->expected_types.push_back(LogicalType::BIGINT); + for (idx_t row = 0; row < get.expressions.size(); row++) { + // Keep all expressions of a VALUES row in one evaluation group. + vector arguments; + for (idx_t column = 0; column < fields.size(); column++) { + auto expression = export_context.ExportExpression(get, expressions, row * fields.size() + column, + expression_context, path); + if (expression.HasError()) { + return LogicalPlanSQLExportResult::Failure(expression); + } + arguments.emplace_back(LogicalPlanSQLExportHelpers::FieldIdentifier(column), + std::move(expression.GetValue())); + } + auto value = SQLExportHelpers::SystemFunction("struct_pack", std::move(arguments)); + if (row + 1 == get.expressions.size()) { + cases->ElseMutable() = std::move(value); + } else { + CaseCheck check; + check.when_expr = make_uniq( + ExpressionType::COMPARE_EQUAL, + make_uniq(LogicalPlanSQLExportHelpers::FieldIdentifier(0), row_alias), + ConstantExpression::FromValue(Value::BIGINT(int64_t(row)))); + check.then_expr = std::move(value); + cases->CaseChecksMutable().push_back(std::move(check)); + } + vector> ordinal; + ordinal.push_back(ConstantExpression::FromValue(Value::BIGINT(int64_t(row)))); + rows->values.push_back(std::move(ordinal)); + } + unique_ptr value; + if (get.expressions.size() == 1) { + value = std::move(cases->ElseMutable()); + } else { + auto join = make_uniq(JoinRefType::CROSS); + join->left = std::move(select->from_table); + join->right = std::move(rows); + select->from_table = std::move(join); + value = std::move(cases); + } + return export_context.ExportRow(std::move(select), std::move(value), std::move(fields)); +} + +LogicalPlanSQLExportResult LogicalPlanSQLExportContext::ExportRow(unique_ptr select, + unique_ptr value, + vector fields) { + // Keep row evaluation below consumers that can filter or limit emitted rows. + vector> list_arguments; + list_arguments.push_back(std::move(value)); + auto list = SQLExportHelpers::SystemFunction("list_value", std::move(list_arguments)); + vector> unnest_arguments; + unnest_arguments.push_back(std::move(list)); + auto unnest = make_uniq(Identifier("unnest"), std::move(unnest_arguments)); + unnest->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(0)); + select->select_list.push_back(std::move(unnest)); + + auto statement = make_uniq(); + statement->node = std::move(select); + auto packed_alias = NextRelationAlias(); + auto packed = make_uniq(std::move(statement), packed_alias); + packed->column_name_alias.push_back(LogicalPlanSQLExportHelpers::FieldIdentifier(0)); + select = make_uniq(); + select->from_table = std::move(packed); + for (idx_t column = 0; column < fields.size(); column++) { + auto value = make_uniq( + ExpressionType::STRUCT_EXTRACT, + make_uniq(LogicalPlanSQLExportHelpers::FieldIdentifier(0), packed_alias), + ConstantExpression::FromValue( + Value(LogicalPlanSQLExportHelpers::FieldIdentifier(column).GetIdentifierName()))); + value->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(column)); + select->select_list.push_back(std::move(value)); + } + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(fields)}); +} + +static LogicalPlanSQLExportResult ExportConstantSource(LogicalOperator &op, bool has_rows, + const LogicalPlanVerificationPath &path) { + D_ASSERT(op.children.empty() && op.expressions.empty()); + auto fields = LogicalPlanSQLExportHelpers::CreateFields(op, path); + if (fields.HasError()) { + return LogicalPlanSQLExportResult::Failure(fields); + } + auto select = make_uniq(); + select->from_table = make_uniq(); + for (idx_t i = 0; i < fields.GetValue().size(); i++) { + auto expression = LogicalPlanSQLExportHelpers::ExportTypedNull(fields.GetValue()[i].type, path); + if (expression.HasError()) { + return LogicalPlanSQLExportResult::Failure(expression); + } + expression.GetValue()->SetAlias(LogicalPlanSQLExportHelpers::FieldIdentifier(i)); + select->select_list.push_back(std::move(expression.GetValue())); + } + if (!has_rows) { + select->where_clause = ConstantExpression::FromValue(Value::BOOLEAN(false)); + } + return LogicalPlanSQLExportResult::Success({std::move(select), std::move(fields.GetValue())}); +} + +LogicalPlanVerificationResult +LogicalDummyScan::ToSQL(LogicalPlanSQLExportContext &context, const LogicalPlanVerificationPath &path) { + return ExportConstantSource(*this, true, path); +} + +LogicalPlanVerificationResult +LogicalEmptyResult::ToSQL(LogicalPlanSQLExportContext &context, const LogicalPlanVerificationPath &path) { + return ExportConstantSource(*this, false, path); +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/sql_export_window.cpp b/src/duckdb/src/planner/sql_export/sql_export_window.cpp new file mode 100644 index 000000000..9f329bb12 --- /dev/null +++ b/src/duckdb/src/planner/sql_export/sql_export_window.cpp @@ -0,0 +1,225 @@ +#include "duckdb/planner/sql_export/bound_expression_sql_exporter_internal.hpp" +#include "duckdb/planner/expression/bound_aggregate_expression.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" +#include "duckdb/parser/expression/cast_expression.hpp" +#include "duckdb/parser/expression/constant_expression.hpp" +#include "duckdb/planner/expression/bound_cast_expression.hpp" +#include "duckdb/planner/expression/bound_constant_expression.hpp" +#include "duckdb/planner/expression/bound_function_expression.hpp" +#include "duckdb/function/window_function.hpp" +#include "duckdb/planner/expression/bound_window_expression.hpp" + +namespace duckdb { + +struct WindowSQLFrame { + optional_ptr order; + optional_ptr start; + optional_ptr end; + unique_ptr start_literal; + unique_ptr end_literal; +}; + +static optional NumericRangeLiteral(const ParsedExpression &literal) { + if (literal.GetExpressionClass() == ExpressionClass::CONSTANT) { + auto value = literal.Cast().GetLiteral().ToValue(); + return value.type().IsNumeric() && !value.IsNull() ? optional(value) : optional(); + } + if (literal.GetExpressionClass() == ExpressionClass::CAST) { + auto &cast = literal.Cast(); + auto type = UnboundType::TryDefaultBind(cast.TargetType()); + auto child = NumericRangeLiteral(cast.Child()); + if (child && type.IsNumeric() && !cast.IsTryCast()) { + return child->DefaultTryCastAs(type); + } + } + return {}; +} + +//! The SQL offset of a RANGE frame endpoint: the retained literal when it still describes the endpoint, +//! otherwise the live input the binder's endpoint arithmetic was derived from +static optional_ptr RangeOffset(const BoundWindowExpression &expression, const WindowSQLFrame &frame, + const unique_ptr &endpoint, WindowBoundary boundary, + const unique_ptr &literal, + const unique_ptr &origin) { + const bool is_range_offset = + boundary == WindowBoundary::EXPR_PRECEDING_RANGE || boundary == WindowBoundary::EXPR_FOLLOWING_RANGE; + if (!is_range_offset) { + return endpoint.get(); + } + if (!origin || !endpoint || !frame.order) { + return nullptr; + } + auto &order = expression.OrderBy()[0]; + const bool origin_matches_endpoint = origin->boundary == boundary && origin->direction == order.type; + if (!origin_matches_endpoint) { + return nullptr; + } + const bool has_literal_offset = literal && literal->GetExpressionClass() == ExpressionClass::BOUND_CONSTANT; + const bool endpoint_is_order = has_literal_offset && Expression::Equals(*endpoint, *order.expression); + if (endpoint_is_order) { + auto &value = literal->Cast().GetValue(); + if (!value.IsNull() && value.type().IsNumeric() && value == Value::Numeric(value.type(), 0)) { + return literal.get(); + } + } + auto input = origin->Match(*endpoint, *frame.order); + if (!input) { + return nullptr; + } + auto &offset = *input; + const bool offset_is_constant = offset.GetExpressionClass() == ExpressionClass::BOUND_CONSTANT; + if (offset_is_constant && offset.Cast().GetValue().IsNull()) { + return nullptr; + } + if (has_literal_offset && offset_is_constant) { + auto &original = literal->Cast().GetValue(); + auto ¤t = offset.Cast().GetValue(); + const bool has_numeric_offsets = original.type().IsNumeric() && current.type().IsNumeric(); + if (!original.IsNull() && has_numeric_offsets && original == current) { + return literal.get(); + } + } + return &offset; +} + +static WindowSQLFrame ReconstructWindowFrame(const BoundWindowExpression &expression) { + WindowSQLFrame frame; + if (expression.OrderBy().size() == 1) { + frame.order = WindowRangeCast::Match(*expression.OrderBy()[0].expression, expression.SQLRangeOrderCasts()); + if (frame.order && expression.SQLRangeOrderType().IsComplete() && + !frame.order->GetReturnType().EqualsIncludingCollation(expression.SQLRangeOrderType())) { + frame.order = nullptr; + } + } + auto retained_offset = [&](const unique_ptr &literal) -> unique_ptr { + auto value = literal ? NumericRangeLiteral(*literal) : optional(); + return value ? make_uniq(*value) : nullptr; + }; + frame.start_literal = retained_offset(expression.SQLRangeStart()); + frame.end_literal = retained_offset(expression.SQLRangeEnd()); + frame.start = RangeOffset(expression, frame, expression.StartExpr(), expression.WindowStart(), frame.start_literal, + expression.SQLRangeStartBoundary()); + frame.end = RangeOffset(expression, frame, expression.EndExpr(), expression.WindowEnd(), frame.end_literal, + expression.SQLRangeEndBoundary()); + return frame; +} + +BoundExpressionSQLExportResult BoundExpressionSQLExportState::ExportWindow(const BoundWindowExpression &expression, + const LogicalPlanVerificationPath &path) { + if (expression.AggregateFunction()) { + return BoundExpressionSQLExportState::PreserveCollation( + expression.GetReturnType(), ExportWindowFunction(expression, *expression.AggregateFunction(), path), path); + } + D_ASSERT(expression.WindowFunction()); + return BoundExpressionSQLExportState::PreserveCollation( + expression.GetReturnType(), ExportWindowFunction(expression, *expression.WindowFunction(), path), path); +} + +template +BoundExpressionSQLExportResult +BoundExpressionSQLExportState::ExportWindowFunction(const BoundWindowExpression &expression, const FUNCTION &function, + const LogicalPlanVerificationPath &path) { + auto &definition = function.GetDefinition(); + D_ASSERT(definition); + auto identity = BoundExpressionSQLExportState::DefinitionFunctionIdentity( + *definition, function.GetLogicalArguments(), function.GetLogicalReturnType()); + if (!identity.IsValid()) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::InternalExpressionInvariant( + path, expression, "Bound window function identity is incomplete")}); + } + if (expression.GetChildren().size() != function.GetLogicalArguments().size()) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The window function does not retain every SQL argument")}); + } + auto name = BoundExpressionSQLExportState::RebindableFunctionName(*definition); + if (!name || !SQLExportHelpers::IsSQLValueType(expression.GetReturnType()) || + definition->GetProperties().GetCaptureArgumentAliases()) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The window function no longer represents its logical SQL signature")}); + } + auto frame = ReconstructWindowFrame(expression); + if ((expression.StartExpr() && !frame.start) || (expression.EndExpr() && !frame.end)) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFeature( + path, "window_range_offset", "The RANGE endpoint does not retain its SQL offset and ordering operand")}); + } + + vector source_children; + for (auto &partition : expression.Partitions()) { + source_children.emplace_back(partition.get()); + } + for (auto &order : expression.OrderBy()) { + source_children.emplace_back(frame.order ? frame.order : order.expression.get()); + } + for (idx_t i = 0; i < expression.GetChildren().size(); i++) { + source_children.emplace_back(expression.GetChildren()[i].get()); + } + if (expression.Filter()) { + source_children.emplace_back(expression.Filter().get(), LogicalType::BOOLEAN); + } + if (expression.StartExpr()) { + source_children.emplace_back(frame.start); + } + if (expression.EndExpr()) { + source_children.emplace_back(frame.end); + } + for (auto &order : expression.ArgOrders()) { + source_children.emplace_back(order.expression.get()); + } + vector> children; + vector issues; + ExportChildren(source_children, path, children, issues); + if (!issues.empty()) { + return BoundExpressionSQLExportResult::Failure(std::move(issues)); + } + auto window = make_uniq(name->Catalog().GetIdentifierName(), name->Schema().GetIdentifierName(), + name->Name().GetIdentifierName()); + window->SetQualifiedName(*name); + idx_t ordinal = 0; + for (idx_t i = 0; i < expression.Partitions().size(); i++) { + window->PartitionsMutable().push_back(std::move(children[ordinal++])); + } + for (auto &order : expression.OrderBy()) { + window->OrderByMutable().emplace_back( + order.type, order.null_order, + SQLExportHelpers::OrderExpression(order.expression->GetReturnType(), std::move(children[ordinal++]))); + } + auto &named_arguments = function.GetNamedArguments(); + auto positional_count = function.GetPositionalArgumentCount(); + if (!named_arguments.empty() && positional_count + named_arguments.size() != expression.GetChildren().size()) { + return BoundExpressionSQLExportResult::Failure({BoundExpressionSQLExportState::UnsupportedFunction( + path, std::move(identity), "The named SQL arguments are incomplete")}); + } + for (idx_t i = 0; i < expression.GetChildren().size(); i++) { + auto argument_name = + !named_arguments.empty() && i >= positional_count ? named_arguments[i - positional_count] : Identifier(); + window->GetArgumentsMutable().emplace_back(std::move(argument_name), std::move(children[ordinal++])); + } + if (expression.Filter()) { + window->FilterMutable() = std::move(children[ordinal++]); + } + if (expression.StartExpr()) { + window->StartExprMutable() = std::move(children[ordinal++]); + } + if (expression.EndExpr()) { + window->EndExprMutable() = std::move(children[ordinal++]); + } + for (auto &order : expression.ArgOrders()) { + window->ArgOrdersMutable().emplace_back( + order.type, order.null_order, + SQLExportHelpers::OrderExpression(order.expression->GetReturnType(), std::move(children[ordinal++]))); + } + window->IgnoreNullsMutable() = expression.IgnoreNulls(); + window->HasIgnoreNullsMutable() = expression.IgnoreNulls(); + window->DistinctMutable() = expression.Distinct(); + window->WindowStartMutable() = expression.WindowStart(); + window->WindowEndMutable() = expression.WindowEnd(); + window->WindowExcludeMutable() = expression.WindowExclude(); + unique_ptr result = std::move(window); + if (SQLExportHelpers::IsSQLRepresentableType(expression.GetReturnType()) && definition->HasBindCallback() && + definition->GetReturnType() != expression.GetReturnType()) { + return RestoreResultType(expression.GetReturnType(), std::move(result), path); + } + return BoundExpressionSQLExportResult::Success(std::move(result)); +} + +} // namespace duckdb diff --git a/src/duckdb/src/planner/sql_export/table_function_sql_export.cpp b/src/duckdb/src/planner/sql_export/table_function_sql_export.cpp new file mode 100644 index 000000000..e032ebfac --- /dev/null +++ b/src/duckdb/src/planner/sql_export/table_function_sql_export.cpp @@ -0,0 +1,341 @@ +#include "duckdb/planner/sql_export/logical_plan_sql_exporter_internal.hpp" +#include "duckdb/catalog/catalog_entry/table_catalog_entry.hpp" +#include "duckdb/catalog/catalog_entry/schema_catalog_entry.hpp" +#include "duckdb/parser/tableref/basetableref.hpp" +#include "duckdb/parser/tableref/at_clause.hpp" +#include "duckdb/parser/expression/operator_expression.hpp" +#include "duckdb/parser/expression/window_expression.hpp" +#include "duckdb/function/table_function.hpp" +#include "duckdb/parser/expression/cast_expression.hpp" +#include "duckdb/parser/expression/columnref_expression.hpp" +#include "duckdb/parser/expression/comparison_expression.hpp" +#include "duckdb/parser/expression/constant_expression.hpp" +#include "duckdb/parser/expression/function_expression.hpp" +#include "duckdb/parser/expression/subquery_expression.hpp" +#include "duckdb/parser/query_node/select_node.hpp" +#include "duckdb/parser/tableref/joinref.hpp" +#include "duckdb/parser/tableref/table_function_ref.hpp" +#include "duckdb/planner/operator/logical_get.hpp" +#include "duckdb/planner/bound_expression_sql_exporter.hpp" +#include "duckdb/planner/expression/bound_constant_expression.hpp" +#include "duckdb/planner/sql_export_helpers.hpp" + +namespace duckdb { + +static string SQLFunctionCallGuard(const LogicalGet &get, bool has_input) { + const auto argument_count = get.function.GetArguments().size(); + const bool has_expected_parameters = + get.function.HasVarArgs() ? get.parameters.size() >= argument_count : get.parameters.size() == argument_count; + if (!has_input && !has_expected_parameters) { + return "positional_parameters"; + } + if (has_input && !get.function.in_out_function) { + return "input_child"; + } + if (has_input && (get.input_table_types.empty() || get.input_table_types.size() > get.children[0]->types.size())) { + return "input_columns"; + } + if (get.ordinality_idx.IsValid() && get.source_ordinality != OrdinalityType::WITH_ORDINALITY) { + return "ordinality"; + } + if (!get.scan_partition_indices.empty()) { + return "scan_partitions"; + } + { + auto name = get.function.GetQualifiedName(); + if (name.Catalog().empty()) { + return "catalog_identifier"; + } + if (name.Schema().empty()) { + return "schema_identifier"; + } + for (auto &component : name.Path()) { + if (component.empty()) { + return "function_identifier"; + } + } + } + return string(); +} + +static bool AppendTableFunctionColumnPath(vector &path, const LogicalType &parent_type, + const ColumnIndex &index, LogicalType &source_type) { + Identifier name; + LogicalType child_type = LogicalType::INVALID; + if (parent_type.id() == LogicalTypeId::STRUCT) { + auto &children = StructType::GetChildTypes(parent_type); + if (index.HasPrimaryIndex()) { + if (index.GetPrimaryIndex() >= children.size()) { + return false; + } + name = children[index.GetPrimaryIndex()].first; + child_type = children[index.GetPrimaryIndex()].second; + } else { + name = Identifier(index.GetFieldName()); + for (auto &child : children) { + if (child.first == name) { + child_type = child.second; + break; + } + } + } + } else if (parent_type.id() == LogicalTypeId::VARIANT && !index.HasPrimaryIndex()) { + name = Identifier(index.GetFieldName()); + child_type = parent_type; + } else { + return false; + } + path.push_back(std::move(name)); + if (index.GetChildIndexes().empty()) { + source_type = std::move(child_type); + return true; + } + if (index.GetChildIndexes().size() != 1 || child_type.id() == LogicalTypeId::INVALID) { + return false; + } + return AppendTableFunctionColumnPath(path, child_type, index.GetChildIndexes()[0], source_type); +} + +static unique_ptr TableScanColumnToSQL(const ColumnIndex &index, const LogicalType &type, + unique_ptr expression) { + if (!index.IsPushdownExtract()) { + if (index.HasType() && index.GetType() != type) { + return make_uniq(index.GetType(), std::move(expression)); + } + return expression; + } + D_ASSERT(index.ChildIndexCount() == 1); + auto &child = index.GetChildIndex(0); + Value key; + LogicalType child_type; + if (child.HasPrimaryIndex()) { + D_ASSERT(type.id() == LogicalTypeId::STRUCT); + auto field = child.GetPrimaryIndex(); + key = StructType::IsUnnamed(type) ? Value::BIGINT(NumericCast(field + 1)) + : Value(StructType::GetChildName(type, field)); + child_type = StructType::GetChildType(type, field); + } else { + D_ASSERT(type.id() == LogicalTypeId::VARIANT); + key = Value(child.GetFieldName()); + child_type = type; + } + auto extract = make_uniq(ExpressionType::ARRAY_EXTRACT); + extract->GetChildrenMutable().push_back(std::move(expression)); + extract->GetChildrenMutable().push_back(ConstantExpression::FromValue(key)); + return TableScanColumnToSQL(child, child_type, std::move(extract)); +} + +static unique_ptr TableFunctionColumn(const LogicalGet &get, const TableRef &function_ref, + const ColumnIndex &index) { + if (get.GetTable()) { + if (index.IsRowNumberColumn()) { + auto row_number = make_uniq("system", "main", "row_number"); + row_number->WindowStartMutable() = WindowBoundary::UNBOUNDED_PRECEDING; + row_number->WindowEndMutable() = WindowBoundary::CURRENT_ROW_RANGE; + return std::move(row_number); + } + if (!index.IsVirtualColumn()) { + auto &column = get.GetTable()->GetColumn(index.ToLogical()); + auto expression = make_uniq(function_ref.column_name_alias[index.GetPrimaryIndex()], + function_ref.alias); + return TableScanColumnToSQL(index, column.Type(), std::move(expression)); + } + } + if (index.IsVirtualColumn()) { + if (index.IsEmptyColumn()) { + return ConstantExpression::FromValue(Value::INTEGER(1)); + } + auto entry = get.virtual_columns.find(index.GetPrimaryIndex()); + if (entry == get.virtual_columns.end()) { + return nullptr; + } + return make_uniq(entry->second.name, function_ref.alias); + } + auto primary_index = index.GetPrimaryIndex(); + if (primary_index >= function_ref.column_name_alias.size() || primary_index >= get.returned_types.size()) { + return nullptr; + } + vector path {function_ref.alias, function_ref.column_name_alias[primary_index]}; + if (!index.IsPushdownExtract()) { + auto column = make_uniq(std::move(path)); + if (index.HasType() && !get.returned_types[primary_index].EqualsIncludingCollation(index.GetScanType())) { + return make_uniq(index.GetScanType(), std::move(column)); + } + return std::move(column); + } + LogicalType source_type; + if (index.GetChildIndexes().size() != 1 || + !AppendTableFunctionColumnPath(path, get.returned_types[primary_index], index.GetChildIndexes()[0], + source_type)) { + return nullptr; + } + auto column = make_uniq(std::move(path)); + if (!source_type.EqualsIncludingCollation(index.GetScanType())) { + return make_uniq(index.GetScanType(), std::move(column)); + } + return std::move(column); +} + +SQLSourceQueryResult LogicalPlanSQLExportHelpers::ReconstructSQLSource(ClientContext &context, const LogicalGet &get, + unique_ptr input, + const Identifier &relation_alias, + bool source_ordinality) { + if (!get.scan_partition_indices.empty()) { + return {nullptr, "scan_partitions"}; + } + unique_ptr source; + if (get.function.to_sql) { + auto result = get.function.to_sql(context, get); + if (!result.source) { + return {nullptr, + result.unsupported_reason.empty() ? "source_callback_declined" : result.unsupported_reason}; + } + if (input) { + return {nullptr, "custom_source_input"}; + } + source = std::move(result.source); + } else if (auto table = get.GetTable()) { + auto reference = make_uniq(); + reference->SetQualifiedName(table->schema.GetQualifiedName(table->name)); + source = std::move(reference); + } else { + if (get.bind_info) { + return {nullptr, "process_local_input"}; + } + auto guard = SQLFunctionCallGuard(get, input != nullptr); + if (!guard.empty()) { + return {nullptr, std::move(guard)}; + } + auto function = make_uniq(); + vector> parameters; + const auto &signature = get.function.GetArguments(); + const bool table_parameter = + std::find(signature.begin(), signature.end(), LogicalType::TABLE) != signature.end(); + if (input && !table_parameter) { + if (input->column_name_alias.size() < get.input_table_types.size()) { + return {nullptr, "input_columns"}; + } + for (idx_t i = 0; i < get.input_table_types.size(); i++) { + parameters.push_back(make_uniq(input->column_name_alias[i], input->alias)); + } + } else { + bool consumed_input = false; + for (idx_t i = 0; i < get.parameters.size(); i++) { + if (i < signature.size() && signature[i] == LogicalType::TABLE) { + const bool has_available_input = input && !consumed_input && get.projected_input.empty(); + const bool has_input_columns = has_available_input && + input->column_name_alias.size() >= get.input_table_types.size() && + get.input_table_names.size() == get.input_table_types.size(); + if (!has_input_columns) { + return {nullptr, "table_parameter"}; + } + auto query = make_uniq(); + for (idx_t column_idx = 0; column_idx < get.input_table_types.size(); column_idx++) { + auto column = + make_uniq(input->column_name_alias[column_idx], input->alias); + column->SetAlias(get.input_table_names[column_idx]); + query->select_list.push_back(std::move(column)); + } + query->from_table = std::move(input); + auto subquery = make_uniq(); + subquery->GetSubqueryTypeMutable() = SubqueryType::SCALAR; + subquery->SubqueryMutable() = make_uniq(); + subquery->SubqueryMutable()->node = std::move(query); + parameters.push_back(std::move(subquery)); + consumed_input = true; + continue; + } + auto exported = BoundExpressionSQLExporter::Export(BoundConstantExpression(get.parameters[i]), {}); + if (exported.HasError()) { + return {nullptr, "positional_parameter_expression"}; + } + parameters.push_back(std::move(exported.GetValue())); + } + if (input) { + return {nullptr, "table_parameter"}; + } + } + for (auto ¶meter : get.named_parameters) { + auto exported = BoundExpressionSQLExporter::Export(BoundConstantExpression(parameter.second), {}); + if (exported.HasError()) { + return {nullptr, "named_parameter_expression"}; + } + parameters.push_back(make_uniq(ExpressionType::COMPARE_EQUAL, + make_uniq(parameter.first), + std::move(exported.GetValue()))); + } + function->function = make_uniq(get.function.GetQualifiedName(), std::move(parameters)); + source = std::move(function); + } + if (get.at_clause) { + if (source->type != TableReferenceType::BASE_TABLE) { + return {nullptr, "at_clause"}; + } + source->Cast().at_clause = + make_uniq(get.at_clause->Unit(), ConstantExpression::FromValue(get.at_clause->GetValue())); + } + auto &function_ref = *source; + if (source->type == TableReferenceType::TABLE_FUNCTION) { + source->Cast().with_ordinality = get.ordinality_idx.IsValid() || source_ordinality + ? OrdinalityType::WITH_ORDINALITY + : OrdinalityType::WITHOUT_ORDINALITY; + } else if (get.ordinality_idx.IsValid() || source_ordinality) { + return {nullptr, "source_ordinality"}; + } + function_ref.alias = relation_alias; + function_ref.column_name_alias.clear(); + for (idx_t i = 0; i < get.returned_types.size(); i++) { + function_ref.column_name_alias.emplace_back("column" + std::to_string(i)); + } + if (source_ordinality) { + function_ref.column_name_alias.emplace_back("column" + std::to_string(get.returned_types.size())); + } + auto select = make_uniq(); + for (auto &index : get.GetColumnIds()) { + auto column = TableFunctionColumn(get, function_ref, index); + if (!column) { + return {nullptr, "to_sql_callback_declined_without_guard"}; + } + select->select_list.push_back(std::move(column)); + } + for (auto index : get.projected_input) { + select->select_list.push_back(make_uniq(input->column_name_alias[index], input->alias)); + } + if (source_ordinality) { + select->select_list.push_back( + make_uniq(function_ref.column_name_alias.back(), function_ref.alias)); + } + if (get.extra_info.file_filter_expressions) { + BoundExpressionSQLExportContext export_context; + export_context.client_context = &context; + export_context.resolve_binding = [&](const ColumnBinding &binding) -> optional { + auto index = binding.column_index.GetIndex(); + if (binding.table_index != TableIndex(ExtraOperatorInfo::FILE_FILTER_TABLE_INDEX) || + index >= get.returned_types.size()) { + return {}; + } + return ResolvedSQLColumnReference {{function_ref.alias, function_ref.column_name_alias[index]}, + get.returned_types[index]}; + }; + for (auto &predicate : *get.extra_info.file_filter_expressions) { + auto exported = BoundExpressionSQLExporter::Export(*predicate, export_context); + if (exported.HasError()) { + return {nullptr, "to_sql_callback_declined_without_guard"}; + } + select->where_clause = + SQLExportHelpers::Conjoin(std::move(select->where_clause), std::move(exported.GetValue())); + } + } + + if (!input) { + select->from_table = std::move(source); + return {std::move(select), {}}; + } + auto join = make_uniq(JoinRefType::CROSS); + join->left = std::move(input); + join->right = std::move(source); + select->from_table = std::move(join); + return {std::move(select), {}}; +} + +} // namespace duckdb diff --git a/src/duckdb/src/storage/serialization/serialize_logical_operator.cpp b/src/duckdb/src/storage/serialization/serialize_logical_operator.cpp index f7d44ca1d..440ba5ba4 100644 --- a/src/duckdb/src/storage/serialization/serialize_logical_operator.cpp +++ b/src/duckdb/src/storage/serialization/serialize_logical_operator.cpp @@ -555,6 +555,7 @@ void LogicalExplain::Serialize(Serializer &serializer) const { serializer.WritePropertyWithDefault(201, "physical_plan", physical_plan); serializer.WritePropertyWithDefault(202, "logical_plan_unopt", logical_plan_unopt); serializer.WritePropertyWithDefault(203, "logical_plan_opt", logical_plan_opt); + serializer.WritePropertyWithDefault>(204, "sql_output_names", sql_output_names, vector()); } unique_ptr LogicalExplain::Deserialize(Deserializer &deserializer) { @@ -563,6 +564,7 @@ unique_ptr LogicalExplain::Deserialize(Deserializer &deserializ deserializer.ReadPropertyWithDefault(201, "physical_plan", result->physical_plan); deserializer.ReadPropertyWithDefault(202, "logical_plan_unopt", result->logical_plan_unopt); deserializer.ReadPropertyWithDefault(203, "logical_plan_opt", result->logical_plan_opt); + deserializer.ReadPropertyWithExplicitDefault>(204, "sql_output_names", result->sql_output_names, vector()); return std::move(result); } diff --git a/src/duckdb/src/storage/serialization/serialize_nodes.cpp b/src/duckdb/src/storage/serialization/serialize_nodes.cpp index e0321920d..8be9790d2 100644 --- a/src/duckdb/src/storage/serialization/serialize_nodes.cpp +++ b/src/duckdb/src/storage/serialization/serialize_nodes.cpp @@ -9,6 +9,7 @@ #include "duckdb/parser/statement/select_statement.hpp" #include "duckdb/parser/query_node.hpp" #include "duckdb/parser/result_modifier.hpp" +#include "duckdb/planner/tableref/bound_at_clause.hpp" #include "duckdb/planner/bound_result_modifier.hpp" #include "duckdb/planner/operator/logical_external_resource.hpp" #include "duckdb/parser/expression/case_expression.hpp" @@ -90,6 +91,18 @@ unique_ptr BaseReservoirSampling::Deserialize(Deserialize return result; } +void BoundAtClause::Serialize(Serializer &serializer) const { + serializer.WritePropertyWithDefault(100, "unit", unit); + serializer.WriteProperty(101, "value", val); +} + +unique_ptr BoundAtClause::Deserialize(Deserializer &deserializer) { + auto unit = deserializer.ReadPropertyWithDefault(100, "unit"); + auto val = deserializer.ReadProperty(101, "value"); + auto result = duckdb::unique_ptr(new BoundAtClause(std::move(unit), val)); + return result; +} + void BoundCaseCheck::Serialize(Serializer &serializer) const { serializer.WritePropertyWithDefault>(100, "when_expr", when_expr); serializer.WritePropertyWithDefault>(101, "then_expr", then_expr); diff --git a/src/duckdb/ub_src_planner.cpp b/src/duckdb/ub_src_planner.cpp index 6676dae76..caf2a8218 100644 --- a/src/duckdb/ub_src_planner.cpp +++ b/src/duckdb/ub_src_planner.cpp @@ -4,8 +4,6 @@ #include "src/planner/binding_alias.cpp" -#include "src/planner/bound_expression_sql_exporter.cpp" - #include "src/planner/bound_parameter_map.cpp" #include "src/planner/bound_result_modifier.cpp" diff --git a/src/duckdb/ub_src_planner_expression.cpp b/src/duckdb/ub_src_planner_expression.cpp index 19a61ea89..a66a4ed61 100644 --- a/src/duckdb/ub_src_planner_expression.cpp +++ b/src/duckdb/ub_src_planner_expression.cpp @@ -36,3 +36,5 @@ #include "src/planner/expression/legacy_bound_comparison_expression.cpp" +#include "src/planner/expression/window_range_info.cpp" + diff --git a/src/duckdb/ub_src_planner_sql_export.cpp b/src/duckdb/ub_src_planner_sql_export.cpp new file mode 100644 index 000000000..3c2c1cbb2 --- /dev/null +++ b/src/duckdb/ub_src_planner_sql_export.cpp @@ -0,0 +1,28 @@ +#include "src/planner/sql_export/bound_expression_sql_exporter.cpp" + +#include "src/planner/sql_export/logical_plan_sql_exporter.cpp" + +#include "src/planner/sql_export/sql_export_constants.cpp" + +#include "src/planner/sql_export/sql_export_cte.cpp" + +#include "src/planner/sql_export/sql_export_functions.cpp" + +#include "src/planner/sql_export/sql_export_joins.cpp" + +#include "src/planner/sql_export/sql_export_limit.cpp" + +#include "src/planner/sql_export/sql_export_pivot.cpp" + +#include "src/planner/sql_export/sql_export_relations.cpp" + +#include "src/planner/sql_export/sql_export_scope.cpp" + +#include "src/planner/sql_export/sql_export_sources.cpp" + +#include "src/planner/sql_export/sql_export_values.cpp" + +#include "src/planner/sql_export/sql_export_window.cpp" + +#include "src/planner/sql_export/table_function_sql_export.cpp" +