#include <algorithm>
#include <unordered_set>
#include "binder/binder.h"
#include "binder/expression/aggregate_function_expression.h"
#include "binder/expression/expression_util.h"
#include "binder/expression/node_rel_expression.h"
#include "binder/expression/property_expression.h"
#include "binder/expression/scalar_function_expression.h"
#include "binder/expression/subquery_expression.h"
#include "binder/expression_binder.h"
#include "binder/expression_visitor.h"
#include "common/exception/runtime.h"
#include "function/aggregate/count.h"
#include "function/list/vector_list_functions.h"
#include "planner/operator/factorization/flatten_resolver.h"
#include "planner/operator/logical_accumulate.h"
#include "planner/operator/logical_aggregate.h"
#include "planner/operator/logical_distinct.h"
#include "planner/operator/logical_flatten.h"
#include "planner/operator/logical_hash_join.h"
#include "planner/operator/logical_projection.h"
#include "planner/operator/logical_unwind.h"
#include "planner/operator/scan/logical_query_primary_key_lookup.h"
#include "planner/planner.h"
#include "storage/storage_manager.h"
#include "storage/table/node_table.h"
using namespace lbug::binder;
using namespace lbug::common;
namespace lbug {
namespace planner {
static expression_vector getDependentExprs(std::shared_ptr<Expression> expr, const Schema& schema) {
auto analyzer = GroupDependencyAnalyzer(true , schema);
analyzer.visit(expr);
return analyzer.getDependentExprs();
}
static bool isNodePrimaryKey(const Expression& expression, const NodeExpression& node,
table_id_t tableID) {
if (expression.expressionType != ExpressionType::PROPERTY) {
return false;
}
auto& property = expression.constCast<PropertyExpression>();
return property.getVariableName() == node.getInternalID()->getVariableName() &&
property.isPrimaryKey(tableID);
}
bool Planner::findCorrelatedPrimaryKeyLookupKey(const QueryGraphCollection& queryGraphCollection,
const expression_vector& predicates, const expression_vector& corrExprs,
const Schema& outerSchema, std::shared_ptr<Expression>& key,
expression_vector& residualPredicates) {
key = nullptr;
residualPredicates.clear();
if (queryGraphCollection.getNumQueryGraphs() != 1) {
return false;
}
auto queryGraph = queryGraphCollection.getQueryGraph(0);
if (queryGraph->getNumQueryNodes() != 1 || queryGraph->getNumQueryRels() != 0) {
return false;
}
auto node = queryGraph->getQueryNode(0);
auto tableIDs = node->getTableIDs();
if (tableIDs.size() != 1) {
return false;
}
auto tableID = tableIDs[0];
auto table = storage::StorageManager::Get(*clientContext)
->getTable(tableID)
->ptrCast<storage::NodeTable>();
if (table->tryGetPrimaryKeyIndex() == nullptr) {
return false;
}
std::unordered_set<std::string> corrNames;
for (auto& expr : corrExprs) {
corrNames.insert(expr->getUniqueName());
}
for (auto& predicate : predicates) {
if (predicate->expressionType != ExpressionType::EQUALS) {
residualPredicates.push_back(predicate);
continue;
}
auto lhs = predicate->getChild(0);
auto rhs = predicate->getChild(1);
if (isNodePrimaryKey(*rhs, *node, tableID)) {
std::swap(lhs, rhs);
}
auto usableKey = false;
if (key == nullptr && isNodePrimaryKey(*lhs, *node, tableID)) {
usableKey = true;
for (auto& dep : getDependentExprs(rhs, outerSchema)) {
if (!corrNames.contains(dep->getUniqueName())) {
usableKey = false;
break;
}
}
if (getDependentExprs(rhs, outerSchema).empty()) {
usableKey = false;
}
}
if (!usableKey) {
residualPredicates.push_back(predicate);
continue;
}
key = rhs;
}
if (key == nullptr) {
return false;
}
return true;
}
bool Planner::tryPlanCorrelatedPrimaryKeyLookup(const QueryGraphCollection& queryGraphCollection,
const expression_vector& predicates, const expression_vector& corrExprs,
const Schema& outerSchema, cardinality_t corrExprsCard, LogicalPlan& rightPlan) {
std::shared_ptr<Expression> key;
expression_vector residualPredicates;
if (!findCorrelatedPrimaryKeyLookupKey(queryGraphCollection, predicates, corrExprs, outerSchema,
key, residualPredicates)) {
return false;
}
auto node = queryGraphCollection.getQueryGraph(0)->getQueryNode(0);
auto tableID = node->getTableIDs()[0];
appendExpressionsScan(corrExprs, rightPlan);
rightPlan.getLastOperator()->setCardinality(corrExprsCard);
appendDistinct(corrExprs, rightPlan);
appendFlattens(rightPlan.getSchema()->getGroupsPosInScope(), rightPlan);
auto properties = getProperties(*node);
properties.erase(std::remove_if(properties.begin(), properties.end(),
[](const std::shared_ptr<Expression>& expression) {
return expression->constCast<PropertyExpression>().isInternalID();
}),
properties.end());
const auto dependentExprs = getDependentExprs(key, *rightPlan.getSchema());
DASSERT(!dependentExprs.empty());
const auto outputGroupPos = rightPlan.getSchema()->getGroupPos(*dependentExprs[0]);
for ([[maybe_unused]] auto& dependentExpr : dependentExprs) {
DASSERT(rightPlan.getSchema()->getGroupPos(*dependentExpr) == outputGroupPos);
}
auto lookup = std::make_shared<LogicalQueryPrimaryKeyLookup>(tableID, node->getInternalID(),
properties, key, outputGroupPos, rightPlan.getLastOperator());
lookup->computeFactorizedSchema();
lookup->setCardinality(rightPlan.getCardinality());
rightPlan.setLastOperator(std::move(lookup));
appendFilters(residualPredicates, rightPlan);
return true;
}
bool Planner::tryPlanQueryPrimaryKeyLookup(const QueryGraphCollection& queryGraphCollection,
const expression_vector& predicates, LogicalPlan& plan) {
if (queryGraphCollection.getNumQueryGraphs() != 1) {
return false;
}
auto queryGraph = queryGraphCollection.getQueryGraph(0);
if (queryGraph->getNumQueryNodes() != 1 || queryGraph->getNumQueryRels() != 0) {
return false;
}
auto node = queryGraph->getQueryNode(0);
auto tableIDs = node->getTableIDs();
if (tableIDs.size() != 1 || plan.getSchema()->isExpressionInScope(*node->getInternalID())) {
return false;
}
auto tableID = tableIDs[0];
auto table = storage::StorageManager::Get(*clientContext)
->getTable(tableID)
->ptrCast<storage::NodeTable>();
if (table->tryGetPrimaryKeyIndex() == nullptr) {
return false;
}
std::shared_ptr<Expression> key;
expression_vector residualPredicates;
for (auto& predicate : predicates) {
if (predicate->expressionType != ExpressionType::EQUALS) {
residualPredicates.push_back(predicate);
continue;
}
auto lhs = predicate->getChild(0);
auto rhs = predicate->getChild(1);
if (isNodePrimaryKey(*rhs, *node, tableID)) {
std::swap(lhs, rhs);
}
if (key == nullptr && isNodePrimaryKey(*lhs, *node, tableID) &&
plan.getSchema()->evaluable(*rhs) &&
!getDependentExprs(rhs, *plan.getSchema()).empty()) {
key = rhs;
} else {
residualPredicates.push_back(predicate);
}
}
if (key == nullptr) {
return false;
}
appendFlattens(plan.getSchema()->getGroupsPosInScope(), plan);
auto properties = getProperties(*node);
properties.erase(std::remove_if(properties.begin(), properties.end(),
[](const std::shared_ptr<Expression>& expression) {
return expression->constCast<PropertyExpression>().isInternalID();
}),
properties.end());
const auto dependentExprs = getDependentExprs(key, *plan.getSchema());
DASSERT(!dependentExprs.empty());
const auto outputGroupPos = plan.getSchema()->getGroupPos(*dependentExprs[0]);
for ([[maybe_unused]] auto& dependentExpr : dependentExprs) {
DASSERT(plan.getSchema()->getGroupPos(*dependentExpr) == outputGroupPos);
}
auto lookup = std::make_shared<LogicalQueryPrimaryKeyLookup>(tableID, node->getInternalID(),
properties, key, outputGroupPos, plan.getLastOperator());
lookup->computeFactorizedSchema();
lookup->setCardinality(plan.getCardinality());
plan.setLastOperator(std::move(lookup));
appendFilters(residualPredicates, plan);
return true;
}
expression_vector Planner::getCorrelatedExprs(const QueryGraphCollection& collection,
const expression_vector& predicates, Schema* outerSchema) {
expression_vector result;
for (auto& predicate : predicates) {
for (auto& expression : getDependentExprs(predicate, *outerSchema)) {
result.push_back(expression);
}
}
for (auto& node : collection.getQueryNodes()) {
if (outerSchema->isExpressionInScope(*node->getInternalID())) {
result.push_back(node->getInternalID());
}
}
return ExpressionUtil::removeDuplication(result);
}
static bool isRoutableConstantEquality(const std::shared_ptr<Expression>& predicate,
const binder::QueryGraphCollection& collection) {
if (predicate->expressionType != common::ExpressionType::EQUALS) {
return false;
}
std::unordered_set<std::string> innerNames;
for (auto& node : collection.getQueryNodes()) {
innerNames.insert(node->getUniqueName());
}
for (auto& rel : collection.getQueryRels()) {
innerNames.insert(rel->getUniqueName());
}
bool hasConstantSide = false;
for (auto i = 0u; i < 2u; ++i) {
auto collector = DependentVarNameCollector();
collector.visit(predicate->getChild(i));
if (collector.getVarNames().empty()) {
hasConstantSide = true;
continue;
}
for (auto& varName : collector.getVarNames()) {
if (!innerNames.contains(varName)) {
return false;
}
}
}
return hasConstantSide;
}
class SubqueryPredicatePullUpAnalyzer {
public:
SubqueryPredicatePullUpAnalyzer(const Schema& schema,
const QueryGraphCollection& queryGraphCollection, bool routeConstantEqualities = true)
: schema{schema}, queryGraphCollection{queryGraphCollection},
routeConstantEqualities{routeConstantEqualities} {}
bool analyze(const expression_vector& predicates) {
expression_vector correlatedPredicates;
for (auto& predicate : predicates) {
if (getDependentExprs(predicate, schema).empty() ||
(routeConstantEqualities &&
isRoutableConstantEquality(predicate, queryGraphCollection))) {
nonCorrelatedPredicates.push_back(predicate);
} else {
correlatedPredicates.push_back(predicate);
}
}
for (auto predicate : correlatedPredicates) {
auto [left, right] = analyze(predicate);
if (left == nullptr) {
return false;
}
joinConditions.emplace_back(left, right);
}
for (auto& node : queryGraphCollection.getQueryNodes()) {
if (schema.isExpressionInScope(*node->getInternalID())) {
joinConditions.emplace_back(node->getInternalID(), node->getInternalID());
}
}
return true;
}
expression_vector getNonCorrelatedPredicates() const { return nonCorrelatedPredicates; }
std::vector<binder::expression_pair> getJoinConditions() const { return joinConditions; }
expression_vector getCorrelatedInternalIDs() const {
expression_vector exprs;
for (auto& node : queryGraphCollection.getQueryNodes()) {
if (schema.isExpressionInScope(*node->getInternalID())) {
exprs.push_back(node->getInternalID());
}
}
return exprs;
}
private:
expression_pair analyze(std::shared_ptr<Expression> predicate) {
if (predicate->expressionType != common::ExpressionType::EQUALS) {
return {nullptr, nullptr};
}
auto left = predicate->getChild(0);
auto right = predicate->getChild(1);
if (isUnnestableJoinCondition(*left, *right)) {
return {left, right};
}
if (isUnnestableJoinCondition(*right, *left)) {
return {right, left};
}
return {nullptr, nullptr};
}
bool isUnnestableJoinCondition(const Expression& left, const Expression& right) {
return right.expressionType == ExpressionType::PROPERTY &&
schema.isExpressionInScope(left) && !schema.isExpressionInScope(right);
}
private:
const Schema& schema;
const QueryGraphCollection& queryGraphCollection;
bool routeConstantEqualities;
expression_vector nonCorrelatedPredicates;
std::vector<binder::expression_pair> joinConditions;
};
void Planner::planOptionalMatch(const QueryGraphCollection& queryGraphCollection,
const expression_vector& predicates, LogicalPlan& leftPlan,
std::shared_ptr<BoundJoinHintNode> hint) {
planOptionalMatch(queryGraphCollection, predicates, nullptr , leftPlan,
std::move(hint));
}
static bool isListContainsFunc(const std::shared_ptr<Expression>& pred) {
if (pred->expressionType != ExpressionType::FUNCTION) {
return false;
}
auto& funcExpr = pred->constCast<ScalarFunctionExpression>();
auto name = funcExpr.getFunction().name;
return name == function::ListContainsFunction::name || name == function::ListHasFunction::name;
}
static bool containsRef(const std::shared_ptr<Expression>& expr, const std::string& uniqueName) {
if (expr->getUniqueName() == uniqueName) {
return true;
}
for (auto& child : expr->getChildren()) {
if (containsRef(child, uniqueName)) {
return true;
}
}
return false;
}
static bool containsSubqueryOrLambda(const std::shared_ptr<Expression>& expr) {
if (expr->expressionType == ExpressionType::SUBQUERY ||
expr->expressionType == ExpressionType::LAMBDA) {
return true;
}
for (auto& child : expr->getChildren()) {
if (containsSubqueryOrLambda(child)) {
return true;
}
}
return false;
}
static planner::LogicalOperator* findCollectAggregate(planner::LogicalOperator* op,
const std::string& collectVarName, std::shared_ptr<Expression>& collectArg,
expression_vector& groupKeys) {
if (op->getOperatorType() == LogicalOperatorType::AGGREGATE) {
auto& aggregate = op->constCast<LogicalAggregate>();
for (auto& aggExpr : aggregate.getAggregates()) {
if (aggExpr->expressionType != ExpressionType::AGGREGATE_FUNCTION) {
continue;
}
auto& aggFunc = aggExpr->constCast<AggregateFunctionExpression>();
if (aggFunc.getFunction().name != function::CollectFunction::name ||
aggFunc.getNumChildren() != 1) {
continue;
}
if (aggExpr->getUniqueName() == collectVarName) {
collectArg = aggFunc.getChild(0);
groupKeys = aggregate.getKeys();
return groupKeys.empty() ? nullptr : op;
}
}
}
for (auto i = 0u; i < op->getNumChildren(); ++i) {
if (auto found = findCollectAggregate(op->getChild(i).get(), collectVarName, collectArg,
groupKeys)) {
return found;
}
}
return nullptr;
}
bool isNullFreeOperator(LogicalOperatorType type) {
switch (type) {
case LogicalOperatorType::AGGREGATE:
case LogicalOperatorType::COUNT_ANTI_EDGE_CHAIN:
case LogicalOperatorType::COUNT_EXTEND_CHAIN:
case LogicalOperatorType::COUNT_REL_TABLE:
case LogicalOperatorType::CROSS_PRODUCT:
case LogicalOperatorType::DISTINCT:
case LogicalOperatorType::DUMMY_SCAN:
case LogicalOperatorType::EMPTY_RESULT:
case LogicalOperatorType::EXTEND:
case LogicalOperatorType::FILTER:
case LogicalOperatorType::FLATTEN:
case LogicalOperatorType::INDEX_LOOK_UP:
case LogicalOperatorType::INTERSECT:
case LogicalOperatorType::LIMIT:
case LogicalOperatorType::MULTIPLICITY_REDUCER:
case LogicalOperatorType::NODE_LABEL_FILTER:
case LogicalOperatorType::ORDER_BY:
case LogicalOperatorType::PACKED_EXTEND:
case LogicalOperatorType::PATH_PROPERTY_PROBE:
case LogicalOperatorType::PROJECTION:
case LogicalOperatorType::QUERY_PRIMARY_KEY_LOOKUP:
case LogicalOperatorType::REACHABLE_COUNT:
case LogicalOperatorType::RECURSIVE_EXTEND:
case LogicalOperatorType::REL_DEGREE_TABLE:
case LogicalOperatorType::SCAN_NODE_TABLE:
case LogicalOperatorType::SEMI_MASKER:
return true;
case LogicalOperatorType::ACCUMULATE:
case LogicalOperatorType::ALTER:
case LogicalOperatorType::ANALYZE:
case LogicalOperatorType::ATTACH_DATABASE:
case LogicalOperatorType::COPY_FROM:
case LogicalOperatorType::COPY_TO:
case LogicalOperatorType::CREATE_GRAPH:
case LogicalOperatorType::CREATE_INDEX:
case LogicalOperatorType::CREATE_MACRO:
case LogicalOperatorType::CREATE_SEQUENCE:
case LogicalOperatorType::CREATE_TABLE:
case LogicalOperatorType::CREATE_TYPE:
case LogicalOperatorType::DELETE:
case LogicalOperatorType::DETACH_DATABASE:
case LogicalOperatorType::DROP:
case LogicalOperatorType::DUMMY_SINK:
case LogicalOperatorType::EXPLAIN:
case LogicalOperatorType::EXPRESSIONS_SCAN:
case LogicalOperatorType::EXTENSION:
case LogicalOperatorType::EXTENSION_CLAUSE:
case LogicalOperatorType::EXPORT_DATABASE:
case LogicalOperatorType::HASH_JOIN:
case LogicalOperatorType::IMPORT_DATABASE:
case LogicalOperatorType::INSERT:
case LogicalOperatorType::MERGE:
case LogicalOperatorType::NOOP:
case LogicalOperatorType::PARTITIONER:
case LogicalOperatorType::SET_PROPERTY:
case LogicalOperatorType::STANDALONE_CALL:
case LogicalOperatorType::TABLE_FUNCTION_CALL:
case LogicalOperatorType::TRANSACTION:
case LogicalOperatorType::UNION_ALL:
case LogicalOperatorType::UNWIND:
case LogicalOperatorType::UNWIND_DEDUPLICATE:
case LogicalOperatorType::USE_DATABASE:
case LogicalOperatorType::USE_GRAPH:
return false;
}
throw common::RuntimeException("Unknown logical operator type.");
}
static bool subtreeHasNullSupply(planner::LogicalOperator* op) {
switch (op->getOperatorType()) {
case LogicalOperatorType::HASH_JOIN: {
auto joinType = op->constCast<LogicalHashJoin>().getJoinType();
if (joinType == JoinType::LEFT || joinType == JoinType::MARK) {
return true;
}
break;
}
case LogicalOperatorType::ACCUMULATE: {
if (op->constCast<LogicalAccumulate>().getAccumulateType() != AccumulateType::REGULAR) {
return true;
}
break;
}
default:
if (!isNullFreeOperator(op->getOperatorType())) {
return true;
}
break;
}
for (auto i = 0u; i < op->getNumChildren(); ++i) {
if (subtreeHasNullSupply(op->getChild(i).get())) {
return true;
}
}
return false;
}
static bool findPath(planner::LogicalOperator* op, planner::LogicalOperator* target,
std::vector<planner::LogicalOperator*>& path) {
if (op == target) {
return true;
}
for (auto i = 0u; i < op->getNumChildren(); ++i) {
if (findPath(op->getChild(i).get(), target, path)) {
path.push_back(op);
return true;
}
}
return false;
}
static bool pathHasFanOut(const std::vector<planner::LogicalOperator*>& path) {
for (auto* op : path) {
switch (op->getOperatorType()) {
case LogicalOperatorType::PROJECTION:
case LogicalOperatorType::FILTER:
case LogicalOperatorType::FLATTEN:
case LogicalOperatorType::DISTINCT:
case LogicalOperatorType::ORDER_BY:
case LogicalOperatorType::LIMIT:
case LogicalOperatorType::INTERSECT:
case LogicalOperatorType::SEMI_MASKER:
case LogicalOperatorType::NODE_LABEL_FILTER:
case LogicalOperatorType::MULTIPLICITY_REDUCER:
break;
case LogicalOperatorType::ACCUMULATE:
if (op->constCast<LogicalAccumulate>().getAccumulateType() != AccumulateType::REGULAR) {
return true;
}
break;
default:
return true;
}
}
return false;
}
static std::shared_ptr<Expression> findInScope(planner::Schema& schema,
const std::string& uniqueName) {
for (auto& expr : schema.getExpressionsInScope()) {
if (expr->getUniqueName() == uniqueName) {
return expr;
}
}
return nullptr;
}
static std::unordered_set<std::string> collectOuterScopeNames(planner::Schema& schema) {
std::unordered_set<std::string> names;
for (auto& expr : schema.getExpressionsInScope()) {
names.insert(expr->getUniqueName());
}
return names;
}
static std::unordered_set<std::string> collectGraphVarNames(
const QueryGraphCollection& collection) {
std::unordered_set<std::string> names;
for (auto& node : collection.getQueryNodes()) {
names.insert(node->getUniqueName());
}
for (auto& rel : collection.getQueryRels()) {
names.insert(rel->getUniqueName());
}
return names;
}
bool Planner::tryUnnestCollectMembership(const QueryGraphCollection& queryGraphCollection,
expression_vector& predicates, LogicalPlan& leftPlan) {
bool anyContains = false;
for (auto& pred : predicates) {
if (isListContainsFunc(pred)) {
anyContains = true;
break;
}
}
if (!anyContains || leftPlan.isEmpty()) {
return false;
}
auto outerScope = collectOuterScopeNames(*leftPlan.getSchema());
auto legVars = collectGraphVarNames(queryGraphCollection);
struct Convertible {
size_t predIdx;
std::shared_ptr<Expression> elemNode;
};
std::vector<Convertible> convertibles;
std::string collectName;
for (auto i = 0u; i < predicates.size(); ++i) {
auto& pred = predicates[i];
if (!isListContainsFunc(pred) || pred->getNumChildren() != 2) {
continue;
}
auto c0 = pred->getChild(0);
auto c1 = pred->getChild(1);
auto c0Outer = outerScope.contains(c0->getUniqueName());
auto c1Outer = outerScope.contains(c1->getUniqueName());
std::shared_ptr<Expression> listObj;
std::shared_ptr<Expression> elemObj;
if (c0Outer && !c1Outer) {
listObj = c0;
elemObj = c1;
} else if (c1Outer && !c0Outer) {
listObj = c1;
elemObj = c0;
} else {
return false;
}
if (!ExpressionUtil::isNodePattern(*elemObj) && !ExpressionUtil::isRelPattern(*elemObj)) {
return false;
}
if (!legVars.contains(elemObj->getUniqueName()) ||
legVars.contains(listObj->getUniqueName())) {
return false;
}
if (collectName.empty()) {
collectName = listObj->getUniqueName();
} else if (collectName != listObj->getUniqueName()) {
return false;
}
convertibles.push_back({i, elemObj});
}
if (convertibles.empty()) {
return false;
}
std::unordered_set<size_t> convertedIdxs;
for (auto& conv : convertibles) {
convertedIdxs.insert(conv.predIdx);
}
for (auto i = 0u; i < predicates.size(); ++i) {
if (convertedIdxs.contains(i)) {
continue;
}
if (containsSubqueryOrLambda(predicates[i]) || containsRef(predicates[i], collectName)) {
return false;
}
}
std::shared_ptr<Expression> collectArg;
expression_vector groupKeys;
auto aggOp =
findCollectAggregate(leftPlan.getLastOperator().get(), collectName, collectArg, groupKeys);
if (aggOp == nullptr || collectArg == nullptr || !ExpressionUtil::isNodePattern(*collectArg)) {
return false;
}
if (subtreeHasNullSupply(aggOp)) {
return false;
}
std::vector<planner::LogicalOperator*> path;
if (!findPath(leftPlan.getLastOperator().get(), aggOp, path) || pathHasFanOut(path)) {
return false;
}
auto& schema = *leftPlan.getSchema();
auto listObj = findInScope(schema, collectName);
if (listObj == nullptr) {
return false;
}
expression_vector distinctKeys;
for (auto& key : groupKeys) {
auto scoped = findInScope(schema, key->getUniqueName());
if (scoped == nullptr) {
return false;
}
distinctKeys.push_back(scoped);
}
Binder binder(clientContext);
ExpressionBinder expressionBinder(&binder, clientContext);
expression_vector newPreds;
for (auto& conv : convertibles) {
newPreds.push_back(
expressionBinder.createEqualityComparisonExpression(conv.elemNode, collectArg));
}
auto elementID = collectArg->constCast<NodeExpression>().getInternalID();
auto unwind =
std::make_shared<LogicalUnwind>(listObj, collectArg, elementID, leftPlan.getLastOperator());
appendFlattens(unwind->getGroupsPosToFlatten(), leftPlan);
unwind->setChild(0, leftPlan.getLastOperator());
unwind->computeFactorizedSchema();
leftPlan.setLastOperator(unwind);
expression_vector payloads;
for (auto& expr : leftPlan.getSchema()->getExpressionsInScope()) {
auto name = expr->getUniqueName();
if (name == collectName || name == collectArg->getUniqueName()) {
continue;
}
bool isKey = name == elementID->getUniqueName();
for (auto& key : distinctKeys) {
isKey = isKey || name == key->getUniqueName();
}
if (!isKey) {
payloads.push_back(expr);
}
}
distinctKeys.push_back(elementID);
auto distinct = std::make_shared<LogicalDistinct>(distinctKeys, leftPlan.getLastOperator());
distinct->setPayloads(std::move(payloads));
appendFlattens(distinct->getGroupsPosToFlatten(), leftPlan);
distinct->setChild(0, leftPlan.getLastOperator());
distinct->computeFactorizedSchema();
leftPlan.setLastOperator(distinct);
for (auto i = 0u; i < convertibles.size(); ++i) {
predicates[convertibles[i].predIdx] = newPreds[i];
}
return true;
}
static common::ExtendDirection chainExtendDirection(const RelExpression& rel,
const NodeExpression& boundNode) {
if (rel.getDirectionType() == binder::RelDirectionType::BOTH) {
return common::ExtendDirection::BOTH;
}
return *rel.getSrcNode() == boundNode ? common::ExtendDirection::FWD :
common::ExtendDirection::BWD;
}
static planner::LogicalOperator* findPropertyCollectAggregate(planner::LogicalOperator* op,
const std::string& collectVarName, std::shared_ptr<Expression>& collectArg,
expression_vector& groupKeys) {
if (op->getOperatorType() == LogicalOperatorType::AGGREGATE) {
auto& aggregate = op->constCast<LogicalAggregate>();
for (auto& aggExpr : aggregate.getAggregates()) {
if (aggExpr->expressionType != ExpressionType::AGGREGATE_FUNCTION) {
continue;
}
auto& aggFunc = aggExpr->constCast<AggregateFunctionExpression>();
if (aggFunc.getFunction().name != function::CollectFunction::name ||
aggFunc.getNumChildren() != 1) {
continue;
}
if (aggExpr->getUniqueName() == collectVarName) {
collectArg = aggFunc.getChild(0);
groupKeys = aggregate.getKeys();
return op;
}
}
}
for (auto i = 0u; i < op->getNumChildren(); ++i) {
if (auto found = findPropertyCollectAggregate(op->getChild(i).get(), collectVarName,
collectArg, groupKeys)) {
return found;
}
}
return nullptr;
}
static bool isVarFree(const std::shared_ptr<Expression>& expr) {
auto collector = DependentVarNameCollector();
collector.visit(expr);
return collector.getVarNames().empty();
}
bool Planner::tryPlanPropertySeededChain(const QueryGraphCollection& queryGraphCollection,
const expression_vector& predicates, LogicalPlan& leftPlan) {
if (leftPlan.isEmpty() || leftPlan.hasUpdate() ||
queryGraphCollection.getNumQueryGraphs() != 1) {
return false;
}
auto queryGraph = queryGraphCollection.getQueryGraph(0);
if (queryGraph->getNumQueryRels() == 0) {
return false;
}
for (auto i = 0u; i < queryGraph->getNumQueryRels(); ++i) {
if (queryGraph->getQueryRel(i)->getRelType() != QueryRelType::NON_RECURSIVE) {
return false;
}
}
auto outerScope = collectOuterScopeNames(*leftPlan.getSchema());
auto legVars = collectGraphVarNames(queryGraphCollection);
int convertIdx = -1;
std::shared_ptr<Expression> listObj;
std::shared_ptr<Expression> elemProp;
for (auto i = 0u; i < predicates.size(); ++i) {
auto& pred = predicates[i];
if (!isListContainsFunc(pred) || pred->getNumChildren() != 2) {
continue;
}
auto c0 = pred->getChild(0);
auto c1 = pred->getChild(1);
auto c0Outer = outerScope.contains(c0->getUniqueName());
auto c1Outer = outerScope.contains(c1->getUniqueName());
std::shared_ptr<Expression> listSide;
std::shared_ptr<Expression> elemSide;
if (c0Outer && !c1Outer) {
listSide = c0;
elemSide = c1;
} else if (c1Outer && !c0Outer) {
listSide = c1;
elemSide = c0;
} else {
return false;
}
if (elemSide->expressionType != ExpressionType::PROPERTY ||
!elemSide->constCast<PropertyExpression>().isSingleLabel()) {
return false;
}
auto& elemPropRef = elemSide->constCast<PropertyExpression>();
if (!legVars.contains(elemPropRef.getVariableName()) ||
legVars.contains(listSide->getUniqueName())) {
return false;
}
if (convertIdx >= 0) {
return false;
}
convertIdx = static_cast<int>(i);
listObj = listSide;
elemProp = elemSide;
}
if (convertIdx < 0) {
return false;
}
auto collectName = listObj->getUniqueName();
for (auto i = 0u; i < predicates.size(); ++i) {
if (static_cast<int>(i) == convertIdx) {
continue;
}
if (containsSubqueryOrLambda(predicates[i]) || containsRef(predicates[i], collectName)) {
return false;
}
}
std::shared_ptr<Expression> collectArg;
expression_vector groupKeys;
auto aggOp = findPropertyCollectAggregate(leftPlan.getLastOperator().get(), collectName,
collectArg, groupKeys);
if (aggOp == nullptr || collectArg == nullptr ||
collectArg->expressionType != ExpressionType::PROPERTY ||
!collectArg->constCast<PropertyExpression>().isSingleLabel()) {
return false;
}
auto& collectProp = collectArg->constCast<PropertyExpression>();
auto& elemPropRef = elemProp->constCast<PropertyExpression>();
if (collectProp.getSingleTableID() != elemPropRef.getSingleTableID() ||
collectProp.getPropertyName() != elemPropRef.getPropertyName()) {
return false;
}
if (subtreeHasNullSupply(aggOp)) {
return false;
}
std::vector<planner::LogicalOperator*> path;
if (!findPath(leftPlan.getLastOperator().get(), aggOp, path) || pathHasFanOut(path)) {
return false;
}
auto& schema = *leftPlan.getSchema();
auto listScoped = findInScope(schema, collectName);
if (listScoped == nullptr) {
return false;
}
expression_vector distinctKeys;
for (auto& key : groupKeys) {
auto scoped = findInScope(schema, key->getUniqueName());
if (scoped == nullptr) {
return false;
}
distinctKeys.push_back(scoped);
}
std::shared_ptr<NodeExpression> seedNode;
for (auto i = 0u; i < queryGraph->getNumQueryNodes(); ++i) {
auto node = queryGraph->getQueryNode(i);
if (node->getUniqueName() == elemPropRef.getVariableName()) {
seedNode = node;
break;
}
}
if (seedNode == nullptr || seedNode->getTableIDs().size() != 1) {
return false;
}
auto seedTableID = seedNode->getTableIDs()[0];
if (!elemPropRef.isPrimaryKey(seedTableID)) {
return false;
}
auto seedTable = storage::StorageManager::Get(*clientContext)
->getTable(seedTableID)
->ptrCast<storage::NodeTable>();
if (seedTable->tryGetPrimaryKeyIndex() == nullptr) {
return false;
}
std::shared_ptr<NodeExpression> anchorNode;
int anchorPredIdx = -1;
for (auto i = 0u; i < predicates.size(); ++i) {
if (static_cast<int>(i) == convertIdx) {
continue;
}
auto& pred = predicates[i];
if (pred->expressionType != ExpressionType::EQUALS) {
if (!isVarFree(pred)) {
auto collector = DependentVarNameCollector();
collector.visit(pred);
for (auto& var : collector.getVarNames()) {
if (legVars.contains(var)) {
return false;
}
}
}
continue;
}
auto lhs = pred->getChild(0);
auto rhs = pred->getChild(1);
bool matched = false;
for (auto n = 0u; n < queryGraph->getNumQueryNodes(); ++n) {
auto node = queryGraph->getQueryNode(n);
if (node->getTableIDs().size() != 1) {
continue;
}
auto tableID = node->getTableIDs()[0];
auto l = lhs;
auto r = rhs;
if (isNodePrimaryKey(*r, *node, tableID)) {
std::swap(l, r);
}
if (isNodePrimaryKey(*l, *node, tableID) && isVarFree(r)) {
if (anchorNode != nullptr || node->getUniqueName() == seedNode->getUniqueName()) {
return false;
}
anchorNode = node;
anchorPredIdx = static_cast<int>(i);
matched = true;
break;
}
}
if (!matched) {
auto collector = DependentVarNameCollector();
collector.visit(pred);
for (auto& var : collector.getVarNames()) {
if (legVars.contains(var)) {
return false;
}
}
}
}
std::unordered_map<std::string, std::shared_ptr<NodeExpression>> nodeByName;
for (auto i = 0u; i < queryGraph->getNumQueryNodes(); ++i) {
auto node = queryGraph->getQueryNode(i);
nodeByName[node->getUniqueName()] = node;
}
struct ChainStep {
std::shared_ptr<RelExpression> rel;
std::shared_ptr<NodeExpression> nbr;
};
std::vector<ChainStep> steps;
std::unordered_set<std::string> usedRels;
std::shared_ptr<RelExpression> meetRel;
auto meetNode = seedNode;
auto cur = seedNode;
auto numRels = queryGraph->getNumQueryRels();
while (true) {
std::vector<ChainStep> cands;
for (auto i = 0u; i < numRels; ++i) {
auto rel = queryGraph->getQueryRel(i);
if (usedRels.contains(rel->getUniqueName())) {
continue;
}
auto srcIsCur = *rel->getSrcNode() == *cur;
auto dstIsCur = *rel->getDstNode() == *cur;
if (srcIsCur && dstIsCur) {
return false;
}
if (!srcIsCur && !dstIsCur) {
continue;
}
auto nbr = srcIsCur ? rel->getDstNode() : rel->getSrcNode();
cands.push_back({rel, nbr});
}
if (cands.size() != 1) {
if (!cands.empty() || anchorNode != nullptr || usedRels.size() != numRels) {
return false;
}
break;
}
auto stepRel = cands[0].rel;
auto stepNbr = cands[0].nbr;
if (anchorNode != nullptr && *stepNbr == *anchorNode) {
meetRel = stepRel;
meetNode = cur;
break;
}
steps.push_back({stepRel, stepNbr});
usedRels.insert(stepRel->getUniqueName());
cur = stepNbr;
if (steps.size() > numRels) {
return false;
}
}
if (meetRel != nullptr && usedRels.size() + 1 != numRels) {
return false;
}
cardinalityEstimator.init(*queryGraph);
Binder binder(clientContext);
auto freshElem = std::make_shared<PropertyExpression>(elemPropRef);
freshElem->setUniqueName(binder.getUniqueExpressionName(elemProp->getUniqueName()));
auto unwind =
std::make_shared<LogicalUnwind>(listScoped, freshElem, nullptr, leftPlan.getLastOperator());
appendFlattens(unwind->getGroupsPosToFlatten(), leftPlan);
unwind->setChild(0, leftPlan.getLastOperator());
unwind->computeFactorizedSchema();
leftPlan.setLastOperator(unwind);
distinctKeys.push_back(freshElem);
expression_vector payloads;
for (auto& expr : leftPlan.getSchema()->getExpressionsInScope()) {
auto name = expr->getUniqueName();
if (name == collectName || name == freshElem->getUniqueName()) {
continue;
}
bool isKey = false;
for (auto& key : distinctKeys) {
isKey = isKey || name == key->getUniqueName();
}
if (!isKey) {
payloads.push_back(expr);
}
}
auto distinct = std::make_shared<LogicalDistinct>(distinctKeys, leftPlan.getLastOperator());
distinct->setPayloads(std::move(payloads));
appendFlattens(distinct->getGroupsPosToFlatten(), leftPlan);
distinct->setChild(0, leftPlan.getLastOperator());
distinct->computeFactorizedSchema();
leftPlan.setLastOperator(distinct);
appendFlattens(leftPlan.getSchema()->getGroupsPosInScope(), leftPlan);
auto seedProps = getProperties(*seedNode);
seedProps.erase(std::remove_if(seedProps.begin(), seedProps.end(),
[](const std::shared_ptr<Expression>& expression) {
return expression->constCast<PropertyExpression>().isInternalID();
}),
seedProps.end());
const auto dependentExprs = getDependentExprs(freshElem, *leftPlan.getSchema());
DASSERT(!dependentExprs.empty());
const auto outputGroupPos = leftPlan.getSchema()->getGroupPos(*dependentExprs[0]);
for ([[maybe_unused]] auto& dependentExpr : dependentExprs) {
DASSERT(leftPlan.getSchema()->getGroupPos(*dependentExpr) == outputGroupPos);
}
auto lookup =
std::make_shared<LogicalQueryPrimaryKeyLookup>(seedTableID, seedNode->getInternalID(),
seedProps, freshElem, outputGroupPos, leftPlan.getLastOperator());
lookup->computeFactorizedSchema();
lookup->setCardinality(leftPlan.getCardinality());
leftPlan.setLastOperator(std::move(lookup));
auto curObj = seedNode;
for (auto& step : steps) {
auto dir = chainExtendDirection(*step.rel, *curObj);
appendExtend(curObj, step.nbr, step.rel, dir, getProperties(*step.rel), leftPlan);
curObj = step.nbr;
}
if (meetRel != nullptr) {
LogicalPlan anchorPlan;
appendScanNodeTable(anchorNode->getInternalID(), anchorNode->getTableIDs(),
getProperties(*anchorNode), anchorPlan, anchorNode.get());
appendFilter(predicates[anchorPredIdx], anchorPlan);
auto meetDir = chainExtendDirection(*meetRel, *anchorNode);
appendExtend(anchorNode, meetNode, meetRel, meetDir, getProperties(*meetRel), anchorPlan);
LogicalPlan propsPlan;
appendScanNodeTable(meetNode->getInternalID(), meetNode->getTableIDs(),
getProperties(*meetNode), propsPlan, meetNode.get());
appendHashJoin(expression_vector{meetNode->getInternalID()}, JoinType::INNER, anchorPlan,
propsPlan, anchorPlan);
appendHashJoin(expression_vector{meetNode->getInternalID()}, JoinType::INNER, leftPlan,
anchorPlan, leftPlan);
}
for (auto i = 0u; i < queryGraph->getNumQueryNodes(); ++i) {
auto node = queryGraph->getQueryNode(i);
expression_vector missing;
for (auto& prop : getProperties(*node)) {
if (!leftPlan.getSchema()->isExpressionInScope(*prop)) {
missing.push_back(prop);
}
}
if (missing.empty()) {
continue;
}
LogicalPlan propsPlan;
appendScanNodeTable(node->getInternalID(), node->getTableIDs(), missing, propsPlan,
node.get());
appendHashJoin(expression_vector{node->getInternalID()}, JoinType::INNER, leftPlan,
propsPlan, leftPlan);
}
for (auto i = 0u; i < predicates.size(); ++i) {
if (static_cast<int>(i) == convertIdx || static_cast<int>(i) == anchorPredIdx) {
continue;
}
appendFilter(predicates[i], leftPlan);
}
return true;
}
static std::unordered_set<std::string> collectOpScopeNames(planner::LogicalOperator* op) {
std::unordered_set<std::string> names;
if (op == nullptr || op->getSchema() == nullptr) {
return names;
}
for (auto& expr : op->getSchema()->getExpressionsInScope()) {
names.insert(expr->getUniqueName());
}
return names;
}
static std::unordered_set<std::string> collectOpDepNames(const std::shared_ptr<Expression>& expr,
planner::LogicalOperator* scopeOp) {
std::unordered_set<std::string> names;
if (scopeOp == nullptr || scopeOp->getSchema() == nullptr) {
return names;
}
for (auto& dep : getDependentExprs(expr, *scopeOp->getSchema())) {
names.insert(dep->getUniqueName());
}
return names;
}
static bool isCountDistinctAgg(const std::shared_ptr<Expression>& expr,
std::shared_ptr<Expression>& argOut) {
if (expr->expressionType != ExpressionType::AGGREGATE_FUNCTION) {
return false;
}
auto& aggFunc = expr->constCast<AggregateFunctionExpression>();
if (aggFunc.getFunction().name != function::CountFunction::name || !aggFunc.isDistinct() ||
aggFunc.getNumChildren() != 1) {
return false;
}
argOut = aggFunc.getChild(0);
return true;
}
static bool isListPackExpr(const std::shared_ptr<Expression>& expr) {
if (expr->expressionType != ExpressionType::FUNCTION) {
return false;
}
return expr->constCast<ScalarFunctionExpression>().getFunction().name ==
function::ListCreationFunction::name;
}
static bool isKDetermined(const std::shared_ptr<Expression>& needExpr,
const std::unordered_set<std::string>& kNames, const std::unordered_set<std::string>& kPackDeps,
planner::LogicalOperator* scopeOp) {
auto uname = needExpr->getUniqueName();
if (kNames.contains(uname) || kPackDeps.contains(uname)) {
return true;
}
std::string varName;
if (needExpr->expressionType == ExpressionType::PROPERTY) {
varName = needExpr->constCast<PropertyExpression>().getVariableName();
} else if (ExpressionUtil::isNodePattern(*needExpr) ||
ExpressionUtil::isRelPattern(*needExpr)) {
varName = uname;
} else {
return false;
}
if (scopeOp == nullptr || scopeOp->getSchema() == nullptr) {
return false;
}
for (auto& expr : scopeOp->getSchema()->getExpressionsInScope()) {
auto ename = expr->getUniqueName();
if (ExpressionUtil::isNodePattern(*expr) || ExpressionUtil::isRelPattern(*expr)) {
if (ename != varName) {
continue;
}
std::shared_ptr<Expression> internalID;
if (ExpressionUtil::isNodePattern(*expr)) {
internalID = expr->constCast<NodeExpression>().getInternalID();
} else {
internalID = expr->constCast<RelExpression>().getInternalID();
}
if (kNames.contains(internalID->getUniqueName()) ||
kPackDeps.contains(internalID->getUniqueName())) {
return true;
}
} else if (expr->expressionType == ExpressionType::PROPERTY) {
auto& prop = expr->constCast<PropertyExpression>();
if (prop.getVariableName() == varName && prop.isPrimaryKey() &&
(kNames.contains(ename) || kPackDeps.contains(ename))) {
return true;
}
}
}
return false;
}
bool Planner::tryPreAggregateDistinctLeftChain(const expression_vector& keys,
const expression_vector& aggregates, LogicalPlan& plan) {
if (plan.isEmpty() || keys.empty()) {
return false;
}
struct CountDistinct {
std::shared_ptr<Expression> agg;
std::shared_ptr<Expression> arg;
};
std::vector<CountDistinct> cds;
for (auto& agg : aggregates) {
CountDistinct cd;
if (!isCountDistinctAgg(agg, cd.arg)) {
return false;
}
cd.agg = agg;
cds.push_back(std::move(cd));
}
auto cur = plan.getLastOperator();
while (cur->getOperatorType() == LogicalOperatorType::PROJECTION ||
cur->getOperatorType() == LogicalOperatorType::FLATTEN) {
if (cur->getNumChildren() != 1) {
return false;
}
cur = cur->getChild(0);
}
struct ChainLeg {
std::shared_ptr<LogicalOperator> joinOp;
std::shared_ptr<LogicalOperator> build;
};
std::vector<ChainLeg> legs; while (cur->getOperatorType() == LogicalOperatorType::HASH_JOIN &&
cur->constCast<LogicalHashJoin>().getJoinType() == JoinType::LEFT) {
if (cur->getNumChildren() != 2) {
return false;
}
legs.push_back({cur, cur->getChild(1)});
cur = cur->getChild(0);
}
if (legs.size() < 2) {
return false;
}
std::reverse(legs.begin(), legs.end());
const auto n = legs.size();
auto baseOp = cur;
auto baseNames = collectOpScopeNames(baseOp.get());
std::unordered_set<std::string> kNames;
std::unordered_set<std::string> kDepNames;
for (auto& key : keys) {
kNames.insert(key->getUniqueName());
for (auto& dep : collectOpDepNames(key, baseOp.get())) {
kDepNames.insert(dep);
}
}
if (kDepNames.empty()) {
return false;
}
for (auto& dep : kDepNames) {
if (!baseNames.contains(dep)) {
return false;
}
}
std::unordered_set<std::string> kPackDeps;
if (keys.size() == 1 && isListPackExpr(keys[0])) {
for (auto& child : keys[0]->getChildren()) {
kPackDeps.insert(child->getUniqueName());
}
}
std::vector<std::vector<std::shared_ptr<Expression>>> legAggs(n);
std::vector<std::unordered_set<std::string>> buildNames(n);
for (auto i = 0u; i < n; ++i) {
buildNames[i] = collectOpScopeNames(legs[i].build.get());
}
for (auto& cd : cds) {
auto argName = cd.arg->getUniqueName();
if (baseNames.contains(argName)) {
return false;
}
int owner = -1;
for (auto i = 0u; i < n; ++i) {
if (buildNames[i].contains(argName)) {
if (owner >= 0) {
return false;
}
owner = static_cast<int>(i);
}
}
if (owner < 0) {
return false;
}
legAggs[owner].push_back(cd.agg);
}
std::unordered_set<std::string> cdNames;
for (auto& cd : cds) {
cdNames.insert(cd.agg->getUniqueName());
}
std::vector<expression_vector> needExprs(n);
for (auto i = 0u; i + 1 < n; ++i) {
std::unordered_set<std::string> names = kDepNames;
for (auto j = i + 1; j < n; ++j) {
auto probeScopeOp = legs[j].joinOp->getChild(0).get();
for (auto& cond : legs[j].joinOp->constCast<LogicalHashJoin>().getJoinConditions()) {
for (auto& dep : collectOpDepNames(cond.first, probeScopeOp)) {
names.insert(dep);
}
}
}
for (auto j = 0u; j <= i; ++j) {
for (auto& aggObj : legAggs[j]) {
names.insert(aggObj->getUniqueName());
}
}
auto stageScopeOp = legs[i].joinOp.get();
if (stageScopeOp->getSchema() == nullptr) {
return false;
}
auto stageNames = collectOpScopeNames(stageScopeOp);
expression_vector exprs;
for (auto& name : names) {
if (kNames.contains(name) || cdNames.contains(name)) {
continue;
}
if (!stageNames.contains(name)) {
return false;
}
auto found = findInScope(*stageScopeOp->getSchema(), name);
if (found == nullptr || !isKDetermined(found, kNames, kPackDeps, stageScopeOp)) {
return false;
}
exprs.push_back(found);
}
for (auto j = 0u; j <= i; ++j) {
for (auto& aggObj : legAggs[j]) {
exprs.push_back(aggObj);
}
}
needExprs[i] = std::move(exprs);
}
auto appendStageAgg = [&](std::shared_ptr<LogicalOperator> childTop,
const expression_vector& stageAggs,
const expression_vector& payloads) {
auto stageAgg = std::make_shared<LogicalAggregate>(keys, stageAggs, childTop);
stageAgg->setDependentKeys(payloads);
for (auto groupPos : stageAgg->getGroupsPosToFlatten()) {
auto child = stageAgg->getChild(0);
auto group = child->getSchema()->getGroup(groupPos);
if (group->isFlat()) {
continue;
}
auto flatten = std::make_shared<LogicalFlatten>(groupPos, child,
cardinalityEstimator.estimateFlatten(*child, groupPos));
flatten->computeFactorizedSchema();
stageAgg->setChild(0, flatten);
}
stageAgg->computeFactorizedSchema();
const auto childTopProbeCard = childTop->getNumChildren() > 0 ?
childTop->getChild(0)->getCardinality() :
stageAgg->getChild(0)->getCardinality();
stageAgg->setCardinality(std::min(cardinalityEstimator.estimateAggregate(*stageAgg),
std::max<cardinality_t>(childTopProbeCard, 1)));
stageAgg->setCardinalityLocked(true);
return stageAgg;
};
if (baseOp->getSchema() != nullptr) {
expression_vector projExprs;
for (auto& expr : baseOp->getSchema()->getExpressionsInScope()) {
projExprs.push_back(expr);
}
bool needProj = false;
for (auto& key : keys) {
bool present = false;
for (auto& expr : projExprs) {
if (expr->getUniqueName() == key->getUniqueName()) {
present = true;
break;
}
}
if (!present) {
projExprs.push_back(key);
needProj = true;
}
}
if (needProj) {
auto keyProj = std::make_shared<LogicalProjection>(projExprs, baseOp);
keyProj->computeFactorizedSchema();
keyProj->setCardinality(baseOp->getCardinality());
legs[0].joinOp->setChild(0, keyProj);
for (auto i = 0u; i < n; ++i) {
legs[i].joinOp->computeFactorizedSchema();
}
}
}
for (auto i = 0u; i + 1 < n; ++i) {
auto stageAgg = appendStageAgg(legs[i].joinOp, legAggs[i], needExprs[i]);
legs[i + 1].joinOp->setChild(0, stageAgg);
}
for (auto i = 0u; i < n; ++i) {
legs[i].joinOp->computeFactorizedSchema();
}
expression_vector topPayloads;
for (auto i = 0u; i + 1 < n; ++i) {
for (auto& aggObj : legAggs[i]) {
topPayloads.push_back(aggObj);
}
}
auto topAgg = appendStageAgg(legs[n - 1].joinOp, legAggs[n - 1], topPayloads);
plan.setLastOperator(topAgg);
return true;
}
void Planner::planOptionalMatch(const QueryGraphCollection& queryGraphCollection,
const expression_vector& predicates, std::shared_ptr<Expression> mark, LogicalPlan& leftPlan,
std::shared_ptr<BoundJoinHintNode> hint) {
expression_vector legPredicates = predicates;
bool legacyPath = mark != nullptr || hint != nullptr;
if (!leftPlan.isEmpty() && !legacyPath) {
tryUnnestCollectMembership(queryGraphCollection, legPredicates, leftPlan);
}
expression_vector correlatedExprs;
if (!leftPlan.isEmpty()) {
correlatedExprs =
getCorrelatedExprs(queryGraphCollection, legPredicates, leftPlan.getSchema());
}
auto info = QueryGraphPlanningInfo();
info.hint = hint;
if (leftPlan.isEmpty()) {
info.predicates = legPredicates;
auto plan = planQueryGraphCollection(queryGraphCollection, info);
leftPlan.setLastOperator(plan.getLastOperator());
appendOptionalAccumulate(mark, leftPlan);
return;
}
if (correlatedExprs.empty()) {
info.predicates = legPredicates;
auto rightPlan = planQueryGraphCollection(queryGraphCollection, info);
if (leftPlan.hasUpdate()) {
appendAccOptionalCrossProduct(mark, leftPlan, rightPlan, leftPlan);
} else {
appendOptionalCrossProduct(mark, leftPlan, rightPlan, leftPlan);
}
return;
}
info.corrExprsCard = leftPlan.getCardinality();
auto analyzer =
SubqueryPredicatePullUpAnalyzer(*leftPlan.getSchema(), queryGraphCollection, !legacyPath);
bool canUnnest = analyzer.analyze(legPredicates);
bool innerSelective = true;
if (canUnnest && !leftPlan.isEmpty()) {
std::unordered_set<std::string> constFilteredVars;
for (auto& pred : legPredicates) {
if (!isRoutableConstantEquality(pred, queryGraphCollection)) {
continue;
}
auto collector = DependentVarNameCollector();
collector.visit(pred);
constFilteredVars.insert(collector.getVarNames().begin(),
collector.getVarNames().end());
}
for (auto& node : queryGraphCollection.getQueryNodes()) {
if (leftPlan.getSchema()->isExpressionInScope(*node->getInternalID()) &&
!constFilteredVars.contains(node->getUniqueName())) {
innerSelective = false;
break;
}
}
}
bool hasRecursiveRel = false;
for (auto& rel : queryGraphCollection.getQueryRels()) {
if (common::QueryRelTypeUtils::isRecursive(rel->getRelType())) {
hasRecursiveRel = true;
break;
}
}
CorrelatedOptionalLegDecision decision;
decision.legacyPath = legacyPath;
decision.canUnnest = canUnnest;
decision.innerSelective = innerSelective;
decision.hasRecursiveRel = hasRecursiveRel;
if (canUnnest) {
for (auto& joinCondition : analyzer.getJoinConditions()) {
if (joinCondition.first->getUniqueName() != joinCondition.second->getUniqueName()) {
decision.hasCrossVariableJoin = true;
break;
}
}
}
std::shared_ptr<Expression> pkLookupKey;
expression_vector pkLookupResiduals;
decision.pkLookupEligible = findCorrelatedPrimaryKeyLookupKey(queryGraphCollection,
legPredicates, correlatedExprs, *leftPlan.getSchema(), pkLookupKey, pkLookupResiduals);
std::vector<expression_pair> joinConditions;
LogicalPlan rightPlan;
const auto strategy = decision.decide();
if (strategy == CorrelatedOptionalLegDecision::Strategy::PRIMARY_KEY_LOOKUP &&
tryPlanCorrelatedPrimaryKeyLookup(queryGraphCollection, legPredicates, correlatedExprs,
*leftPlan.getSchema(), info.corrExprsCard, rightPlan)) {
info.subqueryType = SubqueryPlanningType::CORRELATED;
info.corrExprs = correlatedExprs;
info.predicates = legPredicates;
for (auto& expr : correlatedExprs) {
joinConditions.emplace_back(expr, expr);
}
appendAccumulate(correlatedExprs, leftPlan);
} else if (strategy == CorrelatedOptionalLegDecision::Strategy::UNNEST_LEFT_JOIN) {
info.subqueryType = SubqueryPlanningType::UNNEST_CORRELATED;
info.corrExprs = analyzer.getCorrelatedInternalIDs();
info.predicates = analyzer.getNonCorrelatedPredicates();
rightPlan = planQueryGraphCollectionInNewContext(queryGraphCollection, info);
joinConditions = analyzer.getJoinConditions();
} else {
info.subqueryType = SubqueryPlanningType::CORRELATED;
info.corrExprs = correlatedExprs;
info.predicates = legPredicates;
for (auto& expr : correlatedExprs) {
joinConditions.emplace_back(expr, expr);
}
rightPlan = planQueryGraphCollectionInNewContext(queryGraphCollection, info);
appendAccumulate(correlatedExprs, leftPlan);
}
if (leftPlan.hasUpdate()) {
appendAccHashJoin(joinConditions, JoinType::LEFT, mark, leftPlan, rightPlan, leftPlan);
} else {
appendHashJoin(joinConditions, JoinType::LEFT, mark, leftPlan, rightPlan, leftPlan);
}
}
void Planner::planRegularMatch(const QueryGraphCollection& queryGraphCollection,
const expression_vector& predicates, LogicalPlan& leftPlan,
std::shared_ptr<BoundJoinHintNode> hint) {
if (tryPlanQueryPrimaryKeyLookup(queryGraphCollection, predicates, leftPlan)) {
return;
}
if (tryPlanPropertySeededChain(queryGraphCollection, predicates, leftPlan)) {
return;
}
expression_vector predicatesToPushDown, predicatesToPullUp;
for (auto& predicate : predicates) {
if (getDependentExprs(predicate, *leftPlan.getSchema()).empty()) {
predicatesToPushDown.push_back(predicate);
} else {
predicatesToPullUp.push_back(predicate);
}
}
auto correlatedExprs =
getCorrelatedExprs(queryGraphCollection, predicatesToPushDown, leftPlan.getSchema());
auto joinNodeIDs =
ExpressionUtil::getExpressionsWithDataType(correlatedExprs, LogicalTypeID::INTERNAL_ID);
auto info = QueryGraphPlanningInfo();
info.predicates = predicatesToPushDown;
info.hint = hint;
if (joinNodeIDs.empty()) {
info.subqueryType = SubqueryPlanningType::NONE;
auto rightPlan = planQueryGraphCollectionInNewContext(queryGraphCollection, info);
if (leftPlan.hasUpdate()) {
appendCrossProduct(rightPlan, leftPlan, leftPlan);
} else {
appendCrossProduct(leftPlan, rightPlan, leftPlan);
}
} else {
info.subqueryType = SubqueryPlanningType::UNNEST_CORRELATED;
info.corrExprs = joinNodeIDs;
info.corrExprsCard = leftPlan.getCardinality();
auto rightPlan = planQueryGraphCollectionInNewContext(queryGraphCollection, info);
if (leftPlan.hasUpdate()) {
appendHashJoin(joinNodeIDs, JoinType::INNER, rightPlan, leftPlan, leftPlan);
} else {
appendHashJoin(joinNodeIDs, JoinType::INNER, leftPlan, rightPlan, leftPlan);
}
}
for (auto& predicate : predicatesToPullUp) {
appendFilter(predicate, leftPlan);
}
}
void Planner::planSubquery(const std::shared_ptr<Expression>& expression, LogicalPlan& outerPlan) {
DASSERT(expression->expressionType == ExpressionType::SUBQUERY);
auto subquery = expression->ptrCast<SubqueryExpression>();
auto correlatedExprs = getDependentExprs(expression, *outerPlan.getSchema());
auto predicates = subquery->getPredicatesSplitOnAnd();
LogicalPlan innerPlan;
auto info = QueryGraphPlanningInfo();
info.hint = subquery->getHint();
if (correlatedExprs.empty()) {
info.subqueryType = SubqueryPlanningType::NONE;
info.predicates = predicates;
innerPlan =
planQueryGraphCollectionInNewContext(*subquery->getQueryGraphCollection(), info);
expression_vector emptyHashKeys;
auto projectExprs = expression_vector{subquery->getProjectionExpr()};
switch (subquery->getSubqueryType()) {
case common::SubqueryType::EXISTS: {
auto aggregates = expression_vector{subquery->getCountStarExpr()};
appendAggregate(emptyHashKeys, aggregates, innerPlan);
appendProjection(projectExprs, innerPlan);
} break;
case common::SubqueryType::COUNT: {
appendAggregate(emptyHashKeys, projectExprs, innerPlan);
} break;
default:
UNREACHABLE_CODE;
}
appendCrossProduct(outerPlan, innerPlan, outerPlan);
return;
}
info.corrExprsCard = outerPlan.getCardinality();
auto analyzer = SubqueryPredicatePullUpAnalyzer(*outerPlan.getSchema(),
*subquery->getQueryGraphCollection());
std::vector<expression_pair> joinConditions;
if (analyzer.analyze(predicates)) {
info.subqueryType = SubqueryPlanningType::UNNEST_CORRELATED;
info.corrExprs = analyzer.getCorrelatedInternalIDs();
info.predicates = analyzer.getNonCorrelatedPredicates();
innerPlan =
planQueryGraphCollectionInNewContext(*subquery->getQueryGraphCollection(), info);
joinConditions = analyzer.getJoinConditions();
} else {
info.subqueryType = SubqueryPlanningType::CORRELATED;
info.corrExprs = correlatedExprs;
info.predicates = predicates;
for (auto& expr : correlatedExprs) {
joinConditions.emplace_back(expr, expr);
}
innerPlan =
planQueryGraphCollectionInNewContext(*subquery->getQueryGraphCollection(), info);
appendAccumulate(correlatedExprs, outerPlan);
}
switch (subquery->getSubqueryType()) {
case common::SubqueryType::EXISTS: {
appendMarkJoin(joinConditions, expression, outerPlan, innerPlan, outerPlan);
} break;
case common::SubqueryType::COUNT: {
expression_vector hashKeys;
for (auto& joinCondition : joinConditions) {
hashKeys.push_back(joinCondition.second);
}
appendAggregate(hashKeys, expression_vector{subquery->getProjectionExpr()}, innerPlan);
appendHashJoin(joinConditions, common::JoinType::COUNT, nullptr, outerPlan, innerPlan,
outerPlan);
} break;
default:
UNREACHABLE_CODE;
}
}
void Planner::planSubqueryIfNecessary(std::shared_ptr<Expression> expression, LogicalPlan& plan) {
auto collector = SubqueryExprCollector();
collector.visit(expression);
if (collector.hasSubquery()) {
for (auto& expr : collector.getSubqueryExprs()) {
if (plan.getSchema()->isExpressionInScope(*expr)) {
continue;
}
planSubquery(expr, plan);
}
}
}
} }