#include "optimizer/optimizer.h"
#include <cstdio>
#include <cstdlib>
#include <unordered_set>
#include "binder/expression/expression_util.h"
#include "catalog/catalog.h"
#include "common/enums/extend_direction_util.h"
#include "main/client_context.h"
#include "optimizer/acc_hash_join_optimizer.h"
#include "optimizer/agg_key_dependency_optimizer.h"
#include "optimizer/bool_folding_optimizer.h"
#include "optimizer/cardinality_updater.h"
#include "optimizer/correlated_subquery_unnest_solver.h"
#include "optimizer/count_rel_table_optimizer.h"
#include "optimizer/factorization_rewriter.h"
#include "optimizer/filter_push_down_optimizer.h"
#include "optimizer/foreign_join_push_down_optimizer.h"
#include "optimizer/group_key_predicate_push_down_optimizer.h"
#include "optimizer/limit_push_down_optimizer.h"
#include "optimizer/order_by_push_down_optimizer.h"
#include "optimizer/projection_push_down_optimizer.h"
#include "optimizer/remove_factorization_rewriter.h"
#include "optimizer/remove_unnecessary_distinct_optimizer.h"
#include "optimizer/remove_unnecessary_join_optimizer.h"
#include "optimizer/remove_unnecessary_order_by_optimizer.h"
#include "optimizer/schema_populator.h"
#include "optimizer/top_k_optimizer.h"
#include "optimizer/unwind_dedup_optimizer.h"
#include "planner/operator/extend/logical_extend.h"
#include "planner/operator/extend/logical_recursive_extend.h"
#include "planner/operator/logical_aggregate.h"
#include "planner/operator/logical_distinct.h"
#include "planner/operator/logical_explain.h"
#include "planner/operator/logical_filter.h"
#include "planner/operator/logical_hash_join.h"
#include "planner/operator/logical_intersect.h"
#include "planner/operator/logical_limit.h"
#include "planner/operator/logical_operator.h"
#include "planner/operator/logical_order_by.h"
#include "planner/operator/logical_partitioner.h"
#include "planner/operator/logical_projection.h"
#include "planner/operator/logical_unwind.h"
#include "planner/operator/logical_unwind_deduplicate.h"
#include "planner/operator/scan/logical_count_anti_edge_chain.h"
#include "planner/operator/scan/logical_count_extend_chain.h"
#include "planner/operator/scan/logical_count_rel_table.h"
#include "planner/operator/scan/logical_query_primary_key_lookup.h"
#include "planner/operator/scan/logical_reachable_count.h"
#include "planner/operator/scan/logical_rel_degree_table.h"
#include "planner/operator/scan/logical_scan_node_table.h"
#include "planner/operator/sip/logical_semi_masker.h"
#include "transaction/transaction.h"
namespace lbug {
namespace optimizer {
namespace {
std::string compactExprs(const binder::expression_vector& exprs, size_t maxExprs = 2) {
if (exprs.empty()) {
return "";
}
if (exprs.size() <= maxExprs) {
return binder::ExpressionUtil::toString(exprs);
}
binder::expression_vector head(exprs.begin(), exprs.begin() + maxExprs);
return binder::ExpressionUtil::toString(head) + ", ... (" + std::to_string(exprs.size()) +
" total)";
}
std::string truncate(std::string s, size_t maxLen = 120) {
if (s.size() <= maxLen) {
return s;
}
return s.substr(0, maxLen) + "...";
}
std::string tableNames(catalog::Catalog* catalog, const transaction::Transaction* transaction,
const std::vector<common::table_id_t>& tableIDs, size_t maxTables = 2) {
if (tableIDs.empty()) {
return "";
}
std::string result;
auto n = std::min(tableIDs.size(), maxTables);
for (auto i = 0u; i < n; ++i) {
if (i > 0) {
result += "|";
}
if (catalog != nullptr && transaction != nullptr &&
catalog->containsTable(transaction, tableIDs[i])) {
result += catalog->getTableCatalogEntry(transaction, tableIDs[i])->getName();
} else {
result += "table_" + std::to_string(tableIDs[i]);
}
}
if (tableIDs.size() > maxTables) {
result += ", ... (" + std::to_string(tableIDs.size()) + " total)";
}
return result;
}
void dumpLogicalTree(const planner::LogicalOperator* op, int depth,
std::unordered_set<const planner::LogicalOperator*>& visited, catalog::Catalog* catalog,
const transaction::Transaction* transaction) {
if (depth > 40 || visited.contains(op)) {
for (auto i = 0; i < depth; ++i) {
fprintf(stderr, " ");
}
fprintf(stderr, "...\n");
return;
}
visited.insert(op);
for (auto i = 0; i < depth; ++i) {
fprintf(stderr, " ");
}
fprintf(stderr, "%s",
planner::LogicalOperatorUtils::logicalOperatorTypeToString(op->getOperatorType()).c_str());
if (op->getOperatorType() == planner::LogicalOperatorType::EXTEND ||
op->getOperatorType() == planner::LogicalOperatorType::PACKED_EXTEND) {
auto& ext = op->constCast<planner::LogicalExtend>();
auto relInfo = ext.getRel()->detailsToString();
auto relTables = tableNames(catalog, transaction, ext.getRel()->getTableIDs());
if (!relTables.empty()) {
if (relInfo == ext.getRel()->getVariableName()) {
relInfo += ":" + relTables;
} else {
relInfo += "(" + relTables + ")";
}
}
fprintf(stderr, " [%s %s bound=%s nbr=%s]", relInfo.c_str(),
common::ExtendDirectionUtil::toString(ext.getDirection()).c_str(),
ext.getBoundNode()->getUniqueName().c_str(), ext.getNbrNode()->getUniqueName().c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::FILTER) {
fprintf(stderr, " [%s]",
op->constCast<planner::LogicalFilter>().getPredicate()->toString().c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::HASH_JOIN) {
auto& join = op->constCast<planner::LogicalHashJoin>();
auto jt = join.getJoinType() == common::JoinType::INNER ? "INNER" :
join.getJoinType() == common::JoinType::MARK ? "MARK" :
join.getJoinType() == common::JoinType::LEFT ? "LEFT" :
"COUNT";
fprintf(stderr, " [%s keys=%s", jt,
join.getJoinNodeIDs().size() == 1 ? join.getJoinNodeIDs()[0]->toString().c_str() :
"...");
if (join.hasMark()) {
fprintf(stderr, " mark=%s]", join.getMark()->toString().c_str());
} else {
fprintf(stderr, "]");
}
} else if (op->getOperatorType() == planner::LogicalOperatorType::AGGREGATE) {
auto& agg = op->constCast<planner::LogicalAggregate>();
fprintf(stderr, " [keys=%s aggs=%s]", compactExprs(agg.getKeys()).c_str(),
compactExprs(agg.getAggregates()).c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::PROJECTION) {
auto& proj = op->constCast<planner::LogicalProjection>();
fprintf(stderr, " [%s]", truncate(compactExprs(proj.getExpressionsToProject())).c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::ORDER_BY) {
auto& orderBy = op->constCast<planner::LogicalOrderBy>();
fprintf(stderr, " [%s]", truncate(compactExprs(orderBy.getExpressionsToOrderBy())).c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::LIMIT) {
fprintf(stderr, " [%s]",
op->constCast<planner::LogicalLimit>().getExpressionsForPrinting().c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::UNWIND) {
auto& unwind = op->constCast<planner::LogicalUnwind>();
fprintf(stderr, " [%s -> %s]", unwind.getInExpr()->toString().c_str(),
unwind.getOutExpr()->toString().c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::DISTINCT) {
auto& distinct = op->constCast<planner::LogicalDistinct>();
fprintf(stderr, " [%s]", truncate(compactExprs(distinct.getKeys())).c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::RECURSIVE_EXTEND) {
auto& extend = op->constCast<planner::LogicalRecursiveExtend>();
fprintf(stderr, " [%s]", extend.getFunction().getFunctionName().c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::INTERSECT) {
fprintf(stderr, " [%s]",
op->constCast<planner::LogicalIntersect>().getIntersectNodeID()->toString().c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::SEMI_MASKER) {
fprintf(stderr, " [%s]",
op->constCast<planner::LogicalSemiMasker>().getKey()->toString().c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::SCAN_NODE_TABLE) {
auto& scan = op->constCast<planner::LogicalScanNodeTable>();
auto tables = tableNames(catalog, transaction, scan.getTableIDs());
std::string info = tables.empty() ? scan.getNodeID()->toString() :
tables + " " + scan.getNodeID()->toString();
if (scan.getScanType() == planner::LogicalScanNodeTableType::PRIMARY_KEY_SCAN) {
info += " PK_SCAN";
if (scan.getExtraInfo() != nullptr) {
auto& pkInfo = scan.getExtraInfo()->constCast<planner::PrimaryKeyScanInfo>();
info += pkInfo.isRange ? " range" : (" key=" + pkInfo.key->toString());
}
} else if (scan.getScanType() == planner::LogicalScanNodeTableType::SECONDARY_INDEX_SCAN) {
info += " IDX_SCAN";
if (scan.getExtraInfo() != nullptr) {
info += " " +
scan.getExtraInfo()->constCast<planner::SecondaryIndexScanInfo>().indexName;
}
}
fprintf(stderr, " [%s]", truncate(info).c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::COUNT_REL_TABLE) {
auto& count = op->constCast<planner::LogicalCountRelTable>();
fprintf(stderr, " [%s dir=%s bound=%s]", count.getRelGroupEntry()->getName().c_str(),
common::ExtendDirectionUtil::toString(count.getDirection()).c_str(),
count.getBoundNode()->getUniqueName().c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::COUNT_EXTEND_CHAIN) {
auto& count = op->constCast<planner::LogicalCountExtendChain>();
fprintf(stderr, " [hops=%llu tables=", (unsigned long long)count.getHops().size());
bool first = true;
for (auto& hop : count.getHops()) {
for (auto& spec : hop.relScans) {
fprintf(stderr, "%s%s", first ? "" : ",", spec.relTableName.c_str());
first = false;
}
}
fprintf(stderr, "]");
} else if (op->getOperatorType() == planner::LogicalOperatorType::COUNT_ANTI_EDGE_CHAIN) {
auto& count = op->constCast<planner::LogicalCountAntiEdgeChain>();
fprintf(stderr,
" [%s antiDir=%s hops=%llu id<>=", count.getAntiRelEntry()->getName().c_str(),
common::ExtendDirectionUtil::toString(count.getAntiEdgeDir()).c_str(),
(unsigned long long)count.getSuffixHops().size());
fprintf(stderr, "%s]", count.getHasNotEquals() ? "true" : "false");
} else if (op->getOperatorType() == planner::LogicalOperatorType::REACHABLE_COUNT) {
auto& count = op->constCast<planner::LogicalReachableCount>();
fprintf(stderr, " [%s dir=%s bound=%s nbr=%s range=%u..%u]",
count.getRelGroupEntry()->getName().c_str(),
common::ExtendDirectionUtil::toString(count.getDirection()).c_str(),
count.getBoundNode()->getUniqueName().c_str(),
count.getNbrNode()->getUniqueName().c_str(), count.getLowerBound(),
count.getUpperBound());
} else if (op->getOperatorType() == planner::LogicalOperatorType::REL_DEGREE_TABLE) {
auto& degree = op->constCast<planner::LogicalRelDegreeTable>();
fprintf(stderr, " [%s dir=%s mode=%s]", degree.getRelGroupEntry()->getName().c_str(),
common::ExtendDirectionUtil::toString(degree.getDirection()).c_str(),
degree.getMode() == planner::RelDegreeTableMode::ACTIVE_BOUND_COUNT ?
"ACTIVE_BOUND_COUNT" :
degree.getMode() == planner::RelDegreeTableMode::TOP_K_DEGREES ? "TOP_K_DEGREES" :
"OFFSET_COUNT");
} else if (op->getOperatorType() == planner::LogicalOperatorType::QUERY_PRIMARY_KEY_LOOKUP) {
auto& lookup = op->constCast<planner::LogicalQueryPrimaryKeyLookup>();
fprintf(stderr, " [table=%llu key=%s]", (unsigned long long)lookup.getTableID(),
lookup.getKey()->toString().c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::UNWIND_DEDUPLICATE) {
auto& dedup = op->constCast<planner::LogicalUnwindDeduplicate>();
fprintf(stderr, " [keys=%s]",
binder::ExpressionUtil::toString(dedup.getKeyExpressions()).c_str());
} else if (op->getOperatorType() == planner::LogicalOperatorType::PARTITIONER) {
fprintf(stderr, " [keys=%llu]",
(unsigned long long)op->constCast<planner::LogicalPartitioner>()
.getInfo()
.getNumInfos());
} else {
auto exprs = op->getExpressionsForPrinting();
if (!exprs.empty()) {
fprintf(stderr, " [%s]", truncate(exprs).c_str());
}
}
fprintf(stderr, "\n");
for (auto i = 0u; i < op->getNumChildren(); ++i) {
dumpLogicalTree(op->getChild(i).get(), depth + 1, visited, catalog, transaction);
}
}
void dumpLogicalPlan(const planner::LogicalPlan* plan, const char* label, catalog::Catalog* catalog,
const transaction::Transaction* transaction) {
fprintf(stderr, "=== LOGICAL PLAN (%s) ===\n", label);
auto* root = plan->getLastOperator().get();
if (root == nullptr) {
return;
}
std::unordered_set<const planner::LogicalOperator*> visited;
dumpLogicalTree(root, 0, visited, catalog, transaction);
fprintf(stderr, "=== END LOGICAL PLAN ===\n");
}
}
void Optimizer::optimize(planner::LogicalPlan* plan, main::ClientContext* context,
const planner::CardinalityEstimator& cardinalityEstimator) {
static const bool dumpLogicalEnabled = getenv("LBUG_DUMP_LOGICAL") != nullptr;
if (dumpLogicalEnabled) {
dumpLogicalPlan(plan, "before optimization", catalog::Catalog::Get(*context),
transaction::Transaction::Get(*context));
}
if (context->getClientConfig()->enablePlanOptimizer) {
auto removeFactorizationRewriter = RemoveFactorizationRewriter();
removeFactorizationRewriter.rewrite(plan);
auto correlatedSubqueryUnnestSolver = CorrelatedSubqueryUnnestSolver(nullptr);
correlatedSubqueryUnnestSolver.solve(plan->getLastOperator().get());
auto removeUnnecessaryJoinOptimizer = RemoveUnnecessaryJoinOptimizer();
removeUnnecessaryJoinOptimizer.rewrite(plan);
auto unwindDedupOptimizer = UnwindDedupOptimizer();
unwindDedupOptimizer.rewrite(plan);
auto countRelTableOptimizer = CountRelTableOptimizer(context);
countRelTableOptimizer.rewrite(plan);
auto foreignJoinPushDownOptimizer = ForeignJoinPushDownOptimizer(context);
foreignJoinPushDownOptimizer.rewrite(plan);
auto boolFoldingOptimizer = BoolFoldingOptimizer();
boolFoldingOptimizer.rewrite(plan);
auto groupKeyPredicatePushDownOptimizer = GroupKeyPredicatePushDownOptimizer();
groupKeyPredicatePushDownOptimizer.rewrite(plan);
auto filterPushDownOptimizer = FilterPushDownOptimizer(context, &cardinalityEstimator);
filterPushDownOptimizer.rewrite(plan);
auto projectionPushDownOptimizer =
ProjectionPushDownOptimizer(context->getClientConfig()->recursivePatternSemantic);
projectionPushDownOptimizer.rewrite(plan);
auto orderByPushDownOptimizer = OrderByPushDownOptimizer();
orderByPushDownOptimizer.rewrite(plan);
auto limitPushDownOptimizer = LimitPushDownOptimizer();
limitPushDownOptimizer.rewrite(plan);
if (context->getClientConfig()->enableSemiMask) {
auto hashJoinSIPOptimizer = HashJoinSIPOptimizer();
hashJoinSIPOptimizer.rewrite(plan);
}
auto topKOptimizer = TopKOptimizer();
topKOptimizer.rewrite(plan);
countRelTableOptimizer.rewrite(plan);
auto factorizationRewriter = FactorizationRewriter();
factorizationRewriter.rewrite(plan);
auto aggKeyDependencyOptimizer = AggKeyDependencyOptimizer();
aggKeyDependencyOptimizer.rewrite(plan);
auto removeUnnecessaryDistinctOptimizer = RemoveUnnecessaryDistinctOptimizer();
removeUnnecessaryDistinctOptimizer.rewrite(plan);
auto removeUnnecessaryOrderByOptimizer = RemoveUnnecessaryOrderByOptimizer();
removeUnnecessaryOrderByOptimizer.rewrite(plan);
if (plan->getLastOperatorRef().getOperatorType() == planner::LogicalOperatorType::EXPLAIN) {
const auto& explain = plan->getLastOperatorRef().cast<planner::LogicalExplain>();
if (explain.getExplainType() == common::ExplainType::LOGICAL_PLAN) {
auto cardinalityUpdater = CardinalityUpdater(cardinalityEstimator,
transaction::Transaction::Get(*context));
cardinalityUpdater.rewrite(plan);
}
}
} else {
auto schemaPopulator = SchemaPopulator{};
schemaPopulator.rewrite(plan);
}
if (dumpLogicalEnabled) {
dumpLogicalPlan(plan, "after optimization", catalog::Catalog::Get(*context),
transaction::Transaction::Get(*context));
}
}
} }