diff --git a/src/include/dbconnector/query/query_writer.hpp b/src/include/dbconnector/query/query_writer.hpp index f1a3f2e..42ccb66 100644 --- a/src/include/dbconnector/query/query_writer.hpp +++ b/src/include/dbconnector/query/query_writer.hpp @@ -16,14 +16,18 @@ class QueryWriter { struct Config { char quote = '"'; QuoteEscapeStyle escape_style = QuoteEscapeStyle::DOUBLE_QUOTE; + std::string blob_literal_prefix; + std::string blob_literal_suffix; }; - static Config CreateConfig(char quote, QuoteEscapeStyle escape_style); + static Config CreateConfig(char quote, QuoteEscapeStyle escape_style, + const std::string &blob_literal_prefix = std::string(), + const std::string &blob_literal_suffix = std::string()); static std::string WriteQuotedAndEscaped(const QueryWriter::Config &config, const std::string &text); - static std::string WriteConstant(const duckdb::Value &val); + static std::string WriteConstant(const QueryWriter::Config &config, const duckdb::Value &val); private: - static std::string EncodeBlob(const std::string &val); + static std::string EncodeBlob(const QueryWriter::Config &config, const std::string &val); }; } // namespace query diff --git a/src/include/dbconnector/table_scan/filter_pushdown.hpp b/src/include/dbconnector/table_scan/filter_pushdown.hpp index 1ee8f80..5e7efaf 100644 --- a/src/include/dbconnector/table_scan/filter_pushdown.hpp +++ b/src/include/dbconnector/table_scan/filter_pushdown.hpp @@ -13,21 +13,36 @@ namespace table_scan { class FilterPushdown { struct Config { char identifier_quote = '"'; + char constant_quote = '\''; query::QuoteEscapeStyle escape_style = query::QuoteEscapeStyle::DOUBLE_QUOTE; + std::string blob_literal_prefix; + std::string blob_literal_suffix; }; public: - static Config CreateConfig(char identifier_quote, query::QuoteEscapeStyle escape_style); + static Config CreateConfig(char identifier_quote, char constant_quote, query::QuoteEscapeStyle escape_style, + const std::string &blob_literal_prefix = std::string(), + const std::string &blob_literal_suffix = std::string()); static std::string TransformFilter(const Config &config, const std::string &column_name, - const duckdb::TableFilter &filter); + const duckdb::TableFilter &filter, duckdb::column_t column_id); private: - static std::string TransformExpression(const std::string &column_name, const duckdb::Expression &expr); + static std::string TransformExpression(const query::QueryWriter::Config &identifier_config, + const query::QueryWriter::Config &constant_config, + const std::string &column_name, const duckdb::Expression &expr, + duckdb::column_t column_id); + static std::string TransformExpressionSubject(const query::QueryWriter::Config &identifier_config, + const std::string &column_name, const duckdb::Expression &expr); + static std::string TransformConstantFilter(const query::QueryWriter::Config &constant_config, + const std::string &column_name, duckdb::ExpressionType comparison_type, + const duckdb::Value &constant, duckdb::column_t column_id); static std::string TransformComparison(duckdb::ExpressionType type); - static std::string CreateExpression(const std::string &column_name, + static std::string CreateExpression(const query::QueryWriter::Config &identifier_config, + const query::QueryWriter::Config &constant_config, + const std::string &column_name, const duckdb::vector> &filters, - const std::string &op); + const std::string &op, duckdb::column_t column_id); }; } // namespace table_scan diff --git a/src/optimizer/aggregate_optimizer.cpp b/src/optimizer/aggregate_optimizer.cpp index 057f8c8..1967420 100644 --- a/src/optimizer/aggregate_optimizer.cpp +++ b/src/optimizer/aggregate_optimizer.cpp @@ -191,8 +191,9 @@ static PushedAggregate TryPushAggregateToMySQL(const AggregateOptimizer::Config return res; } auto column_name = get.names[table_col_idx]; - auto scan_config = table_scan::FilterPushdown::CreateConfig('`', config.escape_style); - auto new_filter = table_scan::FilterPushdown::TransformFilter(scan_config, column_name, entry.Filter()); + auto scan_config = table_scan::FilterPushdown::CreateConfig('`', '\'', config.escape_style); + auto new_filter = + table_scan::FilterPushdown::TransformFilter(scan_config, column_name, entry.Filter(), table_col_idx); if (new_filter.empty()) { return res; } diff --git a/src/query/query_writer.cpp b/src/query/query_writer.cpp index 6bbe9f2..b2d9d1d 100644 --- a/src/query/query_writer.cpp +++ b/src/query/query_writer.cpp @@ -5,10 +5,14 @@ namespace dbconnector { namespace query { -QueryWriter::Config QueryWriter::CreateConfig(char quote, QuoteEscapeStyle escape_style) { +QueryWriter::Config QueryWriter::CreateConfig(char quote, QuoteEscapeStyle escape_style, + const std::string &blob_literal_prefix, + const std::string &blob_literal_suffix) { Config res; res.quote = quote; res.escape_style = escape_style; + res.blob_literal_prefix = blob_literal_prefix; + res.blob_literal_suffix = blob_literal_suffix; return res; } @@ -41,25 +45,26 @@ std::string QueryWriter::WriteQuotedAndEscaped(const QueryWriter::Config &config return result; } -std::string QueryWriter::EncodeBlob(const std::string &val) { +std::string QueryWriter::EncodeBlob(const QueryWriter::Config &config, const std::string &val) { char const HEX_DIGITS[] = "0123456789ABCDEF"; - std::string result = "x'"; + std::string result = config.blob_literal_prefix; for (size_t i = 0; i < val.size(); i++) { uint8_t byte_val = static_cast(val[i]); result += HEX_DIGITS[(byte_val >> 4) & 0xf]; result += HEX_DIGITS[byte_val & 0xf]; } result += "'"; + result += config.blob_literal_suffix; return result; } -std::string QueryWriter::WriteConstant(const duckdb::Value &val) { +std::string QueryWriter::WriteConstant(const QueryWriter::Config &config, const duckdb::Value &val) { if (val.type().IsNumeric() || val.type().id() == duckdb::LogicalTypeId::BOOLEAN) { return val.ToSQLString(); } if (val.type().id() == duckdb::LogicalTypeId::BLOB) { - return EncodeBlob(duckdb::StringValue::Get(val)); + return EncodeBlob(config, duckdb::StringValue::Get(val)); } if (val.type().id() == duckdb::LogicalTypeId::TIMESTAMP_TZ) { return val.DefaultCastAs(duckdb::LogicalType::TIMESTAMP) diff --git a/src/table_scan/filter_pushdown.cpp b/src/table_scan/filter_pushdown.cpp index fd40ed7..71aac1a 100644 --- a/src/table_scan/filter_pushdown.cpp +++ b/src/table_scan/filter_pushdown.cpp @@ -1,5 +1,6 @@ #include "dbconnector/table_scan/filter_pushdown.hpp" +#include "duckdb/function/scalar/struct_utils.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" @@ -19,18 +20,27 @@ namespace table_scan { using namespace duckdb; -FilterPushdown::Config FilterPushdown::CreateConfig(char identifier_quote, query::QuoteEscapeStyle escape_style) { +FilterPushdown::Config FilterPushdown::CreateConfig(char identifier_quote, char constant_quote, + query::QuoteEscapeStyle escape_style, + const std::string &blob_literal_prefix, + const std::string &blob_literal_suffix) { Config res; res.identifier_quote = identifier_quote; + res.constant_quote = constant_quote; res.escape_style = escape_style; + res.blob_literal_prefix = blob_literal_prefix; + res.blob_literal_suffix = blob_literal_suffix; return res; } -std::string FilterPushdown::CreateExpression(const std::string &column_name, - const vector> &filters, const std::string &op) { +std::string FilterPushdown::CreateExpression(const query::QueryWriter::Config &identifier_config, + const query::QueryWriter::Config &constant_config, + const std::string &column_name, + const vector> &filters, const std::string &op, + column_t column_id) { vector filter_entries; for (auto &filter : filters) { - auto new_filter = TransformExpression(column_name, *filter); + auto new_filter = TransformExpression(identifier_config, constant_config, column_name, *filter, column_id); if (new_filter.empty()) { continue; } @@ -71,24 +81,76 @@ static bool IsDirectReference(const Expression &expr) { } } -std::string FilterPushdown::TransformExpression(const std::string &column_name, const Expression &expr) { +string FilterPushdown::TransformConstantFilter(const query::QueryWriter::Config &constant_config, + const string &column_name, ExpressionType comparison_type, + const Value &constant, column_t column_id) { + string constant_string; + if (IsVirtualColumn(column_id)) { + return "FALSE"; + } else { + constant_string = query::QueryWriter::WriteConstant(constant_config, constant); + } + auto operator_string = TransformComparison(comparison_type); + string comparison = StringUtil::Format("%s %s %s", column_name, operator_string, constant_string); + if (constant.type().id() == LogicalTypeId::VARCHAR) { + comparison += " COLLATE \"C\""; + } + return comparison; +} + +string FilterPushdown::TransformExpressionSubject(const query::QueryWriter::Config &identifier_config, + const string &column_name, const Expression &expr) { + switch (expr.GetExpressionClass()) { + case ExpressionClass::BOUND_REF: + case ExpressionClass::BOUND_COLUMN_REF: + return column_name; + case ExpressionClass::BOUND_FUNCTION: { + auto &func = expr.Cast(); + idx_t child_idx; + if (!TryGetStructExtractChildIndex(func, child_idx) || func.GetChildren().empty()) { + return string(); + } + auto parent_name = TransformExpressionSubject(identifier_config, column_name, *func.GetChildren()[0]); + if (parent_name.empty()) { + return string(); + } + auto &struct_type = func.GetChildren()[0]->GetReturnType(); + if (struct_type.id() != LogicalTypeId::STRUCT || StructType::IsUnnamed(struct_type)) { + return string(); + } + auto child_name = query::QueryWriter::WriteQuotedAndEscaped(identifier_config, + StructType::GetChildName(struct_type, child_idx)); + return "(" + parent_name + ")." + child_name; + } + default: + return string(); + } +} + +std::string FilterPushdown::TransformExpression(const query::QueryWriter::Config &identifier_config, + const query::QueryWriter::Config &constant_config, + const std::string &column_name, const Expression &expr, + column_t column_id) { if (BoundComparisonExpression::IsComparison(expr)) { auto &comparison = expr.Cast(); auto comparison_type = comparison.GetExpressionType(); auto &left = BoundComparisonExpression::Left(comparison); auto &right = BoundComparisonExpression::Right(comparison); + auto subject = TransformExpressionSubject(identifier_config, column_name, left); const Value *constant = nullptr; - if (IsDirectReference(left) && right.GetExpressionClass() == ExpressionClass::BOUND_CONSTANT) { + if (!subject.empty() && right.GetExpressionClass() == ExpressionClass::BOUND_CONSTANT) { constant = &right.Cast().GetValue(); - } else if (left.GetExpressionClass() == ExpressionClass::BOUND_CONSTANT && IsDirectReference(right)) { - constant = &left.Cast().GetValue(); - comparison_type = FlipComparisonExpression(comparison_type); } else { - return std::string(); + subject = TransformExpressionSubject(identifier_config, column_name, right); + if (!subject.empty() && left.GetExpressionClass() == ExpressionClass::BOUND_CONSTANT) { + constant = &left.Cast().GetValue(); + comparison_type = FlipComparisonExpression(comparison_type); + } } - auto constant_string = query::QueryWriter::WriteConstant(*constant); - auto operator_string = TransformComparison(comparison_type); - return StringUtil::Format("%s %s %s", column_name, operator_string, constant_string); + if (!constant || subject.empty()) { + return string(); + } + return TransformConstantFilter(constant_config, subject, comparison_type, *constant, column_id); } switch (expr.GetExpressionClass()) { @@ -96,29 +158,34 @@ std::string FilterPushdown::TransformExpression(const std::string &column_name, auto &conjunction = expr.Cast(); switch (conjunction.GetExpressionType()) { case ExpressionType::CONJUNCTION_AND: - return CreateExpression(column_name, conjunction.GetChildren(), "AND"); + return CreateExpression(identifier_config, constant_config, column_name, conjunction.GetChildren(), "AND", + column_id); case ExpressionType::CONJUNCTION_OR: - return CreateExpression(column_name, conjunction.GetChildren(), "OR"); + return CreateExpression(identifier_config, constant_config, column_name, conjunction.GetChildren(), "OR", + column_id); default: return std::string(); } } case ExpressionClass::BOUND_OPERATOR: { auto &op = expr.Cast(); + auto subject = op.GetChildren().empty() + ? string() + : TransformExpressionSubject(identifier_config, column_name, *op.GetChildren()[0]); switch (op.GetExpressionType()) { case ExpressionType::OPERATOR_IS_NULL: - if (op.GetChildren().size() == 1 && IsDirectReference(*op.GetChildren()[0])) { - return column_name + " IS NULL"; + if (!subject.empty()) { + return subject + " IS NULL"; } return std::string(); case ExpressionType::OPERATOR_IS_NOT_NULL: - if (op.GetChildren().size() == 1 && IsDirectReference(*op.GetChildren()[0])) { - return column_name + " IS NOT NULL"; + if (!subject.empty()) { + return subject + " IS NOT NULL"; } return std::string(); case ExpressionType::COMPARE_IN: { - if (op.GetChildren().empty() || !IsDirectReference(*op.GetChildren()[0])) { - return std::string(); + if (subject.empty()) { + return string(); } std::string in_list; for (idx_t i = 1; i < op.GetChildren().size(); i++) { @@ -128,10 +195,14 @@ std::string FilterPushdown::TransformExpression(const std::string &column_name, if (!in_list.empty()) { in_list += ", "; } - in_list += - query::QueryWriter::WriteConstant(op.GetChildren()[i]->Cast().GetValue()); + if (IsVirtualColumn(column_id)) { + in_list += "FALSE"; + } else { + in_list += query::QueryWriter::WriteConstant( + identifier_config, op.GetChildren()[i]->Cast().GetValue()); + } } - return column_name + " IN (" + in_list + ")"; + return IsVirtualColumn(column_id) ? "FALSE" : subject + " IN (" + in_list + ")"; } default: return std::string(); @@ -141,11 +212,15 @@ std::string FilterPushdown::TransformExpression(const std::string &column_name, auto &func = expr.Cast(); if (func.Function().GetName() == OptionalFilterScalarFun::NAME && func.BindInfo()) { auto &data = func.BindInfo()->Cast(); - return data.child_filter_expr ? TransformExpression(column_name, *data.child_filter_expr) : std::string(); + return data.child_filter_expr ? TransformExpression(identifier_config, constant_config, column_name, + *data.child_filter_expr, column_id) + : std::string(); } if (func.Function().GetName() == SelectivityOptionalFilterScalarFun::NAME && func.BindInfo()) { auto &data = func.BindInfo()->Cast(); - return data.child_filter_expr ? TransformExpression(column_name, *data.child_filter_expr) : std::string(); + return data.child_filter_expr ? TransformExpression(identifier_config, constant_config, column_name, + *data.child_filter_expr, column_id) + : std::string(); } if (func.Function().GetName() == DynamicFilterScalarFun::NAME) { return std::string(); @@ -158,11 +233,13 @@ std::string FilterPushdown::TransformExpression(const std::string &column_name, } std::string FilterPushdown::TransformFilter(const FilterPushdown::Config &config, const std::string &column_name, - const TableFilter &filter) { - auto query_config = query::QueryWriter::CreateConfig(config.identifier_quote, config.escape_style); - std::string column_name_quoted = query::QueryWriter::WriteQuotedAndEscaped(query_config, column_name); + const TableFilter &filter, column_t column_id) { + auto identifier_config = query::QueryWriter::CreateConfig(config.identifier_quote, config.escape_style); + auto constant_config = query::QueryWriter::CreateConfig(config.constant_quote, config.escape_style, + config.blob_literal_prefix, config.blob_literal_suffix); + std::string column_name_quoted = query::QueryWriter::WriteQuotedAndEscaped(identifier_config, column_name); auto &expr = FilterUtil::GetExpression(filter, "FilterPushdown::TransformFilter"); - return TransformExpression(column_name_quoted, expr); + return TransformExpression(identifier_config, constant_config, column_name_quoted, expr, column_id); } } // namespace table_scan